mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
Add Azure AI MAI image generation support
This commit is contained in:
parent
934ecdca78
commit
bef33bb0a4
8 changed files with 503 additions and 10 deletions
|
|
@ -451,12 +451,32 @@ 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
|
||||
|
||||
api_base = AzureFoundryModelInfo.get_api_base(api_base)
|
||||
api_key = AzureFoundryModelInfo.get_api_key(api_key)
|
||||
if extra_headers is not None:
|
||||
optional_params["extra_headers"] = extra_headers
|
||||
|
||||
litellm_params_dict["api_base"] = api_base
|
||||
|
||||
if image_generation_config is not None and is_mai_image_model(
|
||||
base_model or model
|
||||
):
|
||||
return llm_http_handler.image_generation_handler(
|
||||
api_key=api_key,
|
||||
model=model,
|
||||
prompt=prompt,
|
||||
image_generation_provider_config=image_generation_config,
|
||||
image_generation_optional_request_params=optional_params,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
litellm_params=litellm_params_dict,
|
||||
logging_obj=litellm_logging_obj,
|
||||
timeout=timeout,
|
||||
client=client,
|
||||
_is_async=aimg_generation,
|
||||
)
|
||||
|
||||
default_headers = {
|
||||
"Content-Type": "application/json",
|
||||
}
|
||||
|
|
|
|||
|
|
@ -7,15 +7,22 @@ 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",
|
||||
]
|
||||
|
||||
|
||||
def is_mai_image_model(model: str) -> bool:
|
||||
normalized_model = model.lower().replace("-", "").replace("_", "")
|
||||
return "maiimage2" in normalized_model
|
||||
|
||||
|
||||
def get_azure_ai_image_generation_config(model: str) -> BaseImageGenerationConfig:
|
||||
model = model.lower()
|
||||
model = model.replace("-", "")
|
||||
|
|
@ -26,6 +33,8 @@ def get_azure_ai_image_generation_config(model: str) -> BaseImageGenerationConfi
|
|||
return AzureFoundryDallE3ImageGenerationConfig()
|
||||
elif "flux" in model:
|
||||
return AzureFoundryFluxImageGenerationConfig()
|
||||
elif "maiimage2" in model:
|
||||
return AzureFoundryMAIImageGenerationConfig()
|
||||
else:
|
||||
verbose_logger.debug(
|
||||
f"Using AzureGPTImageGenerationConfig for model: {model}. This follows the gpt-image-1 model format."
|
||||
|
|
|
|||
|
|
@ -1,6 +1,9 @@
|
|||
from typing import Any
|
||||
|
||||
import litellm
|
||||
from litellm.litellm_core_utils.llm_cost_calc.utils import (
|
||||
calculate_image_response_cost_from_usage,
|
||||
)
|
||||
from litellm.types.utils import ImageResponse
|
||||
|
||||
|
||||
|
|
@ -9,19 +12,29 @@ def cost_calculator(
|
|||
image_response: Any,
|
||||
) -> float:
|
||||
"""
|
||||
Recraft image generation cost calculator
|
||||
Cost calculator for Azure AI image generation models.
|
||||
|
||||
Azure AI supports both flat per-image pricing and token-based image pricing.
|
||||
Prefer usage-based calculation when the response provides token usage, then
|
||||
fall back to per-image pricing for models that only expose flat image cost.
|
||||
"""
|
||||
if not isinstance(image_response, ImageResponse):
|
||||
raise ValueError(
|
||||
f"image_response must be of type ImageResponse got type={type(image_response)}"
|
||||
)
|
||||
|
||||
usage_based_cost = calculate_image_response_cost_from_usage(
|
||||
model=model,
|
||||
image_response=image_response,
|
||||
custom_llm_provider=litellm.LlmProviders.AZURE_AI.value,
|
||||
)
|
||||
if usage_based_cost is not None:
|
||||
return usage_based_cost
|
||||
|
||||
_model_info = litellm.get_model_info(
|
||||
model=model,
|
||||
custom_llm_provider=litellm.LlmProviders.AZURE_AI.value,
|
||||
)
|
||||
output_cost_per_image: float = _model_info.get("output_cost_per_image") or 0.0
|
||||
num_images: int = 0
|
||||
if isinstance(image_response, ImageResponse):
|
||||
if image_response.data:
|
||||
num_images = len(image_response.data)
|
||||
return output_cost_per_image * num_images
|
||||
else:
|
||||
raise ValueError(
|
||||
f"image_response must be of type ImageResponse got type={type(image_response)}"
|
||||
)
|
||||
num_images = len(image_response.data) if image_response.data else 0
|
||||
return output_cost_per_image * num_images
|
||||
|
|
|
|||
188
litellm/llms/azure_ai/image_generation/mai_transformation.py
Normal file
188
litellm/llms/azure_ai/image_generation/mai_transformation.py
Normal file
|
|
@ -0,0 +1,188 @@
|
|||
from typing import TYPE_CHECKING, Any, List, Optional, Tuple
|
||||
from urllib.parse import parse_qsl, urlencode, urlsplit, urlunsplit
|
||||
|
||||
import httpx
|
||||
|
||||
from litellm.llms.azure_ai.common_utils import AzureFoundryModelInfo
|
||||
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 Foundry MAI image generation config.
|
||||
|
||||
MAI image models use the Foundry-specific /mai/v1/images/generations endpoint
|
||||
instead of the Azure OpenAI /openai/deployments/.../images/generations route.
|
||||
"""
|
||||
|
||||
def get_supported_openai_params(
|
||||
self, model: str
|
||||
) -> List[OpenAIImageGenerationOptionalParams]:
|
||||
return ["height", "size", "width"]
|
||||
|
||||
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 key, value in non_default_params.items():
|
||||
if key in optional_params:
|
||||
continue
|
||||
if key in supported_params:
|
||||
optional_params[key] = value
|
||||
elif not drop_params:
|
||||
raise ValueError(
|
||||
f"Parameter {key} is not supported for model {model}. "
|
||||
f"Supported parameters are {supported_params}. Set drop_params=True "
|
||||
"to drop unsupported parameters."
|
||||
)
|
||||
|
||||
return optional_params
|
||||
|
||||
def validate_environment(
|
||||
self,
|
||||
headers: dict,
|
||||
model: str,
|
||||
messages: list,
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
api_key: Optional[str] = None,
|
||||
api_base: Optional[str] = None,
|
||||
) -> dict:
|
||||
validated_headers = dict(headers)
|
||||
validated_headers.setdefault("Content-Type", "application/json")
|
||||
if api_key:
|
||||
validated_headers.setdefault("api-key", api_key)
|
||||
return validated_headers
|
||||
|
||||
def get_complete_url(
|
||||
self,
|
||||
api_base: Optional[str],
|
||||
api_key: Optional[str],
|
||||
model: str,
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
stream: Optional[bool] = None,
|
||||
) -> str:
|
||||
resolved_api_base = AzureFoundryModelInfo.get_api_base(api_base)
|
||||
if resolved_api_base is None:
|
||||
raise ValueError("api_base is required for Azure AI MAI image generation")
|
||||
|
||||
api_version = litellm_params.get("api_version") or "preview"
|
||||
return self._append_api_version(
|
||||
f"{resolved_api_base.rstrip('/')}/mai/v1/images/generations",
|
||||
api_version=api_version,
|
||||
)
|
||||
|
||||
def transform_image_generation_request(
|
||||
self,
|
||||
model: str,
|
||||
prompt: str,
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
headers: dict,
|
||||
) -> dict:
|
||||
width, height = self._resolve_dimensions(optional_params)
|
||||
|
||||
return {
|
||||
"model": model,
|
||||
"prompt": prompt,
|
||||
"width": width,
|
||||
"height": height,
|
||||
}
|
||||
|
||||
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( # type: ignore
|
||||
response_object=response,
|
||||
model_response_object=model_response,
|
||||
response_type="image_generation",
|
||||
)
|
||||
|
||||
image_response.size = response.get(
|
||||
"size",
|
||||
f"{request_data['width']}x{request_data['height']}",
|
||||
)
|
||||
image_response.output_format = "png"
|
||||
|
||||
return image_response
|
||||
|
||||
@staticmethod
|
||||
def _append_api_version(url: str, api_version: str) -> str:
|
||||
split_url = urlsplit(url)
|
||||
query_params = dict(parse_qsl(split_url.query))
|
||||
query_params.setdefault("api-version", api_version)
|
||||
return urlunsplit(
|
||||
(
|
||||
split_url.scheme,
|
||||
split_url.netloc,
|
||||
split_url.path,
|
||||
urlencode(query_params),
|
||||
split_url.fragment,
|
||||
)
|
||||
)
|
||||
|
||||
@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)
|
||||
|
||||
@classmethod
|
||||
def _resolve_dimensions(cls, optional_params: dict) -> Tuple[int, int]:
|
||||
width = optional_params.get("width")
|
||||
height = optional_params.get("height")
|
||||
|
||||
if width is not None or height is not None:
|
||||
if width is None or height is None:
|
||||
raise ValueError(
|
||||
"Azure AI MAI image generation requires both width and height when either is provided."
|
||||
)
|
||||
return cls._validate_dimensions(width=int(width), height=int(height))
|
||||
|
||||
size = optional_params.get("size")
|
||||
if size is not None:
|
||||
parsed_width, parsed_height = cls._parse_size(size)
|
||||
return cls._validate_dimensions(width=parsed_width, height=parsed_height)
|
||||
|
||||
return cls._validate_dimensions(width=1024, height=1024)
|
||||
|
||||
@staticmethod
|
||||
def _validate_dimensions(width: int, height: int) -> Tuple[int, int]:
|
||||
if width < 768 or height < 768:
|
||||
raise ValueError(
|
||||
"Azure AI MAI image generation requires width and height to be at least 768."
|
||||
)
|
||||
if width * height > 1_048_576:
|
||||
raise ValueError(
|
||||
"Azure AI MAI image generation supports a maximum of 1,048,576 total pixels."
|
||||
)
|
||||
return width, height
|
||||
|
|
@ -6574,6 +6574,42 @@
|
|||
"supports_reasoning": true,
|
||||
"supports_tool_choice": true
|
||||
},
|
||||
"azure_ai/MAI-Image-2": {
|
||||
"input_cost_per_token": 5e-06,
|
||||
"litellm_provider": "azure_ai",
|
||||
"max_input_tokens": 32000,
|
||||
"max_tokens": 32000,
|
||||
"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",
|
||||
"supported_endpoints": [
|
||||
"/v1/images/generations"
|
||||
],
|
||||
"supported_modalities": [
|
||||
"text"
|
||||
],
|
||||
"supported_output_modalities": [
|
||||
"image"
|
||||
]
|
||||
},
|
||||
"azure_ai/MAI-Image-2e": {
|
||||
"input_cost_per_token": 5e-06,
|
||||
"litellm_provider": "azure_ai",
|
||||
"max_input_tokens": 32000,
|
||||
"max_tokens": 32000,
|
||||
"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",
|
||||
"supported_endpoints": [
|
||||
"/v1/images/generations"
|
||||
],
|
||||
"supported_modalities": [
|
||||
"text"
|
||||
],
|
||||
"supported_output_modalities": [
|
||||
"image"
|
||||
]
|
||||
},
|
||||
"azure_ai/cohere-rerank-v3-english": {
|
||||
"input_cost_per_query": 0.002,
|
||||
"input_cost_per_token": 0.0,
|
||||
|
|
|
|||
|
|
@ -1059,6 +1059,7 @@ OpenAIImageVariationOptionalParams = Literal["n", "size", "response_format", "us
|
|||
|
||||
OpenAIImageGenerationOptionalParams = Literal[
|
||||
"background",
|
||||
"height",
|
||||
"moderation",
|
||||
"n",
|
||||
"output_compression",
|
||||
|
|
@ -1076,6 +1077,7 @@ OpenAIImageGenerationOptionalParams = Literal[
|
|||
"image_url",
|
||||
"image_prompt_strength",
|
||||
"aspect_ratio",
|
||||
"width",
|
||||
]
|
||||
|
||||
OpenAIImageEditOptionalParams = Literal[
|
||||
|
|
|
|||
|
|
@ -6588,6 +6588,42 @@
|
|||
"supports_reasoning": true,
|
||||
"supports_tool_choice": true
|
||||
},
|
||||
"azure_ai/MAI-Image-2": {
|
||||
"input_cost_per_token": 5e-06,
|
||||
"litellm_provider": "azure_ai",
|
||||
"max_input_tokens": 32000,
|
||||
"max_tokens": 32000,
|
||||
"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",
|
||||
"supported_endpoints": [
|
||||
"/v1/images/generations"
|
||||
],
|
||||
"supported_modalities": [
|
||||
"text"
|
||||
],
|
||||
"supported_output_modalities": [
|
||||
"image"
|
||||
]
|
||||
},
|
||||
"azure_ai/MAI-Image-2e": {
|
||||
"input_cost_per_token": 5e-06,
|
||||
"litellm_provider": "azure_ai",
|
||||
"max_input_tokens": 32000,
|
||||
"max_tokens": 32000,
|
||||
"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",
|
||||
"supported_endpoints": [
|
||||
"/v1/images/generations"
|
||||
],
|
||||
"supported_modalities": [
|
||||
"text"
|
||||
],
|
||||
"supported_output_modalities": [
|
||||
"image"
|
||||
]
|
||||
},
|
||||
"azure_ai/cohere-rerank-v3-english": {
|
||||
"input_cost_per_query": 0.002,
|
||||
"input_cost_per_token": 0.0,
|
||||
|
|
|
|||
|
|
@ -0,0 +1,189 @@
|
|||
"""
|
||||
Test Azure AI MAI image generation model metadata and pricing.
|
||||
"""
|
||||
|
||||
import json
|
||||
from importlib.resources import files
|
||||
from unittest.mock import patch
|
||||
|
||||
import pytest
|
||||
|
||||
|
||||
@pytest.fixture(scope="module")
|
||||
def use_local_model_cost_map():
|
||||
monkeypatch = pytest.MonkeyPatch()
|
||||
monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True")
|
||||
|
||||
import litellm
|
||||
from litellm.utils import _invalidate_model_cost_lowercase_map
|
||||
|
||||
original_model_cost = litellm.model_cost
|
||||
litellm.model_cost = json.loads(
|
||||
files("litellm")
|
||||
.joinpath("model_prices_and_context_window_backup.json")
|
||||
.read_text(encoding="utf-8")
|
||||
)
|
||||
litellm.get_model_info.cache_clear()
|
||||
_invalidate_model_cost_lowercase_map()
|
||||
try:
|
||||
yield litellm
|
||||
finally:
|
||||
litellm.model_cost = original_model_cost
|
||||
litellm.get_model_info.cache_clear()
|
||||
_invalidate_model_cost_lowercase_map()
|
||||
monkeypatch.undo()
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"model_name,output_cost_per_image_token",
|
||||
[
|
||||
("azure_ai/MAI-Image-2", 3.3e-05),
|
||||
("azure_ai/MAI-Image-2e", 1.95e-05),
|
||||
],
|
||||
)
|
||||
def test_azure_ai_mai_image_model_info(
|
||||
use_local_model_cost_map, model_name: str, output_cost_per_image_token: float
|
||||
):
|
||||
model_info = use_local_model_cost_map.get_model_info(model=model_name)
|
||||
|
||||
assert model_info["litellm_provider"] == "azure_ai"
|
||||
assert model_info["mode"] == "image_generation"
|
||||
assert model_info["max_input_tokens"] == 32000
|
||||
assert model_info["max_tokens"] == 32000
|
||||
assert model_info["input_cost_per_token"] == pytest.approx(5e-06)
|
||||
assert model_info["output_cost_per_image_token"] == pytest.approx(
|
||||
output_cost_per_image_token
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"model_name", ["azure_ai/MAI-Image-2", "azure_ai/MAI-Image-2e"]
|
||||
)
|
||||
def test_azure_ai_mai_image_raw_model_cost_entry(
|
||||
use_local_model_cost_map, model_name: str
|
||||
):
|
||||
model_info = use_local_model_cost_map.model_cost[model_name]
|
||||
|
||||
assert model_info["supported_endpoints"] == ["/v1/images/generations"]
|
||||
assert model_info["supported_modalities"] == ["text"]
|
||||
assert model_info["supported_output_modalities"] == ["image"]
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"model_name,expected_total_cost",
|
||||
[
|
||||
("MAI-Image-2", 0.038792),
|
||||
("MAI-Image-2e", 0.024968),
|
||||
],
|
||||
)
|
||||
def test_azure_ai_mai_image_cost_calculator(
|
||||
use_local_model_cost_map, model_name: str, expected_total_cost: float
|
||||
):
|
||||
from litellm.llms.azure_ai.image_generation.cost_calculator import cost_calculator
|
||||
from litellm.types.utils import (
|
||||
ImageResponse,
|
||||
ImageUsage,
|
||||
ImageUsageInputTokensDetails,
|
||||
)
|
||||
|
||||
image_response = ImageResponse(
|
||||
created=123,
|
||||
data=[{"b64_json": "abc", "revised_prompt": "A photorealistic mountain lake"}],
|
||||
usage=ImageUsage(
|
||||
input_tokens=1000,
|
||||
input_tokens_details=ImageUsageInputTokensDetails(
|
||||
image_tokens=0,
|
||||
text_tokens=1000,
|
||||
),
|
||||
output_tokens=1024,
|
||||
total_tokens=2024,
|
||||
),
|
||||
)
|
||||
|
||||
total_cost = cost_calculator(model=model_name, image_response=image_response)
|
||||
|
||||
assert total_cost == pytest.approx(expected_total_cost)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("model_name", ["MAI-Image-2", "MAI-Image-2e"])
|
||||
def test_azure_ai_mai_image_generation_config(model_name: str):
|
||||
from litellm.llms.azure_ai.image_generation import (
|
||||
AzureFoundryMAIImageGenerationConfig,
|
||||
get_azure_ai_image_generation_config,
|
||||
)
|
||||
|
||||
config = get_azure_ai_image_generation_config(model_name)
|
||||
|
||||
assert isinstance(config, AzureFoundryMAIImageGenerationConfig)
|
||||
assert (
|
||||
config.get_complete_url(
|
||||
model=model_name,
|
||||
api_base="https://example.services.ai.azure.com",
|
||||
api_key=None,
|
||||
optional_params={},
|
||||
litellm_params={"api_version": "preview"},
|
||||
)
|
||||
== "https://example.services.ai.azure.com/mai/v1/images/generations?api-version=preview"
|
||||
)
|
||||
|
||||
|
||||
def test_azure_ai_mai_image_generation_request_maps_size():
|
||||
from litellm.llms.azure_ai.image_generation.mai_transformation import (
|
||||
AzureFoundryMAIImageGenerationConfig,
|
||||
)
|
||||
|
||||
request = AzureFoundryMAIImageGenerationConfig().transform_image_generation_request(
|
||||
model="MAI-Image-2e",
|
||||
prompt="A glowing jellyfish in a glass ocean",
|
||||
optional_params={"size": "768x1024"},
|
||||
litellm_params={},
|
||||
headers={},
|
||||
)
|
||||
|
||||
assert request == {
|
||||
"model": "MAI-Image-2e",
|
||||
"prompt": "A glowing jellyfish in a glass ocean",
|
||||
"width": 768,
|
||||
"height": 1024,
|
||||
}
|
||||
|
||||
|
||||
@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
|
||||
):
|
||||
import litellm
|
||||
from litellm.images.main import image_generation
|
||||
from litellm.llms.azure_ai.image_generation.mai_transformation import (
|
||||
AzureFoundryMAIImageGenerationConfig,
|
||||
)
|
||||
|
||||
mock_response = litellm.ImageResponse(
|
||||
created=123,
|
||||
data=[{"b64_json": "abc", "revised_prompt": "A bright desert sunrise"}],
|
||||
)
|
||||
mock_image_generation_handler.return_value = mock_response
|
||||
|
||||
response = image_generation(
|
||||
model="azure_ai/MAI-Image-2e",
|
||||
prompt="A bright desert sunrise",
|
||||
api_base="https://example.services.ai.azure.com",
|
||||
api_key="test-key",
|
||||
size="1024x1024",
|
||||
)
|
||||
|
||||
assert response == mock_response
|
||||
mock_azure_image_generation.assert_not_called()
|
||||
mock_image_generation_handler.assert_called_once()
|
||||
call_kwargs = mock_image_generation_handler.call_args.kwargs
|
||||
assert isinstance(
|
||||
call_kwargs["image_generation_provider_config"],
|
||||
AzureFoundryMAIImageGenerationConfig,
|
||||
)
|
||||
assert call_kwargs["image_generation_optional_request_params"] == {
|
||||
"size": "1024x1024"
|
||||
}
|
||||
assert call_kwargs["litellm_params"]["api_base"] == (
|
||||
"https://example.services.ai.azure.com"
|
||||
)
|
||||
Loading…
Add table
Reference in a new issue