From db92956ae33ed4c4e3233d7e1b0c7229817159bf Mon Sep 17 00:00:00 2001 From: Krrish Dholakia Date: Mon, 17 Mar 2025 14:27:36 -0700 Subject: [PATCH] fix(redis_cache.py): add 5s default timeout --- litellm/_redis.py | 5 +---- litellm/caching/redis_cache.py | 5 +++++ litellm/proxy/_new_secret_config.yaml | 10 +++++++++- tests/litellm/caching/test_redis_cache.py | 13 +++++++++++++ 4 files changed, 28 insertions(+), 5 deletions(-) diff --git a/litellm/_redis.py b/litellm/_redis.py index 1e03993c20e..5b2f85b1afc 100644 --- a/litellm/_redis.py +++ b/litellm/_redis.py @@ -182,9 +182,7 @@ def init_redis_cluster(redis_kwargs) -> redis.RedisCluster: "REDIS_CLUSTER_NODES environment variable is not valid JSON. Please ensure it's properly formatted." ) - verbose_logger.debug( - "init_redis_cluster: startup nodes are being initialized." - ) + verbose_logger.debug("init_redis_cluster: startup nodes are being initialized.") from redis.cluster import ClusterNode args = _get_redis_cluster_kwargs() @@ -307,7 +305,6 @@ def get_redis_async_client( return _init_async_redis_sentinel(redis_kwargs) return async_redis.Redis( - socket_timeout=5, **redis_kwargs, ) diff --git a/litellm/caching/redis_cache.py b/litellm/caching/redis_cache.py index 66245e7476d..0571ac9f15f 100644 --- a/litellm/caching/redis_cache.py +++ b/litellm/caching/redis_cache.py @@ -54,6 +54,7 @@ class RedisCache(BaseCache): redis_flush_size: Optional[int] = 100, namespace: Optional[str] = None, startup_nodes: Optional[List] = None, # for redis-cluster + socket_timeout: Optional[float] = 5.0, # default 5 second timeout **kwargs, ): @@ -70,6 +71,9 @@ class RedisCache(BaseCache): redis_kwargs["password"] = password if startup_nodes is not None: redis_kwargs["startup_nodes"] = startup_nodes + if socket_timeout is not None: + redis_kwargs["socket_timeout"] = socket_timeout + ### HEALTH MONITORING OBJECT ### if kwargs.get("service_logger_obj", None) is not None and isinstance( kwargs["service_logger_obj"], ServiceLogging @@ -556,6 +560,7 @@ class RedisCache(BaseCache): ## LOGGING ## end_time = time.time() _duration = end_time - start_time + asyncio.create_task( self.service_logger_obj.async_service_success_hook( service=ServiceTypes.REDIS, diff --git a/litellm/proxy/_new_secret_config.yaml b/litellm/proxy/_new_secret_config.yaml index 9918ce429e2..64100277a80 100644 --- a/litellm/proxy/_new_secret_config.yaml +++ b/litellm/proxy/_new_secret_config.yaml @@ -6,4 +6,12 @@ model_list: api_base: os.environ/AZURE_API_BASE litellm_settings: - callbacks: ["prometheus"] \ No newline at end of file + callbacks: ["prometheus"] + +router_settings: + routing_strategy: usage-based-routing-v2 # 👈 KEY CHANGE + redis_host: os.environ/REDIS_HOST + redis_password: os.environ/REDIS_PASSWORD + redis_port: os.environ/REDIS_PORT + + diff --git a/tests/litellm/caching/test_redis_cache.py b/tests/litellm/caching/test_redis_cache.py index c2549d7fce7..3b7bdc56296 100644 --- a/tests/litellm/caching/test_redis_cache.py +++ b/tests/litellm/caching/test_redis_cache.py @@ -1,9 +1,13 @@ +import asyncio import json import os import sys +import time from unittest.mock import MagicMock, patch +import httpx import pytest +import respx from fastapi.testclient import TestClient sys.path.insert( @@ -39,3 +43,12 @@ async def test_redis_cache_async_increment(namespace): mock_redis_instance.incrbyfloat.assert_called_once_with( name=expected_key, amount=1 ) + + +@pytest.mark.asyncio +async def test_redis_client_init_with_socket_timeout(): + redis_cache = RedisCache(socket_timeout=1.0) + assert redis_cache.redis_kwargs["socket_timeout"] == 1.0 + client = redis_cache.init_async_client() + assert client is not None + assert client.connection_pool.connection_kwargs["socket_timeout"] == 1.0