diff --git a/libs/langchain_v1/langchain/agents/middleware/model_routing.py b/libs/langchain_v1/langchain/agents/middleware/model_routing.py index 315b6049bf..1e21fe9513 100644 --- a/libs/langchain_v1/langchain/agents/middleware/model_routing.py +++ b/libs/langchain_v1/langchain/agents/middleware/model_routing.py @@ -4,7 +4,7 @@ from __future__ import annotations import json import logging -from typing import TYPE_CHECKING, cast +from typing import TYPE_CHECKING from langchain_core._api import beta from langchain_core.messages import BaseMessage, HumanMessage, SystemMessage @@ -72,8 +72,8 @@ class ModelRoutingMiddleware(AgentMiddleware[ModelRoutingState, ContextT, Respon Routes are selected before the first model call and persisted in `model_route`, keeping the same model throughout tool loops and checkpoint resumes. Clear this state field (set it to `None`) to route a new task. Selection is never cached on - the middleware instance. Applications can override `select_route` or - `aselect_route` for application-specific selection policies. + the middleware instance. Plug this middleware into `create_agent`; applications + can customize routing input and backends without invoking selection directly. Place this middleware before model fallback middleware so fallbacks receive the selected model. Candidate models must support the agent's tools and output format. @@ -212,7 +212,7 @@ class ModelRoutingMiddleware(AgentMiddleware[ModelRoutingState, ContextT, Respon msg = "Model routing selection must be a configured route name" raise ValueError(msg) - def select_route(self, state: ModelRoutingState) -> str: + def _select_route(self, state: ModelRoutingState) -> str: """Select or reuse a route without mutating agent state. Args: @@ -233,7 +233,7 @@ class ModelRoutingMiddleware(AgentMiddleware[ModelRoutingState, ContextT, Respon response = self._routing_model.invoke(self._llm_input(inputs), config=config) return self._validate_route(response.get("route") if isinstance(response, dict) else None) - async def aselect_route(self, state: ModelRoutingState) -> str: + async def _aselect_route(self, state: ModelRoutingState) -> str: """Select or reuse a route asynchronously without mutating agent state. Args: @@ -265,7 +265,7 @@ class ModelRoutingMiddleware(AgentMiddleware[ModelRoutingState, ContextT, Respon The selected route state update. """ del runtime - return {"model_route": self.select_route(state)} + return {"model_route": self._select_route(state)} async def abefore_model( self, state: ModelRoutingState, runtime: Runtime[ContextT] @@ -280,7 +280,7 @@ class ModelRoutingMiddleware(AgentMiddleware[ModelRoutingState, ContextT, Respon The selected route state update. """ del runtime - return {"model_route": await self.aselect_route(state)} + return {"model_route": await self._aselect_route(state)} def wrap_model_call( self, @@ -296,7 +296,7 @@ class ModelRoutingMiddleware(AgentMiddleware[ModelRoutingState, ContextT, Respon Returns: The selected model's response. """ - route = self.select_route(cast("ModelRoutingState", request.state)) + route = self._validate_route(request.state.get("model_route")) return handler(request.override(model=self.models[route])) async def awrap_model_call( @@ -313,5 +313,5 @@ class ModelRoutingMiddleware(AgentMiddleware[ModelRoutingState, ContextT, Respon Returns: The selected model's response. """ - route = await self.aselect_route(cast("ModelRoutingState", request.state)) + route = self._validate_route(request.state.get("model_route")) return await handler(request.override(model=self.models[route])) diff --git a/libs/langchain_v1/tests/unit_tests/agents/middleware/implementations/test_model_routing.py b/libs/langchain_v1/tests/unit_tests/agents/middleware/implementations/test_model_routing.py index c09dab962c..aae709659e 100644 --- a/libs/langchain_v1/tests/unit_tests/agents/middleware/implementations/test_model_routing.py +++ b/libs/langchain_v1/tests/unit_tests/agents/middleware/implementations/test_model_routing.py @@ -10,6 +10,7 @@ from langchain_core.language_models import LanguageModelInput from langchain_core.language_models.fake_chat_models import FakeListChatModel from langchain_core.messages import AIMessage, BaseMessage, HumanMessage, SystemMessage from langchain_core.runnables import Runnable, RunnableConfig, RunnableLambda +from langgraph.runtime import Runtime from pydantic import BaseModel, Field from typing_extensions import override @@ -105,6 +106,7 @@ async def test_route_and_preserve_request(backend: str, *, async_mode: bool) -> }, } criteria = {"small": "Simple tasks", "large": "Complex tasks"} + selection_modes: list[str] = [] router = RoutingChatModel(responses=[""]) def classify(inputs: ModelRoutingInput, config: RunnableConfig) -> str: @@ -114,13 +116,21 @@ async def test_route_and_preserve_request(backend: str, *, async_mode: bool) -> "criteria": criteria, } assert config["metadata"].items() >= internal_call_metadata().items() + selection_modes.append("sync") return "large" + async def aclassify(inputs: ModelRoutingInput, config: RunnableConfig) -> str: + result = classify(inputs, config) + selection_modes[-1] = "async" + return result + middleware = ModelRoutingMiddleware( models=models, system_prompt="Custom routing instructions", routing_model=router if backend == "llm" else None, - decision_model=RunnableLambda(classify) if backend == "classifier" else None, + decision_model=RunnableLambda(classify, afunc=aclassify) + if backend == "classifier" + else None, ) def handler(request: ModelRequest) -> ModelResponse: @@ -131,22 +141,22 @@ async def test_route_and_preserve_request(backend: str, *, async_mode: bool) -> check_request(request, original) return ModelResponse(result=[await request.model.ainvoke(request.messages)]) - original = original.override( - state=ModelRoutingState( - messages=original.messages, - model_route=( - await middleware.abefore_model( - ModelRoutingState(messages=original.messages), original.runtime - ) - )["model_route"], - ) + state = ModelRoutingState(messages=original.messages) + update = ( + await middleware.abefore_model(state, Runtime()) + if async_mode + else middleware.before_model(state, Runtime()) ) + state["model_route"] = update["model_route"] + original = original.override(state=state) response = ( await middleware.awrap_model_call(original, ahandler) if async_mode else middleware.wrap_model_call(original, handler) ) assert response.result[0].content == "large response" + if backend == "classifier": + assert selection_modes == ["async" if async_mode else "sync"] if backend == "llm": assert router.routing_schema is not None assert router.routing_schema["properties"]["route"]["enum"] == ["small", "large"] @@ -173,15 +183,19 @@ async def test_invalid_selection_and_explicit_fallback(backend: str, selection: decision_model=RunnableLambda(classify) if backend == "classifier" else None, ) with pytest.raises(ValueError, match="configured route name"): - middleware.select_route(ModelRoutingState(messages=make_request().messages)) + middleware.before_model(ModelRoutingState(messages=make_request().messages), Runtime()) with pytest.raises(ValueError, match="configured route name"): - await middleware.aselect_route(ModelRoutingState(messages=make_request().messages)) + await middleware.abefore_model( + ModelRoutingState(messages=make_request().messages), Runtime() + ) middleware.fallback_route = "large" original = make_request() original = original.override( state=ModelRoutingState( messages=original.messages, - model_route=middleware.select_route(ModelRoutingState(messages=original.messages)), + model_route=middleware.before_model( + ModelRoutingState(messages=original.messages), Runtime() + )["model_route"], ) ) @@ -217,9 +231,18 @@ async def test_custom_input_and_no_cross_request_cache() -> None: messages=[HumanMessage(content="Lookup"), HumanMessage(content="Injected context")] ) second = make_request() - assert middleware.select_route(ModelRoutingState(messages=first.messages)) == "small" - assert await middleware.aselect_route(ModelRoutingState(messages=second.messages)) == "large" - assert await middleware.aselect_route(ModelRoutingState(messages=first.messages)) == "small" + assert ( + middleware.before_model(ModelRoutingState(messages=first.messages), Runtime())[ + "model_route" + ] + == "small" + ) + assert (await middleware.abefore_model(ModelRoutingState(messages=second.messages), Runtime()))[ + "model_route" + ] == "large" + assert (await middleware.abefore_model(ModelRoutingState(messages=first.messages), Runtime()))[ + "model_route" + ] == "small" @pytest.mark.parametrize( @@ -256,11 +279,15 @@ async def test_backend_errors_propagate_and_missing_input_is_explicit() -> None: fallback_route="small", ) with pytest.raises(RuntimeError, match="Classifier unavailable"): - middleware.select_route(ModelRoutingState(messages=make_request().messages)) + middleware.before_model(ModelRoutingState(messages=make_request().messages), Runtime()) with pytest.raises(RuntimeError, match="Classifier unavailable"): - await middleware.aselect_route(ModelRoutingState(messages=make_request().messages)) + await middleware.abefore_model( + ModelRoutingState(messages=make_request().messages), Runtime() + ) with pytest.raises(ValueError, match="No routing messages"): - middleware.select_route(ModelRoutingState(messages=[AIMessage(content="No user")])) + middleware.before_model( + ModelRoutingState(messages=[AIMessage(content="No user")]), Runtime() + ) async def test_agent_uses_selected_model() -> None: @@ -296,8 +323,8 @@ async def test_route_persists_until_cleared(*, async_mode: bool) -> None: async def prepare() -> str: if async_mode: - return (await middleware.abefore_model(state, None))["model_route"] # type: ignore[arg-type] - return middleware.before_model(state, None)["model_route"] # type: ignore[arg-type] + return (await middleware.abefore_model(state, Runtime()))["model_route"] + return middleware.before_model(state, Runtime())["model_route"] state["model_route"] = await prepare() assert state["model_route"] == "small"