mirror of
https://github.com/langchain-ai/langchain.git
synced 2026-10-05 01:15:09 +03:00
chore(langchain): model router review edits
This commit is contained in:
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]))
|
||||
+109
-37
@@ -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")])
|
||||
|
||||
|
||||
Reference in new issue
Block a user