Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
94 changes: 74 additions & 20 deletions core/provider/src/openai/model_clients.py
Original file line number Diff line number Diff line change
@@ -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
Expand All @@ -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
Expand All @@ -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)

Expand All @@ -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":
Expand All @@ -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
Expand Down Expand Up @@ -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 ────────
Expand Down
129 changes: 129 additions & 0 deletions tests/test_openai_image_client.py
Original file line number Diff line number Diff line change
@@ -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})")
Expand Down Expand Up @@ -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()
Loading