feat(redis): add RedisCircuitBreaker and apply guard decorator to all async ops

This commit is contained in:
Ishaan Jaffer 2026-03-19 20:07:19 -07:00
parent 1ea1f99d0e
commit 0f779650ed

View file

@ -10,6 +10,7 @@ Has 4 primary methods:
import ast
import asyncio
import functools
import hashlib
import inspect
import json
@ -19,7 +20,11 @@ from typing import TYPE_CHECKING, Any, List, Optional, Tuple, Union, cast
import litellm
from litellm._logging import print_verbose, verbose_logger
from litellm.constants import DEFAULT_REDIS_MAJOR_VERSION
from litellm.constants import (
DEFAULT_REDIS_MAJOR_VERSION,
REDIS_CIRCUIT_BREAKER_FAILURE_THRESHOLD,
REDIS_CIRCUIT_BREAKER_RECOVERY_TIMEOUT,
)
from litellm.litellm_core_utils.core_helpers import _get_parent_otel_span_from_kwargs
from litellm.litellm_core_utils.coroutine_checker import coroutine_checker
from litellm.types.caching import (
@ -89,6 +94,86 @@ def _get_call_stack_info(num_frames: int = 2) -> str:
return "unknown"
class RedisCircuitBreaker:
"""
Tracks Redis health for a RedisCache instance.
States:
CLOSED - normal, Redis is called
OPEN - Redis is down, raise immediately (no network call)
HALF_OPEN - recovery probe: allow one request through
Transitions:
CLOSED -> OPEN after failure_threshold consecutive failures
OPEN -> HALF_OPEN after recovery_timeout seconds
HALF_OPEN -> CLOSED on success
HALF_OPEN -> OPEN on failure (resets timer)
"""
CLOSED = "closed"
OPEN = "open"
HALF_OPEN = "half_open"
def __init__(self, failure_threshold: int, recovery_timeout: int) -> None:
self.failure_threshold = failure_threshold
self.recovery_timeout = recovery_timeout
self._failure_count = 0
self._opened_at: Optional[float] = None
self._state = self.CLOSED
def is_open(self) -> bool:
"""Returns True if Redis calls should be skipped."""
if self._state == self.OPEN:
if time.time() - (self._opened_at or 0) > self.recovery_timeout:
self._state = self.HALF_OPEN
return False # allow probe through
return True
return False
def record_failure(self) -> None:
self._failure_count += 1
self._opened_at = time.time()
if self._failure_count >= self.failure_threshold:
if self._state != self.OPEN:
verbose_logger.warning(
"Redis circuit breaker OPENED after %d consecutive failures — "
"fast-failing Redis calls for %ds",
self._failure_count,
self.recovery_timeout,
)
self._state = self.OPEN
def record_success(self) -> None:
if self._state == self.HALF_OPEN:
verbose_logger.info("Redis circuit breaker CLOSED — Redis recovered")
self._failure_count = 0
self._state = self.CLOSED
def _redis_circuit_breaker_guard(method): # type: ignore
"""
Decorator for RedisCache async methods.
Checks the circuit breaker before each call; records success/failure after.
Does not apply to ping/disconnect/test_connection (health/teardown must always run).
"""
@functools.wraps(method)
async def wrapper(self, *args, **kwargs): # type: ignore
if self._circuit_breaker.is_open():
raise Exception(
f"Redis circuit breaker is open — skipping {method.__name__}"
)
try:
result = await method(self, *args, **kwargs)
self._circuit_breaker.record_success()
return result
except Exception as e:
self._circuit_breaker.record_failure()
raise
return wrapper
class RedisCache(BaseCache):
# if users don't provider one, use the default litellm cache
@ -150,6 +235,11 @@ class RedisCache(BaseCache):
except Exception:
pass
self._circuit_breaker = RedisCircuitBreaker(
failure_threshold=REDIS_CIRCUIT_BREAKER_FAILURE_THRESHOLD,
recovery_timeout=REDIS_CIRCUIT_BREAKER_RECOVERY_TIMEOUT,
)
self._setup_health_pings()
if litellm.default_redis_ttl is not None:
@ -375,6 +465,7 @@ class RedisCache(BaseCache):
)
raise e
@_redis_circuit_breaker_guard
async def async_scan_iter(self, pattern: str, count: int = 100) -> list:
start_time = time.time()
try:
@ -451,6 +542,7 @@ class RedisCache(BaseCache):
verbose_logger.error(f"Error registering Redis script: {str(e)}")
raise e
@_redis_circuit_breaker_guard
async def async_set_cache(self, key, value, **kwargs):
from redis.asyncio import Redis
@ -560,6 +652,7 @@ class RedisCache(BaseCache):
results = await pipe.execute()
return results
@_redis_circuit_breaker_guard
async def async_set_cache_pipeline(
self, cache_list: List[Tuple[Any, Any]], ttl: Optional[float] = None, **kwargs
):
@ -636,6 +729,7 @@ class RedisCache(BaseCache):
except Exception:
raise
@_redis_circuit_breaker_guard
async def async_set_cache_sadd(
self, key, value: List, ttl: Optional[float], **kwargs
):
@ -708,6 +802,7 @@ class RedisCache(BaseCache):
value,
)
@_redis_circuit_breaker_guard
async def batch_cache_write(self, key, value, **kwargs):
print_verbose(
f"in batch cache writing for redis buffer size={len(self.redis_batch_writing_buffer)}",
@ -717,6 +812,7 @@ class RedisCache(BaseCache):
if len(self.redis_batch_writing_buffer) >= self.redis_flush_size:
await self.flush_cache_buffer() # logging done in here
@_redis_circuit_breaker_guard
async def async_increment(
self,
key,
@ -894,6 +990,7 @@ class RedisCache(BaseCache):
verbose_logger.error(f"Error occurred in batch get cache - {str(e)}")
return key_value_dict
@_redis_circuit_breaker_guard
async def async_get_cache(
self, key, parent_otel_span: Optional[Span] = None, **kwargs
):
@ -944,6 +1041,7 @@ class RedisCache(BaseCache):
f"litellm.caching.caching: async get() - Got exception from REDIS: {str(e)}"
)
@_redis_circuit_breaker_guard
async def async_batch_get_cache(
self,
key_list: Union[List[str], List[Optional[str]]],
@ -1087,6 +1185,7 @@ class RedisCache(BaseCache):
)
raise e
@_redis_circuit_breaker_guard
async def delete_cache_keys(self, keys):
# typed as Any, redis python lib has incomplete type stubs for RedisCluster and does not include `delete`
_redis_client: Any = self.init_async_client()
@ -1151,6 +1250,7 @@ class RedisCache(BaseCache):
"error": str(e),
}
@_redis_circuit_breaker_guard
async def async_delete_cache(self, key: str):
# typed as Any, redis python lib has incomplete type stubs for RedisCluster and does not include `delete`
_redis_client: Any = self.init_async_client()
@ -1184,6 +1284,7 @@ class RedisCache(BaseCache):
)
return [r for r in results if isinstance(r, float)]
@_redis_circuit_breaker_guard
async def async_increment_pipeline(
self, increment_list: List[RedisPipelineIncrementOperation], **kwargs
) -> Optional[List[float]]:
@ -1247,6 +1348,7 @@ class RedisCache(BaseCache):
)
raise e
@_redis_circuit_breaker_guard
async def async_get_ttl(self, key: str) -> Optional[int]:
"""
Get the remaining TTL of a key in Redis
@ -1270,6 +1372,7 @@ class RedisCache(BaseCache):
verbose_logger.debug(f"Redis TTL Error: {e}")
return None
@_redis_circuit_breaker_guard
async def async_rpush(
self,
key: str,
@ -1336,6 +1439,7 @@ class RedisCache(BaseCache):
raise r
return results
@_redis_circuit_breaker_guard
async def async_rpush_pipeline(
self,
rpush_list: List[RedisPipelineRpushOperation],
@ -1405,6 +1509,7 @@ class RedisCache(BaseCache):
return result
@_redis_circuit_breaker_guard
async def async_lpop(
self,
key: str,
@ -1534,6 +1639,7 @@ class RedisCache(BaseCache):
decoded_results.append(None)
return decoded_results
@_redis_circuit_breaker_guard
async def async_lpop_pipeline(
self,
lpop_list: List[RedisPipelineLpopOperation],