diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index 526a47b1b98..d28e1209f4e 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -602,6 +602,15 @@ async def check_api_key_for_custom_headers_or_pass_through_endpoints( return api_key +# Cache sentinel written when a JWT under AUTO_REGISTER resolved to a proxy +# admin via auth_builder. Proxy admins don't need a mapped virtual key (they +# have full access via auth_builder anyway), but without a cache entry every +# subsequent request from the same JWT identity would re-query the DB for a +# non-existent mapping. Sentinel tells _resolve_jwt_to_virtual_key to skip +# the lookup and return None (caller proceeds to auth_builder). +_JWT_PROXY_ADMIN_SENTINEL = "__JWT_PROXY_ADMIN__" + + class _PendingAutoRegister(NamedTuple): """ Signal returned by ``_resolve_jwt_to_virtual_key`` when the JWT's claim is @@ -808,6 +817,12 @@ async def _resolve_jwt_to_virtual_key( cache_key = f"jwt_key_mapping:{virtual_key_claim_field}:{claim_value}" cached_mapping = await user_api_key_cache.async_get_cache(cache_key) + if cached_mapping == _JWT_PROXY_ADMIN_SENTINEL: + # Previously resolved to a proxy admin via auth_builder; skip the + # mapping lookup and let the caller re-run auth_builder. Avoids a + # repeated DB hit on every proxy-admin request under AUTO_REGISTER. + return None + if cached_mapping == "__NO_MAPPING__": behavior = jwt_handler.litellm_jwtauth.unregistered_jwt_client_behavior if behavior == UnregisteredJWTClientBehavior.REJECT: @@ -1159,6 +1174,19 @@ async def _user_api_key_auth_builder( # noqa: PLR0915 jwt_claims = result.get("jwt_claims", None) if is_proxy_admin: + # Proxy admins authenticate via auth_builder (full + # access), not via a mapped virtual key. If + # AUTO_REGISTER was pending, cache a sentinel so + # future requests from this JWT identity skip the + # DB mapping lookup in _resolve_jwt_to_virtual_key. + # Without this, every proxy-admin request under + # AUTO_REGISTER re-hits get_jwt_key_mapping_object. + if pending_auto_register is not None: + await user_api_key_cache.async_set_cache( + key=pending_auto_register.cache_key, + value=_JWT_PROXY_ADMIN_SENTINEL, + ttl=jwt_handler.litellm_jwtauth.virtual_key_mapping_cache_ttl, + ) return UserAPIKeyAuth( api_key=None, user_role=LitellmUserRoles.PROXY_ADMIN, diff --git a/tests/proxy_unit_tests/test_jwt_key_mapping.py b/tests/proxy_unit_tests/test_jwt_key_mapping.py index 15c3834d1a6..4ad7095a9aa 100644 --- a/tests/proxy_unit_tests/test_jwt_key_mapping.py +++ b/tests/proxy_unit_tests/test_jwt_key_mapping.py @@ -27,7 +27,6 @@ from litellm.proxy.management_endpoints.jwt_key_mapping_endpoints import ( from litellm.caching.caching import DualCache from fastapi import HTTPException - # ────────────────────────────────────────────── # Tests: _resolve_jwt_to_virtual_key # ────────────────────────────────────────────── @@ -1174,6 +1173,51 @@ async def test_auto_register_raises_503_when_winner_mapping_vanishes(): assert "concurrently removed" in exc_info.value.detail +@pytest.mark.asyncio +async def test_proxy_admin_sentinel_skips_db_lookup_on_cache_hit(): + """ + When the cache holds the proxy-admin sentinel (written after a prior + request's is_proxy_admin early-return), _resolve_jwt_to_virtual_key must + return None *without* hitting the DB. Caller proceeds to auth_builder. + + Without this, every subsequent proxy-admin request under AUTO_REGISTER + would re-query get_jwt_key_mapping_object — a cache-miss regression + introduced by the deferred-auto-register refactor. + """ + from litellm.proxy._types import UnregisteredJWTClientBehavior + + jwt_handler = JWTHandler() + jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth( + virtual_key_claim_field="sub", + unregistered_jwt_client_behavior=UnregisteredJWTClientBehavior.AUTO_REGISTER, + virtual_key_mapping_cache_ttl=300, + ) + jwt_claims = {"sub": "admin-user"} + + prisma_client = MagicMock() + # Will fail the test if accessed — proves the sentinel short-circuits DB + prisma_client.db.litellm_jwtkeymapping.find_first = AsyncMock( + side_effect=AssertionError("DB must not be hit when sentinel is cached") + ) + + user_api_key_cache = DualCache() + await user_api_key_cache.async_set_cache( + "jwt_key_mapping:sub:admin-user", "__JWT_PROXY_ADMIN__" + ) + + result = await _resolve_jwt_to_virtual_key( + jwt_claims=jwt_claims, + jwt_handler=jwt_handler, + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + parent_otel_span=None, + proxy_logging_obj=None, + ) + + assert result is None + prisma_client.db.litellm_jwtkeymapping.find_first.assert_not_called() + + # ────────────────────────────────────────────── # Tests: AUTO_REGISTER stamps validated identity from auth_builder # ──────────────────────────────────────────────