From d2ce2d064ff59e9be5028a601487e3b92d74a964 Mon Sep 17 00:00:00 2001 From: eugene-yao-zocdoc Date: Wed, 26 Aug 2026 15:45:25 -0400 Subject: [PATCH 1/8] feat(redis): add ElastiCache IAM authentication --- litellm/_redis.py | 57 ++++- litellm/_redis_credential_provider.py | 85 +++++++- litellm/proxy/_types.py | 4 + .../cache_settings_endpoints.py | 36 ++++ .../coordination_redis_endpoints.py | 29 +++ tests/test_litellm/test_redis.py | 196 +++++++++++++++++- .../test_redis_credential_provider.py | 122 +++++++++++ 7 files changed, 520 insertions(+), 9 deletions(-) create mode 100644 tests/test_litellm/test_redis_credential_provider.py diff --git a/litellm/_redis.py b/litellm/_redis.py index 9381357931e..e281bd080a9 100644 --- a/litellm/_redis.py +++ b/litellm/_redis.py @@ -23,6 +23,7 @@ from redis.credentials import CredentialProvider from litellm import get_secret, get_secret_str from litellm._redis_credential_provider import ( AzureADCredentialProvider, + ElastiCacheIAMCredentialProvider, GCPIAMCredentialProvider, _generate_gcp_iam_access_token, ) @@ -58,6 +59,10 @@ def _get_redis_kwargs(): "azure_client_id", "azure_tenant_id", "azure_client_secret", + "aws_iam_auth", + "aws_iam_user_name", + "aws_iam_cache_name", + "aws_iam_region", } available_args: Final = {x for x in arg_spec.args if x not in exclude_args} | include_args @@ -172,6 +177,32 @@ def _redis_kwargs_from_environment(): return return_dict +def _is_true(value: object | None) -> bool: + return value is True or (isinstance(value, str) and value.lower() == "true") + + +def _build_elasticache_iam_provider(redis_kwargs: dict) -> ElastiCacheIAMCredentialProvider | None: + if not _is_true(redis_kwargs.get("aws_iam_auth")): + return None + + required_settings: Final = { + "aws_iam_user_name": redis_kwargs.get("aws_iam_user_name"), + "aws_iam_cache_name": redis_kwargs.get("aws_iam_cache_name"), + "aws_iam_region": redis_kwargs.get("aws_iam_region") + or get_secret_str("AWS_REGION") + or get_secret_str("AWS_DEFAULT_REGION"), + } + missing_settings: Final = tuple(name for name, value in required_settings.items() if not value) + if missing_settings: + raise ValueError("AWS ElastiCache IAM Redis authentication requires: " + ", ".join(missing_settings)) + + return ElastiCacheIAMCredentialProvider( + user_name=str(required_settings["aws_iam_user_name"]), + cache_name=str(required_settings["aws_iam_cache_name"]), + region=str(required_settings["aws_iam_region"]), + ) + + def create_gcp_iam_redis_connect_func( service_account: str, ssl_ca_certs: str | None = None, @@ -442,6 +473,7 @@ def _get_redis_client_logic(**env_overrides): _azure_redis_ad_token: Final = redis_kwargs.get("azure_redis_ad_token") or get_secret("REDIS_AZURE_AD_TOKEN") _azure_ad_enabled: Final = _azure_redis_ad_token is not None and str(_azure_redis_ad_token).lower() == "true" + _aws_iam_enabled: Final = _is_true(redis_kwargs.get("aws_iam_auth")) if _azure_ad_enabled and _gcp_service_account is not None: verbose_logger.warning( @@ -462,13 +494,24 @@ def _get_redis_client_logic(**env_overrides): azure_tenant_id=_azure_tenant_id, azure_client_secret=_azure_client_secret, ) - # Marker for async paths to detect Azure AD auth. The live credential - # object is attached separately as `_azure_credential` by - # `create_azure_ad_redis_connect_func`; the raw client_id/tenant_id/secret - # are intentionally NOT exposed on the function to avoid leaking - # credentials via inspection or logging. redis_kwargs["redis_connect_func"]._azure_redis_ad_token = True + if _aws_iam_enabled and _gcp_service_account is not None: + verbose_logger.warning( + "Both GCP IAM (gcp_service_account) and AWS ElastiCache IAM (aws_iam_auth) are configured " + "for Redis. Using GCP IAM. Remove one to avoid misconfiguration." + ) + elif _aws_iam_enabled and _azure_ad_enabled: + verbose_logger.warning( + "Both Azure AD (azure_redis_ad_token) and AWS ElastiCache IAM (aws_iam_auth) are configured " + "for Redis. Using Azure AD. Remove one to avoid misconfiguration." + ) + elif _aws_iam_enabled: + aws_provider: Final = _build_elasticache_iam_provider(redis_kwargs) + if aws_provider is not None: + verbose_logger.debug("Setting up AWS ElastiCache IAM authentication for Redis.") + redis_kwargs["credential_provider"] = aws_provider + redis_kwargs.pop("gcp_service_account", None) redis_kwargs.pop("gcp_ssl_ca_certs", None) @@ -477,6 +520,10 @@ def _get_redis_client_logic(**env_overrides): redis_kwargs.pop("azure_client_id", None) redis_kwargs.pop("azure_tenant_id", None) redis_kwargs.pop("azure_client_secret", None) + redis_kwargs.pop("aws_iam_auth", None) + redis_kwargs.pop("aws_iam_user_name", None) + redis_kwargs.pop("aws_iam_cache_name", None) + redis_kwargs.pop("aws_iam_region", None) if redis_kwargs.get("credential_provider") is not None: redis_kwargs.pop("redis_connect_func", None) diff --git a/litellm/_redis_credential_provider.py b/litellm/_redis_credential_provider.py index 98fa62629a8..800473bfa54 100644 --- a/litellm/_redis_credential_provider.py +++ b/litellm/_redis_credential_provider.py @@ -1,7 +1,8 @@ import asyncio import threading import time -from typing import Any, Final +from typing import Any, Final, Protocol +from urllib.parse import quote from redis.credentials import CredentialProvider @@ -104,6 +105,88 @@ class GCPIAMCredentialProvider(CredentialProvider): return (token,) +_ELASTICACHE_SERVICE_NAME: Final = "elasticache" +_ELASTICACHE_TOKEN_TTL_SECONDS: Final = 900 + + +class _FrozenBotocoreCredentials(Protocol): + access_key: str + secret_key: str + token: str | None + + +class _BotocoreCredentials(Protocol): + def get_frozen_credentials(self) -> _FrozenBotocoreCredentials: ... + + +class _BotocoreCredentialsResolver(Protocol): + def __call__(self) -> _BotocoreCredentials | None: ... + + +class ElastiCacheIAMCredentialProvider(CredentialProvider): + def __init__( + self, + user_name: str, + cache_name: str, + region: str, + credentials_resolver: _BotocoreCredentialsResolver | None = None, + token_lifetime_seconds: int = _ELASTICACHE_TOKEN_TTL_SECONDS, + ) -> None: + self._user_name = user_name + self._cache_name = cache_name + self._region = region + self._credentials_resolver = credentials_resolver or self._resolve_credentials + self._credentials: _BotocoreCredentials | None = None + self._token_lifetime_seconds = token_lifetime_seconds + + @staticmethod + def _resolve_credentials() -> Any: + try: + import botocore.session + except ImportError as e: + raise ImportError( + "botocore is required for ElastiCache IAM Redis authentication. Install it with: pip install boto3" + ) from e + + return botocore.session.get_session().get_credentials() + + def _get_credentials(self) -> tuple[str, str]: + credentials: Final = self._credentials if self._credentials is not None else self._credentials_resolver() + if credentials is None: + raise RuntimeError("Unable to resolve AWS credentials for ElastiCache IAM Redis authentication") + self._credentials = credentials + + frozen_credentials: Final = credentials.get_frozen_credentials() + if frozen_credentials is None: + raise RuntimeError("Unable to resolve AWS credentials for ElastiCache IAM Redis authentication") + + try: + from botocore.auth import SigV4QueryAuth + from botocore.awsrequest import AWSRequest + except ImportError as e: + raise ImportError( + "botocore is required for ElastiCache IAM Redis authentication. Install it with: pip install boto3" + ) from e + + request: Final = AWSRequest( + method="GET", + url=(f"https://{self._cache_name}/?Action=connect&User={quote(self._user_name, safe='')}"), + ) + SigV4QueryAuth( + frozen_credentials, + _ELASTICACHE_SERVICE_NAME, + self._region, + expires=self._token_lifetime_seconds, + ).add_auth(request) + return self._user_name, request.url.removeprefix("https://") + + def get_credentials(self) -> tuple[str, str]: + return self._get_credentials() + + async def get_credentials_async(self) -> tuple[str, str]: + return await asyncio.to_thread(self._get_credentials) + + class AzureADCredentialProvider(CredentialProvider): """ redis.credentials.CredentialProvider implementation that supplies Azure AD diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 61acac3ff74..740a63829c8 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -2326,6 +2326,10 @@ class CoordinationRedisParams(LiteLLMPydanticObjectBase): ) sentinel_password: str | None = Field(None, description="password for the sentinel nodes") service_name: str | None = Field(None, description="sentinel service name") + aws_iam_auth: bool | str | None = Field(None, description="enable AWS ElastiCache IAM authentication") + aws_iam_user_name: str | None = Field(None, description="AWS ElastiCache IAM user name") + aws_iam_cache_name: str | None = Field(None, description="AWS ElastiCache cache name") + aws_iam_region: str | None = Field(None, description="AWS region for ElastiCache IAM authentication") def has_connection_target(self) -> bool: return any(value is not None for value in (self.host, self.url, self.startup_nodes, self.sentinel_nodes)) diff --git a/litellm/types/management_endpoints/cache_settings_endpoints.py b/litellm/types/management_endpoints/cache_settings_endpoints.py index 2e6d0be9071..04600117ce8 100644 --- a/litellm/types/management_endpoints/cache_settings_endpoints.py +++ b/litellm/types/management_endpoints/cache_settings_endpoints.py @@ -224,4 +224,40 @@ CACHE_SETTINGS_FIELDS: Final[list[CacheSettingsField]] = [ ui_field_name="SSL Check Hostname", redis_type=None, ), + CacheSettingsField( + field_name="aws_iam_auth", + field_type="Boolean", + field_value=None, + field_description="Enable AWS ElastiCache IAM authentication", + field_default=False, + ui_field_name="AWS IAM Authentication", + redis_type=None, + ), + CacheSettingsField( + field_name="aws_iam_user_name", + field_type="String", + field_value=None, + field_description="AWS ElastiCache IAM user name", + field_default=None, + ui_field_name="AWS IAM User Name", + redis_type=None, + ), + CacheSettingsField( + field_name="aws_iam_cache_name", + field_type="String", + field_value=None, + field_description="AWS ElastiCache cache name", + field_default=None, + ui_field_name="AWS IAM Cache Name", + redis_type=None, + ), + CacheSettingsField( + field_name="aws_iam_region", + field_type="String", + field_value=None, + field_description="AWS region for ElastiCache IAM authentication", + field_default=None, + ui_field_name="AWS IAM Region", + redis_type=None, + ), ] diff --git a/litellm/types/management_endpoints/coordination_redis_endpoints.py b/litellm/types/management_endpoints/coordination_redis_endpoints.py index 30033346ed7..86b32752cbb 100644 --- a/litellm/types/management_endpoints/coordination_redis_endpoints.py +++ b/litellm/types/management_endpoints/coordination_redis_endpoints.py @@ -102,4 +102,33 @@ COORDINATION_REDIS_SETTINGS_FIELDS: Final[list[CoordinationRedisSettingsField]] ui_field_name="Service Name", section="sentinel", ), + CoordinationRedisSettingsField( + field_name="aws_iam_auth", + field_type="Boolean", + field_description="Enable AWS ElastiCache IAM authentication", + field_default=False, + ui_field_name="AWS IAM Authentication", + section="connection", + ), + CoordinationRedisSettingsField( + field_name="aws_iam_user_name", + field_type="String", + field_description="AWS ElastiCache IAM user name", + ui_field_name="AWS IAM User Name", + section="connection", + ), + CoordinationRedisSettingsField( + field_name="aws_iam_cache_name", + field_type="String", + field_description="AWS ElastiCache cache name", + ui_field_name="AWS IAM Cache Name", + section="connection", + ), + CoordinationRedisSettingsField( + field_name="aws_iam_region", + field_type="String", + field_description="AWS region for ElastiCache IAM authentication", + ui_field_name="AWS IAM Region", + section="connection", + ), ] diff --git a/tests/test_litellm/test_redis.py b/tests/test_litellm/test_redis.py index 826beb74a27..0c0c257d396 100644 --- a/tests/test_litellm/test_redis.py +++ b/tests/test_litellm/test_redis.py @@ -23,6 +23,7 @@ from litellm._redis import ( ) from litellm._redis_credential_provider import ( AzureADCredentialProvider, + ElastiCacheIAMCredentialProvider, GCPIAMCredentialProvider, _token_cache, ) @@ -77,6 +78,8 @@ def clean_redis_environment(monkeypatch): "REDIS_URL", "REDIS_CLUSTER_NODES", "REDIS_SENTINEL_NODES", + "AWS_REGION", + "AWS_DEFAULT_REGION", *_get_redis_env_kwarg_mapping(), ): monkeypatch.delenv(var, raising=False) @@ -109,6 +112,17 @@ def test_credential_provider_is_not_environment_derived(): assert "credential_provider" not in mapping.values() +def test_aws_iam_settings_are_environment_derived(): + allowed = _get_redis_kwargs() + mapping = _get_redis_env_kwarg_mapping() + + assert {"aws_iam_auth", "aws_iam_user_name", "aws_iam_cache_name", "aws_iam_region"} <= allowed + assert mapping["REDIS_AWS_IAM_AUTH"] == "aws_iam_auth" + assert mapping["REDIS_AWS_IAM_USER_NAME"] == "aws_iam_user_name" + assert mapping["REDIS_AWS_IAM_CACHE_NAME"] == "aws_iam_cache_name" + assert mapping["REDIS_AWS_IAM_REGION"] == "aws_iam_region" + + def test_sync_direct_preserves_credential_provider_identity(clean_redis_environment): provider = _StubCredentialProvider() @@ -299,6 +313,180 @@ def test_gcp_kwargs_never_survive_client_logic(clean_redis_environment, override assert "gcp_ssl_ca_certs" not in redis_kwargs +def test_aws_iam_environment_settings_install_provider(clean_redis_environment, monkeypatch): + monkeypatch.setenv("REDIS_AWS_IAM_AUTH", "true") + monkeypatch.setenv("REDIS_AWS_IAM_USER_NAME", "iam-user") + monkeypatch.setenv("REDIS_AWS_IAM_CACHE_NAME", "cache.example.com") + monkeypatch.setenv("REDIS_AWS_IAM_REGION", "us-east-1") + + redis_kwargs = _get_redis_client_logic(host="cache.example.com", port=6379) + + assert isinstance(redis_kwargs["credential_provider"], ElastiCacheIAMCredentialProvider) + assert not {"aws_iam_auth", "aws_iam_user_name", "aws_iam_cache_name", "aws_iam_region"} & redis_kwargs.keys() + + +def test_aws_iam_settings_are_removed_for_url_and_static_credentials(clean_redis_environment): + redis_kwargs = _get_redis_client_logic( + url="rediss://url-user:url-pass@cache.example.com:6380", + aws_iam_auth=True, + aws_iam_user_name="iam-user", + aws_iam_cache_name="cache.example.com", + aws_iam_region="us-east-1", + username="static-user", + password="static-password", + ) + + assert isinstance(redis_kwargs["credential_provider"], ElastiCacheIAMCredentialProvider) + assert redis_kwargs["url"] == "rediss://cache.example.com:6380" + assert "username" not in redis_kwargs + assert "password" not in redis_kwargs + assert not {"aws_iam_auth", "aws_iam_user_name", "aws_iam_cache_name", "aws_iam_region"} & redis_kwargs.keys() + + +@pytest.mark.parametrize("missing", ["aws_iam_user_name", "aws_iam_cache_name", "aws_iam_region"]) +def test_aws_iam_missing_setting_fails_closed(clean_redis_environment, missing): + settings = { + "aws_iam_auth": True, + "aws_iam_user_name": "iam-user", + "aws_iam_cache_name": "cache.example.com", + "aws_iam_region": "us-east-1", + } + settings[missing] = None + + with pytest.raises(ValueError, match=missing): + _get_redis_client_logic(host="cache.example.com", port=6379, **settings) + + +@pytest.mark.parametrize("region_var", ["AWS_REGION", "AWS_DEFAULT_REGION"]) +def test_aws_iam_region_falls_back_to_environment(clean_redis_environment, monkeypatch, region_var): + monkeypatch.setenv(region_var, "sa-east-1") + + redis_kwargs = _get_redis_client_logic( + host="cache.example.com", + port=6379, + aws_iam_auth=True, + aws_iam_user_name="iam-user", + aws_iam_cache_name="cache.example.com", + ) + + provider = redis_kwargs["credential_provider"] + assert isinstance(provider, ElastiCacheIAMCredentialProvider) + assert provider._region == "sa-east-1" + + +def test_aws_iam_region_prefers_aws_region_over_default_region(clean_redis_environment, monkeypatch): + monkeypatch.setenv("AWS_REGION", "sa-east-1") + monkeypatch.setenv("AWS_DEFAULT_REGION", "eu-west-1") + + redis_kwargs = _get_redis_client_logic( + host="cache.example.com", + port=6379, + aws_iam_auth=True, + aws_iam_user_name="iam-user", + aws_iam_cache_name="cache.example.com", + ) + + assert redis_kwargs["credential_provider"]._region == "sa-east-1" + + +def test_aws_iam_region_prefers_explicit_over_environment(clean_redis_environment, monkeypatch): + monkeypatch.setenv("AWS_REGION", "sa-east-1") + + redis_kwargs = _get_redis_client_logic( + host="cache.example.com", + port=6379, + aws_iam_auth=True, + aws_iam_user_name="iam-user", + aws_iam_cache_name="cache.example.com", + aws_iam_region="explicit-region", + ) + + assert redis_kwargs["credential_provider"]._region == "explicit-region" + + +def test_aws_iam_settings_map_to_distinct_provider_fields(clean_redis_environment): + redis_kwargs = _get_redis_client_logic( + host="cache.example.com", + port=6379, + aws_iam_auth=True, + aws_iam_user_name="iam-user-value", + aws_iam_cache_name="iam-cache-value", + aws_iam_region="iam-region-value", + ) + + provider = redis_kwargs["credential_provider"] + assert isinstance(provider, ElastiCacheIAMCredentialProvider) + assert provider._user_name == "iam-user-value" + assert provider._cache_name == "iam-cache-value" + assert provider._region == "iam-region-value" + + +@pytest.mark.parametrize("aws_iam_auth", [False, "false"]) +def test_aws_iam_auth_disabled_does_not_install_provider(clean_redis_environment, aws_iam_auth): + redis_kwargs = _get_redis_client_logic( + host="cache.example.com", + port=6379, + aws_iam_auth=aws_iam_auth, + aws_iam_user_name="iam-user", + aws_iam_cache_name="cache.example.com", + aws_iam_region="us-east-1", + ) + + assert "credential_provider" not in redis_kwargs + assert not {"aws_iam_auth", "aws_iam_user_name", "aws_iam_cache_name", "aws_iam_region"} & redis_kwargs.keys() + + +def test_explicit_provider_wins_over_aws_iam(clean_redis_environment): + provider = _StubCredentialProvider() + + redis_kwargs = _get_redis_client_logic( + host="cache.example.com", + port=6379, + credential_provider=provider, + aws_iam_auth=True, + aws_iam_user_name="iam-user", + aws_iam_cache_name="cache.example.com", + aws_iam_region="us-east-1", + ) + + assert redis_kwargs["credential_provider"] is provider + assert "aws_iam_auth" not in redis_kwargs + + +def test_gcp_wins_over_aws_iam(clean_redis_environment): + with patch("litellm._redis.create_gcp_iam_redis_connect_func") as mock_gcp: + mock_gcp.return_value = _gcp_marker_callback() + redis_kwargs = _get_redis_client_logic( + host="cache.example.com", + port=6379, + gcp_service_account="sa@example.com", + aws_iam_auth=True, + aws_iam_user_name="iam-user", + aws_iam_cache_name="cache.example.com", + aws_iam_region="us-east-1", + ) + + assert "credential_provider" not in redis_kwargs + assert redis_kwargs["redis_connect_func"] is mock_gcp.return_value + + +def test_azure_wins_over_aws_iam(clean_redis_environment): + with patch("litellm._redis.create_azure_ad_redis_connect_func") as mock_azure: + mock_azure.return_value = MagicMock() + redis_kwargs = _get_redis_client_logic( + host="cache.example.com", + port=6379, + azure_redis_ad_token="true", + aws_iam_auth=True, + aws_iam_user_name="iam-user", + aws_iam_cache_name="cache.example.com", + aws_iam_region="us-east-1", + ) + + assert "credential_provider" not in redis_kwargs + assert redis_kwargs["redis_connect_func"] is mock_azure.return_value + + def test_provider_keeps_the_rest_of_the_url_intact(clean_redis_environment): provider = _StubCredentialProvider() @@ -1463,9 +1651,11 @@ def test_async_sentinel_keeps_the_credential_provider_off_the_monitors(markers, "redis_connect_func": SimpleNamespace(**markers), } - with patch("litellm._redis.async_redis.Sentinel") as mock_sentinel_cls: - with patch("litellm._redis._get_redis_client_logic", return_value=redis_kwargs): - get_redis_async_client() + with ( + patch("litellm._redis.async_redis.Sentinel") as mock_sentinel_cls, + patch("litellm._redis._get_redis_client_logic", return_value=redis_kwargs), + ): + get_redis_async_client() sentinel_kwargs = mock_sentinel_cls.call_args[1]["sentinel_kwargs"] assert sentinel_kwargs["password"] == sentinel_password diff --git a/tests/test_litellm/test_redis_credential_provider.py b/tests/test_litellm/test_redis_credential_provider.py new file mode 100644 index 00000000000..6299d4d42e0 --- /dev/null +++ b/tests/test_litellm/test_redis_credential_provider.py @@ -0,0 +1,122 @@ +import asyncio +from types import SimpleNamespace +from urllib.parse import parse_qs, urlsplit + +import pytest + +from litellm._redis_credential_provider import ( + ElastiCacheIAMCredentialProvider, + _BotocoreCredentials, +) + + +class _FakeCredentials: + def __init__(self, access_key: str) -> None: + self.access_key = access_key + self.secret_key = "synthetic-secret" + self.token = "synthetic-session-token" + + def get_frozen_credentials(self): + return self + + +class _FakeResolver: + def __init__(self, credentials: _BotocoreCredentials | None) -> None: + self.credentials = credentials + self.calls = 0 + + def __call__(self): + self.calls += 1 + return self.credentials + + +class _RotatingFakeCredentials: + def __init__(self) -> None: + self.calls = 0 + + def __bool__(self) -> bool: + return False + + def get_frozen_credentials(self): + self.calls += 1 + return SimpleNamespace( + access_key=f"AKIA-SYNTHETIC-{self.calls}", + secret_key="synthetic-secret", + token="synthetic-session-token", + ) + + +def test_elasticache_provider_signs_expected_query(): + resolver = _FakeResolver(_FakeCredentials("AKIA-SYNTHETIC")) + provider = ElastiCacheIAMCredentialProvider( + user_name="iam-user", + cache_name="cache.example.com", + region="us-east-1", + credentials_resolver=resolver, + ) + + user_name, token = provider.get_credentials() + parsed = urlsplit("https://" + token) + query = parse_qs(parsed.query) + + assert user_name == "iam-user" + assert parsed.netloc == "cache.example.com" + assert query["Action"] == ["connect"] + assert query["User"] == ["iam-user"] + assert query["X-Amz-Expires"] == ["900"] + assert "elasticache" in query["X-Amz-Credential"][0] + assert query["X-Amz-Credential"][0].split("/")[2] == "us-east-1" + assert not token.startswith("https://") + + +def test_elasticache_provider_resolves_credentials_once_but_refreshes_signature(): + rotating_credentials = _RotatingFakeCredentials() + resolver = _FakeResolver(rotating_credentials) + provider = ElastiCacheIAMCredentialProvider( + user_name="iam-user", + cache_name="cache.example.com", + region="us-east-1", + credentials_resolver=resolver, + ) + + first = provider.get_credentials() + second = provider.get_credentials() + async_result = asyncio.run(provider.get_credentials_async()) + + assert first[0] == second[0] == async_result[0] == "iam-user" + assert first[1] != second[1] + assert async_result[1] != second[1] + assert resolver.calls == 1 + assert rotating_credentials.calls == 3 + + +def test_elasticache_provider_reports_missing_credentials(): + provider = ElastiCacheIAMCredentialProvider( + user_name="iam-user", + cache_name="cache.example.com", + region="us-east-1", + credentials_resolver=_FakeResolver(None), + ) + + with pytest.raises(RuntimeError, match="Unable to resolve AWS credentials"): + provider.get_credentials() + + +def test_elasticache_provider_recovers_after_a_failed_resolution(): + resolver = _FakeResolver(None) + provider = ElastiCacheIAMCredentialProvider( + user_name="iam-user", + cache_name="cache.example.com", + region="us-east-1", + credentials_resolver=resolver, + ) + + with pytest.raises(RuntimeError, match="Unable to resolve AWS credentials"): + provider.get_credentials() + + resolver.credentials = _FakeCredentials("AKIA-SYNTHETIC") + user_name, token = provider.get_credentials() + + assert user_name == "iam-user" + assert token + assert resolver.calls == 2 From 578524b841a8a9275173a559541087d22523c08f Mon Sep 17 00:00:00 2001 From: eugene-yao-zocdoc Date: Wed, 26 Aug 2026 15:58:02 -0400 Subject: [PATCH 2/8] test(redis): cover AWS IAM provider install on async cluster nodes No existing test combined aws_iam_auth with startup_nodes, the shape ElastiCache/Valkey Serverless deployments actually use. --- tests/test_litellm/test_redis.py | 14 ++++++++++++++ 1 file changed, 14 insertions(+) diff --git a/tests/test_litellm/test_redis.py b/tests/test_litellm/test_redis.py index 0c0c257d396..3dffaf7e72a 100644 --- a/tests/test_litellm/test_redis.py +++ b/tests/test_litellm/test_redis.py @@ -487,6 +487,20 @@ def test_azure_wins_over_aws_iam(clean_redis_environment): assert redis_kwargs["redis_connect_func"] is mock_azure.return_value +def test_async_cluster_installs_aws_iam_provider(clean_redis_environment): + startup_nodes = [{"host": "cluster-node", "port": 6379}] + + client = get_redis_async_client( + startup_nodes=startup_nodes, + aws_iam_auth=True, + aws_iam_user_name="iam-user", + aws_iam_cache_name="cache.example.com", + aws_iam_region="us-east-1", + ) + + assert isinstance(client.connection_kwargs["credential_provider"], ElastiCacheIAMCredentialProvider) + + def test_provider_keeps_the_rest_of_the_url_intact(clean_redis_environment): provider = _StubCredentialProvider() From a7f007c118d63c54b9c8aeaab6165557d2ebbbc3 Mon Sep 17 00:00:00 2001 From: eugene-yao-zocdoc Date: Wed, 26 Aug 2026 19:19:01 -0400 Subject: [PATCH 3/8] fix(redis): type ElastiCache IAM configuration Generated with AI Co-Authored-By: Claude Code --- litellm/_redis.py | 41 ++++++++++--------- litellm/_redis_credential_provider.py | 40 +++++++++--------- ruff-strict-budget.json | 2 +- .../test_redis_credential_provider.py | 25 +++++------ ui/litellm-dashboard/src/lib/http/schema.d.ts | 20 +++++++++ 5 files changed, 75 insertions(+), 53 deletions(-) diff --git a/litellm/_redis.py b/litellm/_redis.py index e281bd080a9..289bc883ec7 100644 --- a/litellm/_redis.py +++ b/litellm/_redis.py @@ -181,25 +181,24 @@ def _is_true(value: object | None) -> bool: return value is True or (isinstance(value, str) and value.lower() == "true") -def _build_elasticache_iam_provider(redis_kwargs: dict) -> ElastiCacheIAMCredentialProvider | None: - if not _is_true(redis_kwargs.get("aws_iam_auth")): - return None - - required_settings: Final = { - "aws_iam_user_name": redis_kwargs.get("aws_iam_user_name"), - "aws_iam_cache_name": redis_kwargs.get("aws_iam_cache_name"), - "aws_iam_region": redis_kwargs.get("aws_iam_region") - or get_secret_str("AWS_REGION") - or get_secret_str("AWS_DEFAULT_REGION"), - } - missing_settings: Final = tuple(name for name, value in required_settings.items() if not value) +def _build_elasticache_iam_provider( + user_name: object | None, + cache_name: object | None, + region: object | None, +) -> ElastiCacheIAMCredentialProvider: + required_settings: Final = ( + ("aws_iam_user_name", user_name), + ("aws_iam_cache_name", cache_name), + ("aws_iam_region", region), + ) + missing_settings: Final = tuple(name for name, value in required_settings if not value) if missing_settings: raise ValueError("AWS ElastiCache IAM Redis authentication requires: " + ", ".join(missing_settings)) return ElastiCacheIAMCredentialProvider( - user_name=str(required_settings["aws_iam_user_name"]), - cache_name=str(required_settings["aws_iam_cache_name"]), - region=str(required_settings["aws_iam_region"]), + user_name=str(user_name), + cache_name=str(cache_name), + region=str(region), ) @@ -507,10 +506,14 @@ def _get_redis_client_logic(**env_overrides): "for Redis. Using Azure AD. Remove one to avoid misconfiguration." ) elif _aws_iam_enabled: - aws_provider: Final = _build_elasticache_iam_provider(redis_kwargs) - if aws_provider is not None: - verbose_logger.debug("Setting up AWS ElastiCache IAM authentication for Redis.") - redis_kwargs["credential_provider"] = aws_provider + verbose_logger.debug("Setting up AWS ElastiCache IAM authentication for Redis.") + redis_kwargs["credential_provider"] = _build_elasticache_iam_provider( + user_name=redis_kwargs.get("aws_iam_user_name"), + cache_name=redis_kwargs.get("aws_iam_cache_name"), + region=redis_kwargs.get("aws_iam_region") + or get_secret_str("AWS_REGION") + or get_secret_str("AWS_DEFAULT_REGION"), + ) redis_kwargs.pop("gcp_service_account", None) redis_kwargs.pop("gcp_ssl_ca_certs", None) diff --git a/litellm/_redis_credential_provider.py b/litellm/_redis_credential_provider.py index 800473bfa54..9a07d9e62bb 100644 --- a/litellm/_redis_credential_provider.py +++ b/litellm/_redis_credential_provider.py @@ -1,11 +1,19 @@ +from __future__ import annotations + import asyncio import threading import time -from typing import Any, Final, Protocol +from collections.abc import Callable +from typing import TYPE_CHECKING, Any, Final, Protocol from urllib.parse import quote from redis.credentials import CredentialProvider +if TYPE_CHECKING: + from botocore.credentials import Credentials +else: + Credentials = Any # rebind-ok: runtime alias for the type-checking-only botocore import + # Azure AD scope for Redis Cache for Azure. AZURE_REDIS_SCOPE: Final = "https://redis.azure.com/.default" @@ -109,38 +117,24 @@ _ELASTICACHE_SERVICE_NAME: Final = "elasticache" _ELASTICACHE_TOKEN_TTL_SECONDS: Final = 900 -class _FrozenBotocoreCredentials(Protocol): - access_key: str - secret_key: str - token: str | None - - -class _BotocoreCredentials(Protocol): - def get_frozen_credentials(self) -> _FrozenBotocoreCredentials: ... - - -class _BotocoreCredentialsResolver(Protocol): - def __call__(self) -> _BotocoreCredentials | None: ... - - class ElastiCacheIAMCredentialProvider(CredentialProvider): def __init__( self, user_name: str, cache_name: str, region: str, - credentials_resolver: _BotocoreCredentialsResolver | None = None, + credentials_resolver: Callable[[], Credentials | None] | None = None, token_lifetime_seconds: int = _ELASTICACHE_TOKEN_TTL_SECONDS, ) -> None: self._user_name = user_name self._cache_name = cache_name self._region = region self._credentials_resolver = credentials_resolver or self._resolve_credentials - self._credentials: _BotocoreCredentials | None = None + self._credentials: Credentials | None = None self._token_lifetime_seconds = token_lifetime_seconds @staticmethod - def _resolve_credentials() -> Any: + def _resolve_credentials() -> Credentials | None: try: import botocore.session except ImportError as e: @@ -187,6 +181,14 @@ class ElastiCacheIAMCredentialProvider(CredentialProvider): return await asyncio.to_thread(self._get_credentials) +class _AzureToken(Protocol): + token: str + + +class _AzureCredential(Protocol): + def get_token(self, scope: str) -> _AzureToken: ... + + class AzureADCredentialProvider(CredentialProvider): """ redis.credentials.CredentialProvider implementation that supplies Azure AD @@ -198,7 +200,7 @@ class AzureADCredentialProvider(CredentialProvider): fail authentication after the initial token expired (~1 hour TTL). """ - def __init__(self, credential: Any, username: str | None = None) -> None: + def __init__(self, credential: _AzureCredential, username: str | None = None) -> None: self._credential = credential self._username = username diff --git a/ruff-strict-budget.json b/ruff-strict-budget.json index b24335ad112..02d47ed971d 100644 --- a/ruff-strict-budget.json +++ b/ruff-strict-budget.json @@ -24,7 +24,7 @@ "limit": 133 }, "ANN401": { - "limit": 1153 + "limit": 1152 }, "ASYNC230": { "limit": 11 diff --git a/tests/test_litellm/test_redis_credential_provider.py b/tests/test_litellm/test_redis_credential_provider.py index 6299d4d42e0..de4afb7c313 100644 --- a/tests/test_litellm/test_redis_credential_provider.py +++ b/tests/test_litellm/test_redis_credential_provider.py @@ -4,10 +4,7 @@ from urllib.parse import parse_qs, urlsplit import pytest -from litellm._redis_credential_provider import ( - ElastiCacheIAMCredentialProvider, - _BotocoreCredentials, -) +from litellm._redis_credential_provider import ElastiCacheIAMCredentialProvider class _FakeCredentials: @@ -20,16 +17,6 @@ class _FakeCredentials: return self -class _FakeResolver: - def __init__(self, credentials: _BotocoreCredentials | None) -> None: - self.credentials = credentials - self.calls = 0 - - def __call__(self): - self.calls += 1 - return self.credentials - - class _RotatingFakeCredentials: def __init__(self) -> None: self.calls = 0 @@ -46,6 +33,16 @@ class _RotatingFakeCredentials: ) +class _FakeResolver: + def __init__(self, credentials: _FakeCredentials | _RotatingFakeCredentials | None) -> None: + self.credentials = credentials + self.calls = 0 + + def __call__(self): + self.calls += 1 + return self.credentials + + def test_elasticache_provider_signs_expected_query(): resolver = _FakeResolver(_FakeCredentials("AKIA-SYNTHETIC")) provider = ElastiCacheIAMCredentialProvider( diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index c77b6ad2e46..7bfb59969c6 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -24559,6 +24559,26 @@ export interface components { * independently of the response-cache backend in `litellm_settings.cache_params`. */ CoordinationRedisParams: { + /** + * Aws Iam Auth + * @description enable AWS ElastiCache IAM authentication + */ + aws_iam_auth?: boolean | string | null; + /** + * Aws Iam Cache Name + * @description AWS ElastiCache cache name + */ + aws_iam_cache_name?: string | null; + /** + * Aws Iam Region + * @description AWS region for ElastiCache IAM authentication + */ + aws_iam_region?: string | null; + /** + * Aws Iam User Name + * @description AWS ElastiCache IAM user name + */ + aws_iam_user_name?: string | null; /** * Host * @description Redis hostname From a569aca9e8d6c5e76ba3b3901bf4b309fd729f98 Mon Sep 17 00:00:00 2001 From: eugene-yao-zocdoc Date: Wed, 26 Aug 2026 19:48:03 -0400 Subject: [PATCH 4/8] test(redis): cover ElastiCache IAM failures Generated with AI Co-Authored-By: Claude Code --- tests/test_litellm/test_redis.py | 52 ++++++-------- .../test_redis_credential_provider.py | 70 +++++++++++++++++++ 2 files changed, 93 insertions(+), 29 deletions(-) diff --git a/tests/test_litellm/test_redis.py b/tests/test_litellm/test_redis.py index 3dffaf7e72a..37b27adf0ed 100644 --- a/tests/test_litellm/test_redis.py +++ b/tests/test_litellm/test_redis.py @@ -454,37 +454,33 @@ def test_explicit_provider_wins_over_aws_iam(clean_redis_environment): def test_gcp_wins_over_aws_iam(clean_redis_environment): - with patch("litellm._redis.create_gcp_iam_redis_connect_func") as mock_gcp: - mock_gcp.return_value = _gcp_marker_callback() - redis_kwargs = _get_redis_client_logic( - host="cache.example.com", - port=6379, - gcp_service_account="sa@example.com", - aws_iam_auth=True, - aws_iam_user_name="iam-user", - aws_iam_cache_name="cache.example.com", - aws_iam_region="us-east-1", - ) + redis_kwargs = _get_redis_client_logic( + host="cache.example.com", + port=6379, + gcp_service_account="sa@example.com", + aws_iam_auth=True, + aws_iam_user_name="iam-user", + aws_iam_cache_name="cache.example.com", + aws_iam_region="us-east-1", + ) assert "credential_provider" not in redis_kwargs - assert redis_kwargs["redis_connect_func"] is mock_gcp.return_value + assert redis_kwargs["redis_connect_func"]._gcp_service_account == "sa@example.com" def test_azure_wins_over_aws_iam(clean_redis_environment): - with patch("litellm._redis.create_azure_ad_redis_connect_func") as mock_azure: - mock_azure.return_value = MagicMock() - redis_kwargs = _get_redis_client_logic( - host="cache.example.com", - port=6379, - azure_redis_ad_token="true", - aws_iam_auth=True, - aws_iam_user_name="iam-user", - aws_iam_cache_name="cache.example.com", - aws_iam_region="us-east-1", - ) + redis_kwargs = _get_redis_client_logic( + host="cache.example.com", + port=6379, + azure_redis_ad_token="true", + aws_iam_auth=True, + aws_iam_user_name="iam-user", + aws_iam_cache_name="cache.example.com", + aws_iam_region="us-east-1", + ) assert "credential_provider" not in redis_kwargs - assert redis_kwargs["redis_connect_func"] is mock_azure.return_value + assert redis_kwargs["redis_connect_func"]._azure_redis_ad_token is True def test_async_cluster_installs_aws_iam_provider(clean_redis_environment): @@ -1665,11 +1661,9 @@ def test_async_sentinel_keeps_the_credential_provider_off_the_monitors(markers, "redis_connect_func": SimpleNamespace(**markers), } - with ( - patch("litellm._redis.async_redis.Sentinel") as mock_sentinel_cls, - patch("litellm._redis._get_redis_client_logic", return_value=redis_kwargs), - ): - get_redis_async_client() + with patch("litellm._redis.async_redis.Sentinel") as mock_sentinel_cls: + with patch("litellm._redis._get_redis_client_logic", return_value=redis_kwargs): + get_redis_async_client() sentinel_kwargs = mock_sentinel_cls.call_args[1]["sentinel_kwargs"] assert sentinel_kwargs["password"] == sentinel_password diff --git a/tests/test_litellm/test_redis_credential_provider.py b/tests/test_litellm/test_redis_credential_provider.py index de4afb7c313..336e30e12c6 100644 --- a/tests/test_litellm/test_redis_credential_provider.py +++ b/tests/test_litellm/test_redis_credential_provider.py @@ -1,4 +1,6 @@ import asyncio +import builtins +import sys from types import SimpleNamespace from urllib.parse import parse_qs, urlsplit @@ -87,6 +89,41 @@ def test_elasticache_provider_resolves_credentials_once_but_refreshes_signature( assert rotating_credentials.calls == 3 +def test_elasticache_provider_uses_botocore_session_credentials(monkeypatch): + credentials = _FakeCredentials("AKIA-SYNTHETIC") + monkeypatch.setattr("botocore.session.get_session", lambda: SimpleNamespace(get_credentials=lambda: credentials)) + provider = ElastiCacheIAMCredentialProvider( + user_name="iam-user", + cache_name="cache.example.com", + region="us-east-1", + ) + + user_name, token = provider.get_credentials() + + assert user_name == "iam-user" + assert "AKIA-SYNTHETIC" in token + + +def test_elasticache_provider_reports_missing_botocore(monkeypatch): + original_import = builtins.__import__ + + def import_without_botocore(name, *args, **kwargs): + if name == "botocore.session": + raise ImportError("synthetic missing dependency") + return original_import(name, *args, **kwargs) + + monkeypatch.delitem(sys.modules, "botocore.session", raising=False) + monkeypatch.setattr(builtins, "__import__", import_without_botocore) + provider = ElastiCacheIAMCredentialProvider( + user_name="iam-user", + cache_name="cache.example.com", + region="us-east-1", + ) + + with pytest.raises(ImportError, match="pip install boto3"): + provider.get_credentials() + + def test_elasticache_provider_reports_missing_credentials(): provider = ElastiCacheIAMCredentialProvider( user_name="iam-user", @@ -99,6 +136,39 @@ def test_elasticache_provider_reports_missing_credentials(): provider.get_credentials() +def test_elasticache_provider_reports_missing_frozen_credentials(): + provider = ElastiCacheIAMCredentialProvider( + user_name="iam-user", + cache_name="cache.example.com", + region="us-east-1", + credentials_resolver=_FakeResolver(SimpleNamespace(get_frozen_credentials=lambda: None)), + ) + + with pytest.raises(RuntimeError, match="Unable to resolve AWS credentials"): + provider.get_credentials() + + +def test_elasticache_provider_reports_missing_signing_dependency(monkeypatch): + original_import = builtins.__import__ + + def import_without_botocore_auth(name, *args, **kwargs): + if name == "botocore.auth": + raise ImportError("synthetic missing dependency") + return original_import(name, *args, **kwargs) + + monkeypatch.delitem(sys.modules, "botocore.auth", raising=False) + monkeypatch.setattr(builtins, "__import__", import_without_botocore_auth) + provider = ElastiCacheIAMCredentialProvider( + user_name="iam-user", + cache_name="cache.example.com", + region="us-east-1", + credentials_resolver=_FakeResolver(_FakeCredentials("AKIA-SYNTHETIC")), + ) + + with pytest.raises(ImportError, match="pip install boto3"): + provider.get_credentials() + + def test_elasticache_provider_recovers_after_a_failed_resolution(): resolver = _FakeResolver(None) provider = ElastiCacheIAMCredentialProvider( From 69c43cd65d349e6c3db6d0178abd41646cbb3d59 Mon Sep 17 00:00:00 2001 From: eugene-yao-zocdoc Date: Wed, 26 Aug 2026 20:21:03 -0400 Subject: [PATCH 5/8] fix(proxy): validate coordination Redis mappings Generated with AI Co-Authored-By: Claude Code --- .../management_endpoints/coordination_redis_endpoints.py | 2 +- litellm/proxy/proxy_server.py | 4 ++-- 2 files changed, 3 insertions(+), 3 deletions(-) diff --git a/litellm/proxy/management_endpoints/coordination_redis_endpoints.py b/litellm/proxy/management_endpoints/coordination_redis_endpoints.py index 86ce336c7a3..88dc09ab001 100644 --- a/litellm/proxy/management_endpoints/coordination_redis_endpoints.py +++ b/litellm/proxy/management_endpoints/coordination_redis_endpoints.py @@ -140,7 +140,7 @@ def _merge_over_saved( def _validated_params(settings: Mapping[str, object]) -> CoordinationRedisParams: """Validate settings the way startup does: resolve env refs, then require a connection target.""" try: - params: Final = CoordinationRedisParams(**_resolve_env_refs(settings)) + params: Final = CoordinationRedisParams.model_validate(_resolve_env_refs(settings)) except ValidationError as e: invalid_fields: Final = sorted({str(error["loc"][0]) for error in e.errors() if error["loc"]}) raise HTTPException( diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 54abbb22424..55822b4cf95 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -4533,7 +4533,7 @@ class ProxyConfig: if not isinstance(raw_params, dict): raise ValueError("general_settings.coordination_redis must be a mapping of Redis connection params") - coordination_params: Final = CoordinationRedisParams(**_resolve_coordination_redis_env_refs(raw_params)) + coordination_params: Final = CoordinationRedisParams.model_validate(_resolve_coordination_redis_env_refs(raw_params)) if not coordination_params.has_connection_target(): raise ValueError( "general_settings.coordination_redis needs a connection target: " @@ -8661,7 +8661,7 @@ class ProxyStartupEvent: if persisted is None: return None - coordination_params: Final = CoordinationRedisParams(**_resolve_coordination_redis_env_refs(persisted)) + coordination_params: Final = CoordinationRedisParams.model_validate(_resolve_coordination_redis_env_refs(persisted)) if not coordination_params.has_connection_target(): verbose_proxy_logger.warning( "coordination_redis saved in the database names no connection target; ignoring it." From ccb390d5c2c8e5408cff06e530d44b5eb56c2fce Mon Sep 17 00:00:00 2001 From: eugene-yao-zocdoc Date: Wed, 26 Aug 2026 20:26:26 -0400 Subject: [PATCH 6/8] style(proxy): format coordination Redis validation Generated with AI Co-Authored-By: Claude Code --- litellm/proxy/proxy_server.py | 8 ++++++-- 1 file changed, 6 insertions(+), 2 deletions(-) diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 55822b4cf95..d1ef69a8632 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -4533,7 +4533,9 @@ class ProxyConfig: if not isinstance(raw_params, dict): raise ValueError("general_settings.coordination_redis must be a mapping of Redis connection params") - coordination_params: Final = CoordinationRedisParams.model_validate(_resolve_coordination_redis_env_refs(raw_params)) + coordination_params: Final = CoordinationRedisParams.model_validate( + _resolve_coordination_redis_env_refs(raw_params) + ) if not coordination_params.has_connection_target(): raise ValueError( "general_settings.coordination_redis needs a connection target: " @@ -8661,7 +8663,9 @@ class ProxyStartupEvent: if persisted is None: return None - coordination_params: Final = CoordinationRedisParams.model_validate(_resolve_coordination_redis_env_refs(persisted)) + coordination_params: Final = CoordinationRedisParams.model_validate( + _resolve_coordination_redis_env_refs(persisted) + ) if not coordination_params.has_connection_target(): verbose_proxy_logger.warning( "coordination_redis saved in the database names no connection target; ignoring it." From 0e4e6cc0025393ed71abd6ab39cfe1953c7082a7 Mon Sep 17 00:00:00 2001 From: eugene-yao-zocdoc Date: Wed, 26 Aug 2026 21:14:07 -0400 Subject: [PATCH 7/8] ci: retry CodSpeed result upload Generated with AI Co-Authored-By: Claude Code From 3abaa91d284bc5cbd2c8c7da50764d311f542eaa Mon Sep 17 00:00:00 2001 From: eugene-yao-zocdoc Date: Thu, 27 Aug 2026 13:49:12 -0400 Subject: [PATCH 8/8] fix(redis): require TLS for ElastiCache IAM auth --- litellm/_redis.py | 9 ++++ tests/test_litellm/test_redis.py | 78 ++++++++++++++++++++++++++++++++ 2 files changed, 87 insertions(+) diff --git a/litellm/_redis.py b/litellm/_redis.py index 289bc883ec7..baae24531d8 100644 --- a/litellm/_redis.py +++ b/litellm/_redis.py @@ -181,6 +181,13 @@ def _is_true(value: object | None) -> bool: return value is True or (isinstance(value, str) and value.lower() == "true") +def _uses_tls(redis_kwargs: Mapping[str, object]) -> bool: + if redis_kwargs.get("startup_nodes") is not None: + return _is_true(redis_kwargs.get("ssl")) + url: Final = redis_kwargs.get("url") + return urlsplit(url).scheme.lower() == "rediss" if isinstance(url, str) else _is_true(redis_kwargs.get("ssl")) + + def _build_elasticache_iam_provider( user_name: object | None, cache_name: object | None, @@ -506,6 +513,8 @@ def _get_redis_client_logic(**env_overrides): "for Redis. Using Azure AD. Remove one to avoid misconfiguration." ) elif _aws_iam_enabled: + if not _uses_tls(redis_kwargs): + raise ValueError("AWS ElastiCache IAM Redis authentication requires TLS") verbose_logger.debug("Setting up AWS ElastiCache IAM authentication for Redis.") redis_kwargs["credential_provider"] = _build_elasticache_iam_provider( user_name=redis_kwargs.get("aws_iam_user_name"), diff --git a/tests/test_litellm/test_redis.py b/tests/test_litellm/test_redis.py index 37b27adf0ed..f72c00916b5 100644 --- a/tests/test_litellm/test_redis.py +++ b/tests/test_litellm/test_redis.py @@ -318,6 +318,7 @@ def test_aws_iam_environment_settings_install_provider(clean_redis_environment, monkeypatch.setenv("REDIS_AWS_IAM_USER_NAME", "iam-user") monkeypatch.setenv("REDIS_AWS_IAM_CACHE_NAME", "cache.example.com") monkeypatch.setenv("REDIS_AWS_IAM_REGION", "us-east-1") + monkeypatch.setenv("REDIS_SSL", "true") redis_kwargs = _get_redis_client_logic(host="cache.example.com", port=6379) @@ -325,6 +326,77 @@ def test_aws_iam_environment_settings_install_provider(clean_redis_environment, assert not {"aws_iam_auth", "aws_iam_user_name", "aws_iam_cache_name", "aws_iam_region"} & redis_kwargs.keys() +@pytest.mark.parametrize( + "transport", + [ + pytest.param({"host": "cache.example.com", "port": 6379}, id="host_without_ssl"), + pytest.param({"host": "cache.example.com", "port": 6379, "ssl": False}, id="host_ssl_false"), + pytest.param({"url": "redis://cache.example.com:6379", "ssl": True}, id="plaintext_url"), + pytest.param( + {"startup_nodes": [{"host": "cache.example.com", "port": 6379}]}, + id="cluster_without_ssl", + ), + pytest.param( + { + "url": "rediss://cache.example.com:6379", + "startup_nodes": [{"host": "cache.example.com", "port": 6379}], + }, + id="cluster_without_ssl_ignores_url_scheme", + ), + pytest.param( + { + "sentinel_nodes": [("sentinel.example.com", 26379)], + "service_name": "cache", + }, + id="sentinel_without_ssl", + ), + ], +) +def test_aws_iam_auth_rejects_non_tls_connections(clean_redis_environment, transport): + with pytest.raises(ValueError, match="requires TLS"): + _get_redis_client_logic( + **transport, + aws_iam_auth=True, + aws_iam_user_name="iam-user", + aws_iam_cache_name="cache.example.com", + aws_iam_region="us-east-1", + ) + + +@pytest.mark.parametrize( + "transport", + [ + pytest.param({"host": "cache.example.com", "port": 6379, "ssl": True}, id="host"), + pytest.param({"url": "rediss://cache.example.com:6379"}, id="url"), + pytest.param( + { + "startup_nodes": [{"host": "cache.example.com", "port": 6379}], + "ssl": True, + }, + id="cluster", + ), + pytest.param( + { + "sentinel_nodes": [("sentinel.example.com", 26379)], + "service_name": "cache", + "ssl": True, + }, + id="sentinel", + ), + ], +) +def test_aws_iam_auth_accepts_tls_connections(clean_redis_environment, transport): + redis_kwargs = _get_redis_client_logic( + **transport, + aws_iam_auth=True, + aws_iam_user_name="iam-user", + aws_iam_cache_name="cache.example.com", + aws_iam_region="us-east-1", + ) + + assert isinstance(redis_kwargs["credential_provider"], ElastiCacheIAMCredentialProvider) + + def test_aws_iam_settings_are_removed_for_url_and_static_credentials(clean_redis_environment): redis_kwargs = _get_redis_client_logic( url="rediss://url-user:url-pass@cache.example.com:6380", @@ -350,6 +422,7 @@ def test_aws_iam_missing_setting_fails_closed(clean_redis_environment, missing): "aws_iam_user_name": "iam-user", "aws_iam_cache_name": "cache.example.com", "aws_iam_region": "us-east-1", + "ssl": True, } settings[missing] = None @@ -364,6 +437,7 @@ def test_aws_iam_region_falls_back_to_environment(clean_redis_environment, monke redis_kwargs = _get_redis_client_logic( host="cache.example.com", port=6379, + ssl=True, aws_iam_auth=True, aws_iam_user_name="iam-user", aws_iam_cache_name="cache.example.com", @@ -381,6 +455,7 @@ def test_aws_iam_region_prefers_aws_region_over_default_region(clean_redis_envir redis_kwargs = _get_redis_client_logic( host="cache.example.com", port=6379, + ssl=True, aws_iam_auth=True, aws_iam_user_name="iam-user", aws_iam_cache_name="cache.example.com", @@ -395,6 +470,7 @@ def test_aws_iam_region_prefers_explicit_over_environment(clean_redis_environmen redis_kwargs = _get_redis_client_logic( host="cache.example.com", port=6379, + ssl=True, aws_iam_auth=True, aws_iam_user_name="iam-user", aws_iam_cache_name="cache.example.com", @@ -408,6 +484,7 @@ def test_aws_iam_settings_map_to_distinct_provider_fields(clean_redis_environmen redis_kwargs = _get_redis_client_logic( host="cache.example.com", port=6379, + ssl=True, aws_iam_auth=True, aws_iam_user_name="iam-user-value", aws_iam_cache_name="iam-cache-value", @@ -488,6 +565,7 @@ def test_async_cluster_installs_aws_iam_provider(clean_redis_environment): client = get_redis_async_client( startup_nodes=startup_nodes, + ssl=True, aws_iam_auth=True, aws_iam_user_name="iam-user", aws_iam_cache_name="cache.example.com",