mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
test(router): satisfy heuristic v2 quality gate
This commit is contained in:
parent
4bc01c3b47
commit
a4e00349a7
2 changed files with 61 additions and 54 deletions
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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 (
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue