mirror of
https://github.com/langchain-ai/langchain.git
synced 2026-10-05 17:35:28 +03:00
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:
1 parent
e9ac94d9f6
commit
4fe5602ea3
4 files changed
+211
-11
No files matched your search
@@ -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"]
|
||||
+119
-2
@@ -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",
|
||||
]
|
||||
Reference in new issue
Block a user