diff --git a/litellm/secret_managers/get_azure_ad_token_provider.py b/litellm/secret_managers/get_azure_ad_token_provider.py index d7f83855d2d..c2dc09bc65d 100644 --- a/litellm/secret_managers/get_azure_ad_token_provider.py +++ b/litellm/secret_managers/get_azure_ad_token_provider.py @@ -15,6 +15,8 @@ def infer_credential_type_from_environment() -> AzureCredentialType: and os.environ.get("AZURE_TENANT_ID") ): return AzureCredentialType.ClientSecretCredential + elif os.environ.get("AZURE_FEDERATED_TOKEN_FILE"): + return AzureCredentialType.DefaultAzureCredential elif os.environ.get("AZURE_CLIENT_ID"): return AzureCredentialType.ManagedIdentityCredential elif ( 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 f02f59cccc0..cee0da79802 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 @@ -11,6 +11,10 @@ import pytest from litellm.secret_managers.get_azure_ad_token_provider import ( get_azure_ad_token_provider, + infer_credential_type_from_environment, +) +from litellm.types.secret_managers.get_azure_ad_token_provider import ( + AzureCredentialType, ) @@ -215,6 +219,46 @@ class TestGetAzureAdTokenProvider: token = result() assert token == "mock-certificate-token" + @patch.dict( + os.environ, + { + "AZURE_CLIENT_ID": "test-client-id", + "AZURE_TENANT_ID": "test-tenant-id", + "AZURE_FEDERATED_TOKEN_FILE": "/var/run/secrets/azure/tokens/azure-identity-token", + "AZURE_AUTHORITY_HOST": "https://login.microsoftonline.com/", + }, + clear=True, + ) + @patch("azure.identity.get_bearer_token_provider") + @patch("azure.identity.ManagedIdentityCredential") + @patch("azure.identity.DefaultAzureCredential") + def test_get_azure_ad_token_provider_prefers_workload_identity_over_managed_identity( + self, + mock_default_azure_credential, + mock_managed_identity_credential, + mock_get_bearer_token_provider, + ): + """The AKS workload identity webhook injects AZURE_CLIENT_ID, AZURE_TENANT_ID, and + AZURE_FEDERATED_TOKEN_FILE, and never a client secret. Reading the bare client id as a + managed identity sends the pod to IMDS, which has no identity attached to it, so every + token request fails and the federated token is never exchanged. Only + DefaultAzureCredential's chain reaches WorkloadIdentityCredential.""" + mock_credential_instance = MagicMock() + mock_default_azure_credential.return_value = mock_credential_instance + mock_get_bearer_token_provider.return_value = MagicMock( + return_value="mock-workload-identity-token" + ) + + result = get_azure_ad_token_provider() + + assert ( + infer_credential_type_from_environment() + == AzureCredentialType.DefaultAzureCredential + ) + mock_managed_identity_credential.assert_not_called() + mock_default_azure_credential.assert_called_once_with() + assert result() == "mock-workload-identity-token" + @patch.dict(os.environ, {}, clear=True) # Clear all environment variables @patch("azure.identity.get_bearer_token_provider") @patch("azure.identity.DefaultAzureCredential")