diff --git a/litellm/llms/azure/azure.py b/litellm/llms/azure/azure.py index 980b27cda55..38516fbae84 100644 --- a/litellm/llms/azure/azure.py +++ b/litellm/llms/azure/azure.py @@ -1,7 +1,7 @@ import asyncio import json import time -from collections.abc import Callable, Coroutine +from collections.abc import Callable, Coroutine, Mapping, MutableMapping from typing import Any, Final import httpx @@ -1123,6 +1123,40 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM): return base_url_with_deployment + @staticmethod + def _set_azure_ad_auth_header( + headers: MutableMapping[str, str], # mutable-ok: the caller's header dict is updated in place + api_key: str | None, + azure_ad_token: str | None, + azure_ad_token_provider: Callable | None, + azure_client_params: Mapping[str, Any], + ) -> None: + """Put an Entra ID bearer token on the image-generation request headers. + + Image generation does not send its request through the Azure SDK client, + so the credential ``initialize_azure_sdk_client`` resolved never reaches + the wire on its own — only an ``api-key`` header set by the caller does. + A deployment with no API key therefore sends no credential at all and + Azure answers ``401 Access denied due to invalid subscription key``, on + the same account and the same identity that serve chat, embeddings and + transcription key-less. + + ``azure_client_params`` is the fallback source because it is where the + managed-identity / DefaultAzureCredential provider ends up when it comes + from ``litellm.enable_azure_ad_token_refresh`` rather than from explicit + service-principal params. + """ + if api_key is not None or "Authorization" in headers: + return + + provider: Final = azure_ad_token_provider or azure_client_params.get("azure_ad_token_provider") + static_token: Final = azure_ad_token or azure_client_params.get("azure_ad_token") + token: Final = provider() if provider is not None and callable(provider) else static_token + + if token: + headers.pop("api-key", None) + headers["Authorization"] = f"Bearer {token}" + async def aimage_generation( self, data: dict, @@ -1249,12 +1283,6 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM): if not isinstance(max_retries, int): raise AzureOpenAIError(status_code=422, message="max retries must be an int") - if api_key is None and azure_ad_token_provider is not None: - azure_ad_token = azure_ad_token_provider() - if azure_ad_token: - headers.pop("api-key", None) - headers["Authorization"] = f"Bearer {azure_ad_token}" - # init AzureOpenAI Client azure_client_params: Final[dict[str, Any]] = self.initialize_azure_sdk_client( litellm_params=litellm_params or {}, @@ -1264,6 +1292,14 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM): api_base=api_base, is_async=False, ) + + self._set_azure_ad_auth_header( + headers=headers, + api_key=api_key, + azure_ad_token=azure_ad_token, + azure_ad_token_provider=azure_ad_token_provider, + azure_client_params=azure_client_params, + ) if aimg_generation is True: return self.aimage_generation( data=data, diff --git a/tests/test_litellm/llms/azure/image_generation/test_azure_image_generation_ad_token.py b/tests/test_litellm/llms/azure/image_generation/test_azure_image_generation_ad_token.py new file mode 100644 index 00000000000..74b49cdb897 --- /dev/null +++ b/tests/test_litellm/llms/azure/image_generation/test_azure_image_generation_ad_token.py @@ -0,0 +1,148 @@ +"""Image generation must send the Entra ID token when the deployment has no API key. + +Image generation builds its own httpx request instead of going through the +Azure SDK client, so a credential resolved by ``initialize_azure_sdk_client`` +only reaches the wire if it is put on the headers. Without that, a key-less +deployment sends no credential at all and Azure answers +``401 Access denied due to invalid subscription key`` — on the same account +and identity that serve chat, embeddings and transcription key-less. +""" + +from unittest.mock import MagicMock + +import pytest + +from litellm.llms.azure.azure import AzureChatCompletion + + +def _set_header(**kwargs): + headers = kwargs.pop("headers") + AzureChatCompletion._set_azure_ad_auth_header(headers=headers, **kwargs) + return headers + + +class TestSetAzureADAuthHeader: + def test_token_provider_from_client_params_becomes_bearer_header(self): + """The managed-identity provider lands in azure_client_params, not in the + caller's argument, when it comes from enable_azure_ad_token_refresh.""" + headers = _set_header( + headers={"Content-Type": "application/json"}, + api_key=None, + azure_ad_token=None, + azure_ad_token_provider=None, + azure_client_params={"azure_ad_token_provider": lambda: "mi-token"}, + ) + assert headers["Authorization"] == "Bearer mi-token" + + def test_token_from_client_params_becomes_bearer_header(self): + headers = _set_header( + headers={}, + api_key=None, + azure_ad_token=None, + azure_ad_token_provider=None, + azure_client_params={"azure_ad_token": "static-token"}, + ) + assert headers["Authorization"] == "Bearer static-token" + + def test_explicit_provider_wins_over_client_params(self): + headers = _set_header( + headers={}, + api_key=None, + azure_ad_token=None, + azure_ad_token_provider=lambda: "explicit-token", + azure_client_params={"azure_ad_token_provider": lambda: "fallback-token"}, + ) + assert headers["Authorization"] == "Bearer explicit-token" + + def test_stale_api_key_header_is_replaced_by_the_token(self): + headers = _set_header( + headers={"api-key": ""}, + api_key=None, + azure_ad_token=None, + azure_ad_token_provider=lambda: "mi-token", + azure_client_params={}, + ) + assert "api-key" not in headers + assert headers["Authorization"] == "Bearer mi-token" + + def test_api_key_deployment_is_left_alone(self): + """A keyed deployment must keep its api-key header and gain no bearer.""" + headers = _set_header( + headers={"api-key": "sk-azure"}, + api_key="sk-azure", + azure_ad_token=None, + azure_ad_token_provider=lambda: "mi-token", + azure_client_params={"azure_ad_token_provider": lambda: "mi-token"}, + ) + assert headers["api-key"] == "sk-azure" + assert "Authorization" not in headers + + def test_caller_supplied_authorization_is_not_overwritten(self): + headers = _set_header( + headers={"Authorization": "Bearer caller-token"}, + api_key=None, + azure_ad_token=None, + azure_ad_token_provider=lambda: "mi-token", + azure_client_params={}, + ) + assert headers["Authorization"] == "Bearer caller-token" + + def test_no_credential_anywhere_leaves_headers_unchanged(self): + headers = _set_header( + headers={"Content-Type": "application/json"}, + api_key=None, + azure_ad_token=None, + azure_ad_token_provider=None, + azure_client_params={}, + ) + assert headers == {"Content-Type": "application/json"} + + +@pytest.mark.parametrize("provider_key", ["azure_ad_token_provider", "azure_ad_token"]) +def test_image_generation_request_carries_the_bearer_token(monkeypatch, provider_key): + """End to end through image_generation: the outbound request is authenticated.""" + import httpx + + from litellm.types.utils import ImageResponse + + azure = AzureChatCompletion() + captured: dict = {} + + value = (lambda: "mi-token") if provider_key == "azure_ad_token_provider" else "mi-token" + + monkeypatch.setattr( + AzureChatCompletion, + "initialize_azure_sdk_client", + lambda self, **kwargs: { + "azure_endpoint": "https://acct.openai.azure.com", + "api_version": "2025-04-01-preview", + provider_key: value, + }, + ) + + def fake_request(self, *, headers, **kwargs): + captured["headers"] = headers + return httpx.Response( + status_code=200, + json={"created": 1, "data": [{"b64_json": "aGk="}]}, + request=httpx.Request("POST", "https://acct.openai.azure.com"), + ) + + monkeypatch.setattr(AzureChatCompletion, "make_sync_azure_httpx_request", fake_request) + + azure.image_generation( + prompt="a cat", + timeout=60.0, + optional_params={}, + logging_obj=MagicMock(), + headers={"Content-Type": "application/json"}, + model="gpt-image-1", + api_key=None, + api_base="https://acct.openai.azure.com", + api_version="2025-04-01-preview", + model_response=ImageResponse(), + litellm_params={}, + ) + + assert captured["headers"]["Authorization"] == "Bearer mi-token" + assert "api-key" not in captured["headers"]