test(router): satisfy heuristic v2 quality gate

This commit is contained in:
Tin 2026-09-02 15:20:02 -07:00
parent 4bc01c3b47
commit a4e00349a7
2 changed files with 61 additions and 54 deletions

View file

@ -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"

View file

@ -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 (