From e79ea57fb487a92fe5bccdbbec79653d89825d6e Mon Sep 17 00:00:00 2001 From: Tin Chi Lo Date: Tue, 28 Jul 2026 15:30:49 -0700 Subject: [PATCH] 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. --- litellm/constants.py | 1 + litellm/proxy/proxy_server.py | 25 + .../complexity_router/README.md | 39 ++ .../cache_warming/refresher.py | 290 ++++++++++ .../complexity_router/complexity_router.py | 39 ++ .../proxy_server/test_background_health.py | 87 +++ .../cache_warming/test_refresher.py | 545 ++++++++++++++++++ .../router_strategy/test_complexity_router.py | 146 +++++ 8 files changed, 1172 insertions(+) create mode 100644 litellm/router_strategy/complexity_router/cache_warming/refresher.py create mode 100644 tests/test_litellm/router_strategy/complexity_router/cache_warming/test_refresher.py diff --git a/litellm/constants.py b/litellm/constants.py index 1014b472c61..886b8e2132c 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -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)) diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 4486cd7de59..b13537b72f4 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -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. diff --git a/litellm/router_strategy/complexity_router/README.md b/litellm/router_strategy/complexity_router/README.md index a6267453bf7..365ed6afe26 100644 --- a/litellm/router_strategy/complexity_router/README.md +++ b/litellm/router_strategy/complexity_router/README.md @@ -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 diff --git a/litellm/router_strategy/complexity_router/cache_warming/refresher.py b/litellm/router_strategy/complexity_router/cache_warming/refresher.py new file mode 100644 index 00000000000..c905c93d90f --- /dev/null +++ b/litellm/router_strategy/complexity_router/cache_warming/refresher.py @@ -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 diff --git a/litellm/router_strategy/complexity_router/complexity_router.py b/litellm/router_strategy/complexity_router/complexity_router.py index 47a81a3816f..e4bd557912a 100644 --- a/litellm/router_strategy/complexity_router/complexity_router.py +++ b/litellm/router_strategy/complexity_router/complexity_router.py @@ -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 diff --git a/tests/test_litellm/proxy/proxy_server/test_background_health.py b/tests/test_litellm/proxy/proxy_server/test_background_health.py index dca93e137ac..0bdf11eccfe 100644 --- a/tests/test_litellm/proxy/proxy_server/test_background_health.py +++ b/tests/test_litellm/proxy/proxy_server/test_background_health.py @@ -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 # --------------------------------------------------------------------------- diff --git a/tests/test_litellm/router_strategy/complexity_router/cache_warming/test_refresher.py b/tests/test_litellm/router_strategy/complexity_router/cache_warming/test_refresher.py new file mode 100644 index 00000000000..eab053d46f2 --- /dev/null +++ b/tests/test_litellm/router_strategy/complexity_router/cache_warming/test_refresher.py @@ -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 diff --git a/tests/test_litellm/router_strategy/test_complexity_router.py b/tests/test_litellm/router_strategy/test_complexity_router.py index 2d6ca515422..d913308494c 100644 --- a/tests/test_litellm/router_strategy/test_complexity_router.py +++ b/tests/test_litellm/router_strategy/test_complexity_router.py @@ -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"}