mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
fix(auth): serialize concurrent key cache loads
This commit is contained in:
parent
90e4962c81
commit
9a38646f35
1 changed files with 110 additions and 45 deletions
|
|
@ -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(
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue