feat(image-gen): add OpenAI-compatible image adapter with responses fallback

This commit is contained in:
m4
2026-07-23 21:28:38 +08:00
parent b0f9a8d785
commit 12e4f34005
4 changed files with 457 additions and 0 deletions
@@ -0,0 +1,6 @@
"""Image generation provider adapters."""
from .base import ImageGenAdapter, ImageGenError
from .openai import OpenAIImageAdapter
__all__ = ["ImageGenAdapter", "ImageGenError", "OpenAIImageAdapter"]
+104
View File
@@ -0,0 +1,104 @@
"""Shared adapter protocol, errors, and response-parsing helpers."""
from __future__ import annotations
import base64
import binascii
import re
from collections.abc import Awaitable, Callable
from pathlib import Path
from typing import Any, Protocol
from urllib.parse import urlparse
class ImageGenError(Exception):
"""Provider or configuration failure.
The message must be safe to show an agent or a browser: no secrets, no
request headers, no URLs carrying credentials.
"""
class ImageGenAdapter(Protocol):
"""Vendor-specific image generation/editing."""
async def generate(
self, *, prompt: str, size: str, quality: str, background: str, n: int
) -> list[bytes]: ...
async def edit(
self,
*,
image: Path,
mask: Path | None,
prompt: str,
size: str,
quality: str,
) -> list[bytes]: ...
def strip_data_uri(value: str) -> str:
if value.startswith("data:image/"):
return value.split(",", 1)[1]
return value
def looks_like_url(value: str) -> bool:
parsed = urlparse(value)
return parsed.scheme in {"http", "https"} and bool(parsed.netloc)
def looks_like_base64_image(value: str) -> bool:
if value.startswith("data:image/"):
return True
if len(value) < 16 or not re.fullmatch(r"[A-Za-z0-9+/=\s]+", value):
return False
try:
decoded = base64.b64decode(value, validate=False)
except (binascii.Error, ValueError):
return False
return decoded.startswith((b"\x89PNG", b"\xff\xd8\xff", b"RIFF", b"GIF8")) or len(
decoded
) > 256
def collect_image_values(
value: Any, *, b64_values: list[str], urls: list[str]
) -> None:
"""Recursively collect base64 payloads and URLs from a provider response."""
if isinstance(value, dict):
for key, item in value.items():
normalized = key.lower()
if isinstance(item, str):
if normalized in {"b64_json", "image_base64", "base64"}:
b64_values.append(item)
elif normalized == "result" and looks_like_base64_image(item):
b64_values.append(item)
elif normalized in {"url", "image_url"} and looks_like_url(item):
urls.append(item)
elif item.startswith("data:image/"):
b64_values.append(item)
else:
collect_image_values(item, b64_values=b64_values, urls=urls)
elif isinstance(value, list):
for item in value:
collect_image_values(item, b64_values=b64_values, urls=urls)
async def extract_image_bytes_from_values(
b64_values: list[str],
urls: list[str],
*,
download: Callable[[str], Awaitable[bytes]],
) -> list[bytes]:
images: list[bytes] = []
for value in b64_values:
try:
images.append(base64.b64decode(strip_data_uri(value), validate=False))
except (binascii.Error, ValueError):
continue
for url in urls:
images.append(await download(url))
if not images:
raise ImageGenError("Image provider returned no image data")
return images
+203
View File
@@ -0,0 +1,203 @@
"""OpenAI-compatible image generation/editing adapter."""
from __future__ import annotations
import asyncio
import mimetypes
from pathlib import Path
from typing import Any
import httpx
from ..config import ImageModelEntry
from .base import (
ImageGenError,
collect_image_values,
extract_image_bytes_from_values,
)
_TRANSIENT_STATUSES = {429, 500, 502, 503, 504}
_RETRY_DELAYS = (0.0, 1.0, 3.0)
_FALLBACK_MARKERS = (
"requires an image model",
"unsupported model",
"not supported",
"unknown model",
"model_not_found",
"not found",
"invalid endpoint",
)
class OpenAIImageAdapter:
"""Calls ``{base}/v1/images/*``; falls back to the Responses API."""
def __init__(
self,
entry: ImageModelEntry,
*,
client: httpx.AsyncClient | None = None,
timeout: float = 120.0,
) -> None:
self.entry = entry
self._client = client
self._timeout = timeout
# -- helpers -------------------------------------------------------------
def _api_key(self) -> str:
key = self.entry.resolved_api_key()
if not key:
raise ImageGenError(
f"Image model {self.entry.id!r} has no API key; "
"set the environment variable referenced in image_generation."
)
return key
def _endpoint_base(self) -> str:
base = (
self.entry.resolved_base_url() or "https://api.openai.com"
).rstrip("/")
return base if base.endswith("/v1") else f"{base}/v1"
def _raise_for_status(self, resp: httpx.Response) -> None:
if resp.status_code < 400:
return
detail = f"Image provider returned HTTP {resp.status_code}"
try:
payload = resp.json()
error = payload.get("error") if isinstance(payload, dict) else None
message = error.get("message") if isinstance(error, dict) else error
if isinstance(message, str) and message:
detail = f"{detail}: {message[:300]}"
except Exception:
pass
# Never include request headers, the URL, or the API key.
raise ImageGenError(detail)
async def _post_json(self, url: str, payload: dict[str, Any]) -> dict[str, Any]:
headers = {"Authorization": f"Bearer {self._api_key()}"}
last_response: httpx.Response | None = None
for delay in _RETRY_DELAYS:
if delay:
await asyncio.sleep(delay)
try:
if self._client is not None:
resp = await self._client.post(url, headers=headers, json=payload)
else:
async with httpx.AsyncClient(timeout=self._timeout) as client:
resp = await client.post(url, headers=headers, json=payload)
except httpx.TimeoutException:
raise ImageGenError(
f"Image provider request timed out after {self._timeout:g}s"
) from None
except httpx.RequestError as exc:
raise ImageGenError(
f"Image provider request failed: {exc.__class__.__name__}"
) from None
last_response = resp
if resp.status_code not in _TRANSIENT_STATUSES:
self._raise_for_status(resp)
return resp.json()
assert last_response is not None
self._raise_for_status(last_response)
raise ImageGenError("Image provider request failed")
async def _download(self, url: str) -> bytes:
headers = {"Authorization": f"Bearer {self._api_key()}"}
if self._client is not None:
resp = await self._client.get(url, headers=headers)
else:
async with httpx.AsyncClient(timeout=self._timeout) as client:
resp = await client.get(url, headers=headers)
self._raise_for_status(resp)
content_type = resp.headers.get("content-type", "")
if content_type and not content_type.startswith("image/"):
raise ImageGenError("Image provider returned a non-image URL")
return resp.content
async def _extract(self, response: dict[str, Any]) -> list[bytes]:
b64_values: list[str] = []
urls: list[str] = []
collect_image_values(response, b64_values=b64_values, urls=urls)
return await extract_image_bytes_from_values(
b64_values, urls, download=self._download
)
# -- adapter API ---------------------------------------------------------
async def generate(
self, *, prompt: str, size: str, quality: str, background: str, n: int
) -> list[bytes]:
base = self._endpoint_base()
payload: dict[str, Any] = {
"model": self.entry.id,
"prompt": prompt,
"size": size,
"quality": quality,
"n": n,
"response_format": "b64_json",
**self.entry.params,
}
if background and background != "auto":
payload["background"] = background
try:
response = await self._post_json(f"{base}/images/generations", payload)
except ImageGenError as exc:
message = str(exc).lower()
if not any(marker in message for marker in _FALLBACK_MARKERS):
raise
tool: dict[str, Any] = {"type": "image_generation", "size": size}
if quality and quality != "auto":
tool["quality"] = quality
if background and background != "auto":
tool["background"] = background
response = await self._post_json(
f"{base}/responses",
{
"model": self.entry.id,
"input": prompt,
"tools": [tool],
"tool_choice": {"type": "image_generation"},
"metadata": {"image_count": str(n)},
},
)
return await self._extract(response)
async def edit(
self,
*,
image: Path,
mask: Path | None,
prompt: str,
size: str,
quality: str,
) -> list[bytes]:
base = self._endpoint_base()
data = {
"model": self.entry.id,
"prompt": prompt,
"size": size,
"quality": quality,
"response_format": "b64_json",
}
image_mime = mimetypes.guess_type(image.name)[0] or "application/octet-stream"
files: list[tuple[str, tuple[str, Any, str]]] = []
with image.open("rb") as image_fh:
files.append(("image", (image.name, image_fh.read(), image_mime)))
if mask is not None:
mask_mime = mimetypes.guess_type(mask.name)[0] or "application/octet-stream"
with mask.open("rb") as mask_fh:
files.append(("mask", (mask.name, mask_fh.read(), mask_mime)))
headers = {"Authorization": f"Bearer {self._api_key()}"}
if self._client is not None:
resp = await self._client.post(
f"{base}/images/edits", headers=headers, data=data, files=files
)
else:
async with httpx.AsyncClient(timeout=self._timeout) as client:
resp = await client.post(
f"{base}/images/edits", headers=headers, data=data, files=files
)
self._raise_for_status(resp)
return await self._extract(resp.json())
+144
View File
@@ -0,0 +1,144 @@
"""Tests for the OpenAI-compatible image adapter."""
import base64
import httpx
import pytest
from EvoScientist.image_gen.adapters.base import ImageGenError
from EvoScientist.image_gen.adapters.openai import OpenAIImageAdapter
from EvoScientist.image_gen.config import ImageModelEntry
PNG_BYTES = b"\x89PNG\r\n\x1a\n" + b"fake-png-payload"
B64 = base64.b64encode(PNG_BYTES).decode()
def _entry(**overrides):
data = {"id": "gpt-image-2", "provider": "openai", "api_key": "sk-test"}
data.update(overrides)
return ImageModelEntry(**data)
def _adapter(handler, **entry_overrides):
client = httpx.AsyncClient(transport=httpx.MockTransport(handler))
return OpenAIImageAdapter(_entry(**entry_overrides), client=client)
async def test_generate_b64_success():
def handler(request: httpx.Request) -> httpx.Response:
assert request.url.path == "/v1/images/generations"
assert request.headers["Authorization"] == "Bearer sk-test"
import json
payload = json.loads(request.content)
assert payload["model"] == "gpt-image-2"
assert payload["prompt"] == "a cat"
assert payload["response_format"] == "b64_json"
return httpx.Response(200, json={"data": [{"b64_json": B64}]})
adapter = _adapter(handler)
images = await adapter.generate(
prompt="a cat", size="1024x1024", quality="auto", background="auto", n=1
)
assert images == [PNG_BYTES]
async def test_generate_downloads_url_response():
def handler(request: httpx.Request) -> httpx.Response:
if request.url.path == "/v1/images/generations":
return httpx.Response(
200, json={"data": [{"url": "https://files.example/img.png"}]}
)
if request.url.host == "files.example":
return httpx.Response(
200, content=PNG_BYTES, headers={"content-type": "image/png"}
)
return httpx.Response(404)
adapter = _adapter(handler)
images = await adapter.generate(
prompt="a cat", size="1024x1024", quality="auto", background="auto", n=1
)
assert images == [PNG_BYTES]
async def test_generate_falls_back_to_responses_api():
calls: list[str] = []
def handler(request: httpx.Request) -> httpx.Response:
calls.append(request.url.path)
if request.url.path == "/v1/images/generations":
return httpx.Response(
400, json={"error": {"message": "unsupported model for this endpoint"}}
)
if request.url.path == "/v1/responses":
return httpx.Response(
200,
json={
"output": [
{
"type": "image_generation_call",
"result": B64,
}
]
},
)
return httpx.Response(404)
adapter = _adapter(handler)
images = await adapter.generate(
prompt="a cat", size="1024x1024", quality="auto", background="auto", n=1
)
assert images == [PNG_BYTES]
assert calls == ["/v1/images/generations", "/v1/responses"]
async def test_transient_status_retried():
attempts = 0
def handler(request: httpx.Request) -> httpx.Response:
nonlocal attempts
attempts += 1
if attempts < 3:
return httpx.Response(503, json={"error": {"message": "busy"}})
return httpx.Response(200, json={"data": [{"b64_json": B64}]})
adapter = _adapter(handler)
images = await adapter.generate(
prompt="a cat", size="1024x1024", quality="auto", background="auto", n=1
)
assert images == [PNG_BYTES]
assert attempts == 3
async def test_error_message_never_contains_key():
def handler(request: httpx.Request) -> httpx.Response:
return httpx.Response(
401, json={"error": {"message": "invalid authentication"}}
)
adapter = _adapter(handler)
with pytest.raises(ImageGenError) as excinfo:
await adapter.generate(
prompt="a cat", size="1024x1024", quality="auto", background="auto", n=1
)
message = str(excinfo.value)
assert "sk-test" not in message
assert "401" in message
async def test_edit_posts_multipart(tmp_path):
image = tmp_path / "in.png"
image.write_bytes(PNG_BYTES)
def handler(request: httpx.Request) -> httpx.Response:
assert request.url.path == "/v1/images/edits"
assert b'name="image"' in request.content
assert b'name="prompt"' in request.content
return httpx.Response(200, json={"data": [{"b64_json": B64}]})
adapter = _adapter(handler)
images = await adapter.edit(
image=image, mask=None, prompt="make it blue", size="1024x1024", quality="auto"
)
assert images == [PNG_BYTES]