diff --git a/libs/langchain_v1/langchain/agents/middleware/model_routing.py b/libs/langchain_v1/langchain/agents/middleware/model_routing.py index 1e21fe9513..2914a8e35a 100644 --- a/libs/langchain_v1/langchain/agents/middleware/model_routing.py +++ b/libs/langchain_v1/langchain/agents/middleware/model_routing.py @@ -4,11 +4,12 @@ from __future__ import annotations import json import logging -from typing import TYPE_CHECKING +from typing import TYPE_CHECKING, cast from langchain_core._api import beta +from langchain_core.language_models import BaseChatModel from langchain_core.messages import BaseMessage, HumanMessage, SystemMessage -from typing_extensions import NotRequired, TypedDict +from typing_extensions import NotRequired, TypedDict, override from langchain.agents.middleware.internal_call_transformer import ( InternalCallTransformer, @@ -27,13 +28,12 @@ from langchain.chat_models import init_chat_model if TYPE_CHECKING: from collections.abc import Awaitable, Callable, Mapping, Sequence - from langchain_core.language_models import BaseChatModel, LanguageModelInput - from langchain_core.runnables import Runnable, RunnableConfig + from langchain_core.runnables import Runnable from langgraph.runtime import Runtime logger = logging.getLogger(__name__) -DEFAULT_SYSTEM_PROMPT = "Choose the least expensive model likely to complete the user's task." +DEFAULT_INSTRUCTIONS = "Choose the least expensive model likely to complete the user's task." class ModelRoutingConfig(TypedDict): @@ -47,10 +47,20 @@ class ModelRoutingInput(TypedDict): """Beta, experimental classifier input; no compatibility guarantees.""" messages: list[BaseMessage] - system_prompt: str + instructions: str criteria: dict[str, str] +class ModelRoutingOutput(TypedDict): + """Select a model route for the user's task. + + Args: + route: The selected model route. + """ + + route: str + + class ModelRoutingState(AgentState): """Experimental checkpointed route; clear `model_route` to select again.""" @@ -77,7 +87,8 @@ class ModelRoutingMiddleware(AgentMiddleware[ModelRoutingState, ContextT, Respon Place this middleware before model fallback middleware so fallbacks receive the selected model. Candidate models must support the agent's tools and output format. - Example: + ??? example "Route agent calls by task" + ```python middleware = ModelRoutingMiddleware( models={ @@ -87,8 +98,8 @@ class ModelRoutingMiddleware(AgentMiddleware[ModelRoutingState, ContextT, Respon "criteria": "Architectural tradeoffs", }, }, - routing_model=selector_model, - system_prompt="Use the least expensive model that can complete the task safely.", + decision_model=selector_model, + instructions="Use the least expensive model that can complete the task safely.", ) agent = create_agent(model=fast_model, middleware=[middleware]) ``` @@ -101,22 +112,21 @@ class ModelRoutingMiddleware(AgentMiddleware[ModelRoutingState, ContextT, Respon self, *, models: Mapping[str, ModelRoutingConfig], - routing_model: str | BaseChatModel | None = None, - decision_model: Runnable[ModelRoutingInput, str] | None = None, - system_prompt: str = DEFAULT_SYSTEM_PROMPT, + decision_model: str | BaseChatModel | Runnable[ModelRoutingInput, ModelRoutingOutput], + instructions: str = DEFAULT_INSTRUCTIONS, input_extractor: Callable[[ModelRoutingState], Sequence[BaseMessage]] | None = None, fallback_route: str | None = None, ) -> None: - """Initialize routing with exactly one selection backend. + """Initialize routing with a chat model or custom decision runnable. Args: models: Route names mapped to configurations containing a `model` instance or identifier string and its selection `criteria`. - routing_model: Chat model supporting `with_structured_output`. - decision_model: Classification runnable accepting `ModelRoutingInput` and - returning a route name. Adapt provider-specific classifiers with a - runnable; no classifier dependency is required. - system_prompt: Base instructions for selection, separate from the agent prompt. + decision_model: Chat model or model string supporting `with_structured_output`, + or a custom runnable accepting `ModelRoutingInput` and returning + `ModelRoutingOutput`. Chat models use the built-in routing prompt and + output schema. + instructions: Base instructions for selection, separate from the agent prompt. input_extractor: Routing messages extracted from agent state. By default, uses the latest human message. Customize to filter application metadata. fallback_route: Route used for malformed or unknown selections. Without it, @@ -124,16 +134,12 @@ class ModelRoutingMiddleware(AgentMiddleware[ModelRoutingState, ContextT, Respon always propagate; configure retries or fallbacks on the backend runnable. Raises: - ValueError: If routes are empty, the fallback is unknown, or exactly one - selection backend is not supplied. + ValueError: If routes are empty or the fallback is unknown. """ super().__init__() if not models or any(not route for route in models): msg = "models must contain non-empty route names" raise ValueError(msg) - if (routing_model is None) == (decision_model is None): - msg = "Provide exactly one of routing_model or decision_model" - raise ValueError(msg) if fallback_route is not None and fallback_route not in models: msg = "fallback_route must be a configured route name" raise ValueError(msg) @@ -146,16 +152,25 @@ class ModelRoutingMiddleware(AgentMiddleware[ModelRoutingState, ContextT, Respon for route, config in models.items() } self.criteria = {route: config["criteria"] for route, config in models.items()} - self.system_prompt = system_prompt + self.instructions = instructions self.input_extractor = input_extractor self.fallback_route = fallback_route - self.decision_model = decision_model - self._routing_model: Runnable[LanguageModelInput, object] | None = None - if routing_model is not None: - model = ( - init_chat_model(routing_model) if isinstance(routing_model, str) else routing_model - ) - self._routing_model = model.with_structured_output( + self.decision_model: Runnable[ModelRoutingInput, ModelRoutingOutput] + if isinstance(decision_model, (str, BaseChatModel)): + self.decision_model = self._create_decision_model(decision_model) + else: + self.decision_model = decision_model + + def _create_decision_model( + self, model: str | BaseChatModel + ) -> Runnable[ModelRoutingInput, ModelRoutingOutput]: + """Adapt a structured-output chat model to the decision runnable interface.""" + if isinstance(model, str): + model = init_chat_model(model) + return cast( + "Runnable[ModelRoutingInput, ModelRoutingOutput]", + self._llm_input + | model.with_structured_output( { "title": "ModelRoutingResponse", "description": "Select a model route for the user's task.", @@ -170,7 +185,13 @@ class ModelRoutingMiddleware(AgentMiddleware[ModelRoutingState, ContextT, Respon "required": ["route"], "additionalProperties": False, } - ) + ), + ) + + def _extract_route(self, response: object) -> str | None: + """Extract a string route, or return `None` for malformed output.""" + route = response.get("route") if isinstance(response, dict) else None + return route if isinstance(route, str) else None def _routing_input(self, state: ModelRoutingState) -> ModelRoutingInput: """Extract application input without modifying the agent request.""" @@ -190,15 +211,14 @@ class ModelRoutingMiddleware(AgentMiddleware[ModelRoutingState, ContextT, Respon raise ValueError(msg) return { "messages": messages, - "system_prompt": self.system_prompt, + "instructions": self.instructions, "criteria": dict(self.criteria), } def _llm_input(self, inputs: ModelRoutingInput) -> list[BaseMessage]: """Present configured criteria as data alongside base routing instructions.""" prompt = ( - f"{inputs['system_prompt']}\n\nModel routing criteria:\n" - f"{json.dumps(inputs['criteria'])}" + f"{inputs['instructions']}\n\nModel routing criteria:\n{json.dumps(inputs['criteria'])}" ) return [SystemMessage(content=prompt), *inputs["messages"]] @@ -212,90 +232,36 @@ 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: - """Select or reuse a route without mutating agent state. - - Args: - state: Agent state supplying routing input and any persisted route. - - Returns: - A configured route name. - """ - if state.get("model_route") is not None: - return self._validate_route(state["model_route"]) - inputs = self._routing_input(state) - config: RunnableConfig = {"metadata": internal_call_metadata()} - if self.decision_model is not None: - return self._validate_route(self.decision_model.invoke(inputs, config=config)) - if self._routing_model is None: - msg = "No routing backend configured" - raise AssertionError(msg) - 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: - """Select or reuse a route asynchronously without mutating agent state. - - Args: - state: Agent state supplying routing input and any persisted route. - - Returns: - A configured route name. - """ - if state.get("model_route") is not None: - return self._validate_route(state["model_route"]) - inputs = self._routing_input(state) - config: RunnableConfig = {"metadata": internal_call_metadata()} - if self.decision_model is not None: - return self._validate_route(await self.decision_model.ainvoke(inputs, config=config)) - if self._routing_model is None: - msg = "No routing backend configured" - raise AssertionError(msg) - response = await self._routing_model.ainvoke(self._llm_input(inputs), config=config) - return self._validate_route(response.get("route") if isinstance(response, dict) else None) - + @override def before_model(self, state: ModelRoutingState, runtime: Runtime[ContextT]) -> dict[str, str]: - """Persist selection before the model runs. - - Args: - state: Agent state with routing inputs. - runtime: Agent runtime. - - Returns: - The selected route state update. - """ - del runtime - return {"model_route": self._select_route(state)} + """Select or reuse a route before the model runs.""" + route = state.get("model_route") + if route is None: + response = self.decision_model.invoke( + self._routing_input(state), config={"metadata": internal_call_metadata()} + ) + route = self._extract_route(response) + return {"model_route": self._validate_route(route)} + @override async def abefore_model( self, state: ModelRoutingState, runtime: Runtime[ContextT] ) -> dict[str, str]: - """Persist asynchronous selection before the model runs. - - Args: - state: Agent state with routing inputs. - runtime: Agent runtime. - - Returns: - The selected route state update. - """ - del runtime - return {"model_route": await self._aselect_route(state)} + """Select or reuse a route asynchronously before the model runs.""" + route = state.get("model_route") + if route is None: + response = await self.decision_model.ainvoke( + self._routing_input(state), config={"metadata": internal_call_metadata()} + ) + route = self._extract_route(response) + return {"model_route": self._validate_route(route)} def wrap_model_call( self, request: ModelRequest[ContextT], handler: Callable[[ModelRequest[ContextT]], ModelResponse[ResponseT]], ) -> ModelResponse[ResponseT]: - """Call the handler with the selected model, preserving all other request fields. - - Args: - request: Agent model request. - handler: Handler for the selected model request. - - Returns: - The selected model's response. - """ + """Call the handler with the selected model.""" route = self._validate_route(request.state.get("model_route")) return handler(request.override(model=self.models[route])) @@ -304,14 +270,6 @@ class ModelRoutingMiddleware(AgentMiddleware[ModelRoutingState, ContextT, Respon request: ModelRequest[ContextT], handler: Callable[[ModelRequest[ContextT]], Awaitable[ModelResponse[ResponseT]]], ) -> ModelResponse[ResponseT]: - """Call the async handler with the selected model, preserving other request fields. - - Args: - request: Agent model request. - handler: Async handler for the selected model request. - - Returns: - The selected model's response. - """ + """Call the async handler with the selected model.""" 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 bca21ed26a..7e2f974d20 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 @@ -3,6 +3,7 @@ import importlib.util from collections.abc import Sequence from typing import Any +from unittest.mock import patch import pytest from langchain_core._api import LangChainBetaWarning @@ -11,7 +12,7 @@ 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 pydantic import Field from typing_extensions import override from langchain.agents import create_agent @@ -26,6 +27,7 @@ from langchain.agents.middleware.internal_call_transformer import internal_call_ from langchain.agents.middleware.model_routing import ( ModelRoutingConfig, ModelRoutingInput, + ModelRoutingOutput, ModelRoutingState, ) @@ -41,18 +43,19 @@ def test_beta_warning() -> None: models={ "small": {"model": FakeListChatModel(responses=["small"]), "criteria": "All tasks"} }, - decision_model=RunnableLambda(lambda _: "small"), + decision_model=RunnableLambda(lambda _: ModelRoutingOutput(route="small")), ) class RoutingChatModel(FakeListChatModel): - selection: object = "large" + response: object = Field(default_factory=lambda: {"route": "large"}) routing_schema: dict[str, Any] | None = None received_messages: list[BaseMessage] = Field(default_factory=list) + selection_modes: list[str] = Field(default_factory=list) @override def with_structured_output( - self, schema: dict[str, Any] | type[BaseModel], **_kwargs: Any + self, schema: dict[str, Any] | type, **_kwargs: Any ) -> Runnable[LanguageModelInput, Any]: assert isinstance(schema, dict) self.routing_schema = schema @@ -64,9 +67,15 @@ class RoutingChatModel(FakeListChatModel): message for message in messages if isinstance(message, BaseMessage) ] assert config["metadata"].items() >= internal_call_metadata().items() - return {"route": self.selection} + self.selection_modes.append("sync") + return self.response - return RunnableLambda(select) + async def aselect(messages: LanguageModelInput, config: RunnableConfig) -> object: + response = select(messages, config) + self.selection_modes[-1] = "async" + return response + + return RunnableLambda(select, afunc=aselect) def make_request() -> ModelRequest: @@ -93,7 +102,7 @@ def check_request(request: ModelRequest, original: ModelRequest) -> None: assert original.model.invoke("task").content == "original" -@pytest.mark.parametrize("backend", ["llm", "classifier"]) +@pytest.mark.parametrize("backend", ["llm", "llm_string", "classifier"]) @pytest.mark.parametrize("async_mode", [False, True]) async def test_route_and_preserve_request(backend: str, *, async_mode: bool) -> None: original = make_request() @@ -111,29 +120,40 @@ async def test_route_and_preserve_request(backend: str, *, async_mode: bool) -> selection_modes: list[str] = [] router = RoutingChatModel(responses=[""]) - def classify(inputs: ModelRoutingInput, config: RunnableConfig) -> str: + def classify(inputs: ModelRoutingInput, config: RunnableConfig) -> ModelRoutingOutput: assert inputs == { "messages": [original.messages[-1]], - "system_prompt": "Custom routing instructions", + "instructions": "Custom routing instructions", "criteria": criteria, } assert config["metadata"].items() >= internal_call_metadata().items() selection_modes.append("sync") - return "large" + return {"route": "large"} - async def aclassify(inputs: ModelRoutingInput, config: RunnableConfig) -> str: + async def aclassify(inputs: ModelRoutingInput, config: RunnableConfig) -> ModelRoutingOutput: 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, afunc=aclassify) - if backend == "classifier" - else None, - ) + decision_model: str | RoutingChatModel | Runnable[ModelRoutingInput, ModelRoutingOutput] + if backend == "llm_string": + decision_model = "test:router" + elif backend == "llm": + decision_model = router + else: + decision_model = RunnableLambda(classify, afunc=aclassify) + with patch( + "langchain.agents.middleware.model_routing.init_chat_model", return_value=router + ) as init_model: + middleware = ModelRoutingMiddleware( + models=models, + instructions="Custom routing instructions", + decision_model=decision_model, + ) + if backend == "llm_string": + init_model.assert_called_once_with("test:router") + else: + init_model.assert_not_called() def handler(request: ModelRequest) -> ModelResponse: check_request(request, original) @@ -159,8 +179,13 @@ async def test_route_and_preserve_request(backend: str, *, async_mode: bool) -> assert response.result[0].content == "large response" if backend == "classifier": assert selection_modes == ["async" if async_mode else "sync"] - if backend == "llm": + if backend != "classifier": + assert router.selection_modes == ["async" if async_mode else "sync"] assert router.routing_schema is not None + assert router.routing_schema["type"] == "object" + assert router.routing_schema["required"] == ["route"] + assert router.routing_schema["additionalProperties"] is False + assert router.routing_schema["properties"]["route"]["type"] == "string" assert router.routing_schema["properties"]["route"]["enum"] == ["small", "large"] assert router.received_messages[-1] == original.messages[-1] assert "Custom routing instructions" in router.received_messages[0].text @@ -168,21 +193,44 @@ async def test_route_and_preserve_request(backend: str, *, async_mode: bool) -> assert "Complex tasks" in router.received_messages[0].text +def test_routing_schema_preserves_each_instances_route_names() -> None: + route_sets = [["fast-model", "provider:model"], ["_ignore_", "__members__"]] + routers = [ + RoutingChatModel(responses=[""], response={"route": routes[0]}) for routes in route_sets + ] + for routes, router in zip(route_sets, routers, strict=True): + middleware = ModelRoutingMiddleware( + models={ + route: {"model": FakeListChatModel(responses=[route]), "criteria": route} + for route in routes + }, + decision_model=router, + ) + state = ModelRoutingState(messages=[HumanMessage(content="Task")]) + assert middleware.before_model(state, Runtime()) == {"model_route": routes[0]} + + for routes, router in zip(route_sets, routers, strict=True): + assert router.routing_schema is not None + assert router.routing_schema["properties"]["route"]["enum"] == routes + + @pytest.mark.parametrize("backend", ["llm", "classifier"]) -@pytest.mark.parametrize("selection", ["unknown", None, ["large"]]) -async def test_invalid_selection_and_explicit_fallback(backend: str, selection: object) -> None: +@pytest.mark.parametrize( + "response", + [{"route": "unknown"}, {"route": None}, {"route": ["large"]}, {}, None, "large", ["large"]], +) +async def test_invalid_selection_and_explicit_fallback(backend: str, response: object) -> None: models: dict[str, ModelRoutingConfig] = { "large": {"model": FakeListChatModel(responses=["fallback"]), "criteria": "All tasks"} } - router = RoutingChatModel(responses=[""], selection=selection) + router = RoutingChatModel(responses=[""], response=response) def classify(_inputs: ModelRoutingInput) -> Any: - return selection + return response middleware = ModelRoutingMiddleware( models=models, - routing_model=router if backend == "llm" else None, - decision_model=RunnableLambda(classify) if backend == "classifier" else None, + decision_model=RunnableLambda(classify) if backend == "classifier" else router, ) with pytest.raises(ValueError, match="configured route name"): middleware.before_model(ModelRoutingState(messages=make_request().messages), Runtime()) @@ -191,6 +239,9 @@ async def test_invalid_selection_and_explicit_fallback(backend: str, selection: ModelRoutingState(messages=make_request().messages), Runtime() ) middleware.fallback_route = "large" + assert await middleware.abefore_model( + ModelRoutingState(messages=make_request().messages), Runtime() + ) == {"model_route": "large"} original = make_request() original = original.override( state=ModelRoutingState( @@ -210,16 +261,16 @@ async def test_invalid_selection_and_explicit_fallback(backend: str, selection: return ModelResponse(result=[await request.model.ainvoke(request.messages)]) assert middleware.wrap_model_call(original, handler).result[0].content == "fallback" - response = await middleware.awrap_model_call(original, ahandler) - assert response.result[0].content == "fallback" + model_response = await middleware.awrap_model_call(original, ahandler) + assert model_response.result[0].content == "fallback" async def test_custom_input_and_no_cross_request_cache() -> None: def extract(state: ModelRoutingState) -> Sequence[BaseMessage]: return [message for message in state["messages"] if message.text != "Injected context"] - def classify(inputs: ModelRoutingInput) -> str: - return "small" if inputs["messages"][-1].text == "Lookup" else "large" + def classify(inputs: ModelRoutingInput) -> ModelRoutingOutput: + return {"route": "small" if inputs["messages"][-1].text == "Lookup" else "large"} middleware = ModelRoutingMiddleware( models={ @@ -251,8 +302,6 @@ async def test_custom_input_and_no_cross_request_cache() -> None: ("kwargs", "match"), [ ({"models": {}}, "non-empty"), - ({"decision_model": None}, "exactly one"), - ({"routing_model": FakeListChatModel(responses=[""])}, "exactly one"), ({"fallback_route": "unknown"}, "fallback_route"), ], ) @@ -261,15 +310,35 @@ def test_configuration_validation(kwargs: dict[str, Any], match: str) -> None: "models": { "small": {"model": FakeListChatModel(responses=["small"]), "criteria": "All tasks"} }, - "decision_model": RunnableLambda(lambda _: "small"), + "decision_model": RunnableLambda(lambda _: ModelRoutingOutput(route="small")), } options.update(kwargs) with pytest.raises(ValueError, match=match): ModelRoutingMiddleware(**options) +@pytest.mark.parametrize( + ("kwargs", "match"), + [ + ({}, "required keyword-only argument: 'decision_model'"), + ( + {"routing_model": FakeListChatModel(responses=[""])}, + "unexpected keyword argument 'routing_model'", + ), + ], +) +def test_decision_model_required(kwargs: dict[str, Any], match: str) -> None: + with pytest.raises(TypeError, match=match): + ModelRoutingMiddleware( + models={ + "small": {"model": FakeListChatModel(responses=["small"]), "criteria": "All tasks"} + }, + **kwargs, + ) + + async def test_backend_errors_propagate_and_missing_input_is_explicit() -> None: - def classify(_inputs: ModelRoutingInput) -> str: + def classify(_inputs: ModelRoutingInput) -> ModelRoutingOutput: msg = "Classifier unavailable" raise RuntimeError(msg) @@ -292,7 +361,8 @@ async def test_backend_errors_propagate_and_missing_input_is_explicit() -> None: ) -async def test_agent_uses_selected_model() -> None: +@pytest.mark.parametrize("backend", ["llm", "classifier"]) +async def test_agent_uses_selected_model(backend: str) -> None: middleware = ModelRoutingMiddleware( models={ "small": { @@ -300,7 +370,9 @@ async def test_agent_uses_selected_model() -> None: "criteria": "All tasks", } }, - decision_model=RunnableLambda(lambda _: "small"), + decision_model=RoutingChatModel(responses=[""], response={"route": "small"}) + if backend == "llm" + else RunnableLambda(lambda _: ModelRoutingOutput(route="small")), ) agent = create_agent( model=FakeListChatModel(responses=["Original response"]), @@ -319,7 +391,7 @@ async def test_route_persists_until_cleared(*, async_mode: bool) -> None: "small": {"model": FakeListChatModel(responses=["small"]), "criteria": "Lookup"}, "large": {"model": FakeListChatModel(responses=["large"]), "criteria": "Reasoning"}, }, - decision_model=RunnableLambda(lambda _: next(selections)), + decision_model=RunnableLambda(lambda _: ModelRoutingOutput(route=next(selections))), ) state = ModelRoutingState(messages=[HumanMessage(content="Lookup")])