mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-14 23:21:35 +00:00
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:
parent
456b97e635
commit
3f4f88372b
5 changed files with 250 additions and 124 deletions
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -5,4 +5,5 @@ class AzureCredentialType(str, Enum):
|
|||
ClientSecretCredential = "ClientSecretCredential"
|
||||
ManagedIdentityCredential = "ManagedIdentityCredential"
|
||||
CertificateCredential = "CertificateCredential"
|
||||
WorkloadIdentityCredential = "WorkloadIdentityCredential"
|
||||
DefaultAzureCredential = "DefaultAzureCredential"
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue