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()