diff --git a/litellm/proxy/hooks/parallel_request_limiter_v3.py b/litellm/proxy/hooks/parallel_request_limiter_v3.py index 87200905c80..d32e997cd10 100644 --- a/litellm/proxy/hooks/parallel_request_limiter_v3.py +++ b/litellm/proxy/hooks/parallel_request_limiter_v3.py @@ -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 diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index bba6a412e44..47c6ff256f9 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -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: diff --git a/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py b/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py index af91874f76e..69357630ec2 100644 --- a/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py +++ b/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py @@ -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, diff --git a/tests/test_litellm/proxy/proxy_server/test_lifecycle.py b/tests/test_litellm/proxy/proxy_server/test_lifecycle.py index deb7289d2d1..26d4c2c408a 100644 --- a/tests/test_litellm/proxy/proxy_server/test_lifecycle.py +++ b/tests/test_litellm/proxy/proxy_server/test_lifecycle.py @@ -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): """