fix(azure): keep azure_ad_token out of image edit and generation request bodies

image_edit flattened every non-LiteLLM-owned kwarg into the multipart body,
and azure_ad_token is not registered as owned, so Azure rejected edits with
"Unknown parameter: 'azure_ad_token'". The multipart path now skips any
GenericLiteLLMParams field.

image_generation popped azure_ad_token from the top level of optional_params,
but get_optional_params_image_gen nests unknown kwargs under extra_body, so the
token was still sent in the JSON body. Pop it from extra_body, matching the
Azure chat path.
This commit is contained in:
shrey-berri 2026-09-27 00:08:51 +00:00
parent 849f3037b4
commit f6ec6b0c2a
No known key found for this signature in database
3 changed files with 55 additions and 2 deletions

View file

@ -311,7 +311,7 @@ def image_generation(
or get_secret_str("AZURE_API_KEY")
)
azure_ad_token_param: Final = optional_params.pop("azure_ad_token", None)
azure_ad_token_param: Final = optional_params.get("extra_body", {}).pop("azure_ad_token", None)
azure_ad_token: Final = (
azure_ad_token_param
if isinstance(azure_ad_token_param, str) and azure_ad_token_param
@ -861,7 +861,7 @@ def image_edit(
):
image_edit_request_params.update(
flatten_form_field_values(
non_default_params,
{k: v for k, v in non_default_params.items() if k not in GenericLiteLLMParams.model_fields},
extra_body if isinstance(extra_body, dict) else None,
)
)

View file

@ -17,6 +17,7 @@ PNG_BYTES = b"\x89PNG\r\n\x1a\nfakepng"
def _capture_image_edit_request(captured):
def respond(request):
captured["content_type"] = request.headers.get("content-type")
captured["headers"] = request.headers
captured["body"] = request.content
return httpx.Response(200, json={"created": 1712697600, "data": [{"b64_json": "aW1n"}]})
@ -78,6 +79,30 @@ def test_image_edit_keeps_an_internal_prefixed_kwarg_out_of_the_provider_request
assert fields["seed"] == "42"
def test_azure_image_edit_sends_azure_ad_token_as_bearer_header_not_form_field(monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.delenv("AZURE_API_KEY", raising=False)
monkeypatch.delenv("AZURE_OPENAI_API_KEY", raising=False)
captured = {}
client = HTTPHandler(client=httpx.Client(transport=httpx.MockTransport(_capture_image_edit_request(captured))))
litellm.image_edit(
model="azure/gpt-image-deployment",
image=PNG_BYTES,
prompt="add a hat",
api_base="https://resource.services.ai.azure.com",
api_version="2025-04-01-preview",
azure_ad_token="entra-token",
client=client,
seed=42,
)
fields = _multipart_text_fields(captured["content_type"], captured["body"])
assert b"entra-token" not in captured["body"]
assert "azure_ad_token" not in fields
assert fields["seed"] == "42"
assert captured["headers"]["authorization"] == "Bearer entra-token"
def test_image_edit_extra_body_takes_precedence_over_kwargs():
captured = {}
client = HTTPHandler(client=httpx.Client(transport=httpx.MockTransport(_capture_image_edit_request(captured))))

View file

@ -2,6 +2,7 @@ import json
from typing import Final
import httpx
import pytest
import respx
import litellm
@ -27,3 +28,30 @@ def test_image_generation_keeps_an_internal_prefixed_kwarg_out_of_the_provider_r
sent: Final = json.loads(respx_mock.calls[0].request.content)
assert "_litellm_undeclared_sentinel" not in sent, sent
assert sent["prompt"] == "a red circle"
def test_azure_image_generation_sends_azure_ad_token_as_bearer_header_not_body_field(
respx_mock: respx.MockRouter, monkeypatch: pytest.MonkeyPatch
) -> None:
monkeypatch.delenv("AZURE_API_KEY", raising=False)
monkeypatch.delenv("AZURE_OPENAI_API_KEY", raising=False)
api_base: Final = "https://resource.services.ai.azure.com"
mock_route: Final = respx_mock.post(url__regex=rf"{api_base}/openai/deployments/gpt-image-deployment/.*").mock(
return_value=httpx.Response(status_code=200, json={"created": 1712697600, "data": [{"b64_json": "aW1n"}]})
)
litellm.image_generation(
model="azure/gpt-image-deployment",
prompt="a red circle",
api_base=api_base,
api_version="2025-04-01-preview",
azure_ad_token="entra-token",
seed=42,
)
assert mock_route.called
request: Final = respx_mock.calls[0].request
sent: Final = json.loads(request.content)
assert "entra-token" not in request.content.decode()
assert sent["seed"] == 42
assert request.headers["authorization"] == "Bearer entra-token"