feat(typesafe): add choice-based 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-23 21:36:44 +00:00
1 parent e9ac94d9f6
commit 4fe5602ea3
4 files changed
+211 -11

No files matched your search

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