mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
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:
parent
a9cec50960
commit
32b7daf691
7 changed files with 326 additions and 24 deletions
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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%)
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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]}"
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue