diff --git a/litellm/_redis.py b/litellm/_redis.py index b3ec3bf8ef8..6d5b1197b56 100644 --- a/litellm/_redis.py +++ b/litellm/_redis.py @@ -77,19 +77,23 @@ def _get_redis_cluster_kwargs(client=None): # Only allow primitive arguments exclude_args = {"self", "connection_pool", "retry", "host", "port", "startup_nodes"} - available_args = [x for x in arg_spec.args if x not in exclude_args] - available_args.append("password") - available_args.append("username") - available_args.append("ssl") - available_args.append("ssl_cert_reqs") - available_args.append("ssl_check_hostname") - available_args.append("ssl_ca_certs") - available_args.append( - "redis_connect_func" - ) # Needed for sync clusters and IAM detection - available_args.append("gcp_service_account") - available_args.append("gcp_ssl_ca_certs") - available_args.append("max_connections") + available_args = {x for x in arg_spec.args if x not in exclude_args} + available_args |= { + "password", + "username", + "ssl", + "ssl_cert_reqs", + "ssl_check_hostname", + "ssl_ca_certs", + "redis_connect_func", # Needed for sync clusters and IAM detection + "gcp_service_account", + "gcp_ssl_ca_certs", + "max_connections", + "socket_timeout", + "socket_connect_timeout", + "socket_keepalive", + "socket_keepalive_options", + } return available_args @@ -319,6 +323,8 @@ def _init_redis_sentinel(redis_kwargs) -> redis.Redis: service_name = redis_kwargs.get("service_name") connection_kwargs = _get_redis_sentinel_connection_kwargs(redis_kwargs) sentinel_kwargs = dict(connection_kwargs) + connection_kwargs.setdefault("socket_timeout", REDIS_SOCKET_TIMEOUT) + sentinel_kwargs.setdefault("socket_timeout", connection_kwargs["socket_timeout"]) sentinel_kwargs["password"] = sentinel_password if not sentinel_nodes or not service_name: @@ -331,14 +337,12 @@ def _init_redis_sentinel(redis_kwargs) -> redis.Redis: # Set up the Sentinel client sentinel = redis.Sentinel( sentinel_nodes, - socket_timeout=REDIS_SOCKET_TIMEOUT, sentinel_kwargs=sentinel_kwargs, - **connection_kwargs, ) # Return the master instance for the given service - return sentinel.master_for(service_name) + return sentinel.master_for(service_name, **connection_kwargs) def _init_async_redis_sentinel(redis_kwargs) -> async_redis.Redis: @@ -347,6 +351,8 @@ def _init_async_redis_sentinel(redis_kwargs) -> async_redis.Redis: service_name = redis_kwargs.get("service_name") connection_kwargs = _get_redis_sentinel_connection_kwargs(redis_kwargs) sentinel_kwargs = dict(connection_kwargs) + connection_kwargs.setdefault("socket_timeout", REDIS_SOCKET_TIMEOUT) + sentinel_kwargs.setdefault("socket_timeout", connection_kwargs["socket_timeout"]) sentinel_kwargs["password"] = sentinel_password if not sentinel_nodes or not service_name: @@ -359,14 +365,12 @@ def _init_async_redis_sentinel(redis_kwargs) -> async_redis.Redis: # Set up the Sentinel client sentinel = async_redis.Sentinel( sentinel_nodes, - socket_timeout=REDIS_SOCKET_TIMEOUT, sentinel_kwargs=sentinel_kwargs, - **connection_kwargs, ) # Return the master instance for the given service - return sentinel.master_for(service_name) + return sentinel.master_for(service_name, **connection_kwargs) def get_redis_client(**env_overrides): diff --git a/tests/test_litellm/test_redis.py b/tests/test_litellm/test_redis.py index e5a63f6c5c2..4962d5b3fa1 100644 --- a/tests/test_litellm/test_redis.py +++ b/tests/test_litellm/test_redis.py @@ -371,17 +371,19 @@ def test_sync_sentinel_uses_sentinel_password_and_master_password(mock_sentinel_ ssl_check_hostname=True, ssl_ca_certs="/tmp/test-ca.pem", max_connections=17, + socket_timeout=5, ) mock_sentinel_cls.assert_called_once() sentinel_call_kwargs = mock_sentinel_cls.call_args[1] - assert sentinel_call_kwargs["password"] == "redis-secret" - assert sentinel_call_kwargs["username"] == "redis-user" - assert sentinel_call_kwargs["ssl"] is True - assert sentinel_call_kwargs["ssl_cert_reqs"] == "required" - assert sentinel_call_kwargs["ssl_check_hostname"] is True - assert sentinel_call_kwargs["ssl_ca_certs"] == "/tmp/test-ca.pem" - assert sentinel_call_kwargs["max_connections"] == 17 + assert "password" not in sentinel_call_kwargs + assert "username" not in sentinel_call_kwargs + assert "ssl" not in sentinel_call_kwargs + assert "ssl_cert_reqs" not in sentinel_call_kwargs + assert "ssl_check_hostname" not in sentinel_call_kwargs + assert "ssl_ca_certs" not in sentinel_call_kwargs + assert "max_connections" not in sentinel_call_kwargs + assert "socket_timeout" not in sentinel_call_kwargs assert sentinel_call_kwargs["sentinel_kwargs"] == { "password": "sentinel-secret", "username": "redis-user", @@ -390,11 +392,47 @@ def test_sync_sentinel_uses_sentinel_password_and_master_password(mock_sentinel_ "ssl_check_hostname": True, "ssl_ca_certs": "/tmp/test-ca.pem", "max_connections": 17, + "socket_timeout": 5, } assert "service_name" not in sentinel_call_kwargs["sentinel_kwargs"] assert "sentinel_nodes" not in sentinel_call_kwargs["sentinel_kwargs"] assert "sentinel_password" not in sentinel_call_kwargs["sentinel_kwargs"] - mock_sentinel.master_for.assert_called_once_with("mymaster") + mock_sentinel.master_for.assert_called_once_with( + "mymaster", + password="redis-secret", + username="redis-user", + ssl=True, + ssl_cert_reqs="required", + ssl_check_hostname=True, + ssl_ca_certs="/tmp/test-ca.pem", + max_connections=17, + socket_timeout=5, + ) + + +def test_sync_sentinel_socket_timeout_in_connection_kwargs_no_longer_raises(): + """socket_timeout should be applied through master_for without duplicating constructor kwargs.""" + mock_sentinel = MagicMock() + with patch( + "litellm._redis._get_redis_sentinel_connection_kwargs", + return_value={"password": "redis-secret", "socket_timeout": 5}, + ), patch("litellm._redis.redis.Sentinel", return_value=mock_sentinel) as mock_cls: + get_redis_client( + sentinel_nodes=[("sentinel-1", 26379)], + sentinel_password="sentinel-secret", + service_name="mymaster", + ) + + sentinel_call_kwargs = mock_cls.call_args[1] + assert "socket_timeout" not in sentinel_call_kwargs + assert sentinel_call_kwargs["sentinel_kwargs"] == { + "password": "sentinel-secret", + "socket_timeout": 5, + } + assert "password" not in sentinel_call_kwargs + mock_sentinel.master_for.assert_called_once_with( + "mymaster", password="redis-secret", socket_timeout=5 + ) @patch("litellm._redis.async_redis.Sentinel") @@ -416,17 +454,19 @@ def test_async_sentinel_uses_sentinel_password_and_master_password( ssl_check_hostname=True, ssl_ca_certs="/tmp/test-ca.pem", max_connections=17, + socket_timeout=5, ) mock_sentinel_cls.assert_called_once() sentinel_call_kwargs = mock_sentinel_cls.call_args[1] - assert sentinel_call_kwargs["password"] == "redis-secret" - assert sentinel_call_kwargs["username"] == "redis-user" - assert sentinel_call_kwargs["ssl"] is True - assert sentinel_call_kwargs["ssl_cert_reqs"] == "required" - assert sentinel_call_kwargs["ssl_check_hostname"] is True - assert sentinel_call_kwargs["ssl_ca_certs"] == "/tmp/test-ca.pem" - assert sentinel_call_kwargs["max_connections"] == 17 + assert "password" not in sentinel_call_kwargs + assert "username" not in sentinel_call_kwargs + assert "ssl" not in sentinel_call_kwargs + assert "ssl_cert_reqs" not in sentinel_call_kwargs + assert "ssl_check_hostname" not in sentinel_call_kwargs + assert "ssl_ca_certs" not in sentinel_call_kwargs + assert "max_connections" not in sentinel_call_kwargs + assert "socket_timeout" not in sentinel_call_kwargs assert sentinel_call_kwargs["sentinel_kwargs"] == { "password": "sentinel-secret", "username": "redis-user", @@ -435,11 +475,22 @@ def test_async_sentinel_uses_sentinel_password_and_master_password( "ssl_check_hostname": True, "ssl_ca_certs": "/tmp/test-ca.pem", "max_connections": 17, + "socket_timeout": 5, } assert "service_name" not in sentinel_call_kwargs["sentinel_kwargs"] assert "sentinel_nodes" not in sentinel_call_kwargs["sentinel_kwargs"] assert "sentinel_password" not in sentinel_call_kwargs["sentinel_kwargs"] - mock_sentinel.master_for.assert_called_once_with("mymaster") + mock_sentinel.master_for.assert_called_once_with( + "mymaster", + password="redis-secret", + username="redis-user", + ssl=True, + ssl_cert_reqs="required", + ssl_check_hostname=True, + ssl_ca_certs="/tmp/test-ca.pem", + max_connections=17, + socket_timeout=5, + ) @patch("litellm._redis.init_redis_cluster")