From cba4fa403dba0540ec3c69db1466b0419d4bb126 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Thu, 20 Aug 2026 15:42:48 -0700 Subject: [PATCH] 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. --- litellm/_redis.py | 68 ++++++++++++----------------- tests/test_litellm/test_redis.py | 73 ++++++++++++++++++++++++++++++++ 2 files changed, 99 insertions(+), 42 deletions(-) diff --git a/litellm/_redis.py b/litellm/_redis.py index 0acc01fa14f..b78ab285ac5 100644 --- a/litellm/_redis.py +++ b/litellm/_redis.py @@ -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) diff --git a/tests/test_litellm/test_redis.py b/tests/test_litellm/test_redis.py index 896ca2de399..5c02b1f786f 100644 --- a/tests/test_litellm/test_redis.py +++ b/tests/test_litellm/test_redis.py @@ -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)