diff --git a/docs/my-website/docs/caching/azure_redis_passwordless.md b/docs/my-website/docs/caching/azure_redis_passwordless.md new file mode 100644 index 00000000000..7728559a41c --- /dev/null +++ b/docs/my-website/docs/caching/azure_redis_passwordless.md @@ -0,0 +1,83 @@ +--- +title: Azure Managed Redis Passwordless (IAM) +--- + +# Azure Managed Redis Passwordless Authentication + +LiteLLM supports [passwordless authentication to Azure Managed Redis](https://learn.microsoft.com/en-us/azure/azure-cache-for-redis/cache-azure-active-directory-for-authentication) using Azure Active Directory (Microsoft Entra ID). This allows you to securely connect to your Redis cache using Azure Managed Identities or Service Principals, avoiding the need to store static connection passwords. + +## Prerequisites +1. **Azure Cache for Redis** instance with **Azure AD Authentication enabled**. +2. **`azure-identity` package** installed in your LiteLLM environment: + ```bash + pip install azure-identity + ``` +3. Your Azure Identity must have a role assignment (e.g. `Data Owner` or `Data Contributor`) to the Redis cache. + +## Configuration + +Enable passwordless authentication in your LiteLLM configuration using `azure_redis_ad_token: true`. + +### 1. Using System-Assigned Managed Identity + +If running on an Azure service with a System-Assigned Managed Identity (e.g., Azure Container Apps, App Service, AKS), you don't need additional credentials. The `DefaultAzureCredential` will automatically discover the identity. + +**Config (`config.yaml`)**: +```yaml +litellm_settings: + cache: true + cache_params: + type: redis + host: .redis.cache.windows.net + port: 6380 + ssl: true + azure_redis_ad_token: true +``` + +*Note: Azure Managed Redis mandates SSL, so `port: 6380` and `ssl: true` are required.* + +### 2. Using User-Assigned Managed Identity + +If using a User-Assigned Managed Identity, provide your `AZURE_CLIENT_ID` via environment variables. + +**Environment variables**: +```bash +export AZURE_CLIENT_ID="" +# Optional. The Object ID of your identity. Defaults to empty string. +export REDIS_USERNAME="" +``` + +### 3. Using Service Principal + +If authenticating via an Azure Service Principal, set the standard Azure identity environment variables: + +**Environment variables**: +```bash +export AZURE_CLIENT_ID="" +export AZURE_TENANT_ID="" +export AZURE_CLIENT_SECRET="" +# Optional. The Object ID of your Service Principal. Defaults to empty string. +export REDIS_USERNAME="" +``` + +Alternatively, you can provide these directly in the `config.yaml`: + +```yaml +litellm_settings: + cache: true + cache_params: + type: redis + host: .redis.cache.windows.net + port: 6380 + ssl: true + azure_redis_ad_token: true + azure_client_id: os.environ/AZURE_CLIENT_ID + azure_tenant_id: os.environ/AZURE_TENANT_ID + azure_client_secret: os.environ/AZURE_CLIENT_SECRET +``` + +## How It Works + +1. LiteLLM uses `azure-identity` to request short-lived access tokens explicitly scoped for Redis (`https://redis.azure.com/.default`). +2. LiteLLM establishes a secure TLS connection with Redis and sends an `AUTH` command using the generated token. +3. Every time the underlying `redis-py` connection disconnects or reconnects, LiteLLM intercepts the connection attempt to **generate a fresh token**, seamlessly handling token expiration. diff --git a/litellm/_redis.py b/litellm/_redis.py index c61582abd1a..aa3689e8672 100644 --- a/litellm/_redis.py +++ b/litellm/_redis.py @@ -34,7 +34,16 @@ def _get_redis_kwargs(): "retry", } - include_args = ["url", "redis_connect_func", "gcp_service_account", "gcp_ssl_ca_certs"] + include_args = [ + "url", + "redis_connect_func", + "gcp_service_account", + "gcp_ssl_ca_certs", + "azure_redis_ad_token", + "azure_client_id", + "azure_tenant_id", + "azure_client_secret", + ] available_args = [x for x in arg_spec.args if x not in exclude_args] + include_args @@ -78,6 +87,10 @@ def _get_redis_cluster_kwargs(client=None): available_args.append("redis_connect_func") # Needed for sync clusters and IAM detection available_args.append("gcp_service_account") available_args.append("gcp_ssl_ca_certs") + available_args.append("azure_redis_ad_token") + available_args.append("azure_client_id") + available_args.append("azure_tenant_id") + available_args.append("azure_client_secret") available_args.append("max_connections") return available_args @@ -103,10 +116,10 @@ def _redis_kwargs_from_environment(): 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 """ @@ -117,11 +130,11 @@ def _generate_gcp_iam_access_token(service_account: str) -> str: "google-cloud-iam is required for GCP IAM Redis authentication. " "Install it with: pip install google-cloud-iam" ) - + client = iam_credentials_v1.IAMCredentialsClient() request = iam_credentials_v1.GenerateAccessTokenRequest( name=service_account, - scope=['https://www.googleapis.com/auth/cloud-platform'], + scope=["https://www.googleapis.com/auth/cloud-platform"], ) response = client.generate_access_token(request=request) return str(response.access_token) @@ -133,14 +146,15 @@ def create_gcp_iam_redis_connect_func( ) -> Callable: """ Creates a custom Redis connection function for GCP IAM authentication. - + Args: service_account: GCP service account in format 'projects/-/serviceAccounts/name@project.iam.gserviceaccount.com' ssl_ca_certs: Path to SSL CA certificate file for secure connections - + Returns: A connection function that can be used with Redis clients """ + def iam_connect(self): """Initialize the connection and authenticate using GCP IAM""" from redis.exceptions import ( @@ -148,55 +162,163 @@ def create_gcp_iam_redis_connect_func( AuthenticationWrongNumberOfArgsError, ) from redis.utils import str_if_bytes - + self._parser.on_connect(self) - + auth_args = (_generate_gcp_iam_access_token(service_account),) self.send_command("AUTH", *auth_args, check_health=False) - + try: auth_response = self.read_response() except AuthenticationWrongNumberOfArgsError: # Fallback to password auth if IAM fails - if hasattr(self, 'password') and self.password: + if hasattr(self, "password") and self.password: self.send_command("AUTH", self.password, check_health=False) auth_response = self.read_response() else: raise - + if str_if_bytes(auth_response) != "OK": raise AuthenticationError("GCP IAM authentication failed") - + return iam_connect +def _generate_azure_ad_redis_token( + azure_client_id: Optional[str] = None, + azure_tenant_id: Optional[str] = None, + azure_client_secret: Optional[str] = None, +) -> str: + """ + Generate Azure AD access token for Redis authentication. + + Uses the azure-identity SDK to obtain a token scoped to Azure Cache for Redis. + Supports: + - Managed Identity (when running on Azure with AZURE_CLIENT_ID) + - Service Principal (AZURE_CLIENT_ID + AZURE_TENANT_ID + AZURE_CLIENT_SECRET) + - DefaultAzureCredential (automatic discovery) + + Args: + azure_client_id: Optional Azure client ID (overrides AZURE_CLIENT_ID env var) + azure_tenant_id: Optional Azure tenant ID (overrides AZURE_TENANT_ID env var) + azure_client_secret: Optional Azure client secret (overrides AZURE_CLIENT_SECRET env var) + + Returns: + Access token string for Azure Redis authentication + """ + try: + from azure.identity import ( + ClientSecretCredential, + DefaultAzureCredential, + ManagedIdentityCredential, + ) + except ImportError: + raise ImportError( + "azure-identity is required for Azure AD Redis authentication. " + "Install it with: pip install azure-identity" + ) + + AZURE_REDIS_SCOPE = "https://redis.azure.com/.default" + + # Determine credential type + _client_id = azure_client_id or os.environ.get("AZURE_CLIENT_ID") + _tenant_id = azure_tenant_id or os.environ.get("AZURE_TENANT_ID") + _client_secret = azure_client_secret or os.environ.get("AZURE_CLIENT_SECRET") + + if _client_id and _tenant_id and _client_secret: + credential = ClientSecretCredential( + client_id=_client_id, + tenant_id=_tenant_id, + client_secret=_client_secret, + ) + elif _client_id: + credential = ManagedIdentityCredential(client_id=_client_id) + else: + credential = DefaultAzureCredential() + + token = credential.get_token(AZURE_REDIS_SCOPE) + return token.token + + +def create_azure_ad_redis_connect_func( + azure_client_id: Optional[str] = None, + azure_tenant_id: Optional[str] = None, + azure_client_secret: Optional[str] = None, +) -> Callable: + """ + Creates a custom Redis connection function for Azure AD authentication. + + Used for sync Redis clients. Generates a fresh Azure AD token on each + connection/reconnection, ensuring token refresh is handled automatically. + + Args: + azure_client_id: Optional Azure client ID + azure_tenant_id: Optional Azure tenant ID + azure_client_secret: Optional Azure client secret + + Returns: + A connection function that can be used with Redis clients via `redis_connect_func` + """ + + def ad_connect(self): + """Initialize the connection and authenticate using Azure AD""" + from redis.exceptions import ( + AuthenticationError, + AuthenticationWrongNumberOfArgsError, + ) + from redis.utils import str_if_bytes + + self._parser.on_connect(self) + + # Get username from REDIS_USERNAME env var or default principal ID + username = os.environ.get("REDIS_USERNAME", "") + + access_token = _generate_azure_ad_redis_token( + azure_client_id=azure_client_id, + azure_tenant_id=azure_tenant_id, + azure_client_secret=azure_client_secret, + ) + + # Azure Redis expects AUTH (Redis 6+ ACL style) + auth_args = (username, access_token) + self.send_command("AUTH", *auth_args, check_health=False) + + try: + auth_response = self.read_response() + except AuthenticationWrongNumberOfArgsError: + # Fallback: try with just the token + self.send_command("AUTH", access_token, check_health=False) + auth_response = self.read_response() + + if str_if_bytes(auth_response) != "OK": + raise AuthenticationError("Azure AD authentication failed for Redis") + + return ad_connect + + def get_redis_url_from_environment(): if "REDIS_URL" in os.environ: return os.environ["REDIS_URL"] if "REDIS_HOST" not in os.environ or "REDIS_PORT" not in os.environ: - raise ValueError( - "Either 'REDIS_URL' or both 'REDIS_HOST' and 'REDIS_PORT' must be specified for Redis." - ) - + raise ValueError("Either 'REDIS_URL' or both 'REDIS_HOST' and 'REDIS_PORT' must be specified for Redis.") + if "REDIS_SSL" in os.environ and os.environ["REDIS_SSL"].lower() == "true": redis_protocol = "rediss" else: redis_protocol = "redis" - + # Build authentication part of URL auth_part = "" if "REDIS_USERNAME" in os.environ and "REDIS_PASSWORD" in os.environ: auth_part = f"{os.environ['REDIS_USERNAME']}:{os.environ['REDIS_PASSWORD']}@" elif "REDIS_PASSWORD" in os.environ: auth_part = f"{os.environ['REDIS_PASSWORD']}@" - - return ( - f"{redis_protocol}://{auth_part}{os.environ['REDIS_HOST']}:{os.environ['REDIS_PORT']}" - ) + + return f"{redis_protocol}://{auth_part}{os.environ['REDIS_HOST']}:{os.environ['REDIS_PORT']}" -def _get_redis_client_logic(**env_overrides): +def _get_redis_client_logic(**env_overrides): # noqa: PLR0915 """ Common functionality across sync + async redis client implementations """ @@ -226,9 +348,9 @@ def _get_redis_client_logic(**env_overrides): if _sentinel_nodes is not None and isinstance(_sentinel_nodes, str): redis_kwargs["sentinel_nodes"] = json.loads(_sentinel_nodes) - _sentinel_password: Optional[str] = redis_kwargs.get( - "sentinel_password", None - ) or get_secret_str("REDIS_SENTINEL_PASSWORD") + _sentinel_password: Optional[str] = redis_kwargs.get("sentinel_password", None) or get_secret_str( + "REDIS_SENTINEL_PASSWORD" + ) if _sentinel_password is not None: redis_kwargs["sentinel_password"] = _sentinel_password @@ -243,24 +365,49 @@ def _get_redis_client_logic(**env_overrides): # Handle GCP IAM authentication _gcp_service_account = redis_kwargs.get("gcp_service_account") or get_secret_str("REDIS_GCP_SERVICE_ACCOUNT") _gcp_ssl_ca_certs = redis_kwargs.get("gcp_ssl_ca_certs") or get_secret_str("REDIS_GCP_SSL_CA_CERTS") - + if _gcp_service_account is not None: verbose_logger.debug("Setting up GCP IAM authentication for Redis with service account.") redis_kwargs["redis_connect_func"] = create_gcp_iam_redis_connect_func( - service_account=_gcp_service_account, - ssl_ca_certs=_gcp_ssl_ca_certs + service_account=_gcp_service_account, ssl_ca_certs=_gcp_ssl_ca_certs ) # Store GCP service account in redis_connect_func for async cluster access redis_kwargs["redis_connect_func"]._gcp_service_account = _gcp_service_account - + # Remove GCP-specific kwargs that shouldn't be passed to Redis client redis_kwargs.pop("gcp_service_account", None) redis_kwargs.pop("gcp_ssl_ca_certs", None) - + # Only enable SSL if explicitly requested AND SSL CA certs are provided if _gcp_ssl_ca_certs and redis_kwargs.get("ssl", False): redis_kwargs["ssl_ca_certs"] = _gcp_ssl_ca_certs + # Handle Azure AD authentication (after GCP IAM block) + _azure_redis_ad_token = redis_kwargs.get("azure_redis_ad_token") or get_secret("REDIS_AZURE_AD_TOKEN") + + if _azure_redis_ad_token is not None and str(_azure_redis_ad_token).lower() == "true": + _azure_client_id = redis_kwargs.get("azure_client_id") or get_secret_str("AZURE_CLIENT_ID") + _azure_tenant_id = redis_kwargs.get("azure_tenant_id") or get_secret_str("AZURE_TENANT_ID") + _azure_client_secret = redis_kwargs.get("azure_client_secret") or get_secret_str("AZURE_CLIENT_SECRET") + + verbose_logger.debug("Setting up Azure AD authentication for Redis.") + redis_kwargs["redis_connect_func"] = create_azure_ad_redis_connect_func( + azure_client_id=_azure_client_id, + azure_tenant_id=_azure_tenant_id, + azure_client_secret=_azure_client_secret, + ) + # Store Azure config on the function for async cluster access + redis_kwargs["redis_connect_func"]._azure_redis_ad_token = True + redis_kwargs["redis_connect_func"]._azure_client_id = _azure_client_id + redis_kwargs["redis_connect_func"]._azure_tenant_id = _azure_tenant_id + redis_kwargs["redis_connect_func"]._azure_client_secret = _azure_client_secret + + # Remove Azure-specific kwargs that shouldn't be passed to Redis client + redis_kwargs.pop("azure_redis_ad_token", None) + redis_kwargs.pop("azure_client_id", None) + redis_kwargs.pop("azure_tenant_id", None) + redis_kwargs.pop("azure_client_secret", None) + if "url" in redis_kwargs and redis_kwargs["url"] is not None: redis_kwargs.pop("host", None) redis_kwargs.pop("port", None) @@ -268,9 +415,7 @@ def _get_redis_client_logic(**env_overrides): redis_kwargs.pop("password", None) elif "startup_nodes" in redis_kwargs and redis_kwargs["startup_nodes"] is not None: pass - elif ( - "sentinel_nodes" in redis_kwargs and redis_kwargs["sentinel_nodes"] is not None - ): + elif "sentinel_nodes" in redis_kwargs and redis_kwargs["sentinel_nodes"] is not None: pass elif "host" not in redis_kwargs or redis_kwargs["host"] is None: raise ValueError("Either 'host' or 'url' must be specified for redis.") @@ -313,9 +458,7 @@ def _init_redis_sentinel(redis_kwargs) -> redis.Redis: service_name = redis_kwargs.get("service_name") if not sentinel_nodes or not service_name: - raise ValueError( - "Both 'sentinel_nodes' and 'service_name' are required for Redis Sentinel." - ) + raise ValueError("Both 'sentinel_nodes' and 'service_name' are required for Redis Sentinel.") verbose_logger.debug("init_redis_sentinel: sentinel nodes are being initialized.") @@ -337,9 +480,7 @@ def _init_async_redis_sentinel(redis_kwargs) -> async_redis.Redis: service_name = redis_kwargs.get("service_name") if not sentinel_nodes or not service_name: - raise ValueError( - "Both 'sentinel_nodes' and 'service_name' are required for Redis Sentinel." - ) + raise ValueError("Both 'sentinel_nodes' and 'service_name' are required for Redis Sentinel.") verbose_logger.debug("init_redis_sentinel: sentinel nodes are being initialized.") @@ -376,8 +517,9 @@ def get_redis_client(**env_overrides): return redis.Redis(**redis_kwargs) -def get_redis_async_client( - connection_pool: Optional[async_redis.BlockingConnectionPool] = None, **env_overrides, +def get_redis_async_client( # noqa: PLR0915 + connection_pool: Optional[async_redis.BlockingConnectionPool] = None, + **env_overrides, ) -> Union[async_redis.Redis, async_redis.RedisCluster]: redis_kwargs = _get_redis_client_logic(**env_overrides) if "url" in redis_kwargs and redis_kwargs["url"] is not None: @@ -390,9 +532,7 @@ def get_redis_async_client( url_kwargs[arg] = redis_kwargs[arg] else: verbose_logger.debug( - "REDIS: ignoring argument: {}. Not an allowed async_redis.Redis.from_url arg.".format( - arg - ) + "REDIS: ignoring argument: {}. Not an allowed async_redis.Redis.from_url arg.".format(arg) ) return async_redis.Redis.from_url(**url_kwargs) @@ -411,16 +551,20 @@ def get_redis_async_client( # Get GCP service account - first try from redis_connect_func, then from environment gcp_service_account = None - if redis_connect_func and hasattr(redis_connect_func, '_gcp_service_account'): + if redis_connect_func and hasattr(redis_connect_func, "_gcp_service_account"): gcp_service_account = redis_connect_func._gcp_service_account else: gcp_service_account = redis_kwargs.get("gcp_service_account") or get_secret_str("REDIS_GCP_SERVICE_ACCOUNT") - - verbose_logger.debug(f"DEBUG: Redis cluster kwargs: redis_connect_func={redis_connect_func is not None}, gcp_service_account_provided={gcp_service_account is not None}") - + + verbose_logger.debug( + f"DEBUG: Redis cluster kwargs: redis_connect_func={redis_connect_func is not None}, gcp_service_account_provided={gcp_service_account is not None}" + ) + # If GCP IAM is configured (indicated by redis_connect_func), generate access token and use as password if redis_connect_func and gcp_service_account: - verbose_logger.debug("DEBUG: Generating IAM token for service account (value not logged for security reasons)") + verbose_logger.debug( + "DEBUG: Generating IAM token for service account (value not logged for security reasons)" + ) try: # Generate IAM access token using the helper function access_token = _generate_gcp_iam_access_token(gcp_service_account) @@ -429,21 +573,49 @@ def get_redis_async_client( except Exception as e: verbose_logger.error(f"Failed to generate GCP IAM access token: {e}") from redis.exceptions import AuthenticationError + raise AuthenticationError("Failed to generate GCP IAM access token") + # Handle Azure AD authentication for async clusters + elif redis_connect_func and hasattr(redis_connect_func, "_azure_redis_ad_token"): + _az_client_id = getattr(redis_connect_func, "_azure_client_id", None) + _az_tenant_id = getattr(redis_connect_func, "_azure_tenant_id", None) + _az_client_secret = getattr(redis_connect_func, "_azure_client_secret", None) + + verbose_logger.debug("Generating Azure AD token for async Redis cluster") + try: + access_token = _generate_azure_ad_redis_token( + azure_client_id=_az_client_id, + azure_tenant_id=_az_tenant_id, + azure_client_secret=_az_client_secret, + ) + cluster_kwargs["password"] = access_token + # Set username if available + _username = os.environ.get("REDIS_USERNAME", "") + if _username: + cluster_kwargs["username"] = _username + verbose_logger.debug("Successfully generated Azure AD token for async Redis cluster") + except Exception as e: + verbose_logger.error(f"Failed to generate Azure AD access token: {e}") + from redis.exceptions import AuthenticationError + + raise AuthenticationError("Failed to generate Azure AD access token for Redis") else: - verbose_logger.debug(f"DEBUG: Not using GCP IAM auth - redis_connect_func={redis_connect_func is not None}, gcp_service_account_provided={gcp_service_account is not None}") - + verbose_logger.debug( + f"DEBUG: Not using GCP/Azure AD IAM auth - redis_connect_func={redis_connect_func is not None}" + ) + new_startup_nodes: List[ClusterNode] = [] for item in redis_kwargs["startup_nodes"]: new_startup_nodes.append(ClusterNode(**item)) cluster_kwargs.pop("startup_nodes", None) - + # Create async RedisCluster with IAM token as password if available cluster_client = async_redis.RedisCluster( - startup_nodes=new_startup_nodes, **cluster_kwargs # type: ignore + startup_nodes=new_startup_nodes, + **cluster_kwargs, # type: ignore ) - + return cluster_client # Check for Redis Sentinel @@ -463,7 +635,10 @@ def get_redis_connection_pool(**env_overrides): redis_kwargs = _get_redis_client_logic(**env_overrides) verbose_logger.debug("get_redis_connection_pool: redis_kwargs", redis_kwargs) if "url" in redis_kwargs and redis_kwargs["url"] is not None: - pool_kwargs = {"timeout": REDIS_CONNECTION_POOL_TIMEOUT, "url": redis_kwargs["url"]} + pool_kwargs = { + "timeout": REDIS_CONNECTION_POOL_TIMEOUT, + "url": redis_kwargs["url"], + } if "max_connections" in redis_kwargs: try: pool_kwargs["max_connections"] = int(redis_kwargs["max_connections"]) @@ -479,9 +654,8 @@ def get_redis_connection_pool(**env_overrides): redis_kwargs.pop("ssl", None) redis_kwargs["connection_class"] = connection_class redis_kwargs.pop("startup_nodes", None) - return async_redis.BlockingConnectionPool( - timeout=REDIS_CONNECTION_POOL_TIMEOUT, **redis_kwargs - ) + return async_redis.BlockingConnectionPool(timeout=REDIS_CONNECTION_POOL_TIMEOUT, **redis_kwargs) + def _pretty_print_redis_config(redis_kwargs: dict) -> None: """Pretty print the Redis configuration using rich with sensitive data masking""" @@ -492,6 +666,7 @@ def _pretty_print_redis_config(redis_kwargs: dict) -> None: from rich.panel import Panel from rich.table import Table from rich.text import Text + if not verbose_logger.isEnabledFor(logging.DEBUG): return @@ -499,7 +674,7 @@ def _pretty_print_redis_config(redis_kwargs: dict) -> None: # Initialize the sensitive data masker masker = SensitiveDataMasker() - + # Mask sensitive data in redis_kwargs masked_redis_kwargs = masker.mask_dict(redis_kwargs) @@ -531,7 +706,7 @@ def _pretty_print_redis_config(redis_kwargs: dict) -> None: value_str = str(value) else: value_str = str(value) - + config_table.add_row(key, value_str) # Determine connection type @@ -568,4 +743,3 @@ def _pretty_print_redis_config(redis_kwargs: dict) -> None: verbose_logger.info(f"Redis configuration: {masked_redis_kwargs}") except Exception as e: verbose_logger.error(f"Error pretty printing Redis configuration: {e}") - diff --git a/tests/test_litellm/test_utils.py b/tests/test_litellm/test_utils.py index 27e6363e53d..70722c5095c 100644 --- a/tests/test_litellm/test_utils.py +++ b/tests/test_litellm/test_utils.py @@ -6,9 +6,7 @@ from unittest.mock import AsyncMock, MagicMock, patch import pytest from jsonschema import validate -sys.path.insert( - 0, os.path.abspath("../..") -) # Adds the parent directory to the system path +sys.path.insert(0, os.path.abspath("../..")) # Adds the parent directory to the system path import litellm from litellm.proxy.utils import is_valid_api_key @@ -20,11 +18,9 @@ from litellm.types.utils import ( StreamingChoices, ) from litellm.utils import ( - ProviderConfigManager, TextCompletionStreamWrapper, _check_provider_match, _is_streaming_request, - get_llm_provider, get_optional_params_image_gen, is_cached_message, ) @@ -38,22 +34,13 @@ def test_check_provider_match_azure_ai_allows_openai_and_azure(): This is needed for Azure Model Router which can route to OpenAI models. """ # azure_ai should match openai models - assert _check_provider_match( - model_info={"litellm_provider": "openai"}, - custom_llm_provider="azure_ai" - ) is True + assert _check_provider_match(model_info={"litellm_provider": "openai"}, custom_llm_provider="azure_ai") is True # azure_ai should match azure models - assert _check_provider_match( - model_info={"litellm_provider": "azure"}, - custom_llm_provider="azure_ai" - ) is True + assert _check_provider_match(model_info={"litellm_provider": "azure"}, custom_llm_provider="azure_ai") is True # azure_ai should NOT match other providers - assert _check_provider_match( - model_info={"litellm_provider": "anthropic"}, - custom_llm_provider="azure_ai" - ) is False + assert _check_provider_match(model_info={"litellm_provider": "anthropic"}, custom_llm_provider="azure_ai") is False def test_check_provider_match_github_allows_upstream_provider_metadata(): @@ -61,39 +48,38 @@ def test_check_provider_match_github_allows_upstream_provider_metadata(): Test that github provider can match upstream provider metadata. GitHub Models can provide models from multiple providers. """ - assert _check_provider_match( - model_info={"litellm_provider": "openai"}, - custom_llm_provider="github", - ) is True + assert ( + _check_provider_match( + model_info={"litellm_provider": "openai"}, + custom_llm_provider="github", + ) + is True + ) - assert _check_provider_match( - model_info={"litellm_provider": "github"}, - custom_llm_provider="github", - ) is True + assert ( + _check_provider_match( + model_info={"litellm_provider": "github"}, + custom_llm_provider="github", + ) + is True + ) - assert _check_provider_match( - model_info={"litellm_provider": "anthropic"}, - custom_llm_provider="github", - ) is True + assert ( + _check_provider_match( + model_info={"litellm_provider": "anthropic"}, + custom_llm_provider="github", + ) + is True + ) def test_supports_function_calling_github_openai_alias(): assert litellm.utils.supports_function_calling(model="github/gpt-4o-mini") is True - assert ( - litellm.utils.supports_function_calling( - model="gpt-4o-mini", custom_llm_provider="github" - ) - is True - ) + assert litellm.utils.supports_function_calling(model="gpt-4o-mini", custom_llm_provider="github") is True def test_supports_function_calling_github_anthropic_alias(): - assert ( - litellm.utils.supports_function_calling( - model="github/claude-3-5-sonnet-latest" - ) - is True - ) + assert litellm.utils.supports_function_calling(model="github/claude-3-5-sonnet-latest") is True def test_supports_function_calling_deepinfra_llama(): @@ -110,12 +96,7 @@ def test_supports_function_calling_deepinfra_llama(): def test_supports_function_calling_unknown_github_alias_returns_false(): - assert ( - litellm.utils.supports_function_calling( - model="github/non-existent-model-for-capability-check" - ) - is False - ) + assert litellm.utils.supports_function_calling(model="github/non-existent-model-for-capability-check") is False def test_get_optional_params_image_gen(): @@ -168,9 +149,7 @@ def test_get_optional_params_image_gen_vertex_ai_size(): drop_params=True, ) assert optional_params is not None - assert ( - "aspectRatio" not in optional_params - ) # aspectRatio should not be set if size is not provided + assert "aspectRatio" not in optional_params # aspectRatio should not be set if size is not provided assert optional_params["sampleCount"] == 1 @@ -191,26 +170,19 @@ def test_all_model_configs(): VertexAILlama3Config, ) - assert ( - "max_completion_tokens" - in VertexAILlama3Config().get_supported_openai_params(model="llama3") - ) - assert VertexAILlama3Config().map_openai_params( - {"max_completion_tokens": 10}, {}, "llama3", drop_params=False - ) == {"max_tokens": 10} + assert "max_completion_tokens" in VertexAILlama3Config().get_supported_openai_params(model="llama3") + assert VertexAILlama3Config().map_openai_params({"max_completion_tokens": 10}, {}, "llama3", drop_params=False) == { + "max_tokens": 10 + } - assert "max_completion_tokens" in VertexAIAi21Config().get_supported_openai_params( - model="jamba-1.5-mini@001" - ) + assert "max_completion_tokens" in VertexAIAi21Config().get_supported_openai_params(model="jamba-1.5-mini@001") assert VertexAIAi21Config().map_openai_params( {"max_completion_tokens": 10}, {}, "jamba-1.5-mini@001", drop_params=False ) == {"max_tokens": 10} from litellm.llms.fireworks_ai.chat.transformation import FireworksAIConfig - assert "max_completion_tokens" in FireworksAIConfig().get_supported_openai_params( - model="llama3" - ) + assert "max_completion_tokens" in FireworksAIConfig().get_supported_openai_params(model="llama3") assert FireworksAIConfig().map_openai_params( model="llama3", non_default_params={"max_completion_tokens": 10}, @@ -220,9 +192,7 @@ def test_all_model_configs(): from litellm.llms.nvidia_nim.chat.transformation import NvidiaNimConfig - assert "max_completion_tokens" in NvidiaNimConfig().get_supported_openai_params( - model="llama3" - ) + assert "max_completion_tokens" in NvidiaNimConfig().get_supported_openai_params(model="llama3") assert NvidiaNimConfig().map_openai_params( model="llama3", non_default_params={"max_completion_tokens": 10}, @@ -232,9 +202,7 @@ def test_all_model_configs(): from litellm.llms.ollama.chat.transformation import OllamaChatConfig - assert "max_completion_tokens" in OllamaChatConfig().get_supported_openai_params( - model="llama3" - ) + assert "max_completion_tokens" in OllamaChatConfig().get_supported_openai_params(model="llama3") assert OllamaChatConfig().map_openai_params( model="llama3", non_default_params={"max_completion_tokens": 10}, @@ -244,9 +212,7 @@ def test_all_model_configs(): from litellm.llms.predibase.chat.transformation import PredibaseConfig - assert "max_completion_tokens" in PredibaseConfig().get_supported_openai_params( - model="llama3" - ) + assert "max_completion_tokens" in PredibaseConfig().get_supported_openai_params(model="llama3") assert PredibaseConfig().map_openai_params( model="llama3", non_default_params={"max_completion_tokens": 10}, @@ -258,10 +224,7 @@ def test_all_model_configs(): CodestralTextCompletionConfig, ) - assert ( - "max_completion_tokens" - in CodestralTextCompletionConfig().get_supported_openai_params(model="llama3") - ) + assert "max_completion_tokens" in CodestralTextCompletionConfig().get_supported_openai_params(model="llama3") assert CodestralTextCompletionConfig().map_openai_params( model="llama3", non_default_params={"max_completion_tokens": 10}, @@ -273,9 +236,7 @@ def test_all_model_configs(): VolcEngineChatConfig as VolcEngineConfig, ) - assert "max_completion_tokens" in VolcEngineConfig().get_supported_openai_params( - model="llama3" - ) + assert "max_completion_tokens" in VolcEngineConfig().get_supported_openai_params(model="llama3") assert VolcEngineConfig().map_openai_params( model="llama3", non_default_params={"max_completion_tokens": 10}, @@ -285,9 +246,7 @@ def test_all_model_configs(): from litellm.llms.ai21.chat.transformation import AI21ChatConfig - assert "max_completion_tokens" in AI21ChatConfig().get_supported_openai_params( - "jamba-1.5-mini@001" - ) + assert "max_completion_tokens" in AI21ChatConfig().get_supported_openai_params("jamba-1.5-mini@001") assert AI21ChatConfig().map_openai_params( model="jamba-1.5-mini@001", non_default_params={"max_completion_tokens": 10}, @@ -297,9 +256,7 @@ def test_all_model_configs(): from litellm.llms.azure.chat.gpt_transformation import AzureOpenAIConfig - assert "max_completion_tokens" in AzureOpenAIConfig().get_supported_openai_params( - model="gpt-3.5-turbo" - ) + assert "max_completion_tokens" in AzureOpenAIConfig().get_supported_openai_params(model="gpt-3.5-turbo") assert AzureOpenAIConfig().map_openai_params( model="gpt-3.5-turbo", non_default_params={"max_completion_tokens": 10}, @@ -310,11 +267,8 @@ def test_all_model_configs(): from litellm.llms.bedrock.chat.converse_transformation import AmazonConverseConfig - assert ( - "max_completion_tokens" - in AmazonConverseConfig().get_supported_openai_params( - model="anthropic.claude-3-sonnet-20240229-v1:0" - ) + assert "max_completion_tokens" in AmazonConverseConfig().get_supported_openai_params( + model="anthropic.claude-3-sonnet-20240229-v1:0" ) assert AmazonConverseConfig().map_openai_params( model="anthropic.claude-3-sonnet-20240229-v1:0", @@ -327,10 +281,7 @@ def test_all_model_configs(): CodestralTextCompletionConfig, ) - assert ( - "max_completion_tokens" - in CodestralTextCompletionConfig().get_supported_openai_params(model="llama3") - ) + assert "max_completion_tokens" in CodestralTextCompletionConfig().get_supported_openai_params(model="llama3") assert CodestralTextCompletionConfig().map_openai_params( model="llama3", non_default_params={"max_completion_tokens": 10}, @@ -340,11 +291,8 @@ def test_all_model_configs(): from litellm import AmazonAnthropicClaudeConfig, AmazonAnthropicConfig - assert ( - "max_completion_tokens" - in AmazonAnthropicClaudeConfig().get_supported_openai_params( - model="anthropic.claude-3-sonnet-20240229-v1:0" - ) + assert "max_completion_tokens" in AmazonAnthropicClaudeConfig().get_supported_openai_params( + model="anthropic.claude-3-sonnet-20240229-v1:0" ) assert AmazonAnthropicClaudeConfig().map_openai_params( @@ -354,10 +302,7 @@ def test_all_model_configs(): drop_params=False, ) == {"max_tokens": 10} - assert ( - "max_completion_tokens" - in AmazonAnthropicConfig().get_supported_openai_params(model="") - ) + assert "max_completion_tokens" in AmazonAnthropicConfig().get_supported_openai_params(model="") assert AmazonAnthropicConfig().map_openai_params( non_default_params={"max_completion_tokens": 10}, @@ -381,11 +326,8 @@ def test_all_model_configs(): VertexAIAnthropicConfig, ) - assert ( - "max_completion_tokens" - in VertexAIAnthropicConfig().get_supported_openai_params( - model="claude-3-5-sonnet-20240620" - ) + assert "max_completion_tokens" in VertexAIAnthropicConfig().get_supported_openai_params( + model="claude-3-5-sonnet-20240620" ) assert VertexAIAnthropicConfig().map_openai_params( @@ -400,9 +342,7 @@ def test_all_model_configs(): VertexGeminiConfig, ) - assert "max_completion_tokens" in VertexGeminiConfig().get_supported_openai_params( - model="gemini-1.0-pro" - ) + assert "max_completion_tokens" in VertexGeminiConfig().get_supported_openai_params(model="gemini-1.0-pro") assert VertexGeminiConfig().map_openai_params( model="gemini-1.0-pro", @@ -411,12 +351,7 @@ def test_all_model_configs(): drop_params=False, ) == {"max_output_tokens": 10} - assert ( - "max_completion_tokens" - in GoogleAIStudioGeminiConfig().get_supported_openai_params( - model="gemini-1.0-pro" - ) - ) + assert "max_completion_tokens" in GoogleAIStudioGeminiConfig().get_supported_openai_params(model="gemini-1.0-pro") assert GoogleAIStudioGeminiConfig().map_openai_params( model="gemini-1.0-pro", @@ -425,9 +360,7 @@ def test_all_model_configs(): drop_params=False, ) == {"max_output_tokens": 10} - assert "max_completion_tokens" in VertexGeminiConfig().get_supported_openai_params( - model="gemini-1.0-pro" - ) + assert "max_completion_tokens" in VertexGeminiConfig().get_supported_openai_params(model="gemini-1.0-pro") assert VertexGeminiConfig().map_openai_params( model="gemini-1.0-pro", @@ -453,9 +386,7 @@ def test_anthropic_web_search_in_model_info(): model_info = get_model_info(model) assert model_info is not None - assert ( - model_info["supports_web_search"] is True - ), f"Model {model} should support web search" + assert model_info["supports_web_search"] is True, f"Model {model} should support web search" assert ( model_info["search_context_cost_per_query"] is not None ), f"Model {model} should have a search context cost per query" @@ -560,9 +491,7 @@ def validate_model_cost_values(model_data, exceptions=None): continue if isinstance(cost_value, (int, float)) and cost_value > 1: - violations.append( - f"Model '{model_id}' has {field} = {cost_value} which exceeds 1" - ) + violations.append(f"Model '{model_id}' has {field} = {cost_value} which exceeds 1") # Check nested cost fields for field in nested_cost_fields: @@ -646,9 +575,7 @@ def test_aaamodel_prices_and_context_window_json_is_valid(): "input_cost_per_token_cache_hit": {"type": "number"}, "input_cost_per_video_per_second": {"type": "number"}, "input_cost_per_video_per_second_above_8s_interval": {"type": "number"}, - "input_cost_per_video_per_second_above_15s_interval": { - "type": "number" - }, + "input_cost_per_video_per_second_above_15s_interval": {"type": "number"}, "input_cost_per_video_per_second_above_128k_tokens": {"type": "number"}, "input_dbu_cost_per_token": {"type": "number"}, "annotation_cost_per_page": {"type": "number"}, @@ -703,13 +630,9 @@ def test_aaamodel_prices_and_context_window_json_is_valid(): "output_cost_per_token_above_200k_tokens": {"type": "number"}, "output_cost_per_token_above_272k_tokens": {"type": "number"}, "output_cost_per_image_above_1024_and_1024_pixels": {"type": "number"}, - "output_cost_per_image_above_1024_and_1024_pixels_and_premium_image": { - "type": "number" - }, + "output_cost_per_image_above_1024_and_1024_pixels_and_premium_image": {"type": "number"}, "output_cost_per_image_above_512_and_512_pixels": {"type": "number"}, - "output_cost_per_image_above_512_and_512_pixels_and_premium_image": { - "type": "number" - }, + "output_cost_per_image_above_512_and_512_pixels_and_premium_image": {"type": "number"}, "output_cost_per_image_premium_image": {"type": "number"}, "output_cost_per_token_batches": {"type": "number"}, "output_cost_per_reasoning_token": {"type": "number"}, @@ -839,9 +762,7 @@ def test_aaamodel_prices_and_context_window_json_is_valid(): with open(prod_json, "r") as model_prices_file: actual_json = json.load(model_prices_file) assert isinstance(actual_json, dict) - actual_json.pop( - "sample_spec", None - ) # remove the sample, whose schema is inconsistent with the real data + actual_json.pop("sample_spec", None) # remove the sample, whose schema is inconsistent with the real data # Validate schema validate(actual_json, INTENDED_SCHEMA) @@ -877,7 +798,7 @@ def test_max_tokens_consistency(): # Load the model configuration config_path = Path(__file__).parent.parent.parent / "model_prices_and_context_window.json" - with open(config_path, 'r') as f: + with open(config_path, "r") as f: models = json.load(f) inconsistencies = [] @@ -889,22 +810,26 @@ def test_max_tokens_consistency(): # Check if both max_tokens and max_output_tokens exist if isinstance(config, dict): - max_tokens = config.get('max_tokens') - max_output_tokens = config.get('max_output_tokens') + max_tokens = config.get("max_tokens") + max_output_tokens = config.get("max_output_tokens") # Only validate if both exist if max_tokens is not None and max_output_tokens is not None: if max_tokens != max_output_tokens: - inconsistencies.append({ - 'model': model_name, - 'max_tokens': max_tokens, - 'max_output_tokens': max_output_tokens - }) + inconsistencies.append( + { + "model": model_name, + "max_tokens": max_tokens, + "max_output_tokens": max_output_tokens, + } + ) if inconsistencies: error_msg = f"\n\nāŒ Found {len(inconsistencies)} models with max_tokens != max_output_tokens:\n\n" for item in inconsistencies[:10]: # Show first 10 - error_msg += f" {item['model']}: max_tokens={item['max_tokens']}, max_output_tokens={item['max_output_tokens']}\n" + error_msg += ( + f" {item['model']}: max_tokens={item['max_tokens']}, max_output_tokens={item['max_output_tokens']}\n" + ) if len(inconsistencies) > 10: error_msg += f"\n ... and {len(inconsistencies) - 10} more\n" @@ -924,11 +849,11 @@ def test_get_model_info_gemini(): for model, info in model_map.items(): if ( model.startswith("gemini/") - and not "gemma" in model - and not "learnlm" in model - and not "imagen" in model - and not "veo" in model - and not "robotics" in model + and "gemma" not in model + and "learnlm" not in model + and "imagen" not in model + and "veo" not in model + and "robotics" not in model ): assert info.get("tpm") is not None, f"{model} does not have tpm" assert info.get("rpm") is not None, f"{model} does not have rpm" @@ -941,15 +866,10 @@ def test_openai_models_in_model_info(): model_map = litellm.model_cost violated_models = [] for model, info in model_map.items(): - if ( - info.get("litellm_provider") == "openai" - and info.get("supports_vision") is True - ): + if info.get("litellm_provider") == "openai" and info.get("supports_vision") is True: if info.get("supports_pdf_input") is not True: violated_models.append(model) - assert ( - len(violated_models) == 0 - ), f"The following models should support pdf input: {violated_models}" + assert len(violated_models) == 0, f"The following models should support pdf input: {violated_models}" def test_supports_tool_choice_simple_tests(): @@ -957,18 +877,8 @@ def test_supports_tool_choice_simple_tests(): simple sanity checks """ assert litellm.utils.supports_tool_choice(model="gpt-4o") == True - assert ( - litellm.utils.supports_tool_choice( - model="bedrock/anthropic.claude-3-sonnet-20240229-v1:0" - ) - == True - ) - assert ( - litellm.utils.supports_tool_choice( - model="anthropic.claude-3-sonnet-20240229-v1:0" - ) - is True - ) + assert litellm.utils.supports_tool_choice(model="bedrock/anthropic.claude-3-sonnet-20240229-v1:0") == True + assert litellm.utils.supports_tool_choice(model="anthropic.claude-3-sonnet-20240229-v1:0") is True assert ( litellm.utils.supports_tool_choice( @@ -978,17 +888,10 @@ def test_supports_tool_choice_simple_tests(): is True ) + assert litellm.utils.supports_tool_choice(model="us.amazon.nova-micro-v1:0") is False + assert litellm.utils.supports_tool_choice(model="bedrock/us.amazon.nova-micro-v1:0") is False assert ( - litellm.utils.supports_tool_choice(model="us.amazon.nova-micro-v1:0") is False - ) - assert ( - litellm.utils.supports_tool_choice(model="bedrock/us.amazon.nova-micro-v1:0") - is False - ) - assert ( - litellm.utils.supports_tool_choice( - model="us.amazon.nova-micro-v1:0", custom_llm_provider="bedrock_converse" - ) + litellm.utils.supports_tool_choice(model="us.amazon.nova-micro-v1:0", custom_llm_provider="bedrock_converse") is False ) @@ -1081,9 +984,7 @@ def test_supports_computer_use_utility(): try: # Test a model known to support computer_use from backup JSON - supports_cu_anthropic = supports_computer_use( - model="anthropic/claude-4-sonnet-20250514" - ) + supports_cu_anthropic = supports_computer_use(model="anthropic/claude-4-sonnet-20250514") assert supports_cu_anthropic is True # Test a model known not to have the flag or set to false (defaults to False via get_model_info) @@ -1126,9 +1027,7 @@ def test_get_model_info_shows_supports_computer_use(): model_known_not_to_support_computer_use = "gpt-3.5-turbo" info_gpt = litellm.get_model_info(model_known_not_to_support_computer_use) print(f"Info for {model_known_not_to_support_computer_use}: {info_gpt}") - assert ( - info_gpt.get("supports_computer_use") is None - ) # Expecting None due to the default in ModelInfoBase + assert info_gpt.get("supports_computer_use") is None # Expecting None due to the default in ModelInfoBase @pytest.mark.parametrize( @@ -1244,20 +1143,14 @@ class TestProxyFunctionCalling: ("command-nightly", "litellm_proxy/command-nightly", False), ], ) - def test_proxy_function_calling_support_consistency( - self, direct_model, proxy_model, expected_result - ): + def test_proxy_function_calling_support_consistency(self, direct_model, proxy_model, expected_result): """Test that proxy models have the same function calling support as their direct counterparts.""" direct_result = supports_function_calling(direct_model) proxy_result = supports_function_calling(proxy_model) # Both should match the expected result - assert ( - direct_result == expected_result - ), f"Direct model {direct_model} should return {expected_result}" - assert ( - proxy_result == expected_result - ), f"Proxy model {proxy_model} should return {expected_result}" + assert direct_result == expected_result, f"Direct model {direct_model} should return {expected_result}" + assert proxy_result == expected_result, f"Proxy model {proxy_model} should return {expected_result}" # Direct and proxy should be consistent assert ( @@ -1329,9 +1222,7 @@ class TestProxyFunctionCalling: ("litellm_proxy/local-mistral", "ollama/mistral", False), ], ) - def test_proxy_custom_model_names_without_config( - self, proxy_model_name, underlying_model, expected_proxy_result - ): + def test_proxy_custom_model_names_without_config(self, proxy_model_name, underlying_model, expected_proxy_result): """ Test proxy models with custom model names that differ from underlying models. @@ -1342,9 +1233,7 @@ class TestProxyFunctionCalling: # Test the underlying model directly first to establish what it SHOULD return try: underlying_result = supports_function_calling(underlying_model) - print( - f"Underlying model {underlying_model} supports function calling: {underlying_result}" - ) + print(f"Underlying model {underlying_model} supports function calling: {underlying_result}") except Exception as e: print(f"Warning: Could not test underlying model {underlying_model}: {e}") @@ -1366,9 +1255,7 @@ class TestProxyFunctionCalling: # Case 1: Custom model name that cannot be resolved custom_model = "litellm_proxy/my-custom-claude" result = supports_function_calling(custom_model) - assert ( - result is False - ), "Custom model names return False without proxy config context" + assert result is False, "Custom model names return False without proxy config context" # Case 2: Model name that can be resolved (matches pattern) resolvable_model = "litellm_proxy/claude-sonnet-4-5-20250929" @@ -1407,9 +1294,7 @@ class TestProxyFunctionCalling: ), # Hints at Bedrock Claude 3 Sonnet ], ) - def test_proxy_models_with_naming_hints( - self, proxy_model_with_hints, expected_result - ): + def test_proxy_models_with_naming_hints(self, proxy_model_with_hints, expected_result): """ Test proxy models with names that provide hints about the underlying model. @@ -1421,14 +1306,10 @@ class TestProxyFunctionCalling: # Currently these will return False, but we document the expected behavior # In the future, we could implement smarter model name inference - print( - f"Model {proxy_model_with_hints}: current={proxy_result}, desired={expected_result}" - ) + print(f"Model {proxy_model_with_hints}: current={proxy_result}, desired={expected_result}") # For now, we expect False (current behavior), but document the limitation - assert ( - proxy_result is False - ), f"Current limitation: {proxy_model_with_hints} returns False without inference" + assert proxy_result is False, f"Current limitation: {proxy_model_with_hints} returns False without inference" @pytest.mark.parametrize( "proxy_model,expected_result", @@ -1453,9 +1334,7 @@ class TestProxyFunctionCalling: """ try: result = supports_function_calling(model=proxy_model) - assert ( - result == expected_result - ), f"Proxy model {proxy_model} returned {result}, expected {expected_result}" + assert result == expected_result, f"Proxy model {proxy_model} returned {result}, expected {expected_result}" except Exception as e: pytest.fail(f"Error testing proxy model {proxy_model}: {e}") @@ -1495,17 +1374,11 @@ class TestProxyFunctionCalling: parameter explicitly set to None, which is a common usage pattern. """ try: - result = supports_function_calling( - model=model_name, custom_llm_provider=None - ) + result = supports_function_calling(model=model_name, custom_llm_provider=None) # All the models in this test should support function calling - assert ( - result is True - ), f"Model {model_name} should support function calling but returned {result}" + assert result is True, f"Model {model_name} should support function calling but returned {result}" except Exception as e: - pytest.fail( - f"Error testing {model_name} with custom_llm_provider=None: {e}" - ) + pytest.fail(f"Error testing {model_name} with custom_llm_provider=None: {e}") def test_edge_cases_and_malformed_proxy_models(self): """Test edge cases and malformed proxy model names.""" @@ -1541,10 +1414,8 @@ class TestProxyFunctionCalling: direct_result = supports_function_calling(model=direct_model) proxy_result = supports_function_calling(model=proxy_model) - print(f"\nDemonstration of proxy model resolution:") - print( - f"Direct model '{direct_model}' supports function calling: {direct_result}" - ) + print("\nDemonstration of proxy model resolution:") + print(f"Direct model '{direct_model}' supports function calling: {direct_result}") print(f"Proxy model '{proxy_model}' supports function calling: {proxy_result}") # This assertion will currently fail due to the bug @@ -1557,8 +1428,7 @@ class TestProxyFunctionCalling: ) assert direct_result == proxy_result, ( - f"Proxy model resolution issue: {direct_model} -> {direct_result}, " - f"{proxy_model} -> {proxy_result}" + f"Proxy model resolution issue: {direct_model} -> {direct_result}, " f"{proxy_model} -> {proxy_result}" ) @pytest.mark.parametrize( @@ -1735,9 +1605,7 @@ class TestProxyFunctionCalling: underlying_result is True ), f"Claude 3 models should support function calling: {underlying_bedrock_model}" except Exception as e: - print( - f" Warning: Could not test underlying model {underlying_bedrock_model}: {e}" - ) + print(f" Warning: Could not test underlying model {underlying_bedrock_model}: {e}") # Test the proxy model - should return False due to lack of configuration context proxy_result = supports_function_calling(proxy_model_name) @@ -1824,9 +1692,7 @@ class TestProxyFunctionCalling: result = supports_function_calling(model) print(f"Direct test - {model}: {result}") # Claude 3 models should support function calling - assert ( - result is True - ), f"Claude 3 model should support function calling: {model}" + assert result is True, f"Claude 3 model should support function calling: {model}" except Exception as e: print(f"Could not test {model}: {e}") @@ -2004,9 +1870,7 @@ class TestProxyFunctionCalling: underlying_result is True ), f"Claude 3 models should support function calling: {underlying_bedrock_model}" except Exception as e: - print( - f" Warning: Could not test underlying model {underlying_bedrock_model}: {e}" - ) + print(f" Warning: Could not test underlying model {underlying_bedrock_model}: {e}") # Test the proxy model - should return False due to lack of configuration context proxy_result = supports_function_calling(proxy_model_name) @@ -2093,9 +1957,7 @@ class TestProxyFunctionCalling: result = supports_function_calling(model) print(f"Direct test - {model}: {result}") # Claude 3 models should support function calling - assert ( - result is True - ), f"Claude 3 model should support function calling: {model}" + assert result is True, f"Claude 3 model should support function calling: {model}" except Exception as e: print(f"Could not test {model}: {e}") @@ -2273,9 +2135,7 @@ class TestProxyFunctionCalling: underlying_result is True ), f"Claude 3 models should support function calling: {underlying_bedrock_model}" except Exception as e: - print( - f" Warning: Could not test underlying model {underlying_bedrock_model}: {e}" - ) + print(f" Warning: Could not test underlying model {underlying_bedrock_model}: {e}") # Test the proxy model - should return False due to lack of configuration context proxy_result = supports_function_calling(proxy_model_name) @@ -2362,9 +2222,7 @@ class TestProxyFunctionCalling: result = supports_function_calling(model) print(f"Direct test - {model}: {result}") # Claude 3 models should support function calling - assert ( - result is True - ), f"Claude 3 model should support function calling: {model}" + assert result is True, f"Claude 3 model should support function calling: {model}" except Exception as e: print(f"Could not test {model}: {e}") @@ -2377,13 +2235,14 @@ def test_register_model_with_scientific_notation(): # Use a truly unique model name with uuid to avoid conflicts when tests run in parallel test_model_name = f"test-scientific-notation-model-{uuid.uuid4().hex[:12]}" - + # Clear LRU caches that might have stale data from litellm.utils import ( _invalidate_model_cost_lowercase_map, ) + _invalidate_model_cost_lowercase_map() - + model_cost_dict = { test_model_name: { "max_tokens": 8192, @@ -2402,7 +2261,7 @@ def test_register_model_with_scientific_notation(): assert registered_model["output_cost_per_token"] == 6e-07 assert registered_model["litellm_provider"] == "openai" assert registered_model["mode"] == "chat" - + # Clean up after test if test_model_name in litellm.model_cost: del litellm.model_cost[test_model_name] @@ -2614,17 +2473,13 @@ def test_block_key_hashing_logic(): # Additional verification: if it should be hashed, verify it's actually a hash if should_be_hashed: # SHA-256 hashes are 64 characters long and contain only hex digits - assert ( - len(hashed_token) == 64 - ), f"Hash length should be 64, got {len(hashed_token)} for {input_key}" + assert len(hashed_token) == 64, f"Hash length should be 64, got {len(hashed_token)} for {input_key}" assert all( c in "0123456789abcdef" for c in hashed_token ), f"Hash should contain only hex digits for {input_key}" else: # If not hashed, it should be the original string - assert ( - hashed_token == input_key - ), f"Non-hashed key should remain unchanged: {input_key}" + assert hashed_token == input_key, f"Non-hashed key should remain unchanged: {input_key}" print("āœ… All block_key hashing logic tests passed!") @@ -2651,9 +2506,7 @@ def test_generate_gcp_iam_access_token(): mock_iam_credentials_v1.GenerateAccessTokenRequest = Mock() # Test successful token generation by mocking sys.modules - with patch.dict( - "sys.modules", {"google.cloud.iam_credentials_v1": mock_iam_credentials_v1} - ): + with patch.dict("sys.modules", {"google.cloud.iam_credentials_v1": mock_iam_credentials_v1}): from litellm._redis import _generate_gcp_iam_access_token result = _generate_gcp_iam_access_token(service_account) @@ -2692,15 +2545,111 @@ def test_generate_gcp_iam_access_token_import_error(): assert "pip install google-cloud-iam" in str(exc_info.value) +def test_generate_azure_ad_redis_token(): + """Test _generate_azure_ad_redis_token with mocked Azure credential.""" + from unittest.mock import Mock, patch + + expected_token = "azure-access-token-12345" + + mock_token = Mock() + mock_token.token = expected_token + + mock_credential = Mock() + mock_credential.get_token.return_value = mock_token + + with patch("azure.identity.DefaultAzureCredential", return_value=mock_credential): + from litellm._redis import _generate_azure_ad_redis_token + + result = _generate_azure_ad_redis_token() + + assert result == expected_token + mock_credential.get_token.assert_called_once_with("https://redis.azure.com/.default") + + +def test_generate_azure_ad_redis_token_service_principal(): + """Test _generate_azure_ad_redis_token with service principal credentials.""" + from unittest.mock import Mock, patch + + expected_token = "sp-access-token-67890" + + mock_token = Mock() + mock_token.token = expected_token + + mock_credential = Mock() + mock_credential.get_token.return_value = mock_token + + with patch("azure.identity.ClientSecretCredential", return_value=mock_credential) as mock_cls: + from litellm._redis import _generate_azure_ad_redis_token + + result = _generate_azure_ad_redis_token( + azure_client_id="test-client-id", + azure_tenant_id="test-tenant-id", + azure_client_secret="test-secret", + ) + + assert result == expected_token + mock_cls.assert_called_once_with( + client_id="test-client-id", + tenant_id="test-tenant-id", + client_secret="test-secret", + ) + + +def test_generate_azure_ad_redis_token_import_error(): + """Test that _generate_azure_ad_redis_token raises ImportError when azure-identity is missing.""" + from unittest.mock import patch + from litellm._redis import _generate_azure_ad_redis_token + + original_import = __builtins__["__import__"] + + def mock_import(name, *args, **kwargs): + if name == "azure.identity": + raise ImportError("No module named 'azure.identity'") + return original_import(name, *args, **kwargs) + + with patch("builtins.__import__", side_effect=mock_import): + with pytest.raises(ImportError) as exc_info: + _generate_azure_ad_redis_token() + + assert "azure-identity is required" in str(exc_info.value) + + +def test_redis_client_logic_azure_ad_auth(): + """Test that _get_redis_client_logic sets up Azure AD auth when REDIS_AZURE_AD_TOKEN=true.""" + import os + from unittest.mock import patch + + with patch.dict( + os.environ, + { + "REDIS_HOST": "myredis.redis.cache.windows.net", + "REDIS_PORT": "6380", + "REDIS_AZURE_AD_TOKEN": "true", + "REDIS_SSL": "true", + }, + clear=True, + ): + from litellm._redis import _get_redis_client_logic + + redis_kwargs = _get_redis_client_logic() + + # Should have redis_connect_func set + assert "redis_connect_func" in redis_kwargs + assert hasattr(redis_kwargs["redis_connect_func"], "_azure_redis_ad_token") + assert redis_kwargs["redis_connect_func"]._azure_redis_ad_token is True + + # Azure-specific kwargs should be removed + assert "azure_redis_ad_token" not in redis_kwargs + assert "azure_client_id" not in redis_kwargs + + if __name__ == "__main__": # Allow running this test file directly for debugging pytest.main([__file__, "-v"]) def test_model_info_for_vertex_ai_deepseek_model(): - model_info = litellm.get_model_info( - model="vertex_ai/deepseek-ai/deepseek-r1-0528-maas" - ) + model_info = litellm.get_model_info(model="vertex_ai/deepseek-ai/deepseek-r1-0528-maas") assert model_info is not None assert model_info["litellm_provider"] == "vertex_ai-deepseek_models" assert model_info["mode"] == "chat" @@ -2823,9 +2772,7 @@ class TestGetValidModelsWithCLI: ] } - with patch.object( - litellm.module_level_client, "get", return_value=mock_response - ) as mock_get: + with patch.object(litellm.module_level_client, "get", return_value=mock_response) as mock_get: # Test the exact pattern used in cli_token_usage.py result = litellm.get_valid_models( check_provider_endpoint=True, @@ -3063,15 +3010,15 @@ class TestProxyLoggingBudgetAlerts: proxy_logging.slack_alerting_instance.budget_alerts.assert_called_once_with( type=alert_type, user_info=user_info ) - proxy_logging.email_logging_instance.budget_alerts.assert_called_once_with( - type=alert_type, user_info=user_info - ) + proxy_logging.email_logging_instance.budget_alerts.assert_called_once_with(type=alert_type, user_info=user_info) - async def test_budget_alerts_soft_budget_with_alert_emails_bypasses_alerting_none(self): + async def test_budget_alerts_soft_budget_with_alert_emails_bypasses_alerting_none( + self, + ): """ Test that soft_budget alerts with alert_emails bypass the alerting=None check and send emails even when alerting is None. - + This tests the new logic that allows team-specific soft budget email alerts via metadata.soft_budget_alerting_emails to work even when global alerting is disabled. """ @@ -3107,7 +3054,9 @@ class TestProxyLoggingBudgetAlerts: type="soft_budget", user_info=user_info ) - async def test_budget_alerts_soft_budget_without_alert_emails_respects_alerting_none(self): + async def test_budget_alerts_soft_budget_without_alert_emails_respects_alerting_none( + self, + ): """ Test that soft_budget alerts WITHOUT alert_emails still respect alerting=None and do not send emails when alerting is None. @@ -3140,7 +3089,9 @@ class TestProxyLoggingBudgetAlerts: proxy_logging.slack_alerting_instance.budget_alerts.assert_not_called() proxy_logging.email_logging_instance.budget_alerts.assert_not_called() - async def test_budget_alerts_soft_budget_with_empty_alert_emails_respects_alerting_none(self): + async def test_budget_alerts_soft_budget_with_empty_alert_emails_respects_alerting_none( + self, + ): """ Test that soft_budget alerts with empty alert_emails list still respect alerting=None. """ @@ -3277,11 +3228,12 @@ def test_last_assistant_with_tool_calls_has_no_thinking_blocks_issue_18926(): {"role": "user", "content": "Build a feature"}, { "role": "assistant", - "thinking_blocks": [ - {"type": "thinking", "thinking": "Let me analyze the requirements..."} - ], + "thinking_blocks": [{"type": "thinking", "thinking": "Let me analyze the requirements..."}], "tool_calls": [ - {"id": "toolu_1", "function": {"name": "file_editor", "arguments": "{}"}} + { + "id": "toolu_1", + "function": {"name": "file_editor", "arguments": "{}"}, + } ], }, { @@ -3294,7 +3246,10 @@ def test_last_assistant_with_tool_calls_has_no_thinking_blocks_issue_18926(): # NO thinking_blocks - Claude sometimes doesn't include them "content": [{"type": "text", "text": "Let me explore more..."}], "tool_calls": [ - {"id": "toolu_2", "function": {"name": "file_editor", "arguments": "{}"}} + { + "id": "toolu_2", + "function": {"name": "file_editor", "arguments": "{}"}, + } ], }, ] @@ -3307,10 +3262,9 @@ def test_last_assistant_with_tool_calls_has_no_thinking_blocks_issue_18926(): # So we should NOT drop thinking - the combination tells us thinking is in use # The fix uses both checks: only drop if last has none AND no message has any - should_drop_thinking = ( - last_assistant_with_tool_calls_has_no_thinking_blocks(messages) - and not any_assistant_message_has_thinking_blocks(messages) - ) + should_drop_thinking = last_assistant_with_tool_calls_has_no_thinking_blocks( + messages + ) and not any_assistant_message_has_thinking_blocks(messages) assert should_drop_thinking is False @@ -3643,37 +3597,27 @@ class TestMetadataNoneHandling: def test_metadata_none_get_previous_models(self): """kwargs.get("metadata") or {} should return {} when metadata is None.""" kwargs = {"metadata": None} - previous_models = (kwargs.get("metadata") or {}).get( - "previous_models", None - ) + previous_models = (kwargs.get("metadata") or {}).get("previous_models", None) assert previous_models is None def test_metadata_none_model_group_check(self): """'model_group' in (kwargs.get("metadata") or {}) should not raise TypeError.""" kwargs = {"metadata": None} - _is_litellm_router_call = "model_group" in ( - kwargs.get("metadata") or {} - ) + _is_litellm_router_call = "model_group" in (kwargs.get("metadata") or {}) assert _is_litellm_router_call is False def test_metadata_missing_key(self): """Should work when metadata key is completely absent.""" kwargs = {} - previous_models = (kwargs.get("metadata") or {}).get( - "previous_models", None - ) + previous_models = (kwargs.get("metadata") or {}).get("previous_models", None) assert previous_models is None def test_metadata_present_with_values(self): """Should work when metadata has actual values.""" kwargs = {"metadata": {"previous_models": ["model1"], "model_group": "test"}} - previous_models = (kwargs.get("metadata") or {}).get( - "previous_models", None - ) + previous_models = (kwargs.get("metadata") or {}).get("previous_models", None) assert previous_models == ["model1"] - _is_litellm_router_call = "model_group" in ( - kwargs.get("metadata") or {} - ) + _is_litellm_router_call = "model_group" in (kwargs.get("metadata") or {}) assert _is_litellm_router_call is True def test_metadata_none_causes_error_with_old_pattern(self):