From 4fe5602ea365489012de4e53cf6a995eb858e458 Mon Sep 17 00:00:00 2001 From: Thushanth Bengre Date: Wed, 23 Sep 2026 21:36:44 +0000 Subject: [PATCH] feat(typesafe): add choice-based 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 | 83 ++++++++++-- .../middleware/test_tool_selector.py | 121 +++++++++++++++++- 4 files changed, 211 insertions(+), 11 deletions(-) diff --git a/libs/partners/typesafe/README.md b/libs/partners/typesafe/README.md index 7f181dff79..19ce6284e4 100644 --- a/libs/partners/typesafe/README.md +++ b/libs/partners/typesafe/README.md @@ -87,7 +87,21 @@ agent = create_agent( ) ``` -The middleware asks one independent `Noul` question per candidate tool ("is this tool needed next?"), batched into a single TypeSafe request against the latest human message, before every model call. Tools whose probability clears `relevance_threshold` (default `0.5`) are kept, ranked by that probability, and capped at `max_tools` if set. Use `always_include` to keep specific tools regardless of classification. This API is experimental and may change without notice. +`TsToolSelectorMiddleware` asks one independent `Noul` question per candidate tool ("is this tool needed next?"), batched into a single TypeSafe request against the latest human message, before every model call. Tools whose probability clears `relevance_threshold` (default `0.5`) are kept, ranked by that probability, and capped at `max_tools` if set. + +For a single best tool at each model step, use `TsChoiceToolSelectorMiddleware` instead: + +```python +from langchain_typesafe.experimental.middleware import TsChoiceToolSelectorMiddleware + +agent = create_agent( + model, + tools=[tool1, tool2, tool3], + middleware=[TsChoiceToolSelectorMiddleware()], +) +``` + +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. ### LangChain messages as state diff --git a/libs/partners/typesafe/langchain_typesafe/experimental/middleware/__init__.py b/libs/partners/typesafe/langchain_typesafe/experimental/middleware/__init__.py index 986d44e48f..1bb3ca5a86 100644 --- a/libs/partners/typesafe/langchain_typesafe/experimental/middleware/__init__.py +++ b/libs/partners/typesafe/langchain_typesafe/experimental/middleware/__init__.py @@ -6,6 +6,7 @@ from langchain_typesafe.experimental.middleware.skills import ( SkillSource, ) from langchain_typesafe.experimental.middleware.tool_selector import ( + TsChoiceToolSelectorMiddleware, TsToolSelectorMiddleware, ) @@ -13,5 +14,6 @@ __all__ = [ "Skill", "SkillSource", "SkillsMiddleware", + "TsChoiceToolSelectorMiddleware", "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 810d95e293..38809c67bc 100644 --- a/libs/partners/typesafe/langchain_typesafe/experimental/middleware/tool_selector.py +++ b/libs/partners/typesafe/langchain_typesafe/experimental/middleware/tool_selector.py @@ -3,7 +3,7 @@ from __future__ import annotations from dataclasses import dataclass -from typing import TYPE_CHECKING, Any +from typing import TYPE_CHECKING from langchain.agents.middleware.types import ( AgentMiddleware, @@ -14,10 +14,11 @@ from langchain.agents.middleware.types import ( ResponseT, ) from langchain_core.messages import HumanMessage +from langchain_core.runnables import RunnableConfig from typing_extensions import override from langchain_typesafe.classifier import TypeSafeClassifier -from langchain_typesafe.types import Noul +from langchain_typesafe.types import Choice, ClassificationResponse, Noul if TYPE_CHECKING: from collections.abc import Awaitable, Callable @@ -173,9 +174,10 @@ class TsToolSelectorMiddleware( ) def _select_tool_names( - self, response_nouls: dict[str, Any], selection_request: _SelectionRequest + self, response: ClassificationResponse, selection_request: _SelectionRequest ) -> list[str]: """Return classifiable tool names above threshold, ranked by probability.""" + response_nouls = response.nouls scored = [ (name, response_nouls[f"{_TOOL_QUESTION_PREFIX}{name}"].noul) for name in selection_request.valid_tool_names @@ -189,6 +191,10 @@ class TsToolSelectorMiddleware( selected = selected[: self.max_tools] return selected + def _classifier_config(self) -> RunnableConfig: + """Tag classifier calls for tool selection traces.""" + return {"metadata": {"lc_source": "ts_tool_selector"}} + def _process_selection( self, selected_tool_names: list[str], @@ -226,9 +232,9 @@ class TsToolSelectorMiddleware( classifier = self._build_classifier(selection_request.classifiable_tools) response = classifier.invoke( selection_request.last_user_message, - config={"metadata": {"lc_source": "ts_tool_selector"}}, + config=self._classifier_config(), ) - selected_tool_names = self._select_tool_names(response.nouls, selection_request) + selected_tool_names = self._select_tool_names(response, selection_request) modified_request = self._process_selection( selected_tool_names, selection_request, request ) @@ -250,13 +256,74 @@ class TsToolSelectorMiddleware( classifier = self._build_classifier(selection_request.classifiable_tools) response = await classifier.ainvoke( selection_request.last_user_message, - config={"metadata": {"lc_source": "ts_tool_selector"}}, + config=self._classifier_config(), ) - selected_tool_names = self._select_tool_names(response.nouls, selection_request) + selected_tool_names = self._select_tool_names(response, selection_request) modified_request = self._process_selection( selected_tool_names, selection_request, request ) return await handler(modified_request) -__all__ = ["TsToolSelectorMiddleware"] +class TsChoiceToolSelectorMiddleware(TsToolSelectorMiddleware): + """Select one candidate tool for the next model call with a TypeSafe `Choice`. + + Unlike `TsToolSelectorMiddleware`, this variant asks one categorical question + across all available candidate tools, rather than scoring each tool separately. + The choice is made again before each model call. `always_include` tools and + provider-specific tool definitions are passed through in addition to the chosen + tool. Classifier failures and invalid choices raise instead of silently changing + tool availability. + + !!! warning + + This middleware is experimental. Its API may change without notice. + + Args: + always_include: Tool names to include without classification. + + ??? example "Select one tool per step" + + ```python + from langchain_typesafe.experimental.middleware import ( + TsChoiceToolSelectorMiddleware, + ) + + middleware = TsChoiceToolSelectorMiddleware() + ``` + """ + + def __init__(self, *, always_include: list[str] | None = None) -> None: + """Initialize the choice-based tool selector.""" + super().__init__(always_include=always_include) + + 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}, + ) + }, + ) + + 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] + + def _classifier_config(self) -> RunnableConfig: + """Tag choice selector calls separately in traces.""" + return {"metadata": {"lc_source": "ts_choice_tool_selector"}} + + +__all__ = ["TsChoiceToolSelectorMiddleware", "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 8b01dcc33c..a5fc7bc634 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 @@ -11,8 +11,11 @@ from langchain_core.language_models import BaseChatModel from langchain_core.messages import AIMessage, HumanMessage from langchain_core.tools import tool -from langchain_typesafe import NoulAnswer, TypeSafeClassifier -from langchain_typesafe.experimental.middleware import TsToolSelectorMiddleware +from langchain_typesafe import ChoiceAnswer, NoulAnswer, TypeSafeClassifier +from langchain_typesafe.experimental.middleware import ( + TsChoiceToolSelectorMiddleware, + TsToolSelectorMiddleware, +) from langchain_typesafe.experimental.middleware import __all__ as middleware_all from langchain_typesafe.types import ClassificationResponse @@ -45,6 +48,20 @@ def _response(scores: dict[str, float]) -> ClassificationResponse: ) +def _choice_response(choice: str) -> ClassificationResponse: + return ClassificationResponse( + model="jev-latest", + answers={ + "tool": ChoiceAnswer( + type="choice", + choice=choice, + probabilities={choice: 1.0}, + confidence=1.0, + ) + }, + ) + + def _classifier(response: ClassificationResponse) -> MagicMock: classifier = MagicMock(spec=TypeSafeClassifier) classifier.invoke.return_value = response @@ -306,11 +323,111 @@ async def test_async_selection_filters_tools() -> None: assert _tool_names(seen[0].tools) == ["get_weather"] +def test_choice_selects_one_tool_per_model_call() -> None: + """Choose one candidate afresh for every model call, not the entire task.""" + classifier = _classifier(_choice_response("get_weather")) + middleware = TsChoiceToolSelectorMiddleware() + first = _request([get_weather, search_web], [HumanMessage("Find the weather")]) + second = _request( + [get_weather, search_web], + [HumanMessage("Find the weather"), AIMessage("Next search the web")], + ) + 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", + return_value=classifier, + ) as classifier_class: + middleware.wrap_model_call(first, handler) + classifier.invoke.return_value = _choice_response("search_web") + middleware.wrap_model_call(second, handler) + + assert classifier_class.call_count == 2 + question = classifier_class.call_args.kwargs["questions"]["tool"] + assert question.criteria == { + "get_weather": get_weather.description, + "search_web": search_web.description, + } + assert "next" in question.instructions + assert classifier.invoke.call_args.kwargs["config"]["metadata"] == { + "lc_source": "ts_choice_tool_selector" + } + assert _tool_names(seen[0].tools) == ["get_weather"] + assert _tool_names(seen[1].tools) == ["search_web"] + assert first.tools == [get_weather, search_web] + + +@pytest.mark.asyncio +async def test_choice_async_preserves_always_included_and_provider_tools() -> None: + """Keep bypassed tools while asynchronously choosing one candidate.""" + classifier = _classifier(_choice_response("search_web")) + provider_tool = {"type": "web_search"} + request = _request( + [get_weather, search_web, send_email, provider_tool], + [HumanMessage("Search for the latest weather")], + ) + middleware = TsChoiceToolSelectorMiddleware(always_include=["get_weather"]) + 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", + return_value=classifier, + ) as classifier_class: + await middleware.awrap_model_call(request, handler) + + assert classifier.ainvoke.await_count == 1 + assert classifier_class.call_args.kwargs["questions"]["tool"].criteria == { + "search_web": search_web.description, + "send_email": send_email.description, + } + assert _tool_names(seen[0].tools) == ["search_web", "get_weather"] + assert provider_tool in seen[0].tools + + +@pytest.mark.parametrize("response", [_choice_response("unknown"), _response({})]) +def test_choice_rejects_invalid_response(response: ClassificationResponse) -> None: + """Never substitute an unavailable or missing tool silently.""" + classifier = _classifier(response) + request = _request([get_weather], [HumanMessage("Get the weather")]) + + with ( + patch( + "langchain_typesafe.experimental.middleware.tool_selector.TypeSafeClassifier", + return_value=classifier, + ), + pytest.raises(ValueError, match="no valid tool choice"), + ): + TsChoiceToolSelectorMiddleware().wrap_model_call( + request, lambda _req: MagicMock() + ) + + +def test_choice_only_always_included_tools_skips_classifier() -> None: + """Skip classification if no candidate tools remain.""" + request = _request([get_weather], [HumanMessage("Get the weather")]) + with patch( + "langchain_typesafe.experimental.middleware.tool_selector.TypeSafeClassifier" + ) as classifier_class: + TsChoiceToolSelectorMiddleware(always_include=["get_weather"]).wrap_model_call( + request, lambda _req: MagicMock() + ) + classifier_class.assert_not_called() + + def test_experimental_public_interface() -> None: """Expose the tool selector from the experimental middleware namespace.""" assert middleware_all == [ "Skill", "SkillSource", "SkillsMiddleware", + "TsChoiceToolSelectorMiddleware", "TsToolSelectorMiddleware", ]