This commit is contained in:
Pin Zhu 2026-09-30 03:56:56 -04:00 • committed by GitHub
commit 4b9789d6ad
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
3 changed files with 207 additions and 80 deletions

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__ = ()
@ -3028,6 +3037,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,
@ -3048,21 +3138,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(
@ -3880,50 +3974,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(

View file

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

View file

@ -43,6 +43,7 @@ from litellm.proxy.auth.auth_checks import (
LITELLM_SESSION_TOKEN_PREFIX,
ExperimentalUIJWTToken,
_cache_management_object,
_delete_cache_key_object,
_can_object_call_model,
_can_object_call_vector_stores,
_check_agent_access_group_model_access,
@ -661,6 +662,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