refactor(langchain): persist routing through the middleware lifecycle

Co-authored-by: open-swe[bot] <open-swe@users.noreply.github.com>
This commit is contained in:
Sydney Runkleandopen-swe[bot] committed 2026-09-30 23:37:21 +00:00
1 parent e4493d8cb9
commit 04aa56c4b2
3 files changed
+118 -30

No files matched your search

@@ -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",
@@ -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]))
@@ -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"