fix issues
This commit is contained in:
@@ -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
|
||||

|
||||
|
||||
## 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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
|
||||
|
||||
@@ -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
@@ -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
@@ -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}),
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
@@ -0,0 +1,5 @@
|
||||
"""Unit tests for ComfyUI-LogicUtils.
|
||||
|
||||
These tests are intentionally runnable outside a full ComfyUI install.
|
||||
"""
|
||||
|
||||
@@ -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}")
|
||||
|
||||
@@ -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",))
|
||||
@@ -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))
|
||||
@@ -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)
|
||||
@@ -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",))
|
||||
@@ -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))
|
||||
@@ -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)
|
||||
|
||||
@@ -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
@@ -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())
|
||||
|
||||
Reference in New Issue
Block a user