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] 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(