mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
fix: Fix Redis Sentinel client handling to solve authentication error with password protected sentinel (#25625)
* fix Redis Sentinel authentication handling * test: cover Redis Sentinel auth routing * refactor: align Redis Sentinel kwargs threading * fix: avoid duplicate Redis Sentinel socket timeouts * Address review comments
This commit is contained in:
parent
850fe595ac
commit
6518094c1d
2 changed files with 149 additions and 21 deletions
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue