mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
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:
parent
cb4eb82249
commit
cba4fa403d
2 changed files with 99 additions and 42 deletions
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue