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:
Thushanth Bengreandopen-swe[bot] committed 2026-09-24 14:26:04 +00:00
1 parent 82aef99fa5
commit b5ddb60bd2
3 files changed
+70 -14

No files matched your search

+1 -1
View File
@@ -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)
)
@@ -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":