From bc5c853b34a57e6d30eadad463bf3def4c5493fd Mon Sep 17 00:00:00 2001 From: Thushanth Bengre Date: Thu, 24 Sep 2026 02:17:37 +0000 Subject: [PATCH] feat(typesafe): add hybrid tool selector middleware Co-authored-by: open-swe[bot] --- libs/partners/typesafe/README.md | 16 ++ .../experimental/middleware/__init__.py | 2 + .../experimental/middleware/tool_selector.py | 187 +++++++++++-- .../middleware/test_tool_selector.py | 255 ++++++++++++++++++ 4 files changed, 443 insertions(+), 17 deletions(-) diff --git a/libs/partners/typesafe/README.md b/libs/partners/typesafe/README.md index 19ce6284e4..2d855a0149 100644 --- a/libs/partners/typesafe/README.md +++ b/libs/partners/typesafe/README.md @@ -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: diff --git a/libs/partners/typesafe/langchain_typesafe/experimental/middleware/__init__.py b/libs/partners/typesafe/langchain_typesafe/experimental/middleware/__init__.py index 1bb3ca5a86..7ed7b26195 100644 --- a/libs/partners/typesafe/langchain_typesafe/experimental/middleware/__init__.py +++ b/libs/partners/typesafe/langchain_typesafe/experimental/middleware/__init__.py @@ -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", ] diff --git a/libs/partners/typesafe/langchain_typesafe/experimental/middleware/tool_selector.py b/libs/partners/typesafe/langchain_typesafe/experimental/middleware/tool_selector.py index 38809c67bc..5b38f5fd50 100644 --- a/libs/partners/typesafe/langchain_typesafe/experimental/middleware/tool_selector.py +++ b/libs/partners/typesafe/langchain_typesafe/experimental/middleware/tool_selector.py @@ -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", +] diff --git a/libs/partners/typesafe/tests/unit_tests/experimental/middleware/test_tool_selector.py b/libs/partners/typesafe/tests/unit_tests/experimental/middleware/test_tool_selector.py index a5fc7bc634..cce522d58a 100644 --- a/libs/partners/typesafe/tests/unit_tests/experimental/middleware/test_tool_selector.py +++ b/libs/partners/typesafe/tests/unit_tests/experimental/middleware/test_tool_selector.py @@ -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", ]