From 4bd96852d9c56c8ded2c1e919d5dd2c286a5d3a7 Mon Sep 17 00:00:00 2001 From: Taranum Wasu Date: Sat, 4 Jul 2026 23:58:29 +0530 Subject: [PATCH] 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 --- litellm/proxy/auth/user_api_key_auth.py | 90 ++++++++++++------- ...t_end_user_model_max_budget_enforcement.py | 58 ++++++++++++ 2 files changed, 116 insertions(+), 32 deletions(-) create mode 100644 tests/proxy_unit_tests/test_end_user_model_max_budget_enforcement.py diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index 7944bb54d67..3006e3a00ef 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -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 diff --git a/tests/proxy_unit_tests/test_end_user_model_max_budget_enforcement.py b/tests/proxy_unit_tests/test_end_user_model_max_budget_enforcement.py new file mode 100644 index 00000000000..15ae5ec3b7b --- /dev/null +++ b/tests/proxy_unit_tests/test_end_user_model_max_budget_enforcement.py @@ -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()