Compare commits
13
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
1ea2b4cc90 | ||
|
|
f33f39f134 | ||
|
|
5b57d4fc35 | ||
|
|
a1625dddad | ||
|
|
a88232f59a | ||
|
|
6746b86685 | ||
|
|
17af18d397 | ||
|
|
703989599d | ||
|
|
8399fad96b | ||
|
|
0c69abc829 | ||
|
|
e00406747f | ||
|
|
a21e677629 | ||
|
|
13e425959b |
@@ -161,9 +161,8 @@ jobs:
|
||||
assert 'cfg' in input_types['required']
|
||||
print('✓ Sampler Combo interface tests passed')
|
||||
|
||||
# Test return types
|
||||
# RETURN_TYPES[1] is the actual SCHEDULERS list
|
||||
assert node.RETURN_TYPES[0] == 'SAMPLER'
|
||||
# Test return types - Updated to match SAMPLERS list change
|
||||
assert node.RETURN_TYPES[0] == SAMPLERS # Now returns SAMPLERS list
|
||||
assert isinstance(node.RETURN_TYPES[1], list) # SCHEDULERS is a list
|
||||
assert node.RETURN_TYPES[2] == 'INT'
|
||||
assert node.RETURN_TYPES[3] == 'FLOAT'
|
||||
|
||||
@@ -159,6 +159,7 @@ test_images/
|
||||
test_outputs/
|
||||
experiments/
|
||||
.claude/
|
||||
.serena
|
||||
|
||||
# Gemini model cache
|
||||
.gemini_models_cache.json
|
||||
|
||||
@@ -10,7 +10,6 @@ from PIL import Image
|
||||
from PIL.PngImagePlugin import PngInfo
|
||||
import torch
|
||||
from typing import Dict, List, Any, Optional, Tuple
|
||||
import time
|
||||
|
||||
try:
|
||||
import folder_paths
|
||||
@@ -22,9 +21,53 @@ except ImportError:
|
||||
return "./output"
|
||||
|
||||
|
||||
def get_next_counter(output_dir: str, prefix: str) -> int:
|
||||
"""
|
||||
Get next available counter value from persistent counter file
|
||||
|
||||
This prevents file overwrites when the node is called multiple times
|
||||
within the same second by maintaining a persistent counter.
|
||||
|
||||
Args:
|
||||
output_dir: Directory to store counter file
|
||||
prefix: Filename prefix to create unique counter per prefix
|
||||
|
||||
Returns:
|
||||
Next available counter value
|
||||
"""
|
||||
# Create a safe counter filename
|
||||
safe_prefix = "".join(c for c in prefix if c.isalnum() or c in "._-")
|
||||
counter_file = os.path.join(output_dir, f".{safe_prefix}_counter.txt")
|
||||
|
||||
# Read current counter
|
||||
counter = 0
|
||||
if os.path.exists(counter_file):
|
||||
try:
|
||||
with open(counter_file, "r") as f:
|
||||
content = f.read().strip()
|
||||
counter = int(content) if content else 0
|
||||
except (ValueError, IOError):
|
||||
# If file is corrupted or unreadable, start from 0
|
||||
counter = 0
|
||||
|
||||
# Increment counter
|
||||
counter += 1
|
||||
|
||||
# Save updated counter
|
||||
try:
|
||||
with open(counter_file, "w") as f:
|
||||
f.write(str(counter))
|
||||
except IOError:
|
||||
# If we can't write the counter file, continue anyway
|
||||
# Better to risk overwrites than to fail completely
|
||||
pass
|
||||
|
||||
return counter
|
||||
|
||||
|
||||
def get_save_image_path(
|
||||
filename_prefix: str,
|
||||
batch_number: int,
|
||||
counter: int,
|
||||
format_ext: str,
|
||||
output_dir: str,
|
||||
subfolder: str = "",
|
||||
@@ -34,13 +77,13 @@ def get_save_image_path(
|
||||
|
||||
Args:
|
||||
filename_prefix: Base filename prefix
|
||||
batch_number: Batch index for multiple images
|
||||
counter: Persistent counter to ensure unique filenames
|
||||
format_ext: File extension (.png, .jpg, .webp)
|
||||
output_dir: Output directory path
|
||||
subfolder: Optional subfolder within output directory
|
||||
|
||||
Returns:
|
||||
Tuple of (full_path, relative_filename)
|
||||
Tuple of (full_path, preview_filename, relative_subfolder)
|
||||
"""
|
||||
# Split filename_prefix into directory path and actual filename prefix
|
||||
# This allows for directory structures like "kittybear/anime/images/kittybear"
|
||||
@@ -53,9 +96,10 @@ def get_save_image_path(
|
||||
) # Only sanitize problematic chars for filenames
|
||||
safe_prefix = "".join(c for c in safe_prefix if c.isalnum() or c in "._-")
|
||||
|
||||
# Create unique filename with timestamp to avoid conflicts
|
||||
timestamp = int(time.time())
|
||||
filename = f"{safe_prefix}_{timestamp:010d}_{batch_number:05d}{format_ext}"
|
||||
# Create unique filename with counter to avoid conflicts
|
||||
# Using counter instead of timestamp+batch_number prevents overwrites
|
||||
# when multiple images are processed separately
|
||||
filename = f"{safe_prefix}_{counter:05d}{format_ext}"
|
||||
|
||||
# Handle subfolder and prefix directory (but not the filename part)
|
||||
path_components = []
|
||||
@@ -262,13 +306,17 @@ def process_image_batch(
|
||||
results = []
|
||||
enhanced_data = []
|
||||
|
||||
for batch_number, image_tensor in enumerate(images):
|
||||
for image_tensor in images:
|
||||
# Convert tensor to PIL Image
|
||||
img = convert_tensor_to_pil(image_tensor)
|
||||
|
||||
# Generate save path
|
||||
# Get next counter value to ensure unique filenames
|
||||
# This counter persists across node calls, preventing overwrites
|
||||
counter = get_next_counter(output_dir, filename_prefix)
|
||||
|
||||
# Generate save path with persistent counter
|
||||
filepath, preview_filename, relative_subfolder = get_save_image_path(
|
||||
filename_prefix, batch_number, format_ext, output_dir, ""
|
||||
filename_prefix, counter, format_ext, output_dir, ""
|
||||
)
|
||||
|
||||
# Save with format-specific settings
|
||||
|
||||
@@ -136,8 +136,15 @@ class CivitAIDownloader(BaseDownloader):
|
||||
# Validate output path
|
||||
self.validate_output_path(output_path)
|
||||
|
||||
# Validate that URL is from civitai.com domain
|
||||
parsed_url = urlparse(url)
|
||||
if parsed_url.netloc not in ("civitai.com", "www.civitai.com"):
|
||||
raise ValueError(
|
||||
f"Invalid URL: Only civitai.com URLs are supported, got {parsed_url.netloc}"
|
||||
)
|
||||
|
||||
# Convert web URL to API URL if needed
|
||||
if "civitai.com" in url and "/api/download/models/" not in url:
|
||||
if "/api/download/models/" not in url:
|
||||
ids = self._parse_civitai_url(url)
|
||||
|
||||
# If we have a version ID, use it directly
|
||||
|
||||
@@ -67,7 +67,8 @@ class URLDetector:
|
||||
Returns:
|
||||
True if CivitAI URL
|
||||
"""
|
||||
if "civitai.com" not in parsed.netloc:
|
||||
# Validate exact domain match to prevent subdomain attacks
|
||||
if parsed.netloc not in ("civitai.com", "www.civitai.com"):
|
||||
return False
|
||||
|
||||
# Check for API download endpoint
|
||||
@@ -90,8 +91,15 @@ class URLDetector:
|
||||
Returns:
|
||||
True if HuggingFace URL
|
||||
"""
|
||||
# Check main domain and CDN
|
||||
if "huggingface.co" in parsed.netloc:
|
||||
# Validate exact domain match to prevent subdomain attacks
|
||||
# Support both main domain and CDN domains
|
||||
allowed_domains = (
|
||||
"huggingface.co",
|
||||
"www.huggingface.co",
|
||||
"cdn.huggingface.co",
|
||||
"cdn-lfs.huggingface.co",
|
||||
)
|
||||
if parsed.netloc in allowed_domains:
|
||||
return True
|
||||
|
||||
return False
|
||||
|
||||
@@ -144,11 +144,11 @@ class ModelDownloaderNode(ComfyAssetsBaseNode):
|
||||
import time
|
||||
import hashlib
|
||||
|
||||
# Create a unique hash based on inputs and current time
|
||||
input_str = (
|
||||
f"{url}|{save_path}|{filename}|{api_token}|{force_download}|{time.time()}"
|
||||
)
|
||||
return hashlib.md5(input_str.encode()).hexdigest()
|
||||
# Create a unique hash based on non-sensitive inputs and current time
|
||||
# Note: api_token is excluded to avoid sensitive data in hash
|
||||
# The token doesn't affect cache invalidation - URL changes are sufficient
|
||||
input_str = f"{url}|{save_path}|{filename}|{force_download}|{time.time()}"
|
||||
return hashlib.sha256(input_str.encode()).hexdigest()
|
||||
|
||||
|
||||
# Node display name
|
||||
|
||||
@@ -59,7 +59,7 @@ class SamplerComboCompactNode(ComfyAssetsBaseNode):
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("SAMPLER", SCHEDULERS, "INT", "FLOAT")
|
||||
RETURN_TYPES = (SAMPLERS, SCHEDULERS, "INT", "FLOAT")
|
||||
RETURN_NAMES = ("sampler", "scheduler", "steps", "cfg")
|
||||
FUNCTION = "get_combo"
|
||||
CATEGORY = "🫶 ComfyAssets/🌀 Samplers"
|
||||
@@ -82,27 +82,13 @@ class SamplerComboCompactNode(ComfyAssetsBaseNode):
|
||||
try:
|
||||
# Use the same validation logic but with compact interface
|
||||
result = get_sampler_combo(sampler, sched, steps, cfg)
|
||||
# Create the sampler object
|
||||
try:
|
||||
import comfy.samplers
|
||||
|
||||
sampler_obj = comfy.samplers.sampler_object(result[0])
|
||||
except ImportError:
|
||||
# Return sampler name for testing
|
||||
sampler_obj = result[0]
|
||||
return (sampler_obj, result[1], result[2], result[3])
|
||||
# Return the sampler name as string, not object
|
||||
return result
|
||||
|
||||
except Exception as e:
|
||||
# Graceful fallback
|
||||
self.handle_error(f"Error in compact combo: {str(e)}")
|
||||
try:
|
||||
import comfy.samplers
|
||||
|
||||
sampler_obj = comfy.samplers.sampler_object("euler")
|
||||
except ImportError:
|
||||
# Return sampler name for testing
|
||||
sampler_obj = "euler"
|
||||
return (sampler_obj, "normal", 20, 7.0)
|
||||
return ("euler", "normal", 20, 7.0)
|
||||
|
||||
def __str__(self) -> str:
|
||||
"""String representation of the compact node."""
|
||||
|
||||
@@ -64,7 +64,7 @@ class SamplerComboNode(ComfyAssetsBaseNode):
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("SAMPLER", SCHEDULERS, "INT", "FLOAT")
|
||||
RETURN_TYPES = (SAMPLERS, SCHEDULERS, "INT", "FLOAT")
|
||||
RETURN_NAMES = ("sampler_name", "scheduler", "steps", "cfg")
|
||||
FUNCTION = "get_sampler_combo"
|
||||
CATEGORY = "🫶 ComfyAssets/🌀 Samplers"
|
||||
@@ -97,33 +97,18 @@ class SamplerComboNode(ComfyAssetsBaseNode):
|
||||
f"steps={steps}, cfg={cfg}. "
|
||||
f"Using safe defaults: euler, normal, 20 steps, CFG 7.0"
|
||||
)
|
||||
try:
|
||||
import comfy.samplers
|
||||
|
||||
sampler = comfy.samplers.sampler_object("euler")
|
||||
except ImportError:
|
||||
# Return mock object for testing
|
||||
sampler = "euler"
|
||||
return (sampler, "normal", 20, 7.0)
|
||||
return ("euler", "normal", 20, 7.0)
|
||||
|
||||
# Process and return the combo
|
||||
result = get_sampler_combo(sampler_name, scheduler, steps, cfg)
|
||||
|
||||
# Create the sampler object
|
||||
try:
|
||||
import comfy.samplers
|
||||
|
||||
sampler = comfy.samplers.sampler_object(result[0])
|
||||
except ImportError:
|
||||
# Return sampler name for testing
|
||||
sampler = result[0]
|
||||
|
||||
self.log_info(
|
||||
f"Configured sampler combo: {result[0]}, {result[1]}, "
|
||||
f"{result[2]} steps, CFG {result[3]}"
|
||||
)
|
||||
|
||||
return (sampler, result[1], result[2], result[3])
|
||||
# Return the sampler name as string, not object
|
||||
return result
|
||||
|
||||
except Exception as e:
|
||||
# Handle any unexpected errors gracefully
|
||||
@@ -134,14 +119,7 @@ class SamplerComboNode(ComfyAssetsBaseNode):
|
||||
f"{self.__class__.__name__}: Error processing sampler combo: {str(e)}. "
|
||||
f"Using safe defaults: euler, normal, 20 steps, CFG 7.0"
|
||||
)
|
||||
try:
|
||||
import comfy.samplers
|
||||
|
||||
sampler = comfy.samplers.sampler_object("euler")
|
||||
except ImportError:
|
||||
# Return mock object for testing
|
||||
sampler = "euler"
|
||||
return (sampler, "normal", 20, 7.0)
|
||||
return ("euler", "normal", 20, 7.0)
|
||||
|
||||
def validate_inputs(
|
||||
self, sampler_name: str, scheduler: str, steps: int, cfg: float
|
||||
|
||||
+1
-1
@@ -5,7 +5,7 @@ build-backend = "setuptools.build_meta"
|
||||
[project]
|
||||
name = "kikotools"
|
||||
description = "Simple tools for ComfyUI"
|
||||
version = "1.0.23"
|
||||
version = "1.0.24"
|
||||
license = {text = "MIT"}
|
||||
dependencies = []
|
||||
|
||||
|
||||
@@ -18,6 +18,7 @@ from kikotools.tools.kiko_save_image.logic import (
|
||||
save_image_with_format,
|
||||
get_save_image_path,
|
||||
create_png_metadata,
|
||||
get_next_counter,
|
||||
)
|
||||
|
||||
|
||||
@@ -48,26 +49,105 @@ class TestKikoSaveImageLogic:
|
||||
assert pil_image.size == (32, 32)
|
||||
assert pil_image.mode == "RGBA"
|
||||
|
||||
def test_get_save_image_path(self):
|
||||
"""Test save path generation"""
|
||||
def test_get_next_counter_creates_file(self):
|
||||
"""Test counter file creation"""
|
||||
with tempfile.TemporaryDirectory() as temp_dir:
|
||||
# Test basic path generation
|
||||
# First call should create file with counter = 1
|
||||
counter = get_next_counter(temp_dir, "test_prefix")
|
||||
assert counter == 1
|
||||
|
||||
# Verify counter file was created
|
||||
counter_file = os.path.join(temp_dir, ".test_prefix_counter.txt")
|
||||
assert os.path.exists(counter_file)
|
||||
|
||||
# Verify content
|
||||
with open(counter_file, "r") as f:
|
||||
assert f.read().strip() == "1"
|
||||
|
||||
def test_get_next_counter_increments(self):
|
||||
"""Test counter increments correctly"""
|
||||
with tempfile.TemporaryDirectory() as temp_dir:
|
||||
# Multiple calls should increment
|
||||
counter1 = get_next_counter(temp_dir, "test")
|
||||
counter2 = get_next_counter(temp_dir, "test")
|
||||
counter3 = get_next_counter(temp_dir, "test")
|
||||
|
||||
assert counter1 == 1
|
||||
assert counter2 == 2
|
||||
assert counter3 == 3
|
||||
|
||||
def test_get_next_counter_different_prefixes(self):
|
||||
"""Test counters are independent per prefix"""
|
||||
with tempfile.TemporaryDirectory() as temp_dir:
|
||||
# Different prefixes should have separate counters
|
||||
counter_a1 = get_next_counter(temp_dir, "prefix_a")
|
||||
counter_b1 = get_next_counter(temp_dir, "prefix_b")
|
||||
counter_a2 = get_next_counter(temp_dir, "prefix_a")
|
||||
|
||||
assert counter_a1 == 1
|
||||
assert counter_b1 == 1 # Independent counter
|
||||
assert counter_a2 == 2
|
||||
|
||||
def test_get_next_counter_corrupted_file(self):
|
||||
"""Test counter handles corrupted counter files"""
|
||||
with tempfile.TemporaryDirectory() as temp_dir:
|
||||
# Create corrupted counter file
|
||||
counter_file = os.path.join(temp_dir, ".test_counter.txt")
|
||||
with open(counter_file, "w") as f:
|
||||
f.write("not_a_number")
|
||||
|
||||
# Should handle gracefully and start from 1
|
||||
counter = get_next_counter(temp_dir, "test")
|
||||
assert counter == 1
|
||||
|
||||
def test_get_next_counter_empty_file(self):
|
||||
"""Test counter handles empty counter files"""
|
||||
with tempfile.TemporaryDirectory() as temp_dir:
|
||||
# Create empty counter file
|
||||
counter_file = os.path.join(temp_dir, ".test_counter.txt")
|
||||
with open(counter_file, "w") as f:
|
||||
f.write("")
|
||||
|
||||
# Should handle gracefully and start from 1
|
||||
counter = get_next_counter(temp_dir, "test")
|
||||
assert counter == 1
|
||||
|
||||
def test_get_next_counter_sanitizes_prefix(self):
|
||||
"""Test counter sanitizes special characters in prefix"""
|
||||
with tempfile.TemporaryDirectory() as temp_dir:
|
||||
# Prefix with special characters
|
||||
get_next_counter(temp_dir, "test/prefix:with*special")
|
||||
|
||||
# Counter file should be created with sanitized name
|
||||
# Should only contain alphanumeric, dot, dash, underscore
|
||||
counter_files = [
|
||||
f for f in os.listdir(temp_dir) if f.endswith("_counter.txt")
|
||||
]
|
||||
assert len(counter_files) == 1
|
||||
assert "/" not in counter_files[0]
|
||||
assert ":" not in counter_files[0]
|
||||
assert "*" not in counter_files[0]
|
||||
|
||||
def test_get_save_image_path(self):
|
||||
"""Test save path generation with counter"""
|
||||
with tempfile.TemporaryDirectory() as temp_dir:
|
||||
# Test basic path generation with counter
|
||||
full_path, filename, subfolder = get_save_image_path(
|
||||
"test_prefix", 0, ".png", temp_dir
|
||||
"test_prefix", 1, ".png", temp_dir
|
||||
)
|
||||
|
||||
assert full_path.startswith(temp_dir)
|
||||
assert filename.startswith("test_prefix_")
|
||||
assert filename.endswith("_00000.png")
|
||||
assert filename.endswith("00001.png")
|
||||
|
||||
# Test with empty subfolder (standard behavior)
|
||||
# Test with different counter values
|
||||
full_path, filename, subfolder = get_save_image_path(
|
||||
"test", 1, ".jpg", temp_dir, ""
|
||||
"test", 42, ".jpg", temp_dir, ""
|
||||
)
|
||||
|
||||
assert full_path.startswith(temp_dir)
|
||||
assert filename.startswith("test_")
|
||||
assert filename.endswith("_00001.jpg")
|
||||
assert filename.endswith("00042.jpg")
|
||||
|
||||
def test_create_png_metadata(self):
|
||||
"""Test PNG metadata creation"""
|
||||
@@ -550,3 +630,47 @@ class TestIntegration:
|
||||
|
||||
img = Image.open(filepath)
|
||||
assert img.size == (64, 64)
|
||||
|
||||
@patch("kikotools.tools.kiko_save_image.logic.folder_paths")
|
||||
def test_multiple_calls_no_overwrites(self, mock_folder_paths):
|
||||
"""Test that multiple node calls don't overwrite files (bug fix verification)"""
|
||||
with tempfile.TemporaryDirectory() as temp_dir:
|
||||
mock_folder_paths.get_output_directory.return_value = temp_dir
|
||||
|
||||
node = KikoSaveImageNode()
|
||||
|
||||
# Simulate the bug scenario: 6 separate calls with single images
|
||||
# This would have caused overwrites before the counter fix
|
||||
all_filenames = []
|
||||
|
||||
for i in range(6):
|
||||
# Each call processes a single image (like in the bug report)
|
||||
single_image = torch.rand(1, 32, 32, 3)
|
||||
|
||||
result = node.save_images(
|
||||
images=single_image,
|
||||
filename_prefix="KikoSave",
|
||||
format="PNG",
|
||||
)
|
||||
|
||||
# Collect filenames
|
||||
for image_info in result["ui"]["images"]:
|
||||
all_filenames.append(image_info["filename"])
|
||||
|
||||
# Verify all 6 images were saved with unique filenames
|
||||
assert len(all_filenames) == 6
|
||||
assert len(set(all_filenames)) == 6 # All filenames are unique
|
||||
|
||||
# Verify all files actually exist
|
||||
for filename in all_filenames:
|
||||
filepath = os.path.join(temp_dir, filename)
|
||||
assert os.path.exists(filepath), f"File {filename} should exist"
|
||||
|
||||
# Verify filenames follow counter pattern
|
||||
# Should be: KikoSave_00001.png, KikoSave_00002.png, ..., KikoSave_00006.png
|
||||
sorted_filenames = sorted(all_filenames)
|
||||
for i, filename in enumerate(sorted_filenames, start=1):
|
||||
expected_counter = f"{i:05d}"
|
||||
assert (
|
||||
expected_counter in filename
|
||||
), f"Expected counter {expected_counter} in {filename}"
|
||||
|
||||
@@ -294,22 +294,33 @@ class KikoEmbeddingAutocomplete {
|
||||
|
||||
const text = await response.text();
|
||||
|
||||
// Security validation
|
||||
// Security validation - detect potentially dangerous patterns
|
||||
// Note: This is defense-in-depth for text files, not HTML rendering
|
||||
const dangerousPatterns = [
|
||||
/<script[\s\S]*?<\/script>/gi,
|
||||
/<iframe[\s\S]*?>/gi,
|
||||
// Script tags - catch all variations including bypass attempts with whitespace
|
||||
// Matches: <script>, <script >, <script src="x">, </script >, </script>, etc.
|
||||
/<\s*\/?script[^>]*>/gi, // Script tags (opening/closing with any attributes/whitespace)
|
||||
// Other dangerous HTML tags
|
||||
/<\s*iframe[^>]*>/gi,
|
||||
/<\s*embed[^>]*>/gi,
|
||||
/<\s*object[^>]*>/gi,
|
||||
// JavaScript protocol and event handlers
|
||||
/javascript:/gi,
|
||||
/\bon(click|load|error|mouseover|mouseout|focus|blur|change|submit)\s*=/gi,
|
||||
/<embed[\s\S]*?>/gi,
|
||||
/<object[\s\S]*?>/gi,
|
||||
/import\s+[\s\S]*?from/gi,
|
||||
/data:text\/html/gi,
|
||||
/\bon\w+\s*=/gi, // Matches any event handler (onclick, onload, etc.)
|
||||
// Code execution patterns
|
||||
/import\s+/gi,
|
||||
/require\s*\(/gi,
|
||||
/eval\s*\(/gi,
|
||||
/new\s+Function\s*\(/gi,
|
||||
/\.innerHTML\s*=/gi,
|
||||
/document\.\w+/gi,
|
||||
/window\.\w+/gi,
|
||||
/(__proto__|\.prototype\.|\.constructor\s*\()/gi
|
||||
// DOM manipulation
|
||||
/document\./gi,
|
||||
/window\./gi,
|
||||
// Prototype pollution
|
||||
/__proto__/gi,
|
||||
/\.prototype\./gi,
|
||||
/\.constructor\s*\(/gi
|
||||
];
|
||||
|
||||
for (const pattern of dangerousPatterns) {
|
||||
|
||||
Reference in New Issue
Block a user