Files
EvoScientist/tests/test_image_gen_openai_adapter.py
T

145 lines
4.7 KiB
Python

"""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]