From 1e04aee089e2495a2917ab0f7972fffd6959e167 Mon Sep 17 00:00:00 2001 From: Tin Chi Lo Date: Tue, 21 Jul 2026 12:57:10 -0700 Subject: [PATCH] fix(proxy): reject model writes that corrupt an auto-router pseudo-model An auto-router deployment's litellm_params.model (auto_router/...) is the discriminator the router loads it by, but the model management endpoints accepted any client-supplied value verbatim; a doubled or stripped prefix made router init fail on the next load and ignore_invalid_deployments silently dropped the deployment. Validate writes that supply litellm_params.model at all three endpoints against the merged params and reject incoherent values with an actionable 400. Classification is extracted to router_utils/auto_router_model_naming.py so the Router predicates and the validation share one source --- .../model_management_endpoints.py | 59 +++++ litellm/router.py | 23 +- .../router_utils/auto_router_model_naming.py | 101 +++++++++ .../test_model_management_endpoints.py | 213 ++++++++++++++++++ .../test_auto_router_model_naming.py | 77 +++++++ 5 files changed, 457 insertions(+), 16 deletions(-) create mode 100644 litellm/router_utils/auto_router_model_naming.py create mode 100644 tests/test_litellm/router_utils/test_auto_router_model_naming.py diff --git a/litellm/proxy/management_endpoints/model_management_endpoints.py b/litellm/proxy/management_endpoints/model_management_endpoints.py index 28c406edf76..1c0e7211493 100644 --- a/litellm/proxy/management_endpoints/model_management_endpoints.py +++ b/litellm/proxy/management_endpoints/model_management_endpoints.py @@ -57,10 +57,15 @@ from litellm.router import Router from litellm.types.proxy.management_endpoints.model_management_endpoints import ( UpdateUsefulLinksRequest, ) +from litellm.router_utils.auto_router_model_naming import ( + STRATEGY_ROUTER_PARAM_FIELDS, + validate_strategy_router_model_write, +) from litellm.types.router import ( SPECIAL_MODEL_INFO_PARAMS, Deployment, DeploymentTypedDict, + GenericLiteLLMParams, LiteLLMParamsTypedDict, updateDeployment, ) @@ -98,6 +103,45 @@ async def get_db_model(model_id: str, prisma_client: PrismaClient) -> Optional[D return deployment_pydantic_obj +def _strategy_router_write_violation( + incoming_params: GenericLiteLLMParams | None, + existing_params: GenericLiteLLMParams | None, +) -> str | None: + """Reject writes that would corrupt a strategy router's pseudo-model. + + An auto-router deployment's ``litellm_params.model`` (``auto_router/...``) is + the discriminator the router loads it by; a write that mangles it makes the + router drop the deployment silently under ``ignore_invalid_deployments``. + Only writes that supply ``litellm_params.model`` are judged, against the + merged (stored + incoming) params, so partial patches and restores of an + already-corrupted row stay legal. Returns the violation, or None. + """ + if incoming_params is None or incoming_params.model is None: + return None + present_fields = frozenset( + field + for field in STRATEGY_ROUTER_PARAM_FIELDS + for source in (incoming_params, existing_params) + if source is not None and getattr(source, field, None) is not None + ) + return validate_strategy_router_model_write(model=incoming_params.model, present_fields=present_fields) + + +def _raise_on_strategy_router_write_violation( + incoming_params: GenericLiteLLMParams | None, + existing_params: GenericLiteLLMParams | None, +) -> None: + violation = _strategy_router_write_violation(incoming_params=incoming_params, existing_params=existing_params) + if violation is None: + return + raise ProxyException( + message=violation, + type=ProxyErrorTypes.validation_error.value, + code=status.HTTP_400_BAD_REQUEST, + param="litellm_params.model", + ) + + def update_db_model(db_model: Deployment, updated_patch: updateDeployment) -> PrismaCompatibleUpdateDBModel: merged_deployment_dict = DeploymentTypedDict( model_name=db_model.model_name, @@ -255,6 +299,11 @@ async def patch_model( param="blocked", ) + _raise_on_strategy_router_write_violation( + incoming_params=patch_data.litellm_params, + existing_params=db_model.litellm_params, + ) + # Handle team model updates with proper alias management update_data = await _update_team_model_in_db( db_model=db_model, @@ -1292,6 +1341,11 @@ async def add_new_model( premium_user=premium_user, ) + _raise_on_strategy_router_write_violation( + incoming_params=model_params.litellm_params, + existing_params=None, + ) + model_response: Optional[LiteLLM_ProxyModelTable] = None # update DB if store_model_in_db is True: @@ -1446,6 +1500,11 @@ async def update_model( premium_user=premium_user, ) + _raise_on_strategy_router_write_violation( + incoming_params=model_params.litellm_params, + existing_params=deployment.litellm_params, + ) + # update DB if store_model_in_db is True: _existing_litellm_params_dict = dict(_existing_litellm_params.litellm_params) diff --git a/litellm/router.py b/litellm/router.py index d3caa069b28..759d6a0024c 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -110,6 +110,9 @@ from litellm.router_utils.batch_utils import ( replace_model_in_jsonl, should_replace_model_in_jsonl, ) +from litellm.router_utils.auto_router_model_naming import ( + classify_strategy_router_model, +) from litellm.router_utils.client_initalization_utils import InitalizeCachedClient from litellm.router_utils.clientside_credential_handler import ( get_dynamic_litellm_params, @@ -7623,15 +7626,7 @@ class Router: but NOT "auto_router/complexity_router" or "auto_router/adaptive_router" (which use the complexity-router and adaptive-router strategies). """ - if litellm_params.model.startswith("auto_router/complexity_router"): - return False # This is handled by complexity_router - if litellm_params.model.startswith("auto_router/adaptive_router"): - return False # This is handled by adaptive_router - if litellm_params.model.startswith("auto_router/quality_router"): - return False # This is handled by quality_router - if litellm_params.model.startswith("auto_router/"): - return True - return False + return classify_strategy_router_model(litellm_params.model) == "semantic" @staticmethod def _deployment_tags(deployment: Deployment) -> tuple[str, ...]: @@ -7686,9 +7681,7 @@ class Router: Returns True if the litellm_params model starts with "auto_router/complexity_router" """ - if litellm_params.model.startswith("auto_router/complexity_router"): - return True - return False + return classify_strategy_router_model(litellm_params.model) == "complexity" def init_complexity_router_deployment(self, deployment: Deployment): """ @@ -7738,7 +7731,7 @@ class 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.""" - return litellm_params.model.startswith("auto_router/adaptive_router") + return classify_strategy_router_model(litellm_params.model) == "adaptive" def _deployment_participates_in_adaptive_routing(self, litellm_params: LiteLLM_Params) -> bool: """True when this deployment owns an `adaptive_routers` entry once finalized: @@ -7964,9 +7957,7 @@ class Router: Returns True if the litellm_params model starts with "auto_router/quality_router". """ - if litellm_params.model.startswith("auto_router/quality_router"): - return True - return False + return classify_strategy_router_model(litellm_params.model) == "quality" def init_quality_router_deployment(self, deployment: Deployment): """ diff --git a/litellm/router_utils/auto_router_model_naming.py b/litellm/router_utils/auto_router_model_naming.py new file mode 100644 index 00000000000..72865501030 --- /dev/null +++ b/litellm/router_utils/auto_router_model_naming.py @@ -0,0 +1,101 @@ +"""Naming contract for strategy-router (auto-router) pseudo-models. + +A deployment whose ``litellm_params.model`` starts with ``auto_router/`` does not +name a provider model; the string is the discriminator that selects which +pre-routing strategy owns the deployment. This module is the single source of +truth for classifying that string (``Router._is_*_router_deployment`` delegates +here) and for checking that a client-supplied write leaves the deployment +coherent, so management endpoints can reject corruption with a 400 instead of +the router silently dropping the deployment at load time under +``ignore_invalid_deployments``. +""" + +from typing import Literal, Mapping + +AUTO_ROUTER_MODEL_PREFIX = "auto_router/" + +StrategyRouterKind = Literal["semantic", "complexity", "adaptive", "quality"] + +STRATEGY_ROUTER_PARAM_FIELDS: frozenset[str] = frozenset( + { + "auto_router_config", + "auto_router_config_path", + "auto_router_default_model", + "auto_router_embedding_model", + "complexity_router_config", + "complexity_router_default_model", + "adaptive_router_config", + "quality_router_config", + "quality_router_default_model", + } +) + +_REQUIRED_FIELD_GROUPS: Mapping[StrategyRouterKind, tuple[tuple[str, ...], ...]] = { + "semantic": ( + ("auto_router_config", "auto_router_config_path"), + ("auto_router_default_model",), + ("auto_router_embedding_model",), + ), + "complexity": (("complexity_router_config", "complexity_router_default_model"),), + "adaptive": (("adaptive_router_config",),), + "quality": (("quality_router_config", "quality_router_default_model"),), +} + + +def classify_strategy_router_model(model: str) -> StrategyRouterKind | None: + """Classify a ``litellm_params.model`` string the way the Router does. + + Returns None for regular provider models. Mirrors Router registration + exactly: reserved names are matched by prefix, everything else under + ``auto_router/`` is a semantic router. + """ + if not model.startswith(AUTO_ROUTER_MODEL_PREFIX): + return None + remainder = model[len(AUTO_ROUTER_MODEL_PREFIX) :] + if remainder.startswith("complexity_router"): + return "complexity" + if remainder.startswith("adaptive_router"): + return "adaptive" + if remainder.startswith("quality_router"): + return "quality" + return "semantic" + + +def validate_strategy_router_model_write(model: str, present_fields: frozenset[str]) -> str | None: + """Check that writing ``model`` leaves a deployment the router can load. + + ``present_fields`` is the set of strategy-router param fields that are + non-None on the deployment after the write (stored fields merged with the + incoming ones). Returns a human-readable violation, or None when coherent. + """ + kind = classify_strategy_router_model(model) + if kind is None: + offending = sorted(present_fields & STRATEGY_ROUTER_PARAM_FIELDS) + if offending: + return ( + f"litellm_params.model='{model}' does not start with '{AUTO_ROUTER_MODEL_PREFIX}' but the " + f"deployment carries auto-router settings ({', '.join(offending)}), so the router could not " + f"load it. Keep the '{AUTO_ROUTER_MODEL_PREFIX}' prefix; to change the name clients call, " + "edit the public model_name instead." + ) + return None + remainder = model[len(AUTO_ROUTER_MODEL_PREFIX) :] + if remainder.startswith(AUTO_ROUTER_MODEL_PREFIX): + return ( + f"litellm_params.model='{model}' repeats the '{AUTO_ROUTER_MODEL_PREFIX}' prefix, so the router " + f"could not load it. Use '{remainder}'; to change the name clients call, edit the public " + "model_name instead." + ) + if not remainder: + return ( + f"litellm_params.model='{model}' is missing the router name after the '{AUTO_ROUTER_MODEL_PREFIX}' prefix." + ) + missing = tuple( + " or ".join(group) for group in _REQUIRED_FIELD_GROUPS[kind] if not any(f in present_fields for f in group) + ) + if missing: + return ( + f"litellm_params.model='{model}' selects the {kind} router, which requires " + f"{'; '.join(missing)} in litellm_params." + ) + return None 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 3bd83ade20a..1e8add52f74 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 @@ -3330,3 +3330,216 @@ class TestModelInfoAsMapping: assert model_info_as_mapping("{not json") is None assert model_info_as_mapping('["a", "b"]') is None assert model_info_as_mapping(42) is None + + +class TestStrategyRouterWriteValidation: + """Management write paths must reject litellm_params.model values that would + corrupt a strategy router's pseudo-model (LIT-4663). The router loads these + deployments by the auto_router/ discriminator, so a mangled string makes it + drop the deployment silently under ignore_invalid_deployments; the mistake + has to fail loudly at the API boundary instead.""" + + def _stored_complexity_params(self) -> LiteLLM_Params: + return LiteLLM_Params( + model="auto_router/complexity_router", + complexity_router_config={"tiers": {"SIMPLE": "gpt-4o-mini"}}, + ) + + def _db_complexity_router(self, model_id: str) -> Deployment: + return Deployment( + model_name="my-auto-router", + litellm_params=self._stored_complexity_params(), + model_info={"id": model_id}, + ) + + def test_double_prefix_rejected_against_stored_params(self): + from litellm.proxy.management_endpoints.model_management_endpoints import ( + _strategy_router_write_violation, + ) + from litellm.types.router import updateLiteLLMParams + + violation = _strategy_router_write_violation( + incoming_params=updateLiteLLMParams(model="auto_router/auto_router/complexity_router"), + existing_params=self._stored_complexity_params(), + ) + assert violation is not None + assert "repeats" in violation + + def test_prefix_strip_rejected_against_stored_params(self): + from litellm.proxy.management_endpoints.model_management_endpoints import ( + _strategy_router_write_violation, + ) + from litellm.types.router import updateLiteLLMParams + + violation = _strategy_router_write_violation( + incoming_params=updateLiteLLMParams(model="complexity_router"), + existing_params=self._stored_complexity_params(), + ) + assert violation is not None + assert "does not start with" in violation + + def test_patch_without_model_is_not_judged(self): + from litellm.proxy.management_endpoints.model_management_endpoints import ( + _strategy_router_write_violation, + ) + from litellm.types.router import updateLiteLLMParams + + assert ( + _strategy_router_write_violation( + incoming_params=updateLiteLLMParams(rpm=10), + existing_params=self._stored_complexity_params(), + ) + is None + ) + assert _strategy_router_write_violation(incoming_params=None, existing_params=None) is None + + def test_restore_of_corrupted_row_is_allowed(self): + from litellm.proxy.management_endpoints.model_management_endpoints import ( + _strategy_router_write_violation, + ) + from litellm.types.router import updateLiteLLMParams + + corrupted = LiteLLM_Params( + model="auto_router/auto_router/complexity_router", + complexity_router_config={"tiers": {"SIMPLE": "gpt-4o-mini"}}, + ) + assert ( + _strategy_router_write_violation( + incoming_params=updateLiteLLMParams(model="auto_router/complexity_router"), + existing_params=corrupted, + ) + is None + ) + + def test_create_semantic_router_missing_embedding_rejected(self): + from litellm.proxy.management_endpoints.model_management_endpoints import ( + _strategy_router_write_violation, + ) + + violation = _strategy_router_write_violation( + incoming_params=LiteLLM_Params( + model="auto_router/my-router", + auto_router_config="{}", + auto_router_default_model="gpt-4o-mini", + ), + existing_params=None, + ) + assert violation is not None + assert "auto_router_embedding_model" in violation + + @pytest.mark.asyncio + async def test_patch_model_rejects_double_prefix(self): + from litellm.proxy._types import ProxyException + from litellm.proxy.management_endpoints.model_management_endpoints import ( + patch_model, + ) + from litellm.types.router import updateLiteLLMParams + + model_id = "strategy-router-patch-test" + admin = UserAPIKeyAuth(user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN) + + with ( + patch("litellm.proxy.proxy_server.prisma_client", MagicMock()), + patch("litellm.proxy.proxy_server.llm_router", MagicMock()), + patch("litellm.proxy.proxy_server.store_model_in_db", True), + patch("litellm.proxy.proxy_server.premium_user", True), + patch( + "litellm.proxy.management_endpoints.model_management_endpoints.get_db_model", + new=AsyncMock(return_value=self._db_complexity_router(model_id)), + ), + patch( + "litellm.proxy.management_endpoints.model_management_endpoints.ModelManagementAuthChecks.can_user_make_model_call", + new=AsyncMock(return_value=None), + ), + patch( + "litellm.proxy.management_endpoints.model_management_endpoints._update_team_model_in_db", + new=AsyncMock(), + ) as mock_update, + ): + with pytest.raises(ProxyException) as exc_info: + await patch_model( + model_id=model_id, + patch_data=updateDeployment( + litellm_params=updateLiteLLMParams(model="auto_router/auto_router/complexity_router") + ), + user_api_key_dict=admin, + ) + assert "repeats" in str(exc_info.value.message) + mock_update.assert_not_awaited() + + @pytest.mark.asyncio + async def test_add_new_model_rejects_prefixed_model_without_config(self): + from litellm.proxy._types import ProxyException + from litellm.proxy.management_endpoints.model_management_endpoints import ( + add_new_model, + ) + + admin = UserAPIKeyAuth(user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN) + mock_prisma = MagicMock() + + with ( + patch("litellm.proxy.proxy_server.prisma_client", mock_prisma), + patch("litellm.proxy.proxy_server.store_model_in_db", True), + patch("litellm.proxy.proxy_server.premium_user", True), + patch( + "litellm.proxy.management_endpoints.model_management_endpoints.ModelManagementAuthChecks.can_user_make_model_call", + new=AsyncMock(return_value=None), + ), + ): + with pytest.raises(ProxyException) as exc_info: + await add_new_model( + model_params=Deployment( + model_name="my-auto-router", + litellm_params=LiteLLM_Params(model="auto_router/complexity_router"), + model_info={"id": "strategy-router-create-test"}, + ), + user_api_key_dict=admin, + ) + assert "requires" in str(exc_info.value.message) + mock_prisma.db.litellm_proxymodeltable.create.assert_not_called() + + @pytest.mark.asyncio + async def test_update_model_rejects_prefix_strip(self): + from litellm.proxy._types import ProxyException + from litellm.proxy.management_endpoints.model_management_endpoints import ( + update_model, + ) + from litellm.types.router import ModelInfo, updateLiteLLMParams + + model_id = "strategy-router-update-test" + admin = UserAPIKeyAuth(user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN) + + existing_row = MagicMock() + existing_row.model_dump.return_value = { + "model_name": "my-auto-router", + "litellm_params": { + "model": "auto_router/complexity_router", + "complexity_router_config": {"tiers": {"SIMPLE": "gpt-4o-mini"}}, + }, + "model_info": {"id": model_id}, + } + + mock_prisma = MagicMock() + mock_prisma.db.litellm_proxymodeltable.find_unique = AsyncMock(return_value=existing_row) + mock_prisma.db.litellm_proxymodeltable.update = AsyncMock() + + with ( + patch("litellm.proxy.proxy_server.prisma_client", mock_prisma), + patch("litellm.proxy.proxy_server.llm_router", MagicMock()), + patch("litellm.proxy.proxy_server.store_model_in_db", True), + patch("litellm.proxy.proxy_server.premium_user", True), + patch( + "litellm.proxy.management_endpoints.model_management_endpoints.ModelManagementAuthChecks.can_user_make_model_call", + new=AsyncMock(return_value=None), + ), + ): + with pytest.raises(ProxyException) as exc_info: + await update_model( + model_params=updateDeployment( + litellm_params=updateLiteLLMParams(model="complexity_router"), + model_info=ModelInfo(id=model_id), + ), + user_api_key_dict=admin, + ) + assert "does not start with" in str(exc_info.value.message) + mock_prisma.db.litellm_proxymodeltable.update.assert_not_awaited() diff --git a/tests/test_litellm/router_utils/test_auto_router_model_naming.py b/tests/test_litellm/router_utils/test_auto_router_model_naming.py new file mode 100644 index 00000000000..ca290caac0b --- /dev/null +++ b/tests/test_litellm/router_utils/test_auto_router_model_naming.py @@ -0,0 +1,77 @@ +import pytest + +from litellm.router_utils.auto_router_model_naming import ( + classify_strategy_router_model, + validate_strategy_router_model_write, +) + +COMPLEXITY_FIELDS = frozenset({"complexity_router_config"}) +SEMANTIC_FIELDS = frozenset( + {"auto_router_config", "auto_router_default_model", "auto_router_embedding_model"} +) + + +@pytest.mark.parametrize( + "model,expected", + [ + ("anthropic/claude-sonnet-5", None), + ("complexity_router", None), + ("autorouter/complexity_router", None), + ("auto_router/my-router", "semantic"), + ("auto_router/complexity_router", "complexity"), + ("auto_router/complexity_router-eu", "complexity"), + ("auto_router/adaptive_router", "adaptive"), + ("auto_router/quality_router", "quality"), + ("auto_router/auto_router/complexity_router", "semantic"), + ("auto_router/", "semantic"), + ], +) +def test_classify_strategy_router_model(model, expected): + assert classify_strategy_router_model(model) == expected + + +@pytest.mark.parametrize( + "model,present_fields,expected_fragment", + [ + ("auto_router/auto_router/complexity_router", COMPLEXITY_FIELDS, "repeats"), + ("complexity_router", COMPLEXITY_FIELDS, "does not start with"), + ("anthropic/claude-sonnet-5", COMPLEXITY_FIELDS, "does not start with"), + ("auto_router/", frozenset(), "missing the router name"), + ("auto_router/complexity_router", frozenset(), "requires"), + ("auto_router/my-router", frozenset({"auto_router_config"}), "requires"), + ("auto_router/adaptive_router", frozenset(), "requires"), + ("auto_router/quality_router", frozenset(), "requires"), + ], +) +def test_validate_rejects_incoherent_writes(model, present_fields, expected_fragment): + violation = validate_strategy_router_model_write(model=model, present_fields=present_fields) + assert violation is not None + assert expected_fragment in violation + + +@pytest.mark.parametrize( + "model,present_fields", + [ + ("anthropic/claude-sonnet-5", frozenset()), + ("openai/gpt-4o-mini", frozenset({"api_key"})), + ("auto_router/complexity_router", COMPLEXITY_FIELDS), + ("auto_router/complexity_router", frozenset({"complexity_router_default_model"})), + ("auto_router/complexity_router-eu", COMPLEXITY_FIELDS), + ("auto_router/my-router", SEMANTIC_FIELDS), + ( + "auto_router/my-router", + frozenset( + { + "auto_router_config_path", + "auto_router_default_model", + "auto_router_embedding_model", + } + ), + ), + ("auto_router/adaptive_router", frozenset({"adaptive_router_config"})), + ("auto_router/quality_router", frozenset({"quality_router_default_model"})), + ("auto_router/quality_router", frozenset({"quality_router_config"})), + ], +) +def test_validate_accepts_coherent_writes(model, present_fields): + assert validate_strategy_router_model_write(model=model, present_fields=present_fields) is None