fix(caching): judge swallowed Redis failures per admission and throttle worker tracebacks per error type

Swallowed failures are now reported to the breaker when the admitted call exits, so a
call admitted before the breaker opened cannot refresh the open timer or knock out the
recovery probe. The sync batch read raises the typed open-breaker error like its async
twin so DualCache releases its batch reservations, and LoggingWorker throttles tracebacks
per exception class instead of worker-wide

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
mateo 2026-09-10 22:25:03 +00:00
parent df8a9c72ab
commit 7658cd53aa
6 changed files with 162 additions and 38 deletions

View file

@ -206,6 +206,9 @@ class DualCache(BaseCache):
redis_result: Final = self.redis_cache.batch_get_cache(
key_list=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)

View file

@ -285,7 +285,9 @@ class RedisCircuitBreaker:
_RedisCallResult = TypeVar("_RedisCallResult")
_swallowed_redis_failures: Final[ContextVar[int]] = ContextVar("litellm_swallowed_redis_failures", default=0)
_swallowed_redis_failures: Final[ContextVar[tuple[bool, ...]]] = ContextVar(
"litellm_swallowed_redis_failures", default=()
)
def _opaque_kwarg_key(value: object) -> str:
@ -405,8 +407,8 @@ def _breaker_metrics() -> _BreakerMetrics:
return _BreakerMetrics()
def _record_swallowed_redis_failure(breaker: RedisCircuitBreaker, exc: BaseException) -> None:
"""Record a Redis failure that the calling method is about to swallow.
def _record_swallowed_redis_failure(exc: BaseException) -> None:
"""Note a Redis failure that the calling method is about to swallow, for the breaker exit to judge.
The marker is a ContextVar rather than a counter on the breaker because breakers are
shared by every concurrent caller. A plain shared counter cannot tell "my call failed"
@ -416,8 +418,7 @@ def _record_swallowed_redis_failure(breaker: RedisCircuitBreaker, exc: BaseExcep
"""
if not _is_redis_health_failure(exc):
return
breaker.record_failure(is_timeout=_is_redis_timeout_failure(exc))
_swallowed_redis_failures.set(_swallowed_redis_failures.get() + 1)
_swallowed_redis_failures.set((*_swallowed_redis_failures.get(), _is_redis_timeout_failure(exc)))
@dataclass(frozen=True, slots=True)
@ -430,25 +431,42 @@ def _enter_circuit_breaker(breaker: RedisCircuitBreaker, name: str) -> _BreakerA
"""Reject the call if the breaker is open, else snapshot what its outcome will be judged against."""
if breaker.is_open():
raise RedisCircuitBreakerOpenError(f"Redis circuit breaker is open, skipping {name}")
return _BreakerAdmission(swallowed_before=_swallowed_redis_failures.get(), generation=breaker.generation)
return _BreakerAdmission(swallowed_before=len(_swallowed_redis_failures.get()), generation=breaker.generation)
def _take_swallowed_failures(admission: _BreakerAdmission) -> tuple[bool, ...]:
"""Return the is_timeout flag of every failure this call swallowed, and drop them from the context."""
all_swallowed: Final = _swallowed_redis_failures.get()
_swallowed_redis_failures.set(all_swallowed[: admission.swallowed_before])
return all_swallowed[admission.swallowed_before :]
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.
"""Report the call's outcome to the breaker generation it was admitted under.
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.
method that returned is not on its own proof of a healthy Redis. A call admitted
before the breaker opened reports nothing: its failures would refresh the open
timer or knock out the recovery probe, and its success would close it early.
"""
swallowed: Final = _take_swallowed_failures(admission)
if admission.generation != breaker.generation:
return
if _swallowed_redis_failures.get() == admission.swallowed_before:
if not swallowed:
breaker.record_success()
return
for is_timeout in swallowed:
breaker.record_failure(is_timeout=is_timeout)
def _fail_circuit_breaker(breaker: RedisCircuitBreaker, admission: _BreakerAdmission, exc: BaseException) -> None:
if admission.generation != breaker.generation or not _is_redis_health_failure(exc):
swallowed: Final = _take_swallowed_failures(admission)
if admission.generation != breaker.generation:
return
breaker.record_failure(is_timeout=_is_redis_timeout_failure(exc))
for is_timeout in swallowed:
breaker.record_failure(is_timeout=is_timeout)
if _is_redis_health_failure(exc):
breaker.record_failure(is_timeout=_is_redis_timeout_failure(exc))
async def _run_under_circuit_breaker(
@ -1056,7 +1074,7 @@ class RedisCache(BaseCache):
str(e),
value,
)
_record_swallowed_redis_failure(self._circuit_breaker, e)
_record_swallowed_redis_failure(e)
async def _pipeline_helper(
self,
@ -1143,7 +1161,7 @@ class RedisCache(BaseCache):
str(e),
cache_value,
)
_record_swallowed_redis_failure(self._circuit_breaker, e)
_record_swallowed_redis_failure(e)
async def _set_cache_sadd_helper(
self,
@ -1228,7 +1246,7 @@ class RedisCache(BaseCache):
str(e),
value,
)
_record_swallowed_redis_failure(self._circuit_breaker, e)
_record_swallowed_redis_failure(e)
@_redis_circuit_breaker_guard
async def batch_cache_write(self, key, value, **kwargs):
@ -1386,7 +1404,7 @@ class RedisCache(BaseCache):
except Exception as e:
# NON blocking - notify users Redis is throwing an exception
verbose_logger.error("litellm.caching.caching: get() - Got exception from REDIS: %s", e)
_record_swallowed_redis_failure(self._circuit_breaker, e)
_record_swallowed_redis_failure(e)
def _run_redis_mget_operation(self, keys: list[str]) -> Sequence[bytes | str | None]:
"""
@ -1423,9 +1441,9 @@ class RedisCache(BaseCache):
key_value_dict = {}
_key_list: Final = [key for key in key_list if key is not None]
start_time: Final = time.time()
admission: Final = _enter_circuit_breaker(self._circuit_breaker, "batch_get_cache")
try:
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, admission)
@ -1452,8 +1470,6 @@ 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(
@ -1466,7 +1482,7 @@ class RedisCache(BaseCache):
parent_otel_span=parent_otel_span,
)
verbose_logger.error("Error occurred in batch get cache - %s", e)
_record_swallowed_redis_failure(self._circuit_breaker, e)
_fail_circuit_breaker(self._circuit_breaker, admission, e)
return key_value_dict
@_redis_circuit_breaker_guard
@ -1513,7 +1529,7 @@ class RedisCache(BaseCache):
)
)
print_verbose(f"litellm.caching.caching: async get() - Got exception from REDIS: {e}")
_record_swallowed_redis_failure(self._circuit_breaker, e)
_record_swallowed_redis_failure(e)
@_redis_circuit_breaker_guard
async def async_batch_get_cache(
@ -1585,7 +1601,7 @@ class RedisCache(BaseCache):
)
)
verbose_logger.error("Error occurred in async batch get cache - %s", e)
_record_swallowed_redis_failure(self._circuit_breaker, e)
_record_swallowed_redis_failure(e)
return key_value_dict
def sync_ping(self) -> bool:
@ -1837,7 +1853,7 @@ class RedisCache(BaseCache):
return ttl
except Exception as e:
verbose_logger.debug("Redis TTL Error: %s", e)
_record_swallowed_redis_failure(self._circuit_breaker, e)
_record_swallowed_redis_failure(e)
return None
@_redis_circuit_breaker_guard

