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