From 802f71bd469d768f96596acc34cb8bee349b46ee Mon Sep 17 00:00:00 2001 From: m4 Date: Thu, 23 Jul 2026 21:53:40 +0800 Subject: [PATCH] feat(image-gen): add Gemini (Imagen) image adapter --- EvoScientist/image_gen/adapters/__init__.py | 3 +- EvoScientist/image_gen/adapters/gemini.py | 124 ++++++++++++++++++++ tests/test_image_gen_gemini_adapter.py | 81 +++++++++++++ 3 files changed, 207 insertions(+), 1 deletion(-) create mode 100644 EvoScientist/image_gen/adapters/gemini.py create mode 100644 tests/test_image_gen_gemini_adapter.py diff --git a/EvoScientist/image_gen/adapters/__init__.py b/EvoScientist/image_gen/adapters/__init__.py index 7a5271c..7df2a58 100644 --- a/EvoScientist/image_gen/adapters/__init__.py +++ b/EvoScientist/image_gen/adapters/__init__.py @@ -1,6 +1,7 @@ """Image generation provider adapters.""" from .base import ImageGenAdapter, ImageGenError +from .gemini import GeminiImageAdapter from .openai import OpenAIImageAdapter -__all__ = ["ImageGenAdapter", "ImageGenError", "OpenAIImageAdapter"] +__all__ = ["GeminiImageAdapter", "ImageGenAdapter", "ImageGenError", "OpenAIImageAdapter"] diff --git a/EvoScientist/image_gen/adapters/gemini.py b/EvoScientist/image_gen/adapters/gemini.py new file mode 100644 index 0000000..b9c2622 --- /dev/null +++ b/EvoScientist/image_gen/adapters/gemini.py @@ -0,0 +1,124 @@ +"""Gemini (Imagen) image generation adapter. + +Uses the Gemini API ``:predict`` endpoint. The API key travels in the +``x-goog-api-key`` header — never in the URL query — so it cannot leak +through logged URLs. +""" + +from __future__ import annotations + +import base64 +import binascii +from pathlib import Path +from typing import Any + +import httpx + +from ..config import ImageModelEntry +from .base import ImageGenError + +DEFAULT_BASE_URL = "https://generativelanguage.googleapis.com" + +# gpt-image-style sizes -> Imagen aspect ratios. +SIZE_TO_ASPECT = { + "1024x1024": "1:1", + "1536x1024": "16:9", + "1024x1536": "9:16", +} + + +class GeminiImageAdapter: + """Calls ``{base}/v1beta/models/{model}:predict``.""" + + def __init__( + self, + entry: ImageModelEntry, + *, + client: httpx.AsyncClient | None = None, + timeout: float = 120.0, + ) -> None: + self.entry = entry + self._client = client + self._timeout = timeout + + 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: + return (self.entry.resolved_base_url() or DEFAULT_BASE_URL).rstrip("/") + + 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 + raise ImageGenError(detail) + + async def generate( + self, *, prompt: str, size: str, quality: str, background: str, n: int + ) -> list[bytes]: + # `quality`/`background` are not Imagen parameters and are ignored. + url = f"{self._endpoint_base()}/v1beta/models/{self.entry.id}:predict" + payload: dict[str, Any] = { + "instances": [{"prompt": prompt}], + "parameters": { + "sampleCount": n, + "aspectRatio": SIZE_TO_ASPECT.get(size, "1:1"), + **self.entry.params, + }, + } + headers = {"x-goog-api-key": self._api_key()} + 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 + self._raise_for_status(resp) + predictions = resp.json().get("predictions") or [] + images: list[bytes] = [] + for item in predictions: + value = item.get("bytesBase64Encoded") if isinstance(item, dict) else None + if not value: + continue + try: + images.append(base64.b64decode(value, validate=False)) + except (binascii.Error, ValueError): + continue + if not images: + raise ImageGenError("Image provider returned no image data") + return images + + async def edit( + self, + *, + image: Path, + mask: Path | None, + prompt: str, + size: str, + quality: str, + ) -> list[bytes]: + raise ImageGenError( + f"Image model {self.entry.id!r} does not support edit." + ) diff --git a/tests/test_image_gen_gemini_adapter.py b/tests/test_image_gen_gemini_adapter.py new file mode 100644 index 0000000..9ef6640 --- /dev/null +++ b/tests/test_image_gen_gemini_adapter.py @@ -0,0 +1,81 @@ +"""Tests for the Gemini (Imagen) image adapter.""" + +import base64 +import json + +import httpx +import pytest + +from EvoScientist.image_gen.adapters.base import ImageGenError +from EvoScientist.image_gen.adapters.gemini import GeminiImageAdapter +from EvoScientist.image_gen.config import ImageModelEntry + +PNG_BYTES = b"\x89PNG\r\n\x1a\n" + b"fake-gemini-payload" +B64 = base64.b64encode(PNG_BYTES).decode() + + +def _entry(**overrides): + data = { + "id": "imagen-4.0-generate-001", + "provider": "gemini", + "api_key": "gem-key-1", + } + data.update(overrides) + return ImageModelEntry(**data) + + +def _adapter(handler, **entry_overrides): + client = httpx.AsyncClient(transport=httpx.MockTransport(handler)) + return GeminiImageAdapter(_entry(**entry_overrides), client=client) + + +async def test_generate_uses_header_auth_and_predict(): + def handler(request: httpx.Request) -> httpx.Response: + assert ( + request.url.path + == "/v1beta/models/imagen-4.0-generate-001:predict" + ) + # Key must travel in the header, never in the URL. + assert request.headers["x-goog-api-key"] == "gem-key-1" + assert "key" not in request.url.params + payload = json.loads(request.content) + assert payload["instances"] == [{"prompt": "a robot"}] + assert payload["parameters"]["sampleCount"] == 2 + assert payload["parameters"]["aspectRatio"] == "16:9" + return httpx.Response( + 200, json={"predictions": [{"bytesBase64Encoded": B64}] * 2} + ) + + adapter = _adapter(handler) + images = await adapter.generate( + prompt="a robot", size="1536x1024", quality="auto", background="auto", n=2 + ) + assert images == [PNG_BYTES, PNG_BYTES] + + +async def test_default_base_url(): + adapter = GeminiImageAdapter(_entry(), client=None) + assert adapter._endpoint_base() == "https://generativelanguage.googleapis.com" + + +async def test_error_sanitized(): + def handler(request: httpx.Request) -> httpx.Response: + return httpx.Response(403, json={"error": {"message": "forbidden"}}) + + adapter = _adapter(handler) + with pytest.raises(ImageGenError) as excinfo: + await adapter.generate( + prompt="x", size="1024x1024", quality="auto", background="auto", n=1 + ) + assert "gem-key-1" not in str(excinfo.value) + assert "403" in str(excinfo.value) + + +async def test_edit_not_supported(tmp_path): + image = tmp_path / "in.png" + image.write_bytes(PNG_BYTES) + adapter = GeminiImageAdapter(_entry(), client=None) + with pytest.raises(ImageGenError, match="edit"): + await adapter.edit( + image=image, mask=None, prompt="change", size="1024x1024", quality="auto" + )