View file

@ -55,8 +55,8 @@ class LoggingWorker:
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._last_error_traceback_at: dict[type[BaseException], float] = {}
self._errors_since_traceback: dict[type[BaseException], int] = {}
self._queue: asyncio.Queue[LoggingTask] | None = None
self._worker_task: asyncio.Task | None = None
self._running_tasks: set[asyncio.Task] = set()
@ -178,19 +178,21 @@ class LoggingWorker:
sem.release()
def _log_task_error(self, error: Exception) -> None:
"""One traceback per interval: a stalled backend fails every in-flight task at once."""
"""One traceback per error type 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
error_type: Final = type(error)
last_traceback_at: Final = self._last_error_traceback_at.get(error_type)
if last_traceback_at is not None and now - last_traceback_at < self.error_traceback_interval:
self._errors_since_traceback += 1
self._errors_since_traceback[error_type] = self._errors_since_traceback.get(error_type, 0) + 1
return
verbose_logger.exception(
"LoggingWorker error (%d more suppressed since the last traceback): %r",
self._errors_since_traceback,
"LoggingWorker error (%d more %s suppressed since the last traceback): %r",
self._errors_since_traceback.get(error_type, 0),
error_type.__name__,
error,
)
self._last_error_traceback_at = now
self._errors_since_traceback = 0
self._last_error_traceback_at[error_type] = now
self._errors_since_traceback[error_type] = 0
async def _worker_loop(self) -> None:
"""Main worker loop that gets tasks and schedules them to run concurrently."""

View file

@ -653,3 +653,23 @@ def test_open_breaker_is_a_quiet_cache_miss_on_the_sync_read_path(dual_cache_wit
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]}"
def test_open_breaker_does_not_leave_sync_batch_reservations_behind(dual_cache_with_open_breaker, caplog):
"""A batch read skipped by the breaker must not hold its keys for the batch expiry window.
The sync read reserves keys before dialing Redis so concurrent callers do not all hit it.
When the breaker rejects the read, those reservations have to be released, otherwise the
first read after Redis recovers is still throttled for up to default_redis_batch_cache_expiry.
"""
caplog.set_level(logging.DEBUG, logger="LiteLLM")
dual_cache_with_open_breaker.redis_cache.redis_client.mget.side_effect = AssertionError(
"an open breaker must not touch Redis"
)
assert dual_cache_with_open_breaker.batch_get_cache(keys=["lit7468", "lit7460"]) == [None, None]
assert "lit7468" not in dual_cache_with_open_breaker.last_redis_batch_access_time
assert "lit7460" not in dual_cache_with_open_breaker.last_redis_batch_access_time
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

@ -516,14 +516,20 @@ async def test_circuit_breaker_opens_when_method_swallows_redis_failure(call_met
await call_method(cache)
def test_circuit_breaker_open_keeps_sync_batch_get_cache_as_a_miss(sync_batch_redis_cache):
"""An open breaker must preserve the sync batch read's dictionary fallback."""
def test_sync_batch_get_cache_swallowed_failures_open_the_breaker_and_then_fast_fail(sync_batch_redis_cache):
"""The sync batch read hides its Redis error behind an empty dict, but the breaker must still
count it, and once open the read has to raise the typed error like its async twin so DualCache
can release the batch reservations it took before the call.
"""
from litellm.caching.redis_cache import RedisCircuitBreakerOpenError
from litellm.constants import REDIS_CIRCUIT_BREAKER_FAILURE_THRESHOLD
for _ in range(REDIS_CIRCUIT_BREAKER_FAILURE_THRESHOLD):
assert sync_batch_redis_cache.batch_get_cache(key_list=["lit6729"]) == {}
assert sync_batch_redis_cache.batch_get_cache(key_list=["lit6729"]) == {}
with pytest.raises(RedisCircuitBreakerOpenError):
sync_batch_redis_cache.batch_get_cache(key_list=["lit6729"])
assert sync_batch_redis_cache.redis_client.mget.call_count == REDIS_CIRCUIT_BREAKER_FAILURE_THRESHOLD
def test_batch_get_counts_raises_where_batch_get_cache_reports_a_miss(sync_batch_redis_cache):
@ -647,6 +653,7 @@ def test_sync_batch_get_cache_survives_a_service_callback_that_raises(
from concurrent.futures import ThreadPoolExecutor
import litellm
from litellm.caching.redis_cache import RedisCircuitBreakerOpenError
from litellm.constants import REDIS_CIRCUIT_BREAKER_FAILURE_THRESHOLD
cache, service_logger = sync_batch_cache_with_service_logger
@ -661,7 +668,8 @@ def test_sync_batch_get_cache_survives_a_service_callback_that_raises(
with ThreadPoolExecutor(max_workers=1) as pool:
assert pool.submit(cache.batch_get_cache, key_list=["lit6729"]).result() == {}
assert cache.batch_get_cache(key_list=["lit6729"]) == {}
with pytest.raises(RedisCircuitBreakerOpenError):
cache.batch_get_cache(key_list=["lit6729"])
def test_call_stack_info_skips_breaker_guard_frames():
@ -801,7 +809,7 @@ async def test_concurrent_success_is_not_cancelled_by_another_calls_failure():
# starts would leave its snapshot correct and prove nothing.
async def swallows_a_failure():
await asyncio.sleep(0.02)
_record_swallowed_redis_failure(breaker, RedisConnectionError("redis unreachable"))
_record_swallowed_redis_failure(RedisConnectionError("redis unreachable"))
async def succeeds_while_the_other_fails():
await asyncio.sleep(0.05)
@ -1097,6 +1105,53 @@ async def test_success_admitted_before_the_breaker_opened_cannot_close_it():
await _run_under_circuit_breaker(breaker, "op", slow_success)
@pytest.mark.asyncio
async def test_swallowed_failure_admitted_before_the_breaker_opened_cannot_delay_recovery():
"""A stale in-flight call that swallows its Redis error must not restart the open timer.
Guarded methods that return a default instead of raising still report their failure to
the breaker on exit. If that report landed against a newer breaker generation, every
slow call that was already dialing Redis when the breaker opened would push recovery
back by its own socket timeout, and knock out the HALF_OPEN probe if it landed then.
"""
from redis.exceptions import ConnectionError as RedisConnectionError
from litellm.caching.redis_cache import (
RedisCircuitBreaker,
_record_swallowed_redis_failure,
_run_under_circuit_breaker,
)
breaker = RedisCircuitBreaker(failure_threshold=2, recovery_timeout=0.05)
release_stale_call = asyncio.Event()
async def stale_call_that_swallows_its_failure():
await release_stale_call.wait()
_record_swallowed_redis_failure(RedisConnectionError("refused"))
return {}
async def refused():
raise RedisConnectionError("refused")
async def recovered():
return "ok"
in_flight = asyncio.create_task(_run_under_circuit_breaker(breaker, "stale", stale_call_that_swallows_its_failure))
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
await asyncio.sleep(0.04)
release_stale_call.set()
assert await in_flight == {}
await asyncio.sleep(0.03)
assert await _run_under_circuit_breaker(breaker, "probe", recovered) == "ok"
assert breaker.is_open() is False
@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."""

View file

@ -568,5 +568,33 @@ class TestLoggingWorker:
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]
assert "(0 more TimeoutError suppressed" in messages[0]
assert "(4 more TimeoutError suppressed" in messages[1]
@pytest.mark.asyncio
async def test_traceback_throttle_is_per_error_type(self, caplog):
"""A timeout burst from one stalled backend must not hide the first traceback of a different
failure, otherwise a misconfigured callback stays invisible for the whole interval.
"""
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)
async def misconfigured_callback():
raise KeyError("missing api key")
for _ in range(20):
worker.enqueue(stalled_callback())
await asyncio.sleep(0.2)
for _ in range(3):
worker.enqueue(misconfigured_callback())
await asyncio.sleep(0.2)
await worker.stop()
messages = [r.getMessage() for r in caplog.records if "LoggingWorker error" in r.getMessage()]
assert len(messages) == 2, messages
assert "TimeoutError" in messages[0]
assert "KeyError" in messages[1] and "missing api key" in messages[1]