Drive Azure AI MAI routing from model metadata

This commit is contained in:
Emerson Gomes 2026-05-02 15:24:12 -05:00
parent bef33bb0a4
commit 5668e26190
No known key found for this signature in database
GPG key ID: D3DF28AB5D1B5E17
6 changed files with 81 additions and 22 deletions

View file

@ -451,7 +451,9 @@ def image_generation( # noqa: PLR0915
)
elif custom_llm_provider == "azure_ai":
from litellm.llms.azure_ai.common_utils import AzureFoundryModelInfo
from litellm.llms.azure_ai.image_generation import is_mai_image_model
from litellm.llms.azure_ai.image_generation import (
AzureFoundryMAIImageGenerationConfig,
)
api_base = AzureFoundryModelInfo.get_api_base(api_base)
api_key = AzureFoundryModelInfo.get_api_key(api_key)
@ -460,8 +462,8 @@ def image_generation( # noqa: PLR0915
litellm_params_dict["api_base"] = api_base
if image_generation_config is not None and is_mai_image_model(
base_model or model
if isinstance(
image_generation_config, AzureFoundryMAIImageGenerationConfig
):
return llm_http_handler.image_generation_handler(
api_key=api_key,
@ -473,6 +475,7 @@ def image_generation( # noqa: PLR0915
litellm_params=litellm_params_dict,
logging_obj=litellm_logging_obj,
timeout=timeout,
extra_headers=headers,
client=client,
_is_async=aimg_generation,
)

View file

@ -1,3 +1,4 @@
import litellm
from litellm._logging import verbose_logger
from litellm.llms.base_llm.image_generation.transformation import (
BaseImageGenerationConfig,
@ -18,25 +19,41 @@ __all__ = [
]
def is_mai_image_model(model: str) -> bool:
normalized_model = model.lower().replace("-", "").replace("_", "")
return "maiimage2" in normalized_model
def _normalize_model_name(model: str) -> str:
return model.lower().replace("-", "").replace("_", "")
def _supports_mai_endpoint(model: str) -> bool:
model_key = model if model.startswith("azure_ai/") else f"azure_ai/{model}"
try:
resolved_model_info = litellm.get_model_info(model=model_key)
except Exception:
verbose_logger.debug(
"Azure AI model info not found for image model: %s", model_key
)
return False
resolved_key = resolved_model_info.get("key", model_key)
model_info = litellm.model_cost.get(resolved_key, {})
return bool(model_info.get("supports_mai_endpoint"))
def get_azure_ai_image_generation_config(model: str) -> BaseImageGenerationConfig:
model = model.lower()
model = model.replace("-", "")
model = model.replace("_", "")
if model == "" or "dalle2" in model: # empty model is dall-e-2
return AzureFoundryDallE2ImageGenerationConfig()
elif "dalle3" in model:
return AzureFoundryDallE3ImageGenerationConfig()
elif "flux" in model:
return AzureFoundryFluxImageGenerationConfig()
elif "maiimage2" in model:
if _supports_mai_endpoint(model):
return AzureFoundryMAIImageGenerationConfig()
normalized_model = _normalize_model_name(model)
if (
normalized_model == "" or "dalle2" in normalized_model
): # empty model is dall-e-2
return AzureFoundryDallE2ImageGenerationConfig()
elif "dalle3" in normalized_model:
return AzureFoundryDallE3ImageGenerationConfig()
elif "flux" in normalized_model:
return AzureFoundryFluxImageGenerationConfig()
else:
verbose_logger.debug(
f"Using AzureGPTImageGenerationConfig for model: {model}. This follows the gpt-image-1 model format."
"Using AzureGPTImageGenerationConfig for model: %s. This follows the "
"gpt-image-1 model format.",
normalized_model,
)
return AzureFoundryGPTImageGenerationConfig()

View file

@ -132,7 +132,7 @@ class AzureFoundryMAIImageGenerationConfig(BaseImageGenerationConfig):
"size",
f"{request_data['width']}x{request_data['height']}",
)
image_response.output_format = "png"
image_response.output_format = response.get("output_format", "png")
return image_response
@ -153,8 +153,18 @@ class AzureFoundryMAIImageGenerationConfig(BaseImageGenerationConfig):
@staticmethod
def _parse_size(size: str) -> Tuple[int, int]:
width_str, height_str = size.lower().split("x", maxsplit=1)
return int(width_str), int(height_str)
parts = size.lower().split("x", maxsplit=1)
if len(parts) != 2:
raise ValueError(
f"Invalid size format '{size}'. Expected 'WxH' (e.g. '1024x1024')."
)
width_str, height_str = parts
try:
return int(width_str), int(height_str)
except ValueError as exc:
raise ValueError(
f"Invalid size format '{size}'. Expected integer dimensions like '1024x1024'."
) from exc
@classmethod
def _resolve_dimensions(cls, optional_params: dict) -> Tuple[int, int]:

View file

@ -6582,6 +6582,7 @@
"mode": "image_generation",
"output_cost_per_image_token": 3.3e-05,
"source": "https://techcommunity.microsoft.com/blog/azure-ai-foundry-blog/introducing-mai-transcribe-1-mai-voice-1-and-mai-image-2-in-microsoft-foundry/4507787",
"supports_mai_endpoint": true,
"supported_endpoints": [
"/v1/images/generations"
],
@ -6600,6 +6601,7 @@
"mode": "image_generation",
"output_cost_per_image_token": 1.95e-05,
"source": "https://techcommunity.microsoft.com/blog/azure-ai-foundry-blog/introducing-mai-image-2-efficient-faster-more-efficient-image-generation/4510918",
"supports_mai_endpoint": true,
"supported_endpoints": [
"/v1/images/generations"
],

View file

@ -6596,6 +6596,7 @@
"mode": "image_generation",
"output_cost_per_image_token": 3.3e-05,
"source": "https://techcommunity.microsoft.com/blog/azure-ai-foundry-blog/introducing-mai-transcribe-1-mai-voice-1-and-mai-image-2-in-microsoft-foundry/4507787",
"supports_mai_endpoint": true,
"supported_endpoints": [
"/v1/images/generations"
],
@ -6614,6 +6615,7 @@
"mode": "image_generation",
"output_cost_per_image_token": 1.95e-05,
"source": "https://techcommunity.microsoft.com/blog/azure-ai-foundry-blog/introducing-mai-image-2-efficient-faster-more-efficient-image-generation/4510918",
"supports_mai_endpoint": true,
"supported_endpoints": [
"/v1/images/generations"
],

View file

@ -64,6 +64,7 @@ def test_azure_ai_mai_image_raw_model_cost_entry(
):
model_info = use_local_model_cost_map.model_cost[model_name]
assert model_info["supports_mai_endpoint"] is True
assert model_info["supported_endpoints"] == ["/v1/images/generations"]
assert model_info["supported_modalities"] == ["text"]
assert model_info["supported_output_modalities"] == ["image"]
@ -106,7 +107,9 @@ def test_azure_ai_mai_image_cost_calculator(
@pytest.mark.parametrize("model_name", ["MAI-Image-2", "MAI-Image-2e"])
def test_azure_ai_mai_image_generation_config(model_name: str):
def test_azure_ai_mai_image_generation_config(
use_local_model_cost_map, model_name: str
):
from litellm.llms.azure_ai.image_generation import (
AzureFoundryMAIImageGenerationConfig,
get_azure_ai_image_generation_config,
@ -148,10 +151,30 @@ def test_azure_ai_mai_image_generation_request_maps_size():
}
def test_azure_ai_mai_image_generation_request_invalid_size():
from litellm.llms.azure_ai.image_generation.mai_transformation import (
AzureFoundryMAIImageGenerationConfig,
)
with pytest.raises(
ValueError,
match="Invalid size format 'large'. Expected 'WxH'",
):
AzureFoundryMAIImageGenerationConfig().transform_image_generation_request(
model="MAI-Image-2e",
prompt="A glowing jellyfish in a glass ocean",
optional_params={"size": "large"},
litellm_params={},
headers={},
)
@patch("litellm.images.main.azure_chat_completions.image_generation")
@patch("litellm.images.main.llm_http_handler.image_generation_handler")
def test_azure_ai_mai_image_generation_routes_through_http_handler(
mock_image_generation_handler, mock_azure_image_generation
mock_image_generation_handler,
mock_azure_image_generation,
use_local_model_cost_map,
):
import litellm
from litellm.images.main import image_generation
@ -171,6 +194,7 @@ def test_azure_ai_mai_image_generation_routes_through_http_handler(
api_base="https://example.services.ai.azure.com",
api_key="test-key",
size="1024x1024",
headers={"x-trace-id": "abc123"},
)
assert response == mock_response
@ -184,6 +208,7 @@ def test_azure_ai_mai_image_generation_routes_through_http_handler(
assert call_kwargs["image_generation_optional_request_params"] == {
"size": "1024x1024"
}
assert call_kwargs["extra_headers"] == {"x-trace-id": "abc123"}
assert call_kwargs["litellm_params"]["api_base"] == (
"https://example.services.ai.azure.com"
)