address review: skip budget checks for CLI JWT, fix test hygiene

- Skip token-level budget checks for CLI JWT tokens (they have no real
  key-level budget; enforcement happens at user/team level in common_checks)
- Add comments explaining max_budget semantics for CLI JWT
- Use proper ASGI scope instead of private request._url assignment
- Drain fire-and-forget asyncio tasks before restoring mocks
- Apply Black formatting
This commit is contained in:
michelligabriele 2026-03-26 21:09:42 +01:00
parent 91d13e9cb5
commit 1fde546df7
2 changed files with 68 additions and 32 deletions

View file

@ -922,9 +922,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,
@ -1258,6 +1258,12 @@ async def _user_api_key_auth_builder( # noqa: PLR0915
# Check 2a. Check if model has zero cost - if so, skip all budget checks
model = get_model_from_request(request_data, route)
skip_budget_checks = False
# CLI JWT tokens have no real key-level budget (the $0.25 is a
# session default). Token-level checks are meaningless; real
# enforcement happens at user/team level in common_checks.
if valid_token.token == CLI_JWT_TOKEN_NAME:
skip_budget_checks = True
if model is not None and llm_router is not None:
from litellm.proxy.auth.auth_checks import _is_model_cost_zero
@ -1424,18 +1430,16 @@ async def _user_api_key_auth_builder( # noqa: PLR0915
# CLI JWT tokens carry frozen spend/budget from the encrypted blob.
# Replace with real DB values so response headers are accurate.
# max_budget is set to the team budget (or user budget as fallback)
# because that is the effective constraint for CLI JWT users.
# Token-level budget checks are skipped above; real enforcement
# uses _team_obj / user_obj directly in common_checks.
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
):
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
):
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
@ -1518,9 +1522,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

View file

@ -59,9 +59,9 @@ def reset_constants_module():
# Reload modules before test
importlib.reload(constants)
importlib.reload(auth_checks)
yield
# Reload modules after test to clean up
importlib.reload(constants)
importlib.reload(auth_checks)
@ -154,9 +154,9 @@ def test_experimental_ui_token_ignores_litellm_ui_session_duration(
expires = datetime.fromisoformat(token_data["expires"].replace("Z", "+00:00"))
now = get_utc_datetime()
# Must be ~10 min, NOT 24h. If LITELLM_UI_SESSION_DURATION were incorrectly used, this would fail.
assert expires <= now + timedelta(minutes=11), (
"Experimental UI must use 10-min expiry, not LITELLM_UI_SESSION_DURATION"
)
assert expires <= now + timedelta(
minutes=11
), "Experimental UI must use 10-min expiry, not LITELLM_UI_SESSION_DURATION"
def test_get_experimental_ui_login_jwt_auth_token_invalid(
@ -290,13 +290,15 @@ def test_get_cli_jwt_auth_token_custom_expiration(
# Set custom expiration to 48 hours
monkeypatch.setenv("LITELLM_CLI_JWT_EXPIRATION_HOURS", "48")
# Reload the constants module to pick up the new env var
importlib.reload(constants)
# Also reload auth_checks to pick up the new constant value
importlib.reload(auth_checks)
token = auth_checks.ExperimentalUIJWTToken.get_cli_jwt_auth_token(valid_sso_user_defined_values)
token = auth_checks.ExperimentalUIJWTToken.get_cli_jwt_auth_token(
valid_sso_user_defined_values
)
# Decrypt and verify token contents
decrypted_token = decrypt_value_helper(
@ -315,7 +317,8 @@ def test_get_cli_jwt_auth_token_custom_expiration(
@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
import asyncio
from starlette.requests import Request
from litellm.proxy._types import LiteLLM_TeamTableCachedObj
@ -387,8 +390,15 @@ async def test_cli_jwt_auth_flow_updates_spend_and_budget(monkeypatch):
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")
request = Request(
scope={
"type": "http",
"method": "POST",
"path": "/chat/completions",
"query_string": b"",
"headers": [],
}
)
with patch(
"litellm.proxy.auth.user_api_key_auth.get_key_object",
@ -422,6 +432,9 @@ async def test_cli_jwt_auth_flow_updates_spend_and_budget(monkeypatch):
assert (
result.max_budget == 1000.0
), f"Expected max_budget=1000.0, got {result.max_budget}"
# Drain fire-and-forget tasks spawned by _user_api_key_auth_builder
await asyncio.sleep(0)
finally:
for attr, val in original_values.items():
setattr(_proxy_server_mod, attr, val)
@ -430,7 +443,8 @@ async def test_cli_jwt_auth_flow_updates_spend_and_budget(monkeypatch):
@pytest.mark.asyncio
async def test_cli_jwt_auth_flow_fallback_to_user_budget(monkeypatch):
"""Integration test: CLI JWT falls back to user budget when no team is assigned."""
from starlette.datastructures import URL
import asyncio
from starlette.requests import Request
from litellm.proxy.auth.user_api_key_auth import _user_api_key_auth_builder
@ -491,8 +505,15 @@ async def test_cli_jwt_auth_flow_fallback_to_user_budget(monkeypatch):
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")
request = Request(
scope={
"type": "http",
"method": "POST",
"path": "/chat/completions",
"query_string": b"",
"headers": [],
}
)
with patch(
"litellm.proxy.auth.user_api_key_auth.get_key_object",
@ -525,6 +546,9 @@ async def test_cli_jwt_auth_flow_fallback_to_user_budget(monkeypatch):
assert (
result.max_budget == 200.0
), f"Expected max_budget=200.0 (user fallback), got {result.max_budget}"
# Drain fire-and-forget tasks spawned by _user_api_key_auth_builder
await asyncio.sleep(0)
finally:
for attr, val in original_values.items():
setattr(_proxy_server_mod, attr, val)
@ -650,7 +674,9 @@ async def test_get_user_object_upsert_includes_user_email():
mock_prisma_client.db.litellm_usertable.create.assert_called_once()
creation_args = mock_prisma_client.db.litellm_usertable.create.call_args[1]["data"]
assert "user_email" in creation_args, "user_email should be included when upserting a new user"
assert (
"user_email" in creation_args
), "user_email should be included when upserting a new user"
assert creation_args["user_email"] == "test@example.com"
assert creation_args["user_id"] == "new_test_user"
@ -677,7 +703,9 @@ def test_log_budget_lookup_failure_skips_user_not_found():
@pytest.mark.asyncio
@patch("litellm.proxy.management_endpoints.team_endpoints.new_team", new_callable=AsyncMock)
@patch(
"litellm.proxy.management_endpoints.team_endpoints.new_team", new_callable=AsyncMock
)
async def test_get_team_db_check_calls_new_team_on_upsert(mock_new_team, monkeypatch):
"""
Test that _get_team_db_check correctly calls the `new_team` function
@ -711,8 +739,12 @@ async def test_get_team_db_check_calls_new_team_on_upsert(mock_new_team, monkeyp
@pytest.mark.asyncio
@patch("litellm.proxy.management_endpoints.team_endpoints.new_team", new_callable=AsyncMock)
async def test_get_team_db_check_does_not_call_new_team_if_exists(mock_new_team, monkeypatch):
@patch(
"litellm.proxy.management_endpoints.team_endpoints.new_team", new_callable=AsyncMock
)
async def test_get_team_db_check_does_not_call_new_team_if_exists(
mock_new_team, monkeypatch
):
"""
Test that _get_team_db_check does NOT call the `new_team` function
if the team already exists in the database.