fix(redis): keep Azure AD and GCP IAM auth on URL and pool clients

REDIS_URL-based async clients and every async connection pool dropped the
managed-identity credential the caller configured, so they connected
unauthenticated against an auth-enforcing Redis. The conversion from
redis_connect_func to a CredentialProvider now happens once, before any
branch, and covers the url, sentinel, cluster, and pool paths alike.

Also adds credential_provider to the cluster kwargs allowlist, which
silently filtered it out.
This commit is contained in:
mateo-berri 2026-08-20 15:42:48 -07:00
parent cb4eb82249
commit cba4fa403d
2 changed files with 99 additions and 42 deletions

View file

@ -17,6 +17,7 @@ from typing import Final
import redis
import redis.asyncio as async_redis
from redis.credentials import CredentialProvider
from litellm import get_secret, get_secret_str
from litellm._redis_credential_provider import (
@ -134,6 +135,7 @@ def _get_redis_cluster_kwargs(client=None):
"ssl_check_hostname",
"ssl_ca_certs",
"redis_connect_func", # Needed for sync clusters and IAM detection
"credential_provider",
"gcp_service_account",
"gcp_ssl_ca_certs",
"azure_redis_ad_token",
@ -574,6 +576,20 @@ 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:
"""Async redis-py never calls ``redis_connect_func``; it authenticates through a
``CredentialProvider``, which it consults per connection so the token refreshes."""
gcp_service_account: Final = getattr(redis_connect_func, "_gcp_service_account", None)
if gcp_service_account is not None:
return GCPIAMCredentialProvider(gcp_service_account)
azure_credential: Final = getattr(redis_connect_func, "_azure_credential", None)
if azure_credential is not None:
return AzureADCredentialProvider(azure_credential, username=os.environ.get("REDIS_USERNAME") or None)
return None
def get_redis_client(**env_overrides):
redis_kwargs: Final = _get_redis_client_logic(**env_overrides)
@ -601,6 +617,11 @@ def get_redis_async_client(
**env_overrides,
) -> async_redis.Redis | async_redis.RedisCluster:
redis_kwargs: Final = _get_redis_client_logic(**env_overrides)
credential_provider: Final = _async_credential_provider(redis_kwargs.pop("redis_connect_func", None))
if credential_provider is not None:
redis_kwargs["credential_provider"] = credential_provider
redis_kwargs.pop("username", None)
redis_kwargs.pop("password", None)
if "startup_nodes" in redis_kwargs:
from redis.cluster import ClusterNode
@ -611,23 +632,6 @@ def get_redis_async_client(
if arg in args:
cluster_kwargs[arg] = redis_kwargs[arg]
# Handle GCP IAM authentication for async clusters
redis_connect_func = cluster_kwargs.pop("redis_connect_func", None)
# Use a CredentialProvider so the IAM token is regenerated on every new
# connection — mirrors the sync path where redis_connect_func is invoked
# per connection. Without this, the token would expire after ~1 hour.
if redis_connect_func and hasattr(redis_connect_func, "_gcp_service_account"):
cluster_kwargs["credential_provider"] = GCPIAMCredentialProvider(redis_connect_func._gcp_service_account)
# Handle Azure AD authentication for async clusters via CredentialProvider
# so the credential's internal cache + silent refresh runs per connection
# (mirrors GCP IAM above; avoids static-token-baked-in-pool expiry).
elif redis_connect_func and hasattr(redis_connect_func, "_azure_credential"):
cluster_kwargs["credential_provider"] = AzureADCredentialProvider(
redis_connect_func._azure_credential,
username=os.environ.get("REDIS_USERNAME") or None,
)
new_startup_nodes: Final[list[ClusterNode]] = []
for item in redis_kwargs["startup_nodes"]:
@ -667,19 +671,6 @@ def get_redis_async_client(
if "sentinel_nodes" in redis_kwargs and "service_name" in redis_kwargs:
return _init_async_redis_sentinel(redis_kwargs)
# Wrap GCP / Azure AD auth in a CredentialProvider for the standard async
# Redis client. The async client doesn't support redis_connect_func, but it
# does honour credential_provider — which is called per connection, so the
# underlying SDK can refresh tokens silently before they expire.
redis_connect_func = redis_kwargs.pop("redis_connect_func", None)
if redis_connect_func and hasattr(redis_connect_func, "_azure_credential"):
redis_kwargs["credential_provider"] = AzureADCredentialProvider(
redis_connect_func._azure_credential,
username=os.environ.get("REDIS_USERNAME") or None,
)
elif redis_connect_func and hasattr(redis_connect_func, "_gcp_service_account"):
redis_kwargs["credential_provider"] = GCPIAMCredentialProvider(redis_connect_func._gcp_service_account)
_pretty_print_redis_config(redis_kwargs=redis_kwargs)
if connection_pool is not None:
@ -694,6 +685,11 @@ def get_redis_connection_pool(
**env_overrides,
) -> async_redis.BlockingConnectionPool | None:
redis_kwargs: Final = _get_redis_client_logic(**env_overrides)
credential_provider: Final = _async_credential_provider(redis_kwargs.pop("redis_connect_func", None))
if credential_provider is not None:
redis_kwargs["credential_provider"] = credential_provider
redis_kwargs.pop("username", None)
redis_kwargs.pop("password", None)
verbose_logger.debug("get_redis_connection_pool: redis_kwargs", redis_kwargs)
if "startup_nodes" in redis_kwargs:
@ -714,18 +710,6 @@ def get_redis_connection_pool(
)
return async_redis.BlockingConnectionPool.from_url(**pool_kwargs)
# Wrap GCP / Azure AD auth in a CredentialProvider so pool-managed
# connections re-fetch tokens via the SDK's internal cache + silent refresh
# rather than reusing a single token captured at pool creation.
redis_connect_func: Final = redis_kwargs.pop("redis_connect_func", None)
if redis_connect_func and hasattr(redis_connect_func, "_azure_credential"):
redis_kwargs["credential_provider"] = AzureADCredentialProvider(
redis_connect_func._azure_credential,
username=os.environ.get("REDIS_USERNAME") or None,
)
elif redis_connect_func and hasattr(redis_connect_func, "_gcp_service_account"):
redis_kwargs["credential_provider"] = GCPIAMCredentialProvider(redis_connect_func._gcp_service_account)
if redis_kwargs.pop("ssl", None):
redis_kwargs["connection_class"] = async_redis.SSLConnection
return async_redis.BlockingConnectionPool(timeout=REDIS_CONNECTION_POOL_TIMEOUT, **redis_kwargs)

View file

@ -1,4 +1,5 @@
import json
from types import SimpleNamespace
from unittest.mock import MagicMock, patch
import pytest
@ -14,6 +15,7 @@ from litellm._redis import (
)
from litellm.constants import REDIS_CLUSTER_HEALTH_CHECK_INTERVAL
from litellm._redis_credential_provider import (
AzureADCredentialProvider,
GCPIAMCredentialProvider,
_token_cache,
)
@ -910,3 +912,74 @@ def test_url_allowlist_always_carries_socket_timeouts():
allowed = _get_redis_url_kwargs()
assert "socket_timeout" in allowed
assert "socket_connect_timeout" in allowed
AZURE_AD_CONNECT_FUNC = {"_azure_credential": object()}
GCP_IAM_CONNECT_FUNC = {"_gcp_service_account": "projects/-/serviceAccounts/sa@project.iam.gserviceaccount.com"}
@pytest.mark.parametrize(
"markers, provider_cls",
[
(AZURE_AD_CONNECT_FUNC, AzureADCredentialProvider),
(GCP_IAM_CONNECT_FUNC, GCPIAMCredentialProvider),
],
ids=["azure_ad", "gcp_iam"],
)
def test_async_url_client_authenticates_through_credential_provider(markers, provider_cls):
"""A REDIS_URL config with Azure AD or GCP IAM must still reach the server with a credential.
The async client accepts redis_connect_func as a kwarg but never calls it, so the url
branch has to hand the connection a CredentialProvider or it authenticates with nothing.
"""
redis_kwargs = {
"url": "rediss://redis-host:6380",
"redis_connect_func": SimpleNamespace(**markers),
}
with patch("litellm._redis._get_redis_client_logic", return_value=redis_kwargs):
client = get_redis_async_client()
connection_kwargs = client.connection_pool.connection_kwargs
assert isinstance(connection_kwargs.get("credential_provider"), provider_cls)
assert "redis_connect_func" not in connection_kwargs
@pytest.mark.parametrize(
"markers, provider_cls",
[
(AZURE_AD_CONNECT_FUNC, AzureADCredentialProvider),
(GCP_IAM_CONNECT_FUNC, GCPIAMCredentialProvider),
],
ids=["azure_ad", "gcp_iam"],
)
def test_async_url_connection_pool_authenticates_through_credential_provider(markers, provider_cls):
"""Same for the pool-based path: every connection the pool hands out needs the provider."""
redis_kwargs = {
"url": "rediss://redis-host:6380",
"redis_connect_func": SimpleNamespace(**markers),
}
with patch("litellm._redis._get_redis_client_logic", return_value=redis_kwargs):
pool = get_redis_connection_pool()
assert isinstance(pool.connection_kwargs.get("credential_provider"), provider_cls)
assert "redis_connect_func" not in pool.connection_kwargs
def test_async_url_client_drops_username_alongside_credential_provider():
"""redis-py refuses a connection given both a username and a credential_provider, and
AzureADCredentialProvider already carries REDIS_USERNAME, so the username must be dropped.
"""
redis_kwargs = {
"url": "rediss://redis-host:6380",
"username": "redis-user",
"redis_connect_func": SimpleNamespace(**AZURE_AD_CONNECT_FUNC),
}
with patch("litellm._redis._get_redis_client_logic", return_value=redis_kwargs):
client = get_redis_async_client()
pool = client.connection_pool
assert "username" not in pool.connection_kwargs
pool.connection_class(**pool.connection_kwargs)