diff --git a/litellm-proxy-extras/litellm_proxy_extras/migrations/20260902000000_one_heuristic_v2_router/migration.sql b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260902000000_one_heuristic_v2_router/migration.sql index eaae9a68e20..1d38f7ba04d 100644 --- a/litellm-proxy-extras/litellm_proxy_extras/migrations/20260902000000_one_heuristic_v2_router/migration.sql +++ b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260902000000_one_heuristic_v2_router/migration.sql @@ -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; diff --git a/litellm/proxy/management_endpoints/model_management_endpoints.py b/litellm/proxy/management_endpoints/model_management_endpoints.py index 30f3c4337af..c2c0d6ea146 100644 --- a/litellm/proxy/management_endpoints/model_management_endpoints.py +++ b/litellm/proxy/management_endpoints/model_management_endpoints.py @@ -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 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 6c8c068d8ca..5b9172a02f9 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 @@ -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,