diff --git a/litellm/_redis.py b/litellm/_redis.py index f12afbac297..7985c4549bc 100644 --- a/litellm/_redis.py +++ b/litellm/_redis.py @@ -77,19 +77,19 @@ 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", + } return available_args @@ -303,10 +303,24 @@ def init_redis_cluster(redis_kwargs) -> redis.RedisCluster: return redis.RedisCluster(startup_nodes=new_startup_nodes, **cluster_kwargs) # type: ignore +def _get_redis_sentinel_connection_kwargs(redis_kwargs: dict) -> dict: + connection_kwargs = {} + args = _get_redis_kwargs() + for arg in redis_kwargs: + if arg in args: + connection_kwargs[arg] = redis_kwargs[arg] + + return connection_kwargs + + def _init_redis_sentinel(redis_kwargs) -> redis.Redis: sentinel_nodes = redis_kwargs.get("sentinel_nodes") sentinel_password = redis_kwargs.get("sentinel_password") service_name = redis_kwargs.get("service_name") + connection_kwargs = _get_redis_sentinel_connection_kwargs(redis_kwargs) + connection_kwargs.setdefault("socket_timeout", REDIS_SOCKET_TIMEOUT) + sentinel_kwargs = dict(connection_kwargs) + sentinel_kwargs["password"] = sentinel_password if not sentinel_nodes or not service_name: raise ValueError( @@ -318,19 +332,22 @@ def _init_redis_sentinel(redis_kwargs) -> redis.Redis: # Set up the Sentinel client sentinel = redis.Sentinel( sentinel_nodes, - socket_timeout=REDIS_SOCKET_TIMEOUT, - password=sentinel_password, + sentinel_kwargs=sentinel_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: sentinel_nodes = redis_kwargs.get("sentinel_nodes") sentinel_password = redis_kwargs.get("sentinel_password") service_name = redis_kwargs.get("service_name") + connection_kwargs = _get_redis_sentinel_connection_kwargs(redis_kwargs) + connection_kwargs.setdefault("socket_timeout", REDIS_SOCKET_TIMEOUT) + sentinel_kwargs = dict(connection_kwargs) + sentinel_kwargs["password"] = sentinel_password if not sentinel_nodes or not service_name: raise ValueError( @@ -342,13 +359,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, - password=sentinel_password, + sentinel_kwargs=sentinel_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 3ee82699eb8..c2f67c4eeb6 100644 --- a/tests/test_litellm/test_redis.py +++ b/tests/test_litellm/test_redis.py @@ -163,7 +163,6 @@ def test_get_redis_async_client_with_connection_pool(): with patch("litellm._redis.async_redis.Redis") as mock_redis, patch( "litellm._redis._get_redis_client_logic" ) as mock_logic: - # Configure mock to return basic redis kwargs mock_logic.return_value = {"host": "localhost", "port": 6379, "db": 0} @@ -185,7 +184,6 @@ def test_get_redis_async_client_without_connection_pool(): with patch("litellm._redis.async_redis.Redis") as mock_redis, patch( "litellm._redis._get_redis_client_logic" ) as mock_logic: - # Configure mock to return basic redis kwargs mock_logic.return_value = {"host": "localhost", "port": 6379, "db": 0} @@ -356,6 +354,120 @@ def test_sync_client_prefers_cluster_over_url_via_env_var( assert len(call_kwargs["startup_nodes"]) == 1 +@patch("litellm._redis.redis.Sentinel") +def test_sync_sentinel_uses_sentinel_password_and_master_password(mock_sentinel_cls): + """Sentinel auth must be passed to the sentinel, not the Redis master client.""" + mock_sentinel = MagicMock() + mock_sentinel_cls.return_value = mock_sentinel + + get_redis_client( + sentinel_nodes=[("sentinel-1", 26379)], + sentinel_password="sentinel-secret", + service_name="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, + ) + + mock_sentinel_cls.assert_called_once() + sentinel_call_kwargs = mock_sentinel_cls.call_args[1] + 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", + "ssl": True, + "ssl_cert_reqs": "required", + "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", + 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.async_redis.Sentinel") +def test_async_sentinel_uses_sentinel_password_and_master_password( + mock_sentinel_cls, +): + """Async sentinel auth must mirror the sync sentinel password routing.""" + mock_sentinel = MagicMock() + mock_sentinel_cls.return_value = mock_sentinel + + get_redis_async_client( + sentinel_nodes=[("sentinel-1", 26379)], + sentinel_password="sentinel-secret", + service_name="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, + ) + + mock_sentinel_cls.assert_called_once() + sentinel_call_kwargs = mock_sentinel_cls.call_args[1] + 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", + "ssl": True, + "ssl_cert_reqs": "required", + "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", + 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") def test_sync_client_preserves_password_for_cluster_when_url_also_set( mock_init_cluster, monkeypatch