From 73c7d0ef2b6d8ff1f462e68a82c2ac0c58065651 Mon Sep 17 00:00:00 2001 From: Ishaan Jaffer Date: Sat, 21 Mar 2026 11:03:56 -0700 Subject: [PATCH] test: assert no blocking set_cache call in _user_api_key_auth_builder --- .../proxy/auth/test_user_api_key_auth.py | 302 +++++++++++++----- 1 file changed, 221 insertions(+), 81 deletions(-) diff --git a/tests/test_litellm/proxy/auth/test_user_api_key_auth.py b/tests/test_litellm/proxy/auth/test_user_api_key_auth.py index f3f0ba56cb9..efe78d2ea1e 100644 --- a/tests/test_litellm/proxy/auth/test_user_api_key_auth.py +++ b/tests/test_litellm/proxy/auth/test_user_api_key_auth.py @@ -72,7 +72,7 @@ def test_get_api_key_with_custom_litellm_key_header( def test_team_metadata_with_tags_flows_through_jwt_auth(): """ Test that team_metadata (specifically tags) flows through JWT authentication. - + This is a regression test for the issue where JWT auth was not populating team_metadata, causing team-level tags to be missing in litellm_pre_call_utils.py """ @@ -87,7 +87,7 @@ def test_team_metadata_with_tags_flows_through_jwt_auth(): rpm_limit=100, models=["gpt-4", "gpt-3.5-turbo"], ) - + # Simulate constructing UserAPIKeyAuth like we do in JWT auth # This is the pattern from user_api_key_auth.py lines 552-587 user_api_key_auth = UserAPIKeyAuth( @@ -100,14 +100,16 @@ def test_team_metadata_with_tags_flows_through_jwt_auth(): user_role="internal_user", user_id="test-user", ) - + # Verify team_metadata is set - assert user_api_key_auth.team_metadata is not None, "team_metadata should be populated" + assert ( + user_api_key_auth.team_metadata is not None + ), "team_metadata should be populated" assert user_api_key_auth.team_metadata == team_object.metadata, ( f"team_metadata not correctly mapped. " f"Expected: {team_object.metadata}, Got: {user_api_key_auth.team_metadata}" ) - + # Specifically verify tags are present assert "tags" in user_api_key_auth.team_metadata, "tags should be in team_metadata" assert user_api_key_auth.team_metadata["tags"] == ["production", "high-priority"], ( @@ -118,7 +120,7 @@ def test_team_metadata_with_tags_flows_through_jwt_auth(): def test_route_checks_is_llm_api_route(): """Test RouteChecks.is_llm_api_route() correctly identifies LLM API routes including passthrough endpoints""" - + # Test OpenAI routes openai_routes = [ "/v1/chat/completions", @@ -142,18 +144,22 @@ def test_route_checks_is_llm_api_route(): "/v1/realtime", "/realtime", ] - + for route in openai_routes: - assert RouteChecks.is_llm_api_route(route=route), f"Route {route} should be identified as LLM API route" + assert RouteChecks.is_llm_api_route( + route=route + ), f"Route {route} should be identified as LLM API route" # Test Anthropic routes anthropic_routes = [ "/v1/messages", "/v1/messages/count_tokens", ] - + for route in anthropic_routes: - assert RouteChecks.is_llm_api_route(route=route), f"Route {route} should be identified as LLM API route" + assert RouteChecks.is_llm_api_route( + route=route + ), f"Route {route} should be identified as LLM API route" # Test passthrough routes (this is the key improvement over the old route checking) passthrough_routes = [ @@ -171,9 +177,11 @@ def test_route_checks_is_llm_api_route(): "/vllm/v1/chat/completions", "/mistral/v1/chat/completions", ] - + for route in passthrough_routes: - assert RouteChecks.is_llm_api_route(route=route), f"Route {route} should be identified as LLM API route" + assert RouteChecks.is_llm_api_route( + route=route + ), f"Route {route} should be identified as LLM API route" # Test MCP routes mcp_routes = [ @@ -181,9 +189,11 @@ def test_route_checks_is_llm_api_route(): "/mcp/", "/mcp/test", ] - + for route in mcp_routes: - assert RouteChecks.is_llm_api_route(route=route), f"Route {route} should be identified as LLM API route" + assert RouteChecks.is_llm_api_route( + route=route + ), f"Route {route} should be identified as LLM API route" # Test LiteLLM native RAG routes rag_routes = [ @@ -193,7 +203,9 @@ def test_route_checks_is_llm_api_route(): "/v1/rag/query", ] for route in rag_routes: - assert RouteChecks.is_llm_api_route(route=route), f"Route {route} should be identified as LLM API route" + assert RouteChecks.is_llm_api_route( + route=route + ), f"Route {route} should be identified as LLM API route" # Test routes with placeholders placeholder_routes = [ @@ -206,9 +218,11 @@ def test_route_checks_is_llm_api_route(): "/v1/batches/batch_123", "/batches/batch_123", ] - + for route in placeholder_routes: - assert RouteChecks.is_llm_api_route(route=route), f"Route {route} should be identified as LLM API route" + assert RouteChecks.is_llm_api_route( + route=route + ), f"Route {route} should be identified as LLM API route" # Test Azure OpenAI routes azure_routes = [ @@ -217,9 +231,11 @@ def test_route_checks_is_llm_api_route(): "/engines/gpt-4/chat/completions", "/engines/gpt-3.5-turbo/completions", ] - + for route in azure_routes: - assert RouteChecks.is_llm_api_route(route=route), f"Route {route} should be identified as LLM API route" + assert RouteChecks.is_llm_api_route( + route=route + ), f"Route {route} should be identified as LLM API route" # Test non-LLM routes (should return False) non_llm_routes = [ @@ -236,9 +252,11 @@ def test_route_checks_is_llm_api_route(): "/debug", "/test", ] - + for route in non_llm_routes: - assert not RouteChecks.is_llm_api_route(route=route), f"Route {route} should NOT be identified as LLM API route" + assert not RouteChecks.is_llm_api_route( + route=route + ), f"Route {route} should NOT be identified as LLM API route" # Test invalid inputs invalid_inputs = [ @@ -248,9 +266,11 @@ def test_route_checks_is_llm_api_route(): {}, "", ] - + for invalid_input in invalid_inputs: - assert not RouteChecks.is_llm_api_route(route=invalid_input), f"Invalid input {invalid_input} should return False" + assert not RouteChecks.is_llm_api_route( + route=invalid_input + ), f"Invalid input {invalid_input} should return False" @pytest.mark.asyncio @@ -259,7 +279,7 @@ async def test_proxy_admin_expired_key_from_cache(): Test that PROXY_ADMIN keys retrieved from cache are checked for expiration before being returned. This prevents expired keys from bypassing expiration checks when retrieved from cache (which normally happens at lines 1014-1036). - + Regression test for issue where PROXY_ADMIN keys from cache skipped expiration check. """ from datetime import datetime, timedelta, timezone @@ -280,39 +300,43 @@ async def test_proxy_admin_expired_key_from_cache(): api_key = "sk-test-proxy-admin-key" hashed_key = hash_token(api_key) expired_time = datetime.now(timezone.utc) - timedelta(hours=1) # Expired 1 hour ago - + expired_token = UserAPIKeyAuth( api_key=api_key, user_role=LitellmUserRoles.PROXY_ADMIN, expires=expired_time, token=hashed_key, ) - + # Mock cache to return the expired token mock_cache = AsyncMock() mock_cache.async_get_cache = AsyncMock(return_value=expired_token) mock_cache.delete_cache = MagicMock() - + # Mock proxy_logging_obj mock_proxy_logging_obj = MagicMock() mock_proxy_logging_obj.internal_usage_cache = MagicMock() mock_proxy_logging_obj.internal_usage_cache.dual_cache = AsyncMock() - mock_proxy_logging_obj.internal_usage_cache.dual_cache.async_delete_cache = AsyncMock() + mock_proxy_logging_obj.internal_usage_cache.dual_cache.async_delete_cache = ( + AsyncMock() + ) # Mock post_call_failure_hook as async function returning None (no transformation) mock_proxy_logging_obj.post_call_failure_hook = AsyncMock(return_value=None) - + # Mock prisma_client mock_prisma_client = MagicMock() - + # Mock get_key_object to return expired token from cache with patch( "litellm.proxy.auth.user_api_key_auth.get_key_object", new_callable=AsyncMock, - ) as mock_get_key_object, \ - patch("litellm.proxy.auth.user_api_key_auth._delete_cache_key_object", new_callable=AsyncMock) as mock_delete_cache: - + ) as mock_get_key_object, patch( + "litellm.proxy.auth.user_api_key_auth._delete_cache_key_object", + new_callable=AsyncMock, + ) as mock_delete_cache: + mock_get_key_object.return_value = expired_token - + # Set attributes on proxy_server module (these are imported inside _user_api_key_auth_builder) import litellm.proxy.proxy_server as _proxy_server_mod @@ -331,14 +355,12 @@ async def test_proxy_admin_expired_key_from_cache(): "litellm_proxy_admin_name": "admin", } _original_values = { - attr: getattr(_proxy_server_mod, attr, None) - for attr in _attrs_to_set + 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) - # Create a mock request request = Request(scope={"type": "http"}) request._url = URL(url="/chat/completions") @@ -358,38 +380,41 @@ async def test_proxy_admin_expired_key_from_cache(): ) # Verify that ProxyException was raised with expired_key type - assert hasattr(exc_info.value, "type"), "Exception should have 'type' attribute" - assert exc_info.value.type == ProxyErrorTypes.expired_key, ( - f"Expected expired_key error type, got {exc_info.value.type}" - ) - assert "Expired Key" in str(exc_info.value.message), ( - f"Exception message should mention 'Expired Key', got: {exc_info.value.message}" - ) + assert hasattr( + exc_info.value, "type" + ), "Exception should have 'type' attribute" + assert ( + exc_info.value.type == ProxyErrorTypes.expired_key + ), f"Expected expired_key error type, got {exc_info.value.type}" + assert "Expired Key" in str( + exc_info.value.message + ), f"Exception message should mention 'Expired Key', got: {exc_info.value.message}" # Verify that the param field does NOT leak the full API key (Issue #18731) # The param should be abbreviated like "sk-...XXXX" not the full plaintext key - assert exc_info.value.param is not None, "Exception should have 'param' attribute" + assert ( + exc_info.value.param is not None + ), "Exception should have 'param' attribute" assert exc_info.value.param != api_key, ( f"SECURITY: Full API key should NOT be in param field! " f"Got: {exc_info.value.param}, Expected abbreviated format like 'sk-...XXXX'" ) - assert exc_info.value.param.startswith("sk-..."), ( - f"Param should be abbreviated to 'sk-...XXXX' format. Got: {exc_info.value.param}" - ) + assert exc_info.value.param.startswith( + "sk-..." + ), f"Param should be abbreviated to 'sk-...XXXX' format. Got: {exc_info.value.param}" # Verify that cache deletion was called mock_delete_cache.assert_called_once() call_args = mock_delete_cache.call_args - assert call_args[1]["hashed_token"] == hashed_key, ( - "Cache deletion should be called with the hashed key" - ) + assert ( + call_args[1]["hashed_token"] == hashed_key + ), "Cache deletion should be called with the hashed key" finally: # Restore all module-level attributes so subsequent tests are not affected for attr, val in _original_values.items(): setattr(_proxy_server_mod, attr, val) - @pytest.mark.asyncio async def test_return_user_api_key_auth_obj_user_spend_and_budget(): """ @@ -400,7 +425,7 @@ async def test_return_user_api_key_auth_obj_user_spend_and_budget(): from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.auth.user_api_key_auth import _return_user_api_key_auth_obj - + user_obj = type( "LiteLLM_UserTable", (), @@ -413,7 +438,7 @@ async def test_return_user_api_key_auth_obj_user_spend_and_budget(): "user_role": "internal_user", }, ) - + api_key = "sk-test-key" valid_token_dict = { "user_id": "test-user", @@ -421,10 +446,10 @@ async def test_return_user_api_key_auth_obj_user_spend_and_budget(): } route = "/chat/completions" start_time = datetime.now() - + mock_service_logger = MagicMock() mock_service_logger.async_service_success_hook = AsyncMock() - + with patch( "litellm.proxy.auth.user_api_key_auth.user_api_key_service_logger_obj", new=mock_service_logger, @@ -438,7 +463,7 @@ async def test_return_user_api_key_auth_obj_user_spend_and_budget(): start_time=start_time, user_role=None, ) - + assert isinstance(result, UserAPIKeyAuth) assert result.user_spend == 250.0 assert result.user_max_budget == 1000.0 @@ -470,9 +495,7 @@ def test_proxy_admin_jwt_auth_includes_identity_fields(): user_role=LitellmUserRoles.PROXY_ADMIN, user_id="user-abc", team_id="team-123", - team_alias=( - team_object.team_alias if team_object is not None else None - ), + team_alias=(team_object.team_alias if team_object is not None else None), team_metadata=team_object.metadata if team_object is not None else None, org_id="org-456", end_user_id="end-user-789", @@ -503,9 +526,7 @@ def test_proxy_admin_jwt_auth_handles_no_team_object(): user_role=LitellmUserRoles.PROXY_ADMIN, user_id="admin-user", team_id=None, - team_alias=( - team_object.team_alias if team_object is not None else None - ), + team_alias=(team_object.team_alias if team_object is not None else None), team_metadata=team_object.metadata if team_object is not None else None, org_id=None, end_user_id=None, @@ -534,7 +555,10 @@ class TestJWTOAuth2Coexistence: def test_is_jwt_detects_jwt_tokens(self): """JWT tokens have 3 dot-separated parts.""" assert JWTHandler.is_jwt("header.payload.signature") is True - assert JWTHandler.is_jwt("eyJhbGciOiJSUzI1NiJ9.eyJzdWIiOiJ1c2VyMSJ9.sig123") is True + assert ( + JWTHandler.is_jwt("eyJhbGciOiJSUzI1NiJ9.eyJzdWIiOiJ1c2VyMSJ9.sig123") + is True + ) def test_is_jwt_rejects_opaque_tokens(self): """Opaque OAuth2 tokens do not have 3 dot-separated parts.""" @@ -567,12 +591,20 @@ class TestJWTOAuth2Coexistence: mock_request.headers = {"authorization": f"Bearer {opaque_token}"} mock_request.query_params = {} - with patch("litellm.proxy.proxy_server.general_settings", general_settings), \ - patch("litellm.proxy.proxy_server.premium_user", True), \ - patch("litellm.proxy.proxy_server.master_key", "sk-master"), \ - patch("litellm.proxy.proxy_server.prisma_client", None), \ - patch("litellm.proxy.auth.user_api_key_auth.Oauth2Handler.check_oauth2_token", new_callable=AsyncMock, return_value=mock_oauth2_response) as mock_oauth2, \ - patch("litellm.proxy.auth.user_api_key_auth.JWTAuthManager.auth_builder", new_callable=AsyncMock) as mock_jwt_auth: + with patch( + "litellm.proxy.proxy_server.general_settings", general_settings + ), patch("litellm.proxy.proxy_server.premium_user", True), patch( + "litellm.proxy.proxy_server.master_key", "sk-master" + ), patch( + "litellm.proxy.proxy_server.prisma_client", None + ), patch( + "litellm.proxy.auth.user_api_key_auth.Oauth2Handler.check_oauth2_token", + new_callable=AsyncMock, + return_value=mock_oauth2_response, + ) as mock_oauth2, patch( + "litellm.proxy.auth.user_api_key_auth.JWTAuthManager.auth_builder", + new_callable=AsyncMock, + ) as mock_jwt_auth: litellm.proxy.proxy_server.jwt_handler.update_environment( prisma_client=None, @@ -624,12 +656,20 @@ class TestJWTOAuth2Coexistence: mock_request.headers = {"authorization": f"Bearer {jwt_token}"} mock_request.query_params = {} - with patch("litellm.proxy.proxy_server.general_settings", general_settings), \ - patch("litellm.proxy.proxy_server.premium_user", True), \ - patch("litellm.proxy.proxy_server.master_key", "sk-master"), \ - patch("litellm.proxy.proxy_server.prisma_client", None), \ - patch("litellm.proxy.auth.user_api_key_auth.Oauth2Handler.check_oauth2_token", new_callable=AsyncMock) as mock_oauth2, \ - patch("litellm.proxy.auth.user_api_key_auth.JWTAuthManager.auth_builder", new_callable=AsyncMock, return_value=mock_jwt_result) as mock_jwt_auth: + with patch( + "litellm.proxy.proxy_server.general_settings", general_settings + ), patch("litellm.proxy.proxy_server.premium_user", True), patch( + "litellm.proxy.proxy_server.master_key", "sk-master" + ), patch( + "litellm.proxy.proxy_server.prisma_client", None + ), patch( + "litellm.proxy.auth.user_api_key_auth.Oauth2Handler.check_oauth2_token", + new_callable=AsyncMock, + ) as mock_oauth2, patch( + "litellm.proxy.auth.user_api_key_auth.JWTAuthManager.auth_builder", + new_callable=AsyncMock, + return_value=mock_jwt_result, + ) as mock_jwt_auth: litellm.proxy.proxy_server.jwt_handler.update_environment( prisma_client=None, @@ -671,11 +711,17 @@ class TestJWTOAuth2Coexistence: mock_request.headers = {"authorization": f"Bearer {jwt_like_token}"} mock_request.query_params = {} - with patch("litellm.proxy.proxy_server.general_settings", general_settings), \ - patch("litellm.proxy.proxy_server.premium_user", True), \ - patch("litellm.proxy.proxy_server.master_key", "sk-master"), \ - patch("litellm.proxy.proxy_server.prisma_client", None), \ - patch("litellm.proxy.auth.user_api_key_auth.Oauth2Handler.check_oauth2_token", new_callable=AsyncMock, return_value=mock_oauth2_response) as mock_oauth2: + with patch( + "litellm.proxy.proxy_server.general_settings", general_settings + ), patch("litellm.proxy.proxy_server.premium_user", True), patch( + "litellm.proxy.proxy_server.master_key", "sk-master" + ), patch( + "litellm.proxy.proxy_server.prisma_client", None + ), patch( + "litellm.proxy.auth.user_api_key_auth.Oauth2Handler.check_oauth2_token", + new_callable=AsyncMock, + return_value=mock_oauth2_response, + ) as mock_oauth2: result = await user_api_key_auth( request=mock_request, @@ -685,3 +731,97 @@ class TestJWTOAuth2Coexistence: # OAuth2 should handle it since JWT auth is disabled mock_oauth2.assert_called_once_with(token=jwt_like_token) assert result.user_id == "oauth2-user" + + +@pytest.mark.asyncio +async def test_user_api_key_auth_builder_no_blocking_set_cache(): + """ + Regression test: _user_api_key_auth_builder must never call the synchronous + DualCache.set_cache() on the hot auth path (blocks the event loop). + It should use async_set_cache() instead. + """ + from datetime import datetime, timezone + + from starlette.datastructures import URL + from starlette.requests import Request + + from litellm.caching.dual_cache import DualCache + from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth + from litellm.proxy.auth.user_api_key_auth import _user_api_key_auth_builder + + api_key = "sk-test-no-blocking-cache" + valid_token = UserAPIKeyAuth( + api_key=api_key, + token=api_key, + user_role=LitellmUserRoles.INTERNAL_USER, + team_id="team-abc", + ) + + mock_cache = AsyncMock(spec=DualCache) + mock_cache.async_get_cache = AsyncMock(return_value=valid_token) + mock_cache.async_set_cache = AsyncMock(return_value=None) + # set_cache is sync — wrap it with a MagicMock so we can assert it's never called + mock_cache.set_cache = MagicMock( + side_effect=AssertionError( + "Blocking DualCache.set_cache() called on async hot path — use async_set_cache() instead" + ) + ) + + mock_proxy_logging_obj = MagicMock() + mock_proxy_logging_obj.internal_usage_cache = MagicMock() + mock_proxy_logging_obj.internal_usage_cache.dual_cache = AsyncMock() + mock_proxy_logging_obj.internal_usage_cache.dual_cache.async_delete_cache = ( + AsyncMock() + ) + mock_proxy_logging_obj.post_call_failure_hook = AsyncMock(return_value=None) + + import litellm.proxy.proxy_server as _proxy_server_mod + + _attrs = { + "prisma_client": MagicMock(), + "user_api_key_cache": mock_cache, + "proxy_logging_obj": mock_proxy_logging_obj, + "master_key": "sk-master-key", + "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", + } + _originals = {k: getattr(_proxy_server_mod, k, None) for k in _attrs} + try: + for k, v in _attrs.items(): + setattr(_proxy_server_mod, k, v) + + 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=valid_token, + ), patch( + "litellm.proxy.auth.user_api_key_auth.get_team_object", + new_callable=AsyncMock, + return_value=None, + ): + # Should complete without raising the AssertionError planted in set_cache + await _user_api_key_auth_builder( + request=request, + api_key=f"Bearer {api_key}", + azure_api_key_header="", + anthropic_api_key_header=None, + google_ai_studio_api_key_header=None, + azure_apim_header=None, + request_data={}, + ) + + # Extra belt-and-suspenders: confirm the sync method was truly never called + mock_cache.set_cache.assert_not_called() + + finally: + for k, v in _originals.items(): + setattr(_proxy_server_mod, k, v)