fix issues

This commit is contained in:
AngelBottomless (sleepy)
2026-01-21 04:46:32 +00:00
parent c8e10b174a
commit 44d9a0a2ce
19 changed files with 506 additions and 125 deletions
+11
View File
@@ -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`.
+8 -4
View File
@@ -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)
+29 -9
View File
@@ -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()
+71 -47
View File
@@ -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
+3 -20
View File
@@ -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:
"""
+13 -3
View File
@@ -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
+45 -23
View File
@@ -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
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
+8 -3
View File
@@ -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)
validate(fundamental_classes)
+9 -9
View File
@@ -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}),
},
}
+5
View File
@@ -0,0 +1,5 @@
"""Unit tests for ComfyUI-LogicUtils.
These tests are intentionally runnable outside a full ComfyUI install.
"""
+32
View File
@@ -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}")
+35
View File
@@ -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",))
+38
View File
@@ -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))
+42
View File
@@ -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)
+44
View File
@@ -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",))
+46
View File
@@ -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))
+21
View File
@@ -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)
+38
View File
@@ -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",))
+8 -7
View File
@@ -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())