mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
Merge 8ae022d0c0 into 04fa760bf2
This commit is contained in:
commit
4b9789d6ad
3 changed files with 207 additions and 80 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__ = ()
|
||||
|
||||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue