feat(image-gen): add OpenAI-compatible image adapter with responses fallback
This commit is contained in:
@@ -0,0 +1,6 @@
|
||||
"""Image generation provider adapters."""
|
||||
|
||||
from .base import ImageGenAdapter, ImageGenError
|
||||
from .openai import OpenAIImageAdapter
|
||||
|
||||
__all__ = ["ImageGenAdapter", "ImageGenError", "OpenAIImageAdapter"]
|
||||
@@ -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
|
||||
@@ -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())
|
||||
@@ -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]
|
||||
Reference in New Issue
Block a user