test: assert no blocking set_cache call in _user_api_key_auth_builder

This commit is contained in:
Ishaan Jaffer 2026-03-21 11:03:56 -07:00
parent b3036c66d2
commit 73c7d0ef2b

View file

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