From 4ee31f3ff538928da688df7430ae689b36ef77c3 Mon Sep 17 00:00:00 2001 From: Ishaan Jaff Date: Mon, 10 Feb 2025 13:23:53 -0800 Subject: [PATCH] fix redis cluster --- litellm/caching/redis_cache.py | 14 +++++++++----- litellm/caching/redis_cluster.py | 3 +++ litellm/stubs/redis/asyncio/cluster.pyi | 8 ++++++++ pyrightconfig.json | 1 + 4 files changed, 21 insertions(+), 5 deletions(-) create mode 100644 litellm/stubs/redis/asyncio/cluster.pyi diff --git a/litellm/caching/redis_cache.py b/litellm/caching/redis_cache.py index 170f3f45599..298d57e14c4 100644 --- a/litellm/caching/redis_cache.py +++ b/litellm/caching/redis_cache.py @@ -26,15 +26,17 @@ from .base_cache import BaseCache if TYPE_CHECKING: from opentelemetry.trace import Span as _Span - from redis.asyncio import Redis + from redis.asyncio import Redis, RedisCluster from redis.asyncio.client import Pipeline pipeline = Pipeline async_redis_client = Redis + async_redis_cluster_client = RedisCluster Span = _Span else: pipeline = Any async_redis_client = Any + async_redis_cluster_client = Any Span = Any @@ -122,7 +124,9 @@ class RedisCache(BaseCache): else: super().__init__() # defaults to 60s - def init_async_client(self): + def init_async_client( + self, + ) -> Union[async_redis_client, async_redis_cluster_client]: from .._redis import get_redis_async_client return get_redis_async_client( @@ -385,7 +389,7 @@ class RedisCache(BaseCache): return from redis.asyncio import Redis - _redis_client: Redis = self.init_async_client() + _redis_client = self.init_async_client() start_time = time.time() print_verbose( @@ -740,7 +744,7 @@ class RedisCache(BaseCache): """ Use Redis for bulk read operations """ - _redis_client = await self.init_async_client() + _redis_client = self.init_async_client() key_value_dict = {} start_time = time.time() try: @@ -1001,7 +1005,7 @@ class RedisCache(BaseCache): Redis ref: https://redis.io/docs/latest/commands/ttl/ """ try: - _redis_client = await self.init_async_client() + _redis_client = self.init_async_client() async with _redis_client as redis_client: ttl = await redis_client.ttl(key) if ttl <= -1: # -1 means the key does not exist, -2 key does not exist diff --git a/litellm/caching/redis_cluster.py b/litellm/caching/redis_cluster.py index cec548c782c..04746a73184 100644 --- a/litellm/caching/redis_cluster.py +++ b/litellm/caching/redis_cluster.py @@ -1,5 +1,8 @@ """ Redis Cluster Cache implementation + +Key differences: +- RedisClient NEEDs to be re-used across requests, adds 3000ms latency if it's re-created """ from typing import TYPE_CHECKING, Any, Optional diff --git a/litellm/stubs/redis/asyncio/cluster.pyi b/litellm/stubs/redis/asyncio/cluster.pyi new file mode 100644 index 00000000000..fa139a90d7a --- /dev/null +++ b/litellm/stubs/redis/asyncio/cluster.pyi @@ -0,0 +1,8 @@ +""" +Custom stub for redis.asyncio.cluster to add missing attributes for RedisCluster and ClusterPipeline. +""" + +from redis.asyncio.client import Redis + +class RedisCluster(Redis): + pass diff --git a/pyrightconfig.json b/pyrightconfig.json index 9a43abda78c..d8ef893fca2 100644 --- a/pyrightconfig.json +++ b/pyrightconfig.json @@ -1,5 +1,6 @@ { "ignore": [], + "stubPath": "litellm/stubs", "exclude": ["**/node_modules", "**/__pycache__", "litellm/types/utils.py"], "reportMissingImports": false, "reportPrivateImportUsage": false