mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-21 00:21:49 +00:00
fix(proxy): enforce customer model_max_budget on auth paths
Apply end-user model_max_budget from customer budgets on virtual-key and master-key auth paths, and extract a shared enforcement helper. Fixes #31842 Co-authored-by: Cursor <cursoragent@cursor.com>
This commit is contained in:
parent
4bd579cf6f
commit
4bd96852d9
2 changed files with 116 additions and 32 deletions
|
|
@ -1607,6 +1607,14 @@ async def _user_api_key_auth_builder(
|
|||
valid_token=_user_api_key_obj, end_user_params=end_user_params
|
||||
)
|
||||
|
||||
if RouteChecks.is_llm_api_route(route=route):
|
||||
await _enforce_end_user_model_max_budget_checks(
|
||||
valid_token=_user_api_key_obj,
|
||||
request_data=request_data,
|
||||
route=route,
|
||||
request=request,
|
||||
)
|
||||
|
||||
return _user_api_key_obj
|
||||
|
||||
## IF it's not a master key
|
||||
|
|
@ -1666,10 +1674,9 @@ async def _user_api_key_auth_builder(
|
|||
raise e
|
||||
# update end-user params on valid token
|
||||
# These can change per request - it's important to update them here
|
||||
valid_token.end_user_id = end_user_params.get("end_user_id")
|
||||
valid_token.end_user_tpm_limit = end_user_params.get("end_user_tpm_limit")
|
||||
valid_token.end_user_rpm_limit = end_user_params.get("end_user_rpm_limit")
|
||||
valid_token.allowed_model_region = end_user_params.get("allowed_model_region")
|
||||
valid_token = update_valid_token_with_end_user_params(
|
||||
valid_token=valid_token, end_user_params=end_user_params
|
||||
)
|
||||
# update key budget with temp budget increase
|
||||
valid_token = _update_key_budget_with_temp_budget_increase(
|
||||
valid_token
|
||||
|
|
@ -1885,20 +1892,12 @@ async def _user_api_key_auth_builder(
|
|||
current_models = _get_model_names_for_budget_checks(model=current_model)
|
||||
|
||||
# Check 5b. End-user model max budget
|
||||
end_user_mmb = valid_token.end_user_model_max_budget
|
||||
if (
|
||||
end_user_mmb is not None
|
||||
and isinstance(end_user_mmb, dict)
|
||||
and len(end_user_mmb) > 0
|
||||
and current_models
|
||||
and valid_token.end_user_id is not None
|
||||
):
|
||||
for model_name in current_models:
|
||||
await model_max_budget_limiter.is_end_user_within_model_budget(
|
||||
end_user_id=valid_token.end_user_id,
|
||||
end_user_model_max_budget=end_user_mmb,
|
||||
model=model_name,
|
||||
)
|
||||
await _enforce_end_user_model_max_budget_checks(
|
||||
valid_token=valid_token,
|
||||
request_data=request_data,
|
||||
route=route,
|
||||
request=request,
|
||||
)
|
||||
|
||||
# Check 6: Additional Common Checks across jwt + key auth
|
||||
if valid_token.team_id is not None:
|
||||
|
|
@ -2854,6 +2853,41 @@ def iter_router_fallback_model_names(fallbacks: Any) -> Iterator[str]:
|
|||
yield m["model"]
|
||||
|
||||
|
||||
async def _enforce_end_user_model_max_budget_checks(
|
||||
valid_token: UserAPIKeyAuth,
|
||||
request_data: dict,
|
||||
route: str,
|
||||
request: Request,
|
||||
) -> None:
|
||||
from litellm.proxy.proxy_server import llm_router, model_max_budget_limiter
|
||||
|
||||
end_user_mmb = valid_token.end_user_model_max_budget
|
||||
if (
|
||||
end_user_mmb is None
|
||||
or not isinstance(end_user_mmb, dict)
|
||||
or len(end_user_mmb) == 0
|
||||
or valid_token.end_user_id is None
|
||||
):
|
||||
return
|
||||
|
||||
current_model = _get_model_from_request_context(
|
||||
request_data=request_data,
|
||||
route=route,
|
||||
request=request,
|
||||
llm_router=llm_router,
|
||||
)
|
||||
current_models = _get_model_names_for_budget_checks(model=current_model)
|
||||
if not current_models:
|
||||
return
|
||||
|
||||
for model_name in current_models:
|
||||
await model_max_budget_limiter.is_end_user_within_model_budget(
|
||||
end_user_id=valid_token.end_user_id,
|
||||
end_user_model_max_budget=end_user_mmb,
|
||||
model=model_name,
|
||||
)
|
||||
|
||||
|
||||
async def _run_post_custom_auth_checks(
|
||||
valid_token: UserAPIKeyAuth,
|
||||
request: Request,
|
||||
|
|
@ -2955,20 +2989,12 @@ async def _run_post_custom_auth_checks(
|
|||
current_models = _get_model_names_for_budget_checks(model=current_model)
|
||||
|
||||
# 4. Check end-user model_max_budget
|
||||
end_user_mmb = valid_token.end_user_model_max_budget
|
||||
if (
|
||||
end_user_mmb is not None
|
||||
and isinstance(end_user_mmb, dict)
|
||||
and len(end_user_mmb) > 0
|
||||
and current_models
|
||||
and valid_token.end_user_id is not None
|
||||
):
|
||||
for model_name in current_models:
|
||||
await model_max_budget_limiter.is_end_user_within_model_budget(
|
||||
end_user_id=valid_token.end_user_id,
|
||||
end_user_model_max_budget=end_user_mmb,
|
||||
model=model_name,
|
||||
)
|
||||
await _enforce_end_user_model_max_budget_checks(
|
||||
valid_token=valid_token,
|
||||
request_data=request_data,
|
||||
route=route,
|
||||
request=request,
|
||||
)
|
||||
|
||||
# team / user / end_user / project context objects are fetched by
|
||||
# the centralized common_checks gate in user_api_key_auth after
|
||||
|
|
|
|||
|
|
@ -0,0 +1,58 @@
|
|||
import pytest
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import litellm
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.proxy.auth.user_api_key_auth import update_valid_token_with_end_user_params
|
||||
|
||||
|
||||
def test_update_valid_token_applies_end_user_model_max_budget_from_params():
|
||||
valid_token = UserAPIKeyAuth(token="test-key")
|
||||
end_user_params = {
|
||||
"end_user_id": "customer-1",
|
||||
"end_user_model_max_budget": {
|
||||
"google/gemini-2.5-flash-lite": {"max_budget": 1e-05, "budget_duration": "1d"}
|
||||
},
|
||||
}
|
||||
|
||||
result = update_valid_token_with_end_user_params(valid_token, end_user_params)
|
||||
|
||||
assert result.end_user_id == "customer-1"
|
||||
assert result.end_user_model_max_budget == end_user_params["end_user_model_max_budget"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_enforce_end_user_model_max_budget_raises_when_over_budget():
|
||||
from litellm.proxy.auth.user_api_key_auth import _enforce_end_user_model_max_budget_checks
|
||||
|
||||
valid_token = UserAPIKeyAuth(
|
||||
token="master-key",
|
||||
end_user_id="customer-1",
|
||||
end_user_model_max_budget={
|
||||
"google/gemini-2.5-flash-lite": {"max_budget": 1e-05, "budget_duration": "1d"}
|
||||
},
|
||||
)
|
||||
request = MagicMock()
|
||||
request_data = {"model": "google/gemini-2.5-flash-lite"}
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.auth.user_api_key_auth._get_model_from_request_context",
|
||||
return_value="google/gemini-2.5-flash-lite",
|
||||
):
|
||||
with patch(
|
||||
"litellm.proxy.proxy_server.model_max_budget_limiter.is_end_user_within_model_budget",
|
||||
new_callable=AsyncMock,
|
||||
) as mock_check:
|
||||
mock_check.side_effect = litellm.BudgetExceededError(
|
||||
message="Exceeded budget", current_cost=0.0002, max_budget=1e-05
|
||||
)
|
||||
|
||||
with pytest.raises(litellm.BudgetExceededError):
|
||||
await _enforce_end_user_model_max_budget_checks(
|
||||
valid_token=valid_token,
|
||||
request_data=request_data,
|
||||
route="/v1/chat/completions",
|
||||
request=request,
|
||||
)
|
||||
|
||||
mock_check.assert_awaited_once()
|
||||
Loading…
Add table
Reference in a new issue