Compare commits
26
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
60f30b068e | ||
|
|
1439270fe5 | ||
|
|
f8210d8f69 | ||
|
|
aceb34b9b0 | ||
|
|
e9ce1fd2cf | ||
|
|
9b888443ac | ||
|
|
f4743df3ef | ||
|
|
6a68983ef4 | ||
|
|
fb01fa24ae | ||
|
|
65f68f59a1 | ||
|
|
c3fab5581b | ||
|
|
1ea2b4cc90 | ||
|
|
f33f39f134 | ||
|
|
5b57d4fc35 | ||
|
|
a1625dddad | ||
|
|
a88232f59a | ||
|
|
6746b86685 | ||
|
|
17af18d397 | ||
|
|
703989599d | ||
|
|
8399fad96b | ||
|
|
0c69abc829 | ||
|
|
e00406747f | ||
|
|
a21e677629 | ||
|
|
13e425959b | ||
|
|
0a6ee72748 | ||
|
|
b03f0ecf22 |
@@ -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
|
||||
|
||||
@@ -38,6 +38,7 @@ I’m sharing them here with the community, and I hope you find them as useful a
|
||||
| [🔤 Embedding Autocomplete](#-embedding-autocomplete) | Smart autocomplete for embeddings, LoRAs, and tags | 🔧 Utils |
|
||||
| [🧹 Kiko Purge VRAM](#-kiko-purge-vram) | Intelligent VRAM management with detailed reporting | 🛠️ Utils |
|
||||
| [📂 Local Image Loader](#-local-image-loader) | Visual gallery browser for local media files | 💾 Images |
|
||||
| [🌐 Model Downloader](#-model-downloader) | Download models from CivitAI, HuggingFace, and custom URLs | 🛠️ Utils |
|
||||
|
||||
### 🧰 xyz-helpers Tools
|
||||
|
||||
@@ -378,6 +379,32 @@ Visual gallery browser for loading local images, videos, and audio files directl
|
||||
- Extract prompts and settings from previously generated images
|
||||
- Preview media files before loading into workflow
|
||||
|
||||
#### 🌐 Model Downloader
|
||||
Download models, LoRAs, and other assets directly from CivitAI, HuggingFace, and custom URLs within ComfyUI.
|
||||
|
||||
- **Multi-Platform Support**: CivitAI, HuggingFace, and direct download URLs
|
||||
- **Smart URL Detection**: Automatic detection of download source and file handling
|
||||
- **API Token Support**: Optional authentication for private/gated models
|
||||
- **Progress Reporting**: Real-time download progress with speed indicators
|
||||
- **Resume Support**: Skip existing files or force re-download
|
||||
- **Interrupt Handling**: Respects ComfyUI's "Cancel current run" button
|
||||
- **Automatic Cleanup**: Removes partial downloads on cancellation
|
||||
- **Custom Filenames**: Override auto-detected filenames when needed
|
||||
|
||||
**Platform Features:**
|
||||
- **CivitAI**: Model page URLs, version-specific downloads, API authentication
|
||||
- **HuggingFace**: Blob and resolve URLs, branch/revision support, gated model access
|
||||
- **Custom URLs**: Direct download links with bearer token authentication
|
||||
|
||||
**Use Cases:**
|
||||
- Download models without leaving ComfyUI
|
||||
- Automate asset acquisition in workflows
|
||||
- Access private or gated models with API tokens
|
||||
- Build reproducible workflows with automatic model fetching
|
||||
- Quickly test new models from the community
|
||||
|
||||

|
||||
|
||||
### 🔤 Embedding Autocomplete
|
||||
|
||||
**Intelligent autocomplete for embeddings, LoRAs, and custom tags in text prompts.**
|
||||
@@ -699,6 +726,8 @@ Example workflow available: [xyz_helpers_lora_testing.json](examples/workflows/x
|
||||
| **Sampler Select Helper** | Intelligent sampler selection with recommendations | ✅ Complete | [Docs](examples/documentation/sampler_select_helper.md) |
|
||||
| **Scheduler Select Helper** | Optimal scheduler selection for samplers | ✅ Complete | [Docs](examples/documentation/scheduler_select_helper.md) |
|
||||
| **Text Encode Sampler Params** | Combined text encoding and parameter management | ✅ Complete | [Docs](examples/documentation/text_encode_sampler_params.md) |
|
||||
| **Local Image Loader** | Visual gallery browser for local media files | ✅ Complete | [Docs](examples/documentation/local_image_loader.md) |
|
||||
| **Model Downloader** | Download models from CivitAI, HuggingFace, and custom URLs | ✅ Complete | [Docs](examples/documentation/model_downloader.md) |
|
||||
| **Batch Image Processor** | Process multiple images with consistent settings | 🚧 Planned | Coming Soon |
|
||||
| **Advanced Prompt Utilities** | Enhanced prompt manipulation and generation | 🚧 Planned | Coming Soon |
|
||||
|
||||
@@ -981,7 +1010,7 @@ MIT License - see [LICENSE](LICENSE) file for details.
|
||||
|
||||
## 🏷️ Tags
|
||||
|
||||
`comfyui` `custom-nodes` `image-processing` `ai-tools` `sdxl` `flux` `upscaling` `resolution` `batch-processing` `python` `pytorch`
|
||||
`comfyui` `custom-nodes` `image-processing` `ai-tools` `sdxl` `flux` `upscaling` `resolution` `batch-processing` `model-downloader` `civitai` `huggingface` `python` `pytorch`
|
||||
|
||||
## 🔗 Links
|
||||
|
||||
@@ -992,14 +1021,15 @@ MIT License - see [LICENSE](LICENSE) file for details.
|
||||
|
||||
## 📈 Stats
|
||||
|
||||
- **Nodes**: 19 (13 core tools + 6 xyz-helpers)
|
||||
- **Nodes**: 20 (14 core tools + 6 xyz-helpers)
|
||||
- **Features**: Embedding Autocomplete (settings-based, not a node)
|
||||
- **Categories**: 9 emoji-based categories for better organization
|
||||
- **Download Platforms**: 3 (CivitAI, HuggingFace, Custom URLs)
|
||||
- **Format Support**: 3 (PNG, JPEG, WebP with advanced controls)
|
||||
- **Presets**: 26 curated resolution presets
|
||||
- **Interactive Features**: 8+ (swap buttons, history UI, popup viewers, parameter visualization)
|
||||
- **AI Integration**: Gemini API with 40+ model support
|
||||
- **Test Coverage**: 100% (300+ comprehensive tests)
|
||||
- **Test Coverage**: 100% (470+ comprehensive tests)
|
||||
- **Python Version**: 3.8+
|
||||
- **ComfyUI Compatibility**: Latest
|
||||
- **Dependencies**: Minimal (PyTorch, NumPy, Pillow, google-generativeai for Gemini)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
{
|
||||
"last_path": "/home/vito/ai-apps/ComfyUI-3.12/output",
|
||||
"last_path": "/home/vito/ai-apps/ComfyUI/output",
|
||||
"saved_paths": [
|
||||
"/home/vito/ai-apps/ComfyUI-3.12/output/2025-05-01",
|
||||
"/home/vito/ai-apps/ComfyUI-3.12/output/"
|
||||
|
||||
@@ -8,5 +8,15 @@
|
||||
"image": {
|
||||
"path": "/home/vito/ai-apps/ComfyUI-3.12/output/CharacterName_00016_.png"
|
||||
}
|
||||
},
|
||||
"18": {
|
||||
"image": {
|
||||
"path": "/home/vito/ai-apps/ComfyUI/output/ComfyUI_00002_.png"
|
||||
}
|
||||
},
|
||||
"445": {
|
||||
"image": {
|
||||
"path": "/home/vito/ai-apps/ComfyUI/output/ComfyUI_00002_.png"
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -6,6 +6,13 @@ from typing import Optional, Callable
|
||||
from urllib.parse import urlparse, unquote
|
||||
import os
|
||||
|
||||
try:
|
||||
import comfy.model_management
|
||||
|
||||
COMFY_AVAILABLE = True
|
||||
except ImportError:
|
||||
COMFY_AVAILABLE = False
|
||||
|
||||
|
||||
class BaseDownloader(ABC):
|
||||
"""Abstract base class for all downloaders"""
|
||||
@@ -38,6 +45,15 @@ class BaseDownloader(ABC):
|
||||
if self._progress_callback:
|
||||
self._progress_callback(downloaded, total, message)
|
||||
|
||||
def check_interrupt(self) -> None:
|
||||
"""Check if processing has been interrupted by user
|
||||
|
||||
Raises:
|
||||
comfy.model_management.InterruptProcessingException: If user cancelled
|
||||
"""
|
||||
if COMFY_AVAILABLE:
|
||||
comfy.model_management.throw_exception_if_processing_interrupted()
|
||||
|
||||
def extract_filename(self, url: str, default: str = "downloaded_file") -> str:
|
||||
"""Extract filename from URL
|
||||
|
||||
|
||||
@@ -12,6 +12,18 @@ from urllib.parse import urlparse, parse_qs, unquote
|
||||
|
||||
from .base import BaseDownloader
|
||||
|
||||
try:
|
||||
import comfy.model_management
|
||||
|
||||
COMFY_AVAILABLE = True
|
||||
InterruptProcessingException = comfy.model_management.InterruptProcessingException
|
||||
except ImportError:
|
||||
COMFY_AVAILABLE = False
|
||||
# Fallback exception type that will never be raised
|
||||
InterruptProcessingException = type(
|
||||
"InterruptProcessingException", (Exception,), {}
|
||||
)
|
||||
|
||||
|
||||
CHUNK_SIZE = 1638400
|
||||
USER_AGENT = "Mozilla/5.0 (Windows NT 10.0; Win64; x64) AppleWebKit/537.36"
|
||||
@@ -136,8 +148,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
|
||||
@@ -251,61 +270,72 @@ class CivitAIDownloader(BaseDownloader):
|
||||
print(f"Size: {self.format_size(total_size)}")
|
||||
|
||||
# Download with progress
|
||||
with open(output_file, "wb") as f:
|
||||
downloaded = 0
|
||||
start_time = time.time()
|
||||
try:
|
||||
with open(output_file, "wb") as f:
|
||||
downloaded = 0
|
||||
start_time = time.time()
|
||||
|
||||
while True:
|
||||
chunk_start_time = time.time()
|
||||
buffer = response.read(CHUNK_SIZE)
|
||||
chunk_end_time = time.time()
|
||||
while True:
|
||||
chunk_start_time = time.time()
|
||||
buffer = response.read(CHUNK_SIZE)
|
||||
chunk_end_time = time.time()
|
||||
|
||||
if not buffer:
|
||||
break
|
||||
if not buffer:
|
||||
break
|
||||
|
||||
downloaded += len(buffer)
|
||||
f.write(buffer)
|
||||
chunk_time = chunk_end_time - chunk_start_time
|
||||
downloaded += len(buffer)
|
||||
f.write(buffer)
|
||||
chunk_time = chunk_end_time - chunk_start_time
|
||||
|
||||
# Calculate speed
|
||||
speed = self.calculate_speed(len(buffer), chunk_time)
|
||||
# Check for user cancellation
|
||||
self.check_interrupt()
|
||||
|
||||
# Report progress
|
||||
if total_size is not None:
|
||||
progress = downloaded / total_size
|
||||
sys.stdout.write(
|
||||
f'\r[{"=" * int(progress * 50):<50}] {progress * 100:.2f}% - {speed:.2f} MB/s'
|
||||
)
|
||||
sys.stdout.flush()
|
||||
self.report_progress(downloaded, total_size, f"{speed:.2f} MB/s")
|
||||
else:
|
||||
sys.stdout.write(
|
||||
f"\rDownloaded: {self.format_size(downloaded)} - {speed:.2f} MB/s"
|
||||
)
|
||||
sys.stdout.flush()
|
||||
self.report_progress(downloaded, 0, f"{speed:.2f} MB/s")
|
||||
# Calculate speed
|
||||
speed = self.calculate_speed(len(buffer), chunk_time)
|
||||
|
||||
end_time = time.time()
|
||||
time_taken = end_time - start_time
|
||||
hours, remainder = divmod(time_taken, 3600)
|
||||
minutes, seconds = divmod(remainder, 60)
|
||||
# Report progress
|
||||
if total_size is not None:
|
||||
progress = downloaded / total_size
|
||||
sys.stdout.write(
|
||||
f'\r[{"=" * int(progress * 50):<50}] {progress * 100:.2f}% - {speed:.2f} MB/s'
|
||||
)
|
||||
sys.stdout.flush()
|
||||
self.report_progress(
|
||||
downloaded, total_size, f"{speed:.2f} MB/s"
|
||||
)
|
||||
else:
|
||||
sys.stdout.write(
|
||||
f"\rDownloaded: {self.format_size(downloaded)} - {speed:.2f} MB/s"
|
||||
)
|
||||
sys.stdout.flush()
|
||||
self.report_progress(downloaded, 0, f"{speed:.2f} MB/s")
|
||||
|
||||
if hours > 0:
|
||||
time_str = f"{int(hours)}h {int(minutes)}m {int(seconds)}s"
|
||||
elif minutes > 0:
|
||||
time_str = f"{int(minutes)}m {int(seconds)}s"
|
||||
else:
|
||||
time_str = f"{int(seconds)}s"
|
||||
end_time = time.time()
|
||||
time_taken = end_time - start_time
|
||||
hours, remainder = divmod(time_taken, 3600)
|
||||
minutes, seconds = divmod(remainder, 60)
|
||||
|
||||
sys.stdout.write("\n")
|
||||
print(f"✓ Download completed in {time_str}")
|
||||
print(f"✓ File saved as: {output_file}")
|
||||
if hours > 0:
|
||||
time_str = f"{int(hours)}h {int(minutes)}m {int(seconds)}s"
|
||||
elif minutes > 0:
|
||||
time_str = f"{int(minutes)}m {int(seconds)}s"
|
||||
else:
|
||||
time_str = f"{int(seconds)}s"
|
||||
|
||||
# Verify file size
|
||||
actual_size = os.path.getsize(output_file)
|
||||
if total_size and actual_size != total_size:
|
||||
raise Exception(
|
||||
f"Download incomplete. Expected {total_size} bytes, got {actual_size} bytes"
|
||||
)
|
||||
sys.stdout.write("\n")
|
||||
print(f"✓ Download completed in {time_str}")
|
||||
print(f"✓ File saved as: {output_file}")
|
||||
|
||||
return output_file
|
||||
# Verify file size
|
||||
actual_size = os.path.getsize(output_file)
|
||||
if total_size and actual_size != total_size:
|
||||
raise Exception(
|
||||
f"Download incomplete. Expected {total_size} bytes, got {actual_size} bytes"
|
||||
)
|
||||
|
||||
return output_file
|
||||
except InterruptProcessingException:
|
||||
# Clean up partial download on interrupt
|
||||
if os.path.exists(output_file):
|
||||
os.remove(output_file)
|
||||
raise InterruptProcessingException("Download interrupted")
|
||||
|
||||
@@ -9,6 +9,18 @@ from typing import Optional
|
||||
|
||||
from .base import BaseDownloader
|
||||
|
||||
try:
|
||||
import comfy.model_management
|
||||
|
||||
COMFY_AVAILABLE = True
|
||||
InterruptProcessingException = comfy.model_management.InterruptProcessingException
|
||||
except ImportError:
|
||||
COMFY_AVAILABLE = False
|
||||
# Fallback exception type that will never be raised
|
||||
InterruptProcessingException = type(
|
||||
"InterruptProcessingException", (Exception,), {}
|
||||
)
|
||||
|
||||
|
||||
CHUNK_SIZE = 1638400
|
||||
USER_AGENT = "Mozilla/5.0 (Windows NT 10.0; Win64; x64) AppleWebKit/537.36"
|
||||
@@ -120,62 +132,73 @@ class CustomDownloader(BaseDownloader):
|
||||
print("Size: Unknown")
|
||||
|
||||
# Download with progress
|
||||
with open(output_file, "wb") as f:
|
||||
downloaded = 0
|
||||
start_time = time.time()
|
||||
try:
|
||||
with open(output_file, "wb") as f:
|
||||
downloaded = 0
|
||||
start_time = time.time()
|
||||
|
||||
while True:
|
||||
chunk_start_time = time.time()
|
||||
buffer = response.read(CHUNK_SIZE)
|
||||
chunk_end_time = time.time()
|
||||
while True:
|
||||
chunk_start_time = time.time()
|
||||
buffer = response.read(CHUNK_SIZE)
|
||||
chunk_end_time = time.time()
|
||||
|
||||
if not buffer:
|
||||
break
|
||||
if not buffer:
|
||||
break
|
||||
|
||||
downloaded += len(buffer)
|
||||
f.write(buffer)
|
||||
chunk_time = chunk_end_time - chunk_start_time
|
||||
downloaded += len(buffer)
|
||||
f.write(buffer)
|
||||
chunk_time = chunk_end_time - chunk_start_time
|
||||
|
||||
# Calculate speed
|
||||
speed = self.calculate_speed(len(buffer), chunk_time)
|
||||
# Check for user cancellation
|
||||
self.check_interrupt()
|
||||
|
||||
# Report progress
|
||||
if total_size is not None:
|
||||
progress = downloaded / total_size
|
||||
sys.stdout.write(
|
||||
f'\r[{"=" * int(progress * 50):<50}] {progress * 100:.2f}% - {speed:.2f} MB/s'
|
||||
)
|
||||
sys.stdout.flush()
|
||||
self.report_progress(downloaded, total_size, f"{speed:.2f} MB/s")
|
||||
else:
|
||||
sys.stdout.write(
|
||||
f"\rDownloaded: {self.format_size(downloaded)} - {speed:.2f} MB/s"
|
||||
)
|
||||
sys.stdout.flush()
|
||||
self.report_progress(downloaded, 0, f"{speed:.2f} MB/s")
|
||||
# Calculate speed
|
||||
speed = self.calculate_speed(len(buffer), chunk_time)
|
||||
|
||||
end_time = time.time()
|
||||
time_taken = end_time - start_time
|
||||
hours, remainder = divmod(time_taken, 3600)
|
||||
minutes, seconds = divmod(remainder, 60)
|
||||
# Report progress
|
||||
if total_size is not None:
|
||||
progress = downloaded / total_size
|
||||
sys.stdout.write(
|
||||
f'\r[{"=" * int(progress * 50):<50}] {progress * 100:.2f}% - {speed:.2f} MB/s'
|
||||
)
|
||||
sys.stdout.flush()
|
||||
self.report_progress(
|
||||
downloaded, total_size, f"{speed:.2f} MB/s"
|
||||
)
|
||||
else:
|
||||
sys.stdout.write(
|
||||
f"\rDownloaded: {self.format_size(downloaded)} - {speed:.2f} MB/s"
|
||||
)
|
||||
sys.stdout.flush()
|
||||
self.report_progress(downloaded, 0, f"{speed:.2f} MB/s")
|
||||
|
||||
if hours > 0:
|
||||
time_str = f"{int(hours)}h {int(minutes)}m {int(seconds)}s"
|
||||
elif minutes > 0:
|
||||
time_str = f"{int(minutes)}m {int(seconds)}s"
|
||||
else:
|
||||
time_str = f"{int(seconds)}s"
|
||||
end_time = time.time()
|
||||
time_taken = end_time - start_time
|
||||
hours, remainder = divmod(time_taken, 3600)
|
||||
minutes, seconds = divmod(remainder, 60)
|
||||
|
||||
sys.stdout.write("\n")
|
||||
print(f"✓ Download completed in {time_str}")
|
||||
print(f"✓ File saved as: {output_file}")
|
||||
if hours > 0:
|
||||
time_str = f"{int(hours)}h {int(minutes)}m {int(seconds)}s"
|
||||
elif minutes > 0:
|
||||
time_str = f"{int(minutes)}m {int(seconds)}s"
|
||||
else:
|
||||
time_str = f"{int(seconds)}s"
|
||||
|
||||
# Verify file size if known
|
||||
actual_size = os.path.getsize(output_file)
|
||||
if total_size and actual_size != total_size:
|
||||
print(
|
||||
f"⚠ Warning: Downloaded size ({actual_size} bytes) doesn't match expected size ({total_size} bytes)"
|
||||
)
|
||||
# Don't raise error for custom URLs as size mismatch might be acceptable
|
||||
sys.stdout.write("\n")
|
||||
print(f"✓ Download completed in {time_str}")
|
||||
print(f"✓ File saved as: {output_file}")
|
||||
|
||||
return output_file
|
||||
# Verify file size if known
|
||||
actual_size = os.path.getsize(output_file)
|
||||
if total_size and actual_size != total_size:
|
||||
print(
|
||||
f"⚠ Warning: Downloaded size ({actual_size} bytes) doesn't match expected size ({total_size} bytes)"
|
||||
)
|
||||
# Don't raise error for custom URLs as size mismatch might be acceptable
|
||||
|
||||
return output_file
|
||||
except InterruptProcessingException:
|
||||
# Clean up partial download on interrupt
|
||||
if os.path.exists(output_file):
|
||||
os.remove(output_file)
|
||||
raise
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -10,6 +10,18 @@ from urllib.parse import urlparse, quote
|
||||
|
||||
from .base import BaseDownloader
|
||||
|
||||
try:
|
||||
import comfy.model_management
|
||||
|
||||
COMFY_AVAILABLE = True
|
||||
InterruptProcessingException = comfy.model_management.InterruptProcessingException
|
||||
except ImportError:
|
||||
COMFY_AVAILABLE = False
|
||||
# Fallback exception type that will never be raised
|
||||
InterruptProcessingException = type(
|
||||
"InterruptProcessingException", (Exception,), {}
|
||||
)
|
||||
|
||||
|
||||
CHUNK_SIZE = 1638400
|
||||
USER_AGENT = "Mozilla/5.0 (Windows NT 10.0; Win64; x64) AppleWebKit/537.36"
|
||||
@@ -188,61 +200,72 @@ class HuggingFaceDownloader(BaseDownloader):
|
||||
print(f"Size: {self.format_size(total_size)}")
|
||||
|
||||
# Download with progress
|
||||
with open(output_file, "wb") as f:
|
||||
downloaded = 0
|
||||
start_time = time.time()
|
||||
try:
|
||||
with open(output_file, "wb") as f:
|
||||
downloaded = 0
|
||||
start_time = time.time()
|
||||
|
||||
while True:
|
||||
chunk_start_time = time.time()
|
||||
buffer = response.read(CHUNK_SIZE)
|
||||
chunk_end_time = time.time()
|
||||
while True:
|
||||
chunk_start_time = time.time()
|
||||
buffer = response.read(CHUNK_SIZE)
|
||||
chunk_end_time = time.time()
|
||||
|
||||
if not buffer:
|
||||
break
|
||||
if not buffer:
|
||||
break
|
||||
|
||||
downloaded += len(buffer)
|
||||
f.write(buffer)
|
||||
chunk_time = chunk_end_time - chunk_start_time
|
||||
downloaded += len(buffer)
|
||||
f.write(buffer)
|
||||
chunk_time = chunk_end_time - chunk_start_time
|
||||
|
||||
# Calculate speed
|
||||
speed = self.calculate_speed(len(buffer), chunk_time)
|
||||
# Check for user cancellation
|
||||
self.check_interrupt()
|
||||
|
||||
# Report progress
|
||||
if total_size is not None:
|
||||
progress = downloaded / total_size
|
||||
sys.stdout.write(
|
||||
f'\r[{"=" * int(progress * 50):<50}] {progress * 100:.2f}% - {speed:.2f} MB/s'
|
||||
)
|
||||
sys.stdout.flush()
|
||||
self.report_progress(downloaded, total_size, f"{speed:.2f} MB/s")
|
||||
else:
|
||||
sys.stdout.write(
|
||||
f"\rDownloaded: {self.format_size(downloaded)} - {speed:.2f} MB/s"
|
||||
)
|
||||
sys.stdout.flush()
|
||||
self.report_progress(downloaded, 0, f"{speed:.2f} MB/s")
|
||||
# Calculate speed
|
||||
speed = self.calculate_speed(len(buffer), chunk_time)
|
||||
|
||||
end_time = time.time()
|
||||
time_taken = end_time - start_time
|
||||
hours, remainder = divmod(time_taken, 3600)
|
||||
minutes, seconds = divmod(remainder, 60)
|
||||
# Report progress
|
||||
if total_size is not None:
|
||||
progress = downloaded / total_size
|
||||
sys.stdout.write(
|
||||
f'\r[{"=" * int(progress * 50):<50}] {progress * 100:.2f}% - {speed:.2f} MB/s'
|
||||
)
|
||||
sys.stdout.flush()
|
||||
self.report_progress(
|
||||
downloaded, total_size, f"{speed:.2f} MB/s"
|
||||
)
|
||||
else:
|
||||
sys.stdout.write(
|
||||
f"\rDownloaded: {self.format_size(downloaded)} - {speed:.2f} MB/s"
|
||||
)
|
||||
sys.stdout.flush()
|
||||
self.report_progress(downloaded, 0, f"{speed:.2f} MB/s")
|
||||
|
||||
if hours > 0:
|
||||
time_str = f"{int(hours)}h {int(minutes)}m {int(seconds)}s"
|
||||
elif minutes > 0:
|
||||
time_str = f"{int(minutes)}m {int(seconds)}s"
|
||||
else:
|
||||
time_str = f"{int(seconds)}s"
|
||||
end_time = time.time()
|
||||
time_taken = end_time - start_time
|
||||
hours, remainder = divmod(time_taken, 3600)
|
||||
minutes, seconds = divmod(remainder, 60)
|
||||
|
||||
sys.stdout.write("\n")
|
||||
print(f"✓ Download completed in {time_str}")
|
||||
print(f"✓ File saved as: {output_file}")
|
||||
if hours > 0:
|
||||
time_str = f"{int(hours)}h {int(minutes)}m {int(seconds)}s"
|
||||
elif minutes > 0:
|
||||
time_str = f"{int(minutes)}m {int(seconds)}s"
|
||||
else:
|
||||
time_str = f"{int(seconds)}s"
|
||||
|
||||
# Verify file size
|
||||
actual_size = os.path.getsize(output_file)
|
||||
if total_size and actual_size != total_size:
|
||||
raise Exception(
|
||||
f"Download incomplete. Expected {total_size} bytes, got {actual_size} bytes"
|
||||
)
|
||||
sys.stdout.write("\n")
|
||||
print(f"✓ Download completed in {time_str}")
|
||||
print(f"✓ File saved as: {output_file}")
|
||||
|
||||
return output_file
|
||||
# Verify file size
|
||||
actual_size = os.path.getsize(output_file)
|
||||
if total_size and actual_size != total_size:
|
||||
raise Exception(
|
||||
f"Download incomplete. Expected {total_size} bytes, got {actual_size} bytes"
|
||||
)
|
||||
|
||||
return output_file
|
||||
except InterruptProcessingException:
|
||||
# Clean up partial download on interrupt
|
||||
if os.path.exists(output_file):
|
||||
os.remove(output_file)
|
||||
raise
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -28,11 +28,28 @@ class SeedHistoryNode(ComfyAssetsBaseNode):
|
||||
"default": 12345,
|
||||
"min": 0,
|
||||
"max": 0xFFFFFFFFFFFFFFFF,
|
||||
"control_after_generate": True,
|
||||
"tooltip": "Seed value for generation processes. "
|
||||
"History UI tracks all changes automatically.",
|
||||
"Auto-increments/decrements after each run based on mode.",
|
||||
},
|
||||
),
|
||||
}
|
||||
},
|
||||
"optional": {
|
||||
"mode": (
|
||||
[
|
||||
"",
|
||||
"fixed",
|
||||
"increment",
|
||||
"decrement",
|
||||
"randomize",
|
||||
], # Added empty string for legacy workflows
|
||||
{
|
||||
"default": "fixed",
|
||||
"tooltip": "Seed behavior after generation: "
|
||||
"fixed (no change), increment (+1), decrement (-1), or randomize (new random)",
|
||||
},
|
||||
),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("INT",)
|
||||
@@ -40,17 +57,48 @@ class SeedHistoryNode(ComfyAssetsBaseNode):
|
||||
FUNCTION = "output_seed"
|
||||
CATEGORY = "🫶 ComfyAssets/🌱 Seeds"
|
||||
|
||||
def output_seed(self, seed: int) -> Tuple[int]:
|
||||
@classmethod
|
||||
def VALIDATE_INPUTS(cls, seed, mode="fixed"):
|
||||
"""Validate inputs and handle legacy workflows."""
|
||||
# Handle empty or missing mode from old workflows (legacy support)
|
||||
if mode is None or mode == "" or mode == "undefined":
|
||||
return True # Will use default "fixed" in output_seed
|
||||
|
||||
# Validate mode is in allowed list
|
||||
valid_modes = ["fixed", "increment", "decrement", "randomize"]
|
||||
if mode not in valid_modes:
|
||||
return f"Invalid mode: {mode}. Must be one of {valid_modes}"
|
||||
|
||||
return True
|
||||
|
||||
def output_seed(self, seed: int, mode: str = "fixed") -> Tuple[int]:
|
||||
"""
|
||||
Output the seed value for use in other nodes.
|
||||
|
||||
Args:
|
||||
seed: Input seed value
|
||||
mode: Seed mode (fixed, increment, decrement, randomize) - not used in output,
|
||||
but controls the widget behavior via control_after_generate
|
||||
|
||||
Returns:
|
||||
Tuple containing the seed value
|
||||
"""
|
||||
try:
|
||||
# Handle empty mode from old workflows
|
||||
if not mode or mode == "":
|
||||
mode = "fixed"
|
||||
|
||||
# Validate mode is in allowed list
|
||||
valid_modes = ["fixed", "increment", "decrement", "randomize"]
|
||||
if mode not in valid_modes:
|
||||
import logging
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
logger.warning(
|
||||
f"{self.__class__.__name__}: Invalid mode '{mode}'. Using 'fixed'."
|
||||
)
|
||||
mode = "fixed"
|
||||
|
||||
# Validate and sanitize the seed
|
||||
if not validate_seed_value(seed):
|
||||
# Log the validation error but don't raise
|
||||
@@ -65,6 +113,10 @@ class SeedHistoryNode(ComfyAssetsBaseNode):
|
||||
|
||||
clean_seed = sanitize_seed_value(seed)
|
||||
|
||||
# Note: The mode parameter controls the widget's control_after_generate behavior
|
||||
# The actual increment/decrement/randomize happens automatically in the UI
|
||||
# based on the control_after_generate setting and the mode dropdown value
|
||||
|
||||
return (clean_seed,)
|
||||
|
||||
except Exception as e:
|
||||
|
||||
@@ -78,7 +78,37 @@ PRESET_METADATA: Dict[str, PresetMetadata] = {
|
||||
"Portrait",
|
||||
"SDXL portrait 5:12 - very tall portrait",
|
||||
),
|
||||
"704×1408": PresetMetadata(
|
||||
704,
|
||||
1408,
|
||||
"1:2",
|
||||
0.5,
|
||||
0.99,
|
||||
"SDXL",
|
||||
"Portrait",
|
||||
"SDXL portrait 1:2 - extreme tall portrait",
|
||||
),
|
||||
"960×1024": PresetMetadata(
|
||||
960,
|
||||
1024,
|
||||
"15:16",
|
||||
0.938,
|
||||
0.98,
|
||||
"SDXL",
|
||||
"Portrait",
|
||||
"SDXL near-square portrait - subtle portrait",
|
||||
),
|
||||
# SDXL Presets - Landscape
|
||||
"1024×960": PresetMetadata(
|
||||
1024,
|
||||
960,
|
||||
"16:15",
|
||||
1.067,
|
||||
0.98,
|
||||
"SDXL",
|
||||
"Landscape",
|
||||
"SDXL near-square landscape - subtle landscape",
|
||||
),
|
||||
"1152×896": PresetMetadata(
|
||||
1152,
|
||||
896,
|
||||
@@ -119,6 +149,16 @@ PRESET_METADATA: Dict[str, PresetMetadata] = {
|
||||
"Landscape",
|
||||
"SDXL landscape 12:5 - very wide landscape",
|
||||
),
|
||||
"1728×576": PresetMetadata(
|
||||
1728,
|
||||
576,
|
||||
"3:1",
|
||||
3.0,
|
||||
1.0,
|
||||
"SDXL",
|
||||
"Landscape",
|
||||
"SDXL landscape 3:1 - extreme wide panoramic",
|
||||
),
|
||||
# FLUX Presets - High Quality
|
||||
"1920×1080": PresetMetadata(
|
||||
1920,
|
||||
|
||||
@@ -152,13 +152,14 @@ class FluxSamplerParamsNode(ComfyAssetsBaseNode):
|
||||
import comfy.samplers
|
||||
import comfy.model_base
|
||||
import comfy.model_management
|
||||
import comfy.utils
|
||||
import torch
|
||||
from comfy_extras.nodes_custom_sampler import (
|
||||
Noise_RandomNoise,
|
||||
BasicScheduler,
|
||||
BasicGuider,
|
||||
SamplerCustomAdvanced,
|
||||
)
|
||||
from comfy_extras.nodes_latent import LatentBatch
|
||||
from comfy_extras.nodes_model_advanced import (
|
||||
ModelSamplingFlux,
|
||||
ModelSamplingAuraFlow,
|
||||
@@ -170,6 +171,33 @@ class FluxSamplerParamsNode(ComfyAssetsBaseNode):
|
||||
self.handle_error(f"Required ComfyUI modules not available: {e}")
|
||||
return (latent_image, [])
|
||||
|
||||
# Local implementation of LatentBatch functionality
|
||||
# Copied from nodes_latent.py to avoid V3 schema breaking changes
|
||||
def reshape_latent_to(target_shape, latent, repeat_batch=True):
|
||||
"""Reshape latent tensor to match target shape."""
|
||||
if latent.shape[1:] != target_shape[1:]:
|
||||
latent = comfy.utils.common_upscale(
|
||||
latent, target_shape[-1], target_shape[-2], "bilinear", "center"
|
||||
)
|
||||
if repeat_batch:
|
||||
return comfy.utils.repeat_to_batch_size(latent, target_shape[0])
|
||||
else:
|
||||
return latent
|
||||
|
||||
def batch_latents(samples1, samples2):
|
||||
"""Batch two latent samples together."""
|
||||
samples_out = samples1.copy()
|
||||
s1 = samples1["samples"]
|
||||
s2 = samples2["samples"]
|
||||
|
||||
s2 = reshape_latent_to(s1.shape, s2, repeat_batch=False)
|
||||
s = torch.cat((s1, s2), dim=0)
|
||||
samples_out["samples"] = s
|
||||
samples_out["batch_index"] = samples1.get(
|
||||
"batch_index", [x for x in range(0, s1.shape[0])]
|
||||
) + samples2.get("batch_index", [x for x in range(0, s2.shape[0])])
|
||||
return samples_out
|
||||
|
||||
try:
|
||||
if not validate_flux_params(
|
||||
steps, guidance, max_shift, base_shift, denoise
|
||||
@@ -236,7 +264,6 @@ class FluxSamplerParamsNode(ComfyAssetsBaseNode):
|
||||
basicscheduler = BasicScheduler()
|
||||
basicguider = BasicGuider()
|
||||
samplercustomadvanced = SamplerCustomAdvanced()
|
||||
latentbatch = LatentBatch()
|
||||
modelsampling = (
|
||||
ModelSamplingFlux() if not is_schnell else ModelSamplingAuraFlow()
|
||||
)
|
||||
@@ -364,7 +391,7 @@ class FluxSamplerParamsNode(ComfyAssetsBaseNode):
|
||||
if out_latent is None:
|
||||
out_latent = latent
|
||||
else:
|
||||
out_latent = latentbatch.batch(out_latent, latent)[0]
|
||||
out_latent = batch_latents(out_latent, latent)
|
||||
|
||||
if total_samples > 1:
|
||||
pbar.update(1)
|
||||
|
||||
+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}"
|
||||
|
||||
@@ -1,7 +1,8 @@
|
||||
"""Tests for Flux Sampler Params node."""
|
||||
|
||||
import pytest
|
||||
from unittest.mock import Mock, MagicMock
|
||||
import torch
|
||||
from unittest.mock import Mock, MagicMock, patch
|
||||
from kikotools.tools.xyz_helpers.flux_sampler_params import FluxSamplerParamsNode
|
||||
from kikotools.tools.xyz_helpers.flux_sampler_params.logic import (
|
||||
parse_string_to_list,
|
||||
@@ -192,3 +193,102 @@ class TestFluxSamplerParamsNode:
|
||||
node = FluxSamplerParamsNode()
|
||||
assert node.lora_loader is None
|
||||
assert node.cached_lora == (None, None)
|
||||
|
||||
|
||||
class TestLatentBatchingFunctions:
|
||||
"""Test the local latent batching implementation (copied from nodes_latent.py)."""
|
||||
|
||||
def test_batch_latents_basic(self):
|
||||
"""Test basic latent batching functionality."""
|
||||
# This test verifies the local implementation works correctly
|
||||
# The actual batch_latents function is defined inside process_batch method
|
||||
# so we need to mock the imports and test through the node
|
||||
|
||||
# Create mock latent samples
|
||||
samples1 = {
|
||||
"samples": torch.randn(2, 4, 64, 64), # batch=2
|
||||
"batch_index": [0, 1],
|
||||
}
|
||||
|
||||
samples2 = {
|
||||
"samples": torch.randn(3, 4, 64, 64), # batch=3
|
||||
"batch_index": [0, 1, 2],
|
||||
}
|
||||
|
||||
# We can't directly test batch_latents since it's defined inside process_batch
|
||||
# But we can verify the logic by checking tensor concatenation behavior
|
||||
s1 = samples1["samples"]
|
||||
s2 = samples2["samples"]
|
||||
|
||||
# Verify shapes match for concatenation
|
||||
assert s1.shape[1:] == s2.shape[1:] # channels, height, width match
|
||||
|
||||
# Simulate batching
|
||||
batched = torch.cat((s1, s2), dim=0)
|
||||
|
||||
# Verify output shape
|
||||
assert batched.shape[0] == 5 # 2 + 3
|
||||
assert batched.shape[1:] == s1.shape[1:]
|
||||
|
||||
def test_reshape_latent_logic(self):
|
||||
"""Test the reshape latent to logic."""
|
||||
# Test that tensors with matching shapes don't need reshaping
|
||||
latent = torch.randn(2, 4, 64, 64)
|
||||
target_shape = (2, 4, 64, 64)
|
||||
|
||||
# Verify shapes match
|
||||
assert latent.shape[1:] == target_shape[1:]
|
||||
|
||||
# Test with different batch sizes
|
||||
latent_small = torch.randn(1, 4, 64, 64)
|
||||
target_large = (5, 4, 64, 64)
|
||||
|
||||
# Small latent can be repeated to match larger batch
|
||||
assert latent_small.shape[1:] == target_large[1:]
|
||||
|
||||
def test_batch_index_concatenation(self):
|
||||
"""Test that batch indices are properly concatenated."""
|
||||
# Simulate batch index concatenation logic
|
||||
batch_index1 = [0, 1]
|
||||
batch_index2 = [0, 1, 2]
|
||||
|
||||
combined = batch_index1 + batch_index2
|
||||
|
||||
assert combined == [0, 1, 0, 1, 2]
|
||||
assert len(combined) == 5
|
||||
|
||||
def test_latent_samples_copy(self):
|
||||
"""Test that samples dictionary is properly copied."""
|
||||
samples1 = {
|
||||
"samples": torch.randn(2, 4, 64, 64),
|
||||
"batch_index": [0, 1],
|
||||
"extra_key": "value",
|
||||
}
|
||||
|
||||
# Simulate copy behavior
|
||||
samples_out = samples1.copy()
|
||||
|
||||
# Verify it's a shallow copy
|
||||
assert samples_out is not samples1
|
||||
assert samples_out["samples"] is samples1["samples"] # shallow copy
|
||||
assert samples_out["batch_index"] == samples1["batch_index"]
|
||||
assert samples_out["extra_key"] == samples1["extra_key"]
|
||||
|
||||
def test_reshape_latent_to_logic_verification(self):
|
||||
"""Test reshape_latent_to function logic without ComfyUI dependencies."""
|
||||
# This test verifies the logic without needing actual comfy imports
|
||||
|
||||
# Create test data
|
||||
target_shape = (5, 4, 128, 128)
|
||||
latent = torch.randn(2, 4, 64, 64)
|
||||
|
||||
# Verify the logic conditions that would trigger reshaping:
|
||||
# 1. If shapes don't match (height/width), upscale would be called
|
||||
assert latent.shape[1:] != target_shape[1:]
|
||||
|
||||
# 2. If batch sizes are different, repeat would be called
|
||||
assert latent.shape[0] != target_shape[0]
|
||||
|
||||
# Test case where no reshaping is needed
|
||||
matching_latent = torch.randn(5, 4, 128, 128)
|
||||
assert matching_latent.shape == target_shape
|
||||
|
||||
@@ -33,6 +33,9 @@ class KikoEmbeddingAutocomplete {
|
||||
this.widgetCleanupMap = new WeakMap();
|
||||
this.activeWidgets = new WeakSet();
|
||||
|
||||
// Track pending fetch requests for cleanup
|
||||
this.pendingFetches = new Set();
|
||||
|
||||
// Debounce resource fetching
|
||||
this.fetchResourcesDebounced = this.debounce(() => this.fetchResources(), 500);
|
||||
this.fetchResourcesDebounced();
|
||||
@@ -42,9 +45,6 @@ class KikoEmbeddingAutocomplete {
|
||||
this.loadCustomWords(this.settings.customWordsUrl);
|
||||
}
|
||||
|
||||
// Track pending fetch requests for cleanup
|
||||
this.pendingFetches = new Set();
|
||||
|
||||
// Single suggestion container for all widgets
|
||||
this.suggestionContainer = null;
|
||||
this.currentWidget = null;
|
||||
@@ -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