diff --git a/litellm/_redis.py b/litellm/_redis.py index 3e68d50cf16..3ad6e1885db 100644 --- a/litellm/_redis.py +++ b/litellm/_redis.py @@ -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: diff --git a/tests/test_litellm/test_redis.py b/tests/test_litellm/test_redis.py index a96e8541e06..5425abc06e2 100644 --- a/tests/test_litellm/test_redis.py +++ b/tests/test_litellm/test_redis.py @@ -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