145 lines
4.7 KiB
Python
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]
|