diff --git a/libs/partners/typesafe/README.md b/libs/partners/typesafe/README.md index 2d855a0149..60937913fa 100644 --- a/libs/partners/typesafe/README.md +++ b/libs/partners/typesafe/README.md @@ -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 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 5b38f5fd50..e3ffdff527 100644 --- a/libs/partners/typesafe/langchain_typesafe/experimental/middleware/tool_selector.py +++ b/libs/partners/typesafe/langchain_typesafe/experimental/middleware/tool_selector.py @@ -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) ) 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 cce522d58a..94e70bbc7a 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 @@ -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":