diff --git a/litellm/caching/dual_cache.py b/litellm/caching/dual_cache.py index 996273d558a..4af1edae457 100644 --- a/litellm/caching/dual_cache.py +++ b/litellm/caching/dual_cache.py @@ -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) diff --git a/litellm/caching/redis_batch.py b/litellm/caching/redis_batch.py index f6052192685..fbfe14b5803 100644 --- a/litellm/caching/redis_batch.py +++ b/litellm/caching/redis_batch.py @@ -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) diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index 8fbeaf18460..9a34167ad16 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -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( diff --git a/litellm/proxy/auth/auth_object_prefetch.py b/litellm/proxy/auth/auth_object_prefetch.py index ce55190aa02..14d3e2c07dc 100644 --- a/litellm/proxy/auth/auth_object_prefetch.py +++ b/litellm/proxy/auth/auth_object_prefetch.py @@ -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)) diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index ed4d63fb9fc..ed3ec7b4dde 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -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, diff --git a/litellm/proxy/common_utils/user_api_key_cache.py b/litellm/proxy/common_utils/user_api_key_cache.py index 61d7078ae4c..c99665986dd 100644 --- a/litellm/proxy/common_utils/user_api_key_cache.py +++ b/litellm/proxy/common_utils/user_api_key_cache.py @@ -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``. diff --git a/tests/test_litellm/proxy/auth/test_auth_checks.py b/tests/test_litellm/proxy/auth/test_auth_checks.py index d3cb8e4a645..9c8b95fd7e8 100644 --- a/tests/test_litellm/proxy/auth/test_auth_checks.py +++ b/tests/test_litellm/proxy/auth/test_auth_checks.py @@ -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" ) diff --git a/tests/test_litellm/proxy/auth/test_auth_object_prefetch.py b/tests/test_litellm/proxy/auth/test_auth_object_prefetch.py index 0fd0dda3017..ffac95d6815 100644 --- a/tests/test_litellm/proxy/auth/test_auth_object_prefetch.py +++ b/tests/test_litellm/proxy/auth/test_auth_object_prefetch.py @@ -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 diff --git a/tests/test_litellm/proxy/auth/test_user_api_key_auth.py b/tests/test_litellm/proxy/auth/test_user_api_key_auth.py index 78b281d5c78..d1973e1693b 100644 --- a/tests/test_litellm/proxy/auth/test_user_api_key_auth.py +++ b/tests/test_litellm/proxy/auth/test_user_api_key_auth.py @@ -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(), + ) diff --git a/tests/unit/caching/test_redis_batch.py b/tests/unit/caching/test_redis_batch.py index 9433aeac524..93206efc80f 100644 --- a/tests/unit/caching/test_redis_batch.py +++ b/tests/unit/caching/test_redis_batch.py @@ -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 diff --git a/tests/unit/caching/test_request_redis_batch_pre_call.py b/tests/unit/caching/test_request_redis_batch_pre_call.py index c0834974f26..3569031a3e1 100644 --- a/tests/unit/caching/test_request_redis_batch_pre_call.py +++ b/tests/unit/caching/test_request_redis_batch_pre_call.py @@ -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