mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-28 01:32:17 +00:00
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:
parent
9dbfb060bd
commit
28499f4a0a
2 changed files with 61 additions and 20 deletions
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue