mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-02 02:11:58 +00:00
perf(proxy): refresh auth management objects through the request Redis pipeline (#43776)
Identity objects (key, end user) load through the request MGET and their write-backs, the registry reads and the management-object SETs ride the request pipeline. A team refresh invalidates its alias with a pipelined DEL instead of a synchronous DEL plus a duplicate async one, and an MGET miss is remembered so no per-key GET follows it in the same request. Resolves LIT-9012 Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Co-authored-by: yassin <yassin@berri.ai>
This commit is contained in:
parent
cae179e655
commit
13d004fc5a
11 changed files with 404 additions and 34 deletions
|
|
@ -24,7 +24,7 @@ from litellm.types.caching import RedisPipelineIncrementOperation
|
|||
|
||||
from .base_cache import BaseCache
|
||||
from .in_memory_cache import DEFAULT_MAX_SIZE_IN_MEMORY, InMemoryCache
|
||||
from .redis_batch import BatchResult, RedisBatch, active_post_call_redis_batch
|
||||
from .redis_batch import BatchResult, RedisBatch, active_post_call_redis_batch, active_request_redis_batch
|
||||
from .redis_cache import RedisCache, RedisCircuitBreakerOpenError, log_redis_failure
|
||||
|
||||
if TYPE_CHECKING:
|
||||
|
|
@ -279,6 +279,9 @@ class DualCache(BaseCache):
|
|||
result = in_memory_result
|
||||
|
||||
if result is None and self.redis_cache is not None and local_only is False:
|
||||
request_batch: Final = active_request_redis_batch(self.redis_cache)
|
||||
if request_batch is not None and request_batch.read_as_missing(key):
|
||||
return None
|
||||
# If not found in in-memory cache, try fetching from Redis
|
||||
redis_result: Final = await self.redis_cache.async_get_cache(key, parent_otel_span=parent_otel_span)
|
||||
|
||||
|
|
@ -502,12 +505,29 @@ class DualCache(BaseCache):
|
|||
verbose_logger, logging.ERROR, "LiteLLM Cache: exception in async add_cache", e, with_traceback=True
|
||||
)
|
||||
|
||||
async def async_set_cache_pre_call(self, key: str, value: object, ttl: float | None) -> BatchResult[None] | None:
|
||||
"""Memory now, the Redis SET on the request's pipeline, sent with the next read any caller awaits; None
|
||||
when no pipeline is open, so the caller takes its direct path."""
|
||||
batch: Final = None if self.redis_cache is None else active_request_redis_batch(self.redis_cache)
|
||||
return None if batch is None else await self._set_on_batch(batch, key, value, ttl)
|
||||
|
||||
async def async_set_cache_post_call(self, key: str, value: object, ttl: float | None) -> BatchResult[None] | None:
|
||||
"""Memory now, the Redis SET on the request's post-call pipeline; None when no pipeline is open, so the
|
||||
caller takes its direct path."""
|
||||
batch: Final = None if self.redis_cache is None else active_post_call_redis_batch(self.redis_cache)
|
||||
return None if batch is None else await self._set_on_batch(batch, key, value, ttl)
|
||||
|
||||
async def async_delete_cache_pre_call(self, key: str) -> BatchResult[None] | None:
|
||||
"""Memory now, the Redis DEL on the request's pipeline; None when no pipeline is open, so the caller
|
||||
takes its direct path."""
|
||||
batch: Final = None if self.redis_cache is None else active_request_redis_batch(self.redis_cache)
|
||||
if batch is None:
|
||||
return None
|
||||
if self.in_memory_cache is not None:
|
||||
self.in_memory_cache.delete_cache(key)
|
||||
return batch.delete(key)
|
||||
|
||||
async def _set_on_batch(self, batch: RedisBatch, key: str, value: object, ttl: float | None) -> BatchResult[None]:
|
||||
effective_ttl: Final = self.default_in_memory_ttl if ttl is None else ttl
|
||||
if self.in_memory_cache is not None:
|
||||
await self.in_memory_cache.async_set_cache(key, value, ttl=effective_ttl)
|
||||
|
|
|
|||
|
|
@ -47,6 +47,7 @@ class _RedisPipeline(Protocol):
|
|||
def incrbyfloat(self, name: str, amount: float) -> object: ...
|
||||
def expire(self, name: str, time: timedelta) -> object: ...
|
||||
def set(self, name: str, value: str, ex: timedelta | None = None) -> object: ...
|
||||
def delete(self, *names: str) -> object: ...
|
||||
async def execute(self, raise_on_error: bool = True) -> list[object]: ...
|
||||
|
||||
|
||||
|
|
@ -233,6 +234,27 @@ class _Set(_Op[None]):
|
|||
await self._redis_cache.async_set_cache_pipeline_with_ttls(((self._key, self._value, self._ttl),))
|
||||
|
||||
|
||||
class _Delete(_Op[None]):
|
||||
"""DEL of one key, the pipelined twin of ``async_delete_cache``."""
|
||||
|
||||
__slots__ = ("_key", "_redis_cache")
|
||||
|
||||
def __init__(self, redis_cache: RedisCache, key: str) -> None:
|
||||
super().__init__()
|
||||
self._redis_cache: Final = redis_cache
|
||||
self._key: Final = key
|
||||
|
||||
def enqueue(self, pipe: _RedisPipeline) -> int:
|
||||
pipe.delete(self._redis_cache.check_and_fix_namespace(key=self._key))
|
||||
return 1
|
||||
|
||||
def resolve(self, replies: Sequence[object]) -> None:
|
||||
return None
|
||||
|
||||
async def run_alone(self) -> None:
|
||||
await self._redis_cache.async_delete_cache(self._key)
|
||||
|
||||
|
||||
class BatchResult(Generic[_T]):
|
||||
"""Awaitable handle for one declared operation; awaiting it flushes the batch it belongs to."""
|
||||
|
||||
|
|
@ -269,10 +291,23 @@ class RedisBatch:
|
|||
_pending: list[_Op[object]] = field(default_factory=list) # mutable-ok: drained by flush
|
||||
_flush_hooks: list[Callable[[], None]] = field(default_factory=list) # mutable-ok: append-only registry
|
||||
_lock: asyncio.Lock = field(default_factory=asyncio.Lock)
|
||||
_misses: set[str] = field(default_factory=set) # mutable-ok: keys an MGET of this request read as absent
|
||||
flushes: int = 0
|
||||
|
||||
def mget(self, keys: Sequence[str]) -> BatchResult[Mapping[str, object]]:
|
||||
return self._declare(_MGet(self.redis_cache, keys))
|
||||
op: Final = _MGet(self.redis_cache, keys)
|
||||
op.future.add_done_callback(self._note_misses)
|
||||
return self._declare(op)
|
||||
|
||||
def _note_misses(self, future: asyncio.Future[Mapping[str, object]]) -> None:
|
||||
if future.cancelled() or future.exception() is not None:
|
||||
return
|
||||
self._misses.update(key for key, value in future.result().items() if value is None)
|
||||
|
||||
def read_as_missing(self, key: str) -> bool:
|
||||
"""True when an MGET on this batch already found no value under ``key`` and nothing has set it since,
|
||||
so a per-key GET later in the same request can be answered without another round trip."""
|
||||
return key in self._misses
|
||||
|
||||
def script(
|
||||
self, source: str, run: RegisteredScript, keys: Sequence[str], args: Sequence[_ScriptArg]
|
||||
|
|
@ -283,8 +318,13 @@ class RedisBatch:
|
|||
return self._declare(_Increment(self.redis_cache, key, value, ttl))
|
||||
|
||||
def set(self, key: str, value: object, ttl: float | None = None) -> BatchResult[None]:
|
||||
self._misses.discard(key)
|
||||
return self._declare(_Set(self.redis_cache, key, value, ttl))
|
||||
|
||||
def delete(self, key: str) -> BatchResult[None]:
|
||||
self._misses.add(key)
|
||||
return self._declare(_Delete(self.redis_cache, key))
|
||||
|
||||
def add_flush_hook(self, hook: Callable[[], None]) -> None:
|
||||
"""Called at the start of every flush so lazily bound readers can declare their keys into the same trip."""
|
||||
self._flush_hooks.append(hook)
|
||||
|
|
|
|||
|
|
@ -23,7 +23,7 @@ from typing_extensions import NotRequired, ReadOnly, Required, TypedDict, Unpack
|
|||
|
||||
import litellm
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.caching.dual_cache import LimitedSizeOrderedDict
|
||||
from litellm.caching.dual_cache import DualCache, LimitedSizeOrderedDict
|
||||
from litellm.constants import (
|
||||
CLI_JWT_EXPIRATION_HOURS,
|
||||
CLI_SESSION_KEY_PREFIX,
|
||||
|
|
@ -1784,11 +1784,12 @@ async def _load_bounded_registry(
|
|||
if not isinstance(cached, _RegistryNotCached):
|
||||
return cached
|
||||
|
||||
waited_for_another_load: Final = load_lock.locked()
|
||||
async with load_lock:
|
||||
# The request that held the lock has since cached an answer for everyone waiting on it.
|
||||
cached_after_wait: Final = await _cached_registry(cache_key, overflow_sentinel, user_api_key_cache)
|
||||
if not isinstance(cached_after_wait, _RegistryNotCached):
|
||||
return cached_after_wait
|
||||
if waited_for_another_load:
|
||||
cached_after_wait: Final = await _cached_registry(cache_key, overflow_sentinel, user_api_key_cache)
|
||||
if not isinstance(cached_after_wait, _RegistryNotCached):
|
||||
return cached_after_wait
|
||||
|
||||
return await _fetch_and_cache_registry(
|
||||
cache_key=cache_key,
|
||||
|
|
@ -2782,17 +2783,12 @@ async def _cache_team_object(
|
|||
team_table.last_refreshed_at = time.time()
|
||||
|
||||
key: Final = f"team_id:{team_id}"
|
||||
usage_cache: Final = None if proxy_logging_obj is None else proxy_logging_obj.internal_usage_cache.dual_cache
|
||||
# On a shared Redis the write below replaces the team entry and the alias DEL below removes the alias entry
|
||||
# for both caches, so the usage cache only has its own memory to clear.
|
||||
redis_shared: Final = usage_cache is not None and usage_cache.redis_cache is user_api_key_cache.redis_cache
|
||||
|
||||
if proxy_logging_obj is not None:
|
||||
try:
|
||||
await proxy_logging_obj.internal_usage_cache.dual_cache.async_delete_cache(key=key)
|
||||
except Exception as e: # noqa: BLE001 # best-effort invalidation: any cache backend error must not fail the write
|
||||
verbose_proxy_logger.warning(
|
||||
"Failed to invalidate internal usage cache entry %s; "
|
||||
"a stale team object may be served until its TTL expires: %s",
|
||||
key,
|
||||
e,
|
||||
)
|
||||
await _invalidate_usage_cache_entry(usage_cache, key, redis_shared=redis_shared, stale="team object")
|
||||
|
||||
# team_id is the table primary key — guaranteed unique, safe to write.
|
||||
await _cache_management_object(
|
||||
|
|
@ -2819,9 +2815,11 @@ async def _cache_team_object(
|
|||
if team_table.team_alias:
|
||||
alias_key: Final = f"team_alias:{team_table.team_alias}"
|
||||
try:
|
||||
user_api_key_cache.delete_cache(key=alias_key)
|
||||
if proxy_logging_obj is not None:
|
||||
await proxy_logging_obj.internal_usage_cache.dual_cache.async_delete_cache(key=alias_key)
|
||||
pipelined_delete: Final = await user_api_key_cache.async_delete_cache_pre_call(alias_key)
|
||||
if pipelined_delete is None:
|
||||
await user_api_key_cache.async_delete_cache(key=alias_key)
|
||||
else:
|
||||
await pipelined_delete
|
||||
except Exception as e: # noqa: BLE001 # best-effort invalidation: any cache backend error must not fail the mutation
|
||||
verbose_proxy_logger.warning(
|
||||
"Failed to invalidate cached team alias entry %s; "
|
||||
|
|
@ -2829,6 +2827,30 @@ async def _cache_team_object(
|
|||
alias_key,
|
||||
e,
|
||||
)
|
||||
await _invalidate_usage_cache_entry(usage_cache, alias_key, redis_shared=redis_shared, stale="team alias")
|
||||
|
||||
|
||||
async def _invalidate_usage_cache_entry(
|
||||
usage_cache: DualCache | None,
|
||||
key: str,
|
||||
*,
|
||||
redis_shared: bool,
|
||||
stale: str,
|
||||
) -> None:
|
||||
if usage_cache is None:
|
||||
return
|
||||
try:
|
||||
if redis_shared:
|
||||
usage_cache.in_memory_cache.delete_cache(key)
|
||||
else:
|
||||
await usage_cache.async_delete_cache(key=key)
|
||||
except Exception as e: # noqa: BLE001 # best-effort invalidation: any cache backend error must not fail the write
|
||||
verbose_proxy_logger.warning(
|
||||
"Failed to invalidate internal usage cache entry %s; a stale %s may be served until its TTL expires: %s",
|
||||
key.replace("\r", "").replace("\n", ""),
|
||||
stale,
|
||||
e,
|
||||
)
|
||||
|
||||
|
||||
async def invalidate_team_member_spend_state(
|
||||
|
|
|
|||
|
|
@ -15,7 +15,7 @@ from pydantic import BaseModel, TypeAdapter, ValidationError
|
|||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.caching.redis_batch import active_request_redis_batch
|
||||
from litellm.caching.redis_cache import RedisCache
|
||||
from litellm.constants import DEFAULT_IN_MEMORY_TTL
|
||||
from litellm.constants import DEFAULT_IN_MEMORY_TTL, REGISTRY_ERROR_NEGATIVE_CACHE_TTL
|
||||
from litellm.models.organization import LiteLLM_OrganizationTable
|
||||
from litellm.models.team import LiteLLM_TeamTableCachedObj
|
||||
from litellm.models.team_membership import LiteLLM_TeamMembership
|
||||
|
|
@ -324,3 +324,35 @@ async def prefetch_auth_objects(
|
|||
await _fill_from_db(refs, _missing_in_memory(missing, memory), user_api_key_cache, prisma_client)
|
||||
except Exception as e: # noqa: BLE001 # warm-up only; the getters enforce and fail closed on their own
|
||||
verbose_proxy_logger.warning("auth prefetch skipped, falling back to per-object lookups: %s", e)
|
||||
|
||||
|
||||
def _identity_memory_ttl(value: object, management_ttl: float) -> float:
|
||||
"""A registry stored as a string is a sentinel, written with the shorter of the two registry TTLs."""
|
||||
return min(REGISTRY_ERROR_NEGATIVE_CACHE_TTL, management_ttl) if isinstance(value, str) else management_ttl
|
||||
|
||||
|
||||
async def prefetch_identity_keys(cache_keys: Sequence[str], user_api_key_cache: UserApiKeyCache) -> None:
|
||||
"""Warm the entries auth reads before it knows the key's owners (the key object, the end user and the two
|
||||
registries) in one MGET on the request pipeline. Keys the MGET finds absent stay noted on the pipeline, so the
|
||||
per-key getters that follow go to the database without a GET of their own. Best effort, like the
|
||||
owner prefetch: the getters read and enforce on their own."""
|
||||
try:
|
||||
redis_cache: Final = user_api_key_cache.redis_cache
|
||||
if redis_cache is None:
|
||||
return
|
||||
missing: Final = tuple(
|
||||
key
|
||||
for key in dict.fromkeys(cache_keys)
|
||||
if user_api_key_cache.in_memory_cache_for(key).get_cache(key=key) is None
|
||||
)
|
||||
if not missing:
|
||||
return
|
||||
found: Final = _RowValues.validate_python(await _read_redis_rows(sorted(missing), redis_cache))
|
||||
management_ttl: Final = get_management_object_ttl(user_api_key_cache)
|
||||
except Exception as e: # noqa: BLE001 # warm-up only; the getters read Redis and the database on their own
|
||||
verbose_proxy_logger.warning("auth identity prefetch skipped, falling back to per-key lookups: %s", e)
|
||||
return
|
||||
for key, value in ((key, found.get(key)) for key in missing):
|
||||
if value is not None:
|
||||
memory: _InMemoryCache = user_api_key_cache.in_memory_cache_for(key)
|
||||
_set_in_memory(memory, key, value, _identity_memory_ttl(value, management_ttl))
|
||||
|
|
|
|||
|
|
@ -73,7 +73,7 @@ from litellm.proxy.auth.auth_checks import (
|
|||
)
|
||||
from litellm.proxy.auth.auth_exception_handler import UserAPIKeyAuthExceptionHandler
|
||||
from litellm.proxy.auth.auth_method import AuthMethod
|
||||
from litellm.proxy.auth.auth_object_prefetch import AuthObjectRefs, prefetch_auth_objects
|
||||
from litellm.proxy.auth.auth_object_prefetch import AuthObjectRefs, prefetch_auth_objects, prefetch_identity_keys
|
||||
from litellm.proxy.auth.auth_utils import (
|
||||
abbreviate_api_key,
|
||||
get_end_user_id_from_request_body,
|
||||
|
|
@ -120,6 +120,9 @@ from litellm.proxy.common_utils.model_listing_utils import claude_code_requested
|
|||
from litellm.proxy.common_utils.realtime_utils import _realtime_request_body
|
||||
from litellm.proxy.common_utils.user_api_key_cache import (
|
||||
UserApiKeyCache,
|
||||
end_user_cache_key,
|
||||
end_user_restricted_registry_cache_key,
|
||||
model_access_group_registry_cache_key,
|
||||
team_membership_auth_cache_key,
|
||||
)
|
||||
from litellm.proxy.db.db_lookup_gate import bounded_db_lookup
|
||||
|
|
@ -1892,6 +1895,11 @@ async def _user_api_key_auth_builder(
|
|||
proxy_logging_obj=proxy_logging_obj,
|
||||
route=route,
|
||||
)
|
||||
if prisma_client is not None:
|
||||
await prefetch_identity_keys(
|
||||
_identity_cache_keys(api_key, end_user_id=end_user_id, key_is_resolved=valid_token is not None),
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
)
|
||||
if end_user_id:
|
||||
try:
|
||||
end_user_params["end_user_id"] = end_user_id
|
||||
|
|
@ -3248,6 +3256,21 @@ def _spend_counter_redis_cache() -> RedisCache | None:
|
|||
return spend_counter_cache.redis_cache
|
||||
|
||||
|
||||
def _identity_cache_keys(api_key: str, *, end_user_id: str | None, key_is_resolved: bool) -> tuple[str, ...]:
|
||||
"""Cache keys auth reads before it knows the key's owners, all known from the request alone. A key object is
|
||||
cached under the hash of the bearer, so the bearer itself never reaches Redis."""
|
||||
return tuple(
|
||||
key
|
||||
for key in (
|
||||
None if key_is_resolved else hash_token(api_key),
|
||||
None if not end_user_id else end_user_cache_key(end_user_id),
|
||||
None if not end_user_id else end_user_restricted_registry_cache_key(),
|
||||
model_access_group_registry_cache_key(),
|
||||
)
|
||||
if key is not None
|
||||
)
|
||||
|
||||
|
||||
async def _prefetch_referenced_auth_objects(
|
||||
valid_token: UserAPIKeyAuth,
|
||||
end_user_id: str | None,
|
||||
|
|
|
|||
|
|
@ -17,6 +17,8 @@ from litellm.proxy.common_utils.cache_pydantic_utils import CacheCodec
|
|||
if TYPE_CHECKING:
|
||||
from opentelemetry.trace import Span
|
||||
|
||||
from litellm.caching.redis_batch import BatchResult
|
||||
|
||||
T = TypeVar("T", bound=BaseModel)
|
||||
|
||||
_HASHED_TOKEN_CACHE_KEY: Final = re.compile(r"[0-9a-f]{64}")
|
||||
|
|
@ -27,6 +29,9 @@ def is_user_key_cache_key(key: str) -> bool:
|
|||
return _HASHED_TOKEN_CACHE_KEY.fullmatch(key) is not None
|
||||
|
||||
|
||||
_PIPELINED_SET_OPTIONS: Final = frozenset(("ttl",))
|
||||
|
||||
|
||||
class UserApiKeyCache(DualCache):
|
||||
"""
|
||||
DualCache wrapper for UserAPIKeyAuth-like payloads.
|
||||
|
|
@ -208,10 +213,23 @@ class UserApiKeyCache(DualCache):
|
|||
return super().set_cache(key=key, value=payload, local_only=local_only, **kwargs)
|
||||
|
||||
async def async_set_cache(self, key: str | None, value: object, local_only: bool = False, **kwargs: object):
|
||||
"""Inside a request the Redis SET rides the request's pipeline (memory is written at once); anywhere
|
||||
else, or with options the pipeline does not carry, it goes to Redis directly as before."""
|
||||
model_type: Final = cast(type[BaseModel] | None, kwargs.pop("model_type", None))
|
||||
payload: Final[object] = CacheCodec.serialize(value, model_type=model_type)
|
||||
ttl: Final = kwargs.get("ttl")
|
||||
pipelined: Final = (
|
||||
key is not None
|
||||
and not local_only
|
||||
and kwargs.keys() <= _PIPELINED_SET_OPTIONS
|
||||
and (ttl is None or isinstance(ttl, (int, float)))
|
||||
)
|
||||
if key is not None and is_user_key_cache_key(key):
|
||||
if pipelined and await self.key_object_cache.async_set_cache_pre_call(key, payload, ttl) is not None:
|
||||
return None
|
||||
return await self.key_object_cache.async_set_cache(key=key, value=payload, local_only=local_only, **kwargs)
|
||||
if pipelined and await super().async_set_cache_pre_call(key, payload, ttl) is not None:
|
||||
return None
|
||||
return await super().async_set_cache(key=key, value=payload, local_only=local_only, **kwargs)
|
||||
|
||||
def delete_cache(self, key: str) -> None:
|
||||
|
|
@ -226,6 +244,11 @@ class UserApiKeyCache(DualCache):
|
|||
return
|
||||
await super().async_delete_cache(key)
|
||||
|
||||
async def async_delete_cache_pre_call(self, key: str) -> BatchResult[None] | None:
|
||||
if is_user_key_cache_key(key):
|
||||
return await self.key_object_cache.async_delete_cache_pre_call(key)
|
||||
return await super().async_delete_cache_pre_call(key)
|
||||
|
||||
async def async_delete_cache_keys(self, keys: Sequence[str]) -> None:
|
||||
"""Batch twin of ``async_delete_cache``, partitioned like
|
||||
``async_set_cache_pipeline``.
|
||||
|
|
|
|||
|
|
@ -5621,7 +5621,8 @@ async def test_cache_team_object_writes_team_id_and_invalidates_team_alias():
|
|||
team_table = LiteLLM_TeamTableCachedObj(**base_team_row)
|
||||
cache = MagicMock()
|
||||
cache.async_set_cache = AsyncMock()
|
||||
cache.delete_cache = MagicMock()
|
||||
cache.async_delete_cache = AsyncMock()
|
||||
cache.async_delete_cache_pre_call = AsyncMock(return_value=None) # no request pipeline open
|
||||
logging_obj = MagicMock()
|
||||
logging_obj.internal_usage_cache.dual_cache.async_delete_cache = AsyncMock()
|
||||
|
||||
|
|
@ -5642,9 +5643,9 @@ async def test_cache_team_object_writes_team_id_and_invalidates_team_alias():
|
|||
written_value = cache.async_set_cache.await_args.kwargs.get("value") or cache.async_set_cache.await_args.args[1]
|
||||
assert written_value is team_table
|
||||
|
||||
# (2) team_alias-keyed entry is deleted in BOTH the in-memory cache
|
||||
# and the Redis dual cache (mirrors _delete_cache_key_object pattern).
|
||||
cache.delete_cache.assert_called_once_with(key="team_alias:H-Capacity")
|
||||
# (2) team_alias-keyed entry is deleted in BOTH the in-memory cache and the Redis dual cache, on the
|
||||
# async path: a Redis DEL must never run synchronously on the event loop.
|
||||
cache.async_delete_cache.assert_awaited_once_with(key="team_alias:H-Capacity")
|
||||
|
||||
# (4) internal usage cache: team_id entry deleted BEFORE the fresh
|
||||
# write, alias entry deleted as before.
|
||||
|
|
@ -5658,7 +5659,7 @@ async def test_cache_team_object_writes_team_id_and_invalidates_team_alias():
|
|||
aliasless = LiteLLM_TeamTableCachedObj(**{**base_team_row, "team_alias": None})
|
||||
cache2 = MagicMock()
|
||||
cache2.async_set_cache = AsyncMock()
|
||||
cache2.delete_cache = MagicMock()
|
||||
cache2.async_delete_cache = AsyncMock()
|
||||
logging_obj2 = MagicMock()
|
||||
logging_obj2.internal_usage_cache.dual_cache.async_delete_cache = AsyncMock()
|
||||
|
||||
|
|
@ -5669,7 +5670,7 @@ async def test_cache_team_object_writes_team_id_and_invalidates_team_alias():
|
|||
proxy_logging_obj=logging_obj2,
|
||||
)
|
||||
|
||||
cache2.delete_cache.assert_not_called()
|
||||
cache2.async_delete_cache.assert_not_awaited()
|
||||
logging_obj2.internal_usage_cache.dual_cache.async_delete_cache.assert_awaited_once_with(
|
||||
key="team_id:team-no-alias"
|
||||
)
|
||||
|
|
|
|||
|
|
@ -18,13 +18,18 @@ from litellm.proxy._types import (
|
|||
UserAPIKeyAuth,
|
||||
)
|
||||
from litellm.proxy.auth.auth_checks import (
|
||||
get_end_user_object,
|
||||
get_org_object,
|
||||
get_team_membership,
|
||||
get_team_object,
|
||||
get_user_object,
|
||||
)
|
||||
from litellm.proxy.auth.auth_object_prefetch import AuthObjectRefs, prefetch_auth_objects
|
||||
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
|
||||
from litellm.proxy.auth.auth_object_prefetch import AuthObjectRefs, prefetch_auth_objects, prefetch_identity_keys
|
||||
from litellm.proxy.common_utils.user_api_key_cache import (
|
||||
UserApiKeyCache,
|
||||
end_user_cache_key,
|
||||
end_user_restricted_registry_cache_key,
|
||||
)
|
||||
|
||||
USER_ID = "prefetch-user"
|
||||
TEAM_ID = "prefetch-team"
|
||||
|
|
@ -336,3 +341,29 @@ async def test_no_redis_goes_straight_to_one_query():
|
|||
|
||||
assert prisma.db.query_first.await_count == 1
|
||||
assert cache.in_memory_cache.get_cache(f"team_membership:{USER_ID}:{TEAM_ID}") is not None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_identity_prefetch_warms_the_end_user_so_its_getter_needs_neither_redis_nor_the_database():
|
||||
end_user_key = end_user_cache_key("eu-1")
|
||||
redis = CountingRedis({end_user_key: json.dumps({"user_id": "eu-1", "blocked": False, "spend": 0.0})})
|
||||
cache = _cache(redis)
|
||||
prisma = _prisma()
|
||||
|
||||
await prefetch_identity_keys([end_user_key, end_user_restricted_registry_cache_key()], cache)
|
||||
end_user = await get_end_user_object(end_user_id="eu-1", prisma_client=prisma, user_api_key_cache=cache)
|
||||
|
||||
assert end_user is not None and end_user.user_id == "eu-1"
|
||||
assert redis.commands == [f"MGET {end_user_key} {end_user_restricted_registry_cache_key()}"]
|
||||
assert prisma.db.mock_calls == []
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_identity_prefetch_does_not_cache_an_absent_entry_as_present():
|
||||
redis = CountingRedis({})
|
||||
cache = _cache(redis)
|
||||
|
||||
await prefetch_identity_keys([end_user_cache_key("eu-absent")], cache)
|
||||
|
||||
assert redis.round_trips == 1
|
||||
assert cache.in_memory_cache.get_cache(end_user_cache_key("eu-absent")) is None
|
||||
|
|
|
|||
|
|
@ -9354,3 +9354,30 @@ async def test_admission_and_budget_reservation_read_the_key_spend_counter_with_
|
|||
], "admission and reservation share one snapshot, and read-then-write callers go to Redis once it closes"
|
||||
assert redis.async_batch_get_cache.await_count == 1
|
||||
assert "spend:key:hashed" in redis.async_batch_get_cache.await_args.kwargs["key_list"]
|
||||
|
||||
|
||||
def test_identity_prefetch_keys_match_what_auth_reads_for_the_request():
|
||||
from litellm.proxy.auth.user_api_key_auth import _identity_cache_keys
|
||||
from litellm.proxy.common_utils.user_api_key_cache import (
|
||||
end_user_cache_key,
|
||||
end_user_restricted_registry_cache_key,
|
||||
model_access_group_registry_cache_key,
|
||||
)
|
||||
from litellm.proxy.utils import hash_token
|
||||
|
||||
assert _identity_cache_keys("sk-1234", end_user_id="eu-1", key_is_resolved=False) == (
|
||||
hash_token("sk-1234"),
|
||||
end_user_cache_key("eu-1"),
|
||||
end_user_restricted_registry_cache_key(),
|
||||
model_access_group_registry_cache_key(),
|
||||
)
|
||||
assert _identity_cache_keys("a" * 64, end_user_id=None, key_is_resolved=False) == (
|
||||
hash_token("a" * 64),
|
||||
model_access_group_registry_cache_key(),
|
||||
)
|
||||
master_key_keys = _identity_cache_keys("my-master-key", end_user_id=None, key_is_resolved=False)
|
||||
assert master_key_keys == (hash_token("my-master-key"), model_access_group_registry_cache_key())
|
||||
assert "my-master-key" not in master_key_keys, "a bearer that is not an sk- key must not be sent to Redis as is"
|
||||
assert _identity_cache_keys("sk-1234", end_user_id=None, key_is_resolved=True) == (
|
||||
model_access_group_registry_cache_key(),
|
||||
)
|
||||
|
|
|
|||
|
|
@ -58,6 +58,10 @@ class FakePipeline:
|
|||
self.commands.append(("SET", name, value, None if ex is None else int(ex.total_seconds())))
|
||||
return self
|
||||
|
||||
def delete(self, *names: str) -> FakePipeline:
|
||||
self.commands.append(("DEL", *names))
|
||||
return self
|
||||
|
||||
async def execute(self, raise_on_error: bool = True) -> list[Any]:
|
||||
assert raise_on_error is False
|
||||
self.executed = True
|
||||
|
|
@ -101,6 +105,14 @@ class FakeRedisCache(RedisCache):
|
|||
self.store[key] = float(self.store.get(key, 0.0)) + value
|
||||
return self.store[key]
|
||||
|
||||
async def async_set_cache(self, key: str, value: object, **kwargs: object) -> None: # pyright: ignore[reportIncompatibleMethodOverride] # fake, no server
|
||||
self.alone.append(("SET", key, value))
|
||||
self.store[key] = value
|
||||
|
||||
async def async_delete_cache(self, key: str) -> None: # pyright: ignore[reportIncompatibleMethodOverride] # records the direct delete
|
||||
self.alone.append(("DEL", key))
|
||||
self.store.pop(key, None)
|
||||
|
||||
async def async_set_cache_pipeline_with_ttls(self, cache_list: Sequence[tuple[str, object, float | None]]) -> None:
|
||||
self.alone.append(("SET_PIPELINE", tuple(cache_list)))
|
||||
for key, value, _ttl in cache_list:
|
||||
|
|
@ -124,6 +136,8 @@ def replies(command: tuple[Any, ...]) -> Any:
|
|||
return 1
|
||||
case "SET":
|
||||
return True
|
||||
case "DEL":
|
||||
return 1
|
||||
raise AssertionError(command)
|
||||
|
||||
|
||||
|
|
@ -297,3 +311,51 @@ def test_request_scope_hands_out_one_batch_per_backend_and_nests() -> None:
|
|||
assert active_request_redis_batch(cache_a) is first
|
||||
assert len(batches.batches) == 2
|
||||
assert active_request_redis_batch(cache_a) is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_a_key_an_mget_read_as_absent_stays_known_missing_until_something_sets_it() -> None:
|
||||
cache, client = make()
|
||||
batch = RedisBatch(cache)
|
||||
values = await batch.mget(["a-hit", "b-miss"])
|
||||
assert values == {"a-hit": {"k": "a-hit"}, "b-miss": None}
|
||||
assert batch.read_as_missing("b-miss") is True
|
||||
assert batch.read_as_missing("a-hit") is False
|
||||
assert batch.read_as_missing("never-read") is False
|
||||
batch.set("b-miss", "now-present")
|
||||
assert batch.read_as_missing("b-miss") is False
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_a_delete_rides_the_pipeline_under_the_namespace_and_reads_as_missing_afterwards() -> None:
|
||||
cache, client = make(namespace="ns")
|
||||
batch = RedisBatch(cache)
|
||||
gone = batch.delete("team_alias:x")
|
||||
got = batch.mget(["a-hit"])
|
||||
assert await gone is None
|
||||
assert await got == {"a-hit": {"k": "ns:a-hit"}}
|
||||
assert len(client.pipelines) == 1
|
||||
assert client.pipelines[0].commands[0] == ("DEL", "ns:team_alias:x")
|
||||
assert batch.read_as_missing("team_alias:x") is True
|
||||
assert cache.alone == []
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_a_delete_on_a_cluster_cache_runs_as_its_own_del() -> None:
|
||||
client = FakeClient(replies)
|
||||
cache = FakeClusterCache(client)
|
||||
cache.store["team_alias:x"] = "stale"
|
||||
batch = RedisBatch(cache)
|
||||
assert await batch.delete("team_alias:x") is None
|
||||
assert cache.alone == [("DEL", "team_alias:x")]
|
||||
assert "team_alias:x" not in cache.store
|
||||
assert client.pipelines == []
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_a_failed_mget_marks_nothing_as_missing() -> None:
|
||||
cache, client = make(fail=ConnectionError("down"))
|
||||
batch = RedisBatch(cache)
|
||||
with pytest.raises(ConnectionError):
|
||||
await batch.mget(["b-miss"])
|
||||
assert batch.read_as_missing("b-miss") is False
|
||||
|
|
|
|||
|
|
@ -7,15 +7,16 @@ import asyncio
|
|||
import hashlib
|
||||
import json
|
||||
from typing import Any, Final
|
||||
from unittest.mock import AsyncMock
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
import pytest
|
||||
|
||||
from litellm import Router
|
||||
from litellm.caching.dual_cache import DualCache
|
||||
from litellm.caching.redis_batch import active_request_redis_batches, request_redis_batch_scope
|
||||
from litellm.proxy._types import LiteLLM_UserTable
|
||||
from litellm.proxy.auth.auth_object_prefetch import _CacheEntry, _write_back
|
||||
from litellm.proxy._types import LiteLLM_TeamTableCachedObj, LiteLLM_UserTable
|
||||
from litellm.proxy.auth.auth_checks import _cache_team_object
|
||||
from litellm.proxy.auth.auth_object_prefetch import _CacheEntry, _write_back, prefetch_identity_keys
|
||||
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
|
||||
from litellm.proxy.hooks.parallel_request_limiter_v3 import (
|
||||
CHECK_AND_INCREMENT_BY_N_SCRIPT,
|
||||
|
|
@ -27,7 +28,7 @@ from litellm.proxy.utils import InternalUsageCache
|
|||
from litellm.router_utils.cooldown_cache import CooldownCache
|
||||
from litellm.router_utils.routing_read_batch import RoutingPrefetch
|
||||
|
||||
from .test_redis_batch import FakeClient, FakeRedisCache
|
||||
from .test_redis_batch import FakeClient, FakeRedisCache, replies
|
||||
|
||||
_MODEL_GROUP = "claude"
|
||||
_FAR_FUTURE = 4_102_444_800.0 # 2100-01-01, a cooldown stamped then is still active
|
||||
|
|
@ -528,3 +529,91 @@ async def test_auth_write_back_outside_a_scope_writes_through_as_before():
|
|||
assert [(op[0], [(key, ttl) for key, _value, ttl in op[1]]) for op in redis_cache.alone] == [
|
||||
("SET_PIPELINE", [("user-1", 42)])
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_a_key_the_request_mget_read_as_absent_is_not_read_again_by_a_per_key_get():
|
||||
client = FakeClient(_lua_ok_replies)
|
||||
redis_cache = FakeRedisCache(client)
|
||||
cache = UserApiKeyCache(redis_cache=redis_cache)
|
||||
with request_redis_batch_scope() as request:
|
||||
assert await request.batch(redis_cache).mget(["absent-key"]) == {"absent-key": None}
|
||||
assert await cache.async_get_cache("absent-key") is None
|
||||
assert redis_cache.alone == [] and len(client.pipelines) == 1
|
||||
await cache.async_set_cache("absent-key", {"v": 1}, ttl=5)
|
||||
await request.flush_all()
|
||||
assert [c[:2] for c in client.pipelines[1].commands] == [("SET", "absent-key")]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_management_object_writes_inside_a_request_ride_its_pipeline_and_write_through_outside():
|
||||
client = FakeClient(_lua_ok_replies)
|
||||
redis_cache = FakeRedisCache(client)
|
||||
cache = UserApiKeyCache(redis_cache=redis_cache)
|
||||
with request_redis_batch_scope() as request:
|
||||
await cache.async_set_cache("team_id:t1", {"team_id": "t1"}, ttl=60)
|
||||
await cache.async_set_cache("hashed-key-object", {"token": "hashed-key-object"}, ttl=60)
|
||||
assert client.pipelines == []
|
||||
assert cache.in_memory_cache.get_cache("team_id:t1") == {"team_id": "t1"}
|
||||
assert await cache.async_get_cache("hashed-key-object") == {"token": "hashed-key-object"}
|
||||
await request.flush_all()
|
||||
assert sorted((c[0], c[1], c[3]) for c in client.pipelines[0].commands) == [
|
||||
("SET", "hashed-key-object", 60),
|
||||
("SET", "team_id:t1", 60),
|
||||
]
|
||||
await cache.async_set_cache("team_id:t2", {"team_id": "t2"}, ttl=60)
|
||||
assert len(client.pipelines) == 1
|
||||
assert redis_cache.alone == [("SET", "team_id:t2", {"team_id": "t2"})]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_a_team_refresh_inside_a_request_sends_its_set_and_alias_del_in_one_pipeline_before_returning():
|
||||
client = FakeClient(_lua_ok_replies)
|
||||
redis_cache = FakeRedisCache(client)
|
||||
cache = UserApiKeyCache(redis_cache=redis_cache)
|
||||
usage_cache = DualCache(redis_cache=redis_cache)
|
||||
usage_cache.in_memory_cache.set_cache("team_id:t1", "stale team")
|
||||
usage_cache.in_memory_cache.set_cache("team_alias:alpha", "stale alias")
|
||||
cache.in_memory_cache.set_cache("team_alias:alpha", "stale alias")
|
||||
proxy_logging_obj = MagicMock()
|
||||
proxy_logging_obj.internal_usage_cache = InternalUsageCache(dual_cache=usage_cache)
|
||||
team = LiteLLM_TeamTableCachedObj(team_id="t1", team_alias="alpha")
|
||||
with request_redis_batch_scope() as request:
|
||||
await _cache_team_object("t1", team, cache, proxy_logging_obj)
|
||||
assert [c[:2] for c in client.pipelines[0].commands] == [("SET", "team_id:t1"), ("DEL", "team_alias:alpha")], (
|
||||
"the alias DEL must reach Redis before the refresh returns, or another request can refill memory from it"
|
||||
)
|
||||
assert redis_cache.alone == []
|
||||
assert usage_cache.in_memory_cache.get_cache("team_id:t1") is None
|
||||
assert usage_cache.in_memory_cache.get_cache("team_alias:alpha") is None
|
||||
assert cache.in_memory_cache.get_cache("team_alias:alpha") is None
|
||||
assert cache.in_memory_cache.get_cache("team_id:t1")["team_id"] == "t1"
|
||||
await request.flush_all()
|
||||
assert len(client.pipelines) == 1 and redis_cache.alone == []
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_a_pipelined_management_write_without_a_ttl_expires_in_redis_like_the_direct_path():
|
||||
client = FakeClient(_lua_ok_replies)
|
||||
redis_cache = FakeRedisCache(client)
|
||||
cache = UserApiKeyCache(redis_cache=redis_cache)
|
||||
cache.update_cache_ttl(default_in_memory_ttl=5, default_redis_ttl=None)
|
||||
with request_redis_batch_scope() as request:
|
||||
await cache.async_set_cache("team_id:t1", {"team_id": "t1"})
|
||||
await request.flush_all()
|
||||
assert [(c[0], c[1], c[3]) for c in client.pipelines[0].commands] == [("SET", "team_id:t1", 5)]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_identity_prefetch_is_one_mget_after_which_hits_and_misses_alike_cost_no_read():
|
||||
client = FakeClient(replies)
|
||||
redis_cache = FakeRedisCache(client)
|
||||
cache = UserApiKeyCache(redis_cache=redis_cache)
|
||||
with request_redis_batch_scope():
|
||||
await prefetch_identity_keys(["key-hit", "end_user_id:eu-miss", "key-hit"], cache)
|
||||
assert [c[0] for c in client.pipelines[0].commands] == ["MGET"]
|
||||
assert sorted(client.pipelines[0].commands[0][1:]) == ["end_user_id:eu-miss", "key-hit"]
|
||||
assert await cache.async_get_cache("key-hit") == {"k": "key-hit"}
|
||||
assert await cache.async_get_cache("end_user_id:eu-miss") is None
|
||||
assert len(client.pipelines) == 1 and redis_cache.alone == []
|
||||
assert cache.in_memory_cache.get_cache("end_user_id:eu-miss") is None
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue