Files
comfyui_o1key/tests/test_nano_banana_dynamic_inputs.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

335 lines
14 KiB
Python

"""Offline tests for Nano Banana's dynamic reference-image inputs."""
import asyncio
import os
import sys
import unittest
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.environ.get("COMFYUI_ROOT", r"D:\o1key\ComfyUI")
sys.path.insert(0, COMFY_ROOT)
sys.path.insert(0, CUSTOM_NODES_DIR)
from comfyui_o1key.nodes.nano_banana import ( # noqa: E402
MAX_REFERENCE_IMAGES,
NanoBanana,
_collect_autogrow_inputs,
)
from comfyui_o1key.utils.nano_banana_models import ( # noqa: E402
NANO_BANANA_MODEL_OPTIONS,
NANO_BANANA_MODEL_MATRIX,
NANO_BANANA_ROUTE_OPTIONS,
resolve_nano_banana_model,
)
from comfyui_o1key.nodes.batch_nano_banana import ( # noqa: E402
BatchNanoBananaPro,
_MAX_FIXED_REFERENCE_IMAGES,
_MAX_FOLDER_INPUTS,
_generate_single_async,
_normalize_image_quality,
_path_count_from_label,
)
from comfyui_o1key.utils.file_utils import ImageInfo # noqa: E402
class NanoBananaDynamicInputTests(unittest.TestCase):
def test_schema_uses_empty_prompt_and_autogrow_reference_images(self):
schema = NanoBanana.define_schema()
schema.validate()
for item in schema.inputs:
if item.io_type == "COMBO":
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)
inputs = {item.id: item for item in schema.inputs}
self.assertEqual(inputs["prompt"].default, "")
self.assertIn("参考图组", inputs)
self.assertNotIn("网络", inputs)
self.assertFalse(inputs["模型线路"].advanced)
self.assertEqual(inputs["模型"].options, NANO_BANANA_MODEL_OPTIONS)
self.assertEqual(inputs["模型线路"].options, NANO_BANANA_ROUTE_OPTIONS)
self.assertEqual(inputs["分辨率"].options, ["1K", "2K", "4K"])
self.assertEqual(
[item.id for item in schema.inputs[:3]],
["prompt", "模型", "模型线路"],
)
self.assertFalse(inputs["seed"].advanced)
self.assertEqual(inputs["缩放图片"].options, ["不缩放", "智能缩放"])
self.assertEqual(inputs["缩放图片"].default, "不缩放")
self.assertNotIn("色彩纠正", inputs)
self.assertFalse(inputs["缩放图片"].advanced)
template = inputs["参考图组"].as_dict()["template"]
self.assertEqual(template["min"], 0)
self.assertEqual(
template["names"],
[f"参考图{i}" for i in range(1, MAX_REFERENCE_IMAGES + 1)],
)
self.assertTrue(schema.accept_all_inputs)
def test_single_node_rejects_removed_512_resolution(self):
with self.assertRaisesRegex(ValueError, "分辨率 '512' 无效"):
NanoBanana._validate_model_config("Nano Banana 2", "1:1", "512")
def test_collects_only_connected_autogrow_slots_in_order(self):
first = object()
third = object()
self.assertEqual(
_collect_autogrow_inputs({
"参考图1": first,
"参考图2": None,
"参考图3": third,
}),
[first, third],
)
def test_tolerates_legacy_single_value(self):
image = object()
self.assertEqual(_collect_autogrow_inputs(image), [image])
self.assertEqual(_collect_autogrow_inputs(None), [])
def test_model_route_matrix_matches_api_model_names(self):
expected = {
("Nano Banana Pro", "畅速"): "gemini-3-pro-image-c-sp",
("Nano Banana 2", "畅速"): "gemini-3.1-flash-image-c-sp",
("Nano Banana 2 Lite", "畅速"): "gemini-3.1-flash-lite-image-c-sp",
("Nano Banana", "畅速"): "nano-banana",
("Nano Banana Pro", "直连"): "gemini-3-pro-image-c-sd",
("Nano Banana 2", "直连"): "gemini-3.1-flash-image-c-sd",
("Nano Banana 2 Lite", "直连"): "gemini-3.1-flash-lite-image-c-sd",
("Nano Banana", "直连"): "nano-banana",
("Nano Banana Pro", "专线"): "gemini-3-pro-image",
("Nano Banana 2", "专线"): "gemini-3.1-flash-image",
("Nano Banana 2 Lite", "专线"): "gemini-3.1-flash-lite-image",
("Nano Banana", "专线"): "gemini-2.5-flash-image",
}
self.assertEqual(NANO_BANANA_MODEL_MATRIX, expected)
for key, model_id in expected.items():
self.assertEqual(resolve_nano_banana_model(*key), model_id)
def test_legacy_billing_values_resolve_to_equivalent_routes(self):
self.assertEqual(
resolve_nano_banana_model("Nano Banana Pro", "特价"),
"gemini-3-pro-image-c-sp",
)
self.assertEqual(
resolve_nano_banana_model("Nano Banana Pro", "官方"),
"gemini-3-pro-image",
)
class BatchNanoBananaDynamicInputTests(unittest.TestCase):
def test_schema_uses_dynamic_folders_autogrow_and_visible_options(self):
schema = BatchNanoBananaPro.define_schema()
schema.validate()
for item in schema.inputs:
if item.io_type == "COMBO":
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)
inputs = {item.id: item for item in schema.inputs}
self.assertEqual(schema.display_name, "Nano Banana 批量跑图")
self.assertEqual(
[item.id for item in schema.inputs[:7]],
["prompt", "模型", "模型线路", "思考等级", "分辨率", "宽高比", "图片路径数量"],
)
self.assertEqual(
[item.id for item in schema.inputs[-7:]],
["参考图组", "缩放图片", "图片输出格式", "图片质量", "图片保存命名规则", "图片保存路径", "seed"],
)
self.assertNotIn("图片路径1", inputs)
self.assertNotIn("参考图1", inputs)
self.assertTrue(schema.accept_all_inputs)
folder_options = inputs["图片路径数量"].as_dict()["options"]
self.assertEqual(
[option["key"] for option in folder_options],
[f"{count}个路径" for count in range(1, _MAX_FOLDER_INPUTS + 1)],
)
one_folder = folder_options[0]["inputs"]["required"]
five_folders = folder_options[-1]["inputs"]["required"]
self.assertEqual(list(one_folder), ["参考图1(主图)"])
self.assertEqual(
list(five_folders),
[
"参考图1(主图)", "参考图2", "参考图3", "参考图4", "参考图5",
"图片配对模式",
],
)
self.assertNotIn("图片随机抽取", one_folder)
self.assertNotIn("图片随机抽取", five_folders)
self.assertEqual(
five_folders["图片配对模式"][1]["options"],
["相同文件名", "同序号", "全匹配", "不配对"],
)
reference_template = inputs["参考图组"].as_dict()["template"]
self.assertEqual(reference_template["min"], 0)
self.assertEqual(
reference_template["names"],
[f"参考图{index}" for index in range(1, _MAX_FIXED_REFERENCE_IMAGES + 1)],
)
self.assertFalse(inputs["模型线路"].advanced)
self.assertEqual(inputs["分辨率"].options, ["1K", "2K", "4K"])
self.assertEqual(inputs["图片质量"].io_type, "INT")
self.assertEqual(inputs["图片质量"].default, 95)
self.assertEqual(inputs["图片质量"].min, 1)
self.assertEqual(inputs["图片质量"].max, 100)
for input_name in (
"seed", "图片输出格式", "图片质量", "图片保存命名规则", "图片保存路径",
"缩放图片",
):
self.assertFalse(inputs[input_name].advanced)
self.assertEqual(inputs["缩放图片"].options, ["不缩放", "智能缩放"])
self.assertNotIn("色彩纠正", inputs)
def test_batch_quality_accepts_serialized_string_and_returns_integer(self):
self.assertEqual(_normalize_image_quality("95"), 95)
self.assertIsInstance(_normalize_image_quality("95"), int)
self.assertEqual(_normalize_image_quality("87"), 87)
def test_batch_node_rejects_removed_512_resolution(self):
with self.assertRaisesRegex(ValueError, "分辨率 '512' 无效"):
BatchNanoBananaPro._validate_model_config("Nano Banana 2", "1:1", "512")
def test_execute_unpacks_dynamic_folders_and_autogrow_references(self):
first_reference = object()
third_reference = object()
output = object()
folder_group = {
"图片路径数量": "3个路径",
"参考图1(主图)": r"D:\images\one",
"参考图2": r"D:\images\two",
"参考图3": r"D:\images\three",
"图片配对模式": "相同文件名",
}
with patch.object(BatchNanoBananaPro, "process_batch", return_value=(output,)) as process:
result = BatchNanoBananaPro.execute(
prompt="test",
模型="Nano Banana 2",
图片路径数量=folder_group,
参考图组={
"参考图1": first_reference,
"参考图2": None,
"参考图3": third_reference,
},
缩放图片="智能缩放",
)
self.assertIs(result.result[0], output)
arguments = process.call_args.kwargs
self.assertEqual(
[arguments[f"文件夹{index}"] for index in range(1, 6)],
[r"D:\images\one", r"D:\images\two", r"D:\images\three", "", ""],
)
self.assertEqual(arguments["图片配对模式"], "相同文件名")
self.assertNotIn("随机抽取路径", arguments)
self.assertIs(arguments["参考图1"], first_reference)
self.assertIs(arguments["参考图2"], third_reference)
self.assertEqual(arguments["缩放图片"], "智能缩放")
self.assertNotIn("色彩纠正", arguments)
def test_batch_request_forwards_resize_mode_to_shared_nano_transport(self):
async def scenario():
with patch(
"comfyui_o1key.nodes.batch_nano_banana.generate_nano_banana_async",
new=AsyncMock(return_value=([], {})),
) as generate:
await _generate_single_async(
session=object(),
base_url="https://example.invalid",
api_key="test",
prompt="test",
model="nano-test",
resolution="2K",
aspect_ratio="1:1",
resize_mode="智能缩放",
)
self.assertEqual(generate.await_args.kwargs["resize_mode"], "智能缩放")
asyncio.run(scenario())
def test_batch_task_saves_generated_image_without_colour_postprocessing(self):
source = Image.new("RGB", (4, 4), color=(12, 34, 56))
generated = Image.new("RGB", (4, 4), color=(90, 80, 70))
info = ImageInfo(
image=source,
filename="source",
extension=".png",
source_path="",
)
async def scenario():
node = BatchNanoBananaPro()
with (
patch(
"comfyui_o1key.nodes.batch_nano_banana._generate_single_async",
new=AsyncMock(return_value=[generated]),
),
patch("comfyui_o1key.nodes.batch_nano_banana._save_generated_image") as save,
):
result = await node._generate_single_task(
session=object(),
base_url="https://example.invalid",
api_key="test",
prompt="test",
model="nano-test",
resolution="2K",
aspect_ratio="1:1",
thinking_level=None,
images=[info],
output_folder=os.path.abspath(os.curdir),
task_index=0,
)
self.assertTrue(result["success"])
self.assertIs(save.call_args.args[0], generated)
asyncio.run(scenario())
def test_path_count_parser_accepts_new_and_legacy_labels(self):
self.assertEqual(_path_count_from_label("3个路径"), 3)
self.assertEqual(_path_count_from_label("4个文件夹"), 4)
self.assertEqual(_path_count_from_label("invalid"), 1)
def test_execute_still_accepts_legacy_numbered_inputs(self):
reference = object()
output = object()
with patch.object(BatchNanoBananaPro, "process_batch", return_value=(output,)) as process:
BatchNanoBananaPro.execute(
prompt="legacy",
模型="Nano Banana Pro",
图片路径1=r"D:\legacy\one",
图片路径4=r"D:\legacy\four",
图片配对模式="1*N",
图片随机抽取="4",
参考图1=reference,
)
arguments = process.call_args.kwargs
self.assertEqual(arguments["文件夹1"], r"D:\legacy\one")
self.assertEqual(arguments["文件夹4"], r"D:\legacy\four")
self.assertEqual(arguments["图片配对模式"], "全匹配")
self.assertNotIn("随机抽取路径", arguments)
self.assertIs(arguments["参考图1"], reference)
def test_same_index_pairing_uses_shortest_folder_length(self):
first = [object(), object(), object()]
second = [object(), object()]
pairs = BatchNanoBananaPro()._create_pairs([first, second], "同序号")
self.assertEqual(
pairs,
[(first[0], second[0]), (first[1], second[1])],
)
if __name__ == "__main__":
unittest.main(verbosity=2)