mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
The ElastiCache IAM gate read `aws_iam_auth` and `ssl` with a helper that only accepted the literal string "true", while the kwarg coercion that runs later accepts "true", "1" and "yes". Type coercion happens after the gate, so `REDIS_AWS_IAM_AUTH=1` silently skipped IAM auth and `REDIS_SSL=1` made the "requires TLS" check fail closed on a connection that was in fact TLS. Both helpers now share `_str_to_bool`. AWS signs serverless cache tokens with an extra `ResourceType=ServerlessCache` query parameter, so tokens minted for a serverless cache were rejected. Adds an `aws_iam_serverless` setting (`REDIS_AWS_IAM_SERVERLESS`) that puts the parameter into the signed URL, and lowercases the cache name because ElastiCache lowercases it at creation time.
228 lines
8.5 KiB
Python
228 lines
8.5 KiB
Python
from __future__ import annotations
|
|
|
|
import asyncio
|
|
import threading
|
|
import time
|
|
from collections.abc import Callable
|
|
from typing import TYPE_CHECKING, Final, Protocol
|
|
from urllib.parse import urlencode
|
|
|
|
from redis.credentials import CredentialProvider
|
|
|
|
if TYPE_CHECKING:
|
|
from botocore.credentials import Credentials
|
|
|
|
# Azure AD scope for Redis Cache for Azure.
|
|
AZURE_REDIS_SCOPE: Final = "https://redis.azure.com/.default"
|
|
|
|
# GCP IAM tokens are valid for 1 hour. Cache for 55 minutes to refresh before expiry.
|
|
_GCP_IAM_TOKEN_TTL_SECONDS: Final = 3300
|
|
|
|
# Module-level cache shared across all GCPIAMCredentialProvider instances for the
|
|
# same service account, so multiple Redis connections on the same pod share one token.
|
|
# Keyed by service_account → (token, expiry_monotonic_timestamp).
|
|
_token_cache: Final[dict[str, tuple[str, float]]] = {}
|
|
_token_cache_lock: Final = threading.Lock()
|
|
|
|
|
|
class AzureAccessToken(Protocol):
|
|
"""The ``azure.core.credentials.AccessToken`` shape this module reads."""
|
|
|
|
@property
|
|
def token(self) -> str: ...
|
|
|
|
|
|
class AzureCredential(Protocol):
|
|
"""The ``azure-identity`` credential surface this module calls."""
|
|
|
|
def get_token(self, *scopes: str) -> AzureAccessToken: ...
|
|
|
|
|
|
def _generate_gcp_iam_access_token(service_account: str) -> str:
|
|
"""
|
|
Generate GCP IAM access token for Redis authentication.
|
|
|
|
Args:
|
|
service_account: GCP service account in format 'projects/-/serviceAccounts/name@project.iam.gserviceaccount.com'
|
|
|
|
Returns:
|
|
Access token string for GCP IAM authentication
|
|
"""
|
|
try:
|
|
from google.cloud import iam_credentials_v1
|
|
except ImportError:
|
|
raise ImportError(
|
|
"google-cloud-iam is required for GCP IAM Redis authentication. "
|
|
"Install it with: pip install google-cloud-iam"
|
|
)
|
|
|
|
client: Final = iam_credentials_v1.IAMCredentialsClient()
|
|
request: Final = iam_credentials_v1.GenerateAccessTokenRequest(
|
|
name=service_account,
|
|
scope=["https://www.googleapis.com/auth/cloud-platform"],
|
|
)
|
|
response: Final = client.generate_access_token(request=request)
|
|
return str(response.access_token)
|
|
|
|
|
|
def _get_cached_gcp_iam_token(service_account: str) -> str:
|
|
"""
|
|
Return a cached GCP IAM token, refreshing only when expired.
|
|
|
|
Uses a module-level cache shared across all GCPIAMCredentialProvider
|
|
instances for the same service account. The threading.Lock ensures only
|
|
one thread performs the network round-trip on expiry; all others wait
|
|
briefly and read the fresh token (double-checked locking pattern).
|
|
|
|
This avoids N concurrent blocking IAM refreshes when N Redis connections
|
|
are established simultaneously (e.g. during health checks or pool warm-up),
|
|
which would otherwise serialise inside Python's async event loop and cause
|
|
cascading request latency.
|
|
"""
|
|
cached = _token_cache.get(service_account)
|
|
if cached is not None:
|
|
token, expiry = cached
|
|
if time.monotonic() < expiry:
|
|
return token
|
|
|
|
with _token_cache_lock:
|
|
# Re-check inside the lock: another thread may have refreshed already.
|
|
cached = _token_cache.get(service_account)
|
|
if cached is not None:
|
|
token, expiry = cached
|
|
if time.monotonic() < expiry:
|
|
return token
|
|
|
|
token = _generate_gcp_iam_access_token(service_account)
|
|
_token_cache[service_account] = (
|
|
token,
|
|
time.monotonic() + _GCP_IAM_TOKEN_TTL_SECONDS,
|
|
)
|
|
return token
|
|
|
|
|
|
class GCPIAMCredentialProvider(CredentialProvider):
|
|
"""
|
|
redis.credentials.CredentialProvider implementation that supplies GCP IAM tokens
|
|
for Redis authentication, with module-level caching per service account.
|
|
|
|
Tokens are cached for _GCP_IAM_TOKEN_TTL_SECONDS (55 min) so that repeated
|
|
connection establishments — e.g. during connection pool warm-up or health checks —
|
|
do not each trigger a synchronous network round-trip that would block Python's
|
|
async event loop and cause cascading request latency.
|
|
"""
|
|
|
|
def __init__(self, gcp_service_account: str) -> None:
|
|
self._gcp_service_account = gcp_service_account
|
|
|
|
def get_credentials(self) -> tuple[str]:
|
|
token: Final = _get_cached_gcp_iam_token(self._gcp_service_account)
|
|
return (token,)
|
|
|
|
async def get_credentials_async(self) -> tuple[str]:
|
|
token: Final = await asyncio.to_thread(_get_cached_gcp_iam_token, self._gcp_service_account)
|
|
return (token,)
|
|
|
|
|
|
_ELASTICACHE_SERVICE_NAME: Final = "elasticache"
|
|
_ELASTICACHE_TOKEN_TTL_SECONDS: Final = 900
|
|
_ELASTICACHE_SERVERLESS_RESOURCE_TYPE: Final = "ServerlessCache"
|
|
|
|
|
|
class ElastiCacheIAMCredentialProvider(CredentialProvider):
|
|
def __init__(
|
|
self,
|
|
user_name: str,
|
|
cache_name: str,
|
|
region: str,
|
|
is_serverless: bool = False,
|
|
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.lower()
|
|
self._region = region
|
|
self._is_serverless = is_serverless
|
|
self._credentials_resolver = credentials_resolver or self._resolve_credentials
|
|
self._credentials: Credentials | None = None
|
|
self._token_lifetime_seconds = token_lifetime_seconds
|
|
|
|
@staticmethod
|
|
def _resolve_credentials() -> Credentials | None:
|
|
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()
|
|
|
|
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
|
|
|
|
query: Final = urlencode(
|
|
(
|
|
("Action", "connect"),
|
|
("User", self._user_name),
|
|
*((("ResourceType", _ELASTICACHE_SERVERLESS_RESOURCE_TYPE),) if self._is_serverless else ()),
|
|
)
|
|
)
|
|
request: Final = AWSRequest(method="GET", url=f"https://{self._cache_name}/?{query}")
|
|
SigV4QueryAuth(
|
|
frozen_credentials,
|
|
_ELASTICACHE_SERVICE_NAME,
|
|
self._region,
|
|
expires=self._token_lifetime_seconds,
|
|
).add_auth(request)
|
|
signed_url: Final = request.url
|
|
if signed_url is None:
|
|
raise RuntimeError("Unable to generate AWS ElastiCache IAM credentials")
|
|
return self._user_name, signed_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
|
|
tokens for Redis authentication.
|
|
|
|
Wraps an azure-identity credential object so the Azure SDK's internal token
|
|
cache and silent refresh are honoured on every Redis connection. This avoids
|
|
the static-token-baked-in-pool issue where pool-managed connections would
|
|
fail authentication after the initial token expired (~1 hour TTL).
|
|
"""
|
|
|
|
def __init__(self, credential: AzureCredential, username: str | None = None) -> None:
|
|
self._credential = credential
|
|
self._username = username
|
|
|
|
def get_credentials(self) -> tuple[str] | tuple[str, str]:
|
|
token: Final = self._credential.get_token(AZURE_REDIS_SCOPE).token
|
|
if self._username:
|
|
return (self._username, token)
|
|
return (token,)
|
|
|
|
async def get_credentials_async(self) -> tuple[str] | tuple[str, str]:
|
|
token_obj: Final = await asyncio.to_thread(self._credential.get_token, AZURE_REDIS_SCOPE)
|
|
if self._username:
|
|
return (self._username, token_obj.token)
|
|
return (token_obj.token,)
|