fix(caching): make an open Redis circuit breaker a quiet cache miss

An open breaker raised a generic Exception on every skipped call, and DualCache caught
it and logged a full ERROR traceback each time. Under load that became hundreds of
traceback formats per second on every replica and pinned the proxies at 100% CPU.

Raise a typed RedisCircuitBreakerOpenError instead and have DualCache return its
in-memory result without logging for it. Classify redis-py pool exhaustion
(ConnectionError chained from TimeoutError) as a timeout so a latency blip goes
through the duration gate. Track a breaker generation so a call admitted before the
breaker opened cannot close it, leaving that to the HALF_OPEN probe. Rate limit the
LoggingWorker callback-error traceback to one per interval so a stalled logging
backend cannot start a second traceback storm.

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
mateo 2026-09-10 20:58:41 +00:00
parent a9cec50960
commit 32b7daf691
7 changed files with 326 additions and 24 deletions

View file

@ -23,7 +23,7 @@ from litellm.constants import DEFAULT_MAX_REDIS_BATCH_CACHE_SIZE
from .base_cache import BaseCache
from .in_memory_cache import InMemoryCache
from .redis_cache import RedisCache
from .redis_cache import RedisCache, RedisCircuitBreakerOpenError
if TYPE_CHECKING:
from opentelemetry.trace import Span as _Span
@ -250,6 +250,8 @@ class DualCache(BaseCache):
print_verbose(f"get cache: cache result: {result}")
return result
except RedisCircuitBreakerOpenError:
return None
except Exception:
verbose_logger.error(traceback.format_exc())
@ -319,6 +321,9 @@ class DualCache(BaseCache):
redis_result: Final = await self.redis_cache.async_batch_get_cache(
sublist_keys, parent_otel_span=parent_otel_span
)
except RedisCircuitBreakerOpenError:
self._rollback_redis_batch_key_reservations(previous_access_times)
return result
except Exception:
# Do not throttle subsequent callers if the Redis read fails.
self._rollback_redis_batch_key_reservations(previous_access_times)
@ -352,6 +357,8 @@ class DualCache(BaseCache):
if self.redis_cache is not None and local_only is False:
await self.redis_cache.async_set_cache(key, value, **kwargs)
except RedisCircuitBreakerOpenError:
return
except Exception as e:
verbose_logger.exception("LiteLLM Cache: Excepton async add_cache: %s", e)
@ -371,6 +378,8 @@ class DualCache(BaseCache):
await self.redis_cache.async_set_cache_pipeline(
cache_list=cache_list, ttl=kwargs.pop("ttl", None), **kwargs
)
except RedisCircuitBreakerOpenError:
return
except Exception as e:
verbose_logger.exception("LiteLLM Cache: Excepton async add_cache: %s", e)
@ -408,6 +417,8 @@ class DualCache(BaseCache):
refresh_ttl=refresh_ttl,
)
return result
except RedisCircuitBreakerOpenError:
return result
except Exception as e:
verbose_logger.warning(
@ -437,6 +448,8 @@ class DualCache(BaseCache):
parent_otel_span=parent_otel_span,
)
return result
except RedisCircuitBreakerOpenError:
return result
except Exception as e:
verbose_logger.warning(

View file

@ -17,6 +17,7 @@ import json
import time
from collections.abc import Awaitable, Callable, Sequence
from contextvars import ContextVar
from dataclasses import dataclass
from datetime import timedelta
from typing import TYPE_CHECKING, Any, Final, Protocol, TypeVar, cast
@ -148,6 +149,10 @@ def _get_call_stack_info(num_frames: int = 2) -> str:
return "unknown"
class RedisCircuitBreakerOpenError(Exception):
"""Expected fast-fail while the breaker is open; optional-cache callers treat it as a miss."""
class RedisCircuitBreaker:
"""
Tracks Redis health for a RedisCache instance.
@ -163,8 +168,12 @@ class RedisCircuitBreaker:
(no success or hard failure in between) that reaches
failure_threshold and spans timeout_min_duration seconds
OPEN -> HALF_OPEN after recovery_timeout seconds
HALF_OPEN -> CLOSED on success
HALF_OPEN -> OPEN on failure (resets timer)
HALF_OPEN -> CLOSED on the recovery probe's success
HALF_OPEN -> OPEN on the recovery probe's failure (resets timer)
Every OPEN transition starts a new generation. A call reports its outcome only for the
generation it was admitted under, so a success from a call that was already in flight
when the breaker opened cannot close it and the HALF_OPEN probe is the only call that can.
Timeouts are accounted separately from hard connectivity failures because the async
Redis timeout includes time waiting for the worker event loop to resume: one loop
@ -194,9 +203,14 @@ class RedisCircuitBreaker:
self._timeout_count = 0
self._timeout_streak_started_at: float | None = None
self._opened_at: float | None = None
self._generation = 0
self._state = self.CLOSED
_breaker_metrics().record_state_change(None, self._state)
@property
def generation(self) -> int:
return self._generation
def is_open(self) -> bool:
"""Returns True if Redis calls should be skipped."""
if not self.enabled:
@ -241,18 +255,19 @@ class RedisCircuitBreaker:
if self._state != self.OPEN:
verbose_logger.warning(
"Redis circuit breaker OPENED after %d consecutive failures"
" (%d hard connectivity) — fast-failing Redis calls for %ds",
" (%d hard connectivity), fast-failing Redis calls for %ds",
self._failure_count,
self._hard_failure_count,
self.recovery_timeout,
)
self._generation += 1
self._set_state(self.OPEN)
def record_success(self) -> None:
if not self.enabled:
return
if self._state == self.HALF_OPEN:
verbose_logger.info("Redis circuit breaker CLOSED — Redis recovered")
verbose_logger.info("Redis circuit breaker CLOSED, Redis recovered")
self._failure_count = 0
self._hard_failure_count = 0
self._timeout_count = 0
@ -320,7 +335,10 @@ def _redis_timeout_error_types() -> tuple[type, ...]:
def _is_redis_timeout_failure(exc: BaseException) -> bool:
return isinstance(exc, _redis_timeout_error_types())
"""Follows __cause__: a blocking pool wait raises ConnectionError from asyncio.TimeoutError."""
if isinstance(exc, _redis_timeout_error_types()):
return True
return exc.__cause__ is not None and _is_redis_timeout_failure(exc.__cause__)
class _BreakerMetrics:
@ -391,23 +409,37 @@ def _record_swallowed_redis_failure(breaker: RedisCircuitBreaker, exc: BaseExcep
_swallowed_redis_failures.set(_swallowed_redis_failures.get() + 1)
def _enter_circuit_breaker(breaker: RedisCircuitBreaker, name: str) -> int:
"""Reject the call if the breaker is open, else return the swallowed-failure count to compare against."""
@dataclass(frozen=True, slots=True)
class _BreakerAdmission:
swallowed_before: int
generation: int
def _enter_circuit_breaker(breaker: RedisCircuitBreaker, name: str) -> _BreakerAdmission:
"""Reject the call if the breaker is open, else snapshot what its outcome will be judged against."""
if breaker.is_open():
raise Exception(f"Redis circuit breaker is open — skipping {name}")
return _swallowed_redis_failures.get()
raise RedisCircuitBreakerOpenError(f"Redis circuit breaker is open, skipping {name}")
return _BreakerAdmission(swallowed_before=_swallowed_redis_failures.get(), generation=breaker.generation)
def _exit_circuit_breaker(breaker: RedisCircuitBreaker, swallowed_before: int) -> None:
"""Record success only when nothing failed while the call ran.
def _exit_circuit_breaker(breaker: RedisCircuitBreaker, admission: _BreakerAdmission) -> None:
"""Record success only when nothing failed while the call ran and the breaker has not opened since.
Several Redis methods catch their own connection errors and return a default, so a
method that returned is not on its own proof of a healthy Redis.
"""
if _swallowed_redis_failures.get() == swallowed_before:
if admission.generation != breaker.generation:
return
if _swallowed_redis_failures.get() == admission.swallowed_before:
breaker.record_success()
def _fail_circuit_breaker(breaker: RedisCircuitBreaker, admission: _BreakerAdmission, exc: BaseException) -> None:
if admission.generation != breaker.generation or not _is_redis_health_failure(exc):
return
breaker.record_failure(is_timeout=_is_redis_timeout_failure(exc))
async def _run_under_circuit_breaker(
breaker: RedisCircuitBreaker,
name: str,
@ -418,14 +450,13 @@ async def _run_under_circuit_breaker(
Shared by the method decorator and the Lua script executor so both feed the same
health signal.
"""
swallowed_before: Final = _enter_circuit_breaker(breaker, name)
admission: Final = _enter_circuit_breaker(breaker, name)
try:
result: Final = await call()
except Exception as e:
if _is_redis_health_failure(e):
breaker.record_failure(is_timeout=_is_redis_timeout_failure(e))
_fail_circuit_breaker(breaker, admission, e)
raise
_exit_circuit_breaker(breaker, swallowed_before)
_exit_circuit_breaker(breaker, admission)
return result
@ -435,14 +466,13 @@ def _run_under_circuit_breaker_sync(
call: Callable[[], _RedisCallResult],
) -> _RedisCallResult:
"""Run one blocking Redis call under a circuit breaker, feeding the same health signal as the async path."""
swallowed_before: Final = _enter_circuit_breaker(breaker, name)
admission: Final = _enter_circuit_breaker(breaker, name)
try:
result: Final = call()
except Exception as e:
if _is_redis_health_failure(e):
breaker.record_failure()
_fail_circuit_breaker(breaker, admission, e)
raise
_exit_circuit_breaker(breaker, swallowed_before)
_exit_circuit_breaker(breaker, admission)
return result
@ -1382,10 +1412,10 @@ class RedisCache(BaseCache):
start_time: Final = time.time()
try:
swallowed_before: Final = _enter_circuit_breaker(self._circuit_breaker, "batch_get_cache")
admission: Final = _enter_circuit_breaker(self._circuit_breaker, "batch_get_cache")
_keys: Final = [self.check_and_fix_namespace(key=cache_key or "") for cache_key in _key_list]
results: Final = self._run_redis_mget_operation(keys=_keys)
_exit_circuit_breaker(self._circuit_breaker, swallowed_before)
_exit_circuit_breaker(self._circuit_breaker, admission)
end_time: Final = time.time()
_duration: Final = end_time - start_time
self.service_logger_obj.service_success_hook(
@ -1409,6 +1439,8 @@ class RedisCache(BaseCache):
decoded_results[k] = v
return decoded_results
except RedisCircuitBreakerOpenError:
return key_value_dict
except Exception as e:
failed_at: Final = time.time()
self.service_logger_obj.service_failure_hook(

View file

@ -571,6 +571,7 @@ ANTHROPIC_MESSAGES_MAX_DETACHED_STREAM_DRAINS: Final = int(
LOGGING_WORKER_CONCURRENCY: Final = int(os.getenv("LOGGING_WORKER_CONCURRENCY", 100)) # Must be above 0
LOGGING_WORKER_MAX_QUEUE_SIZE: Final = int(os.getenv("LOGGING_WORKER_MAX_QUEUE_SIZE", 50_000))
LOGGING_WORKER_MAX_TIME_PER_COROUTINE: Final = float(os.getenv("LOGGING_WORKER_MAX_TIME_PER_COROUTINE", 20.0))
LOGGING_WORKER_ERROR_TRACEBACK_INTERVAL_SECONDS: Final = 60.0
LOGGING_WORKER_CLEAR_PERCENTAGE: Final = int(
os.getenv("LOGGING_WORKER_CLEAR_PERCENTAGE", 50)
) # Percentage of queue to clear (default: 50%)

View file

@ -6,6 +6,7 @@ import atexit
import contextvars
import inspect
import logging
import time
from collections.abc import Coroutine, Iterator
from typing import Final
@ -16,6 +17,7 @@ from litellm.constants import (
LOGGING_WORKER_AGGRESSIVE_CLEAR_COOLDOWN_SECONDS,
LOGGING_WORKER_CLEAR_PERCENTAGE,
LOGGING_WORKER_CONCURRENCY,
LOGGING_WORKER_ERROR_TRACEBACK_INTERVAL_SECONDS,
LOGGING_WORKER_MAX_QUEUE_SIZE,
LOGGING_WORKER_MAX_TIME_PER_COROUTINE,
MAX_ITERATIONS_TO_CLEAR_QUEUE,
@ -47,10 +49,14 @@ class LoggingWorker:
timeout: float = LOGGING_WORKER_MAX_TIME_PER_COROUTINE,
max_queue_size: int = LOGGING_WORKER_MAX_QUEUE_SIZE,
concurrency: int = LOGGING_WORKER_CONCURRENCY,
error_traceback_interval: float = LOGGING_WORKER_ERROR_TRACEBACK_INTERVAL_SECONDS,
):
self.timeout = timeout
self.max_queue_size = max_queue_size
self.concurrency = concurrency
self.error_traceback_interval = error_traceback_interval
self._last_error_traceback_at: float | None = None
self._errors_since_traceback: int = 0
self._queue: asyncio.Queue[LoggingTask] | None = None
self._worker_task: asyncio.Task | None = None
self._running_tasks: set[asyncio.Task] = set()
@ -163,7 +169,7 @@ class LoggingWorker:
timeout=self.timeout,
)
except Exception as e:
verbose_logger.exception("LoggingWorker error: %s", e)
self._log_task_error(e)
finally:
self._untrack_dequeued(task)
self._queue.task_done()
@ -171,6 +177,21 @@ class LoggingWorker:
# Always release semaphore, even if queue is None
sem.release()
def _log_task_error(self, error: Exception) -> None:
"""One traceback per interval: a stalled backend fails every in-flight task at once."""
now: Final = time.monotonic()
last_traceback_at: Final = self._last_error_traceback_at
if last_traceback_at is not None and now - last_traceback_at < self.error_traceback_interval:
self._errors_since_traceback += 1
return
verbose_logger.exception(
"LoggingWorker error (%d more suppressed since the last traceback): %r",
self._errors_since_traceback,
error,
)
self._last_error_traceback_at = now
self._errors_since_traceback = 0
async def _worker_loop(self) -> None:
"""Main worker loop that gets tasks and schedules them to run concurrently."""
try:

View file

@ -1,4 +1,5 @@
import asyncio
import logging
import time
import uuid
from unittest.mock import AsyncMock, MagicMock, patch
@ -576,3 +577,64 @@ async def test_dual_cache_late_attach_redis_wires_writes_and_ttl_async():
assert mock_redis.async_set_cache.call_args[0][:2] == (key_after, val_after)
assert in_memory.get_cache(key_after) == val_after
@pytest.fixture
def dual_cache_with_open_breaker():
"""A DualCache whose Redis tier is behind an already-open circuit breaker.
The Redis client is a mock that fails the test if anything reaches it, so every
guarded call has to be short-circuited by the breaker.
"""
from redis.exceptions import ConnectionError as RedisConnectionError
from litellm.caching.redis_cache import _is_redis_timeout_failure
from litellm.constants import REDIS_CIRCUIT_BREAKER_FAILURE_THRESHOLD
with (
patch("asyncio.get_running_loop", side_effect=RuntimeError("No running event loop")),
patch( # test-quality-ok: RedisCache.__init__ builds its client eagerly, with no injection point
"litellm._redis.get_redis_client", return_value=MagicMock()
),
):
redis_cache = RedisCache(host="127.0.0.1", port=6379)
unreachable = AsyncMock()
unreachable.get.side_effect = AssertionError("an open breaker must not touch Redis")
unreachable.mget.side_effect = AssertionError("an open breaker must not touch Redis")
unreachable.set.side_effect = AssertionError("an open breaker must not touch Redis")
unreachable.pipeline.side_effect = AssertionError("an open breaker must not touch Redis")
for _ in range(REDIS_CIRCUIT_BREAKER_FAILURE_THRESHOLD):
redis_cache._circuit_breaker.record_failure(is_timeout=_is_redis_timeout_failure(RedisConnectionError("refused")))
assert redis_cache._circuit_breaker.is_open() is True
with patch.object(redis_cache, "init_async_client", return_value=unreachable):
yield DualCache(in_memory_cache=InMemoryCache(), redis_cache=redis_cache)
@pytest.mark.asyncio
@pytest.mark.parametrize(
"call, expected",
[
pytest.param(lambda c: c.async_get_cache("lit7468"), lambda n: None, id="async_get_cache"),
pytest.param(
lambda c: c.async_batch_get_cache(["lit7468", "lit7460"]), lambda n: [None, None], id="async_batch_get_cache"
),
pytest.param(lambda c: c.async_set_cache("lit7468", "v"), lambda n: None, id="async_set_cache"),
pytest.param(lambda c: c.async_set_cache_pipeline([("lit7468", "v")]), lambda n: None, id="async_set_cache_pipeline"),
pytest.param(lambda c: c.async_increment_cache("lit7468", 1.0, ttl=60), float, id="async_increment_cache"),
],
)
async def test_open_breaker_is_a_quiet_cache_miss(dual_cache_with_open_breaker, call, expected, caplog):
"""While the breaker is open, every cache operation must degrade to the in-memory result
without logging anything above DEBUG.
Before this, each skipped call raised a generic exception that DualCache caught and logged
as a full ERROR traceback. Under production request rates that was hundreds of stack
formats per second per replica, enough to pin every proxy at 100% CPU on a Redis blip.
"""
caplog.set_level(logging.DEBUG, logger="LiteLLM")
for n in range(1, 51):
assert await call(dual_cache_with_open_breaker) == expected(n)
noisy = [r for r in caplog.records if r.levelno > logging.DEBUG]
assert noisy == [], f"an open breaker must be silent per call, got {[r.getMessage() for r in noisy]}"

View file

@ -1013,3 +1013,131 @@ async def test_breaker_metrics_track_state_and_failure_class():
breaker.record_success()
assert sample("litellm_redis_circuit_breaker_state", {"state": "open"}) == open_gauge_before
assert sample("litellm_redis_circuit_breaker_state", {"state": "closed"}) == closed_gauge_before + 1
@pytest.mark.asyncio
async def test_pool_exhaustion_counts_as_a_timeout_not_a_hard_failure():
"""redis-py reports a blocking pool that waited out its timeout as a ConnectionError
chained from the underlying TimeoutError. That is Redis being slow, the same signal as
a read timeout, so a burst of them must go through the duration gate instead of opening
the breaker on the fifth one as though Redis had refused the connection.
"""
from redis.exceptions import ConnectionError as RedisConnectionError
from litellm.caching.redis_cache import (
RedisCircuitBreaker,
_is_redis_timeout_failure,
_run_under_circuit_breaker,
)
def pool_exhausted() -> RedisConnectionError:
"""Built the way redis-py's BlockingConnectionPool.get_connection raises it."""
try:
try:
raise asyncio.TimeoutError()
except asyncio.TimeoutError as err:
raise RedisConnectionError("No connection available.") from err
except RedisConnectionError as chained:
return chained
assert _is_redis_timeout_failure(pool_exhausted()) is True
assert _is_redis_timeout_failure(RedisConnectionError("Connection refused")) is False
breaker = RedisCircuitBreaker(failure_threshold=3, recovery_timeout=60, timeout_min_duration=5.0)
async def pool_exhausted_call():
raise pool_exhausted()
for _ in range(breaker.failure_threshold + 1):
with pytest.raises(RedisConnectionError, match="No connection available"):
await _run_under_circuit_breaker(breaker, "op", pool_exhausted_call)
assert breaker.is_open() is False, "an instantaneous burst of pool waits must not open the breaker"
@pytest.mark.asyncio
async def test_success_admitted_before_the_breaker_opened_cannot_close_it():
"""Only the HALF_OPEN recovery probe may close the breaker.
A call that was already in flight when the breaker opened knows nothing about whether
Redis has recovered. Letting its late success close the breaker made the state flap
OPEN -> CLOSED -> OPEN under load, and every OPEN transition re-logged the warning while
the next five failures each paid the full socket timeout again.
"""
from redis.exceptions import ConnectionError as RedisConnectionError
from litellm.caching.redis_cache import (
RedisCircuitBreaker,
RedisCircuitBreakerOpenError,
_run_under_circuit_breaker,
)
breaker = RedisCircuitBreaker(failure_threshold=2, recovery_timeout=60)
release_slow_success = asyncio.Event()
async def slow_success():
await release_slow_success.wait()
return "ok"
async def refused():
raise RedisConnectionError("refused")
in_flight = asyncio.create_task(_run_under_circuit_breaker(breaker, "slow", slow_success))
await asyncio.sleep(0)
for _ in range(breaker.failure_threshold):
with pytest.raises(RedisConnectionError):
await _run_under_circuit_breaker(breaker, "op", refused)
assert breaker.is_open() is True
release_slow_success.set()
assert await in_flight == "ok"
assert breaker.is_open() is True, "a pre-open success is not a recovery probe"
with pytest.raises(RedisCircuitBreakerOpenError):
await _run_under_circuit_breaker(breaker, "op", slow_success)
@pytest.mark.asyncio
async def test_open_breaker_raises_its_own_exception_type():
"""Callers with an optional cache need to tell the expected fast-fail apart from a real error."""
from redis.exceptions import ConnectionError as RedisConnectionError
from litellm.caching.redis_cache import (
RedisCircuitBreaker,
RedisCircuitBreakerOpenError,
_run_under_circuit_breaker,
)
breaker = RedisCircuitBreaker(failure_threshold=1, recovery_timeout=60)
async def refused():
raise RedisConnectionError("refused")
with pytest.raises(RedisConnectionError):
await _run_under_circuit_breaker(breaker, "op", refused)
with pytest.raises(RedisCircuitBreakerOpenError, match="circuit breaker is open"):
await _run_under_circuit_breaker(breaker, "op", refused)
@pytest.mark.asyncio
async def test_recovery_probe_still_closes_the_breaker():
from redis.exceptions import ConnectionError as RedisConnectionError
from litellm.caching.redis_cache import RedisCircuitBreaker, _run_under_circuit_breaker
breaker = RedisCircuitBreaker(failure_threshold=1, recovery_timeout=0.05)
async def refused():
raise RedisConnectionError("refused")
async def recovered():
return "ok"
with pytest.raises(RedisConnectionError):
await _run_under_circuit_breaker(breaker, "op", refused)
assert breaker.is_open() is True
await asyncio.sleep(0.06)
assert await _run_under_circuit_breaker(breaker, "probe", recovered) == "ok"
assert breaker.is_open() is False

View file

@ -525,3 +525,48 @@ class TestLoggingWorker:
asyncio.run(rebind_on_second_loop())
assert sorted(executed) == [0, 1, 2, 3, 4]
@pytest.mark.asyncio
async def test_callback_timeout_burst_logs_one_traceback_per_interval(self, caplog):
"""A slow logging backend times out every in-flight callback at once. Logging a full
traceback for each of them turns that stall into a CPU-bound log storm on every replica,
so the worker must emit one traceback per interval and count the rest.
"""
caplog.set_level(logging.DEBUG, logger="LiteLLM")
worker = LoggingWorker(timeout=0.05, max_queue_size=200, concurrency=100, error_traceback_interval=60.0)
worker.start()
async def stalled_callback():
await asyncio.sleep(10)
for _ in range(40):
worker.enqueue(stalled_callback())
await asyncio.sleep(0.5)
await worker.stop()
errors = [r for r in caplog.records if r.levelno >= logging.ERROR and "LoggingWorker error" in r.getMessage()]
assert len(errors) == 1, f"expected one traceback for the burst, got {len(errors)}"
assert errors[0].exc_info is not None
@pytest.mark.asyncio
async def test_error_traceback_resumes_after_interval_with_suppressed_count(self, caplog):
caplog.set_level(logging.DEBUG, logger="LiteLLM")
worker = LoggingWorker(timeout=0.05, max_queue_size=200, concurrency=100, error_traceback_interval=0.2)
worker.start()
async def stalled_callback():
await asyncio.sleep(10)
for _ in range(5):
worker.enqueue(stalled_callback())
await asyncio.sleep(0.3)
for _ in range(3):
worker.enqueue(stalled_callback())
await asyncio.sleep(0.3)
await worker.stop()
messages = [r.getMessage() for r in caplog.records if "LoggingWorker error" in r.getMessage()]
assert len(messages) == 2, messages
assert "(0 more suppressed" in messages[0]
assert "(4 more suppressed" in messages[1]