diff --git a/core/provider/src/openai/model_clients.py b/core/provider/src/openai/model_clients.py index 4db5f445..4cc4ea67 100644 --- a/core/provider/src/openai/model_clients.py +++ b/core/provider/src/openai/model_clients.py @@ -1,8 +1,11 @@ -from openai import AsyncOpenAI, APIStatusError, APITimeoutError, APIConnectionError +import base64 +import binascii import re import time from typing import Optional, Union +from openai import AsyncOpenAI, APIStatusError, APITimeoutError, APIConnectionError + from core.provider import ModelInfo from core.provider import LLMModelClient, ImageModelClient, EmbeddingModelClient from core.provider.llm_model import LLMRequest, LLMResponse @@ -15,8 +18,9 @@ class OpenAIImageClient(ImageModelClient): """Image generation client with two on-the-wire shapes. - Original ``/v1/images/generations`` is the OpenAI standard image - endpoint. ``/v1/chat/completions`` mode handles the increasingly + The OpenAI image mode uses ``/v1/images/generations`` for text-to-image + and ``/v1/images/edits`` for image-to-image. ``/v1/chat/completions`` + mode handles the increasingly common pattern where multimodal chat models (e.g. some Tongyi / Doubao / 3rd-party-compatible APIs) return images embedded in chat content — either as Markdown ``![alt](url)``, data URIs, or @@ -38,6 +42,12 @@ class OpenAIImageClient(ImageModelClient): re.DOTALL, ) + _IMAGE_EXTENSIONS = { + "image/png": ".png", + "image/jpeg": ".jpg", + "image/webp": ".webp", + } + def __init__(self, model: ModelInfo): super().__init__(model) @@ -58,6 +68,50 @@ def _build_client(self) -> AsyncOpenAI: timeout=timeout, ) + @staticmethod + def _image_from_response(images_response, operation: str) -> Image: + """Return the first image from either URL or Base64 API output.""" + if not images_response.data: + raise ValueError(f"{operation} API returned empty data") + + result = images_response.data[0] + url = getattr(result, "url", None) + if url: + return Image(image=url) + + b64_json = getattr(result, "b64_json", None) + if b64_json: + return Image(image=b64_json) + + raise ValueError(f"{operation} API returned neither url nor b64_json") + + @classmethod + def _image_file_from_data_url( + cls, data_url: str, index: int + ) -> tuple[str, bytes, str]: + """Convert a Base64 image Data URL to an OpenAI upload tuple.""" + header, separator, encoded_data = data_url.partition(",") + if not separator or not header.startswith("data:image/"): + raise ValueError("Image input must be an image Data URL") + + mime_type, *parameters = header.removeprefix("data:").split(";") + mime_type = mime_type.lower() + if "base64" not in {parameter.lower() for parameter in parameters}: + raise ValueError("Image Data URL must contain Base64 data") + + extension = cls._IMAGE_EXTENSIONS.get(mime_type) + if extension is None: + raise ValueError(f"Unsupported image MIME type: {mime_type}") + + try: + image_bytes = base64.b64decode(encoded_data, validate=True) + except (binascii.Error, ValueError) as e: + raise ValueError("Image Data URL contains invalid Base64 data") from e + if not image_bytes: + raise ValueError("Image Data URL contains empty image data") + + return f"image_{index}{extension}", image_bytes, mime_type + async def text_to_image(self, prompt) -> Image: endpoint = self.model.model_config.get("endpoint", "v1/image") if endpoint == "v1/chat": @@ -74,12 +128,8 @@ async def _text_to_image_via_generations(self, prompt: str) -> Image: model=self.model.model_id, prompt=prompt, size=image_size if image_size else None, - response_format="url", - extra_body={"watermark": False}, ) - if not images_response.data: - raise ValueError("Image generation API returned empty data") - return Image(image=images_response.data[0].url) + return self._image_from_response(images_response, "Image generation") except (APIStatusError, APITimeoutError, APIConnectionError) as e: logger.error(f"Image generation API error: {e}") raise @@ -242,30 +292,34 @@ async def image_to_image(self, prompt: str, image: Union[Image, list[Image]]) -> endpoint = self.model.model_config.get("endpoint", "v1/image") if endpoint == "v1/chat": return await self._image_to_image_via_chat(prompt, image) - return await self._image_to_image_via_generations(prompt, image) + return await self._image_to_image_via_edits(prompt, image) - # ──────── image-to-image Mode A: /v1/images/generations ──────── + # ──────── image-to-image Mode A: /v1/images/edits ──────── - async def _image_to_image_via_generations(self, prompt: str, images: list[Image]) -> Image: + async def _image_to_image_via_edits(self, prompt: str, images: list[Image]) -> Image: client = self._build_client() image_size = self.model.model_config.get("size", None) - image_data_urls = [await img.to_data_url() for img in images] try: - images_response = await client.images.generate( + image_files = [ + self._image_file_from_data_url(await image.to_data_url(), index) + for index, image in enumerate(images) + ] + if not image_files: + raise ValueError("Image edit requires at least one input image") + image_input = image_files[0] if len(image_files) == 1 else image_files + + images_response = await client.images.edit( model=self.model.model_id, prompt=prompt, + image=image_input, size=image_size if image_size else None, - response_format="url", - extra_body={"watermark": False, "image": image_data_urls}, ) - if not images_response.data: - raise ValueError("Image-to-image generation API returned empty data") - return Image(image=images_response.data[0].url) + return self._image_from_response(images_response, "Image edit") except (APIStatusError, APITimeoutError, APIConnectionError) as e: - logger.error(f"Image-to-image generation API error: {e}") + logger.error(f"Image edit API error: {e}") raise except Exception as e: - logger.error(f"Image-to-image generation error: {e}") + logger.error(f"Image edit error: {e}") raise # ──────── image-to-image Mode B: /v1/chat/completions ──────── diff --git a/tests/test_openai_image_client.py b/tests/test_openai_image_client.py index da7bfe55..b41676f5 100644 --- a/tests/test_openai_image_client.py +++ b/tests/test_openai_image_client.py @@ -1,12 +1,31 @@ +import base64 from types import SimpleNamespace +from unittest.mock import AsyncMock +import pytest + +from core.chat.message_elements import Image from core.provider.src.openai.model_clients import OpenAIImageClient +PNG_BYTES = b"\x89PNG\r\n\x1a\n" +PNG_BASE64 = base64.b64encode(PNG_BYTES).decode() + + def _client(): return object.__new__(OpenAIImageClient) +def _configured_client(images_api): + client = _client() + client.model = SimpleNamespace( + model_id="gpt-image-1", + model_config={"endpoint": "v1/image", "size": "1024x1024"}, + ) + client._build_client = lambda: SimpleNamespace(images=images_api) + return client + + def test_extracts_markdown_data_url_from_chat_content(): data_url = "data:image/png;base64,iVBORw0KGgoAAAA" message = SimpleNamespace(content=f"![image]({data_url})") @@ -38,3 +57,113 @@ def test_extracts_markdown_https_url_from_chat_content(): assert image is not None assert image.image == "https://example.com/image.png" assert image.image_type == "url" + + +@pytest.mark.asyncio +async def test_v1_image_text_to_image_accepts_base64_response(): + generate = AsyncMock( + return_value=SimpleNamespace( + data=[SimpleNamespace(url=None, b64_json=PNG_BASE64)] + ) + ) + client = _configured_client(SimpleNamespace(generate=generate)) + + result = await client.text_to_image("draw a blue square") + + assert result.image == PNG_BASE64 + assert result.image_type == "base64" + call = generate.await_args.kwargs + assert call == { + "model": "gpt-image-1", + "prompt": "draw a blue square", + "size": "1024x1024", + } + + +@pytest.mark.asyncio +async def test_v1_image_uses_edits_and_accepts_base64_response(): + edit = AsyncMock( + return_value=SimpleNamespace( + data=[SimpleNamespace(url=None, b64_json=PNG_BASE64)] + ) + ) + client = _configured_client(SimpleNamespace(edit=edit)) + + result = await client.image_to_image( + "make it blue", + Image(image=f"data:image/png;base64,{PNG_BASE64}"), + ) + + assert result.image == PNG_BASE64 + assert result.image_type == "base64" + edit.assert_awaited_once() + call = edit.await_args.kwargs + assert call["model"] == "gpt-image-1" + assert call["prompt"] == "make it blue" + assert call["size"] == "1024x1024" + assert "response_format" not in call + assert "extra_body" not in call + assert call["image"] == ("image_0.png", PNG_BYTES, "image/png") + + +@pytest.mark.asyncio +async def test_v1_image_edit_accepts_url_response_and_multiple_images(): + edit = AsyncMock( + return_value=SimpleNamespace( + data=[ + SimpleNamespace( + url="https://example.com/edited.webp", + b64_json=None, + ) + ] + ) + ) + client = _configured_client(SimpleNamespace(edit=edit)) + jpeg_base64 = base64.b64encode(b"jpeg").decode() + webp_base64 = base64.b64encode(b"webp").decode() + + result = await client.image_to_image( + "combine them", + [ + Image(image=f"data:image/jpeg;base64,{jpeg_base64}"), + Image(image=f"data:image/webp;base64,{webp_base64}"), + ], + ) + + assert result.image == "https://example.com/edited.webp" + image_files = edit.await_args.kwargs["image"] + assert image_files == [ + ("image_0.jpg", b"jpeg", "image/jpeg"), + ("image_1.webp", b"webp", "image/webp"), + ] + + +@pytest.mark.parametrize( + "data_url, message", + [ + ("data:image/png,cG5n", "must contain Base64"), + ("data:image/png;base64,%%%", "invalid Base64"), + ("data:image/png;base64,", "empty image data"), + ("data:image/gif;base64,R0lG", "Unsupported image MIME type"), + ], +) +def test_image_file_from_data_url_rejects_invalid_input(data_url, message): + with pytest.raises(ValueError, match=message): + _client()._image_file_from_data_url(data_url, 0) + + +@pytest.mark.parametrize("data", [[], [SimpleNamespace(url=None, b64_json=None)]]) +def test_image_response_rejects_missing_image_data(data): + with pytest.raises(ValueError, match="returned"): + _client()._image_from_response(SimpleNamespace(data=data), "Image edit") + + +@pytest.mark.asyncio +async def test_image_edit_requires_input_image(): + edit = AsyncMock() + client = _configured_client(SimpleNamespace(edit=edit)) + + with pytest.raises(ValueError, match="at least one input image"): + await client.image_to_image("edit it", []) + + edit.assert_not_awaited()