feat(image-gen): add Gemini (Imagen) image adapter
This commit is contained in:
@@ -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"]
|
||||
|
||||
@@ -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."
|
||||
)
|
||||
@@ -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"
|
||||
)
|
||||
Reference in New Issue
Block a user