feat(azure_ai): add MAI-Image-2.5 image generation support

Route azure_ai MAI models to /mai/v1/images/generations and map OpenAI size to width/height for the serverless API.

Co-authored-by: Cursor <cursoragent@cursor.com>
This commit is contained in:
Sameer Kankute 2026-06-04 22:39:39 +05:30
parent ed073d382d
commit f2c1ff1e74
No known key found for this signature in database
5 changed files with 271 additions and 0 deletions

View file

@ -1101,6 +1101,7 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM):
) -> str:
from litellm.llms.azure_ai.image_generation import (
AzureFoundryFluxImageGenerationConfig,
AzureFoundryMAIImageGenerationConfig,
)
api_base: str = azure_client_params.get(
@ -1112,6 +1113,13 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM):
if model is None:
model = ""
# MAI image models: /mai/v1/images/generations (serverless Azure AI)
if AzureFoundryMAIImageGenerationConfig.is_mai_model(model):
return AzureFoundryMAIImageGenerationConfig.get_mai_image_generation_url(
api_base=api_base,
api_version=api_version,
)
# Handle FLUX 2 models on Azure AI which use a different URL pattern
# e.g., /providers/blackforestlabs/v1/flux-2-pro instead of /openai/deployments/{model}/images/generations
if AzureFoundryFluxImageGenerationConfig.is_flux2_model(model):

View file

@ -7,12 +7,14 @@ from .dall_e_2_transformation import AzureFoundryDallE2ImageGenerationConfig
from .dall_e_3_transformation import AzureFoundryDallE3ImageGenerationConfig
from .flux_transformation import AzureFoundryFluxImageGenerationConfig
from .gpt_transformation import AzureFoundryGPTImageGenerationConfig
from .mai_transformation import AzureFoundryMAIImageGenerationConfig
__all__ = [
"AzureFoundryFluxImageGenerationConfig",
"AzureFoundryGPTImageGenerationConfig",
"AzureFoundryDallE2ImageGenerationConfig",
"AzureFoundryDallE3ImageGenerationConfig",
"AzureFoundryMAIImageGenerationConfig",
]
@ -24,6 +26,8 @@ def get_azure_ai_image_generation_config(model: str) -> BaseImageGenerationConfi
return AzureFoundryDallE2ImageGenerationConfig()
elif "dalle3" in model:
return AzureFoundryDallE3ImageGenerationConfig()
elif AzureFoundryMAIImageGenerationConfig.is_mai_model(model):
return AzureFoundryMAIImageGenerationConfig()
elif "flux" in model:
return AzureFoundryFluxImageGenerationConfig()
else:

View file

@ -0,0 +1,138 @@
from typing import TYPE_CHECKING, Any, List, Optional
import httpx
from litellm.llms.base_llm.image_generation.transformation import (
BaseImageGenerationConfig,
)
from litellm.types.llms.openai import OpenAIImageGenerationOptionalParams
from litellm.types.utils import ImageResponse
from litellm.utils import convert_to_model_response_object
if TYPE_CHECKING:
from litellm.litellm_core_utils.logging import Logging as LiteLLMLoggingObj
class AzureFoundryMAIImageGenerationConfig(BaseImageGenerationConfig):
"""Azure AI Foundry MAI image generation (e.g. MAI-Image-2.5)."""
DEFAULT_WIDTH = 1024
DEFAULT_HEIGHT = 1024
@staticmethod
def get_mai_image_generation_url(
api_base: Optional[str],
api_version: Optional[str],
) -> str:
if api_base is None:
raise ValueError("api_base is required for Azure AI MAI image generation")
api_base = api_base.rstrip("/")
api_version = api_version or "preview"
if "/mai/" in api_base:
if "?" in api_base:
return api_base
return f"{api_base}?api-version={api_version}"
return f"{api_base}/mai/v1/images/generations?api-version={api_version}"
@staticmethod
def is_mai_model(model: str) -> bool:
model_normalized = model.lower().replace("-", "").replace("_", "")
return "maiimage" in model_normalized
def get_supported_openai_params(
self, model: str
) -> List[OpenAIImageGenerationOptionalParams]:
return ["n", "size"]
def map_openai_params(
self,
non_default_params: dict,
optional_params: dict,
model: str,
drop_params: bool,
) -> dict:
supported_params = self.get_supported_openai_params(model)
for k, v in non_default_params.items():
if k in optional_params:
continue
if k in supported_params:
if k == "size" and v:
self._map_size_param(v, optional_params)
else:
optional_params[k] = v
elif k in ("width", "height"):
optional_params[k] = v
elif not drop_params:
raise ValueError(
f"Parameter {k} is not supported for model {model}. "
f"Supported parameters are {supported_params} and width/height. "
f"Set drop_params=True to drop unsupported parameters."
)
if "width" not in optional_params and "height" not in optional_params:
optional_params["width"] = self.DEFAULT_WIDTH
optional_params["height"] = self.DEFAULT_HEIGHT
optional_params.pop("size", None)
return optional_params
def _map_size_param(self, size: str, optional_params: dict) -> None:
size_mapping = {
"1024x1024": (1024, 1024),
"1792x1024": (1792, 1024),
"1024x1792": (1024, 1792),
"512x512": (512, 512),
"256x256": (256, 256),
}
if size in size_mapping:
width, height = size_mapping[size]
optional_params["width"] = width
optional_params["height"] = height
elif "x" in size:
try:
width, height = map(int, size.lower().split("x"))
optional_params["width"] = width
optional_params["height"] = height
except ValueError:
raise ValueError(
f"Invalid size format: '{size}'. Expected format 'WIDTHxHEIGHT' (e.g., '1024x1024')."
)
def transform_image_generation_response(
self,
model: str,
raw_response: httpx.Response,
model_response: ImageResponse,
logging_obj: "LiteLLMLoggingObj",
request_data: dict,
optional_params: dict,
litellm_params: dict,
encoding: Any,
api_key: Optional[str] = None,
json_mode: Optional[bool] = None,
) -> ImageResponse:
response = raw_response.json()
logging_obj.post_call(
input=request_data.get("prompt", ""),
api_key=api_key,
additional_args={"complete_input_dict": request_data},
original_response=response,
)
image_response: ImageResponse = convert_to_model_response_object(
response_object=response,
model_response_object=model_response,
response_type="image_generation",
)
width = optional_params.get("width", self.DEFAULT_WIDTH)
height = optional_params.get("height", self.DEFAULT_HEIGHT)
image_response.size = f"{width}x{height}"
return image_response

