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.
This commit is contained in:
Tin Chi Lo 2026-07-30 11:16:38 -07:00
parent c0adf8653d
commit 049c1abef5
4 changed files with 36 additions and 70 deletions

View file

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

View file

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

View file

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

View file

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