mirror of
https://github.com/langchain-ai/langchain.git
synced 2026-10-05 09:25:14 +03:00
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:
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]))
|
||||
+54
-10
@@ -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"
|
||||
Reference in new issue
Block a user