mirror of
https://github.com/langchain-ai/langchain.git
synced 2026-10-05 09:25:14 +03:00
feat(typesafe): add hybrid tool selector middleware
Co-authored-by: open-swe[bot] <open-swe@users.noreply.github.com>
This commit is contained in:
1 parent
35a8757e92
commit
bc5c853b34
4 files changed
+443
-17
No files matched your search
@@ -103,6 +103,22 @@ agent = create_agent(
|
||||
|
||||
This variant asks one `Choice` question over the candidate tools and exposes only the chosen tool for the next model call; it chooses again on subsequent calls. Both variants accept `always_include` to keep named tools without classification, and preserve provider-specific tool definitions. This API is experimental and may change without notice.
|
||||
|
||||
### Experimental hybrid tool selector middleware
|
||||
|
||||
Use `TsHybridToolSelectorMiddleware` when a step might need no tool, a single tool, or several tools:
|
||||
|
||||
```python
|
||||
from langchain_typesafe.experimental.middleware import TsHybridToolSelectorMiddleware
|
||||
|
||||
agent = create_agent(
|
||||
model,
|
||||
tools=[tool1, tool2, tool3],
|
||||
middleware=[TsHybridToolSelectorMiddleware(relevance_threshold=0.5, max_tools=3)],
|
||||
)
|
||||
```
|
||||
|
||||
Before each model call, a `Choice` question against the latest human message picks `none`, `single`, or `multiple`. `none` hides candidate tools without a second classifier call; `single` chooses one tool as in `TsChoiceToolSelectorMiddleware`; `multiple` uses the thresholded, probability-ranked `Noul` batch from `TsToolSelectorMiddleware`. `always_include` tools and provider-specific tool definitions remain available in every mode. Classifier errors and invalid choices raise. This API is experimental and may change without notice.
|
||||
|
||||
### LangChain messages as state
|
||||
|
||||
`BaseMessage` objects and message sequences can appear at the root or anywhere inside JSON state. The integration recursively converts them to objects with `role` and `content` fields while preserving surrounding application data:
|
||||
|
||||
@@ -7,6 +7,7 @@ from langchain_typesafe.experimental.middleware.skills import (
|
||||
)
|
||||
from langchain_typesafe.experimental.middleware.tool_selector import (
|
||||
TsChoiceToolSelectorMiddleware,
|
||||
TsHybridToolSelectorMiddleware,
|
||||
TsToolSelectorMiddleware,
|
||||
)
|
||||
|
||||
@@ -15,5 +16,6 @@ __all__ = [
|
||||
"SkillSource",
|
||||
"SkillsMiddleware",
|
||||
"TsChoiceToolSelectorMiddleware",
|
||||
"TsHybridToolSelectorMiddleware",
|
||||
"TsToolSelectorMiddleware",
|
||||
]
|
||||
+170
-17
@@ -265,6 +265,32 @@ class TsToolSelectorMiddleware(
|
||||
return await handler(modified_request)
|
||||
|
||||
|
||||
def _build_choice_classifier(tools: list[BaseTool]) -> TypeSafeClassifier:
|
||||
"""Build one categorical question for all candidate tools."""
|
||||
return TypeSafeClassifier(
|
||||
questions={
|
||||
"tool": Choice(
|
||||
instructions=(
|
||||
"Which one of these tools is needed next to make progress on "
|
||||
"the user's current request? Select exactly one tool."
|
||||
),
|
||||
criteria={tool.name: tool.description for tool in tools},
|
||||
)
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
def _select_choice_tool_names(
|
||||
response: ClassificationResponse, selection_request: _SelectionRequest
|
||||
) -> list[str]:
|
||||
"""Require that the chosen tool belongs to the candidates."""
|
||||
answer = response.choices.get("tool")
|
||||
if answer is None or answer.choice not in selection_request.valid_tool_names:
|
||||
msg = "TypeSafe returned no valid tool choice for the current request"
|
||||
raise ValueError(msg)
|
||||
return [answer.choice]
|
||||
|
||||
|
||||
class TsChoiceToolSelectorMiddleware(TsToolSelectorMiddleware):
|
||||
"""Select one candidate tool for the next model call with a TypeSafe `Choice`.
|
||||
|
||||
@@ -299,31 +325,158 @@ class TsChoiceToolSelectorMiddleware(TsToolSelectorMiddleware):
|
||||
|
||||
def _build_classifier(self, tools: list[BaseTool]) -> TypeSafeClassifier:
|
||||
"""Build one categorical question for all candidate tools."""
|
||||
return TypeSafeClassifier(
|
||||
questions={
|
||||
"tool": Choice(
|
||||
instructions=(
|
||||
"Which one of these tools is needed next to make progress on "
|
||||
"the user's current request? Select exactly one tool."
|
||||
),
|
||||
criteria={tool.name: tool.description for tool in tools},
|
||||
)
|
||||
},
|
||||
)
|
||||
return _build_choice_classifier(tools)
|
||||
|
||||
def _select_tool_names(
|
||||
self, response: ClassificationResponse, selection_request: _SelectionRequest
|
||||
) -> list[str]:
|
||||
"""Return only the chosen tool if it is one of the candidates."""
|
||||
answer = response.choices.get("tool")
|
||||
if answer is None or answer.choice not in selection_request.valid_tool_names:
|
||||
msg = "TypeSafe returned no valid tool choice for the current request"
|
||||
raise ValueError(msg)
|
||||
return [answer.choice]
|
||||
return _select_choice_tool_names(response, selection_request)
|
||||
|
||||
def _classifier_config(self) -> RunnableConfig:
|
||||
"""Tag choice selector calls separately in traces."""
|
||||
return {"metadata": {"lc_source": "ts_choice_tool_selector"}}
|
||||
|
||||
|
||||
__all__ = ["TsChoiceToolSelectorMiddleware", "TsToolSelectorMiddleware"]
|
||||
class TsHybridToolSelectorMiddleware(TsToolSelectorMiddleware):
|
||||
"""Choose whether to expose no tools, one tool, or several tools per model call.
|
||||
|
||||
A first `Choice` question decides the selection shape from the latest human
|
||||
message. A second classifier selects tools only for `single` or `multiple`:
|
||||
`single` uses the choice selector, while `multiple` uses the thresholded
|
||||
per-tool `Noul` selector. Classifier errors and invalid choices propagate.
|
||||
|
||||
!!! warning
|
||||
|
||||
This middleware is experimental. Its API may change without notice.
|
||||
|
||||
Args:
|
||||
relevance_threshold: Minimum probability for tools in `multiple` mode.
|
||||
max_tools: Maximum number of classified tools in `multiple` mode.
|
||||
always_include: Tool names to pass through without classification.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
relevance_threshold: float = 0.5,
|
||||
max_tools: int | None = None,
|
||||
always_include: list[str] | None = None,
|
||||
) -> None:
|
||||
"""Initialize the hybrid selector with multi-tool selection settings."""
|
||||
super().__init__(
|
||||
relevance_threshold=relevance_threshold,
|
||||
max_tools=max_tools,
|
||||
always_include=always_include,
|
||||
)
|
||||
|
||||
def _build_shape_classifier(self) -> TypeSafeClassifier:
|
||||
"""Build the question that decides how many tools the next step needs."""
|
||||
return TypeSafeClassifier(
|
||||
questions={
|
||||
"shape": Choice(
|
||||
instructions=(
|
||||
"Does the user's current request need a tool for the next "
|
||||
"step? Choose none if no tool is needed, single if one tool "
|
||||
"is needed, or multiple if several tools are needed."
|
||||
),
|
||||
criteria={
|
||||
"none": "No tool is needed for the next step.",
|
||||
"single": "Exactly one tool is needed for the next step.",
|
||||
"multiple": "Several tools may be needed for the next step.",
|
||||
},
|
||||
)
|
||||
}
|
||||
)
|
||||
|
||||
def _shape(self, response: ClassificationResponse) -> str:
|
||||
"""Require a valid answer before changing the tools available to the model."""
|
||||
answer = response.choices.get("shape")
|
||||
if answer is None or answer.choice not in {"none", "single", "multiple"}:
|
||||
msg = (
|
||||
"TypeSafe returned no valid tool selection shape "
|
||||
"for the current request"
|
||||
)
|
||||
raise ValueError(msg)
|
||||
return answer.choice
|
||||
|
||||
def _classifier_config(self) -> RunnableConfig:
|
||||
"""Identify second-stage tool selection separately in traces."""
|
||||
return {"metadata": {"lc_source": "ts_hybrid_tool_selector_stage_2"}}
|
||||
|
||||
def _stage_one_config(self) -> RunnableConfig:
|
||||
"""Identify the shape classification separately in traces."""
|
||||
return {"metadata": {"lc_source": "ts_hybrid_tool_selector_stage_1"}}
|
||||
|
||||
@override
|
||||
def wrap_model_call(
|
||||
self,
|
||||
request: ModelRequest[None],
|
||||
handler: Callable[[ModelRequest[None]], ModelResponse[ResponseT]],
|
||||
) -> ModelResponse[ResponseT]:
|
||||
"""Classify the shape and select tools before invoking the model."""
|
||||
selection_request = self._prepare_selection_request(request)
|
||||
if selection_request is None:
|
||||
return handler(request)
|
||||
|
||||
shape_response = self._build_shape_classifier().invoke(
|
||||
selection_request.last_user_message, config=self._stage_one_config()
|
||||
)
|
||||
shape = self._shape(shape_response)
|
||||
selected = []
|
||||
if shape != "none":
|
||||
classifier = (
|
||||
_build_choice_classifier(selection_request.classifiable_tools)
|
||||
if shape == "single"
|
||||
else self._build_classifier(selection_request.classifiable_tools)
|
||||
)
|
||||
response = classifier.invoke(
|
||||
selection_request.last_user_message, config=self._classifier_config()
|
||||
)
|
||||
selected = (
|
||||
_select_choice_tool_names(response, selection_request)
|
||||
if shape == "single"
|
||||
else self._select_tool_names(response, selection_request)
|
||||
)
|
||||
return handler(self._process_selection(selected, selection_request, request))
|
||||
|
||||
@override
|
||||
async def awrap_model_call(
|
||||
self,
|
||||
request: ModelRequest[None],
|
||||
handler: Callable[[ModelRequest[None]], Awaitable[ModelResponse[ResponseT]]],
|
||||
) -> ModelResponse[ResponseT]:
|
||||
"""Classify the shape and select tools asynchronously."""
|
||||
selection_request = self._prepare_selection_request(request)
|
||||
if selection_request is None:
|
||||
return await handler(request)
|
||||
|
||||
shape_response = await self._build_shape_classifier().ainvoke(
|
||||
selection_request.last_user_message, config=self._stage_one_config()
|
||||
)
|
||||
shape = self._shape(shape_response)
|
||||
selected = []
|
||||
if shape != "none":
|
||||
classifier = (
|
||||
_build_choice_classifier(selection_request.classifiable_tools)
|
||||
if shape == "single"
|
||||
else self._build_classifier(selection_request.classifiable_tools)
|
||||
)
|
||||
response = await classifier.ainvoke(
|
||||
selection_request.last_user_message, config=self._classifier_config()
|
||||
)
|
||||
selected = (
|
||||
_select_choice_tool_names(response, selection_request)
|
||||
if shape == "single"
|
||||
else self._select_tool_names(response, selection_request)
|
||||
)
|
||||
return await handler(
|
||||
self._process_selection(selected, selection_request, request)
|
||||
)
|
||||
|
||||
|
||||
__all__ = [
|
||||
"TsChoiceToolSelectorMiddleware",
|
||||
"TsHybridToolSelectorMiddleware",
|
||||
"TsToolSelectorMiddleware",
|
||||
]
|
||||
+255
@@ -14,6 +14,7 @@ from langchain_core.tools import tool
|
||||
from langchain_typesafe import ChoiceAnswer, NoulAnswer, TypeSafeClassifier
|
||||
from langchain_typesafe.experimental.middleware import (
|
||||
TsChoiceToolSelectorMiddleware,
|
||||
TsHybridToolSelectorMiddleware,
|
||||
TsToolSelectorMiddleware,
|
||||
)
|
||||
from langchain_typesafe.experimental.middleware import __all__ as middleware_all
|
||||
@@ -62,6 +63,20 @@ def _choice_response(choice: str) -> ClassificationResponse:
|
||||
)
|
||||
|
||||
|
||||
def _shape_response(shape: str) -> ClassificationResponse:
|
||||
return ClassificationResponse(
|
||||
model="jev-latest",
|
||||
answers={
|
||||
"shape": ChoiceAnswer(
|
||||
type="choice",
|
||||
choice=shape,
|
||||
probabilities={shape: 1.0},
|
||||
confidence=1.0,
|
||||
)
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
def _classifier(response: ClassificationResponse) -> MagicMock:
|
||||
classifier = MagicMock(spec=TypeSafeClassifier)
|
||||
classifier.invoke.return_value = response
|
||||
@@ -422,6 +437,245 @@ def test_choice_only_always_included_tools_skips_classifier() -> None:
|
||||
classifier_class.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.parametrize("shape", ["none", "single", "multiple"])
|
||||
def test_hybrid_routes_sync(shape: str) -> None:
|
||||
"""Classify shape once, then invoke only the needed selection stage."""
|
||||
first = _classifier(_shape_response(shape))
|
||||
second = _classifier(
|
||||
_choice_response("search_web")
|
||||
if shape == "single"
|
||||
else _response({"search_web": 0.8, "send_email": 0.1})
|
||||
)
|
||||
request = _request(
|
||||
[get_weather, search_web, send_email, {"type": "web_search"}],
|
||||
[HumanMessage("Help me"), AIMessage("Next step")],
|
||||
)
|
||||
middleware = TsHybridToolSelectorMiddleware(
|
||||
max_tools=1, always_include=["get_weather"]
|
||||
)
|
||||
seen: list[ModelRequest[Any]] = []
|
||||
|
||||
def handler(modified: ModelRequest[Any]) -> ModelResponse[Any]:
|
||||
seen.append(modified)
|
||||
return cast("ModelResponse[Any]", MagicMock())
|
||||
|
||||
with patch(
|
||||
"langchain_typesafe.experimental.middleware.tool_selector.TypeSafeClassifier",
|
||||
side_effect=[first, second],
|
||||
) as classifier_class:
|
||||
middleware.wrap_model_call(request, handler)
|
||||
|
||||
assert classifier_class.call_count == (1 if shape == "none" else 2)
|
||||
assert classifier_class.call_args_list[0].kwargs["questions"]["shape"].criteria == {
|
||||
"none": "No tool is needed for the next step.",
|
||||
"single": "Exactly one tool is needed for the next step.",
|
||||
"multiple": "Several tools may be needed for the next step.",
|
||||
}
|
||||
first.invoke.assert_called_once_with(
|
||||
request.messages[0],
|
||||
config={"metadata": {"lc_source": "ts_hybrid_tool_selector_stage_1"}},
|
||||
)
|
||||
assert _tool_names(seen[0].tools) == (
|
||||
["get_weather"] if shape == "none" else ["search_web", "get_weather"]
|
||||
)
|
||||
assert seen[0].tools[-1] is request.tools[-1]
|
||||
assert request.tools[0] is get_weather
|
||||
if shape == "none":
|
||||
second.invoke.assert_not_called()
|
||||
else:
|
||||
second.invoke.assert_called_once_with(
|
||||
request.messages[0],
|
||||
config={"metadata": {"lc_source": "ts_hybrid_tool_selector_stage_2"}},
|
||||
)
|
||||
questions = classifier_class.call_args.kwargs["questions"]
|
||||
if shape == "single":
|
||||
assert questions["tool"].criteria == {
|
||||
"search_web": search_web.description,
|
||||
"send_email": send_email.description,
|
||||
}
|
||||
else:
|
||||
assert set(questions) == {"tool::search_web", "tool::send_email"}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("shape", ["none", "single", "multiple"])
|
||||
async def test_hybrid_routes_async(shape: str) -> None:
|
||||
"""Use the asynchronous classifier path for every shape."""
|
||||
first = _classifier(_shape_response(shape))
|
||||
second = _classifier(
|
||||
_choice_response("search_web")
|
||||
if shape == "single"
|
||||
else _response({"search_web": 0.8, "send_email": 0.1})
|
||||
)
|
||||
request = _request([get_weather, search_web, send_email], [HumanMessage("Help me")])
|
||||
seen: list[ModelRequest[Any]] = []
|
||||
|
||||
async def handler(modified: ModelRequest[Any]) -> ModelResponse[Any]:
|
||||
seen.append(modified)
|
||||
return cast("ModelResponse[Any]", MagicMock())
|
||||
|
||||
with patch(
|
||||
"langchain_typesafe.experimental.middleware.tool_selector.TypeSafeClassifier",
|
||||
side_effect=[first, second],
|
||||
) as classifier_class:
|
||||
await TsHybridToolSelectorMiddleware(max_tools=1).awrap_model_call(
|
||||
request, handler
|
||||
)
|
||||
|
||||
assert classifier_class.call_count == (1 if shape == "none" else 2)
|
||||
first.ainvoke.assert_awaited_once()
|
||||
assert _tool_names(seen[0].tools) == ([] if shape == "none" else ["search_web"])
|
||||
if shape == "none":
|
||||
second.ainvoke.assert_not_awaited()
|
||||
else:
|
||||
second.ainvoke.assert_awaited_once()
|
||||
assert second.ainvoke.call_args.kwargs["config"]["metadata"] == {
|
||||
"lc_source": "ts_hybrid_tool_selector_stage_2"
|
||||
}
|
||||
|
||||
|
||||
@pytest.mark.parametrize("response", [_shape_response("unknown"), _response({})])
|
||||
def test_hybrid_rejects_invalid_shape(response: ClassificationResponse) -> None:
|
||||
"""Reject unknown or missing shape answers before tool selection."""
|
||||
request = _request([get_weather], [HumanMessage("Help me")])
|
||||
with (
|
||||
patch(
|
||||
"langchain_typesafe.experimental.middleware.tool_selector.TypeSafeClassifier",
|
||||
return_value=_classifier(response),
|
||||
) as classifier_class,
|
||||
pytest.raises(ValueError, match="no valid tool selection shape"),
|
||||
):
|
||||
TsHybridToolSelectorMiddleware().wrap_model_call(
|
||||
request, lambda _req: MagicMock()
|
||||
)
|
||||
classifier_class.assert_called_once()
|
||||
|
||||
|
||||
@pytest.mark.parametrize("stage", [1, 2])
|
||||
def test_hybrid_sync_classifier_failure_propagates(stage: int) -> None:
|
||||
"""Never continue with guessed tools after a classifier error."""
|
||||
first = _classifier(_shape_response("single"))
|
||||
second = _classifier(_choice_response("get_weather"))
|
||||
(first if stage == 1 else second).invoke.side_effect = RuntimeError("unavailable")
|
||||
request = _request([get_weather], [HumanMessage("Help me")])
|
||||
with (
|
||||
patch(
|
||||
"langchain_typesafe.experimental.middleware.tool_selector.TypeSafeClassifier",
|
||||
side_effect=[first, second],
|
||||
),
|
||||
pytest.raises(RuntimeError, match="unavailable"),
|
||||
):
|
||||
TsHybridToolSelectorMiddleware().wrap_model_call(
|
||||
request, lambda _req: MagicMock()
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("stage", [1, 2])
|
||||
async def test_hybrid_async_classifier_failure_propagates(stage: int) -> None:
|
||||
"""Never continue after an asynchronous classifier error."""
|
||||
first = _classifier(_shape_response("multiple"))
|
||||
second = _classifier(_response({"get_weather": 0.9}))
|
||||
(first if stage == 1 else second).ainvoke.side_effect = RuntimeError("unavailable")
|
||||
request = _request([get_weather], [HumanMessage("Help me")])
|
||||
with (
|
||||
patch(
|
||||
"langchain_typesafe.experimental.middleware.tool_selector.TypeSafeClassifier",
|
||||
side_effect=[first, second],
|
||||
),
|
||||
pytest.raises(RuntimeError, match="unavailable"),
|
||||
):
|
||||
await TsHybridToolSelectorMiddleware().awrap_model_call(request, AsyncMock())
|
||||
|
||||
|
||||
@pytest.mark.parametrize("response", [_choice_response("invalid"), _response({})])
|
||||
def test_hybrid_single_rejects_invalid_tool(response: ClassificationResponse) -> None:
|
||||
"""Preserve the choice selector's fail-closed invalid-tool behavior."""
|
||||
first = _classifier(_shape_response("single"))
|
||||
second = _classifier(response)
|
||||
request = _request([get_weather], [HumanMessage("Help me")])
|
||||
with (
|
||||
patch(
|
||||
"langchain_typesafe.experimental.middleware.tool_selector.TypeSafeClassifier",
|
||||
side_effect=[first, second],
|
||||
),
|
||||
pytest.raises(ValueError, match="no valid tool choice"),
|
||||
):
|
||||
TsHybridToolSelectorMiddleware().wrap_model_call(
|
||||
request, lambda _req: MagicMock()
|
||||
)
|
||||
|
||||
|
||||
def test_hybrid_skips_classification_without_candidates() -> None:
|
||||
"""Preserve all bypassed tools without calling either classifier."""
|
||||
provider_tool = {"type": "web_search"}
|
||||
request = _request([get_weather, provider_tool], [HumanMessage("Help me")])
|
||||
seen: list[ModelRequest[Any]] = []
|
||||
|
||||
def handler(modified: ModelRequest[Any]) -> ModelResponse[Any]:
|
||||
seen.append(modified)
|
||||
return cast("ModelResponse[Any]", MagicMock())
|
||||
|
||||
with patch(
|
||||
"langchain_typesafe.experimental.middleware.tool_selector.TypeSafeClassifier"
|
||||
) as classifier_class:
|
||||
TsHybridToolSelectorMiddleware(always_include=["get_weather"]).wrap_model_call(
|
||||
request, handler
|
||||
)
|
||||
classifier_class.assert_not_called()
|
||||
assert seen[0] is request
|
||||
assert provider_tool in seen[0].tools
|
||||
|
||||
|
||||
def test_hybrid_rejects_missing_always_include() -> None:
|
||||
"""Validate bypassed tool names before either selection stage."""
|
||||
request = _request([get_weather], [HumanMessage("Help me")])
|
||||
with pytest.raises(ValueError, match="not found in request"):
|
||||
TsHybridToolSelectorMiddleware(always_include=["send_email"]).wrap_model_call(
|
||||
request, lambda _req: MagicMock()
|
||||
)
|
||||
|
||||
|
||||
def test_hybrid_validates_threshold() -> None:
|
||||
"""Apply the same threshold validation as the multi-tool selector."""
|
||||
with pytest.raises(ValueError, match="relevance_threshold"):
|
||||
TsHybridToolSelectorMiddleware(relevance_threshold=1.1)
|
||||
|
||||
|
||||
def test_hybrid_multiple_can_keep_zero_tools() -> None:
|
||||
"""Respect the Noul threshold even when the shape is multiple."""
|
||||
first = _classifier(_shape_response("multiple"))
|
||||
second = _classifier(_response({"get_weather": 0.1}))
|
||||
request = _request([get_weather], [HumanMessage("Help me")])
|
||||
seen: list[ModelRequest[Any]] = []
|
||||
|
||||
def handler(modified: ModelRequest[Any]) -> ModelResponse[Any]:
|
||||
seen.append(modified)
|
||||
return cast("ModelResponse[Any]", MagicMock())
|
||||
|
||||
with patch(
|
||||
"langchain_typesafe.experimental.middleware.tool_selector.TypeSafeClassifier",
|
||||
side_effect=[first, second],
|
||||
):
|
||||
TsHybridToolSelectorMiddleware().wrap_model_call(request, handler)
|
||||
assert seen[0].tools == []
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_hybrid_async_rejects_invalid_shape() -> None:
|
||||
"""Reject an invalid stage-one answer before async tool selection."""
|
||||
request = _request([get_weather], [HumanMessage("Help me")])
|
||||
with (
|
||||
patch(
|
||||
"langchain_typesafe.experimental.middleware.tool_selector.TypeSafeClassifier",
|
||||
return_value=_classifier(_shape_response("unexpected")),
|
||||
) as classifier_class,
|
||||
pytest.raises(ValueError, match="no valid tool selection shape"),
|
||||
):
|
||||
await TsHybridToolSelectorMiddleware().awrap_model_call(request, AsyncMock())
|
||||
classifier_class.assert_called_once()
|
||||
|
||||
|
||||
def test_experimental_public_interface() -> None:
|
||||
"""Expose the tool selector from the experimental middleware namespace."""
|
||||
assert middleware_all == [
|
||||
@@ -429,5 +683,6 @@ def test_experimental_public_interface() -> None:
|
||||
"SkillSource",
|
||||
"SkillsMiddleware",
|
||||
"TsChoiceToolSelectorMiddleware",
|
||||
"TsHybridToolSelectorMiddleware",
|
||||
"TsToolSelectorMiddleware",
|
||||
]
|
||||
Reference in new issue
Block a user