215 lines
7.6 KiB
Python
215 lines
7.6 KiB
Python
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").glob("*.meta.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()
|