refactor(secret_managers): reuse shared Azure credential builder for Key Vault + add workload identity

Co-Authored-By: Ishaan Jaffer <155045088+ishaan-berri@users.noreply.github.com>
This commit is contained in:
Devin AI 2026-07-16 14:26:35 +00:00
parent 456b97e635
commit 3f4f88372b
5 changed files with 250 additions and 124 deletions

View file

@ -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)

View file

@ -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)

View file

@ -5,4 +5,5 @@ class AzureCredentialType(str, Enum):
ClientSecretCredential = "ClientSecretCredential"
ManagedIdentityCredential = "ManagedIdentityCredential"
CertificateCredential = "CertificateCredential"
WorkloadIdentityCredential = "WorkloadIdentityCredential"
DefaultAzureCredential = "DefaultAzureCredential"

View file

@ -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
# ---------------------------------------------------------------------------

View file

@ -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()