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 85e55a5c30d..f02f59cccc0 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 @@ -214,3 +214,32 @@ class TestGetAzureAdTokenProvider: # Test that the returned callable works token = result() assert token == "mock-certificate-token" + + @patch.dict(os.environ, {}, clear=True) # Clear all environment variables + @patch("azure.identity.get_bearer_token_provider") + @patch("azure.identity.DefaultAzureCredential") + def test_get_azure_ad_token_provider_defaults_to_default_azure_credential( + self, mock_default_azure_credential, mock_get_bearer_token_provider + ): + """Test get_azure_ad_token_provider defaults to DefaultAzureCredential when no credentials are present.""" + # Mock the Azure identity credential instance + mock_credential_instance = MagicMock() + mock_default_azure_credential.return_value = mock_credential_instance + + # Mock the bearer token provider + mock_token_provider = MagicMock(return_value="mock-default-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_default_azure_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-default-token"