From 712e042aa4e6b01ed21b1392a326fe3d64660288 Mon Sep 17 00:00:00 2001 From: Shuai Zhang Date: Sat, 31 May 2025 01:03:08 -0700 Subject: [PATCH] fix(secret-managers): Break AzureCredentialType restriction on AZURE_CREDENTIAL (#11272) --- .../get_azure_ad_token_provider.py | 4 +-- .../test_get_azure_ad_token_provider.py | 34 +++++++++++++++++++ 2 files changed, 35 insertions(+), 3 deletions(-) diff --git a/litellm/secret_managers/get_azure_ad_token_provider.py b/litellm/secret_managers/get_azure_ad_token_provider.py index d546745a0e..d46330443b 100644 --- a/litellm/secret_managers/get_azure_ad_token_provider.py +++ b/litellm/secret_managers/get_azure_ad_token_provider.py @@ -34,9 +34,7 @@ def get_azure_ad_token_provider(azure_scope: Optional[str] = None) -> Callable[[ if azure_scope is None: azure_scope = os.environ.get("AZURE_SCOPE", "https://cognitiveservices.azure.com/.default") - cred: Union[AzureCredentialType, str] = AzureCredentialType( - os.environ.get("AZURE_CREDENTIAL", AzureCredentialType.ClientSecretCredential) - ) + cred: str = os.environ.get("AZURE_CREDENTIAL", AzureCredentialType.ClientSecretCredential) credential: Optional[ Union[ ClientSecretCredential, diff --git a/tests/test_litellm/secret_managers/test_get_azure_ad_token_provider.py b/tests/test_litellm/secret_managers/test_get_azure_ad_token_provider.py index 7fd427dae4..44d745afa6 100644 --- a/tests/test_litellm/secret_managers/test_get_azure_ad_token_provider.py +++ b/tests/test_litellm/secret_managers/test_get_azure_ad_token_provider.py @@ -136,3 +136,37 @@ class TestGetAzureAdTokenProvider: # Test that the returned callable works token = result() assert token == "mock-certificate-token" + + @patch.dict( + os.environ, + { + "AZURE_CREDENTIAL": "DefaultAzureCredential", + }, + ) + @patch("azure.identity.get_bearer_token_provider") + @patch("azure.identity.DefaultAzureCredential") + def test_get_azure_ad_token_provider_default_azure_credential( + self, mock_certificate_credential, mock_get_bearer_token_provider + ): + """Test get_azure_ad_token_provider with DefaultAzureCredential.""" + # Mock the Azure identity credential instance + mock_credential_instance = MagicMock() + mock_certificate_credential.return_value = mock_credential_instance + + # Mock the bearer token provider + mock_token_provider = MagicMock(return_value="mock-certificate-token") + mock_get_bearer_token_provider.return_value = mock_token_provider + + # Call the function + result = get_azure_ad_token_provider() + + # Assertions + assert callable(result) + mock_certificate_credential.assert_called_once_with() + mock_get_bearer_token_provider.assert_called_once_with( + mock_credential_instance, "https://cognitiveservices.azure.com/.default" + ) + + # Test that the returned callable works + token = result() + assert token == "mock-certificate-token"