fix(azure): restrict the storage credential chain to deployment identities (#39637)

* fix(azure): restrict the storage credential chain to deployment identities

The keyless Azure Storage path walks the full DefaultAzureCredential chain, so a
proxy with no storage service principal authenticates as whichever identity the
host happens to carry: an operator's az login on a workstation, or the
AZURE_CLIENT_ID/AZURE_CLIENT_SECRET service principal set for Azure OpenAI.
Neither is the identity granted Storage Blob Data Contributor.

Narrow the chain to workload identity and managed identity, the two credentials
a deployment legitimately holds. Azure OpenAI, Postgres IAM auth and the other
callers of get_azure_ad_token_provider keep the full chain.

* test(azure): read the credential chain off the mock instead of an accumulator

* chore: drop a stray launch traceback committed at the repo root

* fix(azure): let the storage chain reach a system assigned managed identity

DefaultAzureCredential keeps one managed identity link and pins it to
AZURE_CLIENT_ID, so a host that sets that variable for Azure OpenAI and runs as
a system assigned identity never got asked for a storage token. Build the chain
from the three credentials a deployment can carry instead of subtracting the
ones it cannot.
This commit is contained in:
yucheng-berri 2026-09-03 18:29:32 -07:00 committed by GitHub
parent 4a537e2c19
commit 4e18c0f63a
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
5 changed files with 184 additions and 2 deletions

View file

@ -36,7 +36,7 @@ AZURE_STORAGE_TOKEN_SCOPE: Final = "https://storage.azure.com/.default"
def _cached_credential_chain_token_provider() -> Callable[[], str]:
return get_azure_ad_token_provider(
azure_scope=AZURE_STORAGE_TOKEN_SCOPE,
azure_credential=AzureCredentialType.DefaultAzureCredential,
azure_credential=AzureCredentialType.DeploymentIdentityCredential,
)

View file

@ -57,9 +57,11 @@ def get_azure_ad_token_provider(
from azure import identity
from azure.identity import (
CertificateCredential,
ChainedTokenCredential,
ClientSecretCredential,
DefaultAzureCredential,
ManagedIdentityCredential,
WorkloadIdentityCredential,
get_bearer_token_provider,
)
@ -101,6 +103,28 @@ def get_azure_ad_token_provider(
# DefaultAzureCredential doesn't require explicit environment variables
# It automatically discovers credentials from the environment (managed identity, CLI, etc.)
credential = DefaultAzureCredential()
elif cred == AzureCredentialType.DeploymentIdentityCredential:
# DefaultAzureCredential cannot express this: excluding its developer credentials still
# leaves one managed identity link, which AZURE_CLIENT_ID pins to a user assigned identity,
# so a host running as a system assigned identity never gets asked
workload_client_id: Final = os.environ.get("AZURE_CLIENT_ID")
workload_tenant_id: Final = os.environ.get("AZURE_TENANT_ID")
workload_token_file: Final = os.environ.get("AZURE_FEDERATED_TOKEN_FILE")
credential = ChainedTokenCredential(
*(
(
WorkloadIdentityCredential(
client_id=workload_client_id,
tenant_id=workload_tenant_id,
token_file_path=workload_token_file,
),
)
if workload_client_id and workload_tenant_id and workload_token_file
else ()
),
*((ManagedIdentityCredential(client_id=workload_client_id),) if workload_client_id else ()),
ManagedIdentityCredential(),
)
else:
cred_cls: Final = getattr(identity, cred)
credential = cred_cls()

View file

@ -6,3 +6,4 @@ class AzureCredentialType(str, Enum):
ManagedIdentityCredential = "ManagedIdentityCredential"
CertificateCredential = "CertificateCredential"
DefaultAzureCredential = "DefaultAzureCredential"
DeploymentIdentityCredential = "DeploymentIdentityCredential"

View file

@ -42,6 +42,7 @@ def workload_identity_env_vars(monkeypatch):
"AZURE_STORAGE_ENDPOINT_SUFFIX",
"AZURE_CLIENT_SECRET",
"AZURE_CREDENTIAL",
"AZURE_TOKEN_CREDENTIALS",
"AZURE_SCOPE",
):
monkeypatch.delenv(unset, raising=False)
@ -206,10 +207,28 @@ def test_default_chain_provider_is_storage_scoped_and_built_once_per_process():
assert first() == "chain-token"
mock_builder.assert_called_once_with(
azure_scope="https://storage.azure.com/.default",
azure_credential=AzureCredentialType.DefaultAzureCredential,
azure_credential=AzureCredentialType.DeploymentIdentityCredential,
)
def test_storage_chain_reaches_only_the_identities_a_deployment_carries(workload_identity_env_vars):
"""
The chain runs on a server, where a developer sign-in is a person and not the deployment, so
the storage token must come from workload identity or managed identity or from nothing
"""
_cached_credential_chain_token_provider.cache_clear()
with patch("azure.identity.get_bearer_token_provider", return_value=lambda: "chain-token") as bearer:
_cached_credential_chain_token_provider()
_cached_credential_chain_token_provider.cache_clear()
bearer.assert_called_once()
with bearer.call_args.args[0] as chain:
assert {type(link).__name__ for link in chain.credentials} == {
"WorkloadIdentityCredential",
"ManagedIdentityCredential",
}
@pytest.mark.asyncio
async def test_chain_tokens_are_read_from_the_provider_on_every_refresh(
workload_identity_env_vars,

View file

@ -6,6 +6,7 @@ from unittest.mock import MagicMock, patch
# Adds the grandparent directory to sys.path to allow importing project modules
import pytest
from azure.core.exceptions import ClientAuthenticationError
from litellm.secret_managers.get_azure_ad_token_provider import (
get_azure_ad_token_provider,
@ -16,6 +17,143 @@ from litellm.types.secret_managers.get_azure_ad_token_provider import (
)
class TestDeploymentIdentityCredential:
@staticmethod
def _chain_for(credential_type):
with patch("azure.identity.get_bearer_token_provider", return_value=lambda: "token") as bearer:
get_azure_ad_token_provider(
azure_scope="https://storage.azure.com/.default",
azure_credential=credential_type,
)
bearer.assert_called_once()
with bearer.call_args.args[0] as chain:
return {type(link).__name__ for link in chain.credentials}
@staticmethod
def _managed_identity_client_ids(credential_type):
with patch("azure.identity.get_bearer_token_provider", return_value=lambda: "token") as bearer:
get_azure_ad_token_provider(
azure_scope="https://storage.azure.com/.default",
azure_credential=credential_type,
)
with bearer.call_args.args[0] as chain:
return [
(link._credential._settings or {}).get("client_id")
for link in chain.credentials
if type(link).__name__ == "ManagedIdentityCredential"
]
@patch.dict(
os.environ,
{
"AZURE_CLIENT_ID": "workload-identity-client-id",
"AZURE_TENANT_ID": "workload-identity-tenant-id",
"AZURE_FEDERATED_TOKEN_FILE": "/var/run/secrets/azure/tokens/azure-identity-token",
},
clear=True,
)
def test_deployment_identity_reaches_workload_and_managed_identity_only(self):
assert self._chain_for(AzureCredentialType.DeploymentIdentityCredential) == {
"WorkloadIdentityCredential",
"ManagedIdentityCredential",
}
@patch.dict(
os.environ,
{
"AZURE_CLIENT_ID": "workload-identity-client-id",
"AZURE_TENANT_ID": "workload-identity-tenant-id",
"AZURE_FEDERATED_TOKEN_FILE": "/var/run/secrets/azure/tokens/azure-identity-token",
"AZURE_TOKEN_CREDENTIALS": "dev",
},
clear=True,
)
def test_deployment_identity_survives_a_developer_only_token_credentials_setting(self):
"""AZURE_TOKEN_CREDENTIALS=dev asks the SDK for developer credentials only, which is every
credential this chain drops, so the deployment's own identity has to win over it"""
assert self._chain_for(AzureCredentialType.DeploymentIdentityCredential) == {
"WorkloadIdentityCredential",
"ManagedIdentityCredential",
}
@patch.dict(
os.environ,
{
"AZURE_CLIENT_ID": "azure-openai-client-id",
"AZURE_CLIENT_SECRET": "azure-openai-client-secret",
"AZURE_TENANT_ID": "azure-openai-tenant-id",
},
clear=True,
)
def test_default_azure_credential_keeps_its_full_chain(self):
"""Azure OpenAI callers pass DefaultAzureCredential and must be unaffected by the
narrowing that the storage callback asks for"""
full_chain = self._chain_for(AzureCredentialType.DefaultAzureCredential)
assert "EnvironmentCredential" in full_chain
assert "AzureCliCredential" in full_chain
assert "EnvironmentCredential" not in self._chain_for(
AzureCredentialType.DeploymentIdentityCredential
)
@patch.dict(
os.environ,
{
"AZURE_CLIENT_ID": "azure-openai-client-id",
"AZURE_CLIENT_SECRET": "azure-openai-client-secret",
"AZURE_TENANT_ID": "azure-openai-tenant-id",
},
clear=True,
)
def test_deployment_identity_refuses_to_mint_a_token_for_a_configured_service_principal(self):
"""A host carrying only an Azure OpenAI client secret must get no token at all, and the
refusal must name the identities that were actually tried"""
provider = get_azure_ad_token_provider(
azure_scope="https://storage.azure.com/.default",
azure_credential=AzureCredentialType.DeploymentIdentityCredential,
)
with pytest.raises(ClientAuthenticationError) as refusal:
provider()
assert "ManagedIdentityCredential" in str(refusal.value)
assert "EnvironmentCredential" not in str(refusal.value)
assert "AzureCliCredential" not in str(refusal.value)
assert "azure-openai-client-secret" not in str(refusal.value)
@patch.dict(
os.environ,
{
"AZURE_CLIENT_ID": "azure-openai-client-id",
"AZURE_CLIENT_SECRET": "azure-openai-client-secret",
"AZURE_TENANT_ID": "azure-openai-tenant-id",
},
clear=True,
)
def test_deployment_identity_still_reaches_a_system_assigned_managed_identity(self):
"""AZURE_CLIENT_ID names one identity for the whole proxy, and pointing it at Azure OpenAI
must not hide the system assigned identity the host runs as"""
client_ids = self._managed_identity_client_ids(AzureCredentialType.DeploymentIdentityCredential)
assert "azure-openai-client-id" in client_ids
assert None in client_ids
@patch.dict(
os.environ,
{
"AZURE_CLIENT_ID": "user-assigned-identity-client-id",
"AZURE_TOKEN_CREDENTIALS": "dev",
},
clear=True,
)
def test_deployment_identity_keeps_the_user_assigned_identity_under_a_dev_only_setting(self):
"""AZURE_TOKEN_CREDENTIALS=dev asks the SDK for developer credentials only, and the
identity a host actually runs as has to survive that"""
assert "user-assigned-identity-client-id" in self._managed_identity_client_ids(
AzureCredentialType.DeploymentIdentityCredential
)
class TestGetAzureAdTokenProvider:
@patch.dict(
os.environ,