Files
EvoScientist/tests/test_image_gen_openai_adapter.py
m4 cc9dfb1cc9 fix(image-gen): address review findings in OpenAI image adapter
Scope download Authorization header to the provider origin, translate
httpx errors in _download/edit into safe ImageGenError messages, and
prevent entry params from clobbering core payload keys; also harden
strip_data_uri against malformed values and drop dead code.

Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com>
2026-07-23 21:44:57 +08:00

254 lines
8.5 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]
async def test_download_to_foreign_host_omits_auth_header():
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":
assert "Authorization" not in request.headers
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_download_to_same_host_sends_auth_header():
def handler(request: httpx.Request) -> httpx.Response:
if request.url.path == "/v1/images/generations":
return httpx.Response(
200,
json={"data": [{"url": "https://api.example.com/files/img.png"}]},
)
if request.url.path == "/files/img.png":
assert request.headers["Authorization"] == "Bearer sk-test"
return httpx.Response(
200, content=PNG_BYTES, headers={"content-type": "image/png"}
)
return httpx.Response(404)
adapter = _adapter(handler, base_url="https://api.example.com")
images = await adapter.generate(
prompt="a cat", size="1024x1024", quality="auto", background="auto", n=1
)
assert images == [PNG_BYTES]
async def test_download_network_error_wrapped():
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"}]}
)
raise httpx.ConnectError("boom", request=request)
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 "https://" not in message
assert "files.example" not in message
async def test_edit_network_error_wrapped(tmp_path):
image = tmp_path / "in.png"
image.write_bytes(PNG_BYTES)
def handler(request: httpx.Request) -> httpx.Response:
raise httpx.ConnectError("boom", request=request)
adapter = _adapter(handler)
with pytest.raises(ImageGenError) as excinfo:
await adapter.edit(
image=image,
mask=None,
prompt="make it blue",
size="1024x1024",
quality="auto",
)
assert "https://" not in str(excinfo.value)
async def test_params_cannot_clobber_core_keys():
import json
def handler(request: httpx.Request) -> httpx.Response:
payload = json.loads(request.content)
assert payload["model"] == "gpt-image-2"
assert payload["response_format"] == "b64_json"
return httpx.Response(200, json={"data": [{"b64_json": B64}]})
adapter = _adapter(
handler, params={"model": "evil", "response_format": "url"}
)
images = await adapter.generate(
prompt="a cat", size="1024x1024", quality="auto", background="auto", n=1
)
assert images == [PNG_BYTES]
async def test_malformed_data_uri_ignored():
def handler(request: httpx.Request) -> httpx.Response:
return httpx.Response(
200, json={"data": [{"b64_json": "data:image/png"}, {"b64_json": B64}]}
)
adapter = _adapter(handler)
images = await adapter.generate(
prompt="a cat", size="1024x1024", quality="auto", background="auto", n=1
)
assert PNG_BYTES in images