82 lines
2.7 KiB
Python
82 lines
2.7 KiB
Python
"""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"
|
|
)
|