mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-09 22:31:41 +00:00
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:
parent
f9d29e4e4e
commit
4666d9f08a
2 changed files with 204 additions and 0 deletions
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue