perf(router): fetch cooldown state and usage counters in one Redis round trip (#43320)

* perf(router): fetch cooldown state and usage counters in one Redis round trip

The cooldown filter (CooldownCache) and usage-based-routing-v2 selection
(LowestTPMLoggingHandler_v2) each issued their own MGET on every request
because they live in different objects. RoutingReadBatch fetches both key
sets through DualCache.async_batch_get_cache_shared while the healthy
deployments are resolved and hands the usage slice to the strategy, so
selection does not read again. Each cache keeps its own memory tier,
throttling, reservation rollback and circuit-breaker handling, and the
strategy falls back to its own read when the prefetch does not cover its
keys. simple-shuffle keeps reading only cooldowns.

aresponses no longer issues a second, blocking response-cache read from
the worker thread that runs the sync wrapper.

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(caching): keep per-cache tier failures inside the shared batch read

Wrap the memory-tier prepare and backfill steps of DualCache.async_batch_get_cache_shared
so a failing tier degrades that cache's read to None the way async_batch_get_cache does,
instead of escaping into routing. Drop the aresponses sync-cache guard: for native
Responses models the worker-thread read is the one whose key matches the write, so
skipping it broke cached /v1/responses replays.

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* refactor(router): rename usage key builder so the async cache-call check reads it as a key helper

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(alerting): narrow daily-report cache values before numeric comparison

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* refactor(caching): type the shared batch-read helpers and merge Redis results without mutation

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* style: fix import sort in test_dual_cache

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* refactor(caching): flatten shared batch read keys without a stacked comprehension

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

---------

Co-authored-by: yassin <yassin@berri.ai>
Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
devin-ai-integration[bot] 2026-09-29 13:52:20 -07:00 • committed by GitHub
parent 27c110cb71
commit d2a574b791
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
9 changed files with 622 additions and 102 deletions

View file

@ -8,9 +8,11 @@ Has 4 primary methods:
- async_get_cache
"""
import itertools
import logging
import time
from collections.abc import Sequence
from dataclasses import dataclass
from threading import Lock
from typing import TYPE_CHECKING, Any, Final
@ -47,6 +49,16 @@ class LimitedSizeOrderedDict(OrderedDict):
super().__setitem__(key, value)
@dataclass(frozen=True)
class PendingBatchRead:
"""A batch read that has consulted the in-memory tier and reserved its Redis keys, but not hit Redis yet."""
keys: list[str]
result: list[object | None]
redis_keys: list[str]
previous_access_times: dict[str, float | None]
class DualCache(BaseCache):
"""
DualCache is a cache implementation that updates both Redis and an in-memory cache simultaneously.
@ -301,6 +313,37 @@ class DualCache(BaseCache):
else:
self.last_redis_batch_access_time[key] = previous_time
async def _prepare_batch_get(self, keys: list[str], local_only: bool, **kwargs: object) -> PendingBatchRead:
result: list[object | None] = [None] * len(keys)
if self.in_memory_cache is not None:
in_memory_result: Final = await self.in_memory_cache.async_batch_get_cache(keys, **kwargs)
if in_memory_result is not None:
result = in_memory_result
redis_keys: list[str] = []
previous_access_times: dict[str, float | None] = {}
if None in result and self.redis_cache is not None and local_only is False:
redis_keys, previous_access_times = self._reserve_redis_batch_keys(time.time(), keys, result)
return PendingBatchRead(
keys=keys, result=result, redis_keys=redis_keys, previous_access_times=previous_access_times
)
async def _apply_batch_get(
self, pending: PendingBatchRead, redis_result: dict[str, object] | None, **kwargs: object
) -> list[object | None]:
if redis_result is None or all(v is None for v in redis_result.values()):
return pending.result
merged: Final[list[object | None]] = [
redis_result.get(key, value) for key, value in zip(pending.keys, pending.result)
]
if self.in_memory_cache is not None:
for key, value in redis_result.items():
if value is not None:
await self.in_memory_cache.async_set_cache(key, value, **self._backfill_kwargs(kwargs))
return merged
async def async_batch_get_cache(
self,
keys: list,
@ -309,51 +352,22 @@ class DualCache(BaseCache):
**kwargs,
):
try:
result = [None] * len(keys)
if self.in_memory_cache is not None:
in_memory_result: Final = await self.in_memory_cache.async_batch_get_cache(keys, **kwargs)
if in_memory_result is not None:
result = in_memory_result
if None in result and self.redis_cache is not None and local_only is False:
"""
- for the none values in the result
- check the redis cache
"""
current_time: Final = time.time()
sublist_keys, previous_access_times = self._reserve_redis_batch_keys(current_time, keys, result)
# Only hit Redis if enough time has passed since last access.
if len(sublist_keys) > 0:
try:
# If not found in in-memory cache, try fetching from Redis
redis_result: Final = await self.redis_cache.async_batch_get_cache(
sublist_keys, parent_otel_span=parent_otel_span
)
except Exception as e:
# Do not throttle subsequent callers if the Redis read fails.
self._rollback_redis_batch_key_reservations(previous_access_times)
if isinstance(e, RedisCircuitBreakerOpenError):
verbose_logger.debug("LiteLLM Cache: async_batch_get_cache served from memory only: %s", e)
return result
raise
# Short-circuit if redis_result is None or contains only None values
if redis_result is None or all(v is None for v in redis_result.values()):
return result
# Pre-compute key-to-index mapping for O(1) lookup
key_to_index: Final = {key: i for i, key in enumerate(keys)}
# Update both result and in-memory cache in a single loop
for key, value in redis_result.items():
result[key_to_index[key]] = value
if value is not None and self.in_memory_cache is not None:
await self.in_memory_cache.async_set_cache(key, value, **self._backfill_kwargs(kwargs))
return result
pending: Final = await self._prepare_batch_get(keys, local_only, **kwargs)
# Only hit Redis for keys memory could not serve and enough time has passed since last access.
if not pending.redis_keys or self.redis_cache is None:
return pending.result
try:
redis_result: Final = await self.redis_cache.async_batch_get_cache(
pending.redis_keys, parent_otel_span=parent_otel_span
)
except Exception as e:
# Do not throttle subsequent callers if the Redis read fails.
self._rollback_redis_batch_key_reservations(pending.previous_access_times)
if isinstance(e, RedisCircuitBreakerOpenError):
verbose_logger.debug("LiteLLM Cache: async_batch_get_cache served from memory only: %s", e)
return pending.result
raise
return await self._apply_batch_get(pending, redis_result, **kwargs)
except Exception as e:
log_redis_failure(
verbose_logger,
@ -363,6 +377,74 @@ class DualCache(BaseCache):
with_traceback=True,
)
@staticmethod
async def async_batch_get_cache_shared(
reads: Sequence[tuple["DualCache", list[str]]],
parent_otel_span: Span | None = None,
) -> list[list[object | None] | None]:
"""
`async_batch_get_cache` for several caches in one Redis round trip.
Each cache still serves what it can from its own in-memory tier, applies its own Redis read
throttle and backfills its own memory; only the Redis MGET is shared. A failed MGET is reported
to every cache that took part in it exactly as its own failed `async_batch_get_cache` would be:
None when the read raised, the in-memory result when the circuit breaker is open. A cache whose
Redis client is not the one the first cache uses falls back to its own read.
"""
results: Final[list[list[object | None] | None]] = [None] * len(reads)
shared_redis: Final = reads[0][0].redis_cache if reads else None
pendings: Final[list[tuple[int, DualCache, PendingBatchRead]]] = []
for index, (cache, keys) in enumerate(reads):
if shared_redis is None or cache.redis_cache is not shared_redis:
results[index] = await cache.async_batch_get_cache(keys=keys, parent_otel_span=parent_otel_span)
continue
try:
pending = await cache._prepare_batch_get(keys, local_only=False)
except Exception as e:
DualCache._log_shared_batch_get_failure(e)
continue
pendings.append((index, cache, pending))
results[index] = pending.result
redis_keys: Final = list(
dict.fromkeys(itertools.chain.from_iterable(pending.redis_keys for _, _, pending in pendings))
)
if shared_redis is None or not redis_keys:
return results
try:
redis_result: Final = await shared_redis.async_batch_get_cache(
redis_keys, parent_otel_span=parent_otel_span
)
except Exception as e:
for index, cache, pending in pendings:
cache._rollback_redis_batch_key_reservations(pending.previous_access_times)
if pending.redis_keys and not isinstance(e, RedisCircuitBreakerOpenError):
results[index] = None
if isinstance(e, RedisCircuitBreakerOpenError):
verbose_logger.debug("LiteLLM Cache: async_batch_get_cache_shared served from memory only: %s", e)
else:
DualCache._log_shared_batch_get_failure(e)
return results
for index, cache, pending in pendings:
own_result = {key: redis_result[key] for key in pending.redis_keys if key in redis_result}
try:
results[index] = await cache._apply_batch_get(pending, own_result)
except Exception as e:
results[index] = None
DualCache._log_shared_batch_get_failure(e)
return results
@staticmethod
def _log_shared_batch_get_failure(e: Exception) -> None:
log_redis_failure(
verbose_logger,
logging.ERROR,
"LiteLLM Cache: exception in async_batch_get_cache_shared",
e,
with_traceback=True,
)
async def async_set_cache(self, key, value, local_only: bool = False, **kwargs):
print_verbose(f"async set cache: cache key: {key}; local_only: {local_only}; value: {value}")
try:

View file

@ -376,8 +376,12 @@ class SlackAlerting(CustomBatchLogger):
if combined_metrics_values is None:
return False
metric_values: Final[list[float | None]] = [
val if isinstance(val, (int, float)) else None for val in combined_metrics_values
]
all_none = True
for val in combined_metrics_values:
for val in metric_values:
if val is not None and val > 0:
all_none = False
break
@ -385,8 +389,8 @@ class SlackAlerting(CustomBatchLogger):
if all_none:
return False
failed_request_values: Final = combined_metrics_values[: len(failed_request_keys)] # # [1, 2, None, ..]
latency_values: Final = combined_metrics_values[len(failed_request_keys) :]
failed_request_values: Final = metric_values[: len(failed_request_keys)] # # [1, 2, None, ..]
latency_values: Final = metric_values[len(failed_request_keys) :]
# find top 5 failed
## Replace None values with a placeholder value (-1 in this case)

View file

@ -134,7 +134,7 @@ from litellm.router_strategy.least_busy import LeastBusyLoggingHandler
from litellm.router_strategy.lowest_cost import LowestCostLoggingHandler
from litellm.router_strategy.lowest_latency import LowestLatencyLoggingHandler
from litellm.router_strategy.lowest_tpm_rpm import LowestTPMLoggingHandler
from litellm.router_strategy.lowest_tpm_rpm_v2 import LowestTPMLoggingHandler_v2
from litellm.router_strategy.lowest_tpm_rpm_v2 import LowestTPMLoggingHandler_v2, PrefetchedUsage
from litellm.router_strategy.simple_shuffle import simple_shuffle
from litellm.router_strategy.tag_based_routing import (
_get_tags_from_request_kwargs,
@ -259,6 +259,7 @@ from litellm.router_utils.routing_groups import (
parse_routing_groups,
validate_routing_strategy,
)
from litellm.router_utils.routing_read_batch import RoutingReadBatch
from litellm.scheduler import FlowItem, Scheduler
from litellm.types.litellm_params import RoutingStrategyName
from litellm.types.llms.openai import (
@ -1789,6 +1790,7 @@ class Router:
messages: list[dict[str, str]] | None,
input: str | list | None,
request_kwargs: dict | None,
prefetched_usage: PrefetchedUsage | None = None,
) -> Any | None:
"""
Asks the strategy selector for a deployment. Caller handles
@ -1814,6 +1816,14 @@ class Router:
messages=messages,
input=input,
)
case "usage-based-routing-v2" if isinstance(selector, LowestTPMLoggingHandler_v2):
return await selector.async_get_available_deployments(
model_group=model,
healthy_deployments=healthy_deployments,
messages=messages,
input=input,
prefetched_usage=prefetched_usage,
)
case "usage-based-routing-v2" | "cost-based-routing":
return await selector.async_get_available_deployments(
model_group=model,
@ -12925,6 +12935,7 @@ class Router:
specific_deployment: bool | None = False,
parent_otel_span: Span | None = None,
health_check_probe: bool = False,
routing_read_batch: RoutingReadBatch | None = None,
) -> list[dict] | dict:
"""
Get the healthy deployments for a model.
@ -12977,8 +12988,14 @@ class Router:
health_check_probe=health_check_probe,
)
cooldown_deployments: Final = await _async_get_cooldown_deployments(
litellm_router_instance=self, parent_otel_span=parent_otel_span
cooldown_deployments: Final = (
await _async_get_cooldown_deployments(litellm_router_instance=self, parent_otel_span=parent_otel_span)
if routing_read_batch is None
else await routing_read_batch.async_get_cooldown_deployments(
litellm_router_instance=self,
healthy_deployments=healthy_deployments,
parent_otel_span=parent_otel_span,
)
)
if verbose_router_logger.isEnabledFor(logging.DEBUG):
verbose_router_logger.debug("cooldown deployments: %s", cooldown_deployments)
@ -13256,6 +13273,7 @@ class Router:
# the hook can replace `model` and routing-group lookup must key
# off the final model name.
strategy, strategy_selector = self._get_routing_context(model, request_kwargs)
routing_read_batch: Final = RoutingReadBatch.for_strategy(strategy, strategy_selector)
healthy_deployments: Final = await self.async_get_healthy_deployments(
model=model,
@ -13264,6 +13282,7 @@ class Router:
input=input,
specific_deployment=specific_deployment,
parent_otel_span=parent_otel_span,
routing_read_batch=routing_read_batch,
)
if isinstance(healthy_deployments, dict):
await self._async_override_selector_pre_call_check(
@ -13294,6 +13313,7 @@ class Router:
messages=messages,
input=input,
request_kwargs=request_kwargs,
prefetched_usage=routing_read_batch.prefetched_usage if routing_read_batch is not None else None,
)
if deployment is None:
exception: Final = await async_raise_no_deployment_exception(

View file

@ -1,7 +1,8 @@
#### What this does ####
# identifies lowest tpm deployment
import random
from collections.abc import Sequence
from collections.abc import Mapping, Sequence
from dataclasses import dataclass
from typing import TYPE_CHECKING, Any, Final
import httpx
@ -31,6 +32,26 @@ class RoutingArgs(LiteLLMPydanticObjectBase):
ttl: int = 1 * 60 # 1min (RPM/TPM expire key)
@dataclass(frozen=True)
class PrefetchedUsage:
"""
tpm/rpm counter values another read of this request already fetched from the router cache.
`values` is None when that read failed, which is what `async_batch_get_cache` returns on failure.
"""
keys: frozenset[str]
values: Mapping[str, object] | None
def covers(self, keys: Sequence[str]) -> bool:
return self.keys.issuperset(keys)
def values_for(self, keys: Sequence[str]) -> list[object | None] | None:
if self.values is None:
return None
return [self.values.get(key) for key in keys]
class LowestTPMLoggingHandler_v2(BaseRoutingStrategy, CustomLogger):
"""
Updated version of TPM/RPM Logging.
@ -412,17 +433,35 @@ class LowestTPMLoggingHandler_v2(BaseRoutingStrategy, CustomLogger):
else:
return None
def usage_counter_keys(self, healthy_deployments: list) -> tuple[list[str], list[str]]:
"""The `<id>:<model>:tpm:<HH-MM>` and `<id>:<model>:rpm:<HH-MM>` counter keys selection reads."""
current_minute: Final = get_utc_datetime().strftime("%H-%M")
tpm_keys: Final[list[str]] = []
rpm_keys: Final[list[str]] = []
for m in healthy_deployments:
if isinstance(m, dict):
id = m.get("model_info", {}).get(
"id"
) # a deployment should always have an 'id'. this is set in router.py
deployment_name = m.get("litellm_params", {}).get("model")
tpm_keys.append(f"{id}:{deployment_name}:tpm:{current_minute}")
rpm_keys.append(f"{id}:{deployment_name}:rpm:{current_minute}")
return tpm_keys, rpm_keys
async def async_get_available_deployments(
self,
model_group: str,
healthy_deployments: list,
messages: list[dict[str, str]] | None = None,
input: str | list | None = None,
prefetched_usage: PrefetchedUsage | None = None,
):
"""
Async implementation of get deployments.
Reduces time to retrieve the tpm/rpm values from cache
Reduces time to retrieve the tpm/rpm values from cache. `prefetched_usage` skips the cache
read when it already holds this request's counters (see `RoutingReadBatch`).
"""
# get list of potential deployments
verbose_router_logger.debug(
@ -431,28 +470,15 @@ class LowestTPMLoggingHandler_v2(BaseRoutingStrategy, CustomLogger):
healthy_deployments,
)
dt: Final = get_utc_datetime()
current_minute: Final = dt.strftime("%H-%M")
tpm_keys: Final = []
rpm_keys: Final = []
for m in healthy_deployments:
if isinstance(m, dict):
id = m.get("model_info", {}).get(
"id"
) # a deployment should always have an 'id'. this is set in router.py
deployment_name = m.get("litellm_params", {}).get("model")
tpm_key = f"{id}:{deployment_name}:tpm:{current_minute}"
rpm_key = f"{id}:{deployment_name}:rpm:{current_minute}"
tpm_keys.append(tpm_key)
rpm_keys.append(rpm_key)
tpm_keys, rpm_keys = self.usage_counter_keys(healthy_deployments)
combined_tpm_rpm_keys: Final = tpm_keys + rpm_keys
combined_tpm_rpm_values: Final = await self.router_cache.async_batch_get_cache(
keys=combined_tpm_rpm_keys
) # [1, 2, None, ..]
if prefetched_usage is not None and prefetched_usage.covers(combined_tpm_rpm_keys):
combined_tpm_rpm_values = prefetched_usage.values_for(combined_tpm_rpm_keys)
else:
combined_tpm_rpm_values = await self.router_cache.async_batch_get_cache(
keys=combined_tpm_rpm_keys
) # [1, 2, None, ..]
if combined_tpm_rpm_values is not None:
tpm_values = combined_tpm_rpm_values[: len(tpm_keys)]

View file

@ -4,7 +4,7 @@ Wrapper around router cache. Meant to handle model cooldown logic
import functools
import time
from collections.abc import Mapping
from collections.abc import Mapping, Sequence
from typing import TYPE_CHECKING, Any, Final
from typing_extensions import TypedDict
@ -163,6 +163,12 @@ class CooldownCache:
keys: Final = [CooldownCache.get_cooldown_cache_key(model_id) for model_id in model_ids]
results: Final = await self.cooldown_store.async_batch_get_cache(keys=keys, parent_otel_span=parent_otel_span)
return self.active_cooldowns_from_results(model_ids, results)
def active_cooldowns_from_results(
self, model_ids: list[str], results: Sequence[object] | None
) -> list[tuple[str, CooldownCacheValue]]:
"""The cooldowns still active in a `cooldown_store` batch read of `get_cooldown_cache_key(model_id)` per id."""
active_cooldowns: Final[list[tuple[str, CooldownCacheValue]]] = []
if results is None or all(v is None for v in results):

View file

@ -0,0 +1,72 @@
"""
One Redis round trip for the reads a request needs before a deployment can be picked.
The cooldown filter (`CooldownCache`, its own `DualCache`) and usage-based selection
(`LowestTPMLoggingHandler_v2`, the router cache) each issue their own MGET because they live in
different objects. `RoutingReadBatch` fetches both key sets in one
`DualCache.async_batch_get_cache_shared` while the healthy deployments are being resolved and hands
the usage slice to the strategy, so selection does not read again.
"""
from typing import TYPE_CHECKING, Any, Final
from litellm._logging import verbose_router_logger
from litellm.caching.dual_cache import DualCache
from litellm.router_strategy.lowest_tpm_rpm_v2 import LowestTPMLoggingHandler_v2, PrefetchedUsage
from litellm.router_utils.cooldown_cache import CooldownCache
if TYPE_CHECKING:
from opentelemetry.trace import Span as _Span
from litellm.router import Router as _Router
LitellmRouter = _Router
Span = _Span
else:
LitellmRouter = Any
Span = Any
class RoutingReadBatch:
def __init__(self, usage_selector: LowestTPMLoggingHandler_v2) -> None:
self.usage_selector: Final = usage_selector
self.prefetched_usage: PrefetchedUsage | None = None
@staticmethod
def for_strategy(strategy: str | None, selector: object) -> "RoutingReadBatch | None":
if strategy == "usage-based-routing-v2" and isinstance(selector, LowestTPMLoggingHandler_v2):
return RoutingReadBatch(usage_selector=selector)
return None
async def async_get_cooldown_deployments(
self,
litellm_router_instance: LitellmRouter,
healthy_deployments: list,
parent_otel_span: Span | None,
) -> list[str]:
"""
`_async_get_cooldown_deployments`, with the strategy's tpm/rpm counters for
`healthy_deployments` fetched in the same MGET and kept as `prefetched_usage`.
"""
model_ids: Final = litellm_router_instance.get_model_ids()
cooldown_keys: Final = [CooldownCache.get_cooldown_cache_key(model_id) for model_id in model_ids]
tpm_keys, rpm_keys = self.usage_selector.usage_counter_keys(healthy_deployments)
usage_keys: Final = tpm_keys + rpm_keys
cooldown_results, usage_values = await DualCache.async_batch_get_cache_shared(
[
(litellm_router_instance.cooldown_cache.cooldown_store, cooldown_keys),
(self.usage_selector.router_cache, usage_keys),
],
parent_otel_span=parent_otel_span,
)
self.prefetched_usage = PrefetchedUsage(
keys=frozenset(usage_keys),
values=None if usage_values is None else dict(zip(usage_keys, usage_values)),
)
cooldown_models: Final = litellm_router_instance.cooldown_cache.active_cooldowns_from_results(
model_ids, cooldown_results
)
verbose_router_logger.debug("retrieve cooldown models: %s", cooldown_models)
return [model_id for model_id, _ in cooldown_models]

View file

@ -6,18 +6,16 @@ from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from litellm.constants import DEFAULT_MAX_REDIS_BATCH_CACHE_SIZE
from litellm.caching.dual_cache import DualCache
from litellm.caching.in_memory_cache import InMemoryCache
from litellm.caching.redis_cache import RedisCache, _redis_circuit_breaker_guard, _redis_circuit_breaker_guard_sync
from litellm.constants import DEFAULT_MAX_REDIS_BATCH_CACHE_SIZE
from litellm.types.caching import RedisPipelineIncrementOperation
@pytest.mark.asyncio
async def test_dual_cache_async_batch_get_cache_coalesces_concurrent_redis_reads():
dual_cache = DualCache(
redis_cache=MagicMock(spec=RedisCache), default_redis_batch_cache_expiry=10
)
dual_cache = DualCache(redis_cache=MagicMock(spec=RedisCache), default_redis_batch_cache_expiry=10)
keys = ["shared_a", "shared_b"]
start_gate = asyncio.Event()
@ -44,9 +42,7 @@ async def test_dual_cache_async_batch_get_cache_coalesces_concurrent_redis_reads
@pytest.mark.asyncio
async def test_dual_cache_async_batch_get_cache_rolls_back_redis_reservation_on_error():
dual_cache = DualCache(
redis_cache=MagicMock(spec=RedisCache), default_redis_batch_cache_expiry=10
)
dual_cache = DualCache(redis_cache=MagicMock(spec=RedisCache), default_redis_batch_cache_expiry=10)
keys = ["shared_a", "shared_b"]
with patch.object(
@ -116,9 +112,7 @@ def test_dual_cache_batch_get_cache_only_reads_missing_keys_from_redis():
def test_dual_cache_batch_get_cache_throttles_repeat_redis_reads():
mock_redis = _redis_mock_for_sync_batch({"absent_key": None})
dual_cache = DualCache(
in_memory_cache=InMemoryCache(), redis_cache=mock_redis, default_redis_batch_cache_expiry=10
)
dual_cache = DualCache(in_memory_cache=InMemoryCache(), redis_cache=mock_redis, default_redis_batch_cache_expiry=10)
first = dual_cache.batch_get_cache(keys=["absent_key"])
second = dual_cache.batch_get_cache(keys=["absent_key"])
@ -131,9 +125,7 @@ def test_dual_cache_batch_get_cache_throttles_repeat_redis_reads():
def test_dual_cache_batch_get_cache_rolls_back_redis_reservation_on_error():
mock_redis = MagicMock(spec=RedisCache)
mock_redis.batch_get_cache.side_effect = RuntimeError("redis unavailable")
dual_cache = DualCache(
in_memory_cache=InMemoryCache(), redis_cache=mock_redis, default_redis_batch_cache_expiry=10
)
dual_cache = DualCache(in_memory_cache=InMemoryCache(), redis_cache=mock_redis, default_redis_batch_cache_expiry=10)
first_result = dual_cache.batch_get_cache(keys=["shared_a"])
second_result = dual_cache.batch_get_cache(keys=["shared_a"])
@ -146,9 +138,7 @@ def test_dual_cache_batch_get_cache_rolls_back_redis_reservation_on_error():
def test_dual_cache_batch_get_cache_returns_memory_only_when_redis_read_is_throttled():
mock_redis = _redis_mock_for_sync_batch({"throttled_key": "redis_value"})
dual_cache = DualCache(
in_memory_cache=InMemoryCache(), redis_cache=mock_redis, default_redis_batch_cache_expiry=10
)
dual_cache = DualCache(in_memory_cache=InMemoryCache(), redis_cache=mock_redis, default_redis_batch_cache_expiry=10)
dual_cache.last_redis_batch_access_time["throttled_key"] = time.time()
result = dual_cache.batch_get_cache(keys=["throttled_key"])
@ -257,9 +247,7 @@ async def test_dual_cache_batch_redis_backfill_injects_default_in_memory_ttl():
default_in_memory_ttl, same as the single-key path."""
in_memory_cache = InMemoryCache(default_ttl=600)
mock_redis = MagicMock(spec=RedisCache)
mock_redis.async_batch_get_cache = AsyncMock(
return_value={"batch_backfill_key": "redis_value"}
)
mock_redis.async_batch_get_cache = AsyncMock(return_value={"batch_backfill_key": "redis_value"})
dual_cache = DualCache(
in_memory_cache=in_memory_cache,
redis_cache=mock_redis,
@ -371,9 +359,7 @@ async def test_circuit_breaker_open_skips_redis():
class FakeRedis:
def __init__(self):
self._circuit_breaker = RedisCircuitBreaker(
failure_threshold=3, recovery_timeout=60
)
self._circuit_breaker = RedisCircuitBreaker(failure_threshold=3, recovery_timeout=60)
self._circuit_breaker._state = "open"
self._circuit_breaker._opened_at = time.time()
self.call_count = 0
@ -426,9 +412,7 @@ def test_circuit_breaker_half_open_concurrent_calls_are_fast_failed():
# All subsequent concurrent callers: HALF_OPEN → fast-fail (return True)
for _ in range(10):
assert (
cb.is_open() is True
), "concurrent callers should be fast-failed in HALF_OPEN"
assert cb.is_open() is True, "concurrent callers should be fast-failed in HALF_OPEN"
def test_circuit_breaker_disabled_never_opens():
@ -472,9 +456,7 @@ async def test_circuit_breaker_disabled_guard_always_calls_method():
class FakeRedis:
def __init__(self):
self._circuit_breaker = RedisCircuitBreaker(
failure_threshold=1, recovery_timeout=60, enabled=False
)
self._circuit_breaker = RedisCircuitBreaker(failure_threshold=1, recovery_timeout=60, enabled=False)
self.call_count = 0
@_redis_circuit_breaker_guard
@ -791,3 +773,125 @@ async def test_async_delete_cache_keys_on_empty_list_touches_no_backend():
await dual_cache.async_delete_cache_keys([])
redis_cache.delete_cache_keys.assert_not_awaited()
def _recording_redis(values: dict) -> MagicMock:
redis = MagicMock(spec=RedisCache)
redis.async_batch_get_cache = AsyncMock(
side_effect=lambda key_list, parent_otel_span=None: {key: values.get(key) for key in key_list}
)
return redis
@pytest.mark.asyncio
async def test_shared_batch_read_issues_one_mget_for_two_caches_and_backfills_each_one_separately():
redis = _recording_redis({"a1": 1, "b2": "x"})
first = DualCache(in_memory_cache=InMemoryCache(), redis_cache=redis, default_redis_batch_cache_expiry=10)
second = DualCache(in_memory_cache=InMemoryCache(), redis_cache=redis, default_redis_batch_cache_expiry=10)
results = await DualCache.async_batch_get_cache_shared([(first, ["a1", "a2"]), (second, ["b1", "b2"])])
assert results == [[1, None], [None, "x"]]
assert redis.async_batch_get_cache.await_count == 1
assert redis.async_batch_get_cache.await_args.args[0] == ["a1", "a2", "b1", "b2"]
assert first.in_memory_cache.get_cache("a1") == 1
assert second.in_memory_cache.get_cache("b2") == "x"
assert first.in_memory_cache.get_cache("b2") is None, "backfill leaked into the other cache"
@pytest.mark.asyncio
async def test_shared_batch_read_serves_memory_hits_and_throttles_like_the_separate_reads():
redis = _recording_redis({"a2": 2, "b1": 3})
first = DualCache(in_memory_cache=InMemoryCache(), redis_cache=redis, default_redis_batch_cache_expiry=10)
first.in_memory_cache.set_cache("a1", 5)
second = DualCache(in_memory_cache=InMemoryCache(), redis_cache=redis, default_redis_batch_cache_expiry=10)
second.in_memory_cache.set_cache("b1", 3)
results = await DualCache.async_batch_get_cache_shared([(first, ["a1", "a2"]), (second, ["b1"])])
assert results == [[5, 2], [3]]
assert redis.async_batch_get_cache.await_args.args[0] == ["a2"], "memory hits must not hit Redis"
first.in_memory_cache.delete_cache("a2")
results = await DualCache.async_batch_get_cache_shared([(first, ["a1", "a2"]), (second, ["b1"])])
assert results == [[5, None], [3]]
assert redis.async_batch_get_cache.await_count == 1, "a2 was read within the batch expiry, so it is throttled"
@pytest.mark.asyncio
async def test_shared_batch_read_failure_degrades_exactly_like_two_failed_reads():
redis = MagicMock(spec=RedisCache)
redis.async_batch_get_cache = AsyncMock(side_effect=ConnectionError("redis unavailable"))
first = DualCache(in_memory_cache=InMemoryCache(), redis_cache=redis, default_redis_batch_cache_expiry=10)
second = DualCache(in_memory_cache=InMemoryCache(), redis_cache=redis, default_redis_batch_cache_expiry=10)
third = DualCache(in_memory_cache=InMemoryCache(), redis_cache=redis, default_redis_batch_cache_expiry=10)
third.in_memory_cache.set_cache("c1", "memory")
shared = await DualCache.async_batch_get_cache_shared([(first, ["a1"]), (second, ["b1"]), (third, ["c1"])])
separate = [
await first.async_batch_get_cache(keys=["a1"]),
await second.async_batch_get_cache(keys=["b1"]),
await third.async_batch_get_cache(keys=["c1"]),
]
assert shared == separate == [None, None, ["memory"]]
assert "a1" not in first.last_redis_batch_access_time
assert "b1" not in second.last_redis_batch_access_time
@pytest.mark.asyncio
async def test_shared_batch_read_with_an_open_breaker_keeps_memory_hits_and_releases_reservations():
first = _dual_cache_with_open_breaker_and_a_memory_hit()
second = DualCache(
in_memory_cache=InMemoryCache(), redis_cache=first.redis_cache, default_redis_batch_cache_expiry=10
)
results = await DualCache.async_batch_get_cache_shared([(first, ["k1", "k2"]), (second, ["k3"])])
assert results == [["v1", None], [None]]
assert "k2" not in first.last_redis_batch_access_time
assert "k3" not in second.last_redis_batch_access_time
@pytest.mark.asyncio
async def test_shared_batch_read_falls_back_to_a_caches_own_read_when_its_redis_client_differs():
first_redis = _recording_redis({"a1": 1})
second_redis = _recording_redis({"b1": 2})
first = DualCache(in_memory_cache=InMemoryCache(), redis_cache=first_redis, default_redis_batch_cache_expiry=10)
second = DualCache(in_memory_cache=InMemoryCache(), redis_cache=second_redis, default_redis_batch_cache_expiry=10)
memory_only = DualCache(in_memory_cache=InMemoryCache(), redis_cache=None)
memory_only.in_memory_cache.set_cache("m1", "m")
results = await DualCache.async_batch_get_cache_shared(
[(first, ["a1"]), (second, ["b1"]), (memory_only, ["m1", "m2"])]
)
assert results == [[1], [2], ["m", None]]
assert first_redis.async_batch_get_cache.await_args.args[0] == ["a1"]
assert second_redis.async_batch_get_cache.await_args.args[0] == ["b1"]
@pytest.mark.asyncio
async def test_shared_batch_read_keeps_a_caches_own_tier_failure_to_itself_like_the_separate_read():
redis = _recording_redis({"a1": 1, "b1": 2, "c1": 3})
broken_memory_read = DualCache(
in_memory_cache=InMemoryCache(), redis_cache=redis, default_redis_batch_cache_expiry=10
)
broken_memory_read.in_memory_cache.async_batch_get_cache = AsyncMock(side_effect=RuntimeError("memory read"))
broken_backfill = DualCache(in_memory_cache=InMemoryCache(), redis_cache=redis, default_redis_batch_cache_expiry=10)
broken_backfill.in_memory_cache.async_set_cache = AsyncMock(side_effect=RuntimeError("memory write"))
healthy = DualCache(in_memory_cache=InMemoryCache(), redis_cache=redis, default_redis_batch_cache_expiry=10)
shared = await DualCache.async_batch_get_cache_shared(
[(broken_memory_read, ["a1"]), (broken_backfill, ["b1"]), (healthy, ["c1"])]
)
broken_backfill.last_redis_batch_access_time.clear()
separate = [
await broken_memory_read.async_batch_get_cache(keys=["a1"]),
await broken_backfill.async_batch_get_cache(keys=["b1"]),
await healthy.async_batch_get_cache(keys=["c1"]),
]
assert shared == separate == [None, None, [3]]
assert redis.async_batch_get_cache.await_args_list[0].args[0] == ["b1", "c1"]

View file

@ -1,7 +1,12 @@
from datetime import datetime, timedelta
from typing import Final
from unittest.mock import AsyncMock
import pytest
from litellm import Router
from litellm.caching.dual_cache import DualCache
from litellm.router_strategy.lowest_tpm_rpm_v2 import LowestTPMLoggingHandler_v2, PrefetchedUsage
from litellm.types.router import DeploymentTypedDict, LiteLLMParamsTypedDict
MODEL_GROUP: Final = "lowest-tpm-router"
@ -52,3 +57,26 @@ def test_usage_based_routing_v1_selects_the_lowest_recorded_tpm() -> None:
)
assert deployment["model_info"]["id"] == LOW_USAGE_DEPLOYMENT_ID
@pytest.mark.asyncio
async def test_v2_async_selection_uses_prefetched_counters_only_when_they_cover_its_keys():
router_cache = DualCache()
router_cache.async_batch_get_cache = AsyncMock(return_value=[100, 10, None, None]) # type: ignore[method-assign]
strategy = LowestTPMLoggingHandler_v2(router_cache=router_cache)
deployments = [
{"model_name": "g", "litellm_params": {"model": "m"}, "model_info": {"id": "a"}},
{"model_name": "g", "litellm_params": {"model": "m"}, "model_info": {"id": "b"}},
]
tpm_keys, rpm_keys = strategy.usage_counter_keys(deployments)
keys = tpm_keys + rpm_keys
covering = PrefetchedUsage(keys=frozenset(keys), values=dict(zip(keys, [10, 100, None, None])))
chosen = await strategy.async_get_available_deployments(model_group="g", healthy_deployments=deployments, prefetched_usage=covering)
assert chosen["model_info"]["id"] == "a", "the prefetched counters say a is the lowest"
router_cache.async_batch_get_cache.assert_not_awaited()
stale = PrefetchedUsage(keys=frozenset(keys[:1]), values={keys[0]: 10})
chosen = await strategy.async_get_available_deployments(model_group="g", healthy_deployments=deployments, prefetched_usage=stale)
assert chosen["model_info"]["id"] == "b", "counters that do not cover this minute's keys are read again"
router_cache.async_batch_get_cache.assert_awaited_once_with(keys=keys)

View file

@ -0,0 +1,178 @@
"""
One Redis round trip per request for the router's pre-call reads.
Before `RoutingReadBatch`, `async_get_available_deployment` issued one MGET for the cooldown keys
(`CooldownCache`) and a second one for the tpm/rpm counters (`LowestTPMLoggingHandler_v2`).
"""
import time
from unittest.mock import AsyncMock, MagicMock
import pytest
import litellm
from litellm import Router
from litellm.caching.redis_cache import RedisCache
_MODEL_GROUP = "claude"
_MESSAGES = [{"role": "user", "content": "ping"}]
def _deployment(deployment_id: str) -> dict:
return {
"model_name": _MODEL_GROUP,
"litellm_params": {"model": "anthropic/claude-x", "api_key": "test", "mock_response": "pong"},
"model_info": {"id": deployment_id},
}
def _redis_answering(values_by_key_prefix: dict[str, object]) -> MagicMock:
"""A Redis double that answers each key from its minute-less prefix and records every MGET."""
def _mget(key_list, parent_otel_span=None):
return {key: values_by_key_prefix.get(key.rsplit(":", 1)[0], values_by_key_prefix.get(key)) for key in key_list}
redis = MagicMock(spec=RedisCache)
redis.async_batch_get_cache = AsyncMock(side_effect=_mget)
return redis
def _router(redis: MagicMock, routing_strategy: str) -> Router:
router = Router(
model_list=[_deployment("dep-a"), _deployment("dep-b")],
routing_strategy=routing_strategy,
)
router._update_redis_cache(cache=redis)
return router
def _redis_key_families(redis: MagicMock) -> list[list[str]]:
return [
sorted(key.rsplit(":", 1)[0] if ":tpm:" in key or ":rpm:" in key else key for key in call.args[0])
for call in redis.async_batch_get_cache.await_args_list
]
def _cooldown(seconds: float) -> dict:
return {"exception_received": "429", "status_code": "429", "timestamp": time.time(), "cooldown_time": seconds}
@pytest.mark.asyncio
async def test_usage_based_routing_reads_cooldowns_and_counters_in_one_redis_round_trip():
redis = _redis_answering({})
router = _router(redis, "usage-based-routing-v2")
deployment = await router.async_get_available_deployment(
model=_MODEL_GROUP, request_kwargs={}, messages=_MESSAGES
)
assert deployment["model_info"]["id"] in {"dep-a", "dep-b"}
assert _redis_key_families(redis) == [
[
"dep-a:anthropic/claude-x:rpm",
"dep-a:anthropic/claude-x:tpm",
"dep-b:anthropic/claude-x:rpm",
"dep-b:anthropic/claude-x:tpm",
"deployment:dep-a:cooldown",
"deployment:dep-b:cooldown",
]
], "cooldown state and usage counters must arrive in one MGET"
@pytest.mark.asyncio
async def test_simple_shuffle_still_reads_only_cooldowns():
redis = _redis_answering({})
router = _router(redis, "simple-shuffle")
await router.async_get_available_deployment(model=_MODEL_GROUP, request_kwargs={}, messages=_MESSAGES)
assert _redis_key_families(redis) == [["deployment:dep-a:cooldown", "deployment:dep-b:cooldown"]]
@pytest.mark.asyncio
@pytest.mark.parametrize(
("tpm_a", "tpm_b", "expected"),
[(100, 10, "dep-b"), (10, 100, "dep-a"), (None, 10, "dep-a"), (10, None, "dep-b")],
)
async def test_batched_counters_pick_the_deployment_the_strategy_picks_reading_alone(tpm_a, tpm_b, expected):
counters = {"dep-a:anthropic/claude-x:tpm": tpm_a, "dep-b:anthropic/claude-x:tpm": tpm_b}
routed = _router(_redis_answering(counters), "usage-based-routing-v2")
alone = _router(_redis_answering(counters), "usage-based-routing-v2")
routed_choice = await routed.async_get_available_deployment(
model=_MODEL_GROUP, request_kwargs={}, messages=_MESSAGES
)
alone_choice = await alone.lowesttpm_logger_v2.async_get_available_deployments(
model_group=_MODEL_GROUP, healthy_deployments=alone.model_list, messages=_MESSAGES
)
assert routed_choice["model_info"]["id"] == alone_choice["model_info"]["id"] == expected
@pytest.mark.asyncio
async def test_batched_read_still_excludes_a_cooled_down_deployment():
redis = _redis_answering(
{
"dep-a:anthropic/claude-x:tpm": 100,
"dep-b:anthropic/claude-x:tpm": 10,
"deployment:dep-b:cooldown": _cooldown(seconds=60),
}
)
router = _router(redis, "usage-based-routing-v2")
deployment = await router.async_get_available_deployment(
model=_MODEL_GROUP, request_kwargs={}, messages=_MESSAGES
)
assert deployment["model_info"]["id"] == "dep-a", "dep-b has the lowest tpm but is cooling down"
assert redis.async_batch_get_cache.await_count == 1
@pytest.mark.asyncio
async def test_batched_read_ignores_an_expired_cooldown():
redis = _redis_answering(
{
"dep-a:anthropic/claude-x:tpm": 100,
"dep-b:anthropic/claude-x:tpm": 10,
"deployment:dep-b:cooldown": _cooldown(seconds=-1),
}
)
router = _router(redis, "usage-based-routing-v2")
deployment = await router.async_get_available_deployment(
model=_MODEL_GROUP, request_kwargs={}, messages=_MESSAGES
)
assert deployment["model_info"]["id"] == "dep-b"
@pytest.mark.asyncio
async def test_a_failed_batched_read_degrades_like_the_two_failed_reads_did():
redis = MagicMock(spec=RedisCache)
redis.async_batch_get_cache = AsyncMock(side_effect=ConnectionError("redis unavailable"))
routed = _router(redis, "usage-based-routing-v2")
alone = _router(redis, "usage-based-routing-v2")
with pytest.raises(litellm.RateLimitError, match="No deployments available") as routed_error:
await routed.async_get_available_deployment(model=_MODEL_GROUP, request_kwargs={}, messages=_MESSAGES)
with pytest.raises(litellm.RateLimitError, match="No deployments available") as alone_error:
await alone.lowesttpm_logger_v2.async_get_available_deployments(
model_group=_MODEL_GROUP, healthy_deployments=alone.model_list, messages=_MESSAGES
)
assert str(routed_error.value) == str(alone_error.value)
assert len(routed.cache.last_redis_batch_access_time) == 0, "a failed read must not throttle the next one"
assert len(routed.cooldown_cache.cooldown_store.last_redis_batch_access_time) == 0
@pytest.mark.asyncio
async def test_a_failed_batched_read_leaves_simple_shuffle_routing():
redis = MagicMock(spec=RedisCache)
redis.async_batch_get_cache = AsyncMock(side_effect=ConnectionError("redis unavailable"))
router = _router(redis, "simple-shuffle")
deployment = await router.async_get_available_deployment(
model=_MODEL_GROUP, request_kwargs={}, messages=_MESSAGES
)
assert deployment["model_info"]["id"] in {"dep-a", "dep-b"}