fix(router): count configured heuristic v2 routers

This commit is contained in:
Tin 2026-09-02 15:47:25 -07:00
parent 79c4a53a61
commit 0f85c3e640
3 changed files with 81 additions and 3 deletions

View file

@ -1,4 +1,10 @@
-- Atomically reserve the proxy-wide heuristic_v2 classifier slot
CREATE UNIQUE INDEX IF NOT EXISTS "LiteLLM_ProxyModelTable_one_heuristic_v2_router"
ON "LiteLLM_ProxyModelTable" ((litellm_params #>> '{complexity_router_config,classifier_type}'))
WHERE (litellm_params #>> '{complexity_router_config,classifier_type}') = 'heuristic_v2';
ON "LiteLLM_ProxyModelTable" ((1))
WHERE CASE
WHEN jsonb_typeof(litellm_params) = 'object'
THEN (litellm_params #>> '{complexity_router_config,classifier_type}') = 'heuristic_v2'
WHEN jsonb_typeof(litellm_params) = 'string'
THEN (((litellm_params #>> '{}')::jsonb) #>> '{complexity_router_config,classifier_type}') = 'heuristic_v2'
ELSE FALSE
END;

View file

@ -305,7 +305,7 @@ async def _raise_if_heuristic_v2_slot_taken(
existing_params: GenericLiteLLMParams | None,
current_model_id: str | None = None,
) -> None:
"""Allow at most one persisted heuristic-v2 complexity router per proxy."""
"""Allow at most one 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
@ -320,9 +320,12 @@ async def _raise_if_heuristic_v2_slot_taken(
code=status.HTTP_403_FORBIDDEN,
param="litellm_params.complexity_router_config.classifier_type",
)
from litellm.proxy.proxy_server import llm_router
rows: Final = await _proxy_model_table(prisma_client).find_many(where={})
violation: Final = _heuristic_v2_slot_violation(
persisted_rows=rows,
live_model_list=llm_router.model_list if llm_router is not None else (),
incoming_params=incoming_params,
existing_params=existing_params,
current_model_id=current_model_id,
@ -350,6 +353,7 @@ def _heuristic_v2_admin_violation(*, effective_config: Mapping[str, object] | No
def _heuristic_v2_slot_violation(
*,
persisted_rows: Sequence[_ProxyModelRow],
live_model_list: Sequence[Mapping[str, object]] = (),
incoming_params: GenericLiteLLMParams | None,
existing_params: GenericLiteLLMParams | None,
current_model_id: str | None = None,
@ -358,6 +362,18 @@ def _heuristic_v2_slot_violation(
effective_config: Final = _effective_complexity_router_config(incoming_params, existing_params)
if not uses_heuristic_v2(effective_config):
return None
for deployment in live_model_list:
model_info = deployment.get("model_info")
if isinstance(model_info, Mapping) and model_info.get("db_model") is True:
continue
litellm_params = _litellm_params_mapping(deployment.get("litellm_params"))
live_config = litellm_params.get("complexity_router_config")
if uses_heuristic_v2(live_config if isinstance(live_config, Mapping) else None):
return (
"Only one complexity router can use classifier_type='heuristic_v2' per proxy. "
"A config.yaml deployment already uses the slot; change or remove it before "
"creating a database-backed heuristic_v2 router."
)
for row in persisted_rows:
if current_model_id is not None and row.model_id == current_model_id:
continue

View file

@ -3994,6 +3994,62 @@ class TestStrategyRouterWriteValidation:
is None
)
@pytest.mark.parametrize(
"litellm_params",
[
{"complexity_router_config": {"classifier_type": "heuristic_v2"}},
json.dumps({"complexity_router_config": {"classifier_type": "heuristic_v2"}}),
],
)
def test_config_yaml_heuristic_v2_router_takes_slot(self, litellm_params):
from litellm.proxy.management_endpoints.model_management_endpoints import (
_heuristic_v2_slot_violation,
)
violation = _heuristic_v2_slot_violation(
persisted_rows=(),
live_model_list=[
{
"model_name": "configured-router",
"litellm_params": litellm_params,
"model_info": {"db_model": False},
}
],
incoming_params=LiteLLM_Params(
model="auto_router/complexity_router",
complexity_router_config={"classifier_type": "heuristic_v2"},
),
existing_params=None,
)
assert violation is not None
assert "config.yaml" in violation
def test_db_models_in_live_router_do_not_double_consume_slot(self):
from litellm.proxy.management_endpoints.model_management_endpoints import (
_heuristic_v2_slot_violation,
)
assert (
_heuristic_v2_slot_violation(
persisted_rows=(),
live_model_list=[
{
"model_name": "db-router",
"litellm_params": {
"complexity_router_config": {"classifier_type": "heuristic_v2"}
},
"model_info": {"db_model": True},
}
],
incoming_params=LiteLLM_Params(
model="auto_router/complexity_router",
complexity_router_config={"classifier_type": "heuristic_v2"},
),
existing_params=None,
)
is None
)
def test_heuristic_v1_does_not_consume_v2_slot(self):
from litellm.proxy.management_endpoints.model_management_endpoints import (
_heuristic_v2_slot_violation,