diff --git a/litellm/caching/redis_cluster.py b/litellm/caching/redis_cluster.py new file mode 100644 index 00000000000..cec548c782c --- /dev/null +++ b/litellm/caching/redis_cluster.py @@ -0,0 +1,41 @@ +""" +Redis Cluster Cache implementation +""" + +from typing import TYPE_CHECKING, Any, Optional + +from redis.asyncio import RedisCluster + +from litellm.caching.redis_cache import RedisCache + +if TYPE_CHECKING: + from opentelemetry.trace import Span as _Span + from redis.asyncio import Redis + from redis.asyncio.client import Pipeline + + pipeline = Pipeline + async_redis_client = Redis + Span = _Span +else: + pipeline = Any + async_redis_client = Any + Span = Any + + +class RedisClusterCache(RedisCache): + def __init__(self, *args, **kwargs): + super().__init__(*args, **kwargs) + self.redis_cluster_client: Optional[RedisCluster] = None + + def init_async_client(self): + from .._redis import get_redis_async_client + + if self.redis_cluster_client: + return self.redis_cluster_client + + _redis_client = get_redis_async_client( + connection_pool=self.async_redis_conn_pool, **self.redis_kwargs + ) + if isinstance(_redis_client, RedisCluster): + self.redis_cluster_client = _redis_client + return _redis_client