Author SHA1 Message Date
AngelBottomless (sleepy) 0cb60bc8aa I think this is something env dependent 2026-01-21 07:40:43 +00:00
AngelBottomless (sleepy) 1da10d7a31 fix 2026-01-21 07:28:53 +00:00
AngelBottomless (sleepy) a39013912d enable auto install 2026-01-21 07:18:56 +00:00
AngelBottomless (sleepy) 835e382cd9 add two nodes 2026-01-21 07:03:24 +00:00
AngelBottomless (sleepy) 44d9a0a2ce fix issues 2026-01-21 04:46:32 +00:00
AngelBottomless c8e10b174a bump version, should resolve everything 2026-01-21 13:02:19 +09:00
aria1th 54bee29ced fix bugs 2025-10-19 16:50:40 +09:00
AngelBottomless 9fc69cb6a5 Merge pull request #20 from gshawn3/main
Fix problematic in-place mutations for dicts, lists, and sets
2025-10-19 16:49:54 +09:00
George S d4a1fc5f0b fix problematic in-place mutations 2025-10-12 01:16:40 +00:00
AngelBottomless 1e4c7465eb Add Feature Except Character Name
TODO: implement key based one
2025-08-09 05:12:51 +09:00
aria1th 5992e91930 try optional dependency 2025-05-23 01:16:35 +09:00
aria1th 60f8f1187c optional metadata 2025-05-13 05:16:32 +09:00
aria1th ae00f1ff8b How 2025-05-07 15:00:06 +09:00
AngelBottomless 1136141406 fix bug 2025-03-19 15:11:54 +09:00
AngelBottomless 915e2e9e81 Merge pull request #12 from ComfyNodePRs/update-publish-yaml
Update Github Action for Publishing to Comfy Registry
2025-03-14 17:34:03 +09:00
aria1th d600e0bec7 increase timeout 2025-03-11 11:50:34 +09:00
aria1th 33a0a38ecb fix critical bug 2025-03-10 02:06:31 +09:00
AngelBottomless e07784201f more nodes, cleanups, proper randomness 2025-03-08 05:51:32 +09:00
aria1th f8302d7ac0 uhh do you need preventing MITM 2025-03-08 05:16:29 +09:00
aria1th 9533c09491 add SaveCustomJPGNode 2025-03-07 23:33:48 +09:00
AngelBottomless 71bc2f0cc1 handle batches 2025-03-05 17:32:34 +09:00
AngelBottomless 91e76fc2ee Update io_node.py 2025-03-05 17:32:17 +09:00
AngelBottomless f0c8f9a8bb fix max range 2025-03-05 16:18:08 +09:00
AngelBottomless 09345b8d70 Create README.md 2025-03-04 13:17:29 +09:00
aria1th cbbd61ef40 do we need this? 2025-03-04 13:13:40 +09:00
aria1th 8e0935831a random operation would work always from now 2025-03-04 13:04:34 +09:00
aria1th c1e4c84a70 reworked & fix 2025-03-04 13:03:25 +09:00
aria1th e9e34dd355 some fix 2025-03-04 12:38:52 +09:00
aria1th 2d826bd624 ehhh 2025-03-04 12:34:27 +09:00
aria1th a9a8258c1f fix 2025-03-04 12:22:59 +09:00
aria1th f62821172a optional 2025-03-04 12:21:12 +09:00
aria1th 9624d93ac4 fix behavior 2025-03-04 12:17:57 +09:00
aria1th 46d62f8d6d Do you want madness? 2025-03-04 12:05:42 +09:00
aria1th 247577ae71 add more nodes, fix bug 2025-03-04 11:04:13 +09:00
aria1th 56b5cf8267 increase max int 2025-03-04 03:15:29 +09:00
aria1th 3dbae89658 add DumpJsonl Node 2025-03-04 02:21:22 +09:00
aria1th 2109871a89 add Counters 2025-03-04 02:01:38 +09:00
aria1th 6f138ebd4b Update io_node.py 2025-03-02 16:00:06 +09:00
aria1th ff4602f6e6 Update randomness.py 2025-03-02 15:56:38 +09:00
aria1th 05f05534b9 Update io_node.py 2025-03-02 15:01:48 +09:00
aria1th 20a4170bb8 handle raw image 2025-03-02 14:57:58 +09:00
aria1th 287997120b uhhhh 2025-03-02 14:54:43 +09:00
aria1th 29630d22e4 is this really happening somewhere 2025-03-02 14:52:07 +09:00
aria1th 67b532a1d5 fix? 2025-03-02 14:26:01 +09:00
aria1th 07ad78cdee ensure 2025-03-02 04:43:31 +09:00
aria1th aebcbc8c97 Update randomness.py 2025-03-02 04:33:18 +09:00
aria1th 1bdd981fa1 implement random w/h 2025-03-02 04:26:39 +09:00
AngelBottomless 2451f43fa7 Use pixelate with max tiles 2025-02-07 15:57:06 +09:00
AngelBottomless 6b9d8c9f32 update radius, add pixelate 2025-02-07 15:50:37 +09:00
aria1th 4992749ebb fix random 2025-01-30 21:35:27 +09:00
aria1th 079d9a1e59 prevent partition error 2025-01-30 00:54:53 +09:00
snomiao bd84f2534a chore(publish): update GitHub Actions workflow for node publishing
- Add permissions for issue writing
- Update action version to v1 for publish-node-action
- Add condition to run job only for specific repository owner
2025-01-20 21:27:42 +00:00
aria1th f199fde35c improve file handling 2025-01-16 03:27:42 +09:00
AngelBottomless db84be70ae Enhance error messages 2025-01-09 02:48:51 +09:00
AngelBottomless 17b8937cae skip multiply if dtype is already uint8 2025-01-09 02:45:31 +09:00
AngelBottomless a560f2db7c add rgba handling wrappers 2025-01-09 02:40:29 +09:00
AngelBottomless 269316c01d fix cached state of random-related nodes 2025-01-09 02:11:57 +09:00
AngelBottomless 0e364db321 fix mismatching signatures 2025-01-08 04:47:02 +09:00
AngelBottomless adb5bd9a75 fix missing handlers (ensuring) 2025-01-03 16:18:10 +09:00
AngelBottomless c197f0d66d handle RGBA 2025-01-02 16:58:43 +09:00
aria1th 436196f851 Handle S3 imgs 2024-12-30 22:36:55 +09:00
aria1th a12b2e6657 Bump version
TODO: improve warnings
2024-12-26 14:23:19 +09:00
aria1th 1afe926c21 add ImageFromURL Node (http-limited) 2024-12-26 14:19:05 +09:00
aria1th 382e821e06 TEMP: skip warning 2024-12-26 12:36:12 +09:00
aria1th 6e08078c2a fix is_changed attribute 2024-12-20 02:12:21 +09:00
aria1th 2557c7e3b5 fix is_changed 2024-12-20 01:20:53 +09:00
aria1th eb3a0d0fa4 add to mapping 2024-12-16 19:13:16 +09:00
aria1th 75a3fa7665 Some fixes, add nodes 2024-12-16 19:11:39 +09:00
aria1th a384a92373 add resizing for multiples 2024-12-16 17:40:23 +09:00
aria1th 77e1429843 bump 1.5.2 2024-12-14 10:07:01 +09:00
aria1th b2cc4b0f2b add nsfw-checker and filtering nodes 2024-12-14 10:06:42 +09:00
aria1th 3688c504ed Fix compatibility 2024-12-09 15:36:25 +09:00
AngelBottomless 649de167ec Fix webp save 2024-12-04 14:19:25 +09:00
AngelBottomless d330362006 hotfix 2024-12-04 14:19:06 +09:00
AngelBottomless 2473ef3e84 Update recommended 2024-12-04 14:14:09 +09:00
AngelBottomless b04da7a530 Update io_node.py 2024-12-04 14:13:13 +09:00
AngelBottomless 611253e2ab add SystemRandom 2024-12-04 14:09:22 +09:00
aria1th 52b422128d Create RGBA image from mask+image 2024-12-04 03:45:52 +09:00
aria1th 2e65bc8294 Fix signatures as "Boolean" 2024-12-01 21:15:46 +09:00
aria1th f1f0f4572e Fix and add Base64 handlings 2024-12-01 21:09:42 +09:00
AngelBottomless b3ac53a946 1.3.0-add tagger node 2024-12-01 00:41:36 +09:00
AngelBottomless 656aadfc8b Merge pull request #5 from aria1th/dghs-imgutils
Dghs imgutils taggers node
2024-12-01 00:41:20 +09:00
27 changed files with 4199 additions and 289 deletions
+6 -2
View File
@@ -7,15 +7,19 @@ on:
paths:
- "pyproject.toml"
permissions:
issues: write
jobs:
publish-node:
name: Publish Custom Node to registry
runs-on: ubuntu-latest
if: ${{ github.repository_owner == 'aria1th' }}
steps:
- name: Check out code
uses: actions/checkout@v4
- name: Publish Custom Node
uses: Comfy-Org/publish-node-action@main
uses: Comfy-Org/publish-node-action@v1
with:
## Add your own personal access token to your Github Repository secrets and reference it here.
personal_access_token: ${{ secrets.REGISTRY_ACCESS_TOKEN }}
personal_access_token: ${{ secrets.REGISTRY_ACCESS_TOKEN }}
+16
View File
@@ -0,0 +1,16 @@
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 runs when loaded in ComfyUI. To disable it, set `COMFYUI_LOGICUTILS_SKIP_INSTALL=1`.
- To force CPU/GPU tagger dependency selection, set `COMFYUI_LOGICUTILS_IMGUTILS_VARIANT=cpu` or `COMFYUI_LOGICUTILS_IMGUTILS_VARIANT=gpu`.
+4 -4
View File
@@ -80,18 +80,18 @@ def validate(container):
function_kwargs = list(function_kwargs) + ["self"]
# if args/kwargs are in function kwargs, warn and skip
if "args" in function_kwargs:
print(f"Warning: args in function arguments in {cls.__name__}, skipping argument validation")
#print(f"Warning: args in function arguments in {cls.__name__}, skipping argument validation")
continue
if "kwargs" in function_kwargs:
print(f"Warning: kwargs in function arguments in {cls.__name__}, skipping argument validation")
#print(f"Warning: kwargs in function arguments in {cls.__name__}, skipping argument validation")
continue
# input kwargs are subset of function kwargs
if not set(input_keys).issubset(function_kwargs):
raise Exception(f"INPUT_TYPES and function arguments must match in {cls.__name__}, input_types: {input_keys}, function arguments: {function_kwargs}")
# if not exact match, print warning
if len(set(input_keys)) != len(set(function_kwargs)):
print(f"Warning: INPUT_TYPES and function arguments don't match in {cls.__name__}, input_types: {input_keys}, function arguments: {function_kwargs}")
#print(f"Warning: INPUT_TYPES and function arguments don't match in {cls.__name__}, input_types: {input_keys}, function arguments: {function_kwargs}")
pass
# AllTrue class hijacks the isinstance, issubclass, bool, str, jsonserializable, eq, ne methods to always return True
class AllTrue(str):
def __init__(self, representation=None) -> None:
+152 -3
View File
@@ -1,6 +1,7 @@
from .imgio.converter import PILHandlingHodes
from .autonode import node_wrapper, get_node_names_mappings, validate, anytype, PILImage
from .utils.tagger import get_tags, tagger_keys
from PIL import Image, ImageFilter
auxilary_classes = []
auxilary_node = node_wrapper(auxilary_classes)
@@ -49,6 +50,129 @@ class GetRatingFromTextNode:
}
}
def pixelate(image, pixelation_factor=0.1):
# Downscale the image
small = image.resize(
(int(image.width * pixelation_factor), int(image.height * pixelation_factor)),
resample=Image.NEAREST
)
# Upscale back to original size
return small.resize(image.size, Image.NEAREST)
def pixelate_target_tiles(image, max_tiles=100):
aspect_ratio = image.width / image.height
# solve h * a * h = max_tiles
h = int((max_tiles / aspect_ratio) ** 0.5)
w = int(aspect_ratio * h)
h, w = max(4, h), max(4, w)
# Downscale the image
small = image.resize(
(w, h),
resample=Image.NEAREST,
)
# Upscale back to original size
return small.resize(image.size, Image.NEAREST)
@auxilary_node
class CensorImageByRating:
FUNCTION = "censor_image"
RETURN_TYPES = ("IMAGE",)
CATEGORY = "image"
custom_name = "Censor Image by Rating"
@staticmethod
@PILHandlingHodes.output_wrapper
def censor_image(image, rating_threshold, censor_method, model_name=None):
# Convert input to a PIL image
image = PILHandlingHodes.handle_input(image)
result_dict = get_tags(image, model_name=model_name)
rating = result_dict['rating']
# If rating is general, no censorship required
if rating.lower() == "general":
return (image,)
censor_image = False
if rating_threshold == "general":
# censor if not general
if rating.lower() != "general":
censor_image = True
elif rating_threshold == "sensitive":
# censor if not general or sensitive
if rating.lower() not in ["general", "sensitive"]:
censor_image = True
elif rating_threshold == "questionable":
# censor if not general, sensitive or questionable
if rating.lower() not in ["general", "sensitive", "questionable"]:
censor_image = True
elif rating_threshold == "explicit":
return (image,) # why are you using this?
if censor_image:
if censor_method.lower() == "white":
# Return a white image of the same size
censored_image = Image.new("RGB", image.size, (255, 255, 255))
return (censored_image,)
elif censor_method.lower() == "blur":
# Apply a strong blur (you can adjust radius as needed)
censored_image = image.filter(ImageFilter.GaussianBlur(radius=40))
return (censored_image,)
elif censor_method.lower() == "pixelate":
# first, blur
censored_image = image.filter(ImageFilter.GaussianBlur(radius=20))
censored_image = pixelate_target_tiles(censored_image, max_tiles=100)
return (censored_image,)
# If unknown method is provided, just return the original image
return (image,)
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"image": ("IMAGE",),
"rating_threshold": (["general", "sensitive", "questionable", "explicit"],),
"censor_method": (["blur", "white","pixelate"],),
},
"optional": {
"model_name": (tagger_keys, {"default": tagger_keys[0]}),
}
}
@auxilary_node
class FilterTagsNode:
"""
Filters tags, given a list of tags splitted by ",".
We assume the input text is splittable by "," (or separator). Then, if any tags contain the filter, we remove matching tags.
"""
FUNCTION = "filter_tags"
RETURN_TYPES = ("STRING",)
CATEGORY = "safety"
custom_name = "Filter Tags"
@staticmethod
def filter_tags(tags, filter_tags, separator):
filter_tags = filter_tags.split(",")
filter_tags = [tag.strip() for tag in filter_tags]
tags = tags.split(separator)
tags = [tag.strip() for tag in tags]
filtered = []
for tag in tags:
if all(filter_tag not in tag for filter_tag in filter_tags):
filtered.append(tag)
return (separator.join(filtered), )
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"tags": ("STRING",),
"filter_tags": ("STRING",),
},
# optional separator
"optional": {
"separator": ("STRING", {"default": ","}),
}
}
@auxilary_node
class GetTagsAboveThresholdNode:
FUNCTION = "get_tags_above_threshold"
@@ -96,7 +220,7 @@ class GetTagsAboveThresholdFromTextNode:
"model_name": (tagger_keys, {"default": tagger_keys[0]}),
}
}
@auxilary_node
class GetCharactersAboveThresholdNode:
FUNCTION = "get_tags_above_threshold"
@@ -172,6 +296,31 @@ class GetAllTagsAboveThresholdNode:
"model_name": (tagger_keys, {"default": tagger_keys[0]}),
}
}
@auxilary_node
class GetAllTagsExceptCharacterAboveThresholdNode:
FUNCTION = "get_tags"
RETURN_TYPES = ("STRING",)
CATEGORY = "tagger"
custom_name = "Get All Tags Above Threshold Except Characters"
@staticmethod
def get_tags(image, threshold, replace, model_name):
image = PILHandlingHodes.handle_input(image)
result = get_tags(image, threshold=threshold, replace=replace, model_name=model_name)
result_list = []
result_list.append(result['rating'])
result_list.extend(result['tags'])
return (", ".join(result_list), )
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"image": ("IMAGE",),
},
"optional": {
"threshold": ("FLOAT", {"default": 0.4}),
"replace": ("BOOLEAN", {"default": False}),
"model_name": (tagger_keys, {"default": tagger_keys[0]}),
}
}
CLASS_MAPPINGS, CLASS_NAMES = get_node_names_mappings(auxilary_classes)
validate(auxilary_classes)
validate(auxilary_classes)
+9 -5
View File
@@ -9,8 +9,12 @@ node = node_wrapper(classes)
conversion_operators = {
"Int" : int,
"Float" : float,
"Bool" : bool,
"String" : str
"Boolean" : bool,
"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)
+182
View File
@@ -0,0 +1,182 @@
import os, io
import numpy as np
from .imgio.converter import PILHandlingHodes
from .autonode import node_wrapper, get_node_names_mappings, validate
from PIL import Image
try:
from Crypto.PublicKey import RSA
from Crypto.Cipher import AES, PKCS1_OAEP
from Crypto.Random import get_random_bytes
except ImportError:
print("Crypto library not found. Please install pycryptodome.")
raise
import torch
from base64 import b64encode, b64decode
# List of classes to register
secure_classes = []
secure_node = node_wrapper(secure_classes)
@secure_node
class SecureBase64Encrypt:
"""
Encrypt an image as a base64 string using RSA public key + AES.
- images: Only the first image is used.
- public_key_pem: RSA public key (PEM string).
Outputs: 'encrypted_base64' string that SecureWebPDecrypt can decrypt.
"""
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"images": ("IMAGE",),
"public_key_pem": ("STRING", {"multiline": True, "default": ""}),
}
}
RETURN_TYPES = ("STRING",)
RETURN_NAMES = ("encrypted_base64",)
FUNCTION = "encrypted_base64"
CATEGORY = "image"
custom_name = "Secure Base64 Encrypt"
OUTPUT_NODE = True
RESULT_NODE = True
def encrypted_base64(self, images, public_key_pem):
# Check input
if images is None or not len(images) or images[0] is None:
raise ValueError("No image provided.")
# Load RSA public key
rsa_key = RSA.import_key(public_key_pem)
cipher_rsa = PKCS1_OAEP.new(rsa_key)
# Take first image (if there's a batch dimension, pick index [0])
img_tensor = images[0].clone().detach()
if img_tensor.ndim == 4:
img_tensor = img_tensor[0]
# Convert [0..1] float => [0..255] uint8
img_array = (255.0 * img_tensor.clamp(0, 1).cpu().numpy()).astype("uint8")
# If shape is (C,H,W), transpose to (H,W,C). If shape is (H,W,C), leave as is.
if img_array.ndim == 3:
# If the first dimension is small (1,3,4), interpret as channels-first
if img_array.shape[0] in [1,3,4] and img_array.shape[-1] not in [1,3,4]:
img_array = np.transpose(img_array, (1,2,0))
# The critical fix: ensure array is contiguous
img_array = np.ascontiguousarray(img_array)
# Create a PIL Image from array
pil_img = Image.fromarray(img_array, mode="RGB")
# Save in-memory as lossless WebP
buffer = io.BytesIO()
pil_img.save(buffer, format="WEBP", lossless=True)
image_bytes = buffer.getvalue()
# Generate a random AES session key, encrypt it with RSA
session_key = get_random_bytes(16)
enc_session_key = cipher_rsa.encrypt(session_key)
# Encrypt the image bytes with AES-EAX
cipher_aes = AES.new(session_key, AES.MODE_EAX)
ciphertext, tag = cipher_aes.encrypt_and_digest(image_bytes)
# Build custom envelope
encrypted_blob = (
b"ENCWEBP" +
len(enc_session_key).to_bytes(2, "big") +
enc_session_key +
bytes([len(cipher_aes.nonce)]) + cipher_aes.nonce +
bytes([len(tag)]) + tag +
ciphertext
)
# Base64-encode the blob
encrypted_base64 = b64encode(encrypted_blob).decode("utf-8")
return (encrypted_base64,)
@secure_node
class SecureWebPDecrypt:
"""
Decrypt an encrypted WebP image (or list of them) produced by SecureBase64Encrypt.
Returns a single IMAGE (first one).
"""
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"encrypted_base64": ("STRING", {"multiline": True, "default": ""}),
"private_key_pem": ("STRING", {"multiline": True, "default": ""}),
}
}
RETURN_TYPES = ("IMAGE",)
RETURN_NAMES = ("Decrypted_Image",)
FUNCTION = "decrypt_image"
CATEGORY = "image"
custom_name = "Secure WebP Decrypt"
@PILHandlingHodes.output_wrapper
def decrypt_image(self, encrypted_base64, private_key_pem):
# Convert to list if single string
if encrypted_base64 is None:
encrypted_base64 = []
if isinstance(encrypted_base64, str):
encrypted_base64 = [encrypted_base64]
elif not isinstance(encrypted_base64, (list, tuple)):
raise ValueError("encrypted_base64 must be string or list/tuple.")
# Import RSA private key
if isinstance(private_key_pem, str):
private_key_pem = private_key_pem.encode("utf-8")
rsa_key = RSA.import_key(private_key_pem)
cipher_rsa = PKCS1_OAEP.new(rsa_key)
for b64_item in encrypted_base64:
data = b64decode(b64_item)
if data[:7] != b"ENCWEBP":
raise ValueError("Invalid encrypted WebP data (missing header).")
idx = 7
enc_key_len = int.from_bytes(data[idx : idx + 2], "big")
idx += 2
enc_session_key = data[idx : idx + enc_key_len]
idx += enc_key_len
nonce_len = data[idx]
idx += 1
nonce = data[idx : idx + nonce_len]
idx += nonce_len
tag_len = data[idx]
idx += 1
tag = data[idx : idx + tag_len]
idx += tag_len
ciphertext = data[idx:]
# RSA-decrypt the AES session key
session_key = cipher_rsa.decrypt(enc_session_key)
# AES-EAX decrypt
cipher_aes = AES.new(session_key, AES.MODE_EAX, nonce=nonce)
plaintext = cipher_aes.decrypt(ciphertext)
try:
cipher_aes.verify(tag)
except ValueError:
raise ValueError("Decryption failed: data tampered or wrong key.")
# plaintext -> PIL
pil_img = Image.open(io.BytesIO(plaintext)).convert("RGB")
return (pil_img, )
# Register node classes with ComfyUI
CLASS_MAPPINGS, CLASS_NAMES = get_node_names_mappings(secure_classes)
validate(secure_classes)
+2 -2
View File
@@ -30,7 +30,7 @@ class SDWebuiAPINode:
"height": ("INT", {"default": 1024}),
"hr_scale": ("FLOAT", {"default": 1.5}),
"hr_upscale": ("STRING", {"default": "Latent"}),
"enable_hr": ("BOOL", {"default": False}),
"enable_hr": ("BOOLEAN", {"default": False}),
"cfg_scale": ("INT", {"default": 7}),
}
}
@@ -59,7 +59,7 @@ class SDWebuiAPIFallbackNode:
"height": ("INT", {"default": 1024}),
"hr_scale": ("FLOAT", {"default": 1.5}),
"hr_upscale": ("STRING", {"default": "Latent"}),
"enable_hr": ("BOOL", {"default": False}),
"enable_hr": ("BOOLEAN", {"default": False}),
"cfg_scale": ("INT", {"default": 7}),
}
}
+327 -58
View File
@@ -1,3 +1,4 @@
from typing import Union, List
from PIL import Image
import numpy as np
import base64
@@ -5,11 +6,144 @@ import torch
import requests
import os
from io import BytesIO
import gzip
import re
from urllib.parse import urlparse
def handle_rgba_composite(
image: Image.Image, background_color=(255, 255, 255), as_rgba=False
) -> Image.Image:
"""
Convert RGBA image to RGB image using alpha_composite.
"""
mode = image.mode
if as_rgba:
return image.convert("RGBA") # universal format
if mode == "RGB":
return image
if mode == "RGBA":
# Create a white RGBA background
background = Image.new("RGBA", image.size, (255, 255, 255, 255))
# Composite the original image over the white background
composed = Image.alpha_composite(background, image)
# Convert back to RGB (now that background is flattened)
return composed.convert("RGB")
elif mode == "LA":
# "LA" is 8-bit grayscale + alpha.
rgba_image = image.convert("RGBA")
background = Image.new("RGBA", rgba_image.size, (*background_color, 255))
composed = Image.alpha_composite(background, rgba_image)
return composed.convert("RGB")
# 3. "L" or "1" = Grayscale or Black/White, "P" = Palette
elif mode in ["L", "1", "P"]:
# Simply converting to "RGB" is usually enough.
return image.convert("RGB")
# 4. "CMYK", "YCbCr", "HSV", etc.
elif mode in ["CMYK", "YCbCr", "HSV"]:
# Typically, a .convert("RGB") is enough if you just need an RGB version.
return image.convert("RGB")
print(f"Warning: Unhandled image mode: {mode}. Converting to RGB.")
return image.convert("RGB")
def fetch_image_securely(image_url: str,
allowed_schemes=('http', 'https'),
max_file_size=5_000_000,
request_timeout=30):
"""
Fetches an image from the given URL securely.
This function:
1. Validates the URL scheme (only http/https).
2. Blocks private IP/loopback addresses to prevent SSRF attacks.
3. Streams data to avoid excessive memory usage.
4. Checks MIME type, size limits, and optionally handles form-encoded image data.
:param image_url: URL of the image to retrieve (e.g., an S3-signed URL).
:param allowed_schemes: A tuple of allowed URL schemes (default: ('http', 'https')).
:param max_file_size: Max size (in bytes) of the file to download.
:param request_timeout: Timeout (in seconds) for the request.
:return: PIL Image object if successful, else raises an exception.
"""
# -- 1. Validate scheme to avoid unexpected protocols --
parsed = urlparse(image_url)
if parsed.scheme not in allowed_schemes:
raise ValueError(f"Invalid or disallowed URL scheme: {parsed.scheme}")
# -- 2. Prevent local network (SSRF) attacks by blocking private or loopback addresses --
# This is a simplified check. Consider using a library for robust IP parsing if needed.
ip_like_pattern = r'^(\d{1,3}\.){3}\d{1,3}$'
hostname = parsed.hostname
if (
hostname is None
or hostname.lower() in ("localhost", "127.0.0.1", "::1")
or (re.match(ip_like_pattern, hostname) and hostname.startswith("10."))
or hostname.startswith("192.168.")
or hostname.startswith("172.16.")
or hostname.startswith("172.17.")
or hostname.startswith("172.18.")
or hostname.startswith("172.19.")
or hostname.startswith("172.2") # covers 172.20 - 172.31
or hostname.startswith("172.3")
):
raise ValueError("URL resolves to a private or loopback address, which is disallowed.")
# -- 3. Retrieve the response with a timeout and stream --
# This handles the S3 URL just like any other public HTTPS link.
with requests.get(image_url, timeout=request_timeout, stream=True) as response:
response.raise_for_status()
# -- 4. Check Content-Type in headers --
content_type = response.headers.get('Content-Type', '').lower()
# If it's a direct image...
if content_type.startswith("image/"):
# -- 5. Check Content-Length against max_file_size --
content_length = response.headers.get('Content-Length')
if content_length and int(content_length) > max_file_size:
raise ValueError(
f"File is too large: {int(content_length)} bytes. "
f"Max allowed is {max_file_size} bytes."
)
data = BytesIO()
downloaded = 0
chunk_size = 8192
for chunk in response.iter_content(chunk_size=chunk_size):
downloaded += len(chunk)
if downloaded > max_file_size:
raise ValueError(
f"File exceeded the maximum allowed size of {max_file_size} bytes."
)
data.write(chunk)
# Reset the buffer and open with PIL
data.seek(0)
return Image.open(data)
# If the server reports x-www-form-urlencoded, parse for embedded image data
elif content_type == "application/x-www-form-urlencoded":
# let PIL handle the parsing
try:
return Image.open(BytesIO(response.content))
except Exception as e:
raise ValueError(
f"Failed to parse x-www-form-urlencoded data as image: {e}"
)
else:
# Some other content type we don't handle
raise ValueError(
f"Unsupported Content-Type or not an image: {content_type}"
)
class IOConverter:
"""
Classify the input data type.
Assumes the inputs to be following:
- PIL Image
@@ -17,9 +151,10 @@ class IOConverter:
- torch tensor
- string (path to image)
- base64 string (which can be decoded to bytes and then to image)
- gzip-compressed base64 string
- URL (which can be downloaded to image)
do NOT pass unsafe URLs / base64 strings, as it may cause security issues.
Do NOT pass unsafe URLs / base64 strings, as it may cause security issues.
"""
class InputType:
@@ -28,8 +163,9 @@ class IOConverter:
TORCH = "TORCH"
STRING = "STRING"
BASE64 = "BASE64"
GZIP_BASE64 = "GZIP_BASE64"
URL = "URL"
def __init__(self):
raise Exception("This class should not be instantiated.")
@@ -49,102 +185,209 @@ class IOConverter:
elif input_data.startswith("http://") or input_data.startswith("https://"):
return IOConverter.InputType.URL
else:
raise Exception(f"Invalid string input, {input_data:10}")
# Attempt to detect base64-encoded data
try:
decoded_data = base64.b64decode(input_data, validate=True)
# Check for gzip magic number
if decoded_data[:2] == b'\x1f\x8b':
return IOConverter.InputType.GZIP_BASE64
else:
return IOConverter.InputType.BASE64
except Exception:
raise Exception(f"Invalid string input, cannot be decoded as base64.")
else:
raise Exception(f"Invalid input type, {type(input_data)}")
@staticmethod
def match_dtype(array_or_tensor, is_tensor=False):
# if all value is between 0 and 1, multiply by 255 and convert to uint8
# however already uint8, skip
# check dtype first
if array_or_tensor.dtype == np.uint8 or array_or_tensor.dtype == torch.uint8:
return array_or_tensor
if array_or_tensor.min() >= 0 and array_or_tensor.max() <= 1:
multiplied = array_or_tensor * 255
if not is_tensor:
return multiplied.astype(np.uint8)
else:
return multiplied.to(torch.uint8)
return array_or_tensor
@staticmethod
def convert_to_pil(input_data):
input_type = IOConverter.classify(input_data)
if input_type == IOConverter.InputType.PIL:
return input_data
return handle_rgba_composite(input_data)
elif input_type == IOConverter.InputType.NUMPY:
# [1, 1216, 832, 3], '<f4'] -> [1216, 832, 3], 'uint8'
# if not first element is 1, then it is a batch of images so warning
if input_data.shape[0] != 1:
print("Warning: Batch of images detected, taking first image")
input_data = input_data[0] * 255.0 if input_data.dtype == np.float32 else input_data[0]
input_data = input_data.astype('uint8')
return Image.fromarray(input_data)
result = []
for i in range(input_data.shape[0]):
np_array = IOConverter.match_dtype(input_data[i])
result.append(handle_rgba_composite(Image.fromarray(np_array)))
return result # return list of PIL images
input_data = IOConverter.match_dtype(input_data[0])
return handle_rgba_composite(Image.fromarray(input_data))
elif input_type == IOConverter.InputType.TORCH:
# same as above
if input_data.shape[0] != 1:
print("Warning: Batch of images detected, taking first image")
input_data = input_data[0].cpu().numpy() * 255.0
input_data = input_data.astype('uint8')
return Image.fromarray(input_data)
result = []
for i in range(input_data.shape[0]):
np_array = (
IOConverter.match_dtype(input_data[i], is_tensor=True)
.cpu()
.numpy()
)
result.append(handle_rgba_composite(Image.fromarray(np_array)))
return result
input_data = IOConverter.match_dtype(input_data[0], is_tensor=True)
np_array = input_data.cpu().numpy()
return handle_rgba_composite(Image.fromarray(np_array))
elif input_type == IOConverter.InputType.STRING:
return Image.open(input_data)
elif input_type == IOConverter.InputType.GZIP_BASE64:
decoded_data = IOConverter.read_base64(input_data)
decompressed_data = gzip.decompress(decoded_data)
partial_result = Image.open(BytesIO(decompressed_data))
result = handle_rgba_composite(partial_result)
return result
elif input_type == IOConverter.InputType.BASE64:
return Image.open(BytesIO(base64.b64decode(input_data)))
decoded_data = IOConverter.read_base64(input_data)
partial_result = Image.open(BytesIO(decoded_data))
result = handle_rgba_composite(partial_result)
return result
elif input_type == IOConverter.InputType.URL:
response = requests.get(input_data)
return Image.open(BytesIO(response.content))
partial_result = fetch_image_securely(input_data)
result = handle_rgba_composite(partial_result)
return result
else:
raise Exception(f"Invalid input type, {input_type}")
@staticmethod
def to_tensor(pil_image):
def to_rgb_tensor(pil_image):
if pil_image.mode == "I":
pil_image = pil_image.point(lambda i: i * (1/255)) # convert to float
pil_image = pil_image.point(lambda i: i * (1/255)) # convert to float
pil_image = handle_rgba_composite(pil_image)
np_array = np.array(pil_image).astype(np.float32) / 255.0
image = torch.from_numpy(np_array)[None,][0]
image = image[None,] # to batch
return image
tensor = torch.from_numpy(np_array)
tensor = tensor.unsqueeze(0) # Add batch dimension
# assert 4-dimensional tensor, B,C,H,W
if len(tensor.shape) != 4:
raise Exception(f"Invalid tensor shape, expected 4-dimensional tensor, got {tensor.shape}")
return tensor
@staticmethod
def convert_to_tensor(input_data):
def to_rgba_tensor(pil_image):
if pil_image.mode == "I":
pil_image = pil_image.point(lambda i: i * (1/255)) # convert to float
pil_image = handle_rgba_composite(pil_image, as_rgba=True)
np_array = np.array(pil_image).astype(np.float32) / 255.0
tensor = torch.from_numpy(np_array)
tensor = tensor.unsqueeze(0) # Add batch dimension
# assert 4-dimensional tensor, B,C,H,W
if len(tensor.shape) != 4:
raise Exception(f"Invalid tensor shape, expected 4-dimensional tensor, got {tensor.shape}")
return tensor
@staticmethod
def read_base64(base64_string: str) -> bytes:
return base64.b64decode(base64_string)
@staticmethod
def read_maybe_gzip_base64(base64_string: str) -> bytes:
decoded_data = base64.b64decode(base64_string)
if decoded_data[:2] == b'\x1f\x8b':
result = gzip.decompress(decoded_data)
else:
result = decoded_data
# to string
return result.decode('utf-8')
@staticmethod
def convert_to_rgb_tensor(input_data, rgba=False):
if not rgba:
output_func = IOConverter.to_rgb_tensor
else:
output_func = IOConverter.to_rgba_tensor
input_type = IOConverter.classify(input_data)
if input_type == IOConverter.InputType.PIL:
return IOConverter.to_tensor(input_data)
return output_func(input_data)
elif input_type == IOConverter.InputType.NUMPY:
return torch.from_numpy(input_data)
# if all values are 0~1, skip
if input_data.min() >= 0 and input_data.max() <= 1:
np_array = input_data.astype(np.float32)
else:
np_array = input_data.astype(np.float32) / 255.0
tensor = torch.from_numpy(np_array)
tensor = tensor.unsqueeze(0) # Add batch dimension
return tensor
elif input_type == IOConverter.InputType.TORCH:
return input_data
elif input_type == IOConverter.InputType.STRING:
image = Image.open(input_data)
return IOConverter.to_tensor(image)
return output_func(image)
elif input_type == IOConverter.InputType.GZIP_BASE64:
image = IOConverter.convert_to_pil(input_data)
return output_func(image)
elif input_type == IOConverter.InputType.BASE64:
image = Image.open(BytesIO(base64.b64decode(input_data)))
return IOConverter.to_tensor(image)
image = IOConverter.convert_to_pil(input_data)
return output_func(image)
elif input_type == IOConverter.InputType.URL:
response = requests.get(input_data)
response.raise_for_status()
image = Image.open(BytesIO(response.content))
return IOConverter.to_tensor(image)
else:
raise Exception(f"Invalid input type, {input_type}")
@staticmethod
def convert_to_base64(input_data):
input_type = IOConverter.classify(input_data)
if input_type == IOConverter.InputType.PIL:
buffered = BytesIO()
input_data.save(buffered, format="PNG")
return base64.b64encode(buffered.getvalue()).decode("utf-8")
elif input_type == IOConverter.InputType.NUMPY:
return IOConverter.convert_to_base64(Image.fromarray(input_data))
elif input_type == IOConverter.InputType.TORCH:
return IOConverter.convert_to_base64(Image.fromarray(input_data.cpu().numpy()))
elif input_type == IOConverter.InputType.STRING:
return IOConverter.convert_to_base64(Image.open(input_data))
elif input_type == IOConverter.InputType.BASE64:
return input_data
elif input_type == IOConverter.InputType.URL:
response = requests.get(input_data)
return base64.b64encode(response.content).decode("utf-8")
image = fetch_image_securely(input_data)
return output_func(image)
else:
raise Exception(f"Invalid input type, {input_type}")
@staticmethod
def convert_to_base64(input_data, format="PNG", quality=100, gzip_compress=False):
pil_image = IOConverter.convert_to_pil(input_data)
buffered = BytesIO()
save_params = {'format': format}
if format.upper() in ['JPEG', 'JPG']:
save_params['quality'] = quality
pil_image.save(buffered, **save_params)
buffered.seek(0)
if gzip_compress:
compressed_buffer = BytesIO()
with gzip.GzipFile(fileobj=compressed_buffer, mode='wb') as f:
f.write(buffered.getvalue())
compressed_buffer.seek(0)
base64_data = base64.b64encode(compressed_buffer.getvalue()).decode('utf-8')
else:
base64_data = base64.b64encode(buffered.getvalue()).decode('utf-8')
return base64_data
@staticmethod
def string_to_base64(input_string, gzip_compress=False):
if gzip_compress:
compressed_buffer = BytesIO()
with gzip.GzipFile(fileobj=compressed_buffer, mode='wb') as f:
f.write(input_string.encode())
compressed_buffer.seek(0)
base64_data = base64.b64encode(compressed_buffer.getvalue()).decode('utf-8')
else:
base64_data = base64.b64encode(input_string.encode()).decode('utf-8')
return base64_data
class PILHandlingHodes:
@staticmethod
def handle_input(tensor_or_image):
def handle_input(tensor_or_image) -> Union[Image.Image, List[Image.Image]]:
pil_image = IOConverter.convert_to_pil(tensor_or_image)
return pil_image
@staticmethod
def handle_output_as_pil(pil_image):
def handle_output_as_pil(pil_image: Image.Image) -> Image.Image:
return pil_image
@staticmethod
def handle_output_as_tensor(pil_image):
return IOConverter.convert_to_tensor(pil_image)
def handle_output_as_tensor(pil_image: Image.Image, rgba=False) -> torch.Tensor:
return IOConverter.convert_to_rgb_tensor(pil_image, rgba=rgba)
@staticmethod
def handle_output_as_rgba_tensor(pil_image: Image.Image) -> torch.Tensor:
return IOConverter.convert_to_rgb_tensor(pil_image, rgba=True)
@staticmethod
def output_wrapper(func):
def wrapped(*args, **kwargs):
@@ -158,3 +401,29 @@ class PILHandlingHodes:
return tuple(tuples_collect)
return wrapped
@staticmethod
def rgba_output_wrapper(func):
def wrapped(*args, **kwargs):
outputs = func(*args, **kwargs)
tuples_collect = []
for output in outputs:
if isinstance(output, (Image.Image, torch.Tensor)):
tuples_collect.append(PILHandlingHodes.handle_output_as_rgba_tensor(output))
else:
tuples_collect.append(output)
return tuple(tuples_collect)
return wrapped
@staticmethod
def to_base64(anything, quality=100, format="PNG", gzip_compress=False):
base64_data = IOConverter.convert_to_base64(anything, format=format, quality=quality, gzip_compress=gzip_compress)
return base64_data
@staticmethod
def string_to_base64(input_string, gzip_compress=False):
base64_data = IOConverter.string_to_base64(input_string, gzip_compress=gzip_compress)
return base64_data
@staticmethod
def maybe_gzip_base64_to_string(base64_string):
return IOConverter.read_maybe_gzip_base64(base64_string)
+52 -10
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
@@ -31,30 +32,71 @@ def process_wrap(cmd_str, cwd=None, handler=None):
return process.wait()
if "python_embeded" in sys.executable or "python_embedded" in sys.executable: #standalone python version
pip_install = [sys.executable, '-s', '-m', 'pip', 'install', "-U"]
pip_install = [sys.executable, '-s', '-m', 'pip', 'install', "-U", "--default-timeout=1000"]
else:
pip_install = [sys.executable, '-m', 'pip', 'install', "-U"]
pip_install = [sys.executable, '-m', 'pip', 'install', "-U", "--default-timeout=1000"]
def _imgutils_install_candidates() -> list[str]:
"""
Decide which dghs-imgutils package variant to install first.
- `dghs-imgutils[gpu]` includes optional GPU-related dependencies.
- If GPU detection fails, default to CPU variant first.
"""
variant = os.environ.get("COMFYUI_LOGICUTILS_IMGUTILS_VARIANT", "").strip().lower()
if variant in {"gpu", "cuda"}:
return ["dghs-imgutils[gpu]", "dghs-imgutils"]
if variant in {"cpu"}:
return ["dghs-imgutils", "dghs-imgutils[gpu]"]
# auto
try:
import torch
if getattr(torch, "cuda", None) is not None and torch.cuda.is_available():
return ["dghs-imgutils[gpu]", "dghs-imgutils"]
except Exception:
pass
return ["dghs-imgutils", "dghs-imgutils[gpu]"]
def initialization():
try:
import piexif
except ImportError:
except Exception:
print("piexif not found, installing...")
run_installation("piexif")
try:
import chardet
except ImportError:
except Exception:
print("chardet not found, installing...")
run_installation("chardet")
try:
from imgutils.tagging import get_wd14_tags
except ImportError:
run_installation("dghs-imgutils[gpu]")
except Exception:
print("imgutils not found, installing...")
for candidate in _imgutils_install_candidates():
print(f"Trying to install {candidate}...")
run_installation(candidate)
try:
from imgutils.tagging import get_wd14_tags # noqa: F401
break
except Exception:
continue
try:
from Crypto.PublicKey import RSA
except Exception:
print("pycryptodome not found, installing...")
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()
+1360 -107
View File
File diff suppressed because it is too large Load Diff
+31 -30
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 = ("BOOL",)
@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:
"""
@@ -171,7 +154,7 @@ class LogicGateCompare:
"""
Returns 1 if input1 > input2, 0 otherwise
"""
RETURN_TYPES = ("BOOL",)
RETURN_TYPES = ("BOOLEAN",)
@classmethod
def INPUT_TYPES(s):
return {
@@ -190,7 +173,7 @@ class LogicGateCompareString:
"""
Returns if given regex (1) is found in given string (2)
"""
RETURN_TYPES = ("BOOL",)
RETURN_TYPES = ("BOOLEAN",)
@classmethod
def INPUT_TYPES(s):
return {
@@ -204,6 +187,26 @@ class LogicGateCompareString:
custom_name = "AContainsB(String)"
def compareString(self, regex, input2):
return (True if re.search(regex, input2) else False,)
@node
class GetLengthString:
"""
Returns the length of the input string
"""
RETURN_TYPES = ("INT",)
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"string": ("STRING", {"default": ""}),
}
}
FUNCTION = "lengthString"
CATEGORY = "Logic Gates"
custom_name = "Length of String"
def lengthString(self, string):
return (len(string),)
@node
class StaticNumberInt:
"""
@@ -263,7 +266,7 @@ class LogicGateAnd:
"""
Returns 1 if all inputs are True, 0 otherwise
"""
RETURN_TYPES = ("BOOL",)
RETURN_TYPES = ("BOOLEAN",)
@classmethod
def INPUT_TYPES(s):
return {
@@ -282,7 +285,7 @@ class LogicGateOr:
"""
Returns 1 if any input is True, 0 otherwise
"""
RETURN_TYPES = ("BOOL",)
RETURN_TYPES = ("BOOLEAN",)
@classmethod
def INPUT_TYPES(s):
return {
@@ -295,7 +298,7 @@ class LogicGateOr:
CATEGORY = "Logic Gates"
custom_name = "AOrBGate"
def or_(self, input1, input2):
return (True if input1 and input2 else False,)
return (True if input1 or input2 else False,)
@node
class LogicGateEither:
"""
@@ -365,9 +368,9 @@ class ReplaceString:
def INPUT_TYPES(s):
return {
"required": {
"String": ("STRING", {"default": ""}),
"Regex": ("STRING", {"default": ""}),
"ReplaceWith": ("STRING", {"default": ""}),
"String": ("STRING", {"default": ""}), # input string
"Regex": ("STRING", {"default": ""}), # regex to search for
"ReplaceWith": ("STRING", {"default": ""}), # string to replace with
}
}
FUNCTION = "replace"
@@ -402,7 +405,5 @@ class MemoryNode:
self.memory_value = input1
return (self.memory_value,)
CLASS_MAPPINGS, CLASS_NAMES = get_node_names_mappings(classes)
validate(classes)
+154 -3
View File
@@ -1,3 +1,4 @@
import random
from .autonode import node_wrapper, get_node_names_mappings, validate, anytype
classes = []
node = node_wrapper(classes)
@@ -137,12 +138,162 @@ 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),)
@node
class SigmoidNode:
"""
Returns the sigmoid of a number
"""
RETURN_TYPES = ("FLOAT",)
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"input1": ("FLOAT",),
}
}
FUNCTION = "sigmoid"
CATEGORY = "Math"
custom_name = "Sigmoid"
def sigmoid(self, input1):
return (1 / (1 + math.exp(-input1)),)
def is_prime_small(n: int) -> bool:
"""
Deterministic check for primality for smaller n.
Skips multiples of 2 and 3, then checks i, i+2, i+4 up to sqrt(n).
"""
if n < 2:
return False
if n in (2, 3):
return True
if n % 2 == 0 or n % 3 == 0:
return n == 2 or n == 3
# 6k ± 1 optimization
limit = int(math.isqrt(n)) # integer sqrt
i = 5
while i <= limit:
if n % i == 0 or n % (i + 2) == 0:
return False
i += 6
return True
def miller_rabin_test(d: int, n: int) -> bool:
""" One round of the Miller-Rabin test with a random base 'a'. """
a = random.randrange(2, n - 1)
x = pow(a, d, n) # a^d % n
if x == 1 or x == n - 1:
return True
# Keep squaring x while d does not reach n-1
while d != n - 1:
x = (x * x) % n
d <<= 1 # d *= 2
if x == 1:
return False
if x == n - 1:
return True
return False
def is_prime_miller_rabin(n: int, k: int = 5) -> bool:
"""
Miller-Rabin primality test with k rounds (probabilistic).
Good enough for big integers in practice.
"""
# Handle small or trivial cases
if n < 2:
return False
# check small primes quickly
for small_prime in [2, 3, 5, 7, 11, 13, 17, 19, 23, 29]:
if n == small_prime:
return True
if n % small_prime == 0 and n != small_prime:
return False
# Write n - 1 as d * 2^r
d = n - 1
while d % 2 == 0:
d //= 2
# Witness loop
for _ in range(k):
if not miller_rabin_test(d, n):
return False
return True
@node
class IsPrimeNode:
"""
Checks if an integer is prime.
- If the integer |value| < threshold, uses a deterministic small-check (trial division).
- Otherwise, uses a Miller-Rabin pseudoprime test for a faster check (probabilistic).
Returns a BOOLEAN (True if prime, False if composite).
"""
FUNCTION = "is_prime"
RETURN_TYPES = ("BOOLEAN",)
CATEGORY = "Math"
custom_name = "Is Prime?"
@staticmethod
def is_prime(value: int, threshold: int = 10_000_000, miller_rabin_rounds: int = 5):
# handle negative or zero
if value < 2:
return (False,)
if value < threshold:
# use small prime check
return (is_prime_small(value),)
else:
# use Miller-Rabin
return (is_prime_miller_rabin(value, k=miller_rabin_rounds),)
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"value": ("INT", {"default": 1, "min": -9999999999, "max": 9999999999, "step": 1}),
},
"optional": {
"threshold": ("INT", {"default": 10_000_000, "min": 1, "max": 9999999999, "step": 1}),
"miller_rabin_rounds": ("INT", {"default": 5, "min": 1, "max": 50, "step": 1}),
}
}
@node
class RAMPNode:
"""
Returns the ramp of a number
"""
RETURN_TYPES = ("FLOAT",)
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"input1": ("FLOAT",),
}
}
FUNCTION = "ramp"
CATEGORY = "Math"
custom_name = "RAMP"
def ramp(self, input1):
return (max(0, input1),)
@node
class ModuloNode:
"""
Returns the modulo of a number
+75 -33
View File
@@ -1,33 +1,75 @@
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
NODE_CLASS_MAPPINGS = {
}
NODE_CLASS_MAPPINGS.update(IOMapping)
NODE_CLASS_MAPPINGS.update(LogicMapping)
NODE_CLASS_MAPPINGS.update(RandomMapping)
NODE_CLASS_MAPPINGS.update(ConversionMapping)
NODE_CLASS_MAPPINGS.update(MathMapping)
NODE_CLASS_MAPPINGS.update(ExternalMapping)
NODE_CLASS_MAPPINGS.update(AuxilaryMapping)
NODE_DISPLAY_NAME_MAPPINGS = {
}
NODE_DISPLAY_NAME_MAPPINGS.update(IONames)
NODE_DISPLAY_NAME_MAPPINGS.update(LogicNames)
NODE_DISPLAY_NAME_MAPPINGS.update(RandomNames)
NODE_DISPLAY_NAME_MAPPINGS.update(ConversionNames)
NODE_DISPLAY_NAME_MAPPINGS.update(MathNames)
NODE_DISPLAY_NAME_MAPPINGS.update(ExternalNames)
NODE_DISPLAY_NAME_MAPPINGS.update(AuxilaryNames)
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()
else:
print("Skipping ComfyUI-LogicUtils installation.")
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
NODE_CLASS_MAPPINGS = {
}
NODE_CLASS_MAPPINGS.update(IOMapping)
NODE_CLASS_MAPPINGS.update(LogicMapping)
NODE_CLASS_MAPPINGS.update(RandomMapping)
NODE_CLASS_MAPPINGS.update(ConversionMapping)
NODE_CLASS_MAPPINGS.update(MathMapping)
NODE_CLASS_MAPPINGS.update(ExternalMapping)
NODE_CLASS_MAPPINGS.update(AuxilaryMapping)
NODE_DISPLAY_NAME_MAPPINGS = {
}
NODE_DISPLAY_NAME_MAPPINGS.update(IONames)
NODE_DISPLAY_NAME_MAPPINGS.update(LogicNames)
NODE_DISPLAY_NAME_MAPPINGS.update(RandomNames)
NODE_DISPLAY_NAME_MAPPINGS.update(ConversionNames)
NODE_DISPLAY_NAME_MAPPINGS.update(MathNames)
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 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
+2 -1
View File
@@ -1,7 +1,7 @@
[project]
name = "comfyui-logicutils"
description = "Logical Utils (compare, string, boolean operations) for ComfyUI"
version = "1.2.0"
version = "1.7.3"
license = "MIT"
[project.urls]
@@ -12,3 +12,4 @@ Repository = "https://github.com/aria1th/ComfyUI-LogicUtils"
PublisherId = "angelbottomless"
DisplayName = "ComfyUI-LogicUtils"
Icon = ""
+1031
View File
File diff suppressed because it is too large Load Diff
+419 -27
View File
@@ -1,11 +1,172 @@
import math
import random
import uuid
import time
from .autonode import node_wrapper, get_node_names_mappings, validate
classes = []
node = node_wrapper(classes)
class RandomGuaranteedClass:
OUTPUT_NODE = True
RESULT_NODE = True
@classmethod
def IS_CHANGED(s, *args, **kwargs):
return float("NaN")
@node
class UniformRandomFloat:
class SystemRandomFloat(RandomGuaranteedClass):
"""
Random number generator using system randomness
"""
def __init__(self):
pass
@staticmethod
def generate(min_val=0.0, max_val=1.0, precision=0):
instance = random.SystemRandom(time.time())
value = instance.uniform(min_val, max_val)
if precision > 0:
value = round(value, precision)
return (value,)
RETURN_TYPES = ("FLOAT",)
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"min_val": ("FLOAT", { "default": 0.0, "min": -999999999, "max": 999999999.0, "step": 0.01, "display": "number" }),
"max_val": ("FLOAT", { "default": 1.0, "min": -999999999, "max": 999999999.0, "step": 0.01, "display": "number" }),
"precision": ("INT", { "default": 0, "min": 0, "max": 10, "step": 1, "display": "number" }),
},
}
FUNCTION = "generate"
CATEGORY = "Logic Gates"
custom_name = "System Random Float"
@node
class DimensionSelectorWithSeedNode:
"""
Finds (width, height) such that width*height is near (resolution^2),
ratio = width/height is in [min_ratio, max_ratio],
both are multiples of 'multiples', and uses 'seed' for random tie-break.
"""
RETURN_TYPES = ("INT", "INT")
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"resolution": ("INT", {"default": 1024}),
"min_ratio": ("FLOAT", {"default": 0.6}),
"max_ratio": ("FLOAT", {"default": 1.6}),
"multiples": ("INT", {"default": 32}),
"seed": ("INT", {"default": 0}),
}
}
FUNCTION = "select_dimensions"
CATEGORY = "Logic Gates"
custom_name = "Random Width/Height with Resolution"
def select_dimensions(self, resolution, min_ratio, max_ratio, multiples, seed):
# For reproducible randomness
random.seed(seed)
desired_area = resolution * resolution
ratio = random.uniform(min_ratio, max_ratio)
# width * height = resolution^2, width/height = ratio
# thus h**2 = resolution^2 / ratio
height = int(math.sqrt(desired_area / ratio))
width = int(desired_area / height)
# round to nearest multiple
div_h = height / multiples
div_w = width / multiples
height = round(div_h) * multiples
width = round(div_w) * multiples
# if width * height > resolution^2, reduce width or height
if width * height > desired_area:
if random.choice([True, False]):
width = width - multiples
else:
height = height - multiples
return (width, height)
@node
class SystemRandomInt(RandomGuaranteedClass):
"""
Random number generator using system randomness
Generates an integer value between 0 and 2^32-1
"""
def __init__(self):
pass
@staticmethod
def generate(min_val=0, max_val=2**63-1):
instance = random.SystemRandom(time.time())
value = instance.randint(min_val, max_val)
return (value,)
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"min_val": (
"INT",
{
"default": 0,
"min": -(2**63-1),
"max": 2**63-1,
"step": 1,
"display": "number",
},
),
"max_val": (
"INT",
{
"default": 2**63-1,
"min": -(2**63-1),
"max": 2**63-1,
"step": 1,
"display": "number",
},
),
},
}
RETURN_TYPES = ("INT",)
FUNCTION = "generate"
CATEGORY = "Logic Gates"
custom_name = "System Random Int"
@node
class SystemUUIDGenerator(RandomGuaranteedClass):
"""
Generates a random UUID
"""
def __init__(self):
pass
@staticmethod
def generate(length=36):
value = uuid.uuid4()
value = str(value)
if length < 36:
value = value[:length]
return (value,)
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"length": ("INT", { "default": 36, "min": 1, "max": 36, "step": 1, "display": "number" }),
},
}
RETURN_TYPES = ("STRING",)
FUNCTION = "generate"
CATEGORY = "Logic Gates"
custom_name = "UUID Generator"
@node
class UniformRandomFloat(RandomGuaranteedClass):
"""
Selects a random float from min to max
Fallbacks to default if min is greater than max
@@ -14,7 +175,7 @@ class UniformRandomFloat:
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,...
@@ -28,15 +189,186 @@ class UniformRandomFloat:
"min_val": ("FLOAT", { "default": 0.0, "min": -999999999, "max": 999999999.0, "step": 0.02, "display": "number" }),
"max_val": ("FLOAT", { "default": 1.0, "min": -999999999, "max": 999999999.0, "step": 0.02, "display": "number" }),
"decimal_places": ("INT", { "default": 1, "min": 0, "max": 10, "step": 1, "display": "number" }),
"seed" : ("INT", { "default": 0, "min": 0, "max": 9999999999, "step": 1, "display": "number" }),
"seed" : ("INT", { "default": 0, "min": 0, "max": (2**63-1), "step": 1, "display": "number" }),
},
}
RETURN_TYPES = ("FLOAT",)
FUNCTION = "generate"
CATEGORY = "Logic Gates"
custom_name = "Uniform Random Float"
@node
class UniformRandomInt:
class TriangularRandomFloat(RandomGuaranteedClass):
"""
Selects a random float from min to max
Fallbacks to default if min is greater than max
"""
def __init__(self):
pass
def generate(self, low, high, mode, seed=0):
if low > high:
return (low,)
instance = random.Random(seed)
value = instance.triangular(low, high, mode)
return (value,)
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"low": ("FLOAT", { "default": 0.0, "min": -999999999, "max": 999999999.0, "step": 0.02, "display": "number" }),
"high": ("FLOAT", { "default": 1.0, "min": -999999999, "max": 999999999.0, "step": 0.02, "display": "number" }),
"mode": ("FLOAT", { "default": 0.5, "min": -999999999, "max": 999999999.0, "step": 0.02, "display": "number" }),
"seed" : ("INT", { "default": 0, "min": 0, "max": (2**63-1), "step": 1, "display": "number" }),
},
}
RETURN_TYPES = ("FLOAT",)
FUNCTION = "generate"
CATEGORY = "Logic Gates"
custom_name = "Triangular Random Float"
@node
class WeightedRandomChoice(RandomGuaranteedClass):
"""
Randomly choose one item from a list with weights.
The input string is parsed as "value|weight$value2|weight2..."
Example: "apple|10$banana|1$orange|3"
"""
def __init__(self):
pass
def generate(self, input_string, separator, seed=0):
# Example input: "apple|10$banana|1$orange|3"
# Split by '$' -> ["apple|10", "banana|1", "orange|3"]
items = input_string.split(separator)
choices = []
weights = []
for item in items:
if '|' in item:
val, wt = item.split('|', 1)
choices.append(val)
try:
weights.append(float(wt))
except ValueError:
weights.append(1.0)
else:
# fallback if no weight specified
choices.append(item)
weights.append(1.0)
instance = random.Random(seed)
chosen = instance.choices(population=choices, weights=weights, k=1)[0]
return (chosen,)
RETURN_TYPES = ("STRING",)
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"input_string": ("STRING", {"default": "apple|10$banana|1$orange|3", "display": "text"}),
"separator": ("STRING", {"default": "$", "display": "text"}),
"seed": ("INT", {"default": 0, "min": 0, "max": 2**63-1, "step": 1, "display": "number"}),
}
}
FUNCTION = "generate"
CATEGORY = "Logic Gates"
custom_name = "Weighted Random Choice"
@node
class RandomGaussianFloat(RandomGuaranteedClass):
"""
Generates a random float from a normal (Gaussian) distribution
with specified mean and std_dev.
"""
def __init__(self):
pass
def generate(self, mean, std_dev, decimal_places, seed=0):
instance = random.Random(seed)
value = instance.gauss(mean, std_dev)
value = round(value, decimal_places)
return (value,)
RETURN_TYPES = ("FLOAT",)
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"mean": ("FLOAT", {"default": 0.0, "min": -999999999, "max": 999999999.0, "step": 0.01}),
"std_dev": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 999999999.0, "step": 0.01}),
"decimal_places": ("INT", {"default": 2, "min": 0, "max": 10, "step": 1}),
"seed": ("INT", {"default": 0, "min": 0, "max": 2**63-1, "step": 1}),
},
}
FUNCTION = "generate"
CATEGORY = "Logic Gates"
custom_name = "Random Gaussian Float"
@node
class SystemRandomGaussianFloat(RandomGuaranteedClass):
"""
Generates a random float from a normal (Gaussian) distribution
with specified mean and std_dev.
"""
def __init__(self):
pass
def generate(self, mean, std_dev, decimal_places):
instance = random.SystemRandom(time.time())
value = instance.gauss(mean, std_dev)
value = round(value, decimal_places)
return (value,)
RETURN_TYPES = ("FLOAT",)
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"mean": ("FLOAT", {"default": 0.0, "min": -999999999, "max": 999999999.0, "step": 0.01}),
"std_dev": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 999999999.0, "step": 0.01}),
"decimal_places": ("INT", {"default": 2, "min": 0, "max": 10, "step": 1}),
},
}
FUNCTION = "generate"
CATEGORY = "Logic Gates"
custom_name = "System Random Gaussian Float"
@node
class ProbabilityGate(RandomGuaranteedClass):
"""
Returns TRUE with probability p, FALSE otherwise.
"""
def __init__(self):
pass
def generate(self, probability, seed=0):
instance = random.Random(seed)
value = instance.random() # uniform in [0,1)
return (value < probability,)
RETURN_TYPES = ("BOOLEAN",)
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"probability": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0, "step": 0.01}),
"seed": ("INT", {"default": 0, "min": 0, "max": 2**63-1, "step": 1}),
},
}
FUNCTION = "generate"
CATEGORY = "Logic Gates"
custom_name = "Probability Gate"
@node
class UniformRandomInt(RandomGuaranteedClass):
"""
Selects a random int from min to max
Fallbacks to default if min is greater than max
@@ -45,7 +377,7 @@ class UniformRandomInt:
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}")
@@ -56,7 +388,7 @@ class UniformRandomInt:
"required": {
"min_val": ("INT", { "default": 0, "min": -999999999, "max": 999999999, "step": 1, "display": "number" }),
"max_val": ("INT", { "default": 1, "min": -999999999, "max": 999999999, "step": 1, "display": "number" }),
"seed" : ("INT", { "default": 0, "min": 0, "max": 9999999999, "step": 1, "display": "number" }),
"seed" : ("INT", { "default": 0, "min": 0, "max": (2**63-1), "step": 1, "display": "number" }),
},
}
RETURN_TYPES = ("INT",)
@@ -64,7 +396,7 @@ class UniformRandomInt:
CATEGORY = "Logic Gates"
custom_name = "Uniform Random Int"
@node
class UniformRandomChoice:
class UniformRandomChoice(RandomGuaranteedClass):
"""
Parses input string with separator '$' and returns a random choice
separator can be changed in the input
@@ -82,7 +414,7 @@ class UniformRandomChoice:
"required": {
"input_string": ("STRING", { "default": "a$b$c", "display": "text" }),
"separator": ("STRING", { "default": "$", "display": "text" }),
"seed" : ("INT", { "default": 0, "min": 0, "max": 9999999999, "step": 1, "display": "number" }),
"seed" : ("INT", { "default": 0, "min": 0, "max": (2**63-1), "step": 1, "display": "number" }),
},
}
RETURN_TYPES = ("STRING",)
@@ -108,7 +440,7 @@ class ManualChoiceString:
"required": {
"input_string": ("STRING", { "default": "a$b$c", "display": "text" }),
"separator": ("STRING", { "default": "$", "display": "text" }),
"index": ("INT", { "default": 0, "min": 0, "max": 9999999999, "step": 1, "display": "number" }),
"index": ("INT", { "default": 0, "min": 0, "max": (2**63-1), "step": 1, "display": "number" }),
},
}
RETURN_TYPES = ("STRING",)
@@ -135,7 +467,7 @@ class ManualChoiceInt:
"required": {
"input_string": ("STRING", { "default": "1$2$3", "display": "text" }),
"separator": ("STRING", { "default": "$", "display": "text" }),
"index": ("INT", { "default": 0, "min": 0, "max": 9999999999, "step": 1, "display": "number" }),
"index": ("INT", { "default": 0, "min": 0, "max": (2**63-1), "step": 1, "display": "number" }),
},
}
RETURN_TYPES = ("INT",)
@@ -162,7 +494,7 @@ class ManualChoiceFloat:
"required": {
"input_string": ("STRING", { "default": "1.0$2.0$3.0", "display": "text" }),
"separator": ("STRING", { "default": "$", "display": "text" }),
"index": ("INT", { "default": 0, "min": 0, "max": 9999999999, "step": 1, "display": "number" }),
"index": ("INT", { "default": 0, "min": 0, "max": (2**63-1), "step": 1, "display": "number" }),
},
}
RETURN_TYPES = ("FLOAT",)
@@ -171,7 +503,7 @@ class ManualChoiceFloat:
custom_name = "Manual Choice Float"
@node
class RandomShuffleInt:
class RandomShuffleInt(RandomGuaranteedClass):
"""
Get the shuffled list of integers from start to end
Input types and output types are lists of ints
@@ -190,7 +522,7 @@ class RandomShuffleInt:
"required": {
"input_string": ("STRING", { "default": "1$2$3", "display": "text" }),
"separator": ("STRING", { "default": "$", "display": "text" }),
"seed" : ("INT", { "default": 0, "min": 0, "max": 9999999999, "step": 1, "display": "number" }),
"seed" : ("INT", { "default": 0, "min": 0, "max": (2**63-1), "step": 1, "display": "number" }),
},
}
RETURN_TYPES = ("STRING",)
@@ -198,7 +530,7 @@ class RandomShuffleInt:
CATEGORY = "Logic Gates"
custom_name = "Random Shuffle Int"
@node
class RandomShuffleFloat:
class RandomShuffleFloat(RandomGuaranteedClass):
"""
Get the shuffled list of floats from start to end
Input types and output types are lists of floats
@@ -217,7 +549,7 @@ class RandomShuffleFloat:
"required": {
"input_string": ("STRING", { "default": "1.0$2.0$3.0", "display": "text" }),
"separator": ("STRING", { "default": "$", "display": "text" }),
"seed" : ("INT", { "default": 0, "min": 0, "max": 9999999999, "step": 1, "display": "number" }),
"seed" : ("INT", { "default": 0, "min": 0, "max": (2**63-1), "step": 1, "display": "number" }),
},
}
RETURN_TYPES = ("STRING",)
@@ -225,7 +557,7 @@ class RandomShuffleFloat:
CATEGORY = "Logic Gates"
custom_name = "Random Shuffle Float"
@node
class RandomShuffleString:
class RandomShuffleString(RandomGuaranteedClass):
"""
Get the shuffled list of strings from start to end
Input types and output types are lists of strings
@@ -244,7 +576,7 @@ class RandomShuffleString:
"required": {
"input_string": ("STRING", { "default": "a$b$c", "display": "text" }),
"separator": ("STRING", { "default": "$", "display": "text" }),
"seed" : ("INT", { "default": 0, "min": 0, "max": 9999999999, "step": 1, "display": "number" }),
"seed" : ("INT", { "default": 0, "min": 0, "max": (2**63-1), "step": 1, "display": "number" }),
},
}
RETURN_TYPES = ("STRING",)
@@ -253,13 +585,73 @@ class RandomShuffleString:
custom_name = "Random Shuffle String"
@node
class YieldableIteratorString:
class CounterInteger(RandomGuaranteedClass):
"""
Generates a counter that increments by 1
"""
def __init__(self):
self.counter = None
def generate(self, reset, start):
if self.counter is None:
self.counter = start
if reset:
self.counter = 0
self.counter += 1
return (int(self.counter),)
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"start": ("FLOAT", { "default": 0.0, "min": -(2**63-1), "max": (2**63-1), "step": 1.0, "display": "number" }),
},
"optional": {
"reset": ("BOOLEAN", { "default": False }),
},
}
RETURN_TYPES = ("INT",)
FUNCTION = "generate"
CATEGORY = "Logic Gates"
custom_name = "Counter Integer"
@node
class CounterFloat(RandomGuaranteedClass):
"""
Generates a counter that increments by 1
"""
def __init__(self):
self.counter = None
def generate(self, reset, start, step):
if self.counter is None:
self.counter = start
if reset:
self.counter = start
self.counter += step
return (self.counter,)
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"start": ("FLOAT", { "default": 0.0, "min": -(2**63-1), "max": (2**63-1), "step": 1.0, "display": "number" }),
},
"optional": {
"reset": ("BOOLEAN", {"default": False}),
"step": ("FLOAT", { "default": 1.0, "min": -(2**63-1), "max": (2**63-1), "step": 1.0, "display": "number" }),
},
}
RETURN_TYPES = ("FLOAT",)
FUNCTION = "generate"
CATEGORY = "Logic Gates"
custom_name = "Counter Float"
@node
class YieldableIteratorString(RandomGuaranteedClass):
"""
Yields sequentially from the input list (with separator)
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:
@@ -276,7 +668,7 @@ class YieldableIteratorString:
"required": {
"input_string": ("STRING", { "default": "a$b$c", "display": "text" }),
"separator": ("STRING", { "default": "$", "display": "text" }),
"reset": ("INT", { "default": 0, "min": 0, "max": 1, "step": 1, "display": "number" }),
"reset": ("BOOLEAN", {"default": False}),
},
}
RETURN_TYPES = ("STRING",)
@@ -285,7 +677,7 @@ class YieldableIteratorString:
custom_name = "Yieldable Iterator String"
@node
class YieldableIteratorInt:
class YieldableIteratorInt(RandomGuaranteedClass):
"""
Yields sequentially with start, end, step
Resets if reset is True
@@ -300,21 +692,21 @@ class YieldableIteratorInt:
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
def INPUT_TYPES(s):
return {
"required": {
"start": ("INT", { "default": 0, "min": 0, "max": 9999999999, "step": 1, "display": "number" }),
"end": ("INT", { "default": 10, "min": 0, "max": 9999999999, "step": 1, "display": "number" }),
"step": ("INT", { "default": 1, "min": 0, "max": 9999999999, "step": 1, "display": "number" }),
"reset": ("INT", { "default": 0, "min": 0, "max": 1, "step": 1, "display": "number" }),
"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", {"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)
+57
View File
@@ -0,0 +1,57 @@
import tempfile
import unittest
from pathlib import Path
import torch
from PIL import Image
from import_utils import import_local
class TestIoNodesExtras(unittest.TestCase):
@classmethod
def setUpClass(cls):
cls.io_node = import_local("io_node")
def test_comma_rejoin_defaults(self):
Node = self.io_node.CLASS_MAPPINGS["CommaRejoinNode"]
self.assertEqual(Node.comma_rejoin("a,b , c"), ("a, b, c",))
self.assertEqual(Node.comma_rejoin(" a , b ,c "), ("a, b, c",))
def test_comma_rejoin_custom_separators(self):
Node = self.io_node.CLASS_MAPPINGS["CommaRejoinNode"]
self.assertEqual(
Node.comma_rejoin("a| b |c", split_separator="|", join_separator="| "),
("a| b| c",),
)
def test_random_image_from_folder(self):
Node = self.io_node.CLASS_MAPPINGS["RandomImageFromFolderNode"]
with tempfile.TemporaryDirectory() as tmp:
root = Path(tmp)
sub = root / "sub"
sub.mkdir()
img1 = root / "a.png"
img2 = root / "b.jpg"
img3 = sub / "c.png"
Image.new("RGB", (8, 6), (10, 20, 30)).save(img1)
Image.new("RGB", (8, 6), (40, 50, 60)).save(img2)
Image.new("RGB", (8, 6), (70, 80, 90)).save(img3)
# Non-recursive should only pick from root
tensor, path = Node.random_image_from_folder(
str(root), "all", False, seed=0
)
self.assertTrue(torch.is_tensor(tensor))
self.assertEqual(tuple(tensor.shape), (1, 6, 8, 3))
self.assertIn(Path(path), {img1, img2})
# Recursive should include subfolder images
tensor, path = Node.random_image_from_folder(str(root), "png", True, seed=0)
self.assertTrue(torch.is_tensor(tensor))
self.assertEqual(tuple(tensor.shape), (1, 6, 8, 3))
self.assertIn(Path(path), {img1, img3})
+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",))
+19 -4
View File
@@ -1,9 +1,25 @@
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,
"SwinV2": None,
"ConvNext": None,
"ConvNextV2": None,
"ViT": None,
"MOAT": None,
"SwinV2_v3": None,
"ConvNext_v3": None,
"ViT_v3": None,
}
from typing import Union
from PIL import Image
@@ -27,5 +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
tagger_keys = list(tagger_model_names.keys())