fix(router): support tag-based routing for complexity routers sharing a model_name

This commit is contained in:
Devin AI 2026-07-17 05:20:02 +00:00
parent 4d33964898
commit bcf3c61611
6 changed files with 269 additions and 20 deletions

View file

@ -13,11 +13,16 @@ model/{model_id}/update - PATCH endpoint for model update.
import asyncio
import datetime
import json
from typing import Any, Dict, List, Literal, Optional, Set, Tuple, Union, cast
from typing import TYPE_CHECKING, Any, Dict, List, Literal, Optional, Set, Tuple, Union, cast
from fastapi import APIRouter, Depends, HTTPException, Header, Request, status
from pydantic import BaseModel, ConfigDict, Field
if TYPE_CHECKING:
from litellm.router_strategy.complexity_router.complexity_router import (
ComplexityRouter,
)
from litellm._logging import verbose_proxy_logger
from litellm._uuid import uuid
from litellm.constants import LITELLM_PROXY_ADMIN_NAME
@ -1056,6 +1061,27 @@ def _deployment_name_and_model(deployment: Optional[Union[Deployment, Dict[str,
return deployment.model_name, str(getattr(deployment.litellm_params, "model", "") or "")
def _evict_complexity_router(
complexity_routers: dict[str, list["ComplexityRouter"]],
model_name: str,
model_id: str | None,
) -> None:
"""Drop only the complexity router backed by `model_id`.
Several complexity routers may be registered under one model_name (tag-based
routing), so evict just the deleted deployment's router and keep the siblings; pop
the whole entry only once the last one is gone.
"""
existing = complexity_routers.get(model_name)
if not existing:
return
remaining = [router for router in existing if router.model_id != model_id]
if remaining:
complexity_routers[model_name] = remaining
else:
complexity_routers.pop(model_name, None)
#### [BETA] - This is a beta endpoint, format might change based on user feedback. - https://github.com/BerriAI/litellm/issues/964
@router.post(
"/model/delete",
@ -1137,7 +1163,7 @@ async def delete_model(
deleted_name, deleted_model = _deployment_name_and_model(deleted_deployment)
if deleted_name is not None and deleted_model.startswith("auto_router/"):
llm_router.auto_routers.pop(deleted_name, None)
llm_router.complexity_routers.pop(deleted_name, None)
_evict_complexity_router(llm_router.complexity_routers, deleted_name, model_info.id)
llm_router.adaptive_routers.pop(deleted_name, None)
llm_router.quality_routers.pop(deleted_name, None)

View file

@ -86,7 +86,10 @@ from litellm.router_strategy.lowest_latency import LowestLatencyLoggingHandler
from litellm.router_strategy.lowest_tpm_rpm import LowestTPMLoggingHandler
from litellm.router_strategy.lowest_tpm_rpm_v2 import LowestTPMLoggingHandler_v2
from litellm.router_strategy.simple_shuffle import simple_shuffle
from litellm.router_strategy.tag_based_routing import get_deployments_for_tag
from litellm.router_strategy.tag_based_routing import (
get_deployments_for_tag,
select_index_by_tags,
)
from litellm.router_utils.add_retry_fallback_headers import (
_HiddenParamsHost,
add_fallback_headers_to_response,
@ -177,6 +180,7 @@ from litellm.types.router import (
OptionalPreCallChecks,
RetryPolicy,
RouterCacheEnum,
RouterErrors,
RouterGeneralSettings,
RouterModelGroupAliasItem,
RouterRateLimitError,
@ -488,7 +492,7 @@ class Router:
self.pattern_router = PatternMatchRouter()
self.team_pattern_routers: Dict[str, PatternMatchRouter] = {} # {"TEAM_ID": PatternMatchRouter}
self.auto_routers: Dict[str, "AutoRouter"] = {}
self.complexity_routers: Dict[str, "ComplexityRouter"] = {}
self.complexity_routers: dict[str, list["ComplexityRouter"]] = {}
self.adaptive_routers: Dict[str, "AdaptiveRouter"] = {}
self.quality_routers: Dict[str, "QualityRouter"] = {}
self.routing_plugins: list[RoutingPlugin] = list(plugins) if plugins else []
@ -7657,12 +7661,10 @@ class Router:
default_model=default_model,
litellm_router_instance=self,
complexity_router_config=complexity_router_config,
tags=deployment.litellm_params.tags,
model_id=deployment.model_info.id if deployment.model_info else None,
)
if deployment.model_name in self.complexity_routers:
raise ValueError(
f"Complexity-router deployment {deployment.model_name} already exists. Please use a different model name."
)
self.complexity_routers[deployment.model_name] = complexity_router
self.complexity_routers.setdefault(deployment.model_name, []).append(complexity_router)
def _is_adaptive_router_deployment(self, litellm_params: LiteLLM_Params) -> bool:
"""True when this deployment opts in via the `auto_router/adaptive_router` model prefix."""
@ -7698,12 +7700,13 @@ class Router:
continue
self.init_adaptive_router_deployment(deployment=deployment)
for model_name, complexity_router in self.complexity_routers.items():
if not complexity_router.config.adaptive or model_name in self.adaptive_routers:
continue
adaptive_router = complexity_router._ensure_adaptive_router()
if adaptive_router is not None:
self.adaptive_routers[model_name] = adaptive_router
for model_name, complexity_routers in self.complexity_routers.items():
for complexity_router in complexity_routers:
if not complexity_router.config.adaptive or model_name in self.adaptive_routers:
continue
adaptive_router = complexity_router._ensure_adaptive_router()
if adaptive_router is not None:
self.adaptive_routers[model_name] = adaptive_router
for callback in litellm.logging_callback_manager.get_custom_loggers_for_type(AdaptiveRouterPostCallHook):
litellm.logging_callback_manager.remove_callback_from_all_lists(callback)
@ -10832,9 +10835,10 @@ class Router:
if self.routing_plugins:
await self._run_routing_plugins(model=model, request_kwargs=request_kwargs, messages=messages)
complexity_router = self._select_complexity_router(model=model, request_kwargs=request_kwargs)
router_strategy = (
self.auto_routers.get(model)
or self.complexity_routers.get(model)
or complexity_router
or self.adaptive_routers.get(model)
or self.quality_routers.get(model)
)
@ -10857,15 +10861,65 @@ class Router:
# actual outbound LLM call downstream by litellm.types.utils.all_litellm_params,
# not here.
if pre_routing_hook_response is not None:
alias_index = self.model_name_to_deployment_indices.get(model, [])
if alias_index:
alias_litellm_params = self.model_list[alias_index[0]].get("litellm_params", {})
selected_model_id = complexity_router.model_id if complexity_router is router_strategy else None
alias_index = self._resolve_alias_index(model=model, selected_model_id=selected_model_id)
if alias_index is not None:
alias_litellm_params = self.model_list[alias_index].get("litellm_params", {})
for key, value in alias_litellm_params.items():
if key != "model" and value is not None:
request_kwargs.setdefault(key, value)
return pre_routing_hook_response
def _resolve_alias_index(self, model: str, selected_model_id: str | None) -> int | None:
"""Index into `model_list` of the alias deployment whose params to apply.
When several deployments share `model` (tag-based complexity routing), prefer the
one whose model_info id matches the router the hook actually selected so the right
tags/params are applied, falling back to the first alias.
"""
alias_indices = self.model_name_to_deployment_indices.get(model, [])
if not alias_indices:
return None
if selected_model_id is not None:
for index in alias_indices:
model_info = self.model_list[index].get("model_info") or {}
if model_info.get("id") == selected_model_id:
return index
return alias_indices[0]
def _select_complexity_router(self, model: str, request_kwargs: dict) -> "ComplexityRouter | None":
"""Pick the complexity router for `model`, honouring tag-based routing.
Multiple complexity-router deployments can share a model_name while carrying
different tags. When tag filtering is active, select the router whose tags match
the request's tags (same semantics as deployment tag routing); raise the standard
tag-routing error when nothing matches so callers get a clear 401 instead of being
silently routed through the wrong tier config.
"""
routers = self.complexity_routers.get(model)
if not routers:
return None
if len(routers) == 1:
return routers[0]
request_enable_tag_filtering = request_kwargs.get("enable_tag_filtering")
if request_enable_tag_filtering is not True and self.enable_tag_filtering is not True:
return routers[0]
metadata_variable_name = self._get_metadata_variable_name_from_kwargs(request_kwargs)
request_tags = (request_kwargs.get(metadata_variable_name) or {}).get("tags") or []
selected_index = select_index_by_tags(
tags_per_candidate=[router.tags for router in routers],
request_tags=request_tags,
match_any=self.tag_filtering_match_any,
)
if selected_index is None:
raise ValueError(
f"{RouterErrors.no_deployments_with_tag_routing.value}. Passed model={model} and tags={request_tags}"
)
return routers[selected_index]
def get_available_deployment(
self,
model: str,

View file

@ -142,6 +142,8 @@ class ComplexityRouter(CustomLogger):
litellm_router_instance: Router,
complexity_router_config: dict[str, Any] | None = None,
default_model: str | None = None,
tags: list[str] | None = None,
model_id: str | None = None,
):
"""
Initialize ComplexityRouter.
@ -151,8 +153,14 @@ class ComplexityRouter(CustomLogger):
litellm_router_instance: The LiteLLM Router instance.
complexity_router_config: Optional configuration dict from proxy config.
default_model: Optional default model to use if tier cannot be determined.
tags: The deployment's tags, used to pick the right router when several
complexity-router deployments share a model_name (tag-based routing).
model_id: The deployment's model_info id, used to evict the exact router
when its backing deployment is deleted.
"""
self.model_name = model_name
self.tags = tags
self.model_id = model_id
self.litellm_router_instance = litellm_router_instance
# Parse config - always create a new instance to avoid singleton mutation

View file

@ -114,6 +114,51 @@ def _match_deployment(
return None
def select_index_by_tags(
tags_per_candidate: list[list[str] | None],
request_tags: list[str],
match_any: bool,
) -> int | None:
"""
Pick the single best candidate index for tag-based routing among candidates that
share a model group (e.g. several complexity-router deployments registered under
the same model_name, each carrying different tags).
Mirrors get_deployments_for_tag's selection semantics for the single-pick case:
- `!tag` entries exclude a candidate outright.
- a positive request tag selects the first candidate whose tags match
(match_any / match_all per `match_any`); if none match, a `default`-tagged
candidate wins; otherwise there is no match.
- an untagged request prefers a `default`-tagged candidate, else the first
remaining candidate.
Returns the chosen index, or None when request tags are given but neither a tag
match nor a default candidate exists (caller decides how to surface that).
"""
positive_tags, excluded_patterns = _split_tags(request_tags or [])
excluded_set = frozenset(excluded_patterns)
allowed = [i for i, tags in enumerate(tags_per_candidate) if not excluded_set.intersection(tags or [])]
if not positive_tags:
default_index = next((i for i in allowed if "default" in (tags_per_candidate[i] or [])), None)
if default_index is not None:
return default_index
return allowed[0] if allowed else None
matched_index = next(
(
i
for i in allowed
if tags_per_candidate[i] and is_valid_deployment_tag(tags_per_candidate[i] or [], positive_tags, match_any)
),
None,
)
if matched_index is not None:
return matched_index
return next((i for i in allowed if "default" in (tags_per_candidate[i] or [])), None)
def _split_tags(tags: list[str]) -> tuple[list[str], list[str]]:
positive = [t for t in tags if not t.startswith("!")]
excluded = [tag[1:] for tag in tags if tag.startswith("!") and len(tag) > 1]

View file

@ -668,8 +668,10 @@ class TestDeleteModelClearsRouterRegistry:
"model_info": {"id": model_id},
}
)
deleted_complexity_router = MagicMock()
deleted_complexity_router.model_id = model_id
mock_router.auto_routers = {"smart-router": MagicMock()}
mock_router.complexity_routers = {"smart-router": MagicMock()}
mock_router.complexity_routers = {"smart-router": [deleted_complexity_router]}
_PS = "litellm.proxy.proxy_server"
with (
@ -750,6 +752,27 @@ class TestDeleteModelClearsRouterRegistry:
mock_router.delete_deployment.assert_called_once_with(id=model_id)
assert mock_router.complexity_routers.get("shared-name") is config_router
def test_evict_complexity_router_keeps_sibling_tagged_routers(self):
"""Several complexity routers can share a model_name via different tags. Deleting
one deployment must evict only its router (matched by model_id) and keep the rest,
popping the whole entry only when the last sibling is gone.
"""
from litellm.proxy.management_endpoints.model_management_endpoints import (
_evict_complexity_router,
)
cn_router = MagicMock()
cn_router.model_id = "cn-id"
row_router = MagicMock()
row_router.model_id = "row-id"
registry = {"smart-router": [cn_router, row_router]}
_evict_complexity_router(registry, "smart-router", "cn-id")
assert registry["smart-router"] == [row_router]
_evict_complexity_router(registry, "smart-router", "row-id")
assert "smart-router" not in registry
class TestUpdateModel:
"""

View file

@ -1517,6 +1517,99 @@ class TestRouterPreRoutingAliasOverrides:
assert request_kwargs["drop_params"] is True
class TestComplexityRouterTagBasedRouting:
"""
Regression tests for: two complexity-router deployments sharing a model_name but
carrying different tags used to collapse to the first-registered config (or raise
"already exists"), so a request whose key had the second tag was silently routed
through the wrong tier config. Each tag must resolve to its own config.
"""
def _make_router(self, **router_kwargs) -> Router:
return Router(
model_list=[
{
"model_name": "smart-router",
"model_info": {"id": "cn-router"},
"litellm_params": {
"model": "auto_router/complexity_router",
"tags": ["cn"],
"complexity_router_config": {
"tiers": {"SIMPLE": "cn-simple", "MEDIUM": "cn-medium"}
},
"complexity_router_default_model": "cn-medium",
},
},
{
"model_name": "smart-router",
"model_info": {"id": "row-router"},
"litellm_params": {
"model": "auto_router/complexity_router",
"tags": ["row"],
"complexity_router_config": {
"tiers": {"SIMPLE": "row-simple", "MEDIUM": "row-medium"}
},
"complexity_router_default_model": "row-medium",
},
},
{"model_name": "cn-simple", "litellm_params": {"model": "openai/cn-simple"}},
{"model_name": "row-simple", "litellm_params": {"model": "openai/row-simple"}},
],
**router_kwargs,
)
def test_both_tagged_routers_are_registered(self):
router = self._make_router(enable_tag_filtering=True)
routers = router.complexity_routers["smart-router"]
assert len(routers) == 2
assert {tuple(r.tags or []) for r in routers} == {("cn",), ("row",)}
@pytest.mark.asyncio
async def test_request_tag_selects_matching_config(self):
router = self._make_router(enable_tag_filtering=True)
cn_kwargs: Dict = {"metadata": {"tags": ["cn"]}}
cn_result = await router.async_pre_routing_hook(
model="smart-router",
request_kwargs=cn_kwargs,
messages=[{"role": "user", "content": "hi"}],
)
assert cn_result is not None
assert cn_result.model == "cn-simple"
row_kwargs: Dict = {"metadata": {"tags": ["row"]}}
row_result = await router.async_pre_routing_hook(
model="smart-router",
request_kwargs=row_kwargs,
messages=[{"role": "user", "content": "hi"}],
)
assert row_result is not None
assert row_result.model == "row-simple"
@pytest.mark.asyncio
async def test_unmatched_tag_raises_tag_routing_error(self):
from litellm.types.router import RouterErrors
router = self._make_router(enable_tag_filtering=True)
with pytest.raises(ValueError, match=RouterErrors.no_deployments_with_tag_routing.value):
await router.async_pre_routing_hook(
model="smart-router",
request_kwargs={"metadata": {"tags": ["eu"]}},
messages=[{"role": "user", "content": "hi"}],
)
@pytest.mark.asyncio
async def test_tag_filtering_disabled_falls_back_to_first(self):
router = self._make_router(enable_tag_filtering=False)
result = await router.async_pre_routing_hook(
model="smart-router",
request_kwargs={"metadata": {"tags": ["row"]}},
messages=[{"role": "user", "content": "hi"}],
)
assert result is not None
assert result.model == "cn-simple"
class TestAdaptiveSoftFloors:
def test_adaptive_defaults_use_cost_weighted_cold_policy(self):
config = ComplexityRouterConfig(