From 4666d9f08ae3a059f7bd1ae6e002264a2ee1a0d0 Mon Sep 17 00:00:00 2001 From: michelligabriele Date: Tue, 24 Mar 2026 22:58:44 +0100 Subject: [PATCH] fix(auth): update CLI JWT token spend/budget from DB during auth CLI SSO JWT tokens carry frozen spend/budget from the encrypted blob. Overwrite with real DB values after fetching user_obj and team_obj so response headers report accurate data. --- litellm/proxy/auth/user_api_key_auth.py | 17 ++ .../proxy/auth/test_auth_checks.py | 187 ++++++++++++++++++ 2 files changed, 204 insertions(+) diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index eba787c63b3..e41f926180d 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -22,6 +22,7 @@ from litellm._logging import verbose_logger, verbose_proxy_logger from litellm._service_logger import ServiceLogging from litellm.caching import DualCache from litellm.litellm_core_utils.dd_tracing import tracer +from litellm.constants import CLI_JWT_TOKEN_NAME from litellm.litellm_core_utils.dot_notation_indexing import get_nested_value from litellm.proxy._types import * from litellm.proxy.auth.auth_checks import ( @@ -1421,6 +1422,22 @@ async def _user_api_key_auth_builder( # noqa: PLR0915 key=valid_token.team_id, value=_team_obj ) # save team table in cache - used for tpm/rpm limiting - tpm_rpm_limiter.py + # CLI JWT tokens carry frozen spend/budget from the encrypted blob. + # Replace with real DB values so response headers are accurate. + if valid_token.token == CLI_JWT_TOKEN_NAME: + if user_obj is not None and user_obj.spend is not None: + valid_token.spend = user_obj.spend + if ( + _team_obj is not None + and _team_obj.max_budget is not None + ): + valid_token.max_budget = _team_obj.max_budget + elif ( + user_obj is not None + and user_obj.max_budget is not None + ): + valid_token.max_budget = user_obj.max_budget + # Fetch project object if key belongs to a project _project_obj = None if valid_token.project_id is not None: diff --git a/tests/test_litellm/proxy/auth/test_auth_checks.py b/tests/test_litellm/proxy/auth/test_auth_checks.py index 69188fd200e..ca505cc4fe1 100644 --- a/tests/test_litellm/proxy/auth/test_auth_checks.py +++ b/tests/test_litellm/proxy/auth/test_auth_checks.py @@ -312,6 +312,193 @@ def test_get_cli_jwt_auth_token_custom_expiration( assert expires <= get_utc_datetime() + timedelta(hours=48, minutes=1) +@pytest.mark.asyncio +async def test_cli_jwt_auth_flow_updates_spend_and_budget(monkeypatch): + """Integration test: CLI JWT tokens get real spend/budget through the actual auth builder.""" + from starlette.datastructures import URL + from starlette.requests import Request + + from litellm.proxy._types import LiteLLM_TeamTableCachedObj + from litellm.proxy.auth.user_api_key_auth import _user_api_key_auth_builder + + import litellm.proxy.proxy_server as _proxy_server_mod + + monkeypatch.setenv("EXPERIMENTAL_UI_LOGIN", "True") + + # Create a CLI JWT token for an internal_user with a team + user_info = LiteLLM_UserTable( + user_id="test_user_spend", + user_email="test@example.com", + user_role=LitellmUserRoles.INTERNAL_USER.value, + models=[], + max_budget=100.0, + teams=["test_team_budget"], + ) + token = ExperimentalUIJWTToken.get_cli_jwt_auth_token(user_info) + + # Mock DB objects with real spend/budget + mock_user_obj = LiteLLM_UserTable( + user_id="test_user_spend", + spend=42.50, + max_budget=100.0, + ) + mock_team_obj = LiteLLM_TeamTableCachedObj( + team_id="test_team_budget", + max_budget=1000.0, + spend=150.0, + last_refreshed_at=None, + ) + + # Set up module globals + mock_cache = MagicMock() + mock_cache.async_get_cache = AsyncMock(return_value=None) + mock_cache.async_set_cache = AsyncMock() + mock_cache.set_cache = MagicMock() + + mock_prisma = MagicMock() + + mock_proxy_logging = MagicMock() + mock_proxy_logging.internal_usage_cache = MagicMock() + mock_proxy_logging.internal_usage_cache.dual_cache = MagicMock() + mock_proxy_logging.internal_usage_cache.dual_cache.async_get_cache = AsyncMock( + return_value=None + ) + mock_proxy_logging.budget_alerts = AsyncMock() + mock_proxy_logging.post_call_failure_hook = AsyncMock(return_value=None) + + attrs_to_set = { + "prisma_client": mock_prisma, + "user_api_key_cache": mock_cache, + "proxy_logging_obj": mock_proxy_logging, + "master_key": "sk-master-key-test", + "general_settings": {}, + "llm_model_list": [], + "llm_router": None, + "open_telemetry_logger": None, + "model_max_budget_limiter": MagicMock(), + "user_custom_auth": None, + "jwt_handler": None, + "litellm_proxy_admin_name": "admin", + } + original_values = { + attr: getattr(_proxy_server_mod, attr, None) for attr in attrs_to_set + } + try: + for attr, val in attrs_to_set.items(): + setattr(_proxy_server_mod, attr, val) + + request = Request(scope={"type": "http"}) + request._url = URL(url="/chat/completions") + + with patch( + "litellm.proxy.auth.user_api_key_auth.get_key_object", + new_callable=AsyncMock, + return_value=None, + ), patch( + "litellm.proxy.auth.user_api_key_auth.get_user_object", + new_callable=AsyncMock, + return_value=mock_user_obj, + ), patch( + "litellm.proxy.auth.user_api_key_auth.get_team_object", + new_callable=AsyncMock, + return_value=mock_team_obj, + ), patch( + "litellm.proxy.auth.user_api_key_auth.common_checks", + new_callable=AsyncMock, + return_value=True, + ): + result = await _user_api_key_auth_builder( + request=request, + api_key=f"Bearer {token}", + azure_api_key_header="", + anthropic_api_key_header=None, + google_ai_studio_api_key_header=None, + azure_apim_header=None, + request_data={}, + ) + + # The returned UserAPIKeyAuth should have real DB values + assert result.spend == 42.50, f"Expected spend=42.50, got {result.spend}" + assert ( + result.max_budget == 1000.0 + ), f"Expected max_budget=1000.0, got {result.max_budget}" + finally: + for attr, val in original_values.items(): + setattr(_proxy_server_mod, attr, val) + + +def test_cli_jwt_token_spend_and_budget_from_db(valid_sso_user_defined_values): + """Test that CLI JWT tokens get real spend/budget from DB objects, not frozen blob values.""" + from litellm.constants import CLI_JWT_TOKEN_NAME + + # Generate a CLI JWT token (will have spend=0.0, max_budget=0.25) + token = ExperimentalUIJWTToken.get_cli_jwt_auth_token( + valid_sso_user_defined_values + ) + + # Decrypt to get the UserAPIKeyAuth object + key_obj = ExperimentalUIJWTToken.get_key_object_from_ui_hash_key(token) + assert key_obj is not None + assert key_obj.token == CLI_JWT_TOKEN_NAME + assert key_obj.spend == 0.0 # frozen in blob + assert key_obj.max_budget == litellm.max_ui_session_budget # $0.25 + + # Simulate what the auth flow does after fetching user/team from DB: + # Create mock DB objects with real spend/budget + user_obj = LiteLLM_UserTable( + user_id="test_user", + spend=42.50, + max_budget=100.0, + ) + team_obj = LiteLLM_TeamTable( + team_id="test_team", + max_budget=1000.0, + spend=150.0, + ) + + # Apply the fix logic (same as in user_api_key_auth.py) + if key_obj.token == CLI_JWT_TOKEN_NAME: + if user_obj is not None and user_obj.spend is not None: + key_obj.spend = user_obj.spend + if team_obj is not None and team_obj.max_budget is not None: + key_obj.max_budget = team_obj.max_budget + + # Verify corrected values + assert key_obj.spend == 42.50 # real user spend, not 0.0 + assert key_obj.max_budget == 1000.0 # real team budget, not $0.25 + + +def test_cli_jwt_token_fallback_to_user_budget_when_no_team( + valid_sso_user_defined_values, +): + """Test that CLI JWT falls back to user budget when no team object is available.""" + from litellm.constants import CLI_JWT_TOKEN_NAME + + token = ExperimentalUIJWTToken.get_cli_jwt_auth_token( + valid_sso_user_defined_values + ) + key_obj = ExperimentalUIJWTToken.get_key_object_from_ui_hash_key(token) + assert key_obj is not None + + user_obj = LiteLLM_UserTable( + user_id="test_user", + spend=10.0, + max_budget=200.0, + ) + _team_obj = None + + # Apply the fix logic + if key_obj.token == CLI_JWT_TOKEN_NAME: + if user_obj is not None and user_obj.spend is not None: + key_obj.spend = user_obj.spend + if _team_obj is not None and _team_obj.max_budget is not None: + key_obj.max_budget = _team_obj.max_budget + elif user_obj is not None and user_obj.max_budget is not None: + key_obj.max_budget = user_obj.max_budget + + assert key_obj.spend == 10.0 + assert key_obj.max_budget == 200.0 # user budget, not session $0.25 + @pytest.mark.asyncio async def test_default_internal_user_params_with_get_user_object(monkeypatch):