fix: req changes

This commit is contained in:
Harshit28j 2026-03-07 12:38:31 +05:30
parent 8bdbda0d2c
commit 8c1273fca9
2 changed files with 22 additions and 35 deletions

View file

@ -397,24 +397,18 @@ def _get_redis_client_logic(**env_overrides): # noqa: PLR0915
"REDIS_AZURE_AD_TOKEN"
)
if (
_azure_ad_enabled = (
_azure_redis_ad_token is not None
and str(_azure_redis_ad_token).lower() == "true"
and _gcp_service_account is not None
):
)
if _azure_ad_enabled and _gcp_service_account is not None:
verbose_logger.warning(
"Both GCP IAM (gcp_service_account) and Azure AD (azure_redis_ad_token) are configured for Redis. "
"Using GCP IAM. Remove one to avoid misconfiguration."
)
# Clean up Azure-specific kwargs even though we're not using Azure AD
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)
elif (
_azure_redis_ad_token is not None
and str(_azure_redis_ad_token).lower() == "true"
):
if _azure_ad_enabled and _gcp_service_account is None:
_azure_client_id = redis_kwargs.get("azure_client_id") or get_secret_str(
"AZURE_CLIENT_ID"
)
@ -437,11 +431,11 @@ def _get_redis_client_logic(**env_overrides): # noqa: PLR0915
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)
# Always 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)
@ -606,6 +600,9 @@ def get_redis_async_client( # noqa: PLR0915
)
# If GCP IAM is configured (indicated by redis_connect_func), generate access token and use as password
# NOTE: For async clusters, the token is set once at client creation. If the token expires
# (typically ~1 hour), the client will need to be recreated. The sync path handles this
# automatically via redis_connect_func which refreshes on each connection.
if redis_connect_func and gcp_service_account:
verbose_logger.debug(
"DEBUG: Generating IAM token for service account (value not logged for security reasons)"
@ -623,6 +620,8 @@ def get_redis_async_client( # noqa: PLR0915
raise AuthenticationError("Failed to generate GCP IAM access token")
# Handle Azure AD authentication for async clusters
# NOTE: Same token expiry caveat as GCP IAM above — token is static for the
# lifetime of this cluster client.
elif redis_connect_func and hasattr(
redis_connect_func, "_azure_redis_ad_token"
):

View file

@ -2703,18 +2703,11 @@ def test_generate_gcp_iam_access_token_import_error():
"""
Test that _generate_gcp_iam_access_token raises ImportError when google-cloud-iam is not available.
"""
# Import the function first, before mocking
from litellm._redis import _generate_gcp_iam_access_token
# Mock the import to fail when the function tries to import google.cloud.iam_credentials_v1
original_import = __builtins__["__import__"]
def mock_import(name, *args, **kwargs):
if name == "google.cloud.iam_credentials_v1":
raise ImportError("No module named 'google.cloud.iam_credentials_v1'")
return original_import(name, *args, **kwargs)
with patch("builtins.__import__", side_effect=mock_import):
# Setting module to None in sys.modules forces ImportError on next import attempt,
# even if the module was previously cached
with patch.dict("sys.modules", {"google.cloud.iam_credentials_v1": None}):
with pytest.raises(ImportError) as exc_info:
_generate_gcp_iam_access_token("test-service-account")
@ -2795,14 +2788,9 @@ def test_generate_azure_ad_redis_token_import_error():
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):
# Setting module to None in sys.modules forces ImportError on next import attempt,
# even if the module was previously cached
with patch.dict("sys.modules", {"azure.identity": None}):
with pytest.raises(ImportError) as exc_info:
_generate_azure_ad_redis_token()