mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-04 02:31:27 +00:00
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:
parent
8c30093c5a
commit
02f386a82e
2 changed files with 145 additions and 16 deletions
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue