diff --git a/litellm/proxy/management_endpoints/model_management_endpoints.py b/litellm/proxy/management_endpoints/model_management_endpoints.py index f04efced492..85064693ecc 100644 --- a/litellm/proxy/management_endpoints/model_management_endpoints.py +++ b/litellm/proxy/management_endpoints/model_management_endpoints.py @@ -309,7 +309,34 @@ async def _raise_if_heuristic_v2_slot_taken( if not uses_heuristic_v2(effective_config): return rows: Final = await _proxy_model_table(prisma_client).find_many(where={}) - for row in rows: + violation: Final = _heuristic_v2_slot_violation( + persisted_rows=rows, + incoming_params=incoming_params, + existing_params=existing_params, + current_model_id=current_model_id, + ) + if violation is None: + return + raise ProxyException( + message=violation, + type=ProxyErrorTypes.validation_error.value, + code=status.HTTP_400_BAD_REQUEST, + param="litellm_params.complexity_router_config.classifier_type", + ) + + +def _heuristic_v2_slot_violation( + *, + persisted_rows: Sequence[_ProxyModelRow], + incoming_params: GenericLiteLLMParams | None, + existing_params: GenericLiteLLMParams | None, + current_model_id: str | None = None, +) -> str | None: + """Return the singleton-slot violation for the effective post-write config.""" + effective_config: Final = _effective_complexity_router_config(incoming_params, existing_params) + if not uses_heuristic_v2(effective_config): + return None + for row in persisted_rows: if current_model_id is not None and row.model_id == current_model_id: continue if uses_heuristic_v2( @@ -319,16 +346,12 @@ async def _raise_if_heuristic_v2_slot_taken( ) 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", + return ( + "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." ) + return None ENFORCE_RPM_TPM_ON_MODEL_ADD_SETTING: Final = "enforce_rpm_tpm_on_model_add" 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 09b598ee4f5..22bc5b09843 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,50 +3953,37 @@ 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 + def test_second_heuristic_v2_router_is_rejected(self): from litellm.proxy.management_endpoints.model_management_endpoints import ( - _raise_if_heuristic_v2_slot_taken, + _heuristic_v2_slot_violation, ) - 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, - ) + violation = _heuristic_v2_slot_violation( + persisted_rows=[row], + 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 "Only one complexity router" in violation - @pytest.mark.asyncio - async def test_existing_heuristic_v2_router_can_be_edited(self): + 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, + _heuristic_v2_slot_violation, ) 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(), + assert ( + _heuristic_v2_slot_violation( + persisted_rows=[row], incoming_params=updateLiteLLMParams(rpm=10), existing_params=LiteLLM_Params( model="auto_router/complexity_router", @@ -4004,28 +3991,25 @@ class TestStrategyRouterWriteValidation: ), 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, + is None ) - 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(), + def test_heuristic_v1_does_not_consume_v2_slot(self): + from litellm.proxy.management_endpoints.model_management_endpoints import ( + _heuristic_v2_slot_violation, + ) + + assert ( + _heuristic_v2_slot_violation( + persisted_rows=(), incoming_params=LiteLLM_Params( model="auto_router/complexity_router", complexity_router_config={"classifier_type": "heuristic"}, ), existing_params=None, ) - table.find_many.assert_not_awaited() + is None + ) def test_double_prefix_rejected_against_stored_params(self): from litellm.proxy.management_endpoints.model_management_endpoints import (