mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-14 23:21:35 +00:00
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:
parent
c0adf8653d
commit
049c1abef5
4 changed files with 36 additions and 70 deletions
|
|
@ -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",
|
||||
]
|
||||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue