feat(image-gen): add Gemini (Imagen) image adapter

This commit is contained in:
m4
2026-07-23 21:53:40 +08:00
parent cc9dfb1cc9
commit 802f71bd46
3 changed files with 207 additions and 1 deletions
+2 -1
View File
@@ -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"]
+124
View File
@@ -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."
)
+81
View File
@@ -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"
)