feat(router): limit heuristic v2 to one router

This commit is contained in:
Tin 2026-09-02 14:04:36 -07:00
parent 993766be0e
commit 4bc01c3b47
5 changed files with 225 additions and 0 deletions

View file

@ -92,6 +92,7 @@ from litellm.router_strategy.complexity_router import (
from litellm.router_utils.auto_router_model_naming import (
STRATEGY_ROUTER_PARAM_FIELDS,
carries_complexity_router_settings,
uses_heuristic_v2,
validate_complexity_router_config_placement,
validate_complexity_router_config_write,
validate_strategy_router_model_write,
@ -143,6 +144,9 @@ class _ProxyModelRow(Protocol):
def model_dump_json(self, *, exclude_none: bool = False) -> str: ...
@property
def litellm_params(self) -> object: ...
class _ProxyModelTable(Protocol):
def find_unique(self, *, where: Mapping[str, object]) -> Awaitable[BaseModel | None]: ...
@ -265,6 +269,68 @@ def _raise_on_strategy_router_write_violation(
)
def _litellm_params_mapping(value: object) -> Mapping[str, object]:
"""Normalize Prisma's parsed or JSON-encoded Json column at the typed boundary."""
if isinstance(value, Mapping):
return value
if isinstance(value, str):
try:
parsed: Final = json.loads(value)
except JSONDecodeError:
return MappingProxyType({})
return parsed if isinstance(parsed, Mapping) else MappingProxyType({})
return MappingProxyType({})
def _effective_complexity_router_config(
incoming_params: GenericLiteLLMParams | None,
existing_params: GenericLiteLLMParams | None,
) -> Mapping[str, object] | None:
"""Return the whole config that the write leaves on the deployment."""
config: Final = (
incoming_params.complexity_router_config
if incoming_params is not None and incoming_params.complexity_router_config is not None
else existing_params.complexity_router_config
if existing_params is not None
else None
)
return config if isinstance(config, Mapping) else None
async def _raise_if_heuristic_v2_slot_taken(
*,
prisma_client: PrismaClient,
incoming_params: GenericLiteLLMParams | None,
existing_params: GenericLiteLLMParams | None,
current_model_id: str | None = None,
) -> None:
"""Allow at most one persisted heuristic-v2 complexity router per proxy."""
effective_config: Final = _effective_complexity_router_config(incoming_params, existing_params)
if not uses_heuristic_v2(effective_config):
return
rows: Final = await _proxy_model_table(prisma_client).find_many(where={})
for row in rows:
if current_model_id is not None and row.model_id == current_model_id:
continue
if uses_heuristic_v2(
config
if isinstance(
config := _litellm_params_mapping(row.litellm_params).get("complexity_router_config"), Mapping
)
else None
):
raise ProxyException(
message=(
"Only one complexity router can use classifier_type='heuristic_v2' per proxy. "
"Use classifier_type='heuristic' for this router, or change or delete the existing "
"heuristic_v2 router first."
),
type=ProxyErrorTypes.validation_error.value,
code=status.HTTP_400_BAD_REQUEST,
param="litellm_params.complexity_router_config.classifier_type",
)
ENFORCE_RPM_TPM_ON_MODEL_ADD_SETTING: Final = "enforce_rpm_tpm_on_model_add"
_REQUIRED_RATE_LIMIT_FIELDS: Final = ("rpm", "tpm")
@ -714,6 +780,12 @@ async def patch_model(
incoming_params=patch_data.litellm_params,
existing_params=db_model.litellm_params,
)
await _raise_if_heuristic_v2_slot_taken(
prisma_client=prisma_client,
incoming_params=patch_data.litellm_params,
existing_params=db_model.litellm_params,
current_model_id=model_id,
)
# Handle team model updates with proper alias management
update_data: Final = await _update_team_model_in_db(
@ -1830,6 +1902,11 @@ async def add_new_model(
incoming_params=model_params.litellm_params,
existing_params=None,
)
await _raise_if_heuristic_v2_slot_taken(
prisma_client=prisma_client,
incoming_params=model_params.litellm_params,
existing_params=None,
)
_raise_if_rate_limits_required_but_missing(
litellm_params=model_params.litellm_params,
@ -2008,6 +2085,12 @@ async def update_model(
incoming_params=model_params.litellm_params,
existing_params=deployment.litellm_params,
)
await _raise_if_heuristic_v2_slot_taken(
prisma_client=prisma_client,
incoming_params=model_params.litellm_params,
existing_params=deployment.litellm_params,
current_model_id=_model_id,
)
# update DB
if store_model_in_db is True:

View file

@ -8702,6 +8702,15 @@ class Router:
complexity_router_config: Final[dict | None] = deployment.litellm_params.complexity_router_config
if complexity_router_config and complexity_router_config.get("classifier_type") == "heuristic_v2":
for registered in self.complexity_routers.values():
if any(tagged.strategy.config.classifier_type == "heuristic_v2" for tagged in registered):
raise ValueError(
"Only one complexity router can use classifier_type='heuristic_v2' per proxy. "
"Use classifier_type='heuristic' for this router, or change or remove the existing "
"heuristic_v2 router first."
)
default_model: str | None = deployment.litellm_params.complexity_router_default_model
# If no default model specified, try to get from config tiers. Derived from the

View file

@ -210,6 +210,11 @@ def carries_complexity_router_settings(model: str | None, present_fields: frozen
)
def uses_heuristic_v2(complexity_router_config: Mapping[str, object] | None) -> bool:
"""Whether a complexity-router config selects the singleton v2 classifier."""
return complexity_router_config is not None and complexity_router_config.get("classifier_type") == "heuristic_v2"
def validate_complexity_router_config_placement(litellm_params: Mapping[str, object] | None) -> str | None:
"""Reject a complexity-router setting written beside ``complexity_router_config``.

View file

@ -3953,6 +3953,80 @@ class TestStrategyRouterWriteValidation:
model_info={"id": model_id},
)
@pytest.mark.asyncio
async def test_second_heuristic_v2_router_is_rejected(self):
from litellm.proxy._types import ProxyException
from litellm.proxy.management_endpoints.model_management_endpoints import (
_raise_if_heuristic_v2_slot_taken,
)
table = MagicMock()
row = MagicMock()
row.model_id = "first-v2"
row.litellm_params = {"complexity_router_config": {"classifier_type": "heuristic_v2"}}
table.find_many = AsyncMock(return_value=[row])
with patch(
"litellm.proxy.management_endpoints.model_management_endpoints._proxy_model_table",
return_value=table,
):
with pytest.raises(ProxyException, match="Only one complexity router"):
await _raise_if_heuristic_v2_slot_taken(
prisma_client=MagicMock(),
incoming_params=LiteLLM_Params(
model="auto_router/complexity_router",
complexity_router_config={"classifier_type": "heuristic_v2"},
),
existing_params=None,
)
@pytest.mark.asyncio
async def test_existing_heuristic_v2_router_can_be_edited(self):
from litellm.proxy.management_endpoints.model_management_endpoints import (
_raise_if_heuristic_v2_slot_taken,
)
from litellm.types.router import updateLiteLLMParams
table = MagicMock()
row = MagicMock()
row.model_id = "first-v2"
row.litellm_params = json.dumps({"complexity_router_config": {"classifier_type": "heuristic_v2"}})
table.find_many = AsyncMock(return_value=[row])
with patch(
"litellm.proxy.management_endpoints.model_management_endpoints._proxy_model_table",
return_value=table,
):
await _raise_if_heuristic_v2_slot_taken(
prisma_client=MagicMock(),
incoming_params=updateLiteLLMParams(rpm=10),
existing_params=LiteLLM_Params(
model="auto_router/complexity_router",
complexity_router_config={"classifier_type": "heuristic_v2"},
),
current_model_id="first-v2",
)
@pytest.mark.asyncio
async def test_heuristic_v1_does_not_consume_v2_slot(self):
from litellm.proxy.management_endpoints.model_management_endpoints import (
_raise_if_heuristic_v2_slot_taken,
)
table = MagicMock()
table.find_many = AsyncMock()
with patch(
"litellm.proxy.management_endpoints.model_management_endpoints._proxy_model_table",
return_value=table,
):
await _raise_if_heuristic_v2_slot_taken(
prisma_client=MagicMock(),
incoming_params=LiteLLM_Params(
model="auto_router/complexity_router",
complexity_router_config={"classifier_type": "heuristic"},
),
existing_params=None,
)
table.find_many.assert_not_awaited()
def test_double_prefix_rejected_against_stored_params(self):
from litellm.proxy.management_endpoints.model_management_endpoints import (
_strategy_router_write_violation,

View file

@ -1085,6 +1085,60 @@ class TestRouterComplexityDeploymentMethods:
router.init_complexity_router_deployment(deployment)
assert "auto_router/complexity_router/test-router" in router.complexity_routers
def test_only_one_heuristic_v2_complexity_router_can_register(self):
router = Router(
model_list=[
{
"model_name": "gpt-4o-mini",
"litellm_params": {"model": "openai/gpt-4o-mini"},
}
]
)
def deployment(name: str) -> Deployment:
return Deployment(
model_name=name,
litellm_params=LiteLLM_Params(
model=f"auto_router/complexity_router/{name}",
complexity_router_default_model="gpt-4o-mini",
complexity_router_config={
"classifier_type": "heuristic_v2",
"tiers": {"SIMPLE": "gpt-4o-mini"},
},
),
model_info={"id": name},
)
router.init_complexity_router_deployment(deployment("first"))
with pytest.raises(ValueError, match="Only one complexity router"):
router.init_complexity_router_deployment(deployment("second"))
def test_heuristic_v1_complexity_routers_remain_unlimited(self):
router = Router(
model_list=[
{
"model_name": "gpt-4o-mini",
"litellm_params": {"model": "openai/gpt-4o-mini"},
}
]
)
for name in ("first", "second"):
router.init_complexity_router_deployment(
Deployment(
model_name=name,
litellm_params=LiteLLM_Params(
model=f"auto_router/complexity_router/{name}",
complexity_router_default_model="gpt-4o-mini",
complexity_router_config={
"classifier_type": "heuristic",
"tiers": {"SIMPLE": "gpt-4o-mini"},
},
),
model_info={"id": name},
)
)
assert set(router.complexity_routers) == {"first", "second"}
def test_hybrid_initialization_waits_for_later_pool_deployments(self):
router = Router(
model_list=[