diff --git a/litellm/proxy/management_endpoints/model_management_endpoints.py b/litellm/proxy/management_endpoints/model_management_endpoints.py index ca66640bf46..f04efced492 100644 --- a/litellm/proxy/management_endpoints/model_management_endpoints.py +++ b/litellm/proxy/management_endpoints/model_management_endpoints.py @@ -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: diff --git a/litellm/router.py b/litellm/router.py index 0af514fe8a2..2fc6d13c47b 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -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 diff --git a/litellm/router_utils/auto_router_model_naming.py b/litellm/router_utils/auto_router_model_naming.py index a8aa543d735..2deacc5b4bc 100644 --- a/litellm/router_utils/auto_router_model_naming.py +++ b/litellm/router_utils/auto_router_model_naming.py @@ -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``. 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 4661cc17dbc..09b598ee4f5 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 @@ -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, diff --git a/tests/test_litellm/router_strategy/test_complexity_router.py b/tests/test_litellm/router_strategy/test_complexity_router.py index 941d78085e4..863d2174415 100644 --- a/tests/test_litellm/router_strategy/test_complexity_router.py +++ b/tests/test_litellm/router_strategy/test_complexity_router.py @@ -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=[