"""Offline regression tests for o1key's format-aware save boundary.""" import json import os from io import BytesIO from pathlib import Path from types import SimpleNamespace import sys import tempfile import unittest from PIL import Image ROOT = Path(__file__).resolve().parents[1] CUSTOM_NODES_ROOT = ROOT.parent COMFY_ROOT = CUSTOM_NODES_ROOT.parent sys.path.insert(0, str(COMFY_ROOT)) sys.path.insert(0, str(CUSTOM_NODES_ROOT)) from comfyui_o1key.utils.image_utils import pil_to_tensor # noqa: E402 from comfyui_o1key.utils.o1key_image_save import ( # noqa: E402 detect_image_format, normalize_naming_rule, normalize_save_location, save_temp_images, save_tensor_images, ) def _folder_paths(output_dir: str): return SimpleNamespace( get_save_image_path=lambda prefix, _root, _width, _height: ( output_dir, prefix, 1, "", prefix, ) ) def _encoded_image(image_format: str, *, mode="RGB") -> bytes: image = Image.new(mode, (6, 4), (20, 40, 60, 90) if mode == "RGBA" else (20, 40, 60)) buffer = BytesIO() kwargs = {"format": image_format} if image_format == "JPEG": kwargs.update(quality=91) if image_format == "WEBP": kwargs.update(lossless=True) image.save(buffer, **kwargs) image.close() return buffer.getvalue() class O1keyImageSaveTests(unittest.TestCase): def test_missing_naming_rule_defaults_to_custom_prefix(self): self.assertEqual(normalize_naming_rule(None), "自定义前缀") def test_save_location_supports_output_subfolders_and_absolute_directories(self): raw = _encoded_image("PNG") with tempfile.TemporaryDirectory() as temp_dir, tempfile.TemporaryDirectory() as output_dir: source = os.path.join(temp_dir, "provider.png") Path(source).write_bytes(raw) def get_save_image_path(prefix, root, _width, _height): normalized = prefix.replace("\\", "/") subfolder, filename = normalized.rsplit("/", 1) full_output = os.path.join(root, *subfolder.split("/")) return full_output, filename, 1, subfolder, prefix folder_paths = SimpleNamespace(get_save_image_path=get_save_image_path) results = save_temp_images( [source], "o1key", "原始", output_dir, folder_paths, save_location="project/session-a", naming_rule="自定义前缀", ) self.assertEqual(results[0]["subfolder"].replace("\\", "/"), "project/session-a") self.assertTrue(Path(output_dir, "project", "session-a", results[0]["filename"]).is_file()) self.assertEqual(normalize_save_location(""), "") self.assertEqual(normalize_save_location("output"), "") with tempfile.TemporaryDirectory() as external_dir: self.assertEqual( normalize_save_location(external_dir), os.path.normpath(external_dir), ) with self.assertRaisesRegex(ValueError, r"不能包含 \.\."): normalize_save_location("../outside") with self.assertRaisesRegex(ValueError, "完整绝对路径"): normalize_save_location("C:outside") raw = _encoded_image("PNG") with ( tempfile.TemporaryDirectory() as source_dir, tempfile.TemporaryDirectory() as output_dir, tempfile.TemporaryDirectory() as preview_dir, tempfile.TemporaryDirectory() as external_parent, ): source = os.path.join(source_dir, "provider.png") Path(source).write_bytes(raw) external_dir = os.path.join(external_parent, "new", "destination") def get_save_image_path(prefix, root, _width, _height): subfolder = os.path.dirname(os.path.normpath(prefix)) filename = os.path.basename(os.path.normpath(prefix)) return os.path.join(root, subfolder), filename, 1, subfolder, prefix folder_paths = SimpleNamespace( get_save_image_path=get_save_image_path, get_temp_directory=lambda: preview_dir, ) results = save_temp_images( [source], "external", "原始", output_dir, folder_paths, save_location=external_dir, naming_rule="自定义前缀", ) self.assertTrue(Path(external_dir, "external_00001_.png").is_file()) self.assertEqual(list(Path(output_dir).iterdir()), []) self.assertEqual(results[0]["filename"], "external_00001_.png") self.assertEqual(results[0]["type"], "temp") self.assertIs(results[0]["external_saved"], True) self.assertNotIn(external_dir, results[0]["subfolder"]) self.assertTrue( Path(preview_dir, results[0]["subfolder"], results[0]["filename"]).is_file() ) def test_metadata_bearing_jpeg_uses_native_recoverable_png_container(self): raw = _encoded_image("JPEG") with tempfile.TemporaryDirectory() as temp_dir, tempfile.TemporaryDirectory() as output_dir: source = os.path.join(temp_dir, "provider.jpg") Path(source).write_bytes(raw) workflow = {"nodes": [], "extra": {"note": "x" * 70000}} results = save_temp_images( [source], "o1key", "原始", output_dir, _folder_paths(output_dir), prompt={"1": {"class_type": "O1keyImageGenerator"}}, extra_pnginfo={"workflow": workflow}, naming_rule="自定义前缀", ) saved_path = Path(output_dir, results[0]["filename"]) saved = saved_path.read_bytes() self.assertEqual(results[0]["filename"], "o1key_00001_.png") self.assertEqual(detect_image_format(saved), "PNG") with Image.open(saved_path) as image: self.assertEqual( image.text["prompt"], '{"1": {"class_type": "O1keyImageGenerator"}}', ) self.assertEqual(image.text["workflow"], json.dumps(workflow)) def test_original_jpeg_without_workflow_metadata_keeps_provider_bytes(self): raw = _encoded_image("JPEG") with tempfile.TemporaryDirectory() as temp_dir, tempfile.TemporaryDirectory() as output_dir: source = os.path.join(temp_dir, "provider.jpg") Path(source).write_bytes(raw) results = save_temp_images( [source], "o1key", "原始", output_dir, _folder_paths(output_dir), naming_rule="自定义前缀", ) saved = Path(output_dir, results[0]["filename"]).read_bytes() self.assertEqual(results[0]["filename"], "o1key_00001_.jpg") self.assertEqual(saved, raw) def test_original_png_keeps_idat_and_embeds_native_workflow_fields(self): raw = _encoded_image("PNG", mode="RGBA") with tempfile.TemporaryDirectory() as temp_dir, tempfile.TemporaryDirectory() as output_dir: source = os.path.join(temp_dir, "provider.png") Path(source).write_bytes(raw) results = save_temp_images( [source], "o1key", "原始", output_dir, _folder_paths(output_dir), prompt={"prompt": True}, extra_pnginfo={"workflow": {"last_node_id": 2}}, naming_rule="自定义前缀", ) saved_path = Path(output_dir, results[0]["filename"]) saved = saved_path.read_bytes() self.assertEqual(detect_image_format(saved), "PNG") idat_offset = raw.index(b"IDAT") - 4 idat_length = int.from_bytes(raw[idat_offset : idat_offset + 4], "big") idat_chunk = raw[idat_offset : idat_offset + 12 + idat_length] self.assertIn(idat_chunk, saved) with Image.open(saved_path) as image: self.assertEqual(image.getchannel("A").getextrema(), (90, 90)) self.assertIn("prompt", image.text) self.assertIn("workflow", image.text) def test_explicit_png_jpg_and_webp_are_selected_only_by_the_save_node(self): raw = _encoded_image("JPEG") with tempfile.TemporaryDirectory() as temp_dir, tempfile.TemporaryDirectory() as output_dir: source = os.path.join(temp_dir, "provider.jpg") Path(source).write_bytes(raw) expected = {"png": "PNG", "jpg": "JPEG", "webp": "WEBP"} for save_format, image_format in expected.items(): results = save_temp_images( [source], f"o1key-{save_format}", save_format, output_dir, _folder_paths(output_dir), naming_rule="自定义前缀", ) saved = Path(output_dir, results[0]["filename"]).read_bytes() self.assertEqual(detect_image_format(saved), image_format) with Image.open(BytesIO(raw)) as source_image, Image.open( Path(output_dir, "o1key-webp_00001_.webp") ) as image: self.assertEqual(image.getpixel((0, 0)), source_image.getpixel((0, 0))) def test_original_webp_adds_workflow_exif_without_reencoding_pixels(self): raw = _encoded_image("WEBP", mode="RGBA") with tempfile.TemporaryDirectory() as temp_dir, tempfile.TemporaryDirectory() as output_dir: source = os.path.join(temp_dir, "provider.webp") Path(source).write_bytes(raw) results = save_temp_images( [source], "o1key-webp", "原始", output_dir, _folder_paths(output_dir), extra_pnginfo={"workflow": {"nodes": []}}, naming_rule="自定义前缀", ) saved_path = Path(output_dir, results[0]["filename"]) saved = saved_path.read_bytes() self.assertEqual(detect_image_format(saved), "WEBP") with Image.open(BytesIO(raw)) as before, Image.open(saved_path) as after: self.assertEqual(after.getpixel((0, 0)), before.getpixel((0, 0))) self.assertIn("workflow:", str(after.getexif().get(0x010F, ""))) def test_original_tensor_uses_provider_bytes_and_plain_tensor_falls_back_to_png(self): raw = _encoded_image("WEBP") with Image.open(BytesIO(raw)) as opened: opened.load() source = opened.copy() source.format = "WEBP" source._o1key_original_format = "WEBP" source._o1key_original_bytes = raw tensor = pil_to_tensor([source]) source.close() plain = tensor.clone() with tempfile.TemporaryDirectory() as output_dir: preserved = save_tensor_images( tensor, "preserved", "原始", output_dir, _folder_paths(output_dir), naming_rule="自定义前缀", ) fallback = save_tensor_images( plain, "fallback", "原始", output_dir, _folder_paths(output_dir), naming_rule="自定义前缀", ) self.assertTrue(preserved[0]["filename"].endswith(".webp")) self.assertTrue(fallback[0]["filename"].endswith(".png")) def test_main_image_rule_uses_the_first_reference_stem_and_never_overwrites(self): raw = _encoded_image("PNG") with tempfile.TemporaryDirectory() as temp_dir, tempfile.TemporaryDirectory() as output_dir: source_one = os.path.join(temp_dir, "provider-one.png") source_two = os.path.join(temp_dir, "provider-two.png") Path(source_one).write_bytes(raw) Path(source_two).write_bytes(raw) first = save_temp_images( [source_one, source_two], "ignored", "png", output_dir, _folder_paths(output_dir), naming_rule="和主图一致", main_filename="主图.jpeg", ) second = save_temp_images( [source_one], "ignored", "png", output_dir, _folder_paths(output_dir), naming_rule="和主图一致", main_filename="主图.jpeg", ) self.assertEqual([item["filename"] for item in first], ["主图.png", "主图1.png"]) self.assertEqual(second[0]["filename"], "主图2.png") self.assertEqual(len(list(Path(output_dir).glob("*.png"))), 3) def test_natural_number_rule_continues_after_existing_files(self): raw = _encoded_image("PNG") with tempfile.TemporaryDirectory() as temp_dir, tempfile.TemporaryDirectory() as output_dir: source = os.path.join(temp_dir, "provider.png") Path(source).write_bytes(raw) Path(output_dir, "1.png").write_bytes(b"keep-existing") results = save_temp_images( [source, source], "ignored", "png", output_dir, _folder_paths(output_dir), naming_rule="自然数字", ) self.assertEqual([item["filename"] for item in results], ["2.png", "3.png"]) self.assertEqual(Path(output_dir, "1.png").read_bytes(), b"keep-existing") if __name__ == "__main__": unittest.main(verbosity=2)