From 85f9bdd4129588cdf47c746978fc4b658faec747 Mon Sep 17 00:00:00 2001 From: Krrish Dholakia Date: Sat, 11 Jul 2026 21:38:18 -0700 Subject: [PATCH] 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 --- .../prompt_templates/factory.py | 53 +++++ litellm/router.py | 101 ++++++++ .../complexity_router/complexity_router.py | 38 +-- litellm/types/router.py | 31 ++- .../test_router_routing_plugins.py | 220 ++++++++++++++++++ 5 files changed, 407 insertions(+), 36 deletions(-) create mode 100644 tests/test_litellm/router_strategy/test_router_routing_plugins.py diff --git a/litellm/litellm_core_utils/prompt_templates/factory.py b/litellm/litellm_core_utils/prompt_templates/factory.py index 8bb0e12905e..f7ff4d6b16f 100644 --- a/litellm/litellm_core_utils/prompt_templates/factory.py +++ b/litellm/litellm_core_utils/prompt_templates/factory.py @@ -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 diff --git a/litellm/router.py b/litellm/router.py index 6e773a06c7f..245a50545e7 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -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 diff --git a/litellm/router_strategy/complexity_router/complexity_router.py b/litellm/router_strategy/complexity_router/complexity_router.py index 2138a0112a0..74644f01be8 100644 --- a/litellm/router_strategy/complexity_router/complexity_router.py +++ b/litellm/router_strategy/complexity_router/complexity_router.py @@ -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( diff --git a/litellm/types/router.py b/litellm/types/router.py index 4bac9358392..3bedd97c20c 100644 --- a/litellm/types/router.py +++ b/litellm/types/router.py @@ -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.""" diff --git a/tests/test_litellm/router_strategy/test_router_routing_plugins.py b/tests/test_litellm/router_strategy/test_router_routing_plugins.py new file mode 100644 index 00000000000..e9c12d009e2 --- /dev/null +++ b/tests/test_litellm/router_strategy/test_router_routing_plugins.py @@ -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"]}}, + )