import base64 import importlib.util import json import sys import types import unittest from pathlib import Path from PIL import Image ROOT = Path(__file__).resolve().parents[1] def _load_client(): if importlib.util.find_spec("numpy") is None: numpy_stub = types.ModuleType("numpy") numpy_stub.float32 = float sys.modules["numpy"] = numpy_stub if importlib.util.find_spec("torch") is None: torch_stub = types.ModuleType("torch") torch_stub.Tensor = object sys.modules["torch"] = torch_stub comfy = types.ModuleType("comfy") comfy.__path__ = [] model_management = types.ModuleType("comfy.model_management") model_management.processing_interrupted = lambda: False model_management.InterruptProcessingException = RuntimeError sys.modules[comfy.__name__] = comfy sys.modules[model_management.__name__] = model_management package = types.ModuleType("comfyui_o1key") package.__path__ = [str(ROOT)] clients_package = types.ModuleType("comfyui_o1key.clients") clients_package.__path__ = [str(ROOT / "clients")] utils_package = types.ModuleType("comfyui_o1key.utils") utils_package.__path__ = [str(ROOT / "utils")] sys.modules[package.__name__] = package sys.modules[clients_package.__name__] = clients_package sys.modules[utils_package.__name__] = utils_package config = types.ModuleType("comfyui_o1key.utils.config") config.get_api_key_or_raise = lambda *_args, **_kwargs: "secret" config.get_base_url_by_route = lambda _route=None: "https://api.o1key.cn" sys.modules[config.__name__] = config image_utils = types.ModuleType("comfyui_o1key.utils.image_utils") image_utils.tensor_to_pil = lambda value: [value] sys.modules[image_utils.__name__] = image_utils spec = importlib.util.spec_from_file_location( "comfyui_o1key.clients.grok_image_client", ROOT / "clients" / "grok_image_client.py", ) module = importlib.util.module_from_spec(spec) sys.modules[spec.name] = module spec.loader.exec_module(module) return module GROK = _load_client() class GrokReferenceTests(unittest.TestCase): def _build(self, images): return GROK.GrokImageClient._build_edit_body( prompt="combine references", model="Grok Image Pro", aspect_ratio="auto", resolution="1k", image_list=images, ) def test_single_reference_keeps_legacy_image_field(self): body = self._build([Image.new("RGB", (8, 8), "red")]) self.assertIn("image", body) self.assertNotIn("images", body) self.assertTrue(base64.b64decode(body["image"]).startswith(b"\x89PNG")) def test_three_references_use_multi_image_field(self): body = self._build([ Image.new("RGB", (8, 8), "red"), Image.new("RGB", (8, 8), "green"), Image.new("RGB", (8, 8), "blue"), ]) self.assertNotIn("image", body) self.assertEqual(len(body["images"]), 3) for item in body["images"]: self.assertEqual(item["type"], "image_url") self.assertTrue(item["url"].startswith("data:image/png;base64,")) self.assertLessEqual( len(json.dumps(body, ensure_ascii=False).encode("utf-8")), GROK._MAX_BODY_BYTES, ) def test_more_than_three_references_are_rejected(self): images = [Image.new("RGB", (2, 2)) for _ in range(4)] with self.assertRaisesRegex(ValueError, "最多支持 3 张"): self._build(images) if __name__ == "__main__": unittest.main(verbosity=2)