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:
devin-ai-integration[bot] 2026-09-29 17:56:31 -07:00 • committed by GitHub
parent cae179e655
commit 13d004fc5a
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
11 changed files with 404 additions and 34 deletions

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

@ -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(),
)

View file

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

View file

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