fix(router): give cooldowns their own cache so siblings see a bench in ~1s (#40025)

Cooldown entries rode the router-wide DualCache, which re-reads a key that is
missing from memory at most once every 10s. A deployment benched on one replica
therefore kept taking traffic on its siblings for up to 10 seconds, and the same
shared in-memory tier could evict a live cooldown once 200 unrelated router keys
crowded it out, which sent even the benching replica back to the dead deployment.

CooldownCache now owns a DualCache over the router's Redis with a 1s read
interval and an in-memory tier that only holds cooldown keys. Redis is attached
lazily because the router builds the cooldown cache before it wires Redis up.
This commit is contained in:
Mateo Wang 2026-09-08 10:11:20 -07:00 • committed by GitHub
parent a85c3152ca
commit f769aa4675
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
4 changed files with 212 additions and 39 deletions

View file

@ -73,6 +73,9 @@ DEFAULT_MAX_TOKENS: Final = int(os.getenv("DEFAULT_MAX_TOKENS", 4096))
DEFAULT_ALLOWED_FAILS: Final = int(os.getenv("DEFAULT_ALLOWED_FAILS", 3))
DEFAULT_REDIS_SYNC_INTERVAL: Final = int(os.getenv("DEFAULT_REDIS_SYNC_INTERVAL", 1))
DEFAULT_COOLDOWN_TIME_SECONDS: Final = int(os.getenv("DEFAULT_COOLDOWN_TIME_SECONDS", 5))
DEFAULT_COOLDOWN_REDIS_READ_INTERVAL_SECONDS: Final = float(
os.getenv("DEFAULT_COOLDOWN_REDIS_READ_INTERVAL_SECONDS", "1")
)
DEFAULT_REPLICATE_POLLING_RETRIES: Final = int(os.getenv("DEFAULT_REPLICATE_POLLING_RETRIES", 5))
DEFAULT_REPLICATE_POLLING_DELAY_SECONDS: Final = int(os.getenv("DEFAULT_REPLICATE_POLLING_DELAY_SECONDS", 1))
DEFAULT_IMAGE_TOKEN_COUNT: Final = int(os.getenv("DEFAULT_IMAGE_TOKEN_COUNT", 250))

View file

@ -12,6 +12,7 @@ from typing_extensions import TypedDict
from litellm import verbose_logger
from litellm.caching.caching import DualCache
from litellm.caching.in_memory_cache import InMemoryCache
from litellm.constants import DEFAULT_COOLDOWN_REDIS_READ_INTERVAL_SECONDS
from litellm.litellm_core_utils.sensitive_data_masker import SensitiveDataMasker
if TYPE_CHECKING:
@ -36,10 +37,19 @@ _MAX_CORRECTED_IN_MEMORY_TTL_SECONDS: Final = 60.0
class CooldownCache:
def __init__(self, cache: DualCache, default_cooldown_time: float):
def __init__(
self,
cache: DualCache,
default_cooldown_time: float,
redis_read_interval_seconds: float = DEFAULT_COOLDOWN_REDIS_READ_INTERVAL_SECONDS,
):
self.cache = cache
self.default_cooldown_time = default_cooldown_time
self.in_memory_cache = InMemoryCache()
self._cooldown_store = DualCache(
in_memory_cache=self.in_memory_cache,
default_redis_batch_cache_expiry=redis_read_interval_seconds,
)
# Initialize the masker with custom settings for exception strings
self.exception_masker = SensitiveDataMasker(
visible_prefix=50, # Show first 50 characters
@ -48,6 +58,21 @@ class CooldownCache:
mask_short_values=False, # Truncate long messages only; keep short ones readable
)
@property
def cooldown_store(self) -> DualCache:
"""
The cache cooldown entries live in, with the router's Redis attached on first use.
It is kept separate from the router-wide cache so that a key missing from memory is
re-read from Redis every `redis_read_interval_seconds` rather than on the router
cache's much longer batch interval, which is what lets a sibling replica see a
cooldown another replica wrote, and so that unrelated router keys cannot evict a
cooldown from the in-memory tier before it expires. Redis is attached lazily because
the router builds its cooldown cache before it wires up the shared Redis client.
"""
self._cooldown_store.attach_redis_cache(self.cache.redis_cache)
return self._cooldown_store
def _common_add_cooldown_logic(
self, model_id: str, original_exception, exception_status, cooldown_time: float
) -> tuple[str, CooldownCacheValue]:
@ -93,7 +118,7 @@ class CooldownCache:
)
# Set the cache with a TTL equal to the cooldown time
self.cache.set_cache(
self.cooldown_store.set_cache(
value=cooldown_data,
key=cooldown_key,
ttl=_cooldown_time,
@ -122,13 +147,13 @@ class CooldownCache:
cooldown_cache_value: Final = CooldownCacheValue(**result) # pyright: ignore[reportUnknownArgumentType] - result comes from an untyped cache read, not from our own code
remaining: Final = (cooldown_cache_value["timestamp"] + cooldown_cache_value["cooldown_time"]) - current_time
if remaining <= 0:
self.cache.in_memory_cache.delete_cache(key)
self.in_memory_cache.delete_cache(key)
return None
current_expiry: Final = self.cache.in_memory_cache.ttl_dict.get(key)
current_expiry: Final = self.in_memory_cache.ttl_dict.get(key)
if current_expiry is not None and current_expiry > current_time + remaining + 5:
corrected_ttl: Final = min(remaining, _MAX_CORRECTED_IN_MEMORY_TTL_SECONDS)
self.cache.in_memory_cache.delete_cache(key)
self.cache.in_memory_cache.set_cache(key, result, ttl=corrected_ttl)
self.in_memory_cache.delete_cache(key)
self.in_memory_cache.set_cache(key, result, ttl=corrected_ttl)
return cooldown_cache_value
async def async_get_active_cooldowns(
@ -137,12 +162,7 @@ class CooldownCache:
# Generate the keys for the deployments
keys: Final = [CooldownCache.get_cooldown_cache_key(model_id) for model_id in model_ids]
# Retrieve the values for the keys using mget
## more likely to be none if no models ratelimited. So just check redis every 1s
## each redis call adds ~100ms latency.
## check in memory cache first
results: Final = await self.cache.async_batch_get_cache(keys=keys, parent_otel_span=parent_otel_span)
results: Final = await self.cooldown_store.async_batch_get_cache(keys=keys, parent_otel_span=parent_otel_span)
active_cooldowns: Final[list[tuple[str, CooldownCacheValue]]] = []
if results is None or all(v is None for v in results):
@ -164,7 +184,7 @@ class CooldownCache:
# Generate the keys for the deployments
keys: Final = [CooldownCache.get_cooldown_cache_key(model_id) for model_id in model_ids]
# Retrieve the values for the keys using mget
results: Final = self.cache.batch_get_cache(keys=keys, parent_otel_span=parent_otel_span) or []
results: Final = self.cooldown_store.batch_get_cache(keys=keys, parent_otel_span=parent_otel_span) or []
active_cooldowns: Final = []
current_time: Final = time.time()
@ -184,7 +204,7 @@ class CooldownCache:
keys: Final = [f"deployment:{model_id}:cooldown" for model_id in model_ids]
# Retrieve the values for the keys using mget
results: Final = self.cache.batch_get_cache(keys=keys, parent_otel_span=parent_otel_span) or []
results: Final = self.cooldown_store.batch_get_cache(keys=keys, parent_otel_span=parent_otel_span) or []
min_cooldown_time: float | None = None
# Process the results

View file

@ -246,12 +246,12 @@ class TestCooldownCacheTTLCorrection:
"timestamp": time.time() - 120.0,
"cooldown_time": 60.0,
}
cc.cache.in_memory_cache.set_cache(key, expired_value, ttl=600)
cc.in_memory_cache.set_cache(key, expired_value, ttl=600)
active = cc.get_active_cooldowns(model_ids=[model_id], parent_otel_span=None)
assert active == [], "Expired cooldown entry must not appear in active cooldowns"
assert cc.cache.in_memory_cache.get_cache(key) is None, "Expired entry must be evicted from in-memory cache"
assert cc.in_memory_cache.get_cache(key) is None, "Expired entry must be evicted from in-memory cache"
def test_active_entry_is_returned(self):
"""
@ -267,7 +267,7 @@ class TestCooldownCacheTTLCorrection:
"timestamp": time.time(),
"cooldown_time": 60.0,
}
cc.cache.in_memory_cache.set_cache(key, active_value, ttl=60)
cc.in_memory_cache.set_cache(key, active_value, ttl=60)
active = cc.get_active_cooldowns(model_ids=[model_id], parent_otel_span=None)
@ -290,14 +290,14 @@ class TestCooldownCacheTTLCorrection:
"timestamp": time.time() - (60.0 - remaining),
"cooldown_time": 60.0,
}
cc.cache.in_memory_cache.set_cache(key, value, ttl=600)
cc.in_memory_cache.set_cache(key, value, ttl=600)
before_expiry = cc.cache.in_memory_cache.ttl_dict.get(key)
before_expiry = cc.in_memory_cache.ttl_dict.get(key)
assert before_expiry is not None
cc.get_active_cooldowns(model_ids=[model_id], parent_otel_span=None)
after_expiry = cc.cache.in_memory_cache.ttl_dict.get(key)
after_expiry = cc.in_memory_cache.ttl_dict.get(key)
assert after_expiry is not None
corrected_remaining = after_expiry - time.time()
assert corrected_remaining <= 60.0, "Corrected TTL must not exceed 60s"
@ -318,12 +318,12 @@ class TestCooldownCacheTTLCorrection:
"timestamp": time.time() - 120.0,
"cooldown_time": 60.0,
}
cc.cache.in_memory_cache.set_cache(key, expired_value, ttl=600)
cc.in_memory_cache.set_cache(key, expired_value, ttl=600)
active = await cc.async_get_active_cooldowns(model_ids=[model_id], parent_otel_span=None)
assert active == [], "Expired entry must not appear in async active cooldowns"
assert cc.cache.in_memory_cache.get_cache(key) is None
assert cc.in_memory_cache.get_cache(key) is None
class TestFallbackDeploymentCooldown:

View file

@ -268,12 +268,12 @@ class TestCooldownCacheTTLCorrection:
"timestamp": time.time() - 120.0,
"cooldown_time": 60.0,
}
cc.cache.in_memory_cache.set_cache(key, expired_value, ttl=600)
cc.in_memory_cache.set_cache(key, expired_value, ttl=600)
active = cc.get_active_cooldowns(model_ids=[model_id], parent_otel_span=None)
assert active == [], "Expired cooldown entry must not appear in active cooldowns"
assert cc.cache.in_memory_cache.get_cache(key) is None, "Expired entry must be evicted from in-memory cache"
assert cc.in_memory_cache.get_cache(key) is None, "Expired entry must be evicted from in-memory cache"
def test_active_entry_is_returned(self):
"""
@ -289,7 +289,7 @@ class TestCooldownCacheTTLCorrection:
"timestamp": time.time(),
"cooldown_time": 60.0,
}
cc.cache.in_memory_cache.set_cache(key, active_value, ttl=60)
cc.in_memory_cache.set_cache(key, active_value, ttl=60)
active = cc.get_active_cooldowns(model_ids=[model_id], parent_otel_span=None)
@ -312,14 +312,14 @@ class TestCooldownCacheTTLCorrection:
"timestamp": time.time() - (60.0 - remaining),
"cooldown_time": 60.0,
}
cc.cache.in_memory_cache.set_cache(key, value, ttl=600)
cc.in_memory_cache.set_cache(key, value, ttl=600)
before_expiry = cc.cache.in_memory_cache.ttl_dict.get(key)
before_expiry = cc.in_memory_cache.ttl_dict.get(key)
assert before_expiry is not None
cc.get_active_cooldowns(model_ids=[model_id], parent_otel_span=None)
after_expiry = cc.cache.in_memory_cache.ttl_dict.get(key)
after_expiry = cc.in_memory_cache.ttl_dict.get(key)
assert after_expiry is not None
corrected_remaining = after_expiry - time.time()
assert corrected_remaining <= 60.0, "Corrected TTL must not exceed 60s"
@ -340,12 +340,12 @@ class TestCooldownCacheTTLCorrection:
"timestamp": time.time() - 120.0,
"cooldown_time": 60.0,
}
cc.cache.in_memory_cache.set_cache(key, expired_value, ttl=600)
cc.in_memory_cache.set_cache(key, expired_value, ttl=600)
active = await cc.async_get_active_cooldowns(model_ids=[model_id], parent_otel_span=None)
assert active == [], "Expired entry must not appear in async active cooldowns"
assert cc.cache.in_memory_cache.get_cache(key) is None
assert cc.in_memory_cache.get_cache(key) is None
@pytest.mark.asyncio
async def test_async_active_entry_is_returned(self):
@ -363,7 +363,7 @@ class TestCooldownCacheTTLCorrection:
"timestamp": time.time(),
"cooldown_time": 60.0,
}
cc.cache.in_memory_cache.set_cache(key, active_value, ttl=60)
cc.in_memory_cache.set_cache(key, active_value, ttl=60)
active = await cc.async_get_active_cooldowns(model_ids=[model_id], parent_otel_span=None)
@ -389,18 +389,18 @@ class TestCorrectedActiveCooldown:
cc = self._make_cooldown_cache()
key = "deployment:expired-dep:cooldown"
entry = self._entry(timestamp=time.time() - 120.0, cooldown_time=60.0)
cc.cache.in_memory_cache.set_cache(key, dict(entry), ttl=600)
cc.in_memory_cache.set_cache(key, dict(entry), ttl=600)
result = cc._corrected_active_cooldown(key, dict(entry), current_time=time.time())
assert result is None
assert cc.cache.in_memory_cache.get_cache(key) is None
assert cc.in_memory_cache.get_cache(key) is None
def test_active_entry_within_window_returns_value(self):
cc = self._make_cooldown_cache()
key = "deployment:active-dep:cooldown"
entry = self._entry(timestamp=time.time(), cooldown_time=60.0)
cc.cache.in_memory_cache.set_cache(key, dict(entry), ttl=60)
cc.in_memory_cache.set_cache(key, dict(entry), ttl=60)
result = cc._corrected_active_cooldown(key, dict(entry), current_time=time.time())
@ -412,12 +412,12 @@ class TestCorrectedActiveCooldown:
key = "deployment:backfilled-dep:cooldown"
remaining = 30.0
entry = self._entry(timestamp=time.time() - (60.0 - remaining), cooldown_time=60.0)
cc.cache.in_memory_cache.set_cache(key, dict(entry), ttl=600)
cc.in_memory_cache.set_cache(key, dict(entry), ttl=600)
result = cc._corrected_active_cooldown(key, dict(entry), current_time=time.time())
assert result is not None
corrected_expiry = cc.cache.in_memory_cache.ttl_dict.get(key)
corrected_expiry = cc.in_memory_cache.ttl_dict.get(key)
assert corrected_expiry is not None
assert corrected_expiry - time.time() <= 60.0
@ -425,10 +425,160 @@ class TestCorrectedActiveCooldown:
cc = self._make_cooldown_cache()
key = "deployment:normal-dep:cooldown"
entry = self._entry(timestamp=time.time(), cooldown_time=60.0)
cc.cache.in_memory_cache.set_cache(key, dict(entry), ttl=60)
original_expiry = cc.cache.in_memory_cache.ttl_dict.get(key)
cc.in_memory_cache.set_cache(key, dict(entry), ttl=60)
original_expiry = cc.in_memory_cache.ttl_dict.get(key)
cc._corrected_active_cooldown(key, dict(entry), current_time=time.time())
after_expiry = cc.cache.in_memory_cache.ttl_dict.get(key)
after_expiry = cc.in_memory_cache.ttl_dict.get(key)
assert after_expiry == original_expiry
class SharedRedisDouble:
"""
In-process stand-in for RedisCache, shared by several DualCache instances so that
tests can model two proxy replicas talking to one Redis.
"""
def __init__(self) -> None:
self.store: dict = {} # mutable-ok: stands in for Redis' own mutable keyspace
def set_cache(self, key, value, **kwargs):
self.store[key] = value
async def async_set_cache(self, key, value, **kwargs):
self.store[key] = value
def batch_get_cache(self, key_list, parent_otel_span=None, **kwargs):
return {key: self.store.get(key) for key in key_list}
async def async_batch_get_cache(self, key_list, parent_otel_span=None, **kwargs):
return {key: self.store.get(key) for key in key_list}
class TestCooldownPropagationBetweenReplicas:
"""
A cooldown written by one replica has to reach its siblings quickly. The router's own
DualCache re-reads a key that is missing from memory only every 10s, so cooldown reads
get their own cache with a much shorter Redis read interval.
"""
def _make_replica(self, redis: SharedRedisDouble, read_interval: float | None = None) -> CooldownCache:
router_cache = DualCache(in_memory_cache=InMemoryCache(), redis_cache=redis)
if read_interval is None:
return CooldownCache(cache=router_cache, default_cooldown_time=60.0)
return CooldownCache(
cache=router_cache,
default_cooldown_time=60.0,
redis_read_interval_seconds=read_interval,
)
@pytest.mark.asyncio
async def test_sibling_replica_sees_cooldown_within_configured_read_interval(self):
redis = SharedRedisDouble()
replica_a = self._make_replica(redis, read_interval=0.25)
replica_b = self._make_replica(redis, read_interval=0.25)
model_id = "shared-deployment"
assert await replica_b.async_get_active_cooldowns([model_id], parent_otel_span=None) == []
replica_a.add_deployment_to_cooldown(
model_id=model_id,
original_exception=Exception("Internal server error"),
exception_status=500,
cooldown_time=60.0,
)
time.sleep(0.3)
active = await replica_b.async_get_active_cooldowns([model_id], parent_otel_span=None)
assert [model_id] == [entry[0] for entry in active], (
"sibling replica must pick up a cooldown written by another replica within the read interval"
)
@pytest.mark.asyncio
async def test_sibling_replica_sees_cooldown_within_default_read_interval(self):
redis = SharedRedisDouble()
replica_a = self._make_replica(redis)
replica_b = self._make_replica(redis)
model_id = "default-interval-deployment"
assert await replica_b.async_get_active_cooldowns([model_id], parent_otel_span=None) == []
replica_a.add_deployment_to_cooldown(
model_id=model_id,
original_exception=Exception("Internal server error"),
exception_status=500,
cooldown_time=60.0,
)
time.sleep(1.2)
active = await replica_b.async_get_active_cooldowns([model_id], parent_otel_span=None)
assert [model_id] == [entry[0] for entry in active], (
"the shipped default read interval must let a sibling replica see a cooldown about a second later"
)
def test_sync_read_path_sees_sibling_cooldown_within_read_interval(self):
redis = SharedRedisDouble()
replica_a = self._make_replica(redis, read_interval=0.25)
replica_b = self._make_replica(redis, read_interval=0.25)
model_id = "sync-shared-deployment"
assert replica_b.get_active_cooldowns([model_id], parent_otel_span=None) == []
replica_a.add_deployment_to_cooldown(
model_id=model_id,
original_exception=Exception("Internal server error"),
exception_status=500,
cooldown_time=60.0,
)
time.sleep(0.3)
active = replica_b.get_active_cooldowns([model_id], parent_otel_span=None)
assert [model_id] == [entry[0] for entry in active]
@pytest.mark.asyncio
async def test_redis_attached_after_construction_is_still_used(self):
redis = SharedRedisDouble()
router_cache = DualCache(in_memory_cache=InMemoryCache())
writer = CooldownCache(cache=router_cache, default_cooldown_time=60.0, redis_read_interval_seconds=0.25)
router_cache.attach_redis_cache(redis)
reader = self._make_replica(redis, read_interval=0.25)
model_id = "late-redis-deployment"
writer.add_deployment_to_cooldown(
model_id=model_id,
original_exception=Exception("Internal server error"),
exception_status=500,
cooldown_time=60.0,
)
active = await reader.async_get_active_cooldowns([model_id], parent_otel_span=None)
assert [model_id] == [entry[0] for entry in active], (
"a router that wires Redis after building its cooldown cache must still publish cooldowns to it"
)
class TestCooldownSurvivesUnrelatedCacheTraffic:
@pytest.mark.asyncio
async def test_unrelated_router_cache_writes_do_not_evict_active_cooldown(self):
router_cache = DualCache(in_memory_cache=InMemoryCache())
cc = CooldownCache(cache=router_cache, default_cooldown_time=60.0)
model_id = "busy-router-deployment"
cc.add_deployment_to_cooldown(
model_id=model_id,
original_exception=Exception("Internal server error"),
exception_status=500,
cooldown_time=30.0,
)
for i in range(400):
router_cache.set_cache(key=f"unrelated-router-key-{i}", value={"n": i})
active = await cc.async_get_active_cooldowns([model_id], parent_otel_span=None)
assert [model_id] == [entry[0] for entry in active], (
"unrelated router cache traffic must not evict a cooldown that is still running"
)