fix(proxy): fail open on an open replica circuit breaker, close replica clients at shutdown

This commit is contained in:
michelligabriele 2026-09-15 15:32:35 +02:00
parent 43146036e1
commit 874544496d
No known key found for this signature in database
4 changed files with 101 additions and 4 deletions

View file

@ -1175,8 +1175,21 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
return RateLimitResponse(overall_code=overall_code, statuses=statuses)
async def _read_replica_counters(self, replica: RedisCache, keys: Sequence[str]) -> Mapping[str, object]:
"""`RedisCache` logs and swallows its own failures, so an unreachable replica comes back empty."""
return await replica.async_batch_get_cache(key_list=list(keys))
"""
`async_batch_get_cache` swallows read failures itself, but its circuit-breaker
guard raises `RedisCircuitBreakerOpenError` before the body runs once the
replica's breaker opens. A replica must never fail the request either way.
"""
try:
return await replica.async_batch_get_cache(key_list=list(keys))
except Exception as e: # noqa: BLE001 # any replica failure degrades to local-only enforcement, never a 500
log_redis_failure(
verbose_proxy_logger,
logging.WARNING,
"rate_limit_remote_replicas: replica read failed, enforcing against local counters only",
e,
)
return {}
async def _remote_counter_offsets(
self,
@ -1186,8 +1199,8 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
"""
Sum each counter's value across the remote replicas, counting a replica
only while its copy of that counter's window is still current. A replica
that reads back empty contributes nothing, so a region whose replica link
is down falls back to the per-region enforcement it has today.
that reads back empty or fails contributes nothing, so a region whose
replica link is down falls back to the per-region enforcement it has today.
"""
if not self.remote_replica_caches or not probes:
return _NO_REMOTE_OFFSETS

View file

@ -1002,6 +1002,9 @@ async def proxy_shutdown_event(worker_heartbeat: ProxyWorkerHeartbeat | None = N
if litellm.cache is not None:
await litellm.cache.disconnect()
for replica in rate_limit_remote_replica_caches:
await replica.disconnect()
await jwt_handler.close()
if db_writer_client is not None:

View file

@ -6547,6 +6547,20 @@ class _ReplicaRedis:
return {key: self.snapshot[key] for key in key_list if key in self.snapshot}
class _OpenBreakerReplicaRedis:
"""Replica whose circuit breaker has opened: the guard raises before the read body runs."""
async def async_batch_get_cache(self, key_list, parent_otel_span=None):
from litellm.caching.redis_cache import RedisCircuitBreakerOpenError
raise RedisCircuitBreakerOpenError("Redis circuit breaker is open")
class _ExplodingReplicaRedis:
async def async_batch_get_cache(self, key_list, parent_otel_span=None):
raise ConnectionError("replica unreachable")
class _LuaFailurePrimaryRedis:
"""
Primary Redis whose Lua calls fail. That is how production reaches the in-memory
@ -6755,6 +6769,36 @@ async def test_a_replica_that_reads_back_empty_fails_open(time_controller):
)
@pytest.mark.asyncio
async def test_a_replica_whose_breaker_is_open_fails_open(time_controller):
"""`async_batch_get_cache` swallows its own read failures, but the circuit-breaker guard
wrapping it raises before the body once the replica's breaker opens, which is the state a
replica outage settles into after a few failed reads. That must not fail the request."""
handler, descriptors = _rpm_handler([_OpenBreakerReplicaRedis()], time_controller, 100)
response = await handler.should_rate_limit(descriptors=descriptors)
assert response["overall_code"] == "OK"
assert response["statuses"][0]["limit_remaining"] == 99
@pytest.mark.asyncio
async def test_a_replica_read_that_raises_fails_open_and_warns(time_controller, caplog):
handler, descriptors = _rpm_handler([_ExplodingReplicaRedis()], time_controller, 100)
with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"):
response = await handler.should_rate_limit(descriptors=descriptors)
assert response["overall_code"] == "OK"
assert response["statuses"][0]["limit_remaining"] == 99
replica_warnings = [
record.getMessage()
for record in caplog.records
if record.levelno >= logging.WARNING and "replica read failed" in record.getMessage()
]
assert len(replica_warnings) == 1
@pytest.mark.asyncio
async def test_a_replica_read_missing_the_window_key_contributes_nothing(time_controller):
"""A read that returns the counter but not its window cannot say which window that
@ -6917,6 +6961,20 @@ async def test_reserve_tpm_without_replicas_allows_a_request_that_fits_locally(
assert response["statuses"][0]["limit_remaining"] == 800
@pytest.mark.asyncio
async def test_an_open_replica_breaker_on_the_reservation_path_enforces_local_limits_only(
time_controller,
):
handler, descriptors = _tpm_handler([_OpenBreakerReplicaRedis()], time_controller, 1000)
response = await handler.reserve_tpm_tokens(
descriptors=descriptors, estimated_tokens=200
)
assert response["overall_code"] == "OK"
assert response["statuses"][0]["limit_remaining"] == 800
@pytest.mark.asyncio
async def test_an_empty_replica_read_on_the_reservation_path_enforces_local_limits_only(
time_controller,

View file

@ -130,6 +130,29 @@ async def test_proxy_shutdown_event_disconnects_prisma_and_resets(monkeypatch):
}
@pytest.mark.asyncio
async def test_proxy_shutdown_event_closes_rate_limit_remote_replicas(monkeypatch):
fake_replicas = (MagicMock(), MagicMock())
for replica in fake_replicas:
replica.disconnect = AsyncMock()
monkeypatch.setattr(ps, "rate_limit_remote_replica_caches", fake_replicas, raising=False)
monkeypatch.setattr(ps, "prisma_client", None, raising=False)
fake_jwt = MagicMock()
fake_jwt.close = AsyncMock()
monkeypatch.setattr(ps, "jwt_handler", fake_jwt, raising=False)
monkeypatch.setattr(ps, "db_writer_client", None, raising=False)
import litellm
monkeypatch.setattr(litellm, "cache", None, raising=False)
monkeypatch.setattr(litellm, "success_callback", [], raising=False)
await proxy_shutdown_event()
assert [replica.disconnect.await_count for replica in fake_replicas] == [1, 1]
@pytest.mark.asyncio
async def test_proxy_shutdown_drains_gateway_requests_before_disconnecting(monkeypatch):
"""