diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index 11ba3da46af..8752d4595bc 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -1537,6 +1537,13 @@ async def _user_api_key_auth_builder( if _end_user_object is not None: valid_token.end_user_object_permission = _end_user_object.object_permission + await _maybe_enforce_master_key_end_user_model_max_budget( + valid_token=valid_token, + request_data=request_data, + route=route, + request=request, + ) + return valid_token if valid_token is not None and isinstance(valid_token, UserAPIKeyAuth) and valid_token.team_id is not None: @@ -1594,6 +1601,18 @@ async def _user_api_key_auth_builder( route=route, start_time=start_time, ) + + _user_api_key_obj = update_valid_token_with_end_user_params( + valid_token=_user_api_key_obj, end_user_params=end_user_params + ) + + await _maybe_enforce_master_key_end_user_model_max_budget( + valid_token=_user_api_key_obj, + request_data=request_data, + route=route, + request=request, + ) + asyncio.create_task( _cache_key_object( hashed_token=hash_token(master_key), @@ -1603,18 +1622,6 @@ async def _user_api_key_auth_builder( ) ) - _user_api_key_obj = update_valid_token_with_end_user_params( - valid_token=_user_api_key_obj, end_user_params=end_user_params - ) - - if RouteChecks.is_llm_api_route(route=route) and litellm.enforce_end_user_model_max_budget_on_master_key: - 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 @@ -2853,6 +2860,31 @@ def iter_router_fallback_model_names(fallbacks: Any) -> Iterator[str]: yield m["model"] +def _is_master_key_auth_token(valid_token: UserAPIKeyAuth) -> bool: + return valid_token.api_key == LITELLM_PROXY_MASTER_KEY_ALIAS or valid_token.token == LITELLM_PROXY_MASTER_KEY_ALIAS + + +async def _maybe_enforce_master_key_end_user_model_max_budget( + valid_token: UserAPIKeyAuth, + request_data: dict, + route: str, + request: Request, +) -> None: + if not litellm.enforce_end_user_model_max_budget_on_master_key: + return + if not RouteChecks.is_llm_api_route(route=route): + return + if not _is_master_key_auth_token(valid_token): + return + + await _enforce_end_user_model_max_budget_checks( + valid_token=valid_token, + request_data=request_data, + route=route, + request=request, + ) + + async def _enforce_end_user_model_max_budget_checks( valid_token: UserAPIKeyAuth, request_data: dict, 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 index aabc8bf904d..3ff2904359a 100644 --- 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 @@ -425,3 +425,147 @@ async def test_master_key_auth_enforces_end_user_model_budget_when_flag_enabled( litellm.enforce_end_user_model_max_budget_on_master_key = flag_original for k, v in originals.items(): setattr(proxy_server, k, v) + + +@pytest.mark.asyncio +async def test_cached_master_key_auth_enforces_end_user_model_budget_when_flag_enabled(): + from fastapi import Request + from starlette.datastructures import URL + + import litellm.proxy.proxy_server as proxy_server + from litellm.constants import LITELLM_PROXY_MASTER_KEY_ALIAS + from litellm.proxy._types import LitellmUserRoles + from litellm.proxy.auth.user_api_key_auth import _user_api_key_auth_builder + + cached_master = UserAPIKeyAuth( + api_key=LITELLM_PROXY_MASTER_KEY_ALIAS, + token=LITELLM_PROXY_MASTER_KEY_ALIAS, + user_role=LitellmUserRoles.PROXY_ADMIN, + ) + + async def mock_resolve_key(self, hashed_token: str): + from litellm.proxy.auth.resolvers.store import KeyNotInCacheError + + if self._check_cache_only: + return cached_master + raise KeyNotInCacheError(hashed_token) + + attrs, limiter = _proxy_server_attrs_for_master_key_auth() + originals = {k: getattr(proxy_server, k, None) for k in attrs} + flag_original = litellm.enforce_end_user_model_max_budget_on_master_key + litellm.enforce_end_user_model_max_budget_on_master_key = True + + try: + for k, v in attrs.items(): + setattr(proxy_server, k, v) + + request = Request(scope={"type": "http"}) + request._url = URL(url="/v1/chat/completions") + + with ( + patch( + "litellm.proxy.auth.resolvers.store.IdentityStore._resolve_key", + new=mock_resolve_key, + ), + patch( + "litellm.proxy.auth.user_api_key_auth.resolve_and_validate_end_user_id", + new_callable=AsyncMock, + return_value="customer-1", + ), + patch( + "litellm.proxy.auth.user_api_key_auth.get_end_user_object", + new_callable=AsyncMock, + return_value=_end_user_with_model_budget(), + ), + patch( + "litellm.proxy.auth.user_api_key_auth._get_model_from_request_context", + return_value=MODEL, + ), + ): + result = await _user_api_key_auth_builder( + request=request, + api_key=f"Bearer {attrs['master_key']}", + azure_api_key_header="", + anthropic_api_key_header=None, + google_ai_studio_api_key_header=None, + azure_apim_header=None, + request_data={"user": "customer-1", "model": MODEL}, + ) + + assert result.end_user_model_max_budget == {MODEL: MODEL_BUDGET} + limiter.is_end_user_within_model_budget.assert_awaited() + finally: + litellm.enforce_end_user_model_max_budget_on_master_key = flag_original + for k, v in originals.items(): + setattr(proxy_server, k, v) + + +@pytest.mark.asyncio +async def test_cached_proxy_admin_virtual_key_skips_master_key_budget_enforcement(): + from fastapi import Request + from starlette.datastructures import URL + + import litellm.proxy.proxy_server as proxy_server + from litellm.proxy._types import LitellmUserRoles + from litellm.proxy.auth.user_api_key_auth import _user_api_key_auth_builder + + cached_admin_key = UserAPIKeyAuth( + api_key="sk-admin-virtual", + token="hashed-admin-virtual", + user_role=LitellmUserRoles.PROXY_ADMIN, + ) + + async def mock_resolve_key(self, hashed_token: str): + from litellm.proxy.auth.resolvers.store import KeyNotInCacheError + + if self._check_cache_only: + return cached_admin_key + raise KeyNotInCacheError(hashed_token) + + attrs, limiter = _proxy_server_attrs_for_master_key_auth() + limiter.is_end_user_within_model_budget.side_effect = litellm.BudgetExceededError( + message="Exceeded budget", current_cost=0.0002, max_budget=1e-05 + ) + originals = {k: getattr(proxy_server, k, None) for k in attrs} + flag_original = litellm.enforce_end_user_model_max_budget_on_master_key + litellm.enforce_end_user_model_max_budget_on_master_key = True + + try: + for k, v in attrs.items(): + setattr(proxy_server, k, v) + + request = Request(scope={"type": "http"}) + request._url = URL(url="/v1/chat/completions") + + with ( + patch( + "litellm.proxy.auth.resolvers.store.IdentityStore._resolve_key", + new=mock_resolve_key, + ), + patch( + "litellm.proxy.auth.user_api_key_auth.resolve_and_validate_end_user_id", + new_callable=AsyncMock, + return_value="customer-1", + ), + patch( + "litellm.proxy.auth.user_api_key_auth.get_end_user_object", + new_callable=AsyncMock, + return_value=_end_user_with_model_budget(), + ), + ): + result = await _user_api_key_auth_builder( + request=request, + api_key="Bearer sk-admin-virtual", + azure_api_key_header="", + anthropic_api_key_header=None, + google_ai_studio_api_key_header=None, + azure_apim_header=None, + request_data={"user": "customer-1", "model": MODEL}, + ) + + assert result.user_role == LitellmUserRoles.PROXY_ADMIN + limiter.is_end_user_within_model_budget.assert_not_awaited() + finally: + litellm.enforce_end_user_model_max_budget_on_master_key = flag_original + for k, v in originals.items(): + setattr(proxy_server, k, v)