Compare commits
80
Commits
dghs-imgutils
...
main
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
f77c699543 | ||
|
|
06658072ff | ||
|
|
44d9a0a2ce | ||
|
|
c8e10b174a | ||
|
|
54bee29ced | ||
|
|
9fc69cb6a5 | ||
|
|
d4a1fc5f0b | ||
|
|
1e4c7465eb | ||
|
|
5992e91930 | ||
|
|
60f8f1187c | ||
|
|
ae00f1ff8b | ||
|
|
1136141406 | ||
|
|
915e2e9e81 | ||
|
|
d600e0bec7 | ||
|
|
33a0a38ecb | ||
|
|
e07784201f | ||
|
|
f8302d7ac0 | ||
|
|
9533c09491 | ||
|
|
71bc2f0cc1 | ||
|
|
91e76fc2ee | ||
|
|
f0c8f9a8bb | ||
|
|
09345b8d70 | ||
|
|
cbbd61ef40 | ||
|
|
8e0935831a | ||
|
|
c1e4c84a70 | ||
|
|
e9e34dd355 | ||
|
|
2d826bd624 | ||
|
|
a9a8258c1f | ||
|
|
f62821172a | ||
|
|
9624d93ac4 | ||
|
|
46d62f8d6d | ||
|
|
247577ae71 | ||
|
|
56b5cf8267 | ||
|
|
3dbae89658 | ||
|
|
2109871a89 | ||
|
|
6f138ebd4b | ||
|
|
ff4602f6e6 | ||
|
|
05f05534b9 | ||
|
|
20a4170bb8 | ||
|
|
287997120b | ||
|
|
29630d22e4 | ||
|
|
67b532a1d5 | ||
|
|
07ad78cdee | ||
|
|
aebcbc8c97 | ||
|
|
1bdd981fa1 | ||
|
|
2451f43fa7 | ||
|
|
6b9d8c9f32 | ||
|
|
4992749ebb | ||
|
|
079d9a1e59 | ||
|
|
bd84f2534a | ||
|
|
f199fde35c | ||
|
|
db84be70ae | ||
|
|
17b8937cae | ||
|
|
a560f2db7c | ||
|
|
269316c01d | ||
|
|
0e364db321 | ||
|
|
adb5bd9a75 | ||
|
|
c197f0d66d | ||
|
|
436196f851 | ||
|
|
a12b2e6657 | ||
|
|
1afe926c21 | ||
|
|
382e821e06 | ||
|
|
6e08078c2a | ||
|
|
2557c7e3b5 | ||
|
|
eb3a0d0fa4 | ||
|
|
75a3fa7665 | ||
|
|
a384a92373 | ||
|
|
77e1429843 | ||
|
|
b2cc4b0f2b | ||
|
|
3688c504ed | ||
|
|
649de167ec | ||
|
|
d330362006 | ||
|
|
2473ef3e84 | ||
|
|
b04da7a530 | ||
|
|
611253e2ab | ||
|
|
52b422128d | ||
|
|
2e65bc8294 | ||
|
|
f1f0f4572e | ||
|
|
b3ac53a946 | ||
|
|
656aadfc8b |
@@ -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 }}
|
||||
|
||||
@@ -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
|
||||

|
||||
|
||||
## 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`.
|
||||
+4
-4
@@ -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
@@ -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
@@ -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)
|
||||
|
||||
@@ -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
@@ -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
@@ -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)
|
||||
|
||||
+33
-8
@@ -1,4 +1,5 @@
|
||||
#https://github.com/ltdrdata/ComfyUI-Impact-Pack/blob/Main/install.py
|
||||
import os
|
||||
import sys
|
||||
import subprocess
|
||||
import threading
|
||||
@@ -36,25 +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 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()
|
||||
|
||||
+1259
-105
File diff suppressed because it is too large
Load Diff
+31
-30
@@ -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
@@ -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
|
||||
|
||||
@@ -1,15 +1,39 @@
|
||||
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
|
||||
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
|
||||
|
||||
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 = {
|
||||
}
|
||||
@@ -21,6 +45,8 @@ NODE_CLASS_MAPPINGS.update(MathMapping)
|
||||
NODE_CLASS_MAPPINGS.update(ExternalMapping)
|
||||
NODE_CLASS_MAPPINGS.update(AuxilaryMapping)
|
||||
|
||||
|
||||
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
|
||||
}
|
||||
@@ -31,3 +57,17 @@ 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
|
||||
|
||||
+3
-1
@@ -1,7 +1,7 @@
|
||||
[project]
|
||||
name = "comfyui-logicutils"
|
||||
description = "Logical Utils (compare, string, boolean operations) for ComfyUI"
|
||||
version = "1.2.0"
|
||||
version = "1.8.0"
|
||||
license = "MIT"
|
||||
|
||||
[project.urls]
|
||||
@@ -12,3 +12,5 @@ Repository = "https://github.com/aria1th/ComfyUI-LogicUtils"
|
||||
PublisherId = "angelbottomless"
|
||||
DisplayName = "ComfyUI-LogicUtils"
|
||||
Icon = ""
|
||||
|
||||
|
||||
|
||||
+1031
File diff suppressed because it is too large
Load Diff
+419
-27
@@ -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}),
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
@@ -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",))
|
||||
+19
-4
@@ -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())
|
||||
|
||||
Reference in New Issue
Block a user