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