mirror of
https://github.com/langchain-ai/langchain.git
synced 2026-10-05 01:15:09 +03:00
feat(typesafe): select SemIf for TsHybridToolSelectorMiddleware
Route the shape and both tool-selection stages to the same opt-in classifier model while keeping the Jev default. Co-authored-by: open-swe[bot] <open-swe@users.noreply.github.com>
This commit is contained in:
1 parent
82aef99fa5
commit
b5ddb60bd2
3 files changed
+70
-14
No files matched your search
@@ -117,7 +117,7 @@ agent = create_agent(
|
||||
)
|
||||
```
|
||||
|
||||
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.
|
||||
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. All three selectors default to Jev; to use SemIf, set `classifier_model="semif-qwen3.5-4b"` on the selector and configure `TYPESAFE_BASE_URL` and `TYPESAFE_API_KEY` for a compatible gateway. This API is experimental and may change without notice.
|
||||
|
||||
### LangChain messages as state
|
||||
|
||||
|
||||
@@ -68,6 +68,8 @@ class TsToolSelectorMiddleware(
|
||||
are kept. No limit if not specified.
|
||||
always_include: Tool names to always include regardless of classification.
|
||||
These do not count against `max_tools` and are not sent to TypeSafe.
|
||||
classifier_model: Model used for classification. Set to `semif-qwen3.5-4b`
|
||||
to use SemIf through a compatible gateway; defaults to Jev.
|
||||
|
||||
Raises:
|
||||
ValueError: If `relevance_threshold` is not between `0` and `1`.
|
||||
@@ -94,6 +96,7 @@ class TsToolSelectorMiddleware(
|
||||
relevance_threshold: float = 0.5,
|
||||
max_tools: int | None = None,
|
||||
always_include: list[str] | None = None,
|
||||
classifier_model: str = "jev-latest",
|
||||
) -> None:
|
||||
"""Initialize the tool selector."""
|
||||
super().__init__()
|
||||
@@ -103,6 +106,7 @@ class TsToolSelectorMiddleware(
|
||||
self.relevance_threshold = relevance_threshold
|
||||
self.max_tools = max_tools
|
||||
self.always_include = always_include or []
|
||||
self.classifier_model = classifier_model
|
||||
|
||||
def _prepare_selection_request(
|
||||
self, request: ModelRequest[ContextT]
|
||||
@@ -161,6 +165,7 @@ class TsToolSelectorMiddleware(
|
||||
def _build_classifier(self, tools: list[BaseTool]) -> TypeSafeClassifier:
|
||||
"""Build a classifier with one `Noul` question per candidate tool."""
|
||||
return TypeSafeClassifier(
|
||||
model=self.classifier_model,
|
||||
questions={
|
||||
f"{_TOOL_QUESTION_PREFIX}{tool.name}": Noul(
|
||||
instructions=(
|
||||
@@ -265,9 +270,10 @@ class TsToolSelectorMiddleware(
|
||||
return await handler(modified_request)
|
||||
|
||||
|
||||
def _build_choice_classifier(tools: list[BaseTool]) -> TypeSafeClassifier:
|
||||
def _build_choice_classifier(tools: list[BaseTool], model: str) -> TypeSafeClassifier:
|
||||
"""Build one categorical question for all candidate tools."""
|
||||
return TypeSafeClassifier(
|
||||
model=model,
|
||||
questions={
|
||||
"tool": Choice(
|
||||
instructions=(
|
||||
@@ -307,6 +313,8 @@ class TsChoiceToolSelectorMiddleware(TsToolSelectorMiddleware):
|
||||
|
||||
Args:
|
||||
always_include: Tool names to include without classification.
|
||||
classifier_model: Model used for classification. Set to `semif-qwen3.5-4b`
|
||||
to use SemIf through a compatible gateway; defaults to Jev.
|
||||
|
||||
??? example "Select one tool per step"
|
||||
|
||||
@@ -319,13 +327,20 @@ class TsChoiceToolSelectorMiddleware(TsToolSelectorMiddleware):
|
||||
```
|
||||
"""
|
||||
|
||||
def __init__(self, *, always_include: list[str] | None = None) -> None:
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
always_include: list[str] | None = None,
|
||||
classifier_model: str = "jev-latest",
|
||||
) -> None:
|
||||
"""Initialize the choice-based tool selector."""
|
||||
super().__init__(always_include=always_include)
|
||||
super().__init__(
|
||||
always_include=always_include, classifier_model=classifier_model
|
||||
)
|
||||
|
||||
def _build_classifier(self, tools: list[BaseTool]) -> TypeSafeClassifier:
|
||||
"""Build one categorical question for all candidate tools."""
|
||||
return _build_choice_classifier(tools)
|
||||
return _build_choice_classifier(tools, self.classifier_model)
|
||||
|
||||
def _select_tool_names(
|
||||
self, response: ClassificationResponse, selection_request: _SelectionRequest
|
||||
@@ -354,6 +369,8 @@ class TsHybridToolSelectorMiddleware(TsToolSelectorMiddleware):
|
||||
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.
|
||||
classifier_model: Model used at both selection stages. Set to
|
||||
`semif-qwen3.5-4b` to use SemIf through a compatible gateway.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
@@ -362,17 +379,20 @@ class TsHybridToolSelectorMiddleware(TsToolSelectorMiddleware):
|
||||
relevance_threshold: float = 0.5,
|
||||
max_tools: int | None = None,
|
||||
always_include: list[str] | None = None,
|
||||
classifier_model: str = "jev-latest",
|
||||
) -> None:
|
||||
"""Initialize the hybrid selector with multi-tool selection settings."""
|
||||
super().__init__(
|
||||
relevance_threshold=relevance_threshold,
|
||||
max_tools=max_tools,
|
||||
always_include=always_include,
|
||||
classifier_model=classifier_model,
|
||||
)
|
||||
|
||||
def _build_shape_classifier(self) -> TypeSafeClassifier:
|
||||
"""Build the question that decides how many tools the next step needs."""
|
||||
return TypeSafeClassifier(
|
||||
model=self.classifier_model,
|
||||
questions={
|
||||
"shape": Choice(
|
||||
instructions=(
|
||||
@@ -386,7 +406,7 @@ class TsHybridToolSelectorMiddleware(TsToolSelectorMiddleware):
|
||||
"multiple": "Several tools may be needed for the next step.",
|
||||
},
|
||||
)
|
||||
}
|
||||
},
|
||||
)
|
||||
|
||||
def _shape(self, response: ClassificationResponse) -> str:
|
||||
@@ -426,7 +446,9 @@ class TsHybridToolSelectorMiddleware(TsToolSelectorMiddleware):
|
||||
selected = []
|
||||
if shape != "none":
|
||||
classifier = (
|
||||
_build_choice_classifier(selection_request.classifiable_tools)
|
||||
_build_choice_classifier(
|
||||
selection_request.classifiable_tools, self.classifier_model
|
||||
)
|
||||
if shape == "single"
|
||||
else self._build_classifier(selection_request.classifiable_tools)
|
||||
)
|
||||
@@ -458,7 +480,9 @@ class TsHybridToolSelectorMiddleware(TsToolSelectorMiddleware):
|
||||
selected = []
|
||||
if shape != "none":
|
||||
classifier = (
|
||||
_build_choice_classifier(selection_request.classifiable_tools)
|
||||
_build_choice_classifier(
|
||||
selection_request.classifiable_tools, self.classifier_model
|
||||
)
|
||||
if shape == "single"
|
||||
else self._build_classifier(selection_request.classifiable_tools)
|
||||
)
|
||||
|
||||
+38
-6
@@ -124,6 +124,30 @@ def test_middleware_constructs_classifier_per_call() -> None:
|
||||
assert "current weather" in questions["tool::get_weather"].instructions
|
||||
|
||||
|
||||
@pytest.mark.parametrize("model", ["jev-latest", "semif-qwen3.5-4b"])
|
||||
@pytest.mark.parametrize(
|
||||
"selector", [TsToolSelectorMiddleware, TsChoiceToolSelectorMiddleware]
|
||||
)
|
||||
def test_classifier_uses_requested_model(model: str, selector: type) -> None:
|
||||
"""Route Noul and Choice selection to the configured classifier model."""
|
||||
response = (
|
||||
_choice_response("get_weather")
|
||||
if selector is TsChoiceToolSelectorMiddleware
|
||||
else _response({"get_weather": 0.9})
|
||||
)
|
||||
classifier = _classifier(response)
|
||||
request = _request([get_weather], [HumanMessage("What's the weather?")])
|
||||
middleware = selector(classifier_model=model)
|
||||
|
||||
with patch(
|
||||
"langchain_typesafe.experimental.middleware.tool_selector.TypeSafeClassifier",
|
||||
return_value=classifier,
|
||||
) as classifier_class:
|
||||
middleware.wrap_model_call(request, lambda _req: MagicMock())
|
||||
|
||||
assert classifier_class.call_args.kwargs["model"] == model
|
||||
|
||||
|
||||
def test_sync_selection_filters_tools_above_threshold() -> None:
|
||||
"""Keep only tools whose `Noul` probability clears the threshold."""
|
||||
classifier = _classifier(
|
||||
@@ -438,7 +462,8 @@ def test_choice_only_always_included_tools_skips_classifier() -> None:
|
||||
|
||||
|
||||
@pytest.mark.parametrize("shape", ["none", "single", "multiple"])
|
||||
def test_hybrid_routes_sync(shape: str) -> None:
|
||||
@pytest.mark.parametrize("model", ["jev-latest", "semif-qwen3.5-4b"])
|
||||
def test_hybrid_routes_sync(shape: str, model: str) -> None:
|
||||
"""Classify shape once, then invoke only the needed selection stage."""
|
||||
first = _classifier(_shape_response(shape))
|
||||
second = _classifier(
|
||||
@@ -451,7 +476,7 @@ def test_hybrid_routes_sync(shape: str) -> None:
|
||||
[HumanMessage("Help me"), AIMessage("Next step")],
|
||||
)
|
||||
middleware = TsHybridToolSelectorMiddleware(
|
||||
max_tools=1, always_include=["get_weather"]
|
||||
max_tools=1, always_include=["get_weather"], classifier_model=model
|
||||
)
|
||||
seen: list[ModelRequest[Any]] = []
|
||||
|
||||
@@ -466,6 +491,9 @@ def test_hybrid_routes_sync(shape: str) -> None:
|
||||
middleware.wrap_model_call(request, handler)
|
||||
|
||||
assert classifier_class.call_count == (1 if shape == "none" else 2)
|
||||
assert all(
|
||||
call.kwargs["model"] == model for call in classifier_class.call_args_list
|
||||
)
|
||||
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.",
|
||||
@@ -499,7 +527,8 @@ def test_hybrid_routes_sync(shape: str) -> None:
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("shape", ["none", "single", "multiple"])
|
||||
async def test_hybrid_routes_async(shape: str) -> None:
|
||||
@pytest.mark.parametrize("model", ["jev-latest", "semif-qwen3.5-4b"])
|
||||
async def test_hybrid_routes_async(shape: str, model: str) -> None:
|
||||
"""Use the asynchronous classifier path for every shape."""
|
||||
first = _classifier(_shape_response(shape))
|
||||
second = _classifier(
|
||||
@@ -518,11 +547,14 @@ async def test_hybrid_routes_async(shape: str) -> None:
|
||||
"langchain_typesafe.experimental.middleware.tool_selector.TypeSafeClassifier",
|
||||
side_effect=[first, second],
|
||||
) as classifier_class:
|
||||
await TsHybridToolSelectorMiddleware(max_tools=1).awrap_model_call(
|
||||
request, handler
|
||||
)
|
||||
await TsHybridToolSelectorMiddleware(
|
||||
max_tools=1, classifier_model=model
|
||||
).awrap_model_call(request, handler)
|
||||
|
||||
assert classifier_class.call_count == (1 if shape == "none" else 2)
|
||||
assert all(
|
||||
call.kwargs["model"] == model for call in classifier_class.call_args_list
|
||||
)
|
||||
first.ainvoke.assert_awaited_once()
|
||||
assert _tool_names(seen[0].tools) == ([] if shape == "none" else ["search_web"])
|
||||
if shape == "none":
|
||||
|
||||
Reference in new issue
Block a user