mirror of
https://github.com/langchain-ai/langchain.git
synced 2026-10-05 01:15:09 +03:00
feat(typesafe): select SemIf for TsToolSelectorMiddleware
Allow a classifier model override while retaining Jev as the default. Co-authored-by: open-swe[bot] <open-swe@users.noreply.github.com>
This commit is contained in:
1 parent
e9ac94d9f6
commit
64bdc86998
3 files changed
+22
-1
No files matched your search
@@ -87,7 +87,7 @@ 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.
|
||||
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. To use SemIf instead of the default Jev model, set `classifier_model="semif-qwen3.5-4b"` and configure `TYPESAFE_BASE_URL` and `TYPESAFE_API_KEY` for a gateway that supports SemIf. This API is experimental and may change without notice.
|
||||
|
||||
### LangChain messages as state
|
||||
|
||||
|
||||
@@ -67,6 +67,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`.
|
||||
@@ -93,6 +95,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__()
|
||||
@@ -102,6 +105,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]
|
||||
@@ -160,6 +164,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=(
|
||||
|
||||
@@ -92,6 +92,22 @@ 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"])
|
||||
def test_classifier_uses_requested_model(model: str) -> None:
|
||||
"""Route tool classification to the configured model without changing selection."""
|
||||
classifier = _classifier(_response({"get_weather": 0.9}))
|
||||
request = _request([get_weather], [HumanMessage("What's the weather?")])
|
||||
middleware = TsToolSelectorMiddleware(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(
|
||||
|
||||
Reference in new issue
Block a user