From 64bdc8699844c6671277c663d98aac16bc0f4042 Mon Sep 17 00:00:00 2001 From: Thushanth Bengre Date: Thu, 24 Sep 2026 14:18:43 +0000 Subject: [PATCH] feat(typesafe): select SemIf for `TsToolSelectorMiddleware` Allow a classifier model override while retaining Jev as the default. Co-authored-by: open-swe[bot] --- libs/partners/typesafe/README.md | 2 +- .../experimental/middleware/tool_selector.py | 5 +++++ .../middleware/test_tool_selector.py | 16 ++++++++++++++++ 3 files changed, 22 insertions(+), 1 deletion(-) diff --git a/libs/partners/typesafe/README.md b/libs/partners/typesafe/README.md index 7f181dff79..caae3daedc 100644 --- a/libs/partners/typesafe/README.md +++ b/libs/partners/typesafe/README.md @@ -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 diff --git a/libs/partners/typesafe/langchain_typesafe/experimental/middleware/tool_selector.py b/libs/partners/typesafe/langchain_typesafe/experimental/middleware/tool_selector.py index 810d95e293..e56f44515f 100644 --- a/libs/partners/typesafe/langchain_typesafe/experimental/middleware/tool_selector.py +++ b/libs/partners/typesafe/langchain_typesafe/experimental/middleware/tool_selector.py @@ -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=( diff --git a/libs/partners/typesafe/tests/unit_tests/experimental/middleware/test_tool_selector.py b/libs/partners/typesafe/tests/unit_tests/experimental/middleware/test_tool_selector.py index 8b01dcc33c..2e86c07e04 100644 --- a/libs/partners/typesafe/tests/unit_tests/experimental/middleware/test_tool_selector.py +++ b/libs/partners/typesafe/tests/unit_tests/experimental/middleware/test_tool_selector.py @@ -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(