fix(redis): authenticate sync clusters with IAM credential providers

Signed-off-by: Silu Panda <31051721+SiluPanda@users.noreply.github.com>
This commit is contained in:
Silu Panda 2026-09-07 23:09:14 -07:00
parent 9dbfb060bd
commit 28499f4a0a
2 changed files with 61 additions and 20 deletions

View file

@ -624,11 +624,12 @@ def init_redis_cluster(redis_kwargs) -> redis.RedisCluster:
verbose_logger.debug("init_redis_cluster: startup nodes are being initialized.")
from redis.cluster import ClusterNode
auth_kwargs: Final = _credential_provider_auth_kwargs(redis_kwargs)
args: Final = _get_redis_cluster_kwargs()
cluster_kwargs: Final = {}
for arg in redis_kwargs:
for arg in auth_kwargs:
if arg in args:
cluster_kwargs[arg] = redis_kwargs[arg]
cluster_kwargs[arg] = auth_kwargs[arg]
new_startup_nodes: Final[list[ClusterNode]] = []
@ -706,13 +707,13 @@ def _init_async_redis_sentinel(redis_kwargs) -> async_redis.Redis:
return sentinel.master_for(service_name, **connection_kwargs)
def _async_credential_provider(redis_connect_func: object | None) -> CredentialProvider | None:
"""The Azure AD and GCP IAM connect funcs run their AUTH exchange with the blocking client
API, so on an async connection their ``send_command``/``read_response`` calls return
coroutines nobody awaits and every connect fails. Async paths authenticate through a
``CredentialProvider`` instead, which redis-py consults per connection so the token stays
fresh. Any other ``redis_connect_func`` is left where it is, since redis-py awaits it
itself when it is a coroutine function."""
def _credential_provider_from_connect_func(redis_connect_func: object | None) -> CredentialProvider | None:
"""Translate IAM callbacks for paths that need credentials during the standard handshake.
Async connections cannot run blocking AUTH callbacks. Sync clusters authenticate before
invoking the callback, so they also need the provider during the initial handshake.
redis-py consults the provider for each connection, keeping token refresh intact.
"""
gcp_service_account: Final = getattr(redis_connect_func, "_gcp_service_account", None)
if gcp_service_account is not None:
return GCPIAMCredentialProvider(gcp_service_account)
@ -724,14 +725,13 @@ def _async_credential_provider(redis_connect_func: object | None) -> CredentialP
return None
def _async_auth_kwargs(redis_kwargs: dict) -> dict:
"""Swaps a connect func an async path cannot run for the equivalent credential provider,
which supersedes any static username or password redis-py would otherwise reject it with."""
def _credential_provider_auth_kwargs(redis_kwargs: dict) -> dict:
"""Use a credential provider instead of an IAM callback and conflicting static credentials."""
explicit_provider: Final = redis_kwargs.get("credential_provider")
credential_provider: Final = (
explicit_provider
if explicit_provider is not None
else _async_credential_provider(redis_kwargs.get("redis_connect_func"))
else _credential_provider_from_connect_func(redis_kwargs.get("redis_connect_func"))
)
if credential_provider is None:
return redis_kwargs
@ -769,7 +769,7 @@ def get_redis_async_client(
connection_pool: async_redis.BlockingConnectionPool | None = None,
**env_overrides,
) -> async_redis.Redis | async_redis.RedisCluster:
redis_kwargs: Final = _async_auth_kwargs(_get_redis_client_logic(**env_overrides))
redis_kwargs: Final = _credential_provider_auth_kwargs(_get_redis_client_logic(**env_overrides))
if "startup_nodes" in redis_kwargs:
from redis.cluster import ClusterNode
@ -841,7 +841,7 @@ def get_redis_async_client(
def get_redis_connection_pool(
**env_overrides,
) -> async_redis.BlockingConnectionPool | None:
redis_kwargs: Final = _async_auth_kwargs(_get_redis_client_logic(**env_overrides))
redis_kwargs: Final = _credential_provider_auth_kwargs(_get_redis_client_logic(**env_overrides))
verbose_logger.debug("get_redis_connection_pool: redis_kwargs", redis_kwargs)
if "startup_nodes" in redis_kwargs:

View file

@ -10,7 +10,7 @@ from redis.credentials import CredentialProvider
import litellm
from litellm._redis import (
_async_auth_kwargs,
_credential_provider_auth_kwargs,
_get_redis_client_logic,
_get_redis_cluster_kwargs,
_get_redis_env_kwarg_mapping,
@ -243,6 +243,47 @@ def test_sync_cluster_preserves_credential_provider_identity(clean_redis_environ
assert [(node.host, node.port) for node in cluster_kwargs["startup_nodes"]] == [("cluster-node", 6379)]
def test_sync_cluster_authenticates_with_azure_credentials(clean_redis_environment, monkeypatch):
monkeypatch.setenv("REDIS_USERNAME", "identity-object-id")
credential = MagicMock()
credential.get_token.return_value = SimpleNamespace(token="azure-access-token")
with (
patch("azure.identity.DefaultAzureCredential", return_value=credential),
patch("redis.RedisCluster", autospec=True) as cluster,
):
get_redis_client(
startup_nodes=[{"host": "cluster-node", "port": 6379}],
azure_redis_ad_token=True,
password="stale-password",
)
kwargs = cluster.call_args.kwargs
provider = kwargs.get("credential_provider")
assert isinstance(provider, AzureADCredentialProvider)
assert provider.get_credentials() == ("identity-object-id", "azure-access-token")
assert "username" not in kwargs
assert "password" not in kwargs
assert "redis_connect_func" not in kwargs
credential.get_token.assert_called_once_with("https://redis.azure.com/.default")
def test_sync_cluster_authenticates_with_gcp_credentials(clean_redis_environment):
with patch("redis.RedisCluster", autospec=True) as cluster:
get_redis_client(
startup_nodes=[{"host": "cluster-node", "port": 6379}],
redis_connect_func=_gcp_marker_callback(),
username="stale-user",
password="stale-password",
)
kwargs = cluster.call_args.kwargs
assert isinstance(kwargs.get("credential_provider"), GCPIAMCredentialProvider)
assert "username" not in kwargs
assert "password" not in kwargs
assert "redis_connect_func" not in kwargs
def test_async_cluster_preserves_credential_provider_identity(clean_redis_environment):
provider = _StubCredentialProvider()
startup_nodes = [{"host": "cluster-node", "port": 6379}]
@ -319,10 +360,10 @@ def test_provider_free_url_is_left_untouched(clean_redis_environment):
assert redis_kwargs["url"] == url
def test_async_auth_kwargs_supersedes_credentials_an_explicit_provider_replaces():
def test_credential_provider_auth_kwargs_supersedes_credentials_an_explicit_provider_replaces():
provider = _StubCredentialProvider()
auth_kwargs = _async_auth_kwargs(
auth_kwargs = _credential_provider_auth_kwargs(
{
"host": "redis-host",
"port": 6379,
@ -341,10 +382,10 @@ def test_async_auth_kwargs_supersedes_credentials_an_explicit_provider_replaces(
assert "password" not in auth_kwargs
def test_async_auth_kwargs_leaves_provider_free_kwargs_alone():
def test_credential_provider_auth_kwargs_leaves_provider_free_kwargs_alone():
redis_kwargs = {"host": "redis-host", "port": 6379, "username": "url-user", "password": "url-pass"}
assert _async_auth_kwargs(redis_kwargs) == redis_kwargs
assert _credential_provider_auth_kwargs(redis_kwargs) == redis_kwargs
@pytest.mark.asyncio