From 44d9a0a2cec2ddcc9d264d8ee4da85c44086477f Mon Sep 17 00:00:00 2001 From: "AngelBottomless (sleepy)" Date: Wed, 21 Jan 2026 04:46:32 +0000 Subject: [PATCH] fix issues --- README.md | 11 ++++ conversion.py | 12 ++-- install.py | 38 ++++++++--- io_node.py | 118 ++++++++++++++++++++-------------- logic_gates.py | 23 +------ math_nodes.py | 16 ++++- nodes.py | 68 +++++++++++++------- pystructure.py | 11 +++- randomness.py | 18 +++--- tests/__init__.py | 5 ++ tests/import_utils.py | 32 +++++++++ tests/test_conversion.py | 35 ++++++++++ tests/test_crypto.py | 38 +++++++++++ tests/test_imgio_converter.py | 42 ++++++++++++ tests/test_logic_gates.py | 44 +++++++++++++ tests/test_math_nodes.py | 46 +++++++++++++ tests/test_nodes_import.py | 21 ++++++ tests/test_randomness.py | 38 +++++++++++ utils/tagger.py | 15 +++-- 19 files changed, 506 insertions(+), 125 deletions(-) create mode 100644 tests/__init__.py create mode 100644 tests/import_utils.py create mode 100644 tests/test_conversion.py create mode 100644 tests/test_crypto.py create mode 100644 tests/test_imgio_converter.py create mode 100644 tests/test_logic_gates.py create mode 100644 tests/test_math_nodes.py create mode 100644 tests/test_nodes_import.py create mode 100644 tests/test_randomness.py diff --git a/README.md b/README.md index 0c392f7..5a091a3 100644 --- a/README.md +++ b/README.md @@ -3,3 +3,14 @@ Logic Utilities for someone who wants ~~prime list calculation in comfyui~~ Proper documentation is being prepared, however there are too many nodes ![image](https://github.com/user-attachments/assets/8e388417-6912-41d7-98fa-798b50eacfda) + +## Tests + +Run from the repo root: + +`python -m unittest discover -s tests` + +## Notes + +- Auto-install is opt-in via `COMFYUI_LOGICUTILS_AUTO_INSTALL=1`. +- To force-disable the install hook, set `COMFYUI_LOGICUTILS_SKIP_INSTALL=1`. diff --git a/conversion.py b/conversion.py index 37d4ac2..7ad8d07 100644 --- a/conversion.py +++ b/conversion.py @@ -10,7 +10,11 @@ conversion_operators = { "Int" : int, "Float" : float, "Boolean" : bool, - "String" : str + "String" : str, + "Dict": dict, + "List": list, + "Tuple": tuple, + "Set": set, } def create_class(type_to): class_name = "ConvertAny2{}".format(type_to) @@ -83,9 +87,9 @@ class ConvertComboToString: CATEGORY = "Logic Gates" custom_name = "Convert Combo to String" def convertComboToString(self, combo, separator): - if isinstance(combo, (str, float, int, bool)): - return (combo,) - return (separator.join(combo),) + if isinstance(combo, (list, tuple)): + return (separator.join(str(item) for item in combo),) + return (str(combo),) for type_to in conversion_operators: create_class(type_to) diff --git a/install.py b/install.py index 0e78dbc..df20261 100644 --- a/install.py +++ b/install.py @@ -1,4 +1,5 @@ #https://github.com/ltdrdata/ComfyUI-Impact-Pack/blob/Main/install.py +import os import sys import subprocess import threading @@ -36,30 +37,49 @@ else: pip_install = [sys.executable, '-m', 'pip', 'install', "-U"] def initialization(): + auto_install = os.environ.get("COMFYUI_LOGICUTILS_AUTO_INSTALL", "").strip().lower() in { + "1", + "true", + "yes", + } + if not auto_install: + return + try: import piexif - except ImportError: + except Exception: run_installation("piexif") try: import chardet - except ImportError: + except Exception: run_installation("chardet") try: from imgutils.tagging import get_wd14_tags - except ImportError: - run_installation("dghs-imgutils[gpu]") + except Exception: + # dghs-imgutils currently pins numpy<2, which typically won't have wheels for + # the latest Python releases right away (e.g. Python 3.13 in ComfyUI portable). + if sys.version_info >= (3, 13): + print( + "Skipping auto-install of dghs-imgutils on Python >= 3.13 " + "(tagger nodes will be disabled unless installed manually)." + ) + else: + run_installation("dghs-imgutils[gpu]") try: from Crypto.PublicKey import RSA - except ImportError: + except Exception: run_installation("pycryptodome") def run_installation(pkg_name: str): print(f"Installing {pkg_name}...") - if process_wrap(pip_install + [pkg_name]) == 0: - print(f"Successfully installed {pkg_name}") - else: - print(f"Failed to install {pkg_name}") + try: + if process_wrap(pip_install + [pkg_name]) == 0: + print(f"Successfully installed {pkg_name}") + else: + print(f"Failed to install {pkg_name}") + except Exception as e: + print(f"Failed to install {pkg_name}: {e}") if __name__ == "__main__": initialization() diff --git a/io_node.py b/io_node.py index 5c2a1da..3b77bf2 100644 --- a/io_node.py +++ b/io_node.py @@ -22,8 +22,18 @@ from PIL import Image from PIL import ImageOps from PIL import ImageEnhance from PIL.PngImagePlugin import PngInfo -import folder_paths -from comfy.cli_args import args +try: + import folder_paths +except ModuleNotFoundError: + folder_paths = None +try: + from comfy.cli_args import args +except ModuleNotFoundError: + # Allow importing this module outside a full ComfyUI install (e.g. unit tests). + class _Args: + disable_metadata = True + + args = _Args() import filelock import tempfile @@ -248,20 +258,26 @@ class SaveImageCustomNode: CATEGORY = "image" custom_name = "Save Image Custom Node" - def save_images( - self, - images, - filename_prefix="ComfyUI", - subfolder_dir="", - prompt=None, - extra_pnginfo=None, - ): - if images is None: # sometimes images is empty - images = [] - filename_prefix += self.prefix_append - throw_if_parent_or_root_access(filename_prefix) - throw_if_parent_or_root_access(subfolder_dir) - output_dir = os.path.join(self.output_dir, subfolder_dir) + def save_images( + self, + images, + filename_prefix="ComfyUI", + subfolder_dir="", + prompt=None, + extra_pnginfo=None, + ): + # `images` can be None or empty in some edge cases. + if images is None: + return {"ui": {"images": []}, "outputs": {"images": ""}} + if isinstance(images, torch.Tensor) and images.shape[0] == 0: + return {"ui": {"images": []}, "outputs": {"images": ""}} + if isinstance(images, (list, tuple)) and len(images) == 0: + return {"ui": {"images": []}, "outputs": {"images": ""}} + + filename_prefix += self.prefix_append + throw_if_parent_or_root_access(filename_prefix) + throw_if_parent_or_root_access(subfolder_dir) + output_dir = os.path.join(self.output_dir, subfolder_dir) full_output_folder, filename, counter, subfolder, filename_prefix = ( folder_paths.get_save_image_path( filename_prefix, output_dir, images[0].shape[1], images[0].shape[0] @@ -690,7 +706,7 @@ class ConcatTwoImagesNode: @fundamental_node -class SaveCustomJPGNode: +class SaveCustomJPGNode: def __init__(self): self.output_dir = folder_paths.get_output_directory() self.type = "output" @@ -721,21 +737,25 @@ class SaveCustomJPGNode: CATEGORY = "image" custom_name = "Save Custom JPG Node" - def save_images( - self, - images, - filename_prefix="ComfyUI", - subfolder_dir="", - prompt=None, - extra_pnginfo=None, - quality=95, - optimize=True, - metadata_string="", - ): - if images is None: - images = [] - if not isinstance(images, (list, tuple, torch.Tensor)): - images = [images] + def save_images( + self, + images, + filename_prefix="ComfyUI", + subfolder_dir="", + prompt=None, + extra_pnginfo=None, + quality=95, + optimize=True, + metadata_string="", + ): + if images is None: + return {"ui": {"images": []}, "outputs": {"images": ""}} + if isinstance(images, torch.Tensor) and images.shape[0] == 0: + return {"ui": {"images": []}, "outputs": {"images": ""}} + if isinstance(images, (list, tuple)) and len(images) == 0: + return {"ui": {"images": []}, "outputs": {"images": ""}} + if not isinstance(images, (list, tuple, torch.Tensor)): + images = [images] throw_if_parent_or_root_access(filename_prefix) throw_if_parent_or_root_access(subfolder_dir) @@ -819,7 +839,7 @@ class SaveCustomJPGNode: @fundamental_node -class SaveImageWebpCustomNode: +class SaveImageWebpCustomNode: def __init__(self): self.output_dir = folder_paths.get_output_directory() self.type = "output" @@ -853,24 +873,28 @@ class SaveImageWebpCustomNode: CATEGORY = "image" custom_name = "Save Image Webp Node" - def save_images( - self, - images, - filename_prefix="ComfyUI", - subfolder_dir="", - prompt=None, - extra_pnginfo=None, + def save_images( + self, + images, + filename_prefix="ComfyUI", + subfolder_dir="", + prompt=None, + extra_pnginfo=None, quality=100, lossless=False, compression=4, optimize=False, - metadata_string="", - optional_additional_metadata="", - ): - if images is None: # sometimes images is empty - images = [] - if not isinstance(images, (list, tuple, torch.Tensor)): - images = [images] + metadata_string="", + optional_additional_metadata="", + ): + if images is None: # sometimes images is empty + return {"ui": {"images": []}, "outputs": {"images": ""}} + if isinstance(images, torch.Tensor) and images.shape[0] == 0: + return {"ui": {"images": []}, "outputs": {"images": ""}} + if isinstance(images, (list, tuple)) and len(images) == 0: + return {"ui": {"images": []}, "outputs": {"images": ""}} + if not isinstance(images, (list, tuple, torch.Tensor)): + images = [images] throw_if_parent_or_root_access(filename_prefix) throw_if_parent_or_root_access(subfolder_dir) filename_prefix += self.prefix_append diff --git a/logic_gates.py b/logic_gates.py index 0e14614..e591fda 100644 --- a/logic_gates.py +++ b/logic_gates.py @@ -8,25 +8,6 @@ from .autonode import node_wrapper, get_node_names_mappings, validate, anytype classes = [] node = node_wrapper(classes) -@node -class LogicGateCompare: - """ - Returns 1 if input1 > input2, 0 otherwise - """ - RETURN_TYPES = ("BOOLEAN",) - @classmethod - def INPUT_TYPES(s): - return { - "required": { - "input1": (anytype, {"default": 0.0}), - "input2": (anytype, {"default": 0.0}), - } - } - FUNCTION = "compareFloat" - CATEGORY = "Logic Gates" - custom_name = "ABiggerThanB" - def compareFloat(self, input1, input2): - return (True if input1 > input2 else False,) @node class LogicGateInvertBasic: """ @@ -86,7 +67,9 @@ class LogicGateBitwiseShift: # validate input2 if abs(input2) > 32: raise ValueError("input2 must be between -32 and 32") - return (input1 << input2,) + if input2 >= 0: + return (input1 << input2,) + return (input1 >> abs(input2),) @node class LogicGateBitwiseAnd: """ diff --git a/math_nodes.py b/math_nodes.py index 0fb0336..c0e6365 100644 --- a/math_nodes.py +++ b/math_nodes.py @@ -138,9 +138,18 @@ class PowerNode: CATEGORY = "Math" custom_name = "Power" def power(self, input1, power): - # validate power with log scale, prevent overflow - log_val = math.log(abs(input1), 10) - if log_val * power > 100 or log_val == 0: + abs_input = abs(input1) + # fast paths for values that won't overflow digit-wise + if abs_input == 0: + if power < 0: + raise ZeroDivisionError("0 cannot be raised to a negative power") + return (math.pow(input1, power),) + if abs_input == 1: + return (math.pow(input1, power),) + + # validate power with log10 scale, prevent huge magnitudes + log10_abs = math.log10(abs_input) + if (log10_abs * power) > 100: raise OverflowError("Power is too large, exceeds 100 digits") return (math.pow(input1, power),) @@ -284,6 +293,7 @@ class RAMPNode: def ramp(self, input1): return (max(0, input1),) +@node class ModuloNode: """ Returns the modulo of a number diff --git a/nodes.py b/nodes.py index 22b281d..3846702 100644 --- a/nodes.py +++ b/nodes.py @@ -1,15 +1,37 @@ +import os + +from .install import initialization + + +def _running_in_comfyui() -> bool: + try: + import folder_paths # noqa: F401 + except Exception: + return False + return True + + +_IN_COMFYUI = _running_in_comfyui() +_SKIP_INSTALL = os.environ.get("COMFYUI_LOGICUTILS_SKIP_INSTALL", "").strip().lower() in { + "1", + "true", + "yes", +} + +if _IN_COMFYUI and not _SKIP_INSTALL: + initialization() -from .install import initialization - -initialization() - -from .logic_gates import CLASS_MAPPINGS as LogicMapping, CLASS_NAMES as LogicNames -from .randomness import CLASS_MAPPINGS as RandomMapping, CLASS_NAMES as RandomNames -from .conversion import CLASS_MAPPINGS as ConversionMapping, CLASS_NAMES as ConversionNames -from .math_nodes import CLASS_MAPPINGS as MathMapping, CLASS_NAMES as MathNames -from .io_node import CLASS_MAPPINGS as IOMapping, CLASS_NAMES as IONames -from .auxilary import CLASS_MAPPINGS as AuxilaryMapping, CLASS_NAMES as AuxilaryNames -from .external import CLASS_MAPPINGS as ExternalMapping, CLASS_NAMES as ExternalNames +from .logic_gates import CLASS_MAPPINGS as LogicMapping, CLASS_NAMES as LogicNames +from .randomness import CLASS_MAPPINGS as RandomMapping, CLASS_NAMES as RandomNames +from .conversion import CLASS_MAPPINGS as ConversionMapping, CLASS_NAMES as ConversionNames +from .math_nodes import CLASS_MAPPINGS as MathMapping, CLASS_NAMES as MathNames +if _IN_COMFYUI: + from .io_node import CLASS_MAPPINGS as IOMapping, CLASS_NAMES as IONames +else: + IOMapping = {} + IONames = {} +from .auxilary import CLASS_MAPPINGS as AuxilaryMapping, CLASS_NAMES as AuxilaryNames +from .external import CLASS_MAPPINGS as ExternalMapping, CLASS_NAMES as ExternalNames @@ -37,15 +59,15 @@ NODE_DISPLAY_NAME_MAPPINGS.update(ExternalNames) NODE_DISPLAY_NAME_MAPPINGS.update(AuxilaryNames) -try: - from .pystructure import CLASS_MAPPINGS as PyStructureMapping, CLASS_NAMES as PyStructureNames - NODE_CLASS_MAPPINGS.update(PyStructureMapping) - NODE_DISPLAY_NAME_MAPPINGS.update(PyStructureNames) -except ImportError: - pass -try: - from .crypto import CLASS_MAPPINGS as SecureMapping, CLASS_NAMES as SecureNames - NODE_CLASS_MAPPINGS.update(SecureMapping) - NODE_DISPLAY_NAME_MAPPINGS.update(SecureNames) -except ImportError: - pass \ No newline at end of file +try: + from .pystructure import CLASS_MAPPINGS as PyStructureMapping, CLASS_NAMES as PyStructureNames + NODE_CLASS_MAPPINGS.update(PyStructureMapping) + NODE_DISPLAY_NAME_MAPPINGS.update(PyStructureNames) +except Exception: + pass +try: + from .crypto import CLASS_MAPPINGS as SecureMapping, CLASS_NAMES as SecureNames + NODE_CLASS_MAPPINGS.update(SecureMapping) + NODE_DISPLAY_NAME_MAPPINGS.update(SecureNames) +except Exception: + pass diff --git a/pystructure.py b/pystructure.py index db2a4d3..4b32091 100644 --- a/pystructure.py +++ b/pystructure.py @@ -431,7 +431,7 @@ class GlobalVarGetNode(NewPointer): CATEGORY = "Data" custom_name="Pyobjects/Global Var Get" - def global_var_get(self, key): + def global_var_get(self, key, trigger=None): print("GlobalVarGetNode:", GLOBAL_STORAGE) return (GLOBAL_STORAGE.get(key, None),) @@ -440,7 +440,12 @@ class GlobalVarGetNode(NewPointer): return { "required": { "key": ("STRING", {"default": "my_key"}), - } + }, + "optional": { + # Optional dependency input to force execution ordering. + # Connect a Global Var Set output here to guarantee the Set runs first. + "trigger": (anytype,), + }, } @fundamental_node @@ -1023,4 +1028,4 @@ class SetToListNode(NewPointer): ############################################################################## CLASS_MAPPINGS, CLASS_NAMES = get_node_names_mappings(fundamental_classes) -validate(fundamental_classes) \ No newline at end of file +validate(fundamental_classes) diff --git a/randomness.py b/randomness.py index 18ebf72..27f48ee 100644 --- a/randomness.py +++ b/randomness.py @@ -175,7 +175,7 @@ class UniformRandomFloat(RandomGuaranteedClass): pass def generate(self, min_val, max_val, decimal_places, seed=0): if min_val > max_val: - return min_val + return (min_val,) instance = random.Random(seed) value = instance.uniform(min_val, max_val) # prune to decimal places - 0 = int, 1 = 1 decimal place,... @@ -207,7 +207,7 @@ class TriangularRandomFloat(RandomGuaranteedClass): pass def generate(self, low, high, mode, seed=0): if low > high: - return low + return (low,) instance = random.Random(seed) value = instance.triangular(low, high, mode) return (value,) @@ -377,7 +377,7 @@ class UniformRandomInt(RandomGuaranteedClass): pass def generate(self, min_val, max_val, seed=0): if min_val > max_val: - return min_val + return (min_val,) instance = random.Random(seed) value = instance.randint(min_val, max_val) #print(f"Selected {value} from {min_val} to {max_val}") @@ -635,7 +635,7 @@ class CounterFloat(RandomGuaranteedClass): "start": ("FLOAT", { "default": 0.0, "min": -(2**63-1), "max": (2**63-1), "step": 1.0, "display": "number" }), }, "optional": { - "reset": ("BOOLEAN"), + "reset": ("BOOLEAN", {"default": False}), "step": ("FLOAT", { "default": 1.0, "min": -(2**63-1), "max": (2**63-1), "step": 1.0, "display": "number" }), }, } @@ -651,7 +651,7 @@ class YieldableIteratorString(RandomGuaranteedClass): If reset is True, then it starts from the beginning """ def __init__(self): - self.index = 0 + self.index = -1 def generate(self, input_string, separator, reset): choices = input_string.split(separator) if reset: @@ -668,7 +668,7 @@ class YieldableIteratorString(RandomGuaranteedClass): "required": { "input_string": ("STRING", { "default": "a$b$c", "display": "text" }), "separator": ("STRING", { "default": "$", "display": "text" }), - "reset": ("BOOLEAN"), + "reset": ("BOOLEAN", {"default": False}), }, } RETURN_TYPES = ("STRING",) @@ -692,11 +692,11 @@ class YieldableIteratorInt(RandomGuaranteedClass): if reset: self.iterator = None if self.iterator is None: - self.iterator = range(start, end, step) + self.iterator = iter(range(start, end, step)) try: value = next(self.iterator) except StopIteration: - self.iterator = range(start, end, step) + self.iterator = iter(range(start, end, step)) value = next(self.iterator) return (value,) @classmethod @@ -706,7 +706,7 @@ class YieldableIteratorInt(RandomGuaranteedClass): "start": ("INT", { "default": 0, "min": -(2**63-1), "max": (2**63-1), "step": 1, "display": "number" }), "end": ("INT", { "default": 10, "min": -(2**63-1), "max": (2**63-1), "step": 1, "display": "number" }), "step": ("INT", { "default": 1, "min": -(2**63-1), "max": (2**63-1), "step": 1, "display": "number" }), - "reset": ("BOOLEAN"), + "reset": ("BOOLEAN", {"default": False}), }, } diff --git a/tests/__init__.py b/tests/__init__.py new file mode 100644 index 0000000..fc81267 --- /dev/null +++ b/tests/__init__.py @@ -0,0 +1,5 @@ +"""Unit tests for ComfyUI-LogicUtils. + +These tests are intentionally runnable outside a full ComfyUI install. +""" + diff --git a/tests/import_utils.py b/tests/import_utils.py new file mode 100644 index 0000000..bac9e6c --- /dev/null +++ b/tests/import_utils.py @@ -0,0 +1,32 @@ +from __future__ import annotations + +import importlib +import sys +import types +from pathlib import Path + + +_PKG_NAME = "comfyui_logicutils" + + +def ensure_local_package() -> str: + """Expose the repo's sources as an importable package for unit tests. + + The upstream folder name contains a hyphen, which isn't a valid Python import name. + We create an in-memory package module (with a __path__) so intra-package relative + imports like `from .autonode import ...` work normally. + """ + if _PKG_NAME in sys.modules: + return _PKG_NAME + + repo_root = Path(__file__).resolve().parents[1] + package = types.ModuleType(_PKG_NAME) + package.__path__ = [str(repo_root)] + sys.modules[_PKG_NAME] = package + return _PKG_NAME + + +def import_local(module: str): + pkg = ensure_local_package() + return importlib.import_module(f"{pkg}.{module}") + diff --git a/tests/test_conversion.py b/tests/test_conversion.py new file mode 100644 index 0000000..cd0397f --- /dev/null +++ b/tests/test_conversion.py @@ -0,0 +1,35 @@ +import unittest + +from import_utils import import_local + + +class TestConversion(unittest.TestCase): + @classmethod + def setUpClass(cls): + cls.conversion = import_local("conversion") + + def test_generated_conversion_nodes_exist(self): + expected = { + "ConvertAny2Int", + "ConvertAny2Float", + "ConvertAny2Boolean", + "ConvertAny2String", + "ConvertAny2Dict", + "ConvertAny2List", + "ConvertAny2Tuple", + "ConvertAny2Set", + } + self.assertTrue(expected.issubset(set(self.conversion.CLASS_MAPPINGS.keys()))) + + def test_convert_combo_to_string_always_returns_string(self): + Node = self.conversion.CLASS_MAPPINGS["ConvertComboToString"] + node = Node() + self.assertEqual(node.convertComboToString(["a", "b"], "|"), ("a|b",)) + self.assertEqual(node.convertComboToString([1, 2, 3], ","), ("1,2,3",)) + self.assertEqual(node.convertComboToString(123, "|"), ("123",)) + + def test_string_list_to_combo(self): + Node = self.conversion.CLASS_MAPPINGS["StringListToCombo"] + node = Node() + self.assertEqual(node.stringListToCombo("a$b$c", "$", 1), ("b",)) + self.assertEqual(node.stringListToCombo("abc", "$", 0), ("abc",)) diff --git a/tests/test_crypto.py b/tests/test_crypto.py new file mode 100644 index 0000000..fbb4391 --- /dev/null +++ b/tests/test_crypto.py @@ -0,0 +1,38 @@ +import unittest + +import numpy as np +import torch +from Crypto.PublicKey import RSA + +from import_utils import import_local + + +class TestCryptoNodes(unittest.TestCase): + @classmethod + def setUpClass(cls): + cls.crypto = import_local("crypto") + + def test_encrypt_decrypt_roundtrip(self): + Encrypt = self.crypto.CLASS_MAPPINGS["SecureBase64Encrypt"] + Decrypt = self.crypto.CLASS_MAPPINGS["SecureWebPDecrypt"] + + key = RSA.generate(1024) + private_pem = key.export_key().decode("utf-8") + public_pem = key.publickey().export_key().decode("utf-8") + + # Create a small deterministic RGB image tensor with values aligned to 1/255. + arr = np.zeros((8, 8, 3), dtype=np.uint8) + arr[0, 0] = [255, 0, 0] + arr[0, 1] = [0, 255, 0] + arr[0, 2] = [0, 0, 255] + img = torch.from_numpy(arr.astype(np.float32) / 255.0).unsqueeze(0) + + enc = Encrypt() + encrypted_b64 = enc.encrypted_base64(img, public_pem)[0] + self.assertIsInstance(encrypted_b64, str) + + dec = Decrypt() + decrypted = dec.decrypt_image(encrypted_b64, private_pem)[0] + self.assertTrue(torch.is_tensor(decrypted)) + self.assertEqual(tuple(decrypted.shape), tuple(img.shape)) + self.assertTrue(torch.allclose(decrypted, img, atol=1 / 255, rtol=0)) diff --git a/tests/test_imgio_converter.py b/tests/test_imgio_converter.py new file mode 100644 index 0000000..0e3b21b --- /dev/null +++ b/tests/test_imgio_converter.py @@ -0,0 +1,42 @@ +import unittest + +import numpy as np +import torch +from PIL import Image + +from import_utils import import_local + + +class TestImgIOConverter(unittest.TestCase): + @classmethod + def setUpClass(cls): + cls.converter = import_local("imgio.converter") + + def test_handle_rgba_composite_outputs_rgb(self): + img = Image.new("RGBA", (4, 4), (255, 0, 0, 128)) + out = self.converter.handle_rgba_composite(img) + self.assertEqual(out.mode, "RGB") + + def test_classify_pil_numpy_torch(self): + IOConverter = self.converter.IOConverter + img = Image.new("RGB", (2, 2), (0, 0, 0)) + arr = np.zeros((1, 2, 2, 3), dtype=np.uint8) + ten = torch.zeros((1, 2, 2, 3), dtype=torch.float32) + + self.assertEqual(IOConverter.classify(img), IOConverter.InputType.PIL) + self.assertEqual(IOConverter.classify(arr), IOConverter.InputType.NUMPY) + self.assertEqual(IOConverter.classify(ten), IOConverter.InputType.TORCH) + + def test_base64_roundtrip_pil(self): + IOConverter = self.converter.IOConverter + img = Image.new("RGB", (3, 5), (10, 20, 30)) + b64 = IOConverter.convert_to_base64(img, format="PNG") + out = IOConverter.convert_to_pil(b64) + self.assertEqual(out.size, img.size) + self.assertEqual(out.mode, "RGB") + + def test_gzip_base64_string_roundtrip(self): + IOConverter = self.converter.IOConverter + text = "hello world" + b64 = IOConverter.string_to_base64(text, gzip_compress=True) + self.assertEqual(IOConverter.read_maybe_gzip_base64(b64), text) diff --git a/tests/test_logic_gates.py b/tests/test_logic_gates.py new file mode 100644 index 0000000..5684bff --- /dev/null +++ b/tests/test_logic_gates.py @@ -0,0 +1,44 @@ +import unittest + +from import_utils import import_local + + +class TestLogicGates(unittest.TestCase): + @classmethod + def setUpClass(cls): + cls.logic_gates = import_local("logic_gates") + + def test_mappings_have_unique_keys(self): + module = self.logic_gates + self.assertEqual(len(module.CLASS_MAPPINGS), len(set(module.CLASS_MAPPINGS.keys()))) + self.assertEqual(len(module.CLASS_MAPPINGS), len(module.classes)) + + def test_bitwise_shift_supports_negative_shift(self): + Shift = self.logic_gates.CLASS_MAPPINGS["LogicGateBitwiseShift"] + node = Shift() + self.assertEqual(node.bitwiseShift(8, 1), (16,)) + self.assertEqual(node.bitwiseShift(8, -1), (4,)) + + def test_bitwise_shift_validates_range(self): + Shift = self.logic_gates.CLASS_MAPPINGS["LogicGateBitwiseShift"] + node = Shift() + with self.assertRaises(ValueError): + node.bitwiseShift(1, 33) + + def test_compare_gate(self): + Compare = self.logic_gates.CLASS_MAPPINGS["LogicGateCompare"] + node = Compare() + self.assertEqual(node.compareInt(2, 1), (True,)) + self.assertEqual(node.compareInt(1, 2), (False,)) + + def test_memory_node_flip_flop(self): + Memory = self.logic_gates.CLASS_MAPPINGS["MemoryNode"] + node = Memory() + self.assertEqual(node.memory("a", 0), ("a",)) + self.assertEqual(node.memory("b", 0), ("a",)) + self.assertEqual(node.memory("b", 1), ("b",)) + + def test_replace_string_regex(self): + Replace = self.logic_gates.CLASS_MAPPINGS["ReplaceString"] + node = Replace() + self.assertEqual(node.replace("hello", "l+", "x"), ("hexo",)) diff --git a/tests/test_math_nodes.py b/tests/test_math_nodes.py new file mode 100644 index 0000000..c585b82 --- /dev/null +++ b/tests/test_math_nodes.py @@ -0,0 +1,46 @@ +import unittest +from unittest.mock import patch + +from import_utils import import_local + + +class TestMathNodes(unittest.TestCase): + @classmethod + def setUpClass(cls): + cls.math_nodes = import_local("math_nodes") + + def test_power_allows_one_and_zero(self): + Power = self.math_nodes.CLASS_MAPPINGS["PowerNode"] + node = Power() + self.assertEqual(node.power(1, 5), (1.0,)) + self.assertEqual(node.power(0, 2), (0.0,)) + self.assertEqual(node.power(0, 0), (1.0,)) + with self.assertRaises(ZeroDivisionError): + node.power(0, -1) + + def test_power_overflow_guard(self): + Power = self.math_nodes.CLASS_MAPPINGS["PowerNode"] + node = Power() + with self.assertRaises(OverflowError): + node.power(10, 101) + + def test_modulo_node_is_registered(self): + self.assertIn("ModuloNode", self.math_nodes.CLASS_MAPPINGS) + Modulo = self.math_nodes.CLASS_MAPPINGS["ModuloNode"] + self.assertEqual(Modulo().modulo(10, 3), (1,)) + + def test_is_prime_small(self): + is_prime_small = self.math_nodes.is_prime_small + self.assertFalse(is_prime_small(0)) + self.assertFalse(is_prime_small(1)) + self.assertTrue(is_prime_small(2)) + self.assertTrue(is_prime_small(3)) + self.assertFalse(is_prime_small(4)) + self.assertTrue(is_prime_small(7919)) + + def test_is_prime_miller_rabin_deterministic_seed(self): + is_prime_miller_rabin = self.math_nodes.is_prime_miller_rabin + # Force a deterministic base for stable tests. + with patch("random.randrange", return_value=2): + self.assertTrue(is_prime_miller_rabin(1_000_000_007, k=3)) + self.assertFalse(is_prime_miller_rabin(1517, k=3)) diff --git a/tests/test_nodes_import.py b/tests/test_nodes_import.py new file mode 100644 index 0000000..779caea --- /dev/null +++ b/tests/test_nodes_import.py @@ -0,0 +1,21 @@ +import unittest + +from import_utils import import_local + + +class TestNodesImport(unittest.TestCase): + def test_nodes_imports_outside_comfyui(self): + nodes = import_local("nodes") + self.assertTrue(nodes.NODE_CLASS_MAPPINGS) + self.assertIn("LogicGateCompare", nodes.NODE_CLASS_MAPPINGS) + + try: + import folder_paths # noqa: F401 + in_comfyui = True + except ModuleNotFoundError: + in_comfyui = False + + if not in_comfyui: + # io_node is intentionally skipped outside ComfyUI. + self.assertNotIn("SleepNodeAny", nodes.NODE_CLASS_MAPPINGS) + diff --git a/tests/test_randomness.py b/tests/test_randomness.py new file mode 100644 index 0000000..6d98104 --- /dev/null +++ b/tests/test_randomness.py @@ -0,0 +1,38 @@ +import unittest + +from import_utils import import_local + + +class TestRandomness(unittest.TestCase): + @classmethod + def setUpClass(cls): + cls.randomness = import_local("randomness") + + def test_uniform_random_float_fallback_is_tuple(self): + Node = self.randomness.CLASS_MAPPINGS["UniformRandomFloat"] + self.assertEqual(Node().generate(2.0, 1.0, 1, seed=0), (2.0,)) + + def test_uniform_random_int_fallback_is_tuple(self): + Node = self.randomness.CLASS_MAPPINGS["UniformRandomInt"] + self.assertEqual(Node().generate(2, 1, seed=0), (2,)) + + def test_triangular_random_float_fallback_is_tuple(self): + Node = self.randomness.CLASS_MAPPINGS["TriangularRandomFloat"] + self.assertEqual(Node().generate(2.0, 1.0, 1.5, seed=0), (2.0,)) + + def test_yieldable_iterator_int_is_iterable_and_wraps(self): + Node = self.randomness.CLASS_MAPPINGS["YieldableIteratorInt"] + it = Node() + + self.assertEqual(it.generate(0, 3, 1, True), (0,)) + self.assertEqual(it.generate(0, 3, 1, False), (1,)) + self.assertEqual(it.generate(0, 3, 1, False), (2,)) + self.assertEqual(it.generate(0, 3, 1, False), (0,)) # wraps + + def test_yieldable_iterator_string_starts_at_first(self): + Node = self.randomness.CLASS_MAPPINGS["YieldableIteratorString"] + it = Node() + + self.assertEqual(it.generate("a$b$c", "$", False), ("a",)) + self.assertEqual(it.generate("a$b$c", "$", False), ("b",)) + self.assertEqual(it.generate("a$b$c", "$", True), ("a",)) diff --git a/utils/tagger.py b/utils/tagger.py index 2394bec..e5423c5 100644 --- a/utils/tagger.py +++ b/utils/tagger.py @@ -1,9 +1,13 @@ try: from imgutils.tagging import get_wd14_tags from imgutils.tagging.wd14 import MODEL_NAMES as tagger_model_names -except ImportError: - def get_wd14_tags(image_path): - raise Exception("Tagger feature not available, please install dghs-imgutils") +except Exception as e: + _tagger_import_error = e + + def get_wd14_tags(image_path, model_name=None): + raise RuntimeError( + "Tagger feature not available. Install 'dghs-imgutils' to enable it." + ) from _tagger_import_error tagger_model_names = { "EVA02_Large": None, "ViT_Large": None, @@ -39,7 +43,4 @@ def get_tags(image_path:Union[str, Image.Image], threshold:float = 0.4, replace: result['tags'] = [replace_underscore(tag) for tag in result['tags']] result['chars'] = [replace_underscore(tag) for tag in result['chars']] return result -try: - tagger_keys = list(tagger_model_names.keys()) -except NameError: - tagger_keys = [] +tagger_keys = list(tagger_model_names.keys())