diff --git a/libs/langchain_v1/langchain/agents/middleware/__init__.py b/libs/langchain_v1/langchain/agents/middleware/__init__.py index 1a2ba68b8d..a79ed7315a 100644 --- a/libs/langchain_v1/langchain/agents/middleware/__init__.py +++ b/libs/langchain_v1/langchain/agents/middleware/__init__.py @@ -16,6 +16,7 @@ from langchain.agents.middleware.model_routing import ( ModelRoutingConfig, ModelRoutingInput, ModelRoutingMiddleware, + ModelRoutingState, ) from langchain.agents.middleware.pii import PIIDetectionError, PIIMatch, PIIMiddleware from langchain.agents.middleware.provider_tool_search import ProviderToolSearchMiddleware @@ -79,6 +80,7 @@ __all__ = [ "ModelRoutingConfig", "ModelRoutingInput", "ModelRoutingMiddleware", + "ModelRoutingState", "OutputAgentState", "PIIDetectionError", "PIIMatch", diff --git a/libs/langchain_v1/langchain/agents/middleware/model_routing.py b/libs/langchain_v1/langchain/agents/middleware/model_routing.py index 933ffb4fd0..fe7e3fee46 100644 --- a/libs/langchain_v1/langchain/agents/middleware/model_routing.py +++ b/libs/langchain_v1/langchain/agents/middleware/model_routing.py @@ -8,7 +8,7 @@ from typing import TYPE_CHECKING from langchain_core._api import beta from langchain_core.messages import BaseMessage, HumanMessage, SystemMessage -from typing_extensions import TypedDict +from typing_extensions import NotRequired, TypedDict from langchain.agents.middleware.internal_call_transformer import ( InternalCallTransformer, @@ -29,6 +29,7 @@ if TYPE_CHECKING: from langchain_core.language_models import BaseChatModel, LanguageModelInput from langchain_core.runnables import Runnable, RunnableConfig + from langgraph.runtime import Runtime logger = logging.getLogger(__name__) @@ -50,21 +51,29 @@ class ModelRoutingInput(TypedDict): criteria: dict[str, str] +class ModelRoutingState(AgentState): + """Experimental checkpointed route; clear `model_route` to select again.""" + + model_route: NotRequired[str | None] + + @beta( addendum=( "Experimental API: may change or be removed without notice; no compatibility guarantees." ) ) -class ModelRoutingMiddleware(AgentMiddleware[AgentState[ResponseT], ContextT, ResponseT]): +class ModelRoutingMiddleware(AgentMiddleware[ModelRoutingState, ContextT, ResponseT]): """Route agent model calls using a structured-output LLM or classification runnable. !!! warning "Beta / Experimental" This middleware and its input schema may change or be removed without notice. No compatibility guarantees are provided. - Routes are selected for each model call, without shared or persisted selection state. - Applications requiring one selection per turn can call `select_route` or - `aselect_route` during preparation and persist the result in their own state. + 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. Place this middleware before model fallback middleware so fallbacks receive the selected model. Candidate models must support the agent's tools and output format. @@ -85,6 +94,7 @@ class ModelRoutingMiddleware(AgentMiddleware[AgentState[ResponseT], ContextT, Re ``` """ + state_schema = ModelRoutingState transformers = (InternalCallTransformer,) def __init__( @@ -94,7 +104,7 @@ class ModelRoutingMiddleware(AgentMiddleware[AgentState[ResponseT], ContextT, Re routing_model: str | BaseChatModel | None = None, decision_model: Runnable[ModelRoutingInput, str] | None = None, system_prompt: str = DEFAULT_SYSTEM_PROMPT, - input_extractor: Callable[[ModelRequest[ContextT]], Sequence[BaseMessage]] | None = None, + input_extractor: Callable[[ModelRoutingState], Sequence[BaseMessage]] | None = None, fallback_route: str | None = None, ) -> None: """Initialize routing with exactly one selection backend. @@ -107,7 +117,7 @@ class ModelRoutingMiddleware(AgentMiddleware[AgentState[ResponseT], ContextT, Re 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. - input_extractor: Routing messages extracted from the request. By default, + 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, invalid selections raise `ValueError`. Backend and extractor exceptions @@ -162,15 +172,15 @@ class ModelRoutingMiddleware(AgentMiddleware[AgentState[ResponseT], ContextT, Re } ) - def _routing_input(self, request: ModelRequest[ContextT]) -> ModelRoutingInput: + def _routing_input(self, state: ModelRoutingState) -> ModelRoutingInput: """Extract application input without modifying the agent request.""" if self.input_extractor is not None: - messages = list(self.input_extractor(request)) + messages = list(self.input_extractor(state)) else: messages = next( ( [message] - for message in reversed(request.messages) + for message in reversed(state["messages"]) if isinstance(message, HumanMessage) ), [], @@ -202,16 +212,18 @@ class ModelRoutingMiddleware(AgentMiddleware[AgentState[ResponseT], ContextT, Re msg = "Model routing selection must be a configured route name" raise ValueError(msg) - def select_route(self, request: ModelRequest[ContextT]) -> str: - """Select a route without changing the request. + def select_route(self, state: ModelRoutingState) -> str: + """Select or reuse a route without mutating agent state. Args: - request: Agent request supplying routing input. + state: Agent state supplying routing input and any persisted route. Returns: A configured route name. """ - inputs = self._routing_input(request) + 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)) @@ -221,16 +233,18 @@ class ModelRoutingMiddleware(AgentMiddleware[AgentState[ResponseT], ContextT, Re 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, request: ModelRequest[ContextT]) -> str: - """Select a route asynchronously without changing the request. + async def aselect_route(self, state: ModelRoutingState) -> str: + """Select or reuse a route asynchronously without mutating agent state. Args: - request: Agent request supplying routing input. + state: Agent state supplying routing input and any persisted route. Returns: A configured route name. """ - inputs = self._routing_input(request) + 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)) @@ -240,6 +254,34 @@ class ModelRoutingMiddleware(AgentMiddleware[AgentState[ResponseT], ContextT, Re 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) + 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)} + + 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)} + def wrap_model_call( self, request: ModelRequest[ContextT], @@ -254,7 +296,7 @@ class ModelRoutingMiddleware(AgentMiddleware[AgentState[ResponseT], ContextT, Re Returns: The selected model's response. """ - route = self.select_route(request) + route = self._validate_route(request.state.get("model_route")) return handler(request.override(model=self.models[route])) async def awrap_model_call( @@ -271,5 +313,5 @@ class ModelRoutingMiddleware(AgentMiddleware[AgentState[ResponseT], ContextT, Re Returns: The selected model's response. """ - route = await self.aselect_route(request) + 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 d354c21902..c09dab962c 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 @@ -21,6 +21,7 @@ from langchain.agents.middleware import ( ModelRoutingConfig, ModelRoutingInput, ModelRoutingMiddleware, + ModelRoutingState, model_routing, ) from langchain.agents.middleware.internal_call_transformer import internal_call_metadata @@ -130,6 +131,16 @@ 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"], + ) + ) response = ( await middleware.awrap_model_call(original, ahandler) if async_mode @@ -162,11 +173,17 @@ 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(make_request()) + middleware.select_route(ModelRoutingState(messages=make_request().messages)) with pytest.raises(ValueError, match="configured route name"): - await middleware.aselect_route(make_request()) + await middleware.aselect_route(ModelRoutingState(messages=make_request().messages)) middleware.fallback_route = "large" original = make_request() + original = original.override( + state=ModelRoutingState( + messages=original.messages, + model_route=middleware.select_route(ModelRoutingState(messages=original.messages)), + ) + ) def handler(request: ModelRequest) -> ModelResponse: check_request(request, original) @@ -182,8 +199,8 @@ async def test_invalid_selection_and_explicit_fallback(backend: str, selection: async def test_custom_input_and_no_cross_request_cache() -> None: - def extract(request: ModelRequest) -> Sequence[BaseMessage]: - return [message for message in request.messages if message.text != "Injected context"] + 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" @@ -200,9 +217,9 @@ 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(first) == "small" - assert await middleware.aselect_route(second) == "large" - assert await middleware.aselect_route(first) == "small" + 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" @pytest.mark.parametrize( @@ -239,11 +256,11 @@ 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(make_request()) + middleware.select_route(ModelRoutingState(messages=make_request().messages)) with pytest.raises(RuntimeError, match="Classifier unavailable"): - await middleware.aselect_route(make_request()) + await middleware.aselect_route(ModelRoutingState(messages=make_request().messages)) with pytest.raises(ValueError, match="No routing messages"): - middleware.select_route(make_request().override(messages=[AIMessage(content="No user")])) + middleware.select_route(ModelRoutingState(messages=[AIMessage(content="No user")])) async def test_agent_uses_selected_model() -> None: @@ -263,3 +280,30 @@ async def test_agent_uses_selected_model() -> None: inputs: InputAgentState = {"messages": [HumanMessage(content="Do this task")]} assert agent.invoke(inputs)["messages"][-1].content == "Selected response" assert (await agent.ainvoke(inputs))["messages"][-1].content == "Selected response" + + +@pytest.mark.parametrize("async_mode", [False, True]) +async def test_route_persists_until_cleared(*, async_mode: bool) -> None: + selections = iter(["small", "large"]) + middleware = ModelRoutingMiddleware( + models={ + "small": {"model": FakeListChatModel(responses=["small"]), "criteria": "Lookup"}, + "large": {"model": FakeListChatModel(responses=["large"]), "criteria": "Reasoning"}, + }, + decision_model=RunnableLambda(lambda _: next(selections)), + ) + state = ModelRoutingState(messages=[HumanMessage(content="Lookup")]) + + 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] + + state["model_route"] = await prepare() + assert state["model_route"] == "small" + state["messages"] = [HumanMessage(content="Reasoning")] + state["model_route"] = await prepare() + assert state["model_route"] == "small" + state["model_route"] = None + state["model_route"] = await prepare() + assert state["model_route"] == "large"