Fix Azure MAI image response handling

This commit is contained in:
Cursor Agent 2026-06-08 07:11:15 +00:00
parent 7c8ed6d100
commit e08d51ff2b
No known key found for this signature in database
5 changed files with 76 additions and 2 deletions

View file

@ -43,7 +43,10 @@ from .common_utils import (
process_azure_headers,
select_azure_base_url_or_endpoint,
)
from .image_generation import get_azure_image_generation_config
from .image_generation import (
AzureFoundryMAIImageGenerationConfig,
get_azure_image_generation_config,
)
from .image_generation.http_utils import azure_deployment_image_generation_json_body
@ -1317,6 +1320,21 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM):
data=data,
headers=headers,
)
provider_config = get_azure_image_generation_config(
data.get("model", "dall-e-2")
)
if isinstance(provider_config, AzureFoundryMAIImageGenerationConfig):
return provider_config.transform_image_generation_response(
model=data.get("model", "dall-e-2"),
raw_response=httpx_response,
model_response=model_response or ImageResponse(),
logging_obj=logging_obj,
request_data=data,
optional_params=data,
litellm_params=data,
encoding=litellm.encoding,
)
response = httpx_response.json()
## LOGGING

View file

@ -1,4 +1,5 @@
from litellm._logging import verbose_logger
from litellm.llms.azure_ai.image_generation import AzureFoundryMAIImageGenerationConfig
from litellm.llms.base_llm.image_generation.transformation import (
BaseImageGenerationConfig,
)
@ -24,6 +25,8 @@ def get_azure_image_generation_config(model: str) -> BaseImageGenerationConfig:
return AzureDallE2ImageGenerationConfig()
elif "dalle3" in model:
return AzureDallE3ImageGenerationConfig()
elif AzureFoundryMAIImageGenerationConfig.is_mai_model(model):
return AzureFoundryMAIImageGenerationConfig()
else:
verbose_logger.debug(
f"Using AzureGPTImageGenerationConfig for model: {model}. This follows the gpt-image model format."

View file

@ -49,6 +49,7 @@ class AzureFoundryMAIImageGenerationConfig(BaseImageGenerationConfig):
api_version = api_version or "preview"
if "/mai/" in api_base:
api_base = api_base.replace("/images/generations", "/images/edits")
if "?" in api_base:
return api_base
return f"{api_base}?api-version={api_version}"

View file

@ -1,5 +1,4 @@
import io
import json
import os
import sys
from unittest.mock import MagicMock
@ -29,6 +28,19 @@ class TestAzureMAIImageEdit:
== "https://my-resource.services.ai.azure.com/mai/v1/images/edits?api-version=preview"
)
def test_get_mai_image_edit_url_rewrites_generation_url(self):
url = AzureFoundryMAIImageGenerationConfig.get_mai_image_edit_url(
api_base=(
"https://my-resource.services.ai.azure.com/mai/v1/images/generations"
"?api-version=preview"
),
api_version="preview",
)
assert (
url
== "https://my-resource.services.ai.azure.com/mai/v1/images/edits?api-version=preview"
)
def test_get_azure_ai_image_edit_config_returns_mai(self):
config = get_azure_ai_image_edit_config("MAI-Image-2.5")
assert isinstance(config, AzureFoundryMAIImageEditConfig)

View file

@ -9,6 +9,7 @@ sys.path.insert(0, os.path.abspath("../../../../../.."))
import litellm
from litellm.llms.azure.azure import AzureChatCompletion
from litellm.llms.azure.image_generation import get_azure_image_generation_config
from litellm.llms.azure.image_generation.http_utils import (
azure_deployment_image_generation_json_body,
)
@ -83,6 +84,10 @@ class TestAzureMAIImageGeneration:
config = get_azure_ai_image_generation_config("MAI-Image-2.5")
assert isinstance(config, AzureFoundryMAIImageGenerationConfig)
def test_azure_image_generation_config_returns_mai(self):
config = get_azure_image_generation_config("MAI-Image-2.5")
assert isinstance(config, AzureFoundryMAIImageGenerationConfig)
def test_map_openai_params_size_to_width_height(self):
config = AzureFoundryMAIImageGenerationConfig()
optional_params = config.map_openai_params(
@ -241,6 +246,41 @@ class TestAzureMAIImageGeneration:
assert image_response.usage.input_tokens == 22
assert image_response.usage.total_tokens == 1046
def test_azure_sync_image_generation_uses_mai_response_transform(self):
raw_response = MagicMock(spec=httpx.Response)
raw_response.json.return_value = {
"created": 1780897477,
"data": [{"b64_json": "abc123"}],
"usage": {
"num_output_tokens": 1024,
"num_input_text_tokens": 22,
},
}
class MAIImageGenerationAzureChatCompletion(AzureChatCompletion):
def make_sync_azure_httpx_request(self, **kwargs):
return raw_response
logging_obj = MagicMock()
image_response = MAIImageGenerationAzureChatCompletion().image_generation(
prompt="A red fox",
timeout=60.0,
optional_params={"width": 1792, "height": 1024},
logging_obj=logging_obj,
headers={},
model="MAI-Image-2.5",
api_key="test-key",
api_base="https://my-resource.services.ai.azure.com",
api_version="preview",
litellm_params={},
)
assert image_response.data[0].b64_json == "abc123"
assert image_response.usage.output_tokens == 1024
assert image_response.usage.input_tokens == 22
assert image_response.usage.total_tokens == 1046
assert image_response.size == "1792x1024"
def test_mai_image_cost_calculator_token_based(self):
os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True"
litellm.model_cost = litellm.get_model_cost_map(url="")