chore(langchain): model router review edits

This commit is contained in:
Hunter Lovell committed 2026-10-01 07:36:37 -07:00
1 parent 93fbf66a54
commit 26358138b4
2 files changed
+184 -154

No files matched your search

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