From 79c4a53a61f9ed960d1d0ee6de9b622660e4c06a Mon Sep 17 00:00:00 2001 From: Tin Date: Wed, 2 Sep 2026 15:28:03 -0700 Subject: [PATCH] fix(router): reserve heuristic v2 for proxy admins --- .../migration.sql | 4 ++ .../model_management_endpoints.py | 54 +++++++++++++++++++ .../test_model_management_endpoints.py | 44 +++++++++++++++ 3 files changed, 102 insertions(+) create mode 100644 litellm-proxy-extras/litellm_proxy_extras/migrations/20260902000000_one_heuristic_v2_router/migration.sql 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 new file mode 100644 index 00000000000..eaae9a68e20 --- /dev/null +++ b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260902000000_one_heuristic_v2_router/migration.sql @@ -0,0 +1,4 @@ +-- 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'; diff --git a/litellm/proxy/management_endpoints/model_management_endpoints.py b/litellm/proxy/management_endpoints/model_management_endpoints.py index 85064693ecc..30f3c4337af 100644 --- a/litellm/proxy/management_endpoints/model_management_endpoints.py +++ b/litellm/proxy/management_endpoints/model_management_endpoints.py @@ -300,6 +300,7 @@ def _effective_complexity_router_config( async def _raise_if_heuristic_v2_slot_taken( *, prisma_client: PrismaClient, + user_api_key_dict: UserAPIKeyAuth, incoming_params: GenericLiteLLMParams | None, existing_params: GenericLiteLLMParams | None, current_model_id: str | None = None, @@ -308,6 +309,17 @@ async def _raise_if_heuristic_v2_slot_taken( effective_config: Final = _effective_complexity_router_config(incoming_params, existing_params) if not uses_heuristic_v2(effective_config): return + admin_violation: Final = _heuristic_v2_admin_violation( + effective_config=effective_config, + user_role=user_api_key_dict.user_role, + ) + if admin_violation is not None: + raise ProxyException( + message=admin_violation, + type=ProxyErrorTypes.auth_error.value, + code=status.HTTP_403_FORBIDDEN, + param="litellm_params.complexity_router_config.classifier_type", + ) rows: Final = await _proxy_model_table(prisma_client).find_many(where={}) violation: Final = _heuristic_v2_slot_violation( persisted_rows=rows, @@ -325,6 +337,16 @@ async def _raise_if_heuristic_v2_slot_taken( ) +def _heuristic_v2_admin_violation(*, effective_config: Mapping[str, object] | None, user_role: object) -> str | None: + """Reserve the proxy-wide v2 classifier slot for proxy administrators.""" + if not uses_heuristic_v2(effective_config) or user_role == LitellmUserRoles.PROXY_ADMIN: + return None + return ( + "Only proxy admins can configure classifier_type='heuristic_v2' because it uses a proxy-wide singleton slot. " + "Team admins can use classifier_type='heuristic' or another auto-router type." + ) + + def _heuristic_v2_slot_violation( *, persisted_rows: Sequence[_ProxyModelRow], @@ -354,6 +376,31 @@ def _heuristic_v2_slot_violation( return None +HEURISTIC_V2_SINGLETON_INDEX: Final = "LiteLLM_ProxyModelTable_one_heuristic_v2_router" + + +def _is_heuristic_v2_slot_unique_violation(error: Exception) -> bool: + """Recognize the database backstop when two proxy-admin writes race.""" + return HEURISTIC_V2_SINGLETON_INDEX in str(error) + + +def _heuristic_v2_slot_proxy_exception() -> ProxyException: + return ProxyException( + message=( + "Only one complexity router can use classifier_type='heuristic_v2' per proxy. " + "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", + ) + + +def _raise_if_heuristic_v2_slot_unique_violation(error: Exception) -> None: + if _is_heuristic_v2_slot_unique_violation(error): + raise _heuristic_v2_slot_proxy_exception() from error + + ENFORCE_RPM_TPM_ON_MODEL_ADD_SETTING: Final = "enforce_rpm_tpm_on_model_add" _REQUIRED_RATE_LIMIT_FIELDS: Final = ("rpm", "tpm") @@ -805,6 +852,7 @@ async def patch_model( ) await _raise_if_heuristic_v2_slot_taken( prisma_client=prisma_client, + user_api_key_dict=user_api_key_dict, incoming_params=patch_data.litellm_params, existing_params=db_model.litellm_params, current_model_id=model_id, @@ -867,6 +915,7 @@ async def patch_model( except Exception as e: verbose_proxy_logger.exception("Error in patch_model: %s", e) + _raise_if_heuristic_v2_slot_unique_violation(e) if isinstance(e, (HTTPException, ProxyException)): raise e @@ -1927,6 +1976,7 @@ async def add_new_model( ) await _raise_if_heuristic_v2_slot_taken( prisma_client=prisma_client, + user_api_key_dict=user_api_key_dict, incoming_params=model_params.litellm_params, existing_params=None, ) @@ -1979,6 +2029,7 @@ async def add_new_model( ) except Exception as e: verbose_proxy_logger.exception("Exception in add_new_model: %s", e) + _raise_if_heuristic_v2_slot_unique_violation(e) else: raise HTTPException( @@ -2020,6 +2071,7 @@ async def add_new_model( except Exception as e: verbose_proxy_logger.exception("litellm.proxy.proxy_server.add_new_model(): Exception occured - %s", e) + _raise_if_heuristic_v2_slot_unique_violation(e) if isinstance(e, HTTPException): raise ProxyException( message=getattr(e, "detail", f"Authentication Error({e})"), @@ -2110,6 +2162,7 @@ async def update_model( ) await _raise_if_heuristic_v2_slot_taken( prisma_client=prisma_client, + user_api_key_dict=user_api_key_dict, incoming_params=model_params.litellm_params, existing_params=deployment.litellm_params, current_model_id=_model_id, @@ -2189,6 +2242,7 @@ async def update_model( return model_response except Exception as e: verbose_proxy_logger.exception("litellm.proxy.proxy_server.update_model(): Exception occured - %s", e) + _raise_if_heuristic_v2_slot_unique_violation(e) if isinstance(e, HTTPException): raise ProxyException( message=getattr(e, "detail", f"Authentication Error({e})"), 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 22bc5b09843..6c8c068d8ca 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 @@ -4011,6 +4011,50 @@ class TestStrategyRouterWriteValidation: is None ) + def test_heuristic_v2_is_reserved_for_proxy_admins(self): + from litellm.proxy.management_endpoints.model_management_endpoints import ( + _heuristic_v2_admin_violation, + ) + + config = {"classifier_type": "heuristic_v2"} + violation = _heuristic_v2_admin_violation( + effective_config=config, + user_role=LitellmUserRoles.INTERNAL_USER, + ) + assert violation is not None + assert "Only proxy admins" in violation + assert ( + _heuristic_v2_admin_violation( + effective_config=config, + user_role=LitellmUserRoles.PROXY_ADMIN, + ) + is None + ) + + def test_team_admins_can_still_configure_heuristic_v1(self): + from litellm.proxy.management_endpoints.model_management_endpoints import ( + _heuristic_v2_admin_violation, + ) + + assert ( + _heuristic_v2_admin_violation( + effective_config={"classifier_type": "heuristic"}, + user_role=LitellmUserRoles.INTERNAL_USER, + ) + is None + ) + + def test_database_singleton_violation_is_recognized(self): + from litellm.proxy.management_endpoints.model_management_endpoints import ( + HEURISTIC_V2_SINGLETON_INDEX, + _is_heuristic_v2_slot_unique_violation, + ) + + assert _is_heuristic_v2_slot_unique_violation( + Exception(f'duplicate key violates unique constraint "{HEURISTIC_V2_SINGLETON_INDEX}"') + ) + assert not _is_heuristic_v2_slot_unique_violation(Exception("unrelated database error")) + def test_double_prefix_rejected_against_stored_params(self): from litellm.proxy.management_endpoints.model_management_endpoints import ( _strategy_router_write_violation,