View file

@ -6853,6 +6853,13 @@
"/v1/images/generations"
]
},
"azure_ai/MAI-Image-2.5": {
"litellm_provider": "azure_ai",
"mode": "image_generation",
"supported_endpoints": [
"/v1/images/generations"
]
},
"azure_ai/Llama-3.2-11B-Vision-Instruct": {
"input_cost_per_token": 3.7e-07,
"litellm_provider": "azure_ai",

View file

@ -0,0 +1,114 @@
import os
import sys
sys.path.insert(0, os.path.abspath("../../../../../.."))
from litellm.llms.azure.azure import AzureChatCompletion
from litellm.llms.azure.image_generation.http_utils import (
azure_deployment_image_generation_json_body,
)
from litellm.llms.azure_ai.image_generation import (
AzureFoundryMAIImageGenerationConfig,
get_azure_ai_image_generation_config,
)
from litellm.utils import get_optional_params_image_gen
class TestAzureMAIImageGeneration:
def test_is_mai_model(self):
assert AzureFoundryMAIImageGenerationConfig.is_mai_model("MAI-Image-2.5")
assert AzureFoundryMAIImageGenerationConfig.is_mai_model(
"azure_ai/MAI-Image-2.5"
)
assert not AzureFoundryMAIImageGenerationConfig.is_mai_model("flux.2-pro")
assert not AzureFoundryMAIImageGenerationConfig.is_mai_model("MAI-DS-R1")
def test_get_mai_image_generation_url(self):
url = AzureFoundryMAIImageGenerationConfig.get_mai_image_generation_url(
api_base="https://my-resource.services.ai.azure.com",
api_version="preview",
)
assert (
url
== "https://my-resource.services.ai.azure.com/mai/v1/images/generations?api-version=preview"
)
def test_get_mai_image_generation_url_preserves_full_path(self):
api = (
"https://my-resource.services.ai.azure.com/mai/v1/images/generations"
"?api-version=preview"
)
url = AzureFoundryMAIImageGenerationConfig.get_mai_image_generation_url(
api_base=api,
api_version="preview",
)
assert url == api
def test_get_azure_ai_image_generation_config_returns_mai(self):
config = get_azure_ai_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(
non_default_params={"size": "1024x1024", "n": 1},
optional_params={},
model="MAI-Image-2.5",
drop_params=True,
)
assert optional_params["width"] == 1024
assert optional_params["height"] == 1024
assert optional_params["n"] == 1
assert "size" not in optional_params
def test_map_openai_params_defaults(self):
config = AzureFoundryMAIImageGenerationConfig()
optional_params = config.map_openai_params(
non_default_params={},
optional_params={},
model="MAI-Image-2.5",
drop_params=True,
)
assert optional_params["width"] == 1024
assert optional_params["height"] == 1024
def test_get_optional_params_image_gen_mai(self):
config = AzureFoundryMAIImageGenerationConfig()
optional_params = get_optional_params_image_gen(
model="MAI-Image-2.5",
size="1792x1024",
n=1,
custom_llm_provider="azure_ai",
provider_config=config,
drop_params=True,
)
assert optional_params["width"] == 1792
assert optional_params["height"] == 1024
assert "size" not in optional_params
def test_azure_create_azure_base_url_mai(self):
azure_chat = AzureChatCompletion()
url = azure_chat.create_azure_base_url(
azure_client_params={
"azure_endpoint": "https://my-resource.services.ai.azure.com",
"api_version": "preview",
},
model="MAI-Image-2.5",
)
assert "/mai/v1/images/generations" in url
assert "api-version=preview" in url
def test_mai_json_body_keeps_model(self):
api = (
"https://my-resource.services.ai.azure.com/mai/v1/images/generations"
"?api-version=preview"
)
data = {
"model": "MAI-Image-2.5",
"prompt": "A photograph of a red fox",
"width": 1024,
"height": 1024,
"n": 1,
}
out = azure_deployment_image_generation_json_body(api, data)
assert out == data