diff --git a/litellm/proxy/management_endpoints/model_management_endpoints.py b/litellm/proxy/management_endpoints/model_management_endpoints.py index d458d0f7c4a..a3f2c4c4359 100644 --- a/litellm/proxy/management_endpoints/model_management_endpoints.py +++ b/litellm/proxy/management_endpoints/model_management_endpoints.py @@ -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) diff --git a/litellm/router.py b/litellm/router.py index 78e156801f8..c3d3645fd29 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -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, diff --git a/litellm/router_strategy/complexity_router/complexity_router.py b/litellm/router_strategy/complexity_router/complexity_router.py index fa6f14e9b26..8013afdb9cb 100644 --- a/litellm/router_strategy/complexity_router/complexity_router.py +++ b/litellm/router_strategy/complexity_router/complexity_router.py @@ -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 diff --git a/litellm/router_strategy/tag_based_routing.py b/litellm/router_strategy/tag_based_routing.py index 710c2199107..c16586c0d10 100644 --- a/litellm/router_strategy/tag_based_routing.py +++ b/litellm/router_strategy/tag_based_routing.py @@ -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] diff --git a/tests/test_litellm/proxy/management_endpoints/test_model_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_model_management_endpoints.py index 8c6bdefedae..b59e4312129 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_model_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_model_management_endpoints.py @@ -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: """ diff --git a/tests/test_litellm/router_strategy/test_complexity_router.py b/tests/test_litellm/router_strategy/test_complexity_router.py index 12b2c9abefb..e8f947e0449 100644 --- a/tests/test_litellm/router_strategy/test_complexity_router.py +++ b/tests/test_litellm/router_strategy/test_complexity_router.py @@ -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(