diff --git a/litellm/llms/azure/azure.py b/litellm/llms/azure/azure.py index b111859d74e..7abba583560 100644 --- a/litellm/llms/azure/azure.py +++ b/litellm/llms/azure/azure.py @@ -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 diff --git a/litellm/llms/azure/image_generation/__init__.py b/litellm/llms/azure/image_generation/__init__.py index f60e446f0c4..64636bc689d 100644 --- a/litellm/llms/azure/image_generation/__init__.py +++ b/litellm/llms/azure/image_generation/__init__.py @@ -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." diff --git a/litellm/llms/azure_ai/image_generation/mai_transformation.py b/litellm/llms/azure_ai/image_generation/mai_transformation.py index e80a524e843..04f7a0081e4 100644 --- a/litellm/llms/azure_ai/image_generation/mai_transformation.py +++ b/litellm/llms/azure_ai/image_generation/mai_transformation.py @@ -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}" diff --git a/tests/test_litellm/llms/azure_ai/image_edit/test_mai_image_edit_transformation.py b/tests/test_litellm/llms/azure_ai/image_edit/test_mai_image_edit_transformation.py index 465d11d7448..62126db2475 100644 --- a/tests/test_litellm/llms/azure_ai/image_edit/test_mai_image_edit_transformation.py +++ b/tests/test_litellm/llms/azure_ai/image_edit/test_mai_image_edit_transformation.py @@ -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) diff --git a/tests/test_litellm/llms/azure_ai/image_generation/test_mai_image_generation.py b/tests/test_litellm/llms/azure_ai/image_generation/test_mai_image_generation.py index 0143864f945..c74c863b313 100644 --- a/tests/test_litellm/llms/azure_ai/image_generation/test_mai_image_generation.py +++ b/tests/test_litellm/llms/azure_ai/image_generation/test_mai_image_generation.py @@ -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="")