From 5668e26190945635558e388da01532234ad20b6d Mon Sep 17 00:00:00 2001 From: Emerson Gomes Date: Sat, 2 May 2026 15:24:12 -0500 Subject: [PATCH] Drive Azure AI MAI routing from model metadata --- litellm/images/main.py | 9 ++-- .../azure_ai/image_generation/__init__.py | 45 +++++++++++++------ .../image_generation/mai_transformation.py | 16 +++++-- ...odel_prices_and_context_window_backup.json | 2 + model_prices_and_context_window.json | 2 + ..._azure_ai_mai_image_generation_metadata.py | 29 +++++++++++- 6 files changed, 81 insertions(+), 22 deletions(-) diff --git a/litellm/images/main.py b/litellm/images/main.py index f154955edfc..9cb856bbb60 100644 --- a/litellm/images/main.py +++ b/litellm/images/main.py @@ -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, ) diff --git a/litellm/llms/azure_ai/image_generation/__init__.py b/litellm/llms/azure_ai/image_generation/__init__.py index 631f19e9321..c21323b9ff6 100644 --- a/litellm/llms/azure_ai/image_generation/__init__.py +++ b/litellm/llms/azure_ai/image_generation/__init__.py @@ -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() diff --git a/litellm/llms/azure_ai/image_generation/mai_transformation.py b/litellm/llms/azure_ai/image_generation/mai_transformation.py index c048bb33eda..286149b3126 100644 --- a/litellm/llms/azure_ai/image_generation/mai_transformation.py +++ b/litellm/llms/azure_ai/image_generation/mai_transformation.py @@ -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]: diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index 2cf4c69bd8c..d31f3946e1c 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -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" ], diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index fa58769e90d..e90df3b6d51 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -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" ], diff --git a/tests/test_litellm/llms/azure_ai/image_generation/test_azure_ai_mai_image_generation_metadata.py b/tests/test_litellm/llms/azure_ai/image_generation/test_azure_ai_mai_image_generation_metadata.py index 4c38aeaf26f..c701a085422 100644 --- a/tests/test_litellm/llms/azure_ai/image_generation/test_azure_ai_mai_image_generation_metadata.py +++ b/tests/test_litellm/llms/azure_ai/image_generation/test_azure_ai_mai_image_generation_metadata.py @@ -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" )