zcbot/tests/test_gpt_image.py

215 lines
7.6 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

import base64
import json
import struct
import tempfile
import unittest
import uuid
from pathlib import Path
from unittest.mock import patch
import httpx
from core.ark_client import ArkClient, ArkConfig
from tools.gpt_image import GptImageTool
def _png_stub(width: int = 1536, height: int = 864) -> bytes:
return b"\x89PNG\r\n\x1a\n" + b"\x00\x00\x00\rIHDR" + struct.pack(">II", width, height)
class _FakeArkClient:
json_call = None
multipart_call = None
def __init__(self, *_args, **_kwargs):
pass
def __enter__(self):
return self
def __exit__(self, *_args):
pass
def post_json(self, endpoint, body, *, timeout_s=None):
type(self).json_call = (endpoint, body, timeout_s)
return {
"data": [{"b64_json": base64.b64encode(_png_stub()).decode("ascii")}],
"usage": {"output_tokens": 123},
}
def post_multipart(self, endpoint, data, files, *, timeout_s=None):
type(self).multipart_call = (endpoint, data, files, timeout_s)
return {
"data": [{"b64_json": base64.b64encode(_png_stub()).decode("ascii")}],
}
class ArkClientMultipartTests(unittest.TestCase):
def test_json_and_multipart_set_their_own_content_types(self):
seen = []
def handler(request):
seen.append(request)
return httpx.Response(200, json={"ok": True})
client = ArkClient(
ArkConfig(api_key="test", base_url="https://example.test/v1", raw={})
)
client._client.close()
client._client = httpx.Client(
base_url="https://example.test/v1",
headers={"Authorization": "Bearer test"},
transport=httpx.MockTransport(handler),
)
try:
client.post_json("/images/generations", {"prompt": "draw"})
client.post_multipart(
"/images/edits",
{"prompt": "edit"},
{"image": ("reference.png", b"png-bytes", "image/png")},
)
finally:
client.close()
self.assertEqual(seen[0].headers["content-type"], "application/json")
self.assertTrue(seen[1].headers["content-type"].startswith("multipart/form-data;"))
self.assertIn(b'name="prompt"', seen[1].content)
self.assertIn(b'filename="reference.png"', seen[1].content)
self.assertIn(b"png-bytes", seen[1].content)
class GptImageToolTests(unittest.TestCase):
def setUp(self):
_FakeArkClient.json_call = None
_FakeArkClient.multipart_call = None
self.tmp = tempfile.TemporaryDirectory()
self.root = Path(self.tmp.name)
self.working_dir = self.root / "task"
self.working_dir.mkdir()
self.cfg = {
"model_id": "gpt-image-2",
"endpoint": "/images/generations",
"edit_endpoint": "/images/edits",
"default_size": "auto",
"default_quality": "auto",
"request_timeout_s": 300,
"price_cny_per_image": 0,
}
self.tool = GptImageTool(
gw_cfg=ArkConfig(api_key="test", base_url="https://example.test/v1", raw={}),
image_variant_cfg=self.cfg,
variant_key="gpt_image",
working_dir=self.working_dir,
task_id=uuid.uuid4(),
user_id=uuid.uuid4(),
base_dir=self.working_dir,
user_root=self.root,
daily_limit=0,
)
def tearDown(self):
self.tmp.cleanup()
def _execute(self, **kwargs):
with (
patch("tools.gpt_image.ArkClient", _FakeArkClient),
patch("tools.gpt_image.quota_gate", return_value=""),
patch("tools.gpt_image.record_usage_safe"),
):
return self.tool.execute(**kwargs)
def test_text_to_image_forwards_size_and_quality(self):
result = self._execute(
prompt="draw a materials lab",
size="1536x864",
quality="high",
)
self.assertTrue(result.startswith("[gpt_image]"))
endpoint, body, _timeout = _FakeArkClient.json_call
self.assertEqual(endpoint, "/images/generations")
self.assertEqual(body["size"], "1536x864")
self.assertEqual(body["quality"], "high")
self.assertIn("size=1536x864", result)
self.assertIn("quality=high", result)
self.assertIsNone(_FakeArkClient.multipart_call)
def test_image_edit_uses_multipart_and_records_derivation(self):
reference = self.working_dir / "reference.png"
reference.write_bytes(_png_stub(180, 252))
result = self._execute(
prompt="add a blue border",
reference_images=["reference.png"],
size="1024x1024",
quality="low",
)
self.assertIsNone(_FakeArkClient.json_call)
endpoint, form, files, timeout = _FakeArkClient.multipart_call
self.assertEqual(endpoint, "/images/edits")
self.assertEqual(timeout, 300)
self.assertEqual(form["model"], "gpt-image-2")
self.assertEqual(form["prompt"], "add a blue border")
self.assertEqual(form["n"], "1")
self.assertEqual(form["size"], "1024x1024")
self.assertEqual(form["quality"], "low")
self.assertEqual(form["response_format"], "b64_json")
filename, raw, mime = files["image"]
self.assertEqual(filename, "reference.png")
self.assertEqual(raw, reference.read_bytes())
self.assertEqual(mime, "image/png")
self.assertIn("mode=i2i", result)
self.assertIn("reference=", result)
meta_path = next((self.working_dir / "figures" / ".meta").glob("*.json"))
meta = json.loads(meta_path.read_text(encoding="utf-8"))
self.assertEqual(meta["mode"], "i2i")
self.assertEqual(meta["reference_images"], ["task/reference.png"])
def test_image_edit_rejects_multiple_or_missing_references(self):
result = self._execute(
prompt="edit",
reference_images=["one.png", "two.png"],
)
self.assertIn("仅支持单张", result)
self.assertIsNone(_FakeArkClient.multipart_call)
result = self._execute(prompt="edit", reference_images=["missing.png"])
self.assertIn("图片找不到或越界", result)
self.assertIsNone(_FakeArkClient.multipart_call)
def test_image_edit_accepts_data_url_response_fallback(self):
reference = self.working_dir / "reference.png"
reference.write_bytes(_png_stub(180, 252))
encoded = base64.b64encode(_png_stub()).decode("ascii")
with patch.object(
_FakeArkClient,
"post_multipart",
return_value={"data": [{"url": f"data:image/png;base64,{encoded}"}]},
):
result = self._execute(
prompt="edit",
reference_images=["reference.png"],
)
self.assertTrue(result.startswith("[gpt_image]"))
def test_size_validation(self):
self.assertEqual(self.tool._normalize_size("auto"), ("auto", ""))
self.assertEqual(self.tool._normalize_size("1536×864"), ("1536x864", ""))
for value in ("1000x1000", "4096x1024", "3072x512", "640x640", "bad"):
with self.subTest(value=value):
_size, error = self.tool._normalize_size(value)
self.assertTrue(error.startswith("[Error]"))
def test_quality_validation(self):
result = self._execute(prompt="draw", quality="ultra")
self.assertEqual(result, "[Error] quality 必须是 auto / low / medium / high")
self.assertIsNone(_FakeArkClient.json_call)
if __name__ == "__main__":
unittest.main()