From 9a38646f35b5ade55296e89b1eb13d92af0997f2 Mon Sep 17 00:00:00 2001 From: Pin Zhu <113764596+47Elysia@users.noreply.github.com> Date: Mon, 28 Sep 2026 19:10:06 +0800 Subject: [PATCH 1/3] fix(auth): serialize concurrent key cache loads --- litellm/proxy/auth/auth_checks.py | 155 +++++++++++++++++++++--------- 1 file changed, 110 insertions(+), 45 deletions(-) diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index 12d420141f1..18ee819ee14 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -397,6 +397,15 @@ _TEAM_MEMBERSHIP_INFLIGHT_MAX: Final = 10000 _team_membership_inflight: Final = LimitedSizeOrderedDict(max_size=_TEAM_MEMBERSHIP_INFLIGHT_MAX) +class _KeyObjectLoadReservation: + def __init__(self) -> None: + self.lock: Final = asyncio.Lock() + self.users = 0 + + +_key_object_load_locks: Final[dict[str, _KeyObjectLoadReservation]] = {} + + class _TeamMembershipCacheMiss: __slots__ = () @@ -3018,6 +3027,87 @@ async def _cache_key_object( ) +async def _fetch_and_cache_key_object( + hashed_token: str, + prisma_client: PrismaClient, + user_api_key_cache: UserApiKeyCache, + parent_otel_span: Span | None, + proxy_logging_obj: ProxyLogging | None, +) -> UserAPIKeyAuth | None: + valid_token: Final[BaseModel | None] = await _fetch_key_object_from_db_with_reconnect( + hashed_token=hashed_token, + prisma_client=prisma_client, + parent_otel_span=parent_otel_span, + proxy_logging_obj=proxy_logging_obj, + ) + if valid_token is None: + return None + + response: Final = UserAPIKeyAuth.model_validate(valid_token.model_dump(exclude_none=True)) + if response.object_permission_id and not response.object_permission: + try: + response.object_permission = await get_object_permission( + object_permission_id=response.object_permission_id, + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + parent_otel_span=parent_otel_span, + proxy_logging_obj=proxy_logging_obj, + ) + except Exception as e: + verbose_proxy_logger.debug( + "Failed to load object_permission for key with object_permission_id=%s: %s", + response.object_permission_id, + e, + ) + + await _cache_key_object( + hashed_token=hashed_token, + user_api_key_obj=response, + user_api_key_cache=user_api_key_cache, + proxy_logging_obj=proxy_logging_obj, + ) + return response + + +async def _load_key_object_on_cache_miss( + hashed_token: str, + prisma_client: PrismaClient, + user_api_key_cache: UserApiKeyCache, + parent_otel_span: Span | None, + proxy_logging_obj: ProxyLogging | None, +) -> UserAPIKeyAuth | None: + reservation: Final = _reserve_key_object_load(hashed_token) + try: + async with reservation.lock: + cached: Final = await user_api_key_cache.async_get_cache( + key=hashed_token, + model_type=UserAPIKeyAuth, + ) + if cached is not None: + return _copy_user_api_key_auth_for_cache(user_api_key_obj=cached) + return await _fetch_and_cache_key_object( + hashed_token=hashed_token, + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + parent_otel_span=parent_otel_span, + proxy_logging_obj=proxy_logging_obj, + ) + finally: + _release_key_object_load(hashed_token, reservation) + + +def _reserve_key_object_load(hashed_token: str) -> _KeyObjectLoadReservation: + reservation: Final = _key_object_load_locks.setdefault(hashed_token, _KeyObjectLoadReservation()) + reservation.users += 1 + return reservation + + +def _release_key_object_load(hashed_token: str, reservation: _KeyObjectLoadReservation) -> None: + reservation.users -= 1 + if reservation.users == 0 and _key_object_load_locks.get(hashed_token) is reservation: + _key_object_load_locks.pop(hashed_token, None) + + async def _delete_cache_key_object( hashed_token: str, user_api_key_cache: UserApiKeyCache, @@ -3038,21 +3128,25 @@ async def _delete_cache_key_object( copy's TTL expires. """ key: Final = hashed_token - + reservation: Final = _reserve_key_object_load(key) try: - user_api_key_cache.delete_cache(key=key) + async with reservation.lock: + try: + user_api_key_cache.delete_cache(key=key) - ## UPDATE REDIS CACHE ## - if proxy_logging_obj is not None: - await proxy_logging_obj.internal_usage_cache.dual_cache.async_delete_cache(key=key) - except Exception as e: # noqa: BLE001 # best-effort: a cache error must not fail a committed write - verbose_proxy_logger.warning( - "Failed to invalidate cached key entry %s; a stale key object may be served until its TTL expires: %s", - key, - e, - ) + ## UPDATE REDIS CACHE ## + if proxy_logging_obj is not None: + await proxy_logging_obj.internal_usage_cache.dual_cache.async_delete_cache(key=key) + except Exception as e: # noqa: BLE001 # best-effort: a cache error must not fail a committed write + verbose_proxy_logger.warning( + "Failed to invalidate cached key entry %s; a stale key object may be served until its TTL expires: %s", + key, + e, + ) - await publish_auth_cache_invalidation(cache_key=key) + await publish_auth_cache_invalidation(cache_key=key) + finally: + _release_key_object_load(key, reservation) async def delete_cache_key_objects( @@ -3867,50 +3961,21 @@ async def get_key_object( if check_cache_only: raise Exception(f"Key doesn't exist in cache + check_cache_only=True. key={key}.") - # else, check db - _valid_token: Final[BaseModel | None] = await _fetch_key_object_from_db_with_reconnect( + response: Final = await _load_key_object_on_cache_miss( hashed_token=hashed_token, prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, parent_otel_span=parent_otel_span, proxy_logging_obj=proxy_logging_obj, ) - - if _valid_token is None: + if response is None: raise ProxyException( message=f"Authentication Error, Invalid proxy server token passed. key={hashed_token}, not found in db. Create key via `/key/generate` call.", type=ProxyErrorTypes.token_not_found_in_db, param="key", code=status.HTTP_401_UNAUTHORIZED, ) - - _response: Final = UserAPIKeyAuth.model_validate(_valid_token.model_dump(exclude_none=True)) - - # Load object_permission if object_permission_id exists but object_permission is not loaded - if _response.object_permission_id and not _response.object_permission: - try: - _response.object_permission = await get_object_permission( - object_permission_id=_response.object_permission_id, - prisma_client=prisma_client, - user_api_key_cache=user_api_key_cache, - parent_otel_span=parent_otel_span, - proxy_logging_obj=proxy_logging_obj, - ) - except Exception as e: - verbose_proxy_logger.debug( - "Failed to load object_permission for key with object_permission_id=%s: %s", - _response.object_permission_id, - e, - ) - - # save the key object to cache - await _cache_key_object( - hashed_token=hashed_token, - user_api_key_obj=_response, - user_api_key_cache=user_api_key_cache, - proxy_logging_obj=proxy_logging_obj, - ) - - return _response + return response def _copy_user_api_key_auth_for_cache( From d5b2c586d273905a1f312e8799c5531b7bf3424e Mon Sep 17 00:00:00 2001 From: Pin Zhu <113764596+47Elysia@users.noreply.github.com> Date: Mon, 28 Sep 2026 19:14:53 +0800 Subject: [PATCH 2/3] refactor(auth): share key cache miss loading --- litellm/proxy/auth/resolvers/store.py | 41 ++++----------------------- 1 file changed, 6 insertions(+), 35 deletions(-) diff --git a/litellm/proxy/auth/resolvers/store.py b/litellm/proxy/auth/resolvers/store.py index 832baf03432..7e6fa54016e 100644 --- a/litellm/proxy/auth/resolvers/store.py +++ b/litellm/proxy/auth/resolvers/store.py @@ -3,15 +3,10 @@ from __future__ import annotations from collections.abc import Sequence from typing import TYPE_CHECKING, Final -from pydantic import BaseModel - -from litellm._logging import verbose_proxy_logger from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.auth.auth_checks import ( - _cache_key_object, _copy_user_api_key_auth_for_cache, - _fetch_key_object_from_db_with_reconnect, - get_object_permission, + _load_key_object_on_cache_miss, ) from litellm.proxy.auth.auth_method import AuthMethod from litellm.proxy.auth.network import NetworkContext @@ -34,8 +29,8 @@ from litellm.proxy.auth.resolvers.models import ( from litellm.proxy.auth.roles import TeamRole, map_role, team_role if TYPE_CHECKING: - from litellm.caching.caching import DualCache from litellm.integrations.opentelemetry import Span + from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache from litellm.proxy.utils import PrismaClient, ProxyLogging @@ -59,7 +54,7 @@ class IdentityStore: def __init__( self, prisma_client: PrismaClient | None, - cache: DualCache, + cache: UserApiKeyCache, *, parent_otel_span: Span | None = None, proxy_logging_obj: ProxyLogging | None = None, @@ -110,39 +105,15 @@ class IdentityStore: if self._check_cache_only: raise KeyNotInCacheError(hashed_token) - from_db: Final[BaseModel | None] = await _fetch_key_object_from_db_with_reconnect( + key: Final = await _load_key_object_on_cache_miss( hashed_token=hashed_token, prisma_client=self._prisma, + user_api_key_cache=self._cache, parent_otel_span=self._parent_otel_span, proxy_logging_obj=self._proxy_logging_obj, ) - if from_db is None: + if key is None: raise KeyNotFoundError(hashed_token) - - key: Final = UserAPIKeyAuth.model_validate(from_db.model_dump(exclude_none=True)) - - if key.object_permission_id and not key.object_permission: - try: - key.object_permission = await get_object_permission( - object_permission_id=key.object_permission_id, - prisma_client=self._prisma, - user_api_key_cache=self._cache, - parent_otel_span=self._parent_otel_span, - proxy_logging_obj=self._proxy_logging_obj, - ) - except Exception as e: - verbose_proxy_logger.debug( - "Failed to load object_permission for key with object_permission_id=%s: %s", - key.object_permission_id, - e, - ) - - await _cache_key_object( - hashed_token=hashed_token, - user_api_key_obj=key, - user_api_key_cache=self._cache, - proxy_logging_obj=self._proxy_logging_obj, - ) return key @staticmethod From 8ae022d0c08a4133c07504d8d38ba424dc683708 Mon Sep 17 00:00:00 2001 From: Pin Zhu <113764596+47Elysia@users.noreply.github.com> Date: Mon, 28 Sep 2026 19:23:19 +0800 Subject: [PATCH 3/3] test(auth): cover concurrent key cache misses --- .../proxy/auth/test_auth_checks.py | 91 +++++++++++++++++++ 1 file changed, 91 insertions(+) diff --git a/tests/test_litellm/proxy/auth/test_auth_checks.py b/tests/test_litellm/proxy/auth/test_auth_checks.py index f014e9c26d1..220478423ae 100644 --- a/tests/test_litellm/proxy/auth/test_auth_checks.py +++ b/tests/test_litellm/proxy/auth/test_auth_checks.py @@ -40,6 +40,7 @@ from litellm.types.agents import AgentCaller from litellm.proxy.auth.auth_checks import ( ExperimentalUIJWTToken, _cache_management_object, + _delete_cache_key_object, _can_object_call_model, _can_object_call_vector_stores, _check_agent_access_group_model_access, @@ -611,6 +612,96 @@ async def test_get_key_object_should_raise_if_reconnect_fails_on_db_connection_e assert mock_prisma_client.get_data.await_count == 1 +@pytest.mark.asyncio +async def test_get_key_object_coalesces_parallel_cache_misses(): + started: Final = asyncio.Event() + release: Final = asyncio.Event() + + async def slow_get_data(*args: object, **kwargs: object) -> UserAPIKeyAuth: + started.set() + await release.wait() + return UserAPIKeyAuth(token="hashed-token") + + get_data_mock: Final = AsyncMock(side_effect=slow_get_data) + prisma: Final = MagicMock() + prisma.get_data = get_data_mock + cache: Final = UserApiKeyCache() + + first: Final = asyncio.create_task( + get_key_object( + hashed_token="hashed-token", + prisma_client=prisma, + user_api_key_cache=cache, + ) + ) + await started.wait() + second: Final = asyncio.create_task( + get_key_object( + hashed_token="hashed-token", + prisma_client=prisma, + user_api_key_cache=cache, + ) + ) + await asyncio.sleep(0) + release.set() + + results: Final = await asyncio.gather(first, second) + + assert [result.token for result in results] == ["hashed-token", "hashed-token"] + assert get_data_mock.await_count == 1 + + +@pytest.mark.asyncio +async def test_delete_key_cache_waits_for_inflight_load_before_eviction(): + started: Final = asyncio.Event() + release: Final = asyncio.Event() + rows: Final = iter(("old-model", "new-model")) + + async def get_data(*args: object, **kwargs: object) -> UserAPIKeyAuth: + model: Final = next(rows) + if model == "old-model": + started.set() + await release.wait() + return UserAPIKeyAuth(token="hashed-token", models=[model]) + + get_data_mock: Final = AsyncMock(side_effect=get_data) + prisma: Final = MagicMock() + prisma.get_data = get_data_mock + cache: Final = UserApiKeyCache() + stale_load: Final = asyncio.create_task( + get_key_object( + hashed_token="hashed-token", + prisma_client=prisma, + user_api_key_cache=cache, + ) + ) + await started.wait() + + invalidation: Final = asyncio.create_task( + _delete_cache_key_object( + hashed_token="hashed-token", + user_api_key_cache=cache, + proxy_logging_obj=None, + ) + ) + await asyncio.sleep(0) + assert not invalidation.done() + + release.set() + await invalidation + stale: Final = await stale_load + assert stale.model_dump(mode="json")["models"] == ["old-model"] + assert await cache.async_get_cache(key="hashed-token", model_type=UserAPIKeyAuth) is None + + fresh: Final = await get_key_object( + hashed_token="hashed-token", + prisma_client=prisma, + user_api_key_cache=cache, + ) + assert fresh.model_dump(mode="json")["models"] == ["new-model"] + assert get_data_mock.await_count == 2 + + class _InFlightCountingPrisma: def __init__(self) -> None: self.in_flight = 0