diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index e51e1f0b16b..3f4dd487705 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -126,7 +126,6 @@ from litellm.utils import ( if TYPE_CHECKING: from aiohttp import ClientSession - from azure.core.credentials import TokenCredential from opentelemetry.trace import Span as _Span from litellm.integrations.opentelemetry import OpenTelemetry @@ -2041,28 +2040,6 @@ def _resolve_pydantic_type(typ) -> List: return typs -def _get_azure_key_vault_credential() -> "TokenCredential": - """ - Build the credential used to authenticate to Azure Key Vault. - - When ``AZURE_KEY_VAULT_USE_WORKLOAD_IDENTITY`` is truthy, use - ``WorkloadIdentityCredential`` (federated token auth for AKS workload identity, - no client secret needed); it reads ``AZURE_TENANT_ID``, ``AZURE_CLIENT_ID`` and - ``AZURE_FEDERATED_TOKEN_FILE`` from the environment. Otherwise fall back to - ``DefaultAzureCredential``. - """ - from litellm.secret_managers.main import str_to_bool - - if str_to_bool(os.getenv("AZURE_KEY_VAULT_USE_WORKLOAD_IDENTITY")) is True: - from azure.identity import WorkloadIdentityCredential - - return WorkloadIdentityCredential() - - from azure.identity import DefaultAzureCredential - - return DefaultAzureCredential() - - def load_from_azure_key_vault(use_azure_key_vault: bool = False): if use_azure_key_vault is False: return @@ -2070,13 +2047,17 @@ def load_from_azure_key_vault(use_azure_key_vault: bool = False): try: from azure.keyvault.secrets import SecretClient + from litellm.secret_managers.get_azure_ad_token_provider import ( + get_azure_credential, + ) + # Set your Azure Key Vault URI KVUri = os.getenv("AZURE_KEY_VAULT_URI", None) if KVUri is None: raise Exception("Error when loading keys from Azure Key Vault: AZURE_KEY_VAULT_URI is not set.") - credential = _get_azure_key_vault_credential() + credential = get_azure_credential() # Create the SecretClient using the credential client = SecretClient(vault_url=KVUri, credential=credential) diff --git a/litellm/secret_managers/get_azure_ad_token_provider.py b/litellm/secret_managers/get_azure_ad_token_provider.py index 2c52054964d..f28a37faef2 100644 --- a/litellm/secret_managers/get_azure_ad_token_provider.py +++ b/litellm/secret_managers/get_azure_ad_token_provider.py @@ -1,11 +1,14 @@ import os -from typing import Any, Callable, Optional, Union +from typing import TYPE_CHECKING, Callable, Optional from litellm._logging import verbose_logger from litellm.types.secret_managers.get_azure_ad_token_provider import ( AzureCredentialType, ) +if TYPE_CHECKING: + from azure.core.credentials import TokenCredential + def infer_credential_type_from_environment() -> AzureCredentialType: if ( @@ -14,6 +17,12 @@ def infer_credential_type_from_environment() -> AzureCredentialType: and os.environ.get("AZURE_TENANT_ID") ): return AzureCredentialType.ClientSecretCredential + elif ( + os.environ.get("AZURE_CLIENT_ID") + and os.environ.get("AZURE_TENANT_ID") + and os.environ.get("AZURE_FEDERATED_TOKEN_FILE") + ): + return AzureCredentialType.WorkloadIdentityCredential elif os.environ.get("AZURE_CLIENT_ID"): return AzureCredentialType.ManagedIdentityCredential elif ( @@ -31,6 +40,74 @@ def infer_credential_type_from_environment() -> AzureCredentialType: return AzureCredentialType.DefaultAzureCredential +def get_azure_credential( + azure_credential: Optional[AzureCredentialType] = None, +) -> "TokenCredential": + """ + Build the Azure credential used for AD token acquisition and for authenticating + to Azure secret managers (e.g. Key Vault). + + The credential type is chosen from, in order: the explicit ``azure_credential`` + argument, the ``AZURE_CREDENTIAL`` env var, then inference from the environment. + ``WorkloadIdentityCredential`` supports AKS workload identity federation, reading + ``AZURE_CLIENT_ID``, ``AZURE_TENANT_ID`` and ``AZURE_FEDERATED_TOKEN_FILE``. + + See Also: + https://learn.microsoft.com/en-us/python/api/overview/azure/identity-readme?view=azure-python#service-principal-with-secret; + https://azure.github.io/azure-workload-identity/docs/quick-start.html. + """ + import azure.identity as identity + from azure.identity import ( + CertificateCredential, + ClientSecretCredential, + DefaultAzureCredential, + ManagedIdentityCredential, + WorkloadIdentityCredential, + ) + + cred: str = ( + azure_credential.value + if azure_credential + else None or os.environ.get("AZURE_CREDENTIAL") or infer_credential_type_from_environment() + ) + verbose_logger.info(f"For Azure credential, choosing credential type: {cred}") + + if cred == AzureCredentialType.ClientSecretCredential: + return ClientSecretCredential( + client_id=os.environ["AZURE_CLIENT_ID"], + client_secret=os.environ["AZURE_CLIENT_SECRET"], + tenant_id=os.environ["AZURE_TENANT_ID"], + ) + elif cred == AzureCredentialType.ManagedIdentityCredential: + return ManagedIdentityCredential(client_id=os.environ["AZURE_CLIENT_ID"]) + elif cred == AzureCredentialType.WorkloadIdentityCredential: + return WorkloadIdentityCredential( + client_id=os.environ["AZURE_CLIENT_ID"], + tenant_id=os.environ["AZURE_TENANT_ID"], + token_file_path=os.environ["AZURE_FEDERATED_TOKEN_FILE"], + ) + elif cred == AzureCredentialType.CertificateCredential: + if os.getenv("AZURE_CERTIFICATE_PASSWORD"): + return CertificateCredential( + client_id=os.environ["AZURE_CLIENT_ID"], + tenant_id=os.environ["AZURE_TENANT_ID"], + certificate_path=os.environ["AZURE_CERTIFICATE_PATH"], + password=os.environ["AZURE_CERTIFICATE_PASSWORD"], + ) + return CertificateCredential( + client_id=os.environ["AZURE_CLIENT_ID"], + tenant_id=os.environ["AZURE_TENANT_ID"], + certificate_path=os.environ["AZURE_CERTIFICATE_PATH"], + ) + elif cred == AzureCredentialType.DefaultAzureCredential: + # DefaultAzureCredential doesn't require explicit environment variables + # It automatically discovers credentials from the environment (managed identity, CLI, etc.) + return DefaultAzureCredential() + + cred_cls = getattr(identity, cred) + return cred_cls() + + def get_azure_ad_token_provider( azure_scope: Optional[str] = None, azure_credential: Optional[AzureCredentialType] = None, @@ -51,64 +128,11 @@ def get_azure_ad_token_provider( Returns: Callable that returns a temporary authentication token. """ - import azure.identity as identity - from azure.identity import ( - CertificateCredential, - ClientSecretCredential, - DefaultAzureCredential, - ManagedIdentityCredential, - get_bearer_token_provider, - ) + from azure.identity import get_bearer_token_provider if azure_scope is None: azure_scope = os.environ.get("AZURE_SCOPE") or "https://cognitiveservices.azure.com/.default" - cred: str = ( - azure_credential.value - if azure_credential - else None or os.environ.get("AZURE_CREDENTIAL") or infer_credential_type_from_environment() - ) - verbose_logger.info(f"For Azure AD Token Provider, choosing credential type: {cred}") - credential: Optional[ - Union[ - ClientSecretCredential, - ManagedIdentityCredential, - CertificateCredential, - DefaultAzureCredential, - Any, - ] - ] = None - if cred == AzureCredentialType.ClientSecretCredential: - credential = ClientSecretCredential( - client_id=os.environ["AZURE_CLIENT_ID"], - client_secret=os.environ["AZURE_CLIENT_SECRET"], - tenant_id=os.environ["AZURE_TENANT_ID"], - ) - elif cred == AzureCredentialType.ManagedIdentityCredential: - credential = ManagedIdentityCredential(client_id=os.environ["AZURE_CLIENT_ID"]) - elif cred == AzureCredentialType.CertificateCredential: - if os.getenv("AZURE_CERTIFICATE_PASSWORD"): - credential = CertificateCredential( - client_id=os.environ["AZURE_CLIENT_ID"], - tenant_id=os.environ["AZURE_TENANT_ID"], - certificate_path=os.environ["AZURE_CERTIFICATE_PATH"], - password=os.environ["AZURE_CERTIFICATE_PASSWORD"], - ) - else: - credential = CertificateCredential( - client_id=os.environ["AZURE_CLIENT_ID"], - tenant_id=os.environ["AZURE_TENANT_ID"], - certificate_path=os.environ["AZURE_CERTIFICATE_PATH"], - ) - elif cred == AzureCredentialType.DefaultAzureCredential: - # DefaultAzureCredential doesn't require explicit environment variables - # It automatically discovers credentials from the environment (managed identity, CLI, etc.) - credential = DefaultAzureCredential() - else: - cred_cls = getattr(identity, cred) - credential = cred_cls() - - if credential is None: - raise ValueError("No credential provided") + credential = get_azure_credential(azure_credential) return get_bearer_token_provider(credential, azure_scope) diff --git a/litellm/types/secret_managers/get_azure_ad_token_provider.py b/litellm/types/secret_managers/get_azure_ad_token_provider.py index 5d2f7409f95..04c2663089a 100644 --- a/litellm/types/secret_managers/get_azure_ad_token_provider.py +++ b/litellm/types/secret_managers/get_azure_ad_token_provider.py @@ -5,4 +5,5 @@ class AzureCredentialType(str, Enum): ClientSecretCredential = "ClientSecretCredential" ManagedIdentityCredential = "ManagedIdentityCredential" CertificateCredential = "CertificateCredential" + WorkloadIdentityCredential = "WorkloadIdentityCredential" DefaultAzureCredential = "DefaultAzureCredential" diff --git a/tests/test_litellm/proxy/proxy_server/test_lifecycle.py b/tests/test_litellm/proxy/proxy_server/test_lifecycle.py index e75a6c435d5..40b1ca49759 100644 --- a/tests/test_litellm/proxy/proxy_server/test_lifecycle.py +++ b/tests/test_litellm/proxy/proxy_server/test_lifecycle.py @@ -417,58 +417,45 @@ def test_load_from_azure_key_vault_missing_uri_failure_is_swallowed(monkeypatch) assert result is None -def _install_fake_azure_identity(monkeypatch): - """Install a fake ``azure.identity`` so credential selection can be exercised - without the real Azure SDK installed. Returns the two sentinel classes.""" - import sys - import types +def test_load_from_azure_key_vault_uses_shared_credential_builder(monkeypatch): + """The Key Vault loader must build its credential via the shared + ``get_azure_credential`` helper (so workload identity and other credential + types are honored) rather than hard-coding ``DefaultAzureCredential``.""" + import litellm + import litellm.secret_managers.get_azure_ad_token_provider as azure_cred_mod - class FakeDefaultAzureCredential: - pass + monkeypatch.setenv("AZURE_KEY_VAULT_URI", "https://example.vault.azure.net/") + monkeypatch.setattr(litellm, "secret_manager_client", None, raising=False) - class FakeWorkloadIdentityCredential: - pass + sentinel_credential = object() + sentinel_client = object() + observed = {} - azure_module = types.ModuleType("azure") - identity_module = types.ModuleType("azure.identity") - identity_module.DefaultAzureCredential = FakeDefaultAzureCredential - identity_module.WorkloadIdentityCredential = FakeWorkloadIdentityCredential - azure_module.identity = identity_module + def fake_get_azure_credential(): + observed["credential_built"] = True + return sentinel_credential - monkeypatch.setitem(sys.modules, "azure", azure_module) - monkeypatch.setitem(sys.modules, "azure.identity", identity_module) - return FakeDefaultAzureCredential, FakeWorkloadIdentityCredential + fake_secret_client_module = MagicMock() + def fake_secret_client(vault_url, credential): + observed["vault_url"] = vault_url + observed["credential"] = credential + return sentinel_client -def test_get_azure_key_vault_credential_defaults_to_default_credential(monkeypatch): - default_cls, workload_cls = _install_fake_azure_identity(monkeypatch) - monkeypatch.delenv("AZURE_KEY_VAULT_USE_WORKLOAD_IDENTITY", raising=False) + fake_secret_client_module.SecretClient = fake_secret_client + monkeypatch.setattr(azure_cred_mod, "get_azure_credential", fake_get_azure_credential) + monkeypatch.setitem( + __import__("sys").modules, + "azure.keyvault.secrets", + fake_secret_client_module, + ) - credential = ps._get_azure_key_vault_credential() + load_from_azure_key_vault(use_azure_key_vault=True) - assert isinstance(credential, default_cls) - assert not isinstance(credential, workload_cls) - - -@pytest.mark.parametrize("truthy", ["true", "True", "TRUE"]) -def test_get_azure_key_vault_credential_uses_workload_identity_when_enabled(monkeypatch, truthy): - default_cls, workload_cls = _install_fake_azure_identity(monkeypatch) - monkeypatch.setenv("AZURE_KEY_VAULT_USE_WORKLOAD_IDENTITY", truthy) - - credential = ps._get_azure_key_vault_credential() - - assert isinstance(credential, workload_cls) - assert not isinstance(credential, default_cls) - - -@pytest.mark.parametrize("falsy", ["false", "False", "0", "unset-like"]) -def test_get_azure_key_vault_credential_uses_default_when_disabled(monkeypatch, falsy): - default_cls, workload_cls = _install_fake_azure_identity(monkeypatch) - monkeypatch.setenv("AZURE_KEY_VAULT_USE_WORKLOAD_IDENTITY", falsy) - - credential = ps._get_azure_key_vault_credential() - - assert isinstance(credential, default_cls) + assert observed["credential_built"] is True + assert observed["credential"] is sentinel_credential + assert observed["vault_url"] == "https://example.vault.azure.net/" + assert litellm.secret_manager_client is sentinel_client # --------------------------------------------------------------------------- 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..c178c92598a 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,11 @@ import pytest from litellm.secret_managers.get_azure_ad_token_provider import ( get_azure_ad_token_provider, + get_azure_credential, + infer_credential_type_from_environment, +) +from litellm.types.secret_managers.get_azure_ad_token_provider import ( + AzureCredentialType, ) @@ -243,3 +248,131 @@ class TestGetAzureAdTokenProvider: # Test that the returned callable works token = result() assert token == "mock-default-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_SCOPE": "https://cognitiveservices.azure.com/.default", + "AZURE_CREDENTIAL": "WorkloadIdentityCredential", + }, + ) + @patch("azure.identity.get_bearer_token_provider") + @patch("azure.identity.WorkloadIdentityCredential") + def test_get_azure_ad_token_provider_workload_identity_credential( + self, mock_workload_identity_credential, mock_get_bearer_token_provider + ): + """Test get_azure_ad_token_provider with WorkloadIdentityCredential (AKS federation).""" + mock_credential_instance = MagicMock() + mock_workload_identity_credential.return_value = mock_credential_instance + + mock_token_provider = MagicMock(return_value="mock-workload-identity-token") + mock_get_bearer_token_provider.return_value = mock_token_provider + + result = get_azure_ad_token_provider() + + assert callable(result) + mock_workload_identity_credential.assert_called_once_with( + client_id="test-client-id", + tenant_id="test-tenant-id", + token_file_path="/var/run/secrets/azure/tokens/azure-identity-token", + ) + mock_get_bearer_token_provider.assert_called_once_with( + mock_credential_instance, "https://cognitiveservices.azure.com/.default" + ) + + token = result() + assert token == "mock-workload-identity-token" + + +class TestInferCredentialTypeFromEnvironment: + @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", + }, + clear=True, + ) + def test_infers_workload_identity_from_federated_token_file(self): + """A federated token file plus client/tenant id (and no client secret) implies + workload identity, not managed identity.""" + assert ( + infer_credential_type_from_environment() + == AzureCredentialType.WorkloadIdentityCredential + ) + + @patch.dict( + os.environ, + { + "AZURE_CLIENT_ID": "test-client-id", + "AZURE_CLIENT_SECRET": "test-client-secret", + "AZURE_TENANT_ID": "test-tenant-id", + "AZURE_FEDERATED_TOKEN_FILE": "/var/run/secrets/azure/tokens/azure-identity-token", + }, + clear=True, + ) + def test_client_secret_takes_precedence_over_workload_identity(self): + """When a client secret is present, keep using the service-principal flow.""" + assert ( + infer_credential_type_from_environment() + == AzureCredentialType.ClientSecretCredential + ) + + @patch.dict( + os.environ, + {"AZURE_CLIENT_ID": "test-client-id"}, + clear=True, + ) + def test_client_id_only_infers_managed_identity(self): + assert ( + infer_credential_type_from_environment() + == AzureCredentialType.ManagedIdentityCredential + ) + + +class TestGetAzureCredential: + @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", + }, + clear=True, + ) + @patch("azure.identity.WorkloadIdentityCredential") + def test_returns_workload_identity_credential_object( + self, mock_workload_identity_credential + ): + """get_azure_credential returns the credential object (used by the Key Vault + loader), inferring workload identity from the environment.""" + mock_credential_instance = MagicMock() + mock_workload_identity_credential.return_value = mock_credential_instance + + credential = get_azure_credential() + + assert credential is mock_credential_instance + mock_workload_identity_credential.assert_called_once_with( + client_id="test-client-id", + tenant_id="test-tenant-id", + token_file_path="/var/run/secrets/azure/tokens/azure-identity-token", + ) + + @patch.dict(os.environ, {}, clear=True) + @patch("azure.identity.DefaultAzureCredential") + def test_explicit_argument_overrides_environment( + self, mock_default_azure_credential + ): + mock_credential_instance = MagicMock() + mock_default_azure_credential.return_value = mock_credential_instance + + credential = get_azure_credential( + azure_credential=AzureCredentialType.DefaultAzureCredential + ) + + assert credential is mock_credential_instance + mock_default_azure_credential.assert_called_once_with()