From 91d13e9cb5cb3175f704ddb8750786ab75ee5597 Mon Sep 17 00:00:00 2001 From: michelligabriele Date: Tue, 24 Mar 2026 23:47:46 +0100 Subject: [PATCH] test: replace inline-logic unit tests with integration test Replace two unit tests that copied production logic inline with an integration test that exercises _user_api_key_auth_builder for the no-team fallback case. --- .../proxy/auth/test_auth_checks.py | 156 +++++++++++------- 1 file changed, 93 insertions(+), 63 deletions(-) diff --git a/tests/test_litellm/proxy/auth/test_auth_checks.py b/tests/test_litellm/proxy/auth/test_auth_checks.py index ca505cc4fe1..591abcfb56e 100644 --- a/tests/test_litellm/proxy/auth/test_auth_checks.py +++ b/tests/test_litellm/proxy/auth/test_auth_checks.py @@ -427,77 +427,107 @@ async def test_cli_jwt_auth_flow_updates_spend_and_budget(monkeypatch): 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 +@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 + from starlette.requests import Request - # 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 + 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 WITHOUT a team + user_info = LiteLLM_UserTable( + user_id="test_user_no_team", + user_email="test@example.com", + user_role=LitellmUserRoles.INTERNAL_USER.value, + models=[], + max_budget=200.0, ) + token = ExperimentalUIJWTToken.get_cli_jwt_auth_token(user_info) - # 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", + # Mock DB: user with real spend/budget, no team + mock_user_obj = LiteLLM_UserTable( + user_id="test_user_no_team", 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 + mock_cache = MagicMock() + mock_cache.async_get_cache = AsyncMock(return_value=None) + mock_cache.async_set_cache = AsyncMock() + mock_cache.set_cache = MagicMock() - assert key_obj.spend == 10.0 - assert key_obj.max_budget == 200.0 # user budget, not session $0.25 + 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": MagicMock(), + "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=None, + ), 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={}, + ) + + assert result.spend == 10.0, f"Expected spend=10.0, got {result.spend}" + assert ( + result.max_budget == 200.0 + ), f"Expected max_budget=200.0 (user fallback), got {result.max_budget}" + finally: + for attr, val in original_values.items(): + setattr(_proxy_server_mod, attr, val) @pytest.mark.asyncio