mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-16 23:41:43 +00:00
feat(redis): add ElastiCache IAM authentication
This commit is contained in:
parent
e5da59336d
commit
1273e46b8b
7 changed files with 520 additions and 9 deletions
|
|
@ -24,6 +24,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,
|
||||
)
|
||||
|
|
@ -75,6 +76,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 _unwrapped_init_args(redis.Redis) if x not in exclude_args} | include_args
|
||||
|
|
@ -270,6 +275,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,
|
||||
|
|
@ -540,6 +571,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(
|
||||
|
|
@ -560,13 +592,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)
|
||||
|
||||
|
|
@ -575,6 +618,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)
|
||||
|
|
|
|||
|
|
@ -1,7 +1,8 @@
|
|||
import asyncio
|
||||
import threading
|
||||
import time
|
||||
from typing import Final, Protocol
|
||||
from typing import Any, Final, Protocol
|
||||
from urllib.parse import quote
|
||||
|
||||
from redis.credentials import CredentialProvider
|
||||
|
||||
|
|
@ -117,6 +118,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
|
||||
|
|
|
|||
|
|
@ -2427,6 +2427,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))
|
||||
|
|
|
|||
|
|
@ -237,4 +237,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,
|
||||
),
|
||||
]
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
),
|
||||
]
|
||||
|
|
|
|||
|
|
@ -24,6 +24,7 @@ from litellm._redis import (
|
|||
)
|
||||
from litellm._redis_credential_provider import (
|
||||
AzureADCredentialProvider,
|
||||
ElastiCacheIAMCredentialProvider,
|
||||
GCPIAMCredentialProvider,
|
||||
_token_cache,
|
||||
)
|
||||
|
|
@ -78,6 +79,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)
|
||||
|
|
@ -110,6 +113,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()
|
||||
|
||||
|
|
@ -300,6 +314,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()
|
||||
|
||||
|
|
@ -1530,9 +1718,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
|
||||
|
|
|
|||
122
tests/test_litellm/test_redis_credential_provider.py
Normal file
122
tests/test_litellm/test_redis_credential_provider.py
Normal file
|
|
@ -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
|
||||
Loading…
Add table
Reference in a new issue