fix(auth): serialize concurrent key cache loads

This commit is contained in:
Pin Zhu 2026-09-28 19:10:06 +08:00 • committed by GitHub
parent 90e4962c81
commit 9a38646f35
No known key found for this signature in database
GPG key ID: B5690EEEBB952194

View file

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