Files
comfyui_o1key/tests/test_o1key_image_generator.py
T
Jony ba920f2b66 Publish current ComfyUI O1Key code baseline
Replace the prior release tree with the current plugin, frontend, tests, and documentation. Document retired node IDs and the public Gitea update source.
2026-09-24 19:56:48 +08:00

1215 lines
49 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.
"""Offline tests for the panel-style o1key image generator."""
import asyncio
import base64
import json
import os
import sys
import tempfile
import unittest
from io import BytesIO
from unittest.mock import AsyncMock, patch
import torch
from PIL import Image
PLUGIN_DIR = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
CUSTOM_NODES_DIR = os.path.dirname(PLUGIN_DIR)
COMFY_ROOT = os.path.dirname(CUSTOM_NODES_DIR)
sys.path.insert(0, COMFY_ROOT)
sys.path.insert(0, CUSTOM_NODES_DIR)
from comfy_api.latest import io # noqa: E402
from comfyui_o1key.clients.gpt_image_client import GptImageClient # noqa: E402
from comfyui_o1key.nodes.o1key_image_generator import ( # noqa: E402
O1keyImageGenerator,
O1keyImageSave,
_load_reference_tensors,
_parse_upload_manifest,
)
from comfyui_o1key.utils.image_utils import ( # noqa: E402
IMAGE_BATCH_MODE_CARTESIAN,
IMAGE_BATCH_MODE_GROUP_TO_MODELS,
IMAGE_BATCH_MODE_SINGLE_REFERENCES,
expand_image_generation_tasks,
pil_to_tensor,
tensor_to_pil,
)
from comfyui_o1key.utils.http_error import format_o1key_image_error # noqa: E402
from comfyui_o1key.utils.o1key_image_catalog import ( # noqa: E402
GPT_IMAGE_ASPECT_RATIO_OPTIONS,
MAX_UNIFIED_BATCH_IMAGES,
MAX_UNIFIED_REFERENCE_IMAGES,
UNIFIED_IMAGE_MODEL_OPTIONS,
resolve_gpt_image_size,
resolve_seedream_size,
)
class O1keyImageGeneratorTests(unittest.TestCase):
def test_model_order_matches_panel_and_default(self):
self.assertEqual(
list(UNIFIED_IMAGE_MODEL_OPTIONS),
[
"Nano Banana 2",
"Nano Banana Pro",
"gpt-image-2",
"gpt-image-2.5-sunburst",
"gpt-image-2.5-flare",
"Seedream 5.0 Pro",
"Nano Banana 2 Lite",
"Nano Banana",
],
)
def test_schema_is_panel_node_with_image_output(self):
schema = O1keyImageGenerator.define_schema()
schema.validate()
for item in schema.inputs:
if item.io_type != "COMBO":
continue
self.assertTrue(
all(isinstance(option, str) for option in item.options),
item.id,
)
if item.default is not None:
self.assertIsInstance(item.default, str, item.id)
self.assertTrue(asyncio.iscoroutinefunction(O1keyImageGenerator.execute))
self.assertEqual(schema.node_id, "O1keyImageGenerator")
self.assertEqual(schema.display_name, "o1key 图片生成")
self.assertTrue(schema.not_idempotent)
self.assertEqual(
[item.id for item in schema.outputs],
["IMAGE", "LAYERS", "LAYER_MASKS", "LAYER_INFO"],
)
self.assertEqual(
[item.id for item in schema.inputs],
[
"prompt", "模型", "模型线路", "思考等级", "分辨率",
"宽高比", "生图数量", "seed", "参考图清单",
"质量", "输出格式", "蒙版清单", "缩放图片", "背景",
"批量出图", "批量模式", "模特图清单",
"命名规则", "filename_prefix", "格式", "保存位置", "在线搜索", "图层拆分",
],
)
count_input = next(item for item in schema.inputs if item.id == "生图数量")
self.assertEqual(count_input.options, [str(value) for value in range(1, 10)])
self.assertEqual(count_input.default, "1")
model_input = next(item for item in schema.inputs if item.id == "模型")
self.assertIn("gpt-image-2", model_input.options)
self.assertIn("gpt-image-2.5-sunburst", model_input.options)
self.assertIn("gpt-image-2.5-flare", model_input.options)
self.assertIn("Seedream 5.0 Pro", model_input.options)
thinking_input = next(item for item in schema.inputs if item.id == "思考等级")
self.assertEqual(thinking_input.options, ["低", "高"])
self.assertEqual(thinking_input.default, "低")
resolution_input = next(item for item in schema.inputs if item.id == "分辨率")
self.assertEqual(resolution_input.options[0], "智能")
self.assertEqual(resolution_input.default, "智能")
resize_input = next(item for item in schema.inputs if item.id == "缩放图片")
self.assertEqual(resize_input.options, ["不缩放", "智能缩放"])
self.assertEqual(resize_input.default, "不缩放")
self.assertNotIn("色彩纠正", [item.id for item in schema.inputs])
format_input = next(item for item in schema.inputs if item.id == "输出格式")
self.assertEqual(format_input.options, ["jpeg", "png", "webp"])
self.assertEqual(format_input.default, "jpeg")
background_input = next(item for item in schema.inputs if item.id == "背景")
self.assertEqual(background_input.options, ["auto", "transparent", "opaque"])
self.assertEqual(background_input.default, "auto")
self.assertNotIn("内容审查强度", [item.id for item in schema.inputs])
quality_input = next(item for item in schema.inputs if item.id == "质量")
self.assertEqual(quality_input.options[-2:], ["超高", "最高"])
prompt_input = next(item for item in schema.inputs if item.id == "prompt")
self.assertEqual(prompt_input.display_name, "提示词")
self.assertTrue(prompt_input.socketless)
self.assertIn("总任务数 = 提示词数量 × 生图数量", prompt_input.tooltip)
batch_input = next(item for item in schema.inputs if item.id == "批量出图")
self.assertFalse(batch_input.default)
batch_mode = next(item for item in schema.inputs if item.id == "批量模式")
self.assertEqual(
batch_mode.options,
[
IMAGE_BATCH_MODE_GROUP_TO_MODELS,
IMAGE_BATCH_MODE_CARTESIAN,
IMAGE_BATCH_MODE_SINGLE_REFERENCES,
],
)
naming_input = next(item for item in schema.inputs if item.id == "命名规则")
self.assertEqual(naming_input.options, ["和主图一致", "自然数字", "自定义前缀"])
self.assertEqual(naming_input.default, "自定义前缀")
prefix_input = next(item for item in schema.inputs if item.id == "filename_prefix")
self.assertEqual(prefix_input.display_name, "文件名前缀")
self.assertEqual(prefix_input.default, "o1key")
save_format_input = next(item for item in schema.inputs if item.id == "格式")
self.assertEqual(save_format_input.options, ["原始", "png", "jpg", "webp"])
layer_input = next(item for item in schema.inputs if item.id == "图层拆分")
self.assertFalse(layer_input.default)
self.assertEqual(save_format_input.default, "原始")
location_input = next(item for item in schema.inputs if item.id == "保存位置")
self.assertEqual(location_input.default, "")
search_input = next(item for item in schema.inputs if item.id == "在线搜索")
self.assertEqual(search_input.options, ["关闭", "打开"])
self.assertEqual(search_input.default, "关闭")
def test_batch_task_expansion_covers_group_and_cartesian_pairing(self):
grouped = expand_image_generation_tasks(
"换装",
1,
batch_enabled=True,
batch_mode=IMAGE_BATCH_MODE_GROUP_TO_MODELS,
reference_count=3,
model_reference_count=10,
)
self.assertEqual(len(grouped), 10)
self.assertTrue(all(task["reference_indices"] == (0, 1, 2) for task in grouped))
self.assertEqual(
[task["model_reference_indices"] for task in grouped],
[(index,) for index in range(10)],
)
cartesian = expand_image_generation_tasks(
"换装",
1,
batch_enabled=True,
batch_mode=IMAGE_BATCH_MODE_CARTESIAN,
reference_count=10,
model_reference_count=10,
)
self.assertEqual(len(cartesian), 100)
self.assertEqual(cartesian[0]["reference_indices"], (0,))
self.assertEqual(cartesian[0]["model_reference_indices"], (0,))
self.assertEqual(cartesian[9]["reference_indices"], (0,))
self.assertEqual(cartesian[9]["model_reference_indices"], (9,))
self.assertEqual(cartesian[10]["reference_indices"], (1,))
self.assertEqual(cartesian[10]["model_reference_indices"], (0,))
single_references = expand_image_generation_tasks(
"换动作",
1,
batch_enabled=True,
batch_mode=IMAGE_BATCH_MODE_SINGLE_REFERENCES,
reference_count=3,
model_reference_count=0,
)
self.assertEqual(len(single_references), 3)
self.assertEqual(
[task["reference_indices"] for task in single_references],
[(0,), (1,), (2,)],
)
self.assertTrue(
all(task["model_reference_indices"] == () for task in single_references)
)
fifty_single_references = expand_image_generation_tasks(
"换动作",
1,
batch_enabled=True,
batch_mode=IMAGE_BATCH_MODE_SINGLE_REFERENCES,
reference_count=MAX_UNIFIED_BATCH_IMAGES,
)
self.assertEqual(len(fifty_single_references), 50)
self.assertEqual(fifty_single_references[-1]["reference_indices"], (49,))
def test_reference_limit_is_ten(self):
accepted = [
{"name": f"{index}.png", "subfolder": "", "type": "input"}
for index in range(MAX_UNIFIED_REFERENCE_IMAGES)
]
self.assertEqual(len(_parse_upload_manifest(accepted)), 10)
with self.assertRaisesRegex(ValueError, "最多支持 10 张"):
_parse_upload_manifest([
*accepted,
{"name": "overflow.png", "subfolder": "", "type": "input"},
])
def test_batch_manifest_limit_is_fifty(self):
accepted = [
{"name": f"{index}.png", "subfolder": "", "type": "input"}
for index in range(MAX_UNIFIED_BATCH_IMAGES)
]
self.assertEqual(
len(_parse_upload_manifest(accepted, max_images=MAX_UNIFIED_BATCH_IMAGES)),
50,
)
with self.assertRaisesRegex(ValueError, "最多支持 50 张"):
_parse_upload_manifest(
[
*accepted,
{"name": "overflow.png", "subfolder": "", "type": "input"},
],
max_images=MAX_UNIFIED_BATCH_IMAGES,
)
def test_gpt_resolution_and_ratio_map_to_exact_sizes(self):
self.assertEqual(resolve_gpt_image_size("1K", "1:1"), "1024x1024")
self.assertEqual(resolve_gpt_image_size("2K", "16:9"), "3648x2048")
self.assertEqual(resolve_gpt_image_size("4K", "9:16"), "2160x3840")
self.assertEqual(resolve_gpt_image_size("1K", "智能"), "1024x1024")
self.assertEqual(resolve_gpt_image_size("2K", "智能"), "2048x2048")
self.assertEqual(resolve_gpt_image_size("4K", "智能"), "2880x2880")
self.assertEqual(
resolve_gpt_image_size("3648x20482K 横版 16:9", "1:1"),
"3648x2048",
)
self.assertEqual(
GPT_IMAGE_ASPECT_RATIO_OPTIONS,
["智能", "1:1", "3:2", "2:3", "4:3", "3:4", "16:9", "9:16"],
)
def test_seedream_resolution_and_ratio_map_to_exact_sizes(self):
self.assertEqual(resolve_seedream_size("1K", "智能"), "1024x1024")
self.assertEqual(resolve_seedream_size("1K", "21:9"), "1568x672")
self.assertEqual(resolve_seedream_size("2K", "4:3"), "2368x1776")
self.assertEqual(resolve_seedream_size("2K", "9:16"), "1584x2816")
with self.assertRaisesRegex(ValueError, "Seedream 分辨率"):
resolve_seedream_size("4K", "1:1")
def test_manifest_keeps_only_safe_server_fields(self):
value = json.dumps([
{"name": "one.png", "subfolder": "refs", "type": "input", "url": "ignored"},
])
self.assertEqual(
_parse_upload_manifest(value),
[{"name": "one.png", "subfolder": "refs", "type": "input"}],
)
def test_loads_mixed_size_images_as_independent_tensors(self):
with tempfile.TemporaryDirectory() as temp_dir:
refs_dir = os.path.join(temp_dir, "refs")
os.makedirs(refs_dir)
Image.new("RGB", (8, 6), "red").save(os.path.join(refs_dir, "one.png"))
Image.new("RGB", (5, 9), "blue").save(os.path.join(temp_dir, "two.png"))
manifest = json.dumps([
{"name": "one.png", "subfolder": "refs", "type": "input"},
{"name": "two.png", "subfolder": "", "type": "input"},
])
with patch(
"comfyui_o1key.nodes.o1key_image_generator.folder_paths.get_input_directory",
return_value=temp_dir,
):
tensors = _load_reference_tensors(manifest)
self.assertEqual(tuple(tensors["参考图1"].shape), (1, 6, 8, 3))
self.assertEqual(tuple(tensors["参考图2"].shape), (1, 9, 5, 3))
def test_original_jpeg_format_and_path_survive_tensor_round_trip(self):
with tempfile.TemporaryDirectory() as temp_dir:
image_path = os.path.join(temp_dir, "reference.jpg")
Image.new("RGB", (7, 5), "orange").save(image_path, format="JPEG", quality=92)
manifest = json.dumps([
{"name": "reference.jpg", "subfolder": "", "type": "input"},
])
with patch(
"comfyui_o1key.nodes.o1key_image_generator.folder_paths.get_input_directory",
return_value=temp_dir,
):
tensors = _load_reference_tensors(manifest)
restored = tensor_to_pil(tensors["参考图1"])[0]
self.assertEqual(getattr(restored, "_o1key_original_format", None), "JPEG")
self.assertEqual(getattr(restored, "_o1key_original_path", None), image_path)
def test_rejects_path_traversal(self):
with tempfile.TemporaryDirectory() as temp_dir:
manifest = json.dumps([
{"name": "outside.png", "subfolder": "..", "type": "input"},
])
with patch(
"comfyui_o1key.nodes.o1key_image_generator.folder_paths.get_input_directory",
return_value=temp_dir,
):
with self.assertRaisesRegex(ValueError, "路径不安全"):
_load_reference_tensors(manifest)
def test_execute_delegates_to_nano_banana_without_inline_preview(self):
image = torch.zeros((1, 8, 8, 3), dtype=torch.float32)
with patch(
"comfyui_o1key.nodes.o1key_image_generator.NanoBanana.execute",
return_value=io.NodeOutput(image),
) as execute:
result = asyncio.run(O1keyImageGenerator.execute(
prompt="面板里的提示词",
模型="Nano Banana 2",
模型线路="畅速",
思考等级="高",
分辨率="2K",
宽高比="1:1",
生图数量=2,
seed=42,
参考图清单="[]",
缩放图片="智能缩放",
在线搜索="打开",
))
self.assertIs(result[0], image)
self.assertIsNone(result.ui)
self.assertEqual(execute.call_args.kwargs["prompt"], "面板里的提示词")
self.assertEqual(execute.call_args.kwargs["生图数量"], 2)
self.assertEqual(execute.call_args.kwargs["参考图组"], {})
self.assertEqual(execute.call_args.kwargs["缩放图片"], "智能缩放")
self.assertIs(execute.call_args.kwargs["_o1key_google_search"], True)
self.assertIs(execute.call_args.kwargs["_o1key_unlimited_downloads"], True)
def test_online_search_is_ignored_for_other_models(self):
image = torch.zeros((1, 8, 8, 3), dtype=torch.float32)
with patch(
"comfyui_o1key.nodes.o1key_image_generator.NanoBanana.execute",
return_value=io.NodeOutput(image),
) as execute:
asyncio.run(O1keyImageGenerator.execute(
prompt="一只宇航猫",
模型="Nano Banana Pro",
分辨率="2K",
在线搜索="打开",
))
self.assertIs(execute.call_args.kwargs["_o1key_google_search"], False)
def test_error_wrappers_apply_to_every_model_family(self):
raw = (
"status_code=451, The provided prompt is considered unsafe and it "
"cannot be used to generate content"
)
expected = "提供的提示被认为是不安全的,不能用于生成内容。"
with patch(
"comfyui_o1key.nodes.o1key_image_generator.NanoBanana.execute",
side_effect=RuntimeError(raw),
):
with self.assertRaisesRegex(RuntimeError, expected):
asyncio.run(O1keyImageGenerator.execute(
prompt="测试",
模型="Nano Banana 2",
))
with patch(
"comfyui_o1key.nodes.o1key_image_generator._generate_gpt_images",
new=AsyncMock(side_effect=RuntimeError(raw)),
):
with self.assertRaisesRegex(RuntimeError, expected):
asyncio.run(O1keyImageGenerator.execute(
prompt="测试",
模型="gpt-image-2",
))
def test_o1key_error_formatter_covers_all_node_mappings(self):
cases = {
"status_code=400, content rejected: the image was flagged as unsafe by the content safety system": "内容被拒绝:该图像被内容安全系统标记为不安全。",
"status_code=400, Your request was rejected by the safety system": "您的请求已被安全系统拒绝",
"status_code=403, insufficient balance": "上游额度不足!",
"status_code=502, Image generation returned empty response": "图片生成过程中被内容审查机制拒绝!",
"status_code=451, The provided prompt is considered unsafe and it cannot be used to generate content": "提供的提示被认为是不安全的,不能用于生成内容。",
}
for raw, expected in cases.items():
with self.subTest(raw=raw):
self.assertEqual(format_o1key_image_error(raw), expected)
def test_rejects_unsupported_image_count(self):
with self.assertRaisesRegex(ValueError, "仅支持:1、2、4、9"):
asyncio.run(O1keyImageGenerator.execute(prompt="测试", 生图数量=3))
def test_direct_execution_returns_provider_image_without_color_postprocess(self):
reference = torch.zeros((1, 8, 8, 3), dtype=torch.float32)
reference._o1key_source_metadata = [{"filename": "reference.png"}]
generated = torch.ones((1, 8, 8, 3), dtype=torch.float32)
with (
patch(
"comfyui_o1key.nodes.o1key_image_generator._load_reference_tensors",
return_value={"参考图1": reference},
),
patch(
"comfyui_o1key.nodes.o1key_image_generator.NanoBanana.execute",
return_value=io.NodeOutput(generated),
),
):
result = asyncio.run(O1keyImageGenerator.execute(
prompt="测试",
参考图清单='[{"name":"reference.png","type":"input"}]',
))
self.assertIs(result[0], generated)
self.assertEqual(generated._o1key_main_filename, "reference.png")
self.assertEqual(generated._o1key_save_settings, {
"filename_prefix": "o1key",
"format": "原始",
"save_location": "",
"naming_rule": "自定义前缀",
})
def test_gpt_image_uses_concurrent_single_image_requests(self):
calls = []
active = 0
maximum_active = 0
class FakeClient:
def __init__(self):
self.base_url = ""
self.response_log_enabled = True
self.poll_log_enabled = True
async def generate_image_async(self, **kwargs):
nonlocal active, maximum_active
calls.append(kwargs)
active += 1
maximum_active = max(maximum_active, active)
await __import__("asyncio").sleep(0.01)
active -= 1
return [Image.new("RGB", (6, 4), "green")]
@staticmethod
def _pil_list_to_tensor(images):
return torch.zeros((len(images), 4, 6, 3), dtype=torch.float32)
with (
patch(
"comfyui_o1key.nodes.o1key_image_generator.GptImageClient",
FakeClient,
),
patch(
"comfyui_o1key.nodes.o1key_image_generator.get_base_url_by_route",
return_value="https://example.invalid",
),
):
result = asyncio.run(O1keyImageGenerator.execute(
prompt="一只宇航猫",
模型="gpt-image-2.5-sunburst",
模型线路="畅速",
分辨率="2K",
宽高比="16:9",
生图数量=4,
质量="最高",
输出格式="webp",
背景="transparent",
缩放图片="智能缩放",
))
self.assertEqual(tuple(result[0].shape), (4, 4, 6, 3))
self.assertEqual(len(calls), 4)
self.assertGreater(maximum_active, 1)
self.assertTrue(all(call["n"] == 1 for call in calls))
self.assertTrue(all(call["model"] == "gpt-image-2.5-sunburst-sp" for call in calls))
self.assertTrue(all(call["special_price_parallel"] is False for call in calls))
self.assertTrue(all(call["size"] == "3648x2048" for call in calls))
self.assertTrue(all(call["quality"] == "max" for call in calls))
self.assertTrue(all(call["output_format"] == "webp" for call in calls))
self.assertTrue(all(call["background"] == "transparent" for call in calls))
self.assertTrue(all("moderation" not in call for call in calls))
self.assertTrue(all(call["resize_mode"] == "智能缩放" for call in calls))
def test_seedream_direct_execution_sends_exact_size_and_output_format(self):
calls = []
session = object()
class FakeSessionContext:
async def __aenter__(self):
return session
async def __aexit__(self, _exc_type, _exc, _tb):
return False
class FakeClient:
def __init__(self, **kwargs):
self.settings = kwargs
async def generate_async(self, **kwargs):
calls.append(kwargs)
return [Image.new("RGB", (6, 4), "green")], {}
with (
patch(
"comfyui_o1key.nodes.o1key_image_generator.SeedreamImageClient",
FakeClient,
),
patch(
"comfyui_o1key.nodes.o1key_image_generator.get_api_key_or_raise",
return_value="secret",
),
patch(
"comfyui_o1key.nodes.o1key_image_generator.get_base_url_by_route",
return_value="https://cf-api.o1key.com",
),
patch(
"comfyui_o1key.nodes.o1key_image_generator.create_http_client",
return_value=FakeSessionContext(),
),
):
result = asyncio.run(O1keyImageGenerator.execute(
prompt="一只宇航猫",
模型="Seedream 5.0 Pro",
模型线路="畅速",
分辨率="2K",
宽高比="21:9",
输出格式="png",
格式="webp",
))
self.assertEqual(tuple(result[0].shape), (1, 4, 6, 3))
self.assertEqual(result[0]._o1key_save_settings["format"], "原始")
self.assertEqual(len(calls), 1)
self.assertIs(calls[0]["session"], session)
self.assertEqual(calls[0]["model"], "dola-seedream-5-0-pro-260628-ep")
self.assertEqual(calls[0]["size"], "3136x1344")
self.assertEqual(calls[0]["output_format"], "png")
def test_seedream_direct_execution_validates_all_references_before_client_calls(self):
calls = []
class FakeClient:
def __init__(self, **_kwargs):
pass
async def generate_async(self, **kwargs):
calls.append(kwargs)
return [Image.new("RGB", (6, 4), "green")], {}
invalid_reference = torch.zeros((1, 14, 15, 3), dtype=torch.float32)
with (
patch(
"comfyui_o1key.nodes.o1key_image_generator.SeedreamImageClient",
FakeClient,
),
patch(
"comfyui_o1key.nodes.o1key_image_generator._load_reference_tensors",
return_value={"参考图1": invalid_reference},
),
):
with self.assertRaisesRegex(ValueError, "宽和高都必须大于 14px"):
asyncio.run(O1keyImageGenerator.execute(
prompt="test",
模型="Seedream 5.0 Pro",
参考图清单='[{"name":"invalid.png","type":"input"}]',
))
self.assertEqual(calls, [])
def test_seedream_layer_decomposition_returns_base_layers_masks_and_metadata(self):
calls = []
base = Image.new("RGB", (6, 4), "white")
layer = Image.new("RGBA", (3, 2), (10, 20, 30, 64))
setattr(base, "_o1key_seedream_layer", {"z_index": 0, "name": "底图"})
setattr(layer, "_o1key_seedream_layer", {"z_index": 1, "name": "主体"})
class FakeSessionContext:
async def __aenter__(self):
return object()
async def __aexit__(self, _exc_type, _exc, _tb):
return False
class FakeClient:
def __init__(self, **_kwargs):
pass
async def generate_async(self, **kwargs):
calls.append(kwargs)
return [base, layer], {}
reference = torch.zeros((1, 512, 512, 3), dtype=torch.float32)
with (
patch(
"comfyui_o1key.nodes.o1key_image_generator.SeedreamImageClient",
FakeClient,
),
patch(
"comfyui_o1key.nodes.o1key_image_generator._load_reference_tensors",
return_value={"参考图1": reference},
),
patch(
"comfyui_o1key.nodes.o1key_image_generator.get_api_key_or_raise",
return_value="secret",
),
patch(
"comfyui_o1key.nodes.o1key_image_generator.get_base_url_by_route",
return_value="https://cf-api.o1key.com",
),
patch(
"comfyui_o1key.nodes.o1key_image_generator.create_http_client",
return_value=FakeSessionContext(),
),
):
result = asyncio.run(O1keyImageGenerator.execute(
prompt="",
模型="Seedream 5.0 Pro",
分辨率="1.5K",
生图数量="1",
输出格式="png",
图层拆分=True,
))
self.assertEqual(tuple(result[0].shape), (1, 4, 6, 3))
self.assertEqual(tuple(result[1][0].shape), (1, 2, 3, 3))
self.assertEqual(tuple(result[2][0].shape), (1, 2, 3))
self.assertAlmostEqual(float(result[2][0][0, 0, 0]), 64 / 255)
self.assertEqual(json.loads(result[3]), [{"z_index": 1, "name": "主体"}])
self.assertEqual(calls[0]["size"], "1.5K")
self.assertIs(calls[0]["layer_decomposition"], True)
def test_gpt_batch_prompts_multiply_by_image_count(self):
calls = []
class FakeClient:
def __init__(self):
self.base_url = ""
self.response_log_enabled = True
self.poll_log_enabled = True
async def generate_image_async(self, **kwargs):
calls.append(kwargs)
return [Image.new("RGB", (6, 4), "green")]
@staticmethod
def _pil_list_to_tensor(images):
return torch.zeros((len(images), 4, 6, 3), dtype=torch.float32)
with (
patch(
"comfyui_o1key.nodes.o1key_image_generator.GptImageClient",
FakeClient,
),
patch(
"comfyui_o1key.nodes.o1key_image_generator.get_base_url_by_route",
return_value="https://example.invalid",
),
):
result = asyncio.run(O1keyImageGenerator.execute(
prompt="第一条\n---\n第二条",
模型="gpt-image-2",
分辨率="1K",
宽高比="1:1",
生图数量=2,
))
self.assertEqual(tuple(result[0].shape), (4, 4, 6, 3))
self.assertEqual(
[call["prompt"] for call in calls],
["第一条", "第一条", "第二条", "第二条"],
)
self.assertTrue(all(call["output_format"] == "png" for call in calls))
self.assertEqual(result[0]._o1key_save_settings["format"], "原始")
def test_direct_gpt_cartesian_batch_sends_one_outfit_and_one_model_per_task(self):
outfit_one = torch.zeros((1, 4, 4, 3), dtype=torch.float32)
outfit_two = torch.ones((1, 4, 4, 3), dtype=torch.float32)
model_one = torch.full((1, 4, 4, 3), 0.25, dtype=torch.float32)
model_two = torch.full((1, 4, 4, 3), 0.75, dtype=torch.float32)
calls = []
class FakeClient:
def __init__(self):
self.base_url = ""
self.response_log_enabled = True
self.poll_log_enabled = True
async def generate_image_async(self, **kwargs):
calls.append(kwargs["image_tensor"])
return [Image.new("RGB", (4, 4), "green")]
@staticmethod
def _pil_list_to_tensor(images):
return torch.zeros((len(images), 4, 4, 3), dtype=torch.float32)
def fake_load(_manifest, *, label="参考图", max_images=MAX_UNIFIED_REFERENCE_IMAGES):
if label == "目标图":
return {"模特图1": model_one, "模特图2": model_two}
return {"参考图1": outfit_one, "参考图2": outfit_two}
with (
patch(
"comfyui_o1key.nodes.o1key_image_generator._load_reference_tensors",
side_effect=fake_load,
),
patch(
"comfyui_o1key.nodes.o1key_image_generator.GptImageClient",
FakeClient,
),
patch(
"comfyui_o1key.nodes.o1key_image_generator.get_base_url_by_route",
return_value="https://example.invalid",
),
):
result = asyncio.run(O1keyImageGenerator.execute(
prompt="换装",
模型="gpt-image-2",
分辨率="1K",
批量出图=True,
批量模式=IMAGE_BATCH_MODE_CARTESIAN,
参考图清单="outfits",
模特图清单="models",
))
self.assertEqual(tuple(result[0].shape), (4, 4, 4, 3))
self.assertEqual(len(calls), 4)
expected = [
(outfit_one, model_one),
(outfit_one, model_two),
(outfit_two, model_one),
(outfit_two, model_two),
]
for actual, pair in zip(calls, expected):
self.assertIs(actual[0], pair[0])
self.assertIs(actual[1], pair[1])
def test_gpt_single_reference_batch_sends_each_source_independently(self):
model_one = torch.zeros((1, 4, 4, 3), dtype=torch.float32)
model_two = torch.ones((1, 4, 4, 3), dtype=torch.float32)
calls = []
class FakeClient:
def __init__(self):
self.base_url = ""
self.response_log_enabled = True
self.poll_log_enabled = True
async def generate_image_async(self, **kwargs):
calls.append(kwargs["image_tensor"])
return [Image.new("RGB", (4, 4), "green")]
@staticmethod
def _pil_list_to_tensor(images):
return torch.zeros((len(images), 4, 4, 3), dtype=torch.float32)
def fake_load(_manifest, *, label="参考图", max_images=MAX_UNIFIED_REFERENCE_IMAGES):
self.assertEqual(label, "参考图")
return {"模特图1": model_one, "模特图2": model_two}
with (
patch(
"comfyui_o1key.nodes.o1key_image_generator._load_reference_tensors",
side_effect=fake_load,
),
patch(
"comfyui_o1key.nodes.o1key_image_generator.GptImageClient",
FakeClient,
),
patch(
"comfyui_o1key.nodes.o1key_image_generator.get_base_url_by_route",
return_value="https://example.invalid",
),
):
result = asyncio.run(O1keyImageGenerator.execute(
prompt="换个动作",
模型="gpt-image-2",
分辨率="1K",
批量出图=True,
批量模式=IMAGE_BATCH_MODE_SINGLE_REFERENCES,
参考图清单="models",
模特图清单="ignored-targets",
))
self.assertEqual(tuple(result[0].shape), (2, 4, 4, 3))
self.assertEqual(len(calls), 2)
self.assertEqual([len(call) for call in calls], [1, 1])
self.assertIs(calls[0][0], model_one)
self.assertIs(calls[1][0], model_two)
def test_gpt_rejects_transparent_jpeg_before_paid_request(self):
with patch(
"comfyui_o1key.nodes.o1key_image_generator.GptImageClient",
) as client:
with self.assertRaisesRegex(ValueError, "透明背景仅支持 PNG 或 WebP"):
asyncio.run(O1keyImageGenerator.execute(
prompt="透明图标",
模型="gpt-image-2",
输出格式="jpeg",
背景="transparent",
))
client.assert_not_called()
def test_gpt_rejects_25_quality_on_gpt_2_before_paid_request(self):
with patch(
"comfyui_o1key.nodes.o1key_image_generator.GptImageClient",
) as client:
with self.assertRaisesRegex(ValueError, "仅支持 GPT Image 2.5"):
asyncio.run(O1keyImageGenerator.execute(
prompt="测试",
模型="gpt-image-2",
质量="超高",
))
client.assert_not_called()
def test_gpt_unified_resize_mode_uses_18_mib_policy_and_keeps_mask_aligned(self):
client = object.__new__(GptImageClient)
client._UNIFIED_BODY_LIMIT_BYTES = 24 * 1024
client._SMART_RESIZE_MIN_LONG_EDGE = 16
reference = torch.rand((1, 128, 96, 3), dtype=torch.float32)
mask = torch.zeros((1, 128, 96), dtype=torch.float32)
kwargs = {
"prompt": "测试",
"model": "gpt-image-2",
"quality": "auto",
"size": "auto",
"n": 1,
"image_list": [reference],
"mask_tensor": mask,
"output_format": "webp",
"background": "transparent",
}
with self.assertRaisesRegex(ValueError, "当前设置为“不缩放”"):
client._build_async_generate_body(**kwargs, resize_mode="不缩放")
body = client._build_async_generate_body(**kwargs, resize_mode="智能缩放")
self.assertEqual(body["output_format"], "webp")
self.assertEqual(body["background"], "transparent")
self.assertLessEqual(
client._json_body_size(body),
client._UNIFIED_BODY_LIMIT_BYTES,
)
def _data_url_size(value):
encoded = value.split(",", 1)[1]
with Image.open(BytesIO(base64.b64decode(encoded))) as image:
return image.size
image_size = _data_url_size(body["images"][0])
mask_size = _data_url_size(body["mask"]["image_url"])
self.assertEqual(image_size, mask_size)
self.assertLess(image_size[0], 96)
self.assertLessEqual(abs(image_size[0] * 128 - image_size[1] * 96), 128)
def test_gpt_client_omits_auto_moderation_and_sends_low(self):
client = object.__new__(GptImageClient)
kwargs = {
"prompt": "测试",
"model": "gpt-image-2",
"quality": "auto",
"size": "auto",
"n": 1,
}
automatic = client._build_async_generate_body(**kwargs)
self.assertNotIn("moderation", automatic)
low = client._build_async_generate_body(**kwargs, moderation="low")
self.assertEqual(low["moderation"], "low")
with self.assertRaisesRegex(ValueError, "内容审查强度"):
client._build_async_generate_body(**kwargs, moderation="high")
def test_gpt_client_omits_size_for_smart_resolution(self):
client = object.__new__(GptImageClient)
body = client._build_async_generate_body(
prompt="测试",
model="gpt-image-2",
quality="auto",
size=None,
n=1,
)
self.assertNotIn("size", body)
def test_gpt_client_translates_unsafe_image_rejection(self):
raw = (
"status_code=400, content rejected: the image was flagged as unsafe "
"by the content safety system"
)
expected = "内容被拒绝:该图像被内容安全系统标记为不安全。"
self.assertEqual(GptImageClient._extract_error_message(raw, 400), expected)
self.assertEqual(
GptImageClient._extract_error_message({"error": raw}, 500),
expected,
)
self.assertEqual(
GptImageClient._extract_error_message({"error": {"message": raw}}, 500),
expected,
)
def test_gpt_client_translates_general_safety_rejection(self):
raw = "status_code=400, Your request was rejected by the safety system"
expected = "您的请求已被安全系统拒绝"
self.assertEqual(GptImageClient._extract_error_message(raw, 400), expected)
self.assertEqual(
GptImageClient._extract_error_message({"error": raw}, 500),
expected,
)
def test_gpt_client_translates_insufficient_balance(self):
raw = "status_code=403, insufficient balance"
expected = "上游额度不足!"
self.assertEqual(GptImageClient._extract_error_message(raw, 403), expected)
self.assertEqual(
GptImageClient._extract_error_message({"error": raw}, 500),
expected,
)
def test_gpt_client_translates_empty_image_response(self):
raw = "status_code=502, Image generation returned empty response"
expected = "图片生成过程中被内容审查机制拒绝!"
self.assertEqual(GptImageClient._extract_error_message(raw, 502), expected)
self.assertEqual(
GptImageClient._extract_error_message({"error": raw}, 500),
expected,
)
def test_gpt_client_translates_unsafe_prompt(self):
raw = (
"status_code=451, The provided prompt is considered unsafe and it "
"cannot be used to generate content"
)
expected = "提供的提示被认为是不安全的,不能用于生成内容。"
self.assertEqual(GptImageClient._extract_error_message(raw, 451), expected)
self.assertEqual(
GptImageClient._extract_error_message({"error": raw}, 500),
expected,
)
def test_gpt_client_enforces_n_one_for_every_route(self):
calls = []
active = 0
maximum_active = 0
client = object.__new__(GptImageClient)
async def fake_single_request(**kwargs):
nonlocal active, maximum_active
calls.append(kwargs)
active += 1
maximum_active = max(maximum_active, active)
await __import__("asyncio").sleep(0.01)
active -= 1
return [Image.new("RGB", (4, 4), "blue")]
client._generate_image_task_async = fake_single_request
images = __import__("asyncio").run(client.generate_image_async(
prompt="测试",
model="gpt-image-2",
quality="auto",
size="auto",
n=4,
seed=0,
output_format="webp",
background="transparent",
moderation="low",
special_price_parallel=False,
log_request_start=False,
))
try:
self.assertEqual(len(images), 4)
self.assertEqual(len(calls), 4)
self.assertGreater(maximum_active, 1)
self.assertTrue(all(call["n"] == 1 for call in calls))
self.assertTrue(all(call["output_format"] == "webp" for call in calls))
self.assertTrue(all(call["background"] == "transparent" for call in calls))
self.assertTrue(all(call["moderation"] == "low" for call in calls))
finally:
for image in images:
image.close()
with self.assertRaisesRegex(ValueError, "n 必须为 1"):
client._build_async_generate_body(
prompt="测试",
model="gpt-image-2",
quality="auto",
size="auto",
n=2,
)
with self.assertRaisesRegex(ValueError, "透明背景仅支持 PNG 或 WebP"):
client._build_async_generate_body(
prompt="测试",
model="gpt-image-2",
quality="auto",
size="auto",
n=1,
output_format="jpeg",
background="transparent",
)
def test_matching_save_node_saves_to_output_ui(self):
schema = O1keyImageSave.define_schema()
schema.validate()
self.assertEqual(schema.node_id, "O1keyImageSave")
self.assertEqual(schema.display_name, "o1key 保存图像")
self.assertTrue(schema.is_output_node)
self.assertEqual(
[item.id for item in schema.inputs],
["images"],
)
self.assertEqual([item.id for item in schema.outputs], ["IMAGE"])
self.assertEqual(schema.outputs[0].display_name, "图像")
buffer = BytesIO()
Image.new("RGB", (8, 8), "orange").save(buffer, format="JPEG", quality=92)
raw = buffer.getvalue()
with Image.open(BytesIO(raw)) as opened:
opened.load()
image = opened.copy()
image.format = "JPEG"
image._o1key_original_format = "JPEG"
image._o1key_original_bytes = raw
tensor = pil_to_tensor([image])
tensor._o1key_save_settings = {
"filename_prefix": "custom-prefix",
"format": "原始",
"save_location": "",
"naming_rule": "自定义前缀",
}
image.close()
with tempfile.TemporaryDirectory() as output_dir, patch(
"comfyui_o1key.nodes.o1key_image_generator.folder_paths.get_output_directory",
return_value=output_dir,
), patch(
"comfyui_o1key.nodes.o1key_image_generator.folder_paths.get_save_image_path",
return_value=(output_dir, "custom-prefix", 1, "", "custom-prefix"),
):
result = O1keyImageSave.execute(tensor)
saved = result.ui.as_dict()["images"]
self.assertIs(result[0], tensor)
self.assertEqual(saved[0]["filename"], "custom-prefix_00001_.jpg")
self.assertEqual(saved[0]["type"], "output")
class GptInlineResultRetryTests(unittest.IsolatedAsyncioTestCase):
@staticmethod
def _task_query_session(body, headers):
class Content:
async def iter_chunked(self, _chunk_size):
yield body
class Response:
status = 200
http_version = "HTTP/1.1"
raw_headers = ()
content = Content()
def __init__(self):
self.headers = headers
async def __aenter__(self):
return self
async def __aexit__(self, _exc_type, _exc, _tb):
return None
class Session:
def __init__(self):
self.calls = 0
def get(self, _url, **_kwargs):
self.calls += 1
return Response()
return Session()
async def test_successful_task_query_keeps_transport_trace_silent(self):
body = json.dumps(
{"task_id": "task-trace", "status": "IN_PROGRESS"},
separators=(",", ":"),
).encode("utf-8")
session = self._task_query_session(
body,
{"Content-Length": str(len(body))},
)
client = object.__new__(GptImageClient)
client.response_log_enabled = False
with patch("builtins.print") as print_mock:
payload = await client._get_json_with_response_retry(
session,
"https://example.invalid/task-trace",
{},
"任务 task-trace 查询",
expected_task_id="task-trace",
)
self.assertEqual(payload["status"], "IN_PROGRESS")
rendered = " ".join(str(value) for call in print_mock.call_args_list for value in call.args)
self.assertNotIn("任务查询传输追踪", rendered)
async def test_task_query_retries_when_content_length_is_short(self):
body = b'{"task_id":"task-short","status":"SUCCESS"}'
session = self._task_query_session(
body,
{"Content-Length": str(len(body) + 20)},
)
client = object.__new__(GptImageClient)
client.response_log_enabled = False
with patch(
"comfyui_o1key.clients.gpt_image_client.asyncio.sleep",
new=AsyncMock(),
):
with self.assertRaisesRegex(
RuntimeError,
rf"Content-Length={len(body) + 20}.*received={len(body)}B.*length_check=mismatch",
):
await client._get_json_with_response_retry(
session,
"https://example.invalid/task-short",
{},
"任务 task-short 查询",
expected_task_id="task-short",
)
self.assertEqual(session.calls, 4)
def test_unparseable_response_log_never_prints_partial_base64(self):
client = object.__new__(GptImageClient)
client.response_log_enabled = True
partial_secret = "B" * 4097
with patch("builtins.print") as print_mock:
client._log_original_response_body(
"task response",
'{"data":{"images":[{"b64_json":"' + partial_secret,
)
rendered = " ".join(str(value) for call in print_mock.call_args_list for value in call.args)
self.assertNotIn(partial_secret, rendered)
self.assertIn("content omitted", rendered)
async def test_refetches_same_task_when_inline_base64_is_incomplete(self):
buffer = BytesIO()
Image.new("RGB", (3, 2), "purple").save(buffer, format="PNG")
invalid_task = {
"status": "SUCCESS",
"data": {"images": [{"b64_json": "truncated-base64"}]},
}
valid_task = {
"status": "SUCCESS",
"data": {"images": [{
"b64_json": base64.b64encode(buffer.getvalue()).decode("ascii"),
}]},
}
client = object.__new__(GptImageClient)
session = object()
with (
patch.object(
client,
"_poll_generate_image_task",
new=AsyncMock(return_value=valid_task),
) as poll_task,
patch(
"comfyui_o1key.clients.gpt_image_client.asyncio.sleep",
new=AsyncMock(),
) as retry_sleep,
):
final_task, images = await client._parse_completed_task_images_with_retry(
invalid_task,
session,
"task-2",
log_downloads=False,
)
try:
self.assertIs(final_task, valid_task)
self.assertEqual(len(images), 1)
self.assertEqual(images[0].size, (3, 2))
self.assertEqual(images[0]._o1key_original_format, "PNG")
self.assertEqual(images[0]._o1key_original_bytes, buffer.getvalue())
retry_sleep.assert_awaited_once()
poll_task.assert_awaited_once_with(
session,
"task-2",
progress_callback=None,
initial_delay=False,
)
finally:
for image in images:
image.close()
if __name__ == "__main__":
unittest.main(verbosity=2)