Files
DumasNodes/tests/test_dumas_image_nodes.py

125 lines
3.7 KiB
Python

import importlib
import os
import sys
import tempfile
import types
import unittest
class FakeImageArray:
def __init__(self, width=8, height=6):
self.shape = (height, width, 3)
def cpu(self):
return self
def numpy(self):
return self
def astype(self, _dtype):
return self
def __rmul__(self, _value):
return self
class FakeTensorBatch:
def __init__(self, width=8, height=6):
self.image = FakeImageArray(width=width, height=height)
def __getitem__(self, index):
if index != 0:
raise IndexError(index)
return self.image
class FakePILImage:
saved_paths = []
def save(self, path, compress_level=0):
self.saved_paths.append((path, compress_level))
class DumasImageNodeTests(unittest.TestCase):
@classmethod
def setUpClass(cls):
cls.temp_dir = tempfile.mkdtemp(prefix="dumas-image-node-")
fake_numpy = types.SimpleNamespace(
clip=lambda array, _low, _high: array,
uint8="uint8",
)
fake_pil_image_module = types.SimpleNamespace(fromarray=lambda _array: FakePILImage())
fake_pil_module = types.SimpleNamespace(Image=fake_pil_image_module)
fake_folder_paths = types.SimpleNamespace(
get_temp_directory=lambda: cls.temp_dir,
get_save_image_path=lambda prefix, _out, _width, _height: (
cls.temp_dir,
prefix,
1,
"",
prefix,
),
)
cls._saved_modules = {
name: sys.modules.get(name)
for name in ("numpy", "PIL", "PIL.Image", "folder_paths")
}
sys.modules["numpy"] = fake_numpy
sys.modules["PIL"] = fake_pil_module
sys.modules["PIL.Image"] = fake_pil_image_module
sys.modules["folder_paths"] = fake_folder_paths
cls.image_nodes = importlib.import_module("dumas_image_nodes")
@classmethod
def tearDownClass(cls):
for name, module in cls._saved_modules.items():
if module is None:
sys.modules.pop(name, None)
else:
sys.modules[name] = module
def setUp(self):
FakePILImage.saved_paths = []
def test_compare_images_returns_second_input_as_new_image(self):
node = self.image_nodes.DumasImageCompareNode()
image1 = FakeTensorBatch()
image2 = FakeTensorBatch()
result = node.compare_images(image1=image1, image2=image2)
self.assertIs(result["result"][0], image2)
self.assertEqual([item["slot"] for item in result["ui"]["images"]], [1, 2])
self.assertEqual(len(FakePILImage.saved_paths), 2)
def test_compare_images_falls_back_to_first_image(self):
node = self.image_nodes.DumasImageCompareNode()
image1 = FakeTensorBatch()
result = node.compare_images(image1=image1)
self.assertIs(result["result"][0], image1)
self.assertEqual([item["slot"] for item in result["ui"]["images"]], [1])
def test_compare_images_handles_missing_inputs(self):
node = self.image_nodes.DumasImageCompareNode()
result = node.compare_images()
self.assertIsNone(result["result"][0])
self.assertEqual(result["ui"]["images"], [])
def test_saved_filenames_use_dumas_prefix(self):
node = self.image_nodes.DumasImageCompareNode()
image1 = FakeTensorBatch()
result = node.compare_images(image1=image1)
self.assertTrue(result["ui"]["images"][0]["filename"].startswith("dumas_compare"))
self.assertTrue(os.path.basename(FakePILImage.saved_paths[0][0]).startswith("dumas_compare"))
if __name__ == "__main__":
unittest.main()