Merge pull request #47 from ComfyAssets/bugs/fix-save-and-sampler
Bugs/fix save and sampler
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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}"
|
||||
|
||||
Reference in New Issue
Block a user