feat(complexity_router): consume plane for provider prompt-cache warming

The capture plane (parent commit) records each warming-enabled session's
latest payload in Redis. This adds the consume side so every tier model's
provider prompt cache stays warm while a session is active and tier
switches bill as cache reads:

- CacheWarmingRefresher.run_tick: single-pod via the Redis cron lock;
  when the injected PodLockManager has no Redis attached the refresher
  elects through a fallback PodLockManager bound to the warming store's
  own Redis, so multi-pod duplicate replays cannot happen whenever
  warming itself is active. Enumerates session records, drops idle sessions, skips keys that are deleted, blocked, expired, or at 95% of
  max_budget (fails open on lookup errors; the config-defined master key
  has no token row and stays eligible), and replays each due payload against every
  warm-set model with max_tokens=1, stream=False and response-cache
  bypass, bounded by a 10-replay semaphore. The warm set is warm_models
  or the first member of each tier pool, filtered to anthropic/bedrock
  deployments that support prompt caching; a deployment's declared
  custom_llm_provider wins over inference. Chat payloads replay via
  router.acompletion with internal metadata in the metadata slot;
  anthropic_messages payloads replay via router.aanthropic_messages with
  system and tool_choice forwarded and internal metadata in
  litellm_metadata, since metadata is the provider body field on that
  surface. Replays carry the loop-guard marker, the originating key's
  attribution, and the litellm_cache_warming tag unless tag filtering is
  enabled. Warmth stamps land even for failed attempts so a dead
  deployment is retried once per interval, not every tick.

- proxy_server._complexity_cache_warming_loop: 30s tick modeled on the
  adaptive flusher (sleep-first, resolves llm_router fresh each tick,
  re-raises CancelledError, survives tick exceptions), started
  unconditionally at startup so hot-reloaded routers are covered.

- ComplexityRouter._warm_aware_pick: consulted by _pick_model_for_tier on
  both the plain and plugin-narrowed branches. Prefers pool members whose
  warmth stamp is fresh within refresh_interval_seconds plus 60s slack,
  unions the served model while the session is active, and falls back to
  the existing pick when there is no session, record, store, or
  intersection.

- README: Cache Warming section covering the YAML shape, the Redis and
  store_prompts_in_spend_logs prerequisites, the Redis Cluster
  limitation, spend attribution and tagging, max_sessions semantics, and
  the session_affinity interplay.
This commit is contained in:
Tin Chi Lo 2026-07-28 15:30:49 -07:00
parent be7306b2a5
commit e79ea57fb4
8 changed files with 1172 additions and 0 deletions

View file

@ -1451,6 +1451,7 @@ CLOUDZERO_MAX_FETCHED_DATA_RECORDS = int(os.getenv("CLOUDZERO_MAX_FETCHED_DATA_R
SPEND_LOG_CLEANUP_JOB_NAME = "spend_log_cleanup"
KEY_ROTATION_JOB_NAME = "litellm_key_rotation_job"
EXPIRED_UI_SESSION_KEY_CLEANUP_JOB_NAME = "litellm_expired_ui_session_key_cleanup_job"
CACHE_WARMING_JOB_NAME = "complexity_router_cache_warming"
SPEND_LOG_RUN_LOOPS = int(os.getenv("SPEND_LOG_RUN_LOOPS", 500))
SPEND_LOG_CLEANUP_BATCH_SIZE = int(os.getenv("SPEND_LOG_CLEANUP_BATCH_SIZE", 1000))
SPEND_LOG_CLEANUP_MAX_CONSECUTIVE_BATCH_FAILURES = int(os.getenv("SPEND_LOG_CLEANUP_MAX_CONSECUTIVE_BATCH_FAILURES", 3))

View file

@ -567,6 +567,7 @@ from litellm.router import (
LiteLLM_Params,
ModelGroupInfo,
)
from litellm.router_strategy.complexity_router.cache_warming.refresher import CacheWarmingRefresher
from litellm.scheduler import FlowItem, Scheduler
from litellm.secret_managers.aws_secret_manager import load_aws_kms
from litellm.secret_managers.google_kms import load_google_kms
@ -1106,6 +1107,7 @@ async def proxy_startup_event(app: FastAPI):
await _tagged.strategy.load_state_from_db(prisma_client)
_tagged.strategy._state_loaded = True
asyncio.create_task(_adaptive_router_flusher_loop())
asyncio.create_task(_complexity_cache_warming_loop())
## [Optional] Initialize dd tracer
ProxyStartupEvent._init_dd_tracer()
@ -3318,6 +3320,29 @@ async def _adaptive_router_flusher_loop():
verbose_proxy_logger.exception("adaptive_router flusher iteration failed")
_COMPLEXITY_CACHE_WARMING_TICK_SECONDS = 30
async def _complexity_cache_warming_loop(refresher: CacheWarmingRefresher | None = None):
global llm_router, prisma_client
active_refresher = refresher if refresher is not None else CacheWarmingRefresher()
while True:
try:
await asyncio.sleep(_COMPLEXITY_CACHE_WARMING_TICK_SECONDS)
router = llm_router
if router is None:
continue
await active_refresher.run_tick(
llm_router=router,
pod_lock_manager=proxy_logging_obj.db_spend_update_writer.pod_lock_manager,
prisma_client=prisma_client,
)
except asyncio.CancelledError:
raise
except Exception: # noqa: BLE001 # one failed tick must never kill the warming loop
verbose_proxy_logger.exception("complexity_router cache warming tick failed")
async def _run_background_health_check():
"""
Periodically run health checks in the background on the endpoints.

View file

@ -142,6 +142,45 @@ Technical code keywords are detected case-insensitively and include:
- Infrastructure: `database`, `api`, `endpoint`, `docker`, `kubernetes`
- Actions: `debug`, `implement`, `refactor`, `optimize`
## Cache Warming
Provider prompt caches (Anthropic, Bedrock) are per-model, so a mid-session tier switch pays a fresh cache write on the new model and loses the cache-read discount. `cache_warming` keeps every tier model's prompt cache warm for active sessions: the proxy captures each session's latest payload and a background refresher replays it (`max_tokens=1`) against the other tier models before the provider's ~5 minute cache TTL expires. When the router later switches tiers, the switched-to model already has the session's prefix cached, and the routing pick prefers models whose cache is verifiably warm.
```yaml
model_list:
- model_name: smart-router
litellm_params:
model: auto_router/complexity_router
complexity_router_config:
tiers:
SIMPLE: fast-claude
COMPLEX: smart-claude
session_affinity: false # warming is the alternative to pinning; see below
cache_warming:
enabled: true
refresh_interval_seconds: 270 # keep under the provider cache TTL (Anthropic: 5 min)
session_ttl_seconds: 3600
idle_timeout_seconds: 600 # stop warming a session this long after its last real request
max_sessions: 1000
# warm_models: [fast-claude, smart-claude] # default: first member of each tier pool
general_settings:
store_prompts_in_spend_logs: true # consent gate; warming stores full payloads in Redis
router_settings:
redis_host: localhost
redis_port: 6379
```
Requirements and semantics:
- **Redis is required.** Session payloads, per-model warmth stamps, and a per-router session index live in Redis so all pods share them and a single pod (via a Redis cron lock) runs the replays. Without Redis, warming logs a warning once and no-ops; requests are unaffected. Sessions are tracked through the index rather than keyspace scans, so Redis Cluster is supported.
- **`store_prompts_in_spend_logs: true` is a prerequisite.** Warming persists full request payloads (messages, system, tools) in Redis, so it is gated on the same consent flag that governs storing prompts in spend logs. With the flag off, capture warns once and skips.
- **Only Anthropic and Bedrock models that support prompt caching are warmed.** Other models in the tier pools are left alone. Requests must carry a `metadata.session_id` and exceed the warm set's minimum cacheable token count (`prompt_cache_min_tokens`, default 1024) to be captured.
- **Spend attribution.** Replays run under the originating key: they appear in spend logs attributed to that key/team/user, tagged `litellm_cache_warming` so they are filterable in the Logs UI. A `max_tokens=1` replay of a warm prefix bills roughly 10% of the input cost. Warming stops for keys that are deleted, blocked, expired, or at 95% of their `max_budget` (fails open if the lookup errors).
- **`max_sessions`** caps concurrently warmed sessions per auto-router, enforced atomically at capture; once reached, new sessions are not admitted until existing ones expire.
- **Interplay with `session_affinity`** (default on): affinity pins a session to its first-turn model, so no tier switch happens and warming buys nothing; with affinity on, captured sessions are still warmed but the pin decides routing. Disable `session_affinity` to let per-turn classification switch tiers and have warming make those switches cache hits.
## Performance
- **Classification time**: <1ms typical

View file

@ -0,0 +1,290 @@
import asyncio
import time
from collections.abc import Mapping, Sequence
from datetime import datetime
from typing import TYPE_CHECKING
from pydantic import BaseModel, TypeAdapter
from litellm._logging import verbose_router_logger
from litellm.constants import CACHE_WARMING_JOB_NAME, LITELLM_PROXY_MASTER_KEY_ALIAS
from litellm.router_strategy.complexity_router.cache_warming.eligibility import 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_REPLAY_MARKER_KEY,
CACHE_WARMING_REPLAY_TAG,
CacheWarmingRecord,
decompress_payload,
)
if TYPE_CHECKING:
from litellm.caching.redis_cache import RedisCache
from litellm.proxy.db.db_transaction_queue.pod_lock_manager import PodLockManager
from litellm.proxy.utils import PrismaClient
from litellm.router import Router
from litellm.router_strategy.complexity_router.complexity_router import ComplexityRouter
CACHE_WARMING_MAX_CONCURRENT_REPLAYS = 10
CACHE_WARMING_LOCK_TTL_SECONDS = 60
CACHE_WARMING_BUDGET_STOP_FRACTION = 0.95
_ATTRIBUTION_ADAPTER: TypeAdapter[Mapping[str, str | None]] = TypeAdapter(Mapping[str, str | None])
class _TokenBudgetRow(BaseModel):
token: str
spend: float = 0.0
max_budget: float | None = None
blocked: bool | None = None
expires: datetime | None = None
_TOKEN_ROWS_ADAPTER: TypeAdapter[tuple[_TokenBudgetRow, ...]] = TypeAdapter(tuple[_TokenBudgetRow, ...])
def _excluded_from_warming(row: _TokenBudgetRow, now: float) -> bool:
if row.blocked is True:
return True
if row.expires is not None and row.expires.timestamp() <= now:
return True
return row.max_budget is not None and row.spend >= CACHE_WARMING_BUDGET_STOP_FRACTION * row.max_budget
def collect_warming_enabled_complexity_routers(llm_router: "Router") -> tuple["ComplexityRouter", ...]:
return tuple(
tagged.strategy
for tagged_list in llm_router.complexity_routers.values()
for tagged in tagged_list
if tagged.strategy.config.cache_warming.enabled
)
def _deployment_provider(litellm_params: Mapping[str, object], deployment_model: str) -> str | None:
from litellm.litellm_core_utils.get_llm_provider_logic import get_llm_provider
declared = litellm_params.get("custom_llm_provider")
if isinstance(declared, str) and declared:
return declared
try:
_, provider, _, _ = get_llm_provider(model=deployment_model)
except Exception: # noqa: BLE001 # unroutable deployment just isn't warmable
return None
return provider
def _group_is_cache_warmable(llm_router: "Router", model_group: str) -> bool:
from litellm.utils import supports_prompt_caching
deployments = llm_router.get_model_list(model_name=model_group) or []
for deployment in deployments:
litellm_params = deployment.get("litellm_params") or {} # pyright: ignore[reportUnknownMemberType] # DeploymentTypedDict fields are legacy-untyped
deployment_model = litellm_params.get("model") # pyright: ignore[reportUnknownMemberType] # DeploymentTypedDict fields are legacy-untyped
if not isinstance(deployment_model, str):
continue
provider = _deployment_provider(litellm_params, deployment_model)
if provider not in ("anthropic", "bedrock"):
continue
if supports_prompt_caching(model=deployment_model, custom_llm_provider=provider):
return True
return False
def filter_cache_warmable(llm_router: "Router", model_groups: Sequence[str]) -> tuple[str, ...]:
return tuple(group for group in model_groups if _group_is_cache_warmable(llm_router, group))
class CacheWarmingRefresher:
def __init__(self, max_concurrent_replays: int = CACHE_WARMING_MAX_CONCURRENT_REPLAYS) -> None:
self.max_concurrent_replays = max_concurrent_replays
self._fallback_lock_manager: PodLockManager | None = None
def _resolve_lock_manager(
self, injected: "PodLockManager | None", redis_cache: "RedisCache | None"
) -> "PodLockManager":
if injected is not None and injected.redis_cache is not None:
return injected
if self._fallback_lock_manager is None or self._fallback_lock_manager.redis_cache is not redis_cache:
from litellm.proxy.db.db_transaction_queue.pod_lock_manager import PodLockManager
self._fallback_lock_manager = PodLockManager(redis_cache=redis_cache)
return self._fallback_lock_manager
async def run_tick(
self,
*,
llm_router: "Router",
pod_lock_manager: "PodLockManager | None",
prisma_client: "PrismaClient | None",
) -> None:
warming_routers = collect_warming_enabled_complexity_routers(llm_router)
if not warming_routers:
return
warmable = tuple(
(complexity_router, store)
for complexity_router in warming_routers
if (store := complexity_router.get_cache_warming_store()) is not None and store.redis_cache is not None
)
if not warmable:
return
lock_manager = self._resolve_lock_manager(pod_lock_manager, warmable[0][1].redis_cache)
acquired = await lock_manager.acquire_lock(
cronjob_id=CACHE_WARMING_JOB_NAME, ttl=CACHE_WARMING_LOCK_TTL_SECONDS
)
if not acquired:
return
try:
for complexity_router, store in warmable:
await self._warm_router_sessions(
llm_router=llm_router,
complexity_router=complexity_router,
store=store,
prisma_client=prisma_client,
)
finally:
await lock_manager.release_lock(cronjob_id=CACHE_WARMING_JOB_NAME)
async def _warm_router_sessions(
self,
*,
llm_router: "Router",
complexity_router: "ComplexityRouter",
store: CacheWarmingStore,
prisma_client: "PrismaClient | None",
) -> None:
config = complexity_router.config.cache_warming
session_keys = await store.list_session_keys(max_sessions=config.max_sessions)
if not session_keys:
return
if len(session_keys) >= config.max_sessions:
verbose_router_logger.debug(
"cache_warming: auto-router %s is at its max_sessions cap (%s); "
"new sessions are not admitted until existing ones expire",
complexity_router.model_name,
config.max_sessions,
)
now = time.time()
records = tuple([(key, await store.get_record(key)) for key in session_keys])
active = tuple(
(key, record)
for key, record in records
if record is not None and now - record.last_activity <= config.idle_timeout_seconds
)
if not active:
return
warm_models = filter_cache_warmable(llm_router, resolve_warm_models(complexity_router.config))
if not warm_models:
return
excluded_keys = await self._excluded_key_hashes(
prisma_client,
frozenset(
record.attribution.user_api_key for _, record in active if record.attribution.user_api_key is not None
),
)
semaphore = asyncio.Semaphore(self.max_concurrent_replays)
await asyncio.gather(
*(
self._warm_session(
llm_router=llm_router,
store=store,
session_key=key,
record=record,
warm_models=warm_models,
refresh_interval_seconds=config.refresh_interval_seconds,
session_ttl_seconds=config.session_ttl_seconds,
semaphore=semaphore,
)
for key, record in active
if record.attribution.user_api_key not in excluded_keys
)
)
async def _warm_session(
self,
*,
llm_router: "Router",
store: CacheWarmingStore,
session_key: str,
record: CacheWarmingRecord,
warm_models: tuple[str, ...],
refresh_interval_seconds: int,
session_ttl_seconds: int,
semaphore: asyncio.Semaphore,
) -> None:
warmth = await store.get_warmth(session_key, warm_models)
now = time.time()
due_models = tuple(model for model in warm_models if now - warmth.get(model, 0.0) >= refresh_interval_seconds)
for model_group in due_models:
async with semaphore:
attempted_at = time.time()
try:
await self._replay(llm_router=llm_router, record=record, model_group=model_group)
except Exception: # noqa: BLE001 # one failing replay must not abort the tick
verbose_router_logger.warning(
"cache_warming replay failed for session %s model %s",
session_key,
model_group,
exc_info=True,
)
finally:
await store.mark_warm_attempt(session_key, model_group, attempted_at, session_ttl_seconds)
async def _replay(self, *, llm_router: "Router", record: CacheWarmingRecord, model_group: str) -> None:
payload = decompress_payload(record.payload_compressed)
attribution = _ATTRIBUTION_ADAPTER.validate_python(record.attribution.model_dump())
metadata = { # mutable-ok: request metadata handed to the router call, never retained
CACHE_WARMING_REPLAY_MARKER_KEY: True,
**{key: value for key, value in attribution.items() if value is not None},
**(
{"tags": [CACHE_WARMING_REPLAY_TAG]} if llm_router.enable_tag_filtering is not True else {}
), # mutable-ok: request metadata, never retained
}
messages = [dict(message) for message in payload.messages] # mutable-ok: router call input, never retained
if payload.call_surface == "anthropic_messages":
system = list(payload.system) if isinstance(payload.system, tuple) else payload.system
await llm_router.aanthropic_messages( # pyright: ignore[reportUnknownMemberType] # factory-generated router surface is legacy-untyped
model=model_group,
messages=messages,
system=system,
tools=list(payload.tools) if payload.tools is not None else None,
tool_choice=dict(payload.tool_choice)
if isinstance(payload.tool_choice, Mapping)
else payload.tool_choice, # mutable-ok: router call input, never retained
max_tokens=1,
stream=False,
cache={"no-cache": True}, # mutable-ok: router call input, never retained
litellm_metadata=metadata,
)
return
await llm_router.acompletion( # pyright: ignore[reportUnknownMemberType, reportCallIssue] # router overloads are legacy-untyped
model=model_group,
messages=messages, # pyright: ignore[reportArgumentType] # replay forwards the captured wire shape verbatim
tools=list(payload.tools) if payload.tools is not None else None,
tool_choice=dict(payload.tool_choice)
if isinstance(payload.tool_choice, Mapping)
else payload.tool_choice, # mutable-ok: router call input, never retained
max_tokens=1,
stream=False,
cache={"no-cache": True}, # mutable-ok: router call input, never retained
metadata=metadata,
)
@staticmethod
async def _excluded_key_hashes(prisma_client: "PrismaClient | None", key_hashes: frozenset[str]) -> frozenset[str]:
if prisma_client is None or not key_hashes:
return frozenset()
lookup = frozenset(key for key in key_hashes if key != LITELLM_PROXY_MASTER_KEY_ALIAS)
if not lookup:
return frozenset()
try:
rows = _TOKEN_ROWS_ADAPTER.validate_python(
await prisma_client.db.litellm_verificationtoken.find_many( # pyright: ignore[reportAny] # prisma client is legacy-untyped
where={"token": {"in": list(lookup)}} # mutable-ok: prisma query input, never retained
),
from_attributes=True,
)
except Exception: # noqa: BLE001 # budget stop fails open; warming continues on query errors
verbose_router_logger.warning("cache_warming budget check failed; warming continues", exc_info=True)
return frozenset()
now = time.time()
usable = frozenset(row.token for row in rows if not _excluded_from_warming(row, now))
return lookup - usable

View file

@ -18,6 +18,8 @@ from __future__ import annotations
import asyncio
import random
import re
import time
from collections.abc import Mapping, Sequence
from typing import TYPE_CHECKING, Any, Literal, Union, cast
from pydantic import BaseModel
@ -498,6 +500,9 @@ class ComplexityRouter(CustomLogger):
request_kwargs: dict,
) -> str:
if not self.config.plugins:
warm_pick = await self._warm_aware_pick(self._tier_pools().get(tier.value, []), request_kwargs)
if warm_pick is not None:
return warm_pick
return self.get_model_for_tier(tier)
from litellm.types.router import RoutingContext
@ -520,8 +525,42 @@ class ComplexityRouter(CustomLogger):
# silently bypassed. Raise instead, matching the Router-level plugin
# pipeline's own fail-closed behavior for the same situation.
raise ValueError(f"No candidate models left for tier {tier_key} after routing-plugin filtering")
warm_pick = await self._warm_aware_pick(context.candidate_models, request_kwargs)
if warm_pick is not None:
return warm_pick
return self._pick_from_tier_value(context.candidate_models, tier_key)
async def _warm_aware_pick(self, pool: Sequence[str], request_kwargs: Mapping[str, object]) -> str | None:
config = self.config.cache_warming
if not config.enabled or len(pool) <= 1:
return None
session_id = get_session_id_from_request_kwargs(request_kwargs)
if session_id is None:
return None
store = self.get_cache_warming_store()
if store is None or store.redis_cache is None:
return None
from litellm.router_strategy.complexity_router.cache_warming.types import WARM_FRESHNESS_SLACK_SECONDS
caller_scope = get_user_api_key_hash_from_request_kwargs(request_kwargs) or "unscoped"
record_key = store.record_key(self.model_name, caller_scope, session_id)
record = await store.get_record(record_key)
if record is None:
return None
warmth = await store.get_warmth(record_key, tuple(pool))
now = time.time()
freshness_window = config.refresh_interval_seconds + WARM_FRESHNESS_SLACK_SECONDS
warmed = frozenset(model for model, warmed_at in warmth.items() if now - warmed_at <= freshness_window)
served = (
frozenset((record.served_model,))
if now - record.last_activity <= config.idle_timeout_seconds
else frozenset[str]()
)
candidates = tuple(model for model in pool if model in warmed | served)
if not candidates:
return None
return random.choice(candidates)
def _ensure_adaptive_router(self) -> Any | None:
if not self.config.adaptive:
return None

View file

@ -8,6 +8,7 @@ Pins covered:
- ``_get_endpoint_exception_status``
- ``_write_health_state_to_router_cache``
- ``_adaptive_router_flusher_loop``
- ``_complexity_cache_warming_loop``
- ``_run_background_health_check``
"""
@ -22,6 +23,7 @@ import pytest
import litellm.proxy.proxy_server as proxy_server
from litellm.proxy.proxy_server import (
_adaptive_router_flusher_loop,
_complexity_cache_warming_loop,
_get_endpoint_exception_status,
_get_process_rss_mb,
_run_background_health_check,
@ -427,6 +429,91 @@ async def test_adaptive_router_flusher_loop_times_out_when_sleep_real(monkeypatc
await asyncio.wait_for(_adaptive_router_flusher_loop(), timeout=0.2)
# ---------------------------------------------------------------------------
# _complexity_cache_warming_loop
# ---------------------------------------------------------------------------
class _RecordingRefresher:
def __init__(self, fail_on_call: int | None = None):
self.calls: list[dict] = []
self.fail_on_call = fail_on_call
async def run_tick(self, *, llm_router, pod_lock_manager, prisma_client):
self.calls.append(
{"llm_router": llm_router, "pod_lock_manager": pod_lock_manager, "prisma_client": prisma_client}
)
if self.fail_on_call == len(self.calls):
raise RuntimeError("tick boom")
def _cancel_sleep_after(monkeypatch, iterations: int, on_call=None):
call_count = {"n": 0}
_real_sleep = asyncio.sleep
async def _short_sleep(_seconds):
call_count["n"] += 1
if on_call is not None:
on_call(call_count["n"])
if call_count["n"] > iterations:
raise asyncio.CancelledError()
await _real_sleep(0)
monkeypatch.setattr(proxy_server.asyncio, "sleep", _short_sleep)
@pytest.mark.asyncio
async def test_complexity_cache_warming_loop_survives_tick_exception(monkeypatch):
refresher = _RecordingRefresher(fail_on_call=1)
monkeypatch.setattr(proxy_server, "llm_router", MagicMock())
_cancel_sleep_after(monkeypatch, iterations=2)
with pytest.raises(asyncio.CancelledError):
await _complexity_cache_warming_loop(refresher=refresher)
assert len(refresher.calls) == 2
@pytest.mark.asyncio
async def test_complexity_cache_warming_loop_noop_while_router_none_and_resolves_router_fresh(monkeypatch):
refresher = _RecordingRefresher()
fake_router = MagicMock()
fake_prisma = MagicMock()
monkeypatch.setattr(proxy_server, "llm_router", None)
monkeypatch.setattr(proxy_server, "prisma_client", fake_prisma)
def _swap_router_in(call_number: int):
if call_number == 2:
monkeypatch.setattr(proxy_server, "llm_router", fake_router)
_cancel_sleep_after(monkeypatch, iterations=2, on_call=_swap_router_in)
with pytest.raises(asyncio.CancelledError):
await _complexity_cache_warming_loop(refresher=refresher)
assert len(refresher.calls) == 1
assert refresher.calls[0]["llm_router"] is fake_router
assert refresher.calls[0]["prisma_client"] is fake_prisma
assert (
refresher.calls[0]["pod_lock_manager"]
is proxy_server.proxy_logging_obj.db_spend_update_writer.pod_lock_manager
)
@pytest.mark.asyncio
async def test_complexity_cache_warming_loop_is_infinite(monkeypatch):
monkeypatch.setattr(proxy_server, "llm_router", None)
_real_sleep = asyncio.sleep
async def _instant_sleep(_seconds):
await _real_sleep(0)
monkeypatch.setattr(proxy_server.asyncio, "sleep", _instant_sleep)
with pytest.raises(asyncio.TimeoutError):
await asyncio.wait_for(_complexity_cache_warming_loop(refresher=_RecordingRefresher()), timeout=0.2)
# ---------------------------------------------------------------------------
# _run_background_health_check
# ---------------------------------------------------------------------------

View file

@ -0,0 +1,545 @@
import asyncio
import json
import time
from types import SimpleNamespace
import pytest
from litellm.constants import CACHE_WARMING_JOB_NAME
from litellm.router_strategy.complexity_router.cache_warming.eligibility import resolve_warm_models
from litellm.router_strategy.complexity_router.cache_warming.refresher import (
CacheWarmingRefresher,
collect_warming_enabled_complexity_routers,
filter_cache_warmable,
)
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,
CacheWarmingAttribution,
CacheWarmingPayload,
CacheWarmingRecord,
compress_payload,
)
from litellm.router_strategy.complexity_router.complexity_router import ComplexityRouter
from litellm.router_strategy.complexity_router.config import ComplexityRouterConfig
from litellm.types.router import TaggedPreRoutingStrategy
from tests.test_litellm.router_strategy.complexity_router.cache_warming.test_store import FakeRedisCache
_DEFAULT_DEPLOYMENTS = {
"fast-claude": [{"litellm_params": {"model": "anthropic/claude-haiku-4-5"}}],
"smart-claude": [{"litellm_params": {"model": "anthropic/claude-sonnet-4-5"}}],
"fast-gpt": [{"litellm_params": {"model": "gpt-5-mini"}}],
"titan": [{"litellm_params": {"model": "bedrock/amazon.titan-text-express-v1"}}],
"mystery": [{"litellm_params": {"model": "totally-unknown-model-xyz"}}],
}
class FakeLLMRouter:
def __init__(
self,
redis: FakeRedisCache | None = None,
enable_tag_filtering: bool = False,
deployments: dict | None = None,
replay_delay: float = 0.0,
) -> None:
self.complexity_routers: dict = {}
self.enable_tag_filtering = enable_tag_filtering
self.cache = SimpleNamespace(redis_cache=redis)
self.completion_calls: list[dict] = []
self.anthropic_calls: list[dict] = []
self.max_concurrent = 0
self._in_flight = 0
self.replay_delay = replay_delay
self.failing_message_marker: str | None = None
self._deployments = deployments if deployments is not None else _DEFAULT_DEPLOYMENTS
def get_model_list(self, model_name: str | None = None, team_id: str | None = None):
return self._deployments.get(model_name)
async def acompletion(self, **kwargs: object):
marker = self.failing_message_marker
if marker is not None and marker in json.dumps(kwargs.get("messages")):
raise RuntimeError("provider down")
self._in_flight += 1
self.max_concurrent = max(self.max_concurrent, self._in_flight)
await asyncio.sleep(self.replay_delay)
self._in_flight -= 1
self.completion_calls.append(kwargs)
async def aanthropic_messages(self, **kwargs: object):
self.anthropic_calls.append(kwargs)
class FakePodLockManager:
def __init__(self, acquire_result: bool | None = True, redis_cache: object = "attached") -> None:
self.acquire_result = acquire_result
self.redis_cache = redis_cache
self.acquire_calls: list[tuple[str, int | None]] = []
self.release_calls: list[str] = []
async def acquire_lock(self, cronjob_id: str, ttl: int | None = None) -> bool | None:
self.acquire_calls.append((cronjob_id, ttl))
return self.acquire_result
async def release_lock(self, cronjob_id: str) -> None:
self.release_calls.append(cronjob_id)
class FakePrismaClient:
def __init__(self, rows: tuple = (), raise_error: bool = False) -> None:
self.queries: list[dict] = []
async def find_many(where: dict):
self.queries.append(where)
if raise_error:
raise RuntimeError("db down")
return [row for row in rows if row.token in where["token"]["in"]]
self.db = SimpleNamespace(litellm_verificationtoken=SimpleNamespace(find_many=find_many))
def _warming_rig(
redis: FakeRedisCache | None = None,
enable_tag_filtering: bool = False,
replay_delay: float = 0.0,
**cache_warming_overrides: object,
) -> tuple[FakeLLMRouter, FakeRedisCache | None]:
llm_router = FakeLLMRouter(redis=redis, enable_tag_filtering=enable_tag_filtering, replay_delay=replay_delay)
strategy = ComplexityRouter(
model_name="smart-router",
litellm_router_instance=llm_router,
complexity_router_config={
"tiers": {"SIMPLE": ["fast-claude"], "COMPLEX": ["smart-claude"]},
"cache_warming": {"enabled": True, **cache_warming_overrides},
},
)
llm_router.complexity_routers = {"smart-router": [TaggedPreRoutingStrategy(tags=(), strategy=strategy)]}
return llm_router, redis
def _seed_session(
redis: FakeRedisCache,
session_id: str = "sess-1",
caller_scope: str = "hash-1",
served_model: str = "fast-claude",
last_activity: float | None = None,
warmth: dict | None = None,
user_api_key: str | None = "hash-1",
call_surface: str = "chat_completions",
content: str = "summarize the deployment policy",
tools: tuple | None = None,
tool_choice: object = None,
) -> str:
payload = CacheWarmingPayload(
model=served_model,
messages=({"role": "user", "content": content},),
system="You are a policy assistant" if call_surface == "anthropic_messages" else None,
tools=tools,
tool_choice=tool_choice,
call_surface=call_surface,
)
blob, sha = compress_payload(payload)
record = CacheWarmingRecord(
schema_version=CACHE_WARMING_RECORD_SCHEMA_VERSION,
payload_compressed=blob,
payload_sha256=sha,
token_estimate=2048,
last_activity=last_activity if last_activity is not None else time.time(),
served_model=served_model,
attribution=CacheWarmingAttribution(user_api_key=user_api_key),
auto_router_model_name="smart-router",
)
store = CacheWarmingStore(redis_cache=redis, auto_router_model_name="smart-router")
key = CacheWarmingStore.record_key("smart-router", caller_scope, session_id)
redis.hashes.setdefault(store.sessions_key(), {})[key] = json.dumps(record.model_dump())
redis.zsets.setdefault(store.index_key(), {})[key] = time.time() + 3600
for model_group, stamp in (warmth or {}).items():
redis.data[CacheWarmingStore.warmth_key(key, model_group)] = json.dumps(stamp)
return key
def _warmth_stamp(redis: FakeRedisCache, record_key: str, model_group: str) -> float | None:
raw = redis.data.get(CacheWarmingStore.warmth_key(record_key, model_group))
return json.loads(raw) if raw is not None else None
async def _tick(llm_router, lock=None, prisma=None, refresher: CacheWarmingRefresher | None = None):
await (refresher or CacheWarmingRefresher()).run_tick(
llm_router=llm_router, pod_lock_manager=lock, prisma_client=prisma
)
def _replayed_models(llm_router: FakeLLMRouter) -> list:
return [call["model"] for call in llm_router.completion_calls]
# ---------------------------------------------------------------------------
# warm-set resolution + eligibility
# ---------------------------------------------------------------------------
def test_resolve_warm_models_defaults_to_first_member_per_tier_deduped():
config = ComplexityRouterConfig(
tiers={"SIMPLE": ["fast-claude", "fast-gpt"], "MEDIUM": "fast-claude", "COMPLEX": ["smart-claude"]}
)
assert resolve_warm_models(config) == ("fast-claude", "smart-claude")
def test_resolve_warm_models_prefers_explicit_list():
config = ComplexityRouterConfig(
tiers={"SIMPLE": ["fast-claude"], "COMPLEX": ["smart-claude"]},
cache_warming={"enabled": True, "warm_models": ["smart-claude", "smart-claude", "fast-claude"]},
)
assert resolve_warm_models(config) == ("smart-claude", "fast-claude")
def test_filter_cache_warmable_keeps_only_prompt_cacheable_anthropic_bedrock():
llm_router = FakeLLMRouter()
groups = ["fast-claude", "smart-claude", "fast-gpt", "titan", "mystery", "absent-group"]
assert filter_cache_warmable(llm_router, groups) == ("fast-claude", "smart-claude")
def test_filter_cache_warmable_prefers_declared_custom_llm_provider_over_inference():
deployments = {
"declared-openai": [{"litellm_params": {"model": "anthropic/claude-sonnet-4-5", "custom_llm_provider": "openai"}}]
}
assert filter_cache_warmable(FakeLLMRouter(deployments=deployments), ["declared-openai"]) == ()
def test_collect_warming_enabled_complexity_routers_skips_disabled():
llm_router, _ = _warming_rig(redis=FakeRedisCache())
disabled = ComplexityRouter(
model_name="plain-router",
litellm_router_instance=llm_router,
complexity_router_config={"tiers": {"SIMPLE": ["fast-claude"]}},
)
llm_router.complexity_routers["plain-router"] = [TaggedPreRoutingStrategy(tags=(), strategy=disabled)]
collected = collect_warming_enabled_complexity_routers(llm_router)
assert [strategy.model_name for strategy in collected] == ["smart-router"]
# ---------------------------------------------------------------------------
# pod lock
# ---------------------------------------------------------------------------
@pytest.mark.asyncio
async def test_lock_held_by_other_pod_skips_tick_and_never_releases():
llm_router, redis = _warming_rig(redis=FakeRedisCache())
_seed_session(redis)
lock = FakePodLockManager(acquire_result=False)
await _tick(llm_router, lock=lock)
assert llm_router.completion_calls == []
assert lock.release_calls == []
@pytest.mark.asyncio
async def test_lock_acquired_warms_and_releases():
llm_router, redis = _warming_rig(redis=FakeRedisCache())
_seed_session(redis)
lock = FakePodLockManager(acquire_result=True)
await _tick(llm_router, lock=lock)
assert sorted(_replayed_models(llm_router)) == ["fast-claude", "smart-claude"]
assert lock.acquire_calls == [(CACHE_WARMING_JOB_NAME, 60)]
assert lock.release_calls == [CACHE_WARMING_JOB_NAME]
@pytest.mark.asyncio
async def test_lock_falls_back_to_warming_redis_when_no_injected_manager():
llm_router, redis = _warming_rig(redis=FakeRedisCache())
_seed_session(redis)
await _tick(llm_router, lock=None)
assert sorted(_replayed_models(llm_router)) == ["fast-claude", "smart-claude"]
assert any(key.startswith("cronjob_lock:") for key in redis.data)
@pytest.mark.asyncio
async def test_redisless_injected_manager_is_replaced_by_warming_redis_lock():
llm_router, redis = _warming_rig(redis=FakeRedisCache())
_seed_session(redis)
lock = FakePodLockManager(redis_cache=None)
await _tick(llm_router, lock=lock)
assert lock.acquire_calls == []
assert sorted(_replayed_models(llm_router)) == ["fast-claude", "smart-claude"]
assert any(key.startswith("cronjob_lock:") for key in redis.data)
def test_fallback_lock_manager_is_stable_across_ticks():
redis = FakeRedisCache()
refresher = CacheWarmingRefresher()
first = refresher._resolve_lock_manager(None, redis)
second = refresher._resolve_lock_manager(None, redis)
assert first is second
assert first.redis_cache is redis
@pytest.mark.asyncio
async def test_lock_released_when_tick_raises():
class ExplodingScriptRedis(FakeRedisCache):
def async_register_script(self, script: str):
async def boom(keys: list, args: list):
raise RuntimeError("redis down")
return boom
llm_router, _ = _warming_rig(redis=ExplodingScriptRedis())
lock = FakePodLockManager(acquire_result=True)
with pytest.raises(RuntimeError, match="redis down"):
await _tick(llm_router, lock=lock)
assert lock.release_calls == [CACHE_WARMING_JOB_NAME]
@pytest.mark.asyncio
async def test_no_warming_routers_never_touches_lock():
llm_router = FakeLLMRouter(redis=FakeRedisCache())
lock = FakePodLockManager()
await _tick(llm_router, lock=lock)
assert lock.acquire_calls == []
# ---------------------------------------------------------------------------
# session selection: idle skip + interval pacing
# ---------------------------------------------------------------------------
@pytest.mark.asyncio
async def test_idle_session_not_warmed():
llm_router, redis = _warming_rig(redis=FakeRedisCache())
key = _seed_session(redis, last_activity=time.time() - 601)
sessions_key = CacheWarmingStore(redis_cache=redis, auto_router_model_name="smart-router").sessions_key()
before = redis.hashes[sessions_key][key]
await _tick(llm_router)
assert llm_router.completion_calls == []
assert redis.hashes[sessions_key][key] == before
@pytest.mark.asyncio
async def test_recently_warmed_model_not_replayed_again():
llm_router, redis = _warming_rig(redis=FakeRedisCache())
now = time.time()
_seed_session(redis, warmth={"fast-claude": now - 10, "smart-claude": now - 300})
await _tick(llm_router)
assert _replayed_models(llm_router) == ["smart-claude"]
@pytest.mark.asyncio
async def test_no_redis_store_is_noop():
llm_router, _ = _warming_rig(redis=None)
await _tick(llm_router, lock=FakePodLockManager())
assert llm_router.completion_calls == []
# ---------------------------------------------------------------------------
# replay kwargs
# ---------------------------------------------------------------------------
@pytest.mark.asyncio
async def test_chat_replay_kwargs_exact():
llm_router, redis = _warming_rig(redis=FakeRedisCache())
_seed_session(redis, warmth={"fast-claude": time.time()})
await _tick(llm_router)
assert llm_router.completion_calls == [
{
"model": "smart-claude",
"messages": [{"role": "user", "content": "summarize the deployment policy"}],
"tools": None,
"tool_choice": None,
"max_tokens": 1,
"stream": False,
"cache": {"no-cache": True},
"metadata": {
CACHE_WARMING_REPLAY_MARKER_KEY: True,
"user_api_key": "hash-1",
"tags": [CACHE_WARMING_REPLAY_TAG],
},
}
]
@pytest.mark.asyncio
async def test_anthropic_surface_replays_via_aanthropic_messages_with_system():
llm_router, redis = _warming_rig(redis=FakeRedisCache())
_seed_session(
redis,
call_surface="anthropic_messages",
warmth={"fast-claude": time.time()},
tool_choice={"type": "auto"},
)
await _tick(llm_router)
assert llm_router.completion_calls == []
assert len(llm_router.anthropic_calls) == 1
call = llm_router.anthropic_calls[0]
assert call["model"] == "smart-claude"
assert call["system"] == "You are a policy assistant"
assert call["tool_choice"] == {"type": "auto"}
assert call["max_tokens"] == 1
assert call["stream"] is False
assert call["litellm_metadata"][CACHE_WARMING_REPLAY_MARKER_KEY] is True
assert "metadata" not in call
@pytest.mark.asyncio
async def test_tags_omitted_when_router_tag_filtering_enabled():
llm_router, redis = _warming_rig(redis=FakeRedisCache(), enable_tag_filtering=True)
_seed_session(redis, warmth={"fast-claude": time.time()})
await _tick(llm_router)
metadata = llm_router.completion_calls[0]["metadata"]
assert "tags" not in metadata
assert metadata[CACHE_WARMING_REPLAY_MARKER_KEY] is True
# ---------------------------------------------------------------------------
# failure isolation + attempt timestamping
# ---------------------------------------------------------------------------
@pytest.mark.asyncio
async def test_one_session_failure_does_not_block_other_sessions():
llm_router, redis = _warming_rig(redis=FakeRedisCache())
llm_router.failing_message_marker = "POISON"
_seed_session(redis, session_id="sess-bad", content="POISON payload")
_seed_session(redis, session_id="sess-good")
await _tick(llm_router)
assert sorted(_replayed_models(llm_router)) == ["fast-claude", "smart-claude"]
assert all("POISON" not in json.dumps(call) for call in llm_router.completion_calls)
@pytest.mark.asyncio
async def test_failed_replay_still_stamps_warm_attempt():
llm_router, redis = _warming_rig(redis=FakeRedisCache())
llm_router.failing_message_marker = "POISON"
key = _seed_session(redis, content="POISON payload", warmth={"fast-claude": time.time()})
await _tick(llm_router)
stamp = _warmth_stamp(redis, key, "smart-claude")
assert stamp is not None and stamp > 0
@pytest.mark.asyncio
async def test_successful_replay_stamps_warm_attempt_for_replayed_model_only():
llm_router, redis = _warming_rig(redis=FakeRedisCache())
now = time.time()
key = _seed_session(redis, warmth={"fast-claude": now})
await _tick(llm_router)
smart_stamp = _warmth_stamp(redis, key, "smart-claude")
assert smart_stamp is not None and smart_stamp >= now
assert _warmth_stamp(redis, key, "fast-claude") == pytest.approx(now)
# ---------------------------------------------------------------------------
# budget stop
# ---------------------------------------------------------------------------
@pytest.mark.asyncio
async def test_near_budget_key_sessions_are_not_warmed():
llm_router, redis = _warming_rig(redis=FakeRedisCache())
_seed_session(redis, session_id="sess-broke", user_api_key="broke-key")
_seed_session(redis, session_id="sess-rich", caller_scope="hash-2", user_api_key="rich-key")
prisma = FakePrismaClient(
rows=(
SimpleNamespace(token="broke-key", spend=95.0, max_budget=100.0),
SimpleNamespace(token="rich-key", spend=10.0, max_budget=100.0),
)
)
await _tick(llm_router, prisma=prisma)
assert sorted(_replayed_models(llm_router)) == ["fast-claude", "smart-claude"]
assert sorted(prisma.queries[0]["token"]["in"]) == ["broke-key", "rich-key"]
replayed_keys = {call["metadata"]["user_api_key"] for call in llm_router.completion_calls}
assert replayed_keys == {"rich-key"}
@pytest.mark.asyncio
async def test_budget_stop_skipped_without_prisma():
llm_router, redis = _warming_rig(redis=FakeRedisCache())
_seed_session(redis, user_api_key="broke-key")
await _tick(llm_router, prisma=None)
assert sorted(_replayed_models(llm_router)) == ["fast-claude", "smart-claude"]
@pytest.mark.asyncio
async def test_budget_query_error_fails_open():
llm_router, redis = _warming_rig(redis=FakeRedisCache())
_seed_session(redis, user_api_key="broke-key")
await _tick(llm_router, prisma=FakePrismaClient(raise_error=True))
assert sorted(_replayed_models(llm_router)) == ["fast-claude", "smart-claude"]
@pytest.mark.asyncio
async def test_blocked_and_expired_keys_are_not_warmed():
from datetime import datetime, timedelta, timezone
llm_router, redis = _warming_rig(redis=FakeRedisCache())
_seed_session(redis, session_id="sess-blocked", caller_scope="hash-1", user_api_key="blocked-key")
_seed_session(redis, session_id="sess-expired", caller_scope="hash-2", user_api_key="expired-key")
_seed_session(redis, session_id="sess-live", caller_scope="hash-3", user_api_key="live-key")
prisma = FakePrismaClient(
rows=(
SimpleNamespace(token="blocked-key", spend=0.0, max_budget=None, blocked=True, expires=None),
SimpleNamespace(
token="expired-key",
spend=0.0,
max_budget=None,
blocked=None,
expires=datetime.now(timezone.utc) - timedelta(hours=1),
),
SimpleNamespace(
token="live-key",
spend=0.0,
max_budget=None,
blocked=False,
expires=datetime.now(timezone.utc) + timedelta(hours=1),
),
)
)
await _tick(llm_router, prisma=prisma)
replayed_keys = {call["metadata"]["user_api_key"] for call in llm_router.completion_calls}
assert replayed_keys == {"live-key"}
@pytest.mark.asyncio
async def test_deleted_key_sessions_are_not_warmed():
llm_router, redis = _warming_rig(redis=FakeRedisCache())
_seed_session(redis, session_id="sess-ghost", user_api_key="ghost-key")
await _tick(llm_router, prisma=FakePrismaClient(rows=()))
assert llm_router.completion_calls == []
@pytest.mark.asyncio
async def test_master_key_sessions_warm_without_a_token_row():
from litellm.constants import LITELLM_PROXY_MASTER_KEY_ALIAS
llm_router, redis = _warming_rig(redis=FakeRedisCache())
_seed_session(redis, user_api_key=LITELLM_PROXY_MASTER_KEY_ALIAS)
await _tick(llm_router, prisma=FakePrismaClient(rows=()))
assert sorted(_replayed_models(llm_router)) == ["fast-claude", "smart-claude"]
@pytest.mark.asyncio
async def test_key_without_max_budget_is_warmed():
llm_router, redis = _warming_rig(redis=FakeRedisCache())
_seed_session(redis, user_api_key="unlimited-key")
prisma = FakePrismaClient(rows=(SimpleNamespace(token="unlimited-key", spend=10_000.0, max_budget=None),))
await _tick(llm_router, prisma=prisma)
assert sorted(_replayed_models(llm_router)) == ["fast-claude", "smart-claude"]
# ---------------------------------------------------------------------------
# concurrency bound
# ---------------------------------------------------------------------------
@pytest.mark.asyncio
async def test_replays_bounded_by_semaphore():
llm_router, redis = _warming_rig(redis=FakeRedisCache(), replay_delay=0.02)
now = time.time()
for i in range(6):
_seed_session(
redis, session_id=f"sess-{i}", caller_scope=f"hash-{i}", warmth={"fast-claude": now}
)
await _tick(llm_router, refresher=CacheWarmingRefresher(max_concurrent_replays=2))
assert len(llm_router.completion_calls) == 6
assert llm_router.max_concurrent == 2

View file

@ -3584,3 +3584,149 @@ class TestCacheWarmingDispatcherRegistry:
assert new_strategy._cache_warming_ref != old_ref
assert _WARMING_STRATEGIES.get(new_strategy._cache_warming_ref) is new_strategy
assert _WARMING_STRATEGIES.get(old_ref) is None
class TestWarmAwarePick:
_POOL = ["fast-claude", "smart-claude", "cold-model"]
@staticmethod
def _router(mock_router_instance, redis, **config_overrides):
from types import SimpleNamespace
mock_router_instance.cache = SimpleNamespace(redis_cache=redis)
return ComplexityRouter(
model_name="warm-router",
litellm_router_instance=mock_router_instance,
complexity_router_config={
"tiers": {"SIMPLE": list(TestWarmAwarePick._POOL)},
"session_affinity": False,
"cache_warming": {"enabled": True},
**config_overrides,
},
)
@staticmethod
def _seed(redis, warmth, last_activity=None, served_model="fast-claude"):
import json
import time as time_module
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,
CacheWarmingAttribution,
CacheWarmingPayload,
CacheWarmingRecord,
compress_payload,
)
payload = CacheWarmingPayload(
model=served_model,
messages=({"role": "user", "content": "hello"},),
call_surface="chat_completions",
)
blob, sha = compress_payload(payload)
record = CacheWarmingRecord(
schema_version=CACHE_WARMING_RECORD_SCHEMA_VERSION,
payload_compressed=blob,
payload_sha256=sha,
token_estimate=2048,
last_activity=last_activity if last_activity is not None else time_module.time(),
served_model=served_model,
attribution=CacheWarmingAttribution(user_api_key="hash-w"),
auto_router_model_name="warm-router",
)
store = CacheWarmingStore(redis_cache=redis, auto_router_model_name="warm-router")
key = store.record_key("warm-router", "hash-w", "warm-sess")
redis.hashes.setdefault(store.sessions_key(), {})[key] = json.dumps(record.model_dump())
for model_group, stamp in warmth.items():
redis.data[CacheWarmingStore.warmth_key(key, model_group)] = json.dumps(stamp)
@staticmethod
def _kwargs():
return {"metadata": {"session_id": "warm-sess", "user_api_key_hash": "hash-w"}}
@staticmethod
def _fresh_redis():
from tests.test_litellm.router_strategy.complexity_router.cache_warming.test_store import FakeRedisCache
return FakeRedisCache()
@pytest.mark.asyncio
async def test_pick_restricted_to_warmed_members_and_served_model(self, mock_router_instance):
import time as time_module
redis = self._fresh_redis()
self._seed(redis, warmth={"smart-claude": time_module.time()}, served_model="fast-claude")
router = self._router(mock_router_instance, redis)
picks = {
await router._pick_model_for_tier(ComplexityTier.SIMPLE, None, None, self._kwargs()) for _ in range(20)
}
assert "cold-model" not in picks
assert picks <= {"fast-claude", "smart-claude"}
@pytest.mark.asyncio
async def test_stale_warm_entries_fall_back(self, mock_router_instance):
import time as time_module
redis = self._fresh_redis()
stale = time_module.time() - (270 + 60 + 5)
self._seed(redis, warmth={"smart-claude": stale}, last_activity=time_module.time() - 601)
router = self._router(mock_router_instance, redis)
assert await router._warm_aware_pick(self._POOL, self._kwargs()) is None
@pytest.mark.asyncio
async def test_served_model_counts_as_warm_only_while_session_active(self, mock_router_instance):
import time as time_module
redis = self._fresh_redis()
self._seed(redis, warmth={}, served_model="fast-claude")
router = self._router(mock_router_instance, redis)
picks = {
await router._pick_model_for_tier(ComplexityTier.SIMPLE, None, None, self._kwargs()) for _ in range(20)
}
assert picks == {"fast-claude"}
self._seed(redis, warmth={}, served_model="fast-claude", last_activity=time_module.time() - 601)
assert await router._warm_aware_pick(self._POOL, self._kwargs()) is None
@pytest.mark.asyncio
async def test_falls_back_without_session_or_record_or_redis(self, mock_router_instance):
redis = self._fresh_redis()
router = self._router(mock_router_instance, redis)
assert await router._warm_aware_pick(self._POOL, {}) is None
assert await router._warm_aware_pick(self._POOL, self._kwargs()) is None
no_redis_router = self._router(MagicMock(), None)
assert await no_redis_router._warm_aware_pick(self._POOL, self._kwargs()) is None
pick = await router._pick_model_for_tier(ComplexityTier.SIMPLE, None, None, {})
assert pick in self._POOL
@pytest.mark.asyncio
async def test_disabled_or_single_member_pool_returns_none(self, mock_router_instance):
import time as time_module
redis = self._fresh_redis()
self._seed(redis, warmth={"smart-claude": time_module.time()})
router = self._router(mock_router_instance, redis)
assert await router._warm_aware_pick(["fast-claude"], self._kwargs()) is None
disabled = self._router(MagicMock(), redis, cache_warming={"enabled": False})
assert await disabled._warm_aware_pick(self._POOL, self._kwargs()) is None
@pytest.mark.asyncio
async def test_plugin_narrowed_candidates_get_warm_pick(self, mock_router_instance):
import time as time_module
class ExcludeSmartClaude:
async def run(self, context):
context.candidate_models = [m for m in context.candidate_models if m != "smart-claude"]
return context
redis = self._fresh_redis()
self._seed(redis, warmth={"smart-claude": time_module.time()}, served_model="fast-claude")
router = self._router(mock_router_instance, redis, plugins=[ExcludeSmartClaude()])
picks = {
await router._pick_model_for_tier(
ComplexityTier.SIMPLE, None, None, self._kwargs()
)
for _ in range(20)
}
assert picks == {"fast-claude"}