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"