Merge pull request #34151 from BerriAI/litellm_lit4663_autorouter_prefix

fix(proxy): reject model writes that corrupt an auto-router pseudo-model
This commit is contained in:
tin-berri 2026-07-28 21:01:42 -07:00 • committed by GitHub
commit 9b48bf6084
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
5 changed files with 457 additions and 16 deletions

View file

@ -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)

View file

@ -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):
"""

View file

@ -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

View file

@ -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()

View file

@ -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