feat(typesafe): add hybrid tool selector middleware

Co-authored-by: open-swe[bot] <open-swe@users.noreply.github.com>
This commit is contained in:
Thushanth Bengreandopen-swe[bot] committed 2026-09-24 02:17:37 +00:00
1 parent 35a8757e92
commit bc5c853b34
4 files changed
+443 -17

No files matched your search

+16
View File
@@ -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",
]
@@ -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",
]
@@ -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",
]