From 049c1abef5d56851636bab0475a0afc839987c0f Mon Sep 17 00:00:00 2001 From: Tin Chi Lo Date: Thu, 30 Jul 2026 11:16:38 -0700 Subject: [PATCH] refactor(cache_warming): read a session's record and touched set in one call Fetching the touched models separately added a sequential Redis round trip per live session on every tick, doubling them against the max_sessions cap for data that lives in the same slot. One script now returns both, and get_record is expressed in terms of it so the warm-aware pick keeps its single call too. The eligibility universe also counts a record's served_model as touched, which it is by definition. That matters on upgrade: records captured before the touched set existed would otherwise warm nothing until their next turn rewrote them. Deletes the package __init__ re-exports. Nothing imported those names from the package; every consumer imports the defining module directly, so the block was a second place to list every public name and nothing else. --- .../cache_warming/__init__.py | 33 ----------------- .../cache_warming/refresher.py | 22 +++++++---- .../complexity_router/cache_warming/store.py | 37 +++++++++---------- .../cache_warming/test_store.py | 14 +++---- 4 files changed, 36 insertions(+), 70 deletions(-) diff --git a/litellm/router_strategy/complexity_router/cache_warming/__init__.py b/litellm/router_strategy/complexity_router/cache_warming/__init__.py index 0a2c1f4757a..e69de29bb2d 100644 --- a/litellm/router_strategy/complexity_router/cache_warming/__init__.py +++ b/litellm/router_strategy/complexity_router/cache_warming/__init__.py @@ -1,33 +0,0 @@ -from litellm.router_strategy.complexity_router.cache_warming.capture import capture_session -from litellm.router_strategy.complexity_router.cache_warming.eligibility import ( - min_prompt_cache_tokens_for_warm_set, - resolve_warm_models, -) -from litellm.router_strategy.complexity_router.cache_warming.store import CacheWarmingStore -from litellm.router_strategy.complexity_router.cache_warming.types import ( - CACHE_WARMING_RECORD_SCHEMA_VERSION, - CACHE_WARMING_REPLAY_MARKER_KEY, - CACHE_WARMING_REPLAY_TAG, - WARM_FRESHNESS_SLACK_SECONDS, - CacheWarmingAttribution, - CacheWarmingPayload, - CacheWarmingRecord, - compress_payload, - decompress_payload, -) - -__all__ = [ - "CacheWarmingStore", - "capture_session", - "min_prompt_cache_tokens_for_warm_set", - "resolve_warm_models", - "CACHE_WARMING_RECORD_SCHEMA_VERSION", - "CACHE_WARMING_REPLAY_MARKER_KEY", - "CACHE_WARMING_REPLAY_TAG", - "WARM_FRESHNESS_SLACK_SECONDS", - "CacheWarmingAttribution", - "CacheWarmingPayload", - "CacheWarmingRecord", - "compress_payload", - "decompress_payload", -] diff --git a/litellm/router_strategy/complexity_router/cache_warming/refresher.py b/litellm/router_strategy/complexity_router/cache_warming/refresher.py index 6cb5c640037..02b185eabc3 100644 --- a/litellm/router_strategy/complexity_router/cache_warming/refresher.py +++ b/litellm/router_strategy/complexity_router/cache_warming/refresher.py @@ -508,27 +508,33 @@ class CacheWarmingRefresher: config.max_sessions, ) now = time.time() - records = tuple([(key, await store.get_record(key)) for key in session_keys]) + sessions = tuple([(key, *await store.get_session(key)) for key in session_keys]) active = tuple( - (key, record) - for key, record in records + (key, record, touched) + for key, record, touched in sessions if record is not None and now - record.last_activity <= config.idle_timeout_seconds ) if not active: return - touched = {key: await store.get_touched_models(key) for key, _ in active} allowed = frozenset(resolve_warm_models(complexity_router.config)) warmable = frozenset( filter_cache_warmable( llm_router, - tuple(dict.fromkeys(model for models in touched.values() for model in models if model in allowed)), + tuple( + dict.fromkeys( + model + for _, record, touched in active + for model in (record.served_model, *touched) + if model in allowed + ) + ), ) ) if not warmable: return attributed = frozenset( record.attribution.user_api_key - for _, record in active + for _, record, _ in active if record.attribution.user_api_key is not None and record.attribution.user_api_key != LITELLM_PROXY_MASTER_KEY_ALIAS ) @@ -555,7 +561,7 @@ class CacheWarmingRefresher: session_key=key, record=record, warm_models=tuple( - model for model in dict.fromkeys((record.served_model, *touched[key])) if model in warmable + model for model in dict.fromkeys((record.served_model, *touched)) if model in warmable ), refresh_interval_seconds=config.refresh_interval_seconds, session_ttl_seconds=config.session_ttl_seconds, @@ -566,7 +572,7 @@ class CacheWarmingRefresher: proxy_logging_obj=proxy_logging_obj, lease_lost=lease_lost, ) - for key, record in active + for key, record, touched in active if record.attribution.user_api_key not in excluded_keys ), return_exceptions=True, diff --git a/litellm/router_strategy/complexity_router/cache_warming/store.py b/litellm/router_strategy/complexity_router/cache_warming/store.py index 0504ec2cf5d..83a31da159d 100644 --- a/litellm/router_strategy/complexity_router/cache_warming/store.py +++ b/litellm/router_strategy/complexity_router/cache_warming/store.py @@ -43,8 +43,8 @@ redis.call('EXPIREAT', touched_key, math.ceil(expires_at)) return 1 """ -_TOUCHED_MODELS_SCRIPT = """ -return redis.call('SMEMBERS', KEYS[1]) +_GET_SESSION_SCRIPT = """ +return {redis.call('HGET', KEYS[1], ARGV[1]), redis.call('SMEMBERS', KEYS[2])} """ _LIST_LIVE_SESSIONS_SCRIPT = """ @@ -54,10 +54,6 @@ local limit = tonumber(ARGV[2]) return redis.call('ZRANGEBYSCORE', index_key, '(' .. now, '+inf', 'LIMIT', 0, limit) """ -_GET_RECORD_SCRIPT = """ -return redis.call('HGET', KEYS[1], ARGV[1]) -""" - _MEMBERS_ADAPTER: TypeAdapter[tuple[str | bytes, ...]] = TypeAdapter(tuple[str | bytes, ...]) @@ -121,8 +117,7 @@ class CacheWarmingStore: self._list_live: Callable[..., Awaitable[object]] | None = ( register(_LIST_LIVE_SESSIONS_SCRIPT) if register else None ) - self._get: Callable[..., Awaitable[object]] | None = register(_GET_RECORD_SCRIPT) if register else None - self._touched: Callable[..., Awaitable[object]] | None = register(_TOUCHED_MODELS_SCRIPT) if register else None + self._get: Callable[..., Awaitable[object]] | None = register(_GET_SESSION_SCRIPT) if register else None @staticmethod def record_key(auto_router_model_name: str, caller_scope: str, session_id: str) -> str: @@ -158,21 +153,23 @@ class CacheWarmingStore: return None return self.redis_cache - async def get_record(self, key: str) -> CacheWarmingRecord | None: + async def get_session(self, key: str) -> "tuple[CacheWarmingRecord | None, tuple[str, ...]]": + """A session's record and the models it has been served on, in one round trip. The refresher needs + both for every live session on every tick, so reading them separately would double the sequential + round trips against the cap.""" if self._require_redis() is None or self._get is None: - return None - raw = await self._get(keys=[self.sessions_key()], args=[key]) - return _parse_record(raw) - - async def get_touched_models(self, key: str) -> tuple[str, ...]: - if self._require_redis() is None or self._touched is None: - return () - raw = await self._touched(keys=[self.touched_key(key)], args=[]) + return (None, ()) + raw = await self._get(keys=[self.sessions_key(), self.touched_key(key)], args=[key]) + record_raw, touched_raw = raw if isinstance(raw, (list, tuple)) and len(raw) == 2 else (None, ()) try: - members = _MEMBERS_ADAPTER.validate_python(raw) + members = _MEMBERS_ADAPTER.validate_python(touched_raw) except ValidationError: - return () - return tuple(member.decode() if isinstance(member, bytes) else member for member in members) + members = () + touched = tuple(member.decode() if isinstance(member, bytes) else member for member in members) + return (_parse_record(record_raw), touched) + + async def get_record(self, key: str) -> CacheWarmingRecord | None: + return (await self.get_session(key))[0] async def upsert_session( self, diff --git a/tests/test_litellm/router_strategy/complexity_router/cache_warming/test_store.py b/tests/test_litellm/router_strategy/complexity_router/cache_warming/test_store.py index 9f718a338e6..2fdf381fafe 100644 --- a/tests/test_litellm/router_strategy/complexity_router/cache_warming/test_store.py +++ b/tests/test_litellm/router_strategy/complexity_router/cache_warming/test_store.py @@ -64,18 +64,14 @@ class FakeRedisCache: return 0 return compare_and_delete - if "SMEMBERS" in script: - - async def touched_models(keys: list, args: list) -> list: - return [member.encode("utf-8") for member in sorted(self.sets.get(self._namespaced(keys[0]), set()))] - - return touched_models if "HGET" in script: - async def get_record(keys: list, args: list) -> str | None: - return self.hashes.get(self._namespaced(keys[0]), {}).get(str(args[0])) + async def get_session(keys: list, args: list) -> list: + record = self.hashes.get(self._namespaced(keys[0]), {}).get(str(args[0])) + touched = sorted(self.sets.get(self._namespaced(keys[1]), set())) + return [record, [member.encode("utf-8") for member in touched]] - return get_record + return get_session if "HSET" in script: async def capture(keys: list, args: list) -> int: