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.
This commit is contained in:
michelligabriele 2026-03-24 22:58:44 +01:00
parent f9d29e4e4e
commit 4666d9f08a
2 changed files with 204 additions and 0 deletions

View file

@ -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:

View file

@ -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):