mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
feat(router): add Router(plugins=[...]) routing-plugin pipeline (#32972)
* feat(router): add Router(plugins=[...]) routing-plugin pipeline Runs a sequence of user-supplied plugins before the routing decision is made. Each plugin reads/mutates a RoutingContext (messages, candidate models, metadata, signals); the narrowed candidate list is enforced when picking a deployment, raising rather than silently falling back if a plugin narrows to zero candidates. Prototype for the routing-plugin pipeline discussed in #32168. * fix(router): use ruff-modern typing, add raw/structured messages to RoutingContext - Use dict/list/X|None instead of Dict/List/Optional in new code, staying within the ruff strict-rule budget ratchet - Extract the guardrail-translation message normalization ComplexityRouter already had into a shared resolve_structured_messages() helper (litellm_core_utils/prompt_templates/factory.py), reused by ComplexityRouter and the new routing-plugin pipeline instead of duplicating it - RoutingContext now exposes both raw_messages (as received) and structured_messages (normalized across chat completions / Anthropic messages / Responses API), mirroring CustomGuardrail.apply_guardrail's pattern, per review feedback on #32972 - Add direct unit tests for _run_routing_plugins and _filter_by_routing_plugin_candidates (router_code_coverage gate requires every router.py function be called by name somewhere in tests/) * fix(test): rename to test_router_routing_plugins.py router_code_coverage.py's AST scanner only inspects test files whose filename contains the substring "router" -- test_routing_plugins.py doesn't match (routing != router), so it silently skipped this file and flagged _run_routing_plugins/_filter_by_routing_plugin_candidates as untested despite the direct unit tests added for them. * fix(router): fail closed when plugins are configured but the resolved routing path can't run them Router.completion() (and other sync entry points) resolves deployments via the synchronous get_available_deployment(), which never runs async_pre_routing_hook and therefore never runs the routing-plugin pipeline. async_get_available_deployment() itself falls back to that same synchronous method for routing strategies without an async-native selector (e.g. legacy "usage-based-routing" v1). Both paths would let a policy plugin (e.g. a deny-all rule) be silently bypassed. Raise instead of silently proceeding when self.routing_plugins is configured and the sync path is reached, since applying the pipeline to every selector path is a larger change out of scope for this PR. Per review: https://github.com/BerriAI/litellm/pull/32972/changes/BASE..bdfb583c2c6f8df10004fb249e11629d41ce71fa#r3565373303
This commit is contained in:
parent
c136797805
commit
85f9bdd412
5 changed files with 407 additions and 36 deletions
|
|
@ -5494,3 +5494,56 @@ def has_tool_with_name(tools: Any, tool_name: str) -> bool:
|
|||
elif tool.get("name") == tool_name:
|
||||
return True
|
||||
return False
|
||||
|
||||
|
||||
def resolve_structured_messages(
|
||||
messages: list[dict[str, Any]] | None,
|
||||
request_kwargs: dict[str, Any],
|
||||
) -> list[dict[str, Any]] | None:
|
||||
"""
|
||||
Normalize a request's messages to OpenAI-spec chat-completions shape,
|
||||
regardless of which API surface produced them (chat completions,
|
||||
Anthropic /v1/messages, Responses API ``input``, etc).
|
||||
|
||||
Returns ``messages`` unchanged if already present. Otherwise dispatches
|
||||
through the guardrail translation handlers (the same per-surface
|
||||
conversion logic guardrails use) to convert e.g. Responses API ``input``
|
||||
into a message list. Returns ``None`` if no messages could be resolved.
|
||||
"""
|
||||
if messages:
|
||||
return messages
|
||||
|
||||
from litellm.litellm_core_utils.api_route_to_call_types import (
|
||||
get_call_types_for_route,
|
||||
)
|
||||
from litellm.llms import load_guardrail_translation_mappings
|
||||
from litellm.types.utils import CallTypes
|
||||
|
||||
mappings = load_guardrail_translation_mappings()
|
||||
call_type: CallTypes | None = None
|
||||
|
||||
# 1. Try route-based inference from proxy metadata
|
||||
route = request_kwargs.get("litellm_metadata", {}).get("user_api_key_request_route")
|
||||
if route:
|
||||
call_types_list = get_call_types_for_route(route)
|
||||
if call_types_list:
|
||||
for ct in call_types_list:
|
||||
if ct in mappings:
|
||||
call_type = ct
|
||||
break
|
||||
|
||||
# 2. Fallback: try each mapped handler until one produces messages
|
||||
handlers_to_try: list[Any] = []
|
||||
if call_type is not None and call_type in mappings:
|
||||
handlers_to_try.append(mappings[call_type]())
|
||||
else:
|
||||
handlers_to_try.extend(handler_cls() for handler_cls in mappings.values())
|
||||
|
||||
for handler in handlers_to_try:
|
||||
structured = handler.get_structured_messages(request_kwargs)
|
||||
if structured:
|
||||
return [
|
||||
msg if isinstance(msg, dict) else msg.model_dump() # type: ignore
|
||||
for msg in structured
|
||||
]
|
||||
return None
|
||||
|
|
|
|||
|
|
@ -179,7 +179,9 @@ from litellm.types.router import (
|
|||
RouterModelGroupAliasItem,
|
||||
RouterRateLimitError,
|
||||
RouterRateLimitErrorBasic,
|
||||
RoutingContext,
|
||||
RoutingGroup,
|
||||
RoutingPlugin,
|
||||
RoutingStrategy,
|
||||
SearchToolTypedDict,
|
||||
)
|
||||
|
|
@ -299,6 +301,7 @@ class Router:
|
|||
enable_pre_call_checks: bool = False,
|
||||
enable_tag_filtering: bool = False,
|
||||
tag_filtering_match_any: bool = True,
|
||||
plugins: list[RoutingPlugin] | None = None,
|
||||
retry_after: int = 0, # min time to wait before retrying a failed request
|
||||
retry_policy: Optional[Union[RetryPolicy, dict]] = None, # set custom retries for different exceptions
|
||||
model_group_retry_policy: Dict[str, RetryPolicy] = {}, # set custom retry policies based on model group
|
||||
|
|
@ -477,6 +480,7 @@ class Router:
|
|||
self.complexity_routers: Dict[str, "ComplexityRouter"] = {}
|
||||
self.adaptive_routers: Dict[str, "AdaptiveRouter"] = {}
|
||||
self.quality_routers: Dict[str, "QualityRouter"] = {}
|
||||
self.routing_plugins: list[RoutingPlugin] = list(plugins) if plugins else []
|
||||
|
||||
# Initialize model_group_alias early since it's used in set_model_list
|
||||
self.model_group_alias: Dict[str, Union[str, RouterModelGroupAliasItem]] = (
|
||||
|
|
@ -10321,6 +10325,12 @@ class Router:
|
|||
metadata_variable_name=self._get_metadata_variable_name_from_kwargs(request_kwargs),
|
||||
)
|
||||
|
||||
# narrow to whatever `self.routing_plugins` left in candidate_models
|
||||
healthy_deployments = self._filter_by_routing_plugin_candidates(
|
||||
healthy_deployments=healthy_deployments,
|
||||
request_kwargs=request_kwargs,
|
||||
)
|
||||
|
||||
## ORDER FILTERING ## -> if user set 'order' in deployments, return deployments with lowest order (e.g. order=1 > order=2)
|
||||
_target_order = (request_kwargs or {}).pop("_target_order", None)
|
||||
healthy_deployments = litellm.utils._get_order_filtered_deployments(
|
||||
|
|
@ -10596,6 +10606,76 @@ class Router:
|
|||
)
|
||||
raise e
|
||||
|
||||
async def _run_routing_plugins(
|
||||
self,
|
||||
model: str,
|
||||
request_kwargs: dict,
|
||||
messages: list[dict[str, Any]] | None,
|
||||
) -> RoutingContext:
|
||||
"""
|
||||
Build a RoutingContext for `model`, run it through `self.routing_plugins`
|
||||
in order, then stash the narrowed candidate list and accumulated signals
|
||||
onto `request_kwargs["metadata"]` so `_filter_by_routing_plugin_candidates`
|
||||
(called later, during healthy-deployment filtering) can consume them.
|
||||
"""
|
||||
from litellm.litellm_core_utils.prompt_templates.factory import (
|
||||
resolve_structured_messages,
|
||||
)
|
||||
|
||||
deployments = self.get_model_list(model_name=model) or []
|
||||
candidate_models = [
|
||||
d["litellm_params"]["model"] for d in deployments if d.get("litellm_params", {}).get("model")
|
||||
]
|
||||
|
||||
metadata_key = self._get_metadata_variable_name_from_kwargs(request_kwargs)
|
||||
metadata = request_kwargs.setdefault(metadata_key, {})
|
||||
|
||||
context = RoutingContext(
|
||||
raw_messages=messages or [],
|
||||
structured_messages=resolve_structured_messages(messages=messages, request_kwargs=request_kwargs) or [],
|
||||
candidate_models=candidate_models,
|
||||
metadata=metadata,
|
||||
)
|
||||
|
||||
for plugin in self.routing_plugins:
|
||||
context = await plugin.run(context)
|
||||
|
||||
metadata["routing_plugin_signals"] = context.signals
|
||||
if len(context.candidate_models) < len(candidate_models):
|
||||
metadata["_routing_plugin_candidate_models"] = context.candidate_models
|
||||
|
||||
return context
|
||||
|
||||
def _filter_by_routing_plugin_candidates(
|
||||
self,
|
||||
healthy_deployments: Union[list[dict], dict],
|
||||
request_kwargs: dict,
|
||||
) -> Union[list[dict], dict]:
|
||||
"""
|
||||
Narrow `healthy_deployments` to whatever `self.routing_plugins` left in
|
||||
`context.candidate_models`. Raises rather than silently falling back to
|
||||
the unfiltered pool -- a plugin narrowing to nothing is a policy decision
|
||||
(e.g. no model this tenant's budget allows), not something to bypass.
|
||||
"""
|
||||
if not self.routing_plugins or not isinstance(healthy_deployments, list):
|
||||
return healthy_deployments
|
||||
|
||||
metadata_key = self._get_metadata_variable_name_from_kwargs(request_kwargs)
|
||||
candidate_models = (request_kwargs.get(metadata_key) or {}).get("_routing_plugin_candidate_models")
|
||||
# `is None` (not falsy-check): a plugin narrowing to an empty list must
|
||||
# still hit the "no deployments left" raise below, not be treated the
|
||||
# same as "no plugin ever set this key".
|
||||
if candidate_models is None:
|
||||
return healthy_deployments
|
||||
|
||||
candidate_set = set(candidate_models)
|
||||
filtered = [d for d in healthy_deployments if d.get("litellm_params", {}).get("model") in candidate_set]
|
||||
|
||||
if not filtered:
|
||||
raise ValueError(f"No deployments left after routing-plugin filtering. candidate_models={candidate_models}")
|
||||
|
||||
return filtered
|
||||
|
||||
async def async_pre_routing_hook(
|
||||
self,
|
||||
model: str,
|
||||
|
|
@ -10609,6 +10689,15 @@ class Router:
|
|||
|
||||
Used for the litellm auto-router to modify the request before the routing decision is made.
|
||||
"""
|
||||
#########################################################
|
||||
# Run the routing-plugin pipeline, if any plugins are configured.
|
||||
# Plugins narrow the candidate deployment pool (consumed later by
|
||||
# `_filter_by_routing_plugin_candidates`) and may attach signals for
|
||||
# downstream strategies (auto-router, complexity-router, ...) to read.
|
||||
#########################################################
|
||||
if self.routing_plugins:
|
||||
await self._run_routing_plugins(model=model, request_kwargs=request_kwargs, messages=messages)
|
||||
|
||||
#########################################################
|
||||
# Check if any auto-router should be used
|
||||
#########################################################
|
||||
|
|
@ -10671,6 +10760,18 @@ class Router:
|
|||
"""
|
||||
Returns the deployment based on routing strategy
|
||||
"""
|
||||
if self.routing_plugins:
|
||||
raise ValueError(
|
||||
"Router(plugins=[...]) is configured, but this call resolved to the synchronous "
|
||||
"deployment-selection path, which never runs the routing-plugin pipeline. This "
|
||||
"happens for sync Router methods (e.g. Router.completion()) and for async calls "
|
||||
"with a routing_strategy that has no async-native selector (e.g. legacy "
|
||||
"'usage-based-routing', v1). Silently skipping "
|
||||
"configured plugins would let a policy plugin (e.g. a deny-all rule) be bypassed. "
|
||||
"Use an async Router method with a supported routing_strategy (simple-shuffle, "
|
||||
"usage-based-routing-v2, cost-based-routing, latency-based-routing, least-busy), "
|
||||
"or remove `plugins` from the Router config."
|
||||
)
|
||||
# users need to explicitly call a specific deployment, by setting `specific_deployment = True` as completion()/embedding() kwarg
|
||||
# When this was no explicit we had several issues with fallbacks timing out
|
||||
|
||||
|
|
|
|||
|
|
@ -601,43 +601,11 @@ class ComplexityRouter(CustomLogger):
|
|||
Uses the guardrail translation handler dispatch to convert Responses API
|
||||
``input`` (or other non-chat-completions formats) into OpenAI-spec messages.
|
||||
"""
|
||||
if messages:
|
||||
return messages
|
||||
|
||||
from litellm.litellm_core_utils.api_route_to_call_types import (
|
||||
get_call_types_for_route,
|
||||
from litellm.litellm_core_utils.prompt_templates.factory import (
|
||||
resolve_structured_messages,
|
||||
)
|
||||
from litellm.llms import load_guardrail_translation_mappings
|
||||
from litellm.types.utils import CallTypes
|
||||
|
||||
mappings = load_guardrail_translation_mappings()
|
||||
call_type: Optional[CallTypes] = None
|
||||
|
||||
# 1. Try route-based inference from proxy metadata
|
||||
route = request_kwargs.get("litellm_metadata", {}).get("user_api_key_request_route")
|
||||
if route:
|
||||
call_types_list = get_call_types_for_route(route)
|
||||
if call_types_list:
|
||||
for ct in call_types_list:
|
||||
if ct in mappings:
|
||||
call_type = ct
|
||||
break
|
||||
|
||||
# 2. Fallback: try each mapped handler until one produces messages
|
||||
handlers_to_try: List[Any] = []
|
||||
if call_type is not None and call_type in mappings:
|
||||
handlers_to_try.append(mappings[call_type]())
|
||||
else:
|
||||
handlers_to_try.extend(handler_cls() for handler_cls in mappings.values())
|
||||
|
||||
for handler in handlers_to_try:
|
||||
structured = handler.get_structured_messages(request_kwargs)
|
||||
if structured:
|
||||
return [
|
||||
msg if isinstance(msg, dict) else msg.model_dump() # type: ignore
|
||||
for msg in structured
|
||||
]
|
||||
return None
|
||||
return resolve_structured_messages(messages=messages, request_kwargs=request_kwargs)
|
||||
|
||||
@staticmethod
|
||||
def _extract_user_message_and_system_prompt(
|
||||
|
|
|
|||
|
|
@ -9,7 +9,7 @@ from typing import Any, Dict, List, Literal, Optional, Tuple, Union, get_type_hi
|
|||
|
||||
import httpx
|
||||
from pydantic import BaseModel, ConfigDict, Field, field_validator, model_validator
|
||||
from typing_extensions import Required, TypedDict
|
||||
from typing_extensions import Protocol, Required, TypedDict
|
||||
|
||||
from litellm._uuid import uuid
|
||||
|
||||
|
|
@ -829,6 +829,35 @@ class PreRoutingHookResponse(BaseModel):
|
|||
messages: Optional[List[Dict[str, Any]]]
|
||||
|
||||
|
||||
class RoutingContext(BaseModel):
|
||||
"""
|
||||
Passed through a Router's `plugins` pipeline before the routing decision is made.
|
||||
|
||||
Each plugin reads and mutates this object; the next plugin sees the previous
|
||||
plugin's changes. `candidate_models` narrows as the pipeline runs -- Router
|
||||
only selects a deployment whose `litellm_params.model` survives the pipeline.
|
||||
|
||||
`raw_messages` and `structured_messages` mirror the pattern
|
||||
`CustomGuardrail.apply_guardrail` uses: the message shape differs by API
|
||||
surface (chat completions, Anthropic /v1/messages, Responses API `input`,
|
||||
...), so plugins that need a stable, provider-agnostic shape should read
|
||||
`structured_messages` (normalized to OpenAI chat-completions format);
|
||||
plugins that need the exact original payload can read `raw_messages`.
|
||||
"""
|
||||
|
||||
raw_messages: list[dict[str, Any]]
|
||||
structured_messages: list[dict[str, Any]]
|
||||
candidate_models: list[str]
|
||||
metadata: dict[str, Any] = Field(default_factory=dict)
|
||||
signals: dict[str, Any] = Field(default_factory=dict)
|
||||
|
||||
|
||||
class RoutingPlugin(Protocol):
|
||||
"""Interface a custom routing plugin must implement to run in `Router(plugins=[...])`."""
|
||||
|
||||
async def run(self, context: RoutingContext) -> RoutingContext: ...
|
||||
|
||||
|
||||
class RequestType(str, enum.Enum):
|
||||
"""Fixed v0 taxonomy. User-extensible types come in v1."""
|
||||
|
||||
|
|
|
|||
|
|
@ -0,0 +1,220 @@
|
|||
"""
|
||||
Tests for Router(plugins=[...]) -- a pipeline of routing plugins that run
|
||||
before the routing decision is made, narrowing the candidate deployment pool.
|
||||
|
||||
Discussion: https://github.com/BerriAI/litellm/discussions/32168
|
||||
"""
|
||||
|
||||
import pytest
|
||||
|
||||
from litellm import Router
|
||||
from litellm.types.router import RoutingContext
|
||||
|
||||
|
||||
class LanguageDetector:
|
||||
async def run(self, context: RoutingContext) -> RoutingContext:
|
||||
context.signals["language-detector"] = {"lang": "en"}
|
||||
return context
|
||||
|
||||
|
||||
class DomainClassifier:
|
||||
async def run(self, context: RoutingContext) -> RoutingContext:
|
||||
context.signals["domain-classifier"] = {"domain": "coding", "confidence": 0.93}
|
||||
return context
|
||||
|
||||
|
||||
class TenantPolicy:
|
||||
ALLOWED_PROVIDERS = {"acme-corp": {"openai", "anthropic"}}
|
||||
|
||||
async def run(self, context: RoutingContext) -> RoutingContext:
|
||||
tenant = context.metadata.get("tenant", "default")
|
||||
allowed = self.ALLOWED_PROVIDERS.get(tenant, {"openai", "anthropic", "self-hosted"})
|
||||
context.candidate_models = [m for m in context.candidate_models if m.split("/")[0] in allowed]
|
||||
context.signals["tenant-policy"] = {"tenant": tenant, "allowed_providers": sorted(allowed)}
|
||||
return context
|
||||
|
||||
|
||||
class BudgetPolicy:
|
||||
COST_CAP_PER_TOKEN = 0.000005
|
||||
COST_BY_MODEL = {
|
||||
"openai/gpt-4o-mini": 0.00000015,
|
||||
"anthropic/claude-haiku-4-5": 0.000001,
|
||||
"openai/gpt-5.1": 0.00003,
|
||||
}
|
||||
|
||||
async def run(self, context: RoutingContext) -> RoutingContext:
|
||||
context.candidate_models = [
|
||||
m for m in context.candidate_models if self.COST_BY_MODEL.get(m, 0) <= self.COST_CAP_PER_TOKEN
|
||||
]
|
||||
context.signals["budget-policy"] = {"daily_limit": 100}
|
||||
return context
|
||||
|
||||
|
||||
class BlockEverything:
|
||||
async def run(self, context: RoutingContext) -> RoutingContext:
|
||||
context.candidate_models = []
|
||||
return context
|
||||
|
||||
|
||||
def _smart_router_model_list():
|
||||
return [
|
||||
{
|
||||
"model_name": "smart-router",
|
||||
"litellm_params": {"model": "openai/gpt-4o-mini", "mock_response": "cheap openai"},
|
||||
"model_info": {"tags": ["openai"]},
|
||||
},
|
||||
{
|
||||
"model_name": "smart-router",
|
||||
"litellm_params": {"model": "anthropic/claude-haiku-4-5", "mock_response": "anthropic"},
|
||||
"model_info": {"tags": ["anthropic"]},
|
||||
},
|
||||
{
|
||||
"model_name": "smart-router",
|
||||
"litellm_params": {"model": "openai/gpt-5.1", "mock_response": "expensive openai"},
|
||||
"model_info": {"tags": ["openai"]},
|
||||
},
|
||||
{
|
||||
"model_name": "smart-router",
|
||||
"litellm_params": {"model": "ollama/llama-3-70b", "mock_response": "self hosted"},
|
||||
"model_info": {"tags": ["self-hosted"]},
|
||||
},
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_routing_plugin_pipeline_matches_jeann2013_e2e_scenario():
|
||||
"""
|
||||
https://github.com/BerriAI/litellm/discussions/32168#discussioncomment-17608820
|
||||
|
||||
language plugin -> domain classifier -> tenant policy (openai+anthropic only)
|
||||
-> budget policy (drops over-cap models) -> Router picks the best remaining
|
||||
candidate. Must never land on the self-hosted or over-budget deployment.
|
||||
"""
|
||||
router = Router(
|
||||
model_list=_smart_router_model_list(),
|
||||
plugins=[LanguageDetector(), DomainClassifier(), TenantPolicy(), BudgetPolicy()],
|
||||
)
|
||||
|
||||
response = await router.acompletion(
|
||||
model="smart-router",
|
||||
messages=[{"role": "user", "content": "Write a function to reverse a linked list."}],
|
||||
metadata={"tenant": "acme-corp"},
|
||||
)
|
||||
|
||||
# response.model is the bare model name (litellm strips the provider/ prefix
|
||||
# on the response), so compare against bare names rather than litellm_params.model
|
||||
routed_model = response.model
|
||||
|
||||
assert routed_model in {"gpt-4o-mini", "claude-haiku-4-5"}
|
||||
assert routed_model not in {"llama-3-70b", "gpt-5.1"}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_routing_plugin_narrowing_to_zero_candidates_raises():
|
||||
"""A plugin narrowing to nothing is a policy decision -- must raise, not silently
|
||||
fall back to the unfiltered pool (that would defeat the policy it enforces)."""
|
||||
router = Router(
|
||||
model_list=_smart_router_model_list(),
|
||||
plugins=[BlockEverything()],
|
||||
)
|
||||
|
||||
with pytest.raises(ValueError, match="No deployments left after routing-plugin filtering"):
|
||||
await router.acompletion(
|
||||
model="smart-router",
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
)
|
||||
|
||||
|
||||
def test_sync_get_available_deployment_rejects_configured_plugins():
|
||||
"""
|
||||
Router.completion() (and any other sync entry point) resolves deployments via
|
||||
the synchronous get_available_deployment(), which never runs the routing-plugin
|
||||
pipeline. Silently allowing that would let a deny-all policy plugin be bypassed
|
||||
just by calling the sync API -- must fail closed instead.
|
||||
"""
|
||||
router = Router(model_list=_smart_router_model_list(), plugins=[TenantPolicy()])
|
||||
|
||||
with pytest.raises(ValueError, match="routing-plugin pipeline"):
|
||||
router.get_available_deployment(model="smart-router", messages=[{"role": "user", "content": "hi"}])
|
||||
|
||||
|
||||
def test_sync_router_completion_rejects_configured_plugins():
|
||||
"""End-to-end: Router.completion() (the sync API) must not silently skip plugins either."""
|
||||
router = Router(model_list=_smart_router_model_list(), plugins=[TenantPolicy()])
|
||||
|
||||
with pytest.raises(ValueError, match="routing-plugin pipeline"):
|
||||
router.completion(model="smart-router", messages=[{"role": "user", "content": "hi"}])
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_completion_with_unsupported_strategy_rejects_configured_plugins():
|
||||
"""
|
||||
async_get_available_deployment() itself delegates to the synchronous selector
|
||||
for routing strategies outside {simple-shuffle, usage-based-routing-v2,
|
||||
cost-based-routing, latency-based-routing, least-busy} -- e.g. "usage-based-routing"
|
||||
(v1, not v2) -- which would silently bypass the plugin pipeline on the async path too.
|
||||
"""
|
||||
router = Router(
|
||||
model_list=_smart_router_model_list(),
|
||||
plugins=[TenantPolicy()],
|
||||
routing_strategy="usage-based-routing",
|
||||
)
|
||||
|
||||
with pytest.raises(ValueError, match="routing-plugin pipeline"):
|
||||
await router.acompletion(model="smart-router", messages=[{"role": "user", "content": "hi"}])
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_router_without_plugins_is_unaffected():
|
||||
"""Regression guard: a Router with no `plugins` configured behaves exactly as before."""
|
||||
router = Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "smart-router",
|
||||
"litellm_params": {"model": "openai/gpt-4o-mini", "mock_response": "hi"},
|
||||
},
|
||||
],
|
||||
)
|
||||
response = await router.acompletion(
|
||||
model="smart-router",
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
)
|
||||
assert response.choices[0].message.content == "hi"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_run_routing_plugins_narrows_candidates_and_records_signals():
|
||||
"""Unit-level check of _run_routing_plugins in isolation, independent of acompletion."""
|
||||
router = Router(
|
||||
model_list=_smart_router_model_list(),
|
||||
plugins=[LanguageDetector(), DomainClassifier(), TenantPolicy(), BudgetPolicy()],
|
||||
)
|
||||
request_kwargs = {"metadata": {"tenant": "acme-corp"}}
|
||||
|
||||
context = await router._run_routing_plugins(
|
||||
model="smart-router",
|
||||
request_kwargs=request_kwargs,
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
)
|
||||
|
||||
assert context.candidate_models == ["openai/gpt-4o-mini", "anthropic/claude-haiku-4-5"]
|
||||
assert context.signals["domain-classifier"]["domain"] == "coding"
|
||||
assert request_kwargs["metadata"]["_routing_plugin_candidate_models"] == context.candidate_models
|
||||
|
||||
|
||||
def test_filter_by_routing_plugin_candidates_narrows_and_raises_when_empty():
|
||||
"""Unit-level check of _filter_by_routing_plugin_candidates in isolation."""
|
||||
router = Router(model_list=_smart_router_model_list(), plugins=[TenantPolicy()])
|
||||
healthy_deployments = router.model_list
|
||||
|
||||
narrowed = router._filter_by_routing_plugin_candidates(
|
||||
healthy_deployments=healthy_deployments,
|
||||
request_kwargs={"metadata": {"_routing_plugin_candidate_models": ["openai/gpt-4o-mini"]}},
|
||||
)
|
||||
assert [d["litellm_params"]["model"] for d in narrowed] == ["openai/gpt-4o-mini"]
|
||||
|
||||
with pytest.raises(ValueError, match="No deployments left after routing-plugin filtering"):
|
||||
router._filter_by_routing_plugin_candidates(
|
||||
healthy_deployments=healthy_deployments,
|
||||
request_kwargs={"metadata": {"_routing_plugin_candidate_models": ["nonexistent/model"]}},
|
||||
)
|
||||
Loading…
Add table
Reference in a new issue