fix(proxy): enforce virtual key budgets for JEV test routing

Backport of #41879 to stable/1.101.x.
Cherry-picked from 1e161f516c (main).
This commit is contained in:
moe-berri 2026-09-22 23:15:11 +00:00 • committed by Devin AI
parent 8c30093c5a
commit 02f386a82e
2 changed files with 145 additions and 16 deletions

View file

@ -72,6 +72,11 @@ class AutoRouterRoutingTestRequest(BaseModel):
complexity_router_config: RequestComplexityRouterConfig = Field(
description="The complexity router config to route against, in the shape /model/new accepts",
)
saved_model_id: str | None = Field(
default=None,
min_length=1,
description="Test this saved deployment's server-side configuration instead of the supplied config and default model",
)
default_model: str | None = Field(
default=None,
description="Model to route to when no tier resolves, i.e. complexity_router_default_model",

View file

@ -3,30 +3,46 @@ Unit tests for auto router management endpoints
"""
from collections.abc import Mapping, Sequence
from functools import partial
from pathlib import Path
from typing import Final
from unittest.mock import AsyncMock, MagicMock
import pytest
from fastapi import HTTPException
from fastapi import HTTPException, Request
from pydantic import ValidationError
import litellm
from litellm.proxy import proxy_server
from litellm.proxy._types import (
LitellmUserRoles,
ProxyErrorTypes,
ProxyException,
UserAPIKeyAuth,
)
from litellm.proxy.management_endpoints import auto_router_endpoints
from litellm.proxy.management_endpoints.auto_router_endpoints import (
preview_auto_router_routing,
)
from litellm.router import Router
from litellm.types.utils import Choices, Message, ModelResponse
from litellm.router_strategy.complexity_router import ComplexityRouter
from litellm.router_strategy.complexity_router.jev_classifier import (
JevChoiceAnswer,
JevClassifierClient,
JevSystemOneResponse,
)
from litellm.types.management_endpoints.auto_router_endpoints import (
AutoRouterBenchmarksResponse,
AutoRouterRoutingTestRequest,
)
ROUTING_HTTP_REQUEST: Final = Request(
{"type": "http", "method": "POST", "path": "/auto_router/test_routing", "headers": []}
)
ADMIN = UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN, api_key="sk-test", user_id="admin")
@ -94,7 +110,7 @@ async def _route_body(body: Mapping[str, object], monkeypatch: pytest.MonkeyPatc
import litellm.proxy.proxy_server as proxy_server
monkeypatch.setattr(proxy_server, "llm_router", _router())
return await preview_auto_router_routing(
return await preview_auto_router_routing(http_request=ROUTING_HTTP_REQUEST,
data=_request_from(body, **config_overrides),
user_api_key_dict=ADMIN,
)
@ -121,7 +137,7 @@ async def _classifier_user_payload(body: Mapping[str, object], monkeypatch: pyte
router = RecordingRouter("SIMPLE")
monkeypatch.setattr(proxy_server, "llm_router", router)
await preview_auto_router_routing(
await preview_auto_router_routing(http_request=ROUTING_HTTP_REQUEST,
data=_request_from(body, classifier_type="llm", classifier_llm_config={"model": "classifier-model"}),
user_api_key_dict=ADMIN,
)
@ -198,7 +214,7 @@ async def test_llm_classifier_call_is_billed_to_the_calling_key(monkeypatch: pyt
monkeypatch.setattr(router, "acompletion", fake_acompletion)
monkeypatch.setattr(proxy_server, "llm_router", router)
response = await preview_auto_router_routing(
response = await preview_auto_router_routing(http_request=ROUTING_HTTP_REQUEST,
data=_request(
"what is 2+2",
classifier_type="llm",
@ -359,7 +375,7 @@ async def test_a_key_that_cannot_call_the_classifier_model_is_rejected_before_it
monkeypatch.setattr(proxy_server, "llm_router", router)
with pytest.raises(ProxyException) as exc_info:
await preview_auto_router_routing(
await preview_auto_router_routing(http_request=ROUTING_HTTP_REQUEST,
data=_request("what is 2+2", **config_overrides),
user_api_key_dict=UserAPIKeyAuth(
user_role=LitellmUserRoles.PROXY_ADMIN,
@ -388,7 +404,7 @@ async def test_a_key_over_its_budget_cannot_run_a_classifier_config(monkeypatch:
monkeypatch.setattr(proxy_server, "llm_router", router)
with pytest.raises(ProxyException) as exc_info:
await preview_auto_router_routing(
await preview_auto_router_routing(http_request=ROUTING_HTTP_REQUEST,
data=_request(
"what is 2+2",
classifier_type="llm",
@ -407,20 +423,127 @@ async def test_a_key_over_its_budget_cannot_run_a_classifier_config(monkeypatch:
assert calls == []
@pytest.mark.parametrize(
"max_budget, spend, denied",
(
pytest.param(0.0, 0.0, True, id="zero-budget"),
pytest.param(1.0, 1.0, True, id="budget-reached"),
pytest.param(1.0, 2.0, True, id="budget-exceeded"),
pytest.param(1.0, 0.5, False, id="budget-remaining"),
pytest.param(None, 2.0, False, id="unlimited"),
),
)
@pytest.mark.asyncio
async def test_a_heuristic_config_does_not_need_a_budget(monkeypatch: pytest.MonkeyPatch):
async def test_jev_test_routing_enforces_key_budget_before_provider_invocation(
monkeypatch: pytest.MonkeyPatch, max_budget: float | None, spend: float, denied: bool
) -> None:
client: Final = AsyncMock(spec=JevClassifierClient)
client.evaluate.return_value = JevSystemOneResponse(
model="jev-test",
answers={
"tier": JevChoiceAnswer(type="choice", choice="SIMPLE", probabilities={"SIMPLE": 1.0}, confidence=1.0)
},
)
monkeypatch.setattr(proxy_server, "llm_router", _router())
monkeypatch.setattr(auto_router_endpoints, "ComplexityRouter", partial(ComplexityRouter, jev_client=client))
actor: Final = UserAPIKeyAuth(
user_role=LitellmUserRoles.PROXY_ADMIN,
api_key="sk-jev-budget-test",
user_id="admin",
models=["cheap-model", "typesafe/jev-test"],
max_budget=max_budget,
spend=spend,
)
request: Final = _request(
"what is 2+2",
classifier_type="jev",
jev_classifier_config={"model": "jev-test"},
)
if denied:
with pytest.raises(ProxyException) as exc_info:
await preview_auto_router_routing(http_request=ROUTING_HTTP_REQUEST, data=request, user_api_key_dict=actor)
assert exc_info.value.type == ProxyErrorTypes.budget_exceeded
assert exc_info.value.code == "400"
assert exc_info.value.param is None
assert "Budget has been exceeded!" in exc_info.value.message
client.evaluate.assert_not_called()
return
response: Final = await preview_auto_router_routing(http_request=ROUTING_HTTP_REQUEST,
data=request, user_api_key_dict=actor
)
assert response.routed_model == "cheap-model"
assert response.routing_decision["cause"] == "jev_classifier"
assert response.routing_decision["classifier_model"] == "typesafe/jev-test"
client.evaluate.assert_awaited_once()
@pytest.mark.parametrize(
"max_budget, spend, denied",
((0.0, 0.0, True), (1.0, 2.0, True), (1.0, 0.5, False), (None, 2.0, False)),
)
@pytest.mark.asyncio
async def test_jev_test_routing_hard_blocks_exhausted_throttle_enabled_keys(
monkeypatch: pytest.MonkeyPatch, max_budget: float | None, spend: float, denied: bool
) -> None:
client: Final = AsyncMock(spec=JevClassifierClient)
client.evaluate.return_value = JevSystemOneResponse(
model="jev-test",
answers={
"tier": JevChoiceAnswer(type="choice", choice="SIMPLE", probabilities={"SIMPLE": 1.0}, confidence=1.0)
},
)
monkeypatch.setattr(litellm, "budget_exceeded_throttle_percentage", 0.1)
monkeypatch.setattr(proxy_server, "llm_router", _router())
monkeypatch.setattr(auto_router_endpoints, "ComplexityRouter", partial(ComplexityRouter, jev_client=client))
actor: Final = UserAPIKeyAuth(
user_role=LitellmUserRoles.PROXY_ADMIN,
api_key="sk-jev-throttle-test",
user_id="admin",
models=["cheap-model", "typesafe/jev-test"],
max_budget=max_budget,
spend=spend,
rpm_limit=100,
metadata={"throttle_on_budget_exceeded": True},
)
request: Final = _request(
"what is 2+2",
classifier_type="jev",
jev_classifier_config={"model": "jev-test"},
)
if denied:
with pytest.raises(ProxyException) as exc_info:
await preview_auto_router_routing(http_request=ROUTING_HTTP_REQUEST, data=request, user_api_key_dict=actor)
assert exc_info.value.type == ProxyErrorTypes.budget_exceeded
assert exc_info.value.code == "400"
client.evaluate.assert_not_called()
return
response: Final = await preview_auto_router_routing(http_request=ROUTING_HTTP_REQUEST,
data=request, user_api_key_dict=actor
)
assert response.routing_decision["cause"] == "jev_classifier"
client.evaluate.assert_awaited_once()
@pytest.mark.parametrize("max_budget, spend", ((0.0, 0.0), (1.0, 2.0)))
@pytest.mark.asyncio
async def test_a_heuristic_config_does_not_need_a_budget(
monkeypatch: pytest.MonkeyPatch, max_budget: float, spend: float
):
import litellm.proxy.proxy_server as proxy_server
monkeypatch.setattr(proxy_server, "llm_router", _router())
response = await preview_auto_router_routing(
response = await preview_auto_router_routing(http_request=ROUTING_HTTP_REQUEST,
data=_request("what is 2+2"),
user_api_key_dict=UserAPIKeyAuth(
user_role=LitellmUserRoles.PROXY_ADMIN,
api_key="sk-broke",
user_id="admin",
max_budget=1.0,
spend=2.0,
max_budget=max_budget,
spend=spend,
models=["cheap-model"],
),
)
@ -435,7 +558,7 @@ async def test_no_llm_router_on_the_proxy_is_a_500(monkeypatch: pytest.MonkeyPat
monkeypatch.setattr(proxy_server, "llm_router", None)
with pytest.raises(HTTPException) as exc_info:
await preview_auto_router_routing(data=_request("what is 2+2"), user_api_key_dict=ADMIN)
await preview_auto_router_routing(http_request=ROUTING_HTTP_REQUEST, data=_request("what is 2+2"), user_api_key_dict=ADMIN)
assert exc_info.value.status_code == 500
@ -447,7 +570,7 @@ async def test_non_admin_without_a_team_is_rejected(monkeypatch: pytest.MonkeyPa
monkeypatch.setattr(proxy_server, "llm_router", _router())
with pytest.raises(HTTPException) as exc_info:
await preview_auto_router_routing(
await preview_auto_router_routing(http_request=ROUTING_HTTP_REQUEST,
data=_request("what is 2+2"),
user_api_key_dict=UserAPIKeyAuth(
user_role=LitellmUserRoles.INTERNAL_USER, api_key="sk-user", user_id="user"
@ -794,7 +917,6 @@ class TestAutoRouterBenchmarks:
# ---------------------------------------------------------------------------
from datetime import datetime, timedelta, timezone
from unittest.mock import AsyncMock, MagicMock
from litellm.proxy.management_endpoints.auto_router_endpoints import (
@ -1131,7 +1253,9 @@ async def test_start_shadow_eval_writes_one_leg_per_key_in_one_statement(monkeyp
assert (
len(
{
frozenset((k, tuple(v) if isinstance(v, list) else v) for k, v in row.items() if k not in ("target_id", "id"))
frozenset(
(k, tuple(v) if isinstance(v, list) else v) for k, v in row.items() if k not in ("target_id", "id")
)
for row in rows
}
)
@ -2534,12 +2658,12 @@ async def test_routing_test_never_confirms_models_the_caller_cannot_use(monkeypa
)
monkeypatch.setattr(proxy_server, "prisma_client", _team_prisma("team-probe", models=["mid-model"]))
probing = await preview_auto_router_routing(data=_request("team-probe"), user_api_key_dict=team_admin)
probing = await preview_auto_router_routing(http_request=ROUTING_HTTP_REQUEST, data=_request("team-probe"), user_api_key_dict=team_admin)
assert probing.routed_model == "cheap-model"
assert probing.routed_model_configured is False
monkeypatch.setattr(proxy_server, "prisma_client", _team_prisma("team-grant", models=["cheap-model"]))
granted = await preview_auto_router_routing(data=_request("team-grant"), user_api_key_dict=team_admin)
granted = await preview_auto_router_routing(http_request=ROUTING_HTTP_REQUEST, data=_request("team-grant"), user_api_key_dict=team_admin)
assert granted.routed_model == "cheap-model"
assert granted.routed_model_configured is True