refactor(langchain): keep model selection inside middleware hooks

Co-authored-by: open-swe[bot] <open-swe@users.noreply.github.com>
This commit is contained in:
Sydney Runkleandopen-swe[bot] committed 2026-10-01 02:20:48 +00:00
1 parent b65d30569d
commit 613951a5bd
2 files changed
+57 -30

No files matched your search

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