diff --git a/EvoScientist/image_gen/adapters/__init__.py b/EvoScientist/image_gen/adapters/__init__.py new file mode 100644 index 0000000..7a5271c --- /dev/null +++ b/EvoScientist/image_gen/adapters/__init__.py @@ -0,0 +1,6 @@ +"""Image generation provider adapters.""" + +from .base import ImageGenAdapter, ImageGenError +from .openai import OpenAIImageAdapter + +__all__ = ["ImageGenAdapter", "ImageGenError", "OpenAIImageAdapter"] diff --git a/EvoScientist/image_gen/adapters/base.py b/EvoScientist/image_gen/adapters/base.py new file mode 100644 index 0000000..6f2db03 --- /dev/null +++ b/EvoScientist/image_gen/adapters/base.py @@ -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 diff --git a/EvoScientist/image_gen/adapters/openai.py b/EvoScientist/image_gen/adapters/openai.py new file mode 100644 index 0000000..4a586a8 --- /dev/null +++ b/EvoScientist/image_gen/adapters/openai.py @@ -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()) diff --git a/tests/test_image_gen_openai_adapter.py b/tests/test_image_gen_openai_adapter.py new file mode 100644 index 0000000..5de5813 --- /dev/null +++ b/tests/test_image_gen_openai_adapter.py @@ -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]