From 1ab2814b855ff98a23c24ccf685e9d55cfad2e2c Mon Sep 17 00:00:00 2001 From: Fabian Hoeldin Date: Wed, 8 Apr 2026 15:49:56 +0200 Subject: [PATCH] fix: end-user budget enforcement was bypassed on master key auth Budget checks were performed in get_end_user_object() which ran before model resolution. This had two issues: 1. Master key auth returned early, skipping the budget check entirely 2. Over-budget users were blocked from zero-cost models Moved budget check after model resolution in user_api_key_auth.py to: - Enforce budgets for master key auth paths - Allow over-budget users to access zero-cost models --- litellm/proxy/auth/auth_checks.py | 56 +--- litellm/proxy/auth/user_api_key_auth.py | 94 ++++-- tests/proxy_unit_tests/test_auth_checks.py | 111 +++---- .../test_default_end_user_budget_simple.py | 86 ++--- .../test_user_api_key_auth.py | 301 ++++++++++++++---- 5 files changed, 419 insertions(+), 229 deletions(-) diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index 68bde8434a6..a52f8b38169 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -8,6 +8,7 @@ Run checks for: 2. If user is in budget 3. If end_user ('user' passed to /chat/completions, /embeddings endpoint) is in budget """ + import asyncio import re import time @@ -196,9 +197,7 @@ def _is_model_cost_zero( return True -def _is_cost_explicitly_configured( - model: str, llm_router: "Router" -) -> bool: +def _is_cost_explicitly_configured(model: str, llm_router: "Router") -> bool: """ Check if any deployment in the model group has cost fields explicitly set in its litellm.model_cost entry. @@ -215,10 +214,7 @@ def _is_cost_explicitly_configured( if model_id is None: continue raw_entry = litellm.model_cost.get(model_id, {}) - if ( - "input_cost_per_token" in raw_entry - or "output_cost_per_token" in raw_entry - ): + if "input_cost_per_token" in raw_entry or "output_cost_per_token" in raw_entry: return True return False @@ -456,9 +452,9 @@ async def common_checks( # noqa: PLR0915 model=_model, team_object=team_object, llm_router=llm_router, - team_model_aliases=valid_token.team_model_aliases - if valid_token - else None, + team_model_aliases=( + valid_token.team_model_aliases if valid_token else None + ), ): raise ProxyException( message=f"Team not allowed to access model. Team={team_object.team_id}, Model={_model}. Allowed team models = {team_object.models}", @@ -925,35 +921,6 @@ async def _apply_default_budget_to_end_user( return end_user_obj -def _check_end_user_budget( - end_user_obj: LiteLLM_EndUserTable, - route: str, -) -> None: - """ - Check if end user is within their budget limit. - - Args: - end_user_obj: The end user object to check - route: The request route - - Raises: - litellm.BudgetExceededError: If end user has exceeded their budget - """ - if RouteChecks.is_info_route(route): - return - - if end_user_obj.litellm_budget_table is None: - return - - end_user_budget = end_user_obj.litellm_budget_table.max_budget - if end_user_budget is not None and end_user_obj.spend > end_user_budget: - raise litellm.BudgetExceededError( - current_cost=end_user_obj.spend, - max_budget=end_user_budget, - message=f"ExceededBudget: End User={end_user_obj.user_id} over budget. Spend={end_user_obj.spend}, Budget={end_user_budget}", - ) - - @log_db_metrics async def get_end_user_object( end_user_id: Optional[str], @@ -969,11 +936,14 @@ async def get_end_user_object( If end user exists but has no budget_id, applies the default budget (if configured via litellm.max_end_user_budget_id). + Note: Budget checks are handled in common_checks(), not here. + This allows zero-cost models to skip budget checks properly. + Args: end_user_id: The ID of the end user prisma_client: Database client instance user_api_key_cache: Cache for storing/retrieving data - route: The request route + route: The request route (unused, kept for API compatibility) parent_otel_span: Optional OpenTelemetry span for tracing proxy_logging_obj: Optional proxy logging object @@ -1001,9 +971,6 @@ async def get_end_user_object( parent_otel_span=parent_otel_span, ) - # Check budget limits - _check_end_user_budget(end_user_obj=return_obj, route=route) - return return_obj # Fetch from database @@ -1032,9 +999,6 @@ async def get_end_user_object( key="end_user_id:{}".format(end_user_id), value=_response.dict() ) - # Check budget limits - _check_end_user_budget(end_user_obj=_response, route=route) - return _response except Exception as e: diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index 61c618eeb18..b9fb3567297 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -247,6 +247,50 @@ def _apply_budget_limits_to_end_user_params( verbose_proxy_logger.debug(f"Applied budget limits to end user {end_user_id}") +async def _check_end_user_budget( + end_user_object: Optional[Any], + request_data: dict, + route: str, + llm_router: Optional[Any], +) -> None: + """ + Check if end user has exceeded their budget. Raises BudgetExceededError if so. + Skips check for zero-cost models. + + Args: + end_user_object: The end user object from DB + request_data: Request body data + route: Current route + llm_router: LLM router instance + + Raises: + BudgetExceededError: If end user has exceeded their max_budget + """ + if end_user_object is None: + return + + # Check if model has zero cost - if so, skip budget check + model = get_model_from_request(request_data, route) + if model is not None and llm_router is not None: + from litellm.proxy.auth.auth_checks import _is_model_cost_zero + + if _is_model_cost_zero(model=model, llm_router=llm_router): + verbose_proxy_logger.info( + f"Skipping all budget checks for zero-cost model: {model}" + ) + return + + # Check end-user budget + if end_user_object.litellm_budget_table is not None: + end_user_budget = end_user_object.litellm_budget_table.max_budget + if end_user_budget is not None and end_user_object.spend > end_user_budget: + raise litellm.BudgetExceededError( + current_cost=end_user_object.spend, + max_budget=end_user_budget, + message=f"ExceededBudget: End User={end_user_object.user_id} over budget. Spend={end_user_object.spend}, Budget={end_user_budget}", + ) + + async def user_api_key_auth_websocket(websocket: WebSocket): # Accept the WebSocket connection @@ -314,6 +358,8 @@ def update_valid_token_with_end_user_params( valid_token.end_user_rpm_limit = end_user_params["end_user_rpm_limit"] if end_user_params.get("allowed_model_region") is not None: valid_token.allowed_model_region = end_user_params["allowed_model_region"] + if end_user_params.get("end_user_max_budget") is not None: + valid_token.end_user_max_budget = end_user_params["end_user_max_budget"] if end_user_params.get("end_user_model_max_budget") is not None: valid_token.end_user_model_max_budget = end_user_params[ "end_user_model_max_budget" @@ -696,11 +742,8 @@ async def _user_api_key_auth_builder( # noqa: PLR0915 # Routing uses unverified JWT claims only to choose auth path. # Final authentication is enforced by the selected validator. - route_jwt_to_oauth2 = ( - is_jwt - and _should_route_jwt_to_oauth2_override( - token=api_key, jwt_handler=jwt_handler - ) + route_jwt_to_oauth2 = is_jwt and _should_route_jwt_to_oauth2_override( + token=api_key, jwt_handler=jwt_handler ) # OAuth2 applies for: @@ -716,6 +759,7 @@ async def _user_api_key_auth_builder( # noqa: PLR0915 ) if (should_apply_global_oauth2 and not is_jwt) or should_apply_override_oauth2: from litellm.proxy.proxy_server import premium_user + if premium_user is not True: raise ValueError( "Oauth2 token validation is only available for premium users" @@ -743,10 +787,7 @@ async def _user_api_key_auth_builder( # noqa: PLR0915 if jwt_handler.litellm_jwtauth.virtual_key_claim_field is not None: # Decode JWT to get claims without running full auth_builder jwt_claims: Optional[dict] - if ( - jwt_handler.litellm_jwtauth.oidc_userinfo_enabled - and not is_jwt - ): + if jwt_handler.litellm_jwtauth.oidc_userinfo_enabled and not is_jwt: jwt_claims = await jwt_handler.get_oidc_userinfo(token=api_key) else: jwt_claims = await jwt_handler.auth_jwt(token=api_key) @@ -981,9 +1022,9 @@ async def _user_api_key_auth_builder( # noqa: PLR0915 route=route, ) if _end_user_object is not None: - end_user_params[ - "allowed_model_region" - ] = _end_user_object.allowed_model_region + end_user_params["allowed_model_region"] = ( + _end_user_object.allowed_model_region + ) if _end_user_object.litellm_budget_table is not None: _apply_budget_limits_to_end_user_params( end_user_params=end_user_params, @@ -1078,6 +1119,14 @@ async def _user_api_key_auth_builder( # noqa: PLR0915 _end_user_object.object_permission ) + # Check end-user budget before returning + await _check_end_user_budget( + end_user_object=_end_user_object, + request_data=request_data, + route=route, + llm_router=llm_router, + ) + return valid_token if ( @@ -1156,6 +1205,14 @@ async def _user_api_key_auth_builder( # noqa: PLR0915 valid_token=_user_api_key_obj, end_user_params=end_user_params ) + # Check end-user budget before returning + await _check_end_user_budget( + end_user_object=_end_user_object, + request_data=request_data, + route=route, + llm_router=llm_router, + ) + return _user_api_key_obj ## IF it's not a master key @@ -1537,9 +1594,9 @@ async def _user_api_key_auth_builder( # noqa: PLR0915 if _end_user_object is not None: valid_token_dict.update(end_user_params) - valid_token_dict[ - "end_user_object_permission" - ] = _end_user_object.object_permission + valid_token_dict["end_user_object_permission"] = ( + _end_user_object.object_permission + ) # check if token is from litellm-ui, litellm ui makes keys to allow users to login with sso. These keys can only be used for LiteLLM UI functions # sso/login, ui/login, /key functions and /user functions @@ -1802,12 +1859,9 @@ async def _enforce_key_and_fallback_model_access( if config != {}: model_list = config.get("model_list", []) new_model_list = model_list - verbose_proxy_logger.debug( - f"\n new llm router model list {new_model_list}" - ) + verbose_proxy_logger.debug(f"\n new llm router model list {new_model_list}") elif ( - isinstance(valid_token.models, list) - and "all-team-models" in valid_token.models + isinstance(valid_token.models, list) and "all-team-models" in valid_token.models ): pass else: diff --git a/tests/proxy_unit_tests/test_auth_checks.py b/tests/proxy_unit_tests/test_auth_checks.py index c92ec61b9b2..78378ed7f1e 100644 --- a/tests/proxy_unit_tests/test_auth_checks.py +++ b/tests/proxy_unit_tests/test_auth_checks.py @@ -8,9 +8,7 @@ from dotenv import load_dotenv load_dotenv() import os -sys.path.insert( - 0, os.path.abspath("../..") -) # Adds the parent directory to the system path +sys.path.insert(0, os.path.abspath("../..")) # Adds the parent directory to the system path import pytest, litellm import httpx from litellm.proxy._types import UserAPIKeyAuth @@ -38,7 +36,10 @@ from litellm.proxy.utils import CallInfo async def test_get_end_user_object(customer_spend, customer_budget): """ Scenario 1: normal - Scenario 2: user over budget + Scenario 2: user over budget - budget check moved to user_api_key_auth.py after model resolution + + get_end_user_object() should return the user object regardless of budget status. + Budget enforcement happens later in user_api_key_auth.py to allow zero-cost model access. """ end_user_id = "my-test-customer" _budget = LiteLLM_BudgetTable(max_budget=customer_budget) @@ -51,31 +52,15 @@ async def test_get_end_user_object(customer_spend, customer_budget): _cache = DualCache() _key = "end_user_id:{}".format(end_user_id) _cache.set_cache(key=_key, value=end_user_obj.model_dump()) - try: - await get_end_user_object( - end_user_id=end_user_id, - prisma_client="RANDOM VALUE", # type: ignore - user_api_key_cache=_cache, - route="/v1/chat/completions", - ) - if customer_spend > customer_budget: - pytest.fail( - "Expected call to fail. Customer Spend={}, Customer Budget={}".format( - customer_spend, customer_budget - ) - ) - except Exception as e: - if ( - isinstance(e, litellm.BudgetExceededError) - and customer_spend > customer_budget - ): - pass - else: - pytest.fail( - "Expected call to work. Customer Spend={}, Customer Budget={}, Error={}".format( - customer_spend, customer_budget, str(e) - ) - ) + result = await get_end_user_object( + end_user_id=end_user_id, + prisma_client="RANDOM VALUE", + user_api_key_cache=_cache, + route="/v1/chat/completions", + ) + assert result is not None + assert result.user_id == end_user_id + assert result.spend == customer_spend @pytest.mark.parametrize( @@ -402,16 +387,12 @@ async def test_is_valid_fallback_model(): ) try: - await is_valid_fallback_model( - model="gpt-3.5-turbo", llm_router=router, user_model=None - ) + await is_valid_fallback_model(model="gpt-3.5-turbo", llm_router=router, user_model=None) except Exception as e: pytest.fail(f"Expected is_valid_fallback_model to work, got exception: {e}") try: - await is_valid_fallback_model( - model="gpt-4o", llm_router=router, user_model=None - ) + await is_valid_fallback_model(model="gpt-4o", llm_router=router, user_model=None) pytest.fail("Expected is_valid_fallback_model to fail") except Exception as e: assert "Invalid" in str(e) @@ -426,9 +407,7 @@ async def test_is_valid_fallback_model(): ], ) @pytest.mark.asyncio -async def test_virtual_key_max_budget_check( - token_spend, max_budget, expect_budget_error -): +async def test_virtual_key_max_budget_check(token_spend, max_budget, expect_budget_error): """ Test if virtual key budget checks work as expected: 1. Triggers budget alert for all cases @@ -472,14 +451,10 @@ async def test_virtual_key_max_budget_check( user_obj=user_obj, ) if expect_budget_error: - pytest.fail( - f"Expected BudgetExceededError for spend={token_spend}, max_budget={max_budget}" - ) + pytest.fail(f"Expected BudgetExceededError for spend={token_spend}, max_budget={max_budget}") except litellm.BudgetExceededError as e: if not expect_budget_error: - pytest.fail( - f"Unexpected BudgetExceededError for spend={token_spend}, max_budget={max_budget}" - ) + pytest.fail(f"Unexpected BudgetExceededError for spend={token_spend}, max_budget={max_budget}") assert e.current_cost == token_spend assert e.max_budget == max_budget @@ -546,9 +521,7 @@ async def test_can_team_access_model(model, team_models, expect_to_work): team_model_aliases=None, ) if not expect_to_work: - pytest.fail( - f"Expected model access check to fail for model={model}, team_models={team_models}" - ) + pytest.fail(f"Expected model access check to fail for model={model}, team_models={team_models}") except Exception as e: if expect_to_work: pytest.fail( @@ -601,9 +574,9 @@ async def test_virtual_key_soft_budget_check(spend, soft_budget, expect_alert): await asyncio.sleep(0.1) # Allow time for the alert task to complete - assert ( - alert_triggered == expect_alert - ), f"Expected alert_triggered to be {expect_alert} for spend={spend}, soft_budget={soft_budget}" + assert alert_triggered == expect_alert, ( + f"Expected alert_triggered to be {expect_alert} for spend={spend}, soft_budget={soft_budget}" + ) @pytest.mark.parametrize( @@ -613,9 +586,27 @@ async def test_virtual_key_soft_budget_check(spend, soft_budget, expect_alert): (50, 50, False, None, None), # At soft budget, no metadata - no alert_emails configured, so no alert (25, 50, False, None, None), # Under soft budget (100, None, False, None, None), # No soft budget set - (100, 50, True, {"soft_budget_alerting_emails": ["team1@example.com", "team2@example.com"]}, ["team1@example.com", "team2@example.com"]), # Over soft budget with list of emails - (100, 50, True, {"soft_budget_alerting_emails": "team1@example.com,team2@example.com"}, ["team1@example.com", "team2@example.com"]), # Over soft budget with comma-separated emails - (100, 50, True, {"soft_budget_alerting_emails": ["team1@example.com", "", " ", "team2@example.com"]}, ["team1@example.com", "team2@example.com"]), # Over soft budget with empty strings filtered + ( + 100, + 50, + True, + {"soft_budget_alerting_emails": ["team1@example.com", "team2@example.com"]}, + ["team1@example.com", "team2@example.com"], + ), # Over soft budget with list of emails + ( + 100, + 50, + True, + {"soft_budget_alerting_emails": "team1@example.com,team2@example.com"}, + ["team1@example.com", "team2@example.com"], + ), # Over soft budget with comma-separated emails + ( + 100, + 50, + True, + {"soft_budget_alerting_emails": ["team1@example.com", "", " ", "team2@example.com"]}, + ["team1@example.com", "team2@example.com"], + ), # Over soft budget with empty strings filtered ], ) @pytest.mark.asyncio @@ -667,9 +658,9 @@ async def test_team_soft_budget_check(spend, soft_budget, expect_alert, metadata await asyncio.sleep(0.1) # Allow time for the alert task to complete - assert ( - alert_triggered == expect_alert - ), f"Expected alert_triggered to be {expect_alert} for spend={spend}, soft_budget={soft_budget}" + assert alert_triggered == expect_alert, ( + f"Expected alert_triggered to be {expect_alert} for spend={spend}, soft_budget={soft_budget}" + ) if expect_alert: assert captured_call_info is not None @@ -813,9 +804,7 @@ async def test_get_fuzzy_user_object(): # Test 5: Only email provided (no SSO ID) mock_prisma.db.litellm_usertable.find_first = AsyncMock(return_value=test_user) - result = await _get_fuzzy_user_object( - prisma_client=mock_prisma, user_email="test@example.com" - ) + result = await _get_fuzzy_user_object(prisma_client=mock_prisma, user_email="test@example.com") assert result == test_user mock_prisma.db.litellm_usertable.find_first.assert_called_with( where={"user_email": {"equals": "test@example.com", "mode": "insensitive"}}, @@ -824,9 +813,7 @@ async def test_get_fuzzy_user_object(): # Test 6: Only SSO ID provided (no email) mock_prisma.db.litellm_usertable.find_unique = AsyncMock(return_value=test_user) - result = await _get_fuzzy_user_object( - prisma_client=mock_prisma, sso_user_id="sso_123" - ) + result = await _get_fuzzy_user_object(prisma_client=mock_prisma, sso_user_id="sso_123") assert result == test_user mock_prisma.db.litellm_usertable.find_unique.assert_called_with( where={"sso_user_id": "sso_123"}, include={"organization_memberships": True} diff --git a/tests/proxy_unit_tests/test_default_end_user_budget_simple.py b/tests/proxy_unit_tests/test_default_end_user_budget_simple.py index 92ca1f71703..9fedef3c37b 100644 --- a/tests/proxy_unit_tests/test_default_end_user_budget_simple.py +++ b/tests/proxy_unit_tests/test_default_end_user_budget_simple.py @@ -8,14 +8,14 @@ a default budget to end users without explicit budgets. import sys import os import uuid -from unittest.mock import AsyncMock, MagicMock, patch +from unittest.mock import AsyncMock, MagicMock import pytest sys.path.insert(0, os.path.abspath("../..")) import litellm -from litellm.proxy._types import LiteLLM_BudgetTable, LiteLLM_EndUserTable +from litellm.proxy._types import LiteLLM_BudgetTable from litellm.proxy.auth.auth_checks import get_end_user_object from litellm.caching import DualCache @@ -29,14 +29,14 @@ async def test_default_budget_applied_to_end_user_without_budget(): end_user_id = f"test_user_{uuid.uuid4().hex}" default_budget_id = str(uuid.uuid4()) litellm.max_end_user_budget_id = default_budget_id - + default_budget = LiteLLM_BudgetTable( budget_id=default_budget_id, max_budget=10.0, rpm_limit=2, tpm_limit=10, ) - + # Mock end user in DB without budget mock_end_user_data = { "user_id": end_user_id, @@ -47,7 +47,7 @@ async def test_default_budget_applied_to_end_user_without_budget(): "default_model": None, "blocked": False, } - + mock_prisma_client = MagicMock() mock_prisma_client.db.litellm_endusertable.find_unique = AsyncMock( return_value=MagicMock(dict=lambda: mock_end_user_data) @@ -55,18 +55,18 @@ async def test_default_budget_applied_to_end_user_without_budget(): mock_prisma_client.db.litellm_budgettable.find_unique = AsyncMock( return_value=MagicMock(dict=lambda: default_budget.dict()) ) - + mock_cache = AsyncMock(spec=DualCache) mock_cache.async_get_cache = AsyncMock(return_value=None) mock_cache.async_set_cache = AsyncMock() - + result = await get_end_user_object( end_user_id=end_user_id, prisma_client=mock_prisma_client, user_api_key_cache=mock_cache, route="/chat/completions", ) - + # Verify default budget was applied assert result is not None assert result.litellm_budget_table is not None @@ -74,7 +74,7 @@ async def test_default_budget_applied_to_end_user_without_budget(): assert result.litellm_budget_table.max_budget == 10.0 assert result.litellm_budget_table.rpm_limit == 2 assert result.litellm_budget_table.tpm_limit == 10 - + litellm.max_end_user_budget_id = None @@ -88,13 +88,13 @@ async def test_explicit_budget_not_overridden_by_default(): explicit_budget_id = str(uuid.uuid4()) default_budget_id = str(uuid.uuid4()) litellm.max_end_user_budget_id = default_budget_id - + explicit_budget = LiteLLM_BudgetTable( budget_id=explicit_budget_id, max_budget=100.0, rpm_limit=50, ) - + # Mock end user with explicit budget mock_end_user_data = { "user_id": end_user_id, @@ -105,48 +105,49 @@ async def test_explicit_budget_not_overridden_by_default(): "default_model": None, "blocked": False, } - + mock_prisma_client = MagicMock() mock_prisma_client.db.litellm_endusertable.find_unique = AsyncMock( return_value=MagicMock(dict=lambda: mock_end_user_data) ) - + mock_cache = AsyncMock(spec=DualCache) mock_cache.async_get_cache = AsyncMock(return_value=None) mock_cache.async_set_cache = AsyncMock() - + result = await get_end_user_object( end_user_id=end_user_id, prisma_client=mock_prisma_client, user_api_key_cache=mock_cache, route="/chat/completions", ) - + # Verify explicit budget is kept (not replaced with default) assert result is not None assert result.litellm_budget_table.budget_id == explicit_budget_id assert result.litellm_budget_table.max_budget == 100.0 assert result.litellm_budget_table.rpm_limit == 50 - + litellm.max_end_user_budget_id = None @pytest.mark.asyncio -async def test_budget_enforcement_blocks_over_budget_users(): +async def test_over_budget_user_gets_budget_applied(): """ - Core scenario: Budget limits are actually enforced. - Users who exceed their budget should be blocked. + Budget enforcement moved to user_api_key_auth.py after model resolution. + get_end_user_object() should still apply the budget, but NOT raise error. + This allows zero-cost models to bypass budget checks. """ end_user_id = f"test_user_{uuid.uuid4().hex}" default_budget_id = str(uuid.uuid4()) litellm.max_end_user_budget_id = default_budget_id - + default_budget = LiteLLM_BudgetTable( budget_id=default_budget_id, max_budget=10.0, rpm_limit=2, ) - + # Mock end user who has already spent more than budget mock_end_user_data = { "user_id": end_user_id, @@ -157,7 +158,7 @@ async def test_budget_enforcement_blocks_over_budget_users(): "default_model": None, "blocked": False, } - + mock_prisma_client = MagicMock() mock_prisma_client.db.litellm_endusertable.find_unique = AsyncMock( return_value=MagicMock(dict=lambda: mock_end_user_data) @@ -165,23 +166,25 @@ async def test_budget_enforcement_blocks_over_budget_users(): mock_prisma_client.db.litellm_budgettable.find_unique = AsyncMock( return_value=MagicMock(dict=lambda: default_budget.dict()) ) - + mock_cache = AsyncMock(spec=DualCache) mock_cache.async_get_cache = AsyncMock(return_value=None) mock_cache.async_set_cache = AsyncMock() - - # Should raise BudgetExceededError - with pytest.raises(litellm.BudgetExceededError) as exc_info: - await get_end_user_object( - end_user_id=end_user_id, - prisma_client=mock_prisma_client, - user_api_key_cache=mock_cache, - route="/chat/completions", - ) - - assert "ExceededBudget" in str(exc_info.value) - assert end_user_id in str(exc_info.value) - + + result = await get_end_user_object( + end_user_id=end_user_id, + prisma_client=mock_prisma_client, + user_api_key_cache=mock_cache, + route="/chat/completions", + ) + + # Verify default budget was applied (even though user is over budget) + # Budget check happens later in user_api_key_auth.py after model resolution + assert result is not None + assert result.litellm_budget_table is not None + assert result.litellm_budget_table.max_budget == 10.0 + assert result.spend == 15.0 # Over budget, but no error raised here + litellm.max_end_user_budget_id = None @@ -193,7 +196,7 @@ async def test_system_works_without_default_budget_configured(): """ end_user_id = f"test_user_{uuid.uuid4().hex}" litellm.max_end_user_budget_id = None # Not configured - + # Mock end user without budget mock_end_user_data = { "user_id": end_user_id, @@ -204,25 +207,24 @@ async def test_system_works_without_default_budget_configured(): "default_model": None, "blocked": False, } - + mock_prisma_client = MagicMock() mock_prisma_client.db.litellm_endusertable.find_unique = AsyncMock( return_value=MagicMock(dict=lambda: mock_end_user_data) ) - + mock_cache = AsyncMock(spec=DualCache) mock_cache.async_get_cache = AsyncMock(return_value=None) mock_cache.async_set_cache = AsyncMock() - + result = await get_end_user_object( end_user_id=end_user_id, prisma_client=mock_prisma_client, user_api_key_cache=mock_cache, route="/chat/completions", ) - + # Should work fine, just without budget limits assert result is not None assert result.user_id == end_user_id assert result.litellm_budget_table is None # No budget applied - diff --git a/tests/proxy_unit_tests/test_user_api_key_auth.py b/tests/proxy_unit_tests/test_user_api_key_auth.py index 1a6e2eda9a1..cf12991612a 100644 --- a/tests/proxy_unit_tests/test_user_api_key_auth.py +++ b/tests/proxy_unit_tests/test_user_api_key_auth.py @@ -7,9 +7,7 @@ import sys import litellm.proxy import litellm.proxy.proxy_server -sys.path.insert( - 0, os.path.abspath("../..") -) # Adds the parent directory to the system path +sys.path.insert(0, os.path.abspath("../..")) # Adds the parent directory to the system path from typing import Dict, List, Optional from unittest.mock import MagicMock, patch, AsyncMock @@ -50,9 +48,7 @@ class Request: ), # Request with no client IP should not be allowed ], ) -def test_check_valid_ip( - allowed_ips: Optional[List[str]], client_ip: Optional[str], expected_result: bool -): +def test_check_valid_ip(allowed_ips: Optional[List[str]], client_ip: Optional[str], expected_result: bool): from litellm.proxy.auth.auth_utils import _check_valid_ip request = Request(client_ip) @@ -121,9 +117,7 @@ async def test_check_blocked_team(): last_refreshed_at=time.time(), ) await asyncio.sleep(1) - team_obj = LiteLLM_TeamTableCachedObj( - team_id=_team_id, blocked=False, last_refreshed_at=time.time() - ) + team_obj = LiteLLM_TeamTableCachedObj(team_id=_team_id, blocked=False, last_refreshed_at=time.time()) hashed_token = hash_token(user_key) print(f"STORING TOKEN UNDER KEY={hashed_token}") user_api_key_cache.set_cache(key=hashed_token, value=valid_token) @@ -173,9 +167,7 @@ async def test_team_object_has_object_permission_id(): request = Request(scope={"type": "http"}) request._url = URL(url="/chat/completions") - with patch( - "litellm.proxy.auth.user_api_key_auth.common_checks", new_callable=AsyncMock - ) as mock_common_checks: + with patch("litellm.proxy.auth.user_api_key_auth.common_checks", new_callable=AsyncMock) as mock_common_checks: mock_common_checks.return_value = True await user_api_key_auth(request=request, api_key="Bearer " + user_key) @@ -200,9 +192,7 @@ async def test_returned_user_api_key_auth(user_role, expected_role): from datetime import datetime new_obj = await _return_user_api_key_auth_obj( - user_obj=LiteLLM_UserTable( - user_role=user_role, user_id="", max_budget=None, user_email="" - ), + user_obj=LiteLLM_UserTable(user_role=user_role, user_id="", max_budget=None, user_email=""), api_key="hello-world", parent_otel_span=None, valid_token_dict={}, @@ -253,9 +243,7 @@ async def test_aaauser_personal_budgets(key_ownership): spend=20, ) - user_obj = LiteLLM_UserTable( - user_id=_user_id, spend=11, max_budget=10, user_email="" - ) + user_obj = LiteLLM_UserTable(user_id=_user_id, spend=11, max_budget=10, user_email="") user_api_key_cache.set_cache(key=hash_token(user_key), value=valid_token) user_api_key_cache.set_cache(key="{}".format(_user_id), value=user_obj) @@ -305,9 +293,7 @@ async def test_user_api_key_auth_fails_with_prohibited_params(prohibited_param): request.body = return_body try: - response = await user_api_key_auth( - request=request, api_key="Bearer " + user_key - ) + response = await user_api_key_auth(request=request, api_key="Bearer " + user_key) except Exception as e: print("error str=", str(e)) error_message = str(e.message) @@ -502,9 +488,7 @@ def test_get_api_key_from_custom_header(headers, custom_header_name, expected_ap # Call the function and verify it doesn't raise an exception - api_key = get_api_key_from_custom_header( - request=request, custom_litellm_key_header_name=custom_header_name - ) + api_key = get_api_key_from_custom_header(request=request, custom_litellm_key_header_name=custom_header_name) assert api_key == expected_api_key @@ -521,9 +505,7 @@ from litellm.proxy._types import LitellmUserRoles (LitellmUserRoles.TEAM, "1234", "1234", True), ], ) -def test_allowed_route_inside_route( - user_role, auth_user_id, requested_user_id, expected_result -): +def test_allowed_route_inside_route(user_role, auth_user_id, requested_user_id, expected_result): from litellm.proxy.auth.auth_checks import allowed_route_check_inside_route from litellm.proxy._types import UserAPIKeyAuth, LitellmUserRoles @@ -664,9 +646,7 @@ async def test_soft_budget_alert(): try: # Call user_api_key_auth - response = await user_api_key_auth( - request=request, api_key="Bearer " + user_key - ) + response = await user_api_key_auth(request=request, api_key="Bearer " + user_key) # Assert the request was allowed (no exception raised) assert response is not None @@ -833,10 +813,7 @@ async def test_user_api_key_auth_websocket(): mock_websocket.url = URL(url="/ws") # Mock the return value of `user_api_key_auth` when it's called within the `user_api_key_auth_websocket` function - with patch( - "litellm.proxy.auth.user_api_key_auth.user_api_key_auth", autospec=True - ) as mock_user_api_key_auth: - + with patch("litellm.proxy.auth.user_api_key_auth.user_api_key_auth", autospec=True) as mock_user_api_key_auth: # Make the call to the WebSocket function await user_api_key_auth_websocket(mock_websocket) @@ -845,15 +822,13 @@ async def test_user_api_key_auth_websocket(): # Get the request object that was passed to user_api_key_auth request_arg = mock_user_api_key_auth.call_args.kwargs["request"] - + # Verify that the request has headers set assert hasattr(request_arg, "headers"), "Request object should have headers attribute" assert "authorization" in request_arg.headers, "Request headers should contain authorization" assert request_arg.headers["authorization"] == "Bearer some_api_key" - assert ( - mock_user_api_key_auth.call_args.kwargs["api_key"] == "Bearer some_api_key" - ) + assert mock_user_api_key_auth.call_args.kwargs["api_key"] == "Bearer some_api_key" @pytest.mark.parametrize("enforce_rbac", [True, False]) @@ -1035,14 +1010,10 @@ async def test_jwt_non_admin_team_route_access(monkeypatch): } # Create request - request = Request( - scope={"type": "http", "headers": [(b"authorization", b"Bearer fake.jwt.token")]} - ) + request = Request(scope={"type": "http", "headers": [(b"authorization", b"Bearer fake.jwt.token")]}) request._url = URL(url="/team/new") - monkeypatch.setattr( - litellm.proxy.proxy_server, "general_settings", {"enable_jwt_auth": True} - ) + monkeypatch.setattr(litellm.proxy.proxy_server, "general_settings", {"enable_jwt_auth": True}) # Initialize jwt_handler with a default LiteLLM_JWTAuth so that the # virtual_key_claim_field check in user_api_key_auth doesn't fail with @@ -1059,18 +1030,19 @@ async def test_jwt_non_admin_team_route_access(monkeypatch): # Mock enterprise license check and JWTAuthManager.auth_builder # License check must be mocked to avoid environment variable pollution # in parallel test execution - with patch( - "litellm.proxy.proxy_server.premium_user", - True, - ), patch( - "litellm.proxy.auth.handle_jwt.JWTAuthManager.auth_builder", - return_value=mock_jwt_response, + with ( + patch( + "litellm.proxy.proxy_server.premium_user", + True, + ), + patch( + "litellm.proxy.auth.handle_jwt.JWTAuthManager.auth_builder", + return_value=mock_jwt_response, + ), ): try: await user_api_key_auth(request=request, api_key="Bearer fake.jwt.token") - pytest.fail( - "Expected this call to fail. Non-admin user should not access team routes." - ) + pytest.fail("Expected this call to fail. Non-admin user should not access team routes.") except ProxyException as e: print("e", e) assert "Only proxy admin can be used to generate" in str(e.message) @@ -1101,14 +1073,12 @@ async def test_x_litellm_api_key(): ignored_key = "aj12445" # Create request with headers as bytes - request = Request( - scope={ - "type": "http" - } - ) + request = Request(scope={"type": "http"}) request._url = URL(url="/chat/completions") - valid_token = await user_api_key_auth(request=request, api_key="Bearer " + ignored_key, custom_litellm_key_header=master_key) + valid_token = await user_api_key_auth( + request=request, api_key="Bearer " + ignored_key, custom_litellm_key_header=master_key + ) assert valid_token.token == hash_token(master_key) @@ -1146,3 +1116,216 @@ async def test_user_api_key_from_query_param(): valid_token = await user_api_key_auth(request=request, api_key="") assert valid_token.token == hash_token(user_key) + +class TestEndUserBudgetWithVirtualKey: + """Tests for end user budget enforcement with virtual key authentication.""" + + @pytest.mark.asyncio + async def test_over_budget_end_user_can_access_zero_cost_model_with_virtual_key(self): + """ + Test that an over-budget end user can still access zero-cost models + when authenticating with a virtual key. + + This verifies that budget checks properly skip for zero-cost models + regardless of authentication method. + """ + import time + from fastapi import Request + from starlette.datastructures import URL + from litellm.proxy._types import ( + LiteLLM_BudgetTable, + LiteLLM_EndUserTable, + UserAPIKeyAuth, + ) + from litellm.proxy.auth.user_api_key_auth import user_api_key_auth + from litellm.proxy.proxy_server import hash_token, user_api_key_cache + from litellm.router import Router + + master_key = "sk-master-1234" + user_key = "sk-virtual-key-1234" + hashed_key = hash_token(user_key) + end_user_id = "over-budget-end-user" + + # Create a valid virtual key token + valid_token = UserAPIKeyAuth( + token=hashed_key, + user_id="test-user", + last_refreshed_at=time.time(), + ) + user_api_key_cache.set_cache(key=hashed_key, value=valid_token) + + # Create an over-budget end user + end_user_budget = LiteLLM_BudgetTable( + budget_id="test-budget", + max_budget=10.0, + ) + over_budget_end_user = LiteLLM_EndUserTable( + user_id=end_user_id, + spend=50.0, # Over budget of 10.0 + litellm_budget_table=end_user_budget, + blocked=False, + ) + + # Create a router with zero-cost model + mock_router = Router( + model_list=[ + { + "model_name": "free-model", + "litellm_params": { + "model": "ollama/llama2", + "api_base": "http://localhost:11434", + }, + "model_info": { + "id": "free-model-id", + "input_cost_per_token": 0.0, + "output_cost_per_token": 0.0, + }, + }, + ] + ) + + # Set up the proxy server state + setattr(litellm.proxy.proxy_server, "user_api_key_cache", user_api_key_cache) + setattr(litellm.proxy.proxy_server, "master_key", master_key) + setattr(litellm.proxy.proxy_server, "prisma_client", "mock-prisma") + setattr(litellm.proxy.proxy_server, "llm_router", mock_router) + + # Mock get_end_user_object to return our over-budget user + with patch( + "litellm.proxy.auth.user_api_key_auth.get_end_user_object", + new_callable=AsyncMock, + return_value=over_budget_end_user, + ): + request = Request( + scope={ + "type": "http", + "headers": [(b"x-litellm-end-user-id", end_user_id.encode())], + } + ) + request._url = URL(url="/v1/chat/completions") + + async def return_body(): + return b'{"model": "free-model", "user": "' + end_user_id.encode() + b'"}' + + request.body = return_body + + # Should NOT raise BudgetExceededError for zero-cost model + result = await user_api_key_auth( + request=request, + api_key="Bearer " + user_key, + ) + + assert result is not None + + # Cleanup + setattr(litellm.proxy.proxy_server, "llm_router", None) + + @pytest.mark.asyncio + async def test_over_budget_end_user_blocked_from_paid_model_with_virtual_key(self): + """ + Test that an over-budget end user is blocked from paid models + when authenticating with a virtual key. + + This verifies the budget enforcement works correctly with virtual keys. + """ + import time + from fastapi import Request + from starlette.datastructures import URL + from litellm.proxy._types import ( + LiteLLM_BudgetTable, + LiteLLM_EndUserTable, + UserAPIKeyAuth, + ) + from litellm.proxy.auth.user_api_key_auth import user_api_key_auth + from litellm.proxy.proxy_server import hash_token, user_api_key_cache + from litellm.router import Router + + master_key = "sk-master-1234" + user_key = "sk-virtual-key-5678" + hashed_key = hash_token(user_key) + end_user_id = "over-budget-end-user-2" + + # Create a valid virtual key token + valid_token = UserAPIKeyAuth( + token=hashed_key, + user_id="test-user-2", + last_refreshed_at=time.time(), + ) + user_api_key_cache.set_cache(key=hashed_key, value=valid_token) + + # Create an over-budget end user + end_user_budget = LiteLLM_BudgetTable( + budget_id="test-budget-2", + max_budget=10.0, + ) + over_budget_end_user = LiteLLM_EndUserTable( + user_id=end_user_id, + spend=50.0, # Over budget of 10.0 + litellm_budget_table=end_user_budget, + blocked=False, + ) + + # Create a router with a paid model + mock_router = Router( + model_list=[ + { + "model_name": "paid-model", + "litellm_params": { + "model": "gpt-3.5-turbo", + "api_key": "sk-test", + }, + "model_info": { + "id": "paid-model-id", + }, + }, + ] + ) + + # Set up the proxy server state + setattr(litellm.proxy.proxy_server, "user_api_key_cache", user_api_key_cache) + setattr(litellm.proxy.proxy_server, "master_key", master_key) + setattr(litellm.proxy.proxy_server, "prisma_client", "mock-prisma") + setattr(litellm.proxy.proxy_server, "llm_router", mock_router) + + # Mock get_end_user_object and get_model_info + with ( + patch( + "litellm.proxy.auth.user_api_key_auth.get_end_user_object", + new_callable=AsyncMock, + return_value=over_budget_end_user, + ), + patch("litellm.get_model_info") as mock_get_model_info, + ): + # Mock paid model cost + mock_get_model_info.return_value = { + "input_cost_per_token": 0.0000015, + "output_cost_per_token": 0.000002, + } + + request = Request( + scope={ + "type": "http", + "headers": [(b"x-litellm-end-user-id", end_user_id.encode())], + } + ) + request._url = URL(url="/v1/chat/completions") + + async def return_body(): + return b'{"model": "paid-model", "user": "' + end_user_id.encode() + b'"}' + + request.body = return_body + + # Should raise ProxyException for paid model (budget exceeded) + from litellm.proxy._types import ProxyException + + with pytest.raises(ProxyException) as exc_info: + await user_api_key_auth( + request=request, + api_key="Bearer " + user_key, + ) + + assert "ExceededBudget" in exc_info.value.message + assert end_user_id in exc_info.value.message + + # Cleanup + setattr(litellm.proxy.proxy_server, "llm_router", None)