diff --git a/litellm/caching/redis_cache.py b/litellm/caching/redis_cache.py index 7cec84e0ebb..b0c9535ca28 100644 --- a/litellm/caching/redis_cache.py +++ b/litellm/caching/redis_cache.py @@ -1785,7 +1785,8 @@ class RedisCache(BaseCache): self.redis_client.flushall() async def disconnect(self): - await self.async_redis_conn_pool.disconnect(inuse_connections=True) + if self.async_redis_conn_pool is not None: + await self.async_redis_conn_pool.disconnect(inuse_connections=True) try: self.redis_client.close() except Exception as e: diff --git a/tests/test_litellm/test_redis_disconnect.py b/tests/test_litellm/test_redis_disconnect.py new file mode 100644 index 00000000000..52f8db42c11 --- /dev/null +++ b/tests/test_litellm/test_redis_disconnect.py @@ -0,0 +1,27 @@ +from unittest.mock import AsyncMock, Mock + +import pytest + +from litellm.caching.redis_cache import RedisCache + + +@pytest.mark.asyncio +async def test_disconnect_allows_missing_async_pool(): + cache = RedisCache.__new__(RedisCache) + cache.async_redis_conn_pool = None + cache.redis_client = Mock() + + await cache.disconnect() + + cache.redis_client.close.assert_called_once_with() + + +@pytest.mark.asyncio +async def test_disconnect_closes_present_async_pool(): + cache = RedisCache.__new__(RedisCache) + cache.async_redis_conn_pool = Mock(disconnect=AsyncMock()) + cache.redis_client = Mock() + + await cache.disconnect() + + cache.async_redis_conn_pool.disconnect.assert_awaited_once_with(inuse_connections=True)