This commit is contained in:
negaga53
2025-07-07 23:11:12 +02:00
parent b9e07a3683
commit 9a599a4dfc
3 changed files with 367 additions and 272 deletions
+126 -26
View File
@@ -16,20 +16,47 @@ import base64
import io
import os
import logging
import pathlib
from typing import Tuple, Optional, Union
# Try to import ComfyUI utilities
try:
import folder_paths
import node_helpers
from comfy.model_management import soft_empty_cache
from nodes import PreviewImage, SaveImage
except ImportError:
# Fallback for development/testing
class MockFolderPaths:
@staticmethod
def get_input_directory():
return "input"
@staticmethod
def get_annotated_filepath(filename):
return os.path.join("input", filename)
@staticmethod
def exists_annotated_filepath(filename):
return os.path.exists(os.path.join("input", filename))
class MockNodeHelpers:
@staticmethod
def pillow(func, *args, **kwargs):
return func(*args, **kwargs)
class MockPreviewImage:
def save_images(self, images, filename_prefix, prompt=None, extra_pnginfo=None):
return {"ui": {"images": []}}
class MockSaveImage:
def save_images(self, images, filename_prefix, prompt=None, extra_pnginfo=None):
return {"ui": {"images": []}}
folder_paths = MockFolderPaths()
node_helpers = MockNodeHelpers()
PreviewImage = MockPreviewImage
SaveImage = MockSaveImage
def soft_empty_cache():
pass
@@ -63,8 +90,19 @@ class ImageLoader:
@classmethod
def INPUT_TYPES(cls):
"""Define the input types and their configurations."""
# Get list of supported image files from input directory
input_dir = folder_paths.get_input_directory()
files = []
if os.path.exists(input_dir):
files = [f.name for f in pathlib.Path(input_dir).iterdir() if f.is_file()]
return {
"required": {},
"required": {
"image": (sorted(files), {
"image_upload": True,
"tooltip": "Select an image file from the input directory"
})
},
"optional": {
"filepath": ("STRING", {
"default": "",
@@ -96,16 +134,27 @@ class ImageLoader:
inputs.append(f"{key}:{value}")
return hash(tuple(inputs))
def load_image(self, filepath: str = "", base64: str = "", pasted_base64: str = "") -> Tuple[torch.Tensor, torch.Tensor]:
@classmethod
def VALIDATE_INPUTS(cls, **kwargs):
"""Validate input parameters."""
image = kwargs.get("image", "")
if image and image.strip():
if not folder_paths.exists_annotated_filepath(image):
return f"Invalid image file: {image}"
return True
def load_image(self, image: str = "", filepath: str = "", base64: str = "", pasted_base64: str = "") -> Tuple[torch.Tensor, torch.Tensor]:
"""
Load an image from one of the available sources with precedence handling.
Precedence order:
1. Pasted image (from clipboard via JavaScript)
2. File path
3. Base64 string
2. Image upload (file picker)
3. File path
4. Base64 string
Args:
image: Image file from file picker
filepath: Path to image file
base64: Base64 encoded image string
pasted_base64: Base64 data from clipboard paste (populated by JS)
@@ -114,26 +163,40 @@ class ImageLoader:
Tuple of (image_tensor, mask_tensor)
"""
try:
image_data, source_info = self._get_image_data(filepath, base64, pasted_base64)
image_data, source_info = self._get_image_data(image, filepath, base64, pasted_base64)
if image_data is None:
logger.warning("No valid image source provided")
return self._create_empty_tensors()
image_tensor, mask_tensor = self._create_empty_tensors()
else:
logger.info(f"Loading image from: {source_info}")
image_tensor, mask_tensor = self._process_image_data(image_data)
logger.info(f"Loading image from: {source_info}")
return self._process_image_data(image_data)
# Show the image in the UI
results = self.easySave(image_tensor, "imgloader", "Preview", None, None)
return {"ui": {"images": results},
"result": (image_tensor, mask_tensor)}
except Exception as e:
logger.error(f"Error loading image: {e}")
return self._create_empty_tensors()
image_tensor, mask_tensor = self._create_empty_tensors()
results = self.easySave(image_tensor, "imgloader", "Preview", None, None)
return {"ui": {"images": results},
"result": (image_tensor, mask_tensor)}
finally:
# Clean up GPU memory
soft_empty_cache()
def _get_image_data(self, filepath: str, base64_str: str, pasted_base64: str) -> Tuple[Optional[bytes], str]:
def _get_image_data(self, image: str, filepath: str, base64_str: str, pasted_base64: str) -> Tuple[Optional[bytes], str]:
"""
Extract image data from available sources following precedence rules.
Precedence order (higher precedence overrides lower):
1. Clipboard paste (pasted_base64) - highest precedence
2. File path input (filepath) - overrides image upload
3. Base64 string input (base64_str) - overrides image upload
4. Image upload (image) - lowest precedence
Returns:
Tuple of (image_data_bytes, source_description)
"""
@@ -146,16 +209,16 @@ class ImageLoader:
except Exception as e:
logger.warning(f"Failed to decode pasted image: {e}")
# 2. Medium precedence: File path
# 2. Second precedence: File path input (overrides image upload)
if self._is_valid_input(filepath):
try:
data = self._load_file_data(filepath)
data = self._load_file_data(filepath, use_annotated_path=False)
if data:
return data, f"File: {os.path.basename(filepath)}"
return data, f"File Path: {os.path.basename(filepath)}"
except Exception as e:
logger.warning(f"Failed to load file {filepath}: {e}")
# 3. Lowest precedence: Base64 string input
# 3. Third precedence: Base64 string input (overrides image upload)
if self._is_valid_input(base64_str):
try:
data = self._decode_base64_data(base64_str)
@@ -164,6 +227,15 @@ class ImageLoader:
except Exception as e:
logger.warning(f"Failed to decode base64 string: {e}")
# 4. Lowest precedence: Image upload from file picker
if self._is_valid_input(image):
try:
data = self._load_file_data(image, use_annotated_path=True)
if data:
return data, f"File Upload: {os.path.basename(image)}"
except Exception as e:
logger.warning(f"Failed to load uploaded image {image}: {e}")
return None, "No valid source"
def _is_valid_input(self, value: str) -> bool:
@@ -197,30 +269,46 @@ class ImageLoader:
logger.error(f"Base64 decode error: {e}")
return None
def _load_file_data(self, filepath: str) -> Optional[bytes]:
def _load_file_data(self, filepath: str, use_annotated_path: bool = False) -> Optional[bytes]:
"""
Load image data from a file path.
Args:
filepath: File path (relative paths are resolved to input directory)
use_annotated_path: If True, use ComfyUI's annotated path system (for file picker)
Returns:
File contents as bytes or None if failed
"""
try:
# Resolve relative paths to ComfyUI input directory
if not os.path.isabs(filepath):
filepath = os.path.join(folder_paths.get_input_directory(), filepath)
# Handle empty or None filepath
if not filepath or filepath.strip() == "":
return None
if use_annotated_path:
# Use ComfyUI's annotated filepath system for file picker uploads
full_path = folder_paths.get_annotated_filepath(filepath)
else:
# Handle manual file paths
if not os.path.sep in filepath and not "/" in filepath:
# If filepath doesn't contain path separators, it's likely from the dropdown
# and should be treated as a filename in the input directory
full_path = os.path.join(folder_paths.get_input_directory(), filepath)
elif not os.path.isabs(filepath):
# Resolve relative paths to ComfyUI input directory
full_path = os.path.join(folder_paths.get_input_directory(), filepath)
else:
full_path = filepath
# Validate file exists and is readable
if not os.path.exists(filepath):
raise FileNotFoundError(f"File not found: {filepath}")
if not os.path.exists(full_path):
raise FileNotFoundError(f"File not found: {full_path}")
if not os.path.isfile(filepath):
raise ValueError(f"Path is not a file: {filepath}")
if not os.path.isfile(full_path):
raise ValueError(f"Path is not a file: {full_path}")
# Read file data
with open(filepath, 'rb') as f:
with open(full_path, 'rb') as f:
return f.read()
except Exception as e:
@@ -238,11 +326,11 @@ class ImageLoader:
Tuple of (image_tensor, mask_tensor)
"""
try:
# Open image with PIL
img = Image.open(io.BytesIO(image_data))
# Open image with PIL using node_helpers for better ComfyUI compatibility
img = node_helpers.pillow(Image.open, io.BytesIO(image_data))
# Apply EXIF rotation if present
img = ImageOps.exif_transpose(img)
img = node_helpers.pillow(ImageOps.exif_transpose, img)
# Process RGB image
image_tensor = self._create_image_tensor(img)
@@ -320,6 +408,18 @@ class ImageLoader:
empty_mask = torch.zeros((1, 1, 1), dtype=torch.float32)
return empty_image, empty_mask
def easySave(self, images, filename_prefix, output_type, prompt=None, extra_pnginfo=None):
"""Save or Preview Image"""
if output_type in ["Hide", "None"]:
return list()
elif output_type in ["Preview", "Preview&Choose"]:
filename_prefix = 'easyPreview'
results = PreviewImage().save_images(images, filename_prefix, prompt, extra_pnginfo)
return results['ui']['images']
else:
results = SaveImage().save_images(images, filename_prefix, prompt, extra_pnginfo)
return results['ui']['images']
# Node registration information
+241 -101
View File
@@ -9,7 +9,7 @@
* - Preview functionality
*/
import { app } from "/scripts/app.js";
// Note: app is available globally in ComfyUI context
// Configuration constants
const CONFIG = {
@@ -87,6 +87,7 @@ class ImageLoaderNodeHandler {
* Find and store widget references
*/
findWidgets() {
this.widgets.image = this.node.widgets?.find(w => w.name === "image");
this.widgets.filepath = this.node.widgets?.find(w => w.name === "filepath");
this.widgets.base64 = this.node.widgets?.find(w => w.name === "base64");
this.widgets.pasted = this.node.widgets?.find(w => w.name === "pasted_base64");
@@ -122,6 +123,18 @@ class ImageLoaderNodeHandler {
* Setup event listeners for input handling
*/
setupEventListeners() {
// Image widget (file picker) change handler
if (this.widgets.image) {
const originalCallback = this.widgets.image.callback;
this.widgets.image.callback = (value) => {
if (value && value.trim()) {
this.clearOtherInputs('image');
this.updatePreview('image', value);
}
return originalCallback?.call(this.node, value);
};
}
// Filepath widget change handler
if (this.widgets.filepath) {
const originalCallback = this.widgets.filepath.callback;
@@ -157,108 +170,96 @@ class ImageLoaderNodeHandler {
});
}
// Clipboard paste handler for the entire node
const pasteHandler = (event) => this.handlePaste(event);
// Add paste listener to node element and its container
if (this.node.canvas) {
this.node.canvas.addEventListener('paste', pasteHandler);
this.eventListeners.push({
element: this.node.canvas,
type: 'paste',
handler: pasteHandler
});
}
// Also listen on the node's DOM element if available
if (this.node.element) {
this.node.element.addEventListener('paste', pasteHandler);
this.eventListeners.push({
element: this.node.element,
type: 'paste',
handler: pasteHandler
});
}
// Global paste handler for clipboard images
const globalPasteHandler = (event) => this.handlePaste(event);
document.addEventListener('paste', globalPasteHandler);
this.eventListeners.push({
element: document,
type: 'paste',
handler: globalPasteHandler
});
}
/**
* Setup drag and drop functionality
*/
setupDragAndDrop() {
const nodeElement = this.node.element || this.node.canvas;
// Add drag and drop to the node itself
const nodeElement = this.node;
if (!nodeElement) return;
const dragOverHandler = (event) => {
event.preventDefault();
event.stopPropagation();
// Check if dragged items include images
// Store original handlers to avoid conflicts
const originalOnDragOver = nodeElement.onDragOver;
const originalOnDragLeave = nodeElement.onDragLeave;
const originalOnDrop = nodeElement.onDrop;
nodeElement.onDragOver = (event) => {
if (this.hasImageFiles(event.dataTransfer)) {
event.preventDefault();
event.stopPropagation();
event.dataTransfer.dropEffect = 'copy';
this.setDragState(true);
return true;
}
return originalOnDragOver?.call(nodeElement, event);
};
const dragLeaveHandler = (event) => {
event.preventDefault();
event.stopPropagation();
// Only clear drag state if leaving the node entirely
if (!nodeElement.contains(event.relatedTarget)) {
this.setDragState(false);
}
};
const dropHandler = (event) => {
event.preventDefault();
event.stopPropagation();
nodeElement.onDragLeave = (event) => {
this.setDragState(false);
this.handleDrop(event);
return originalOnDragLeave?.call(nodeElement, event);
};
// Add drag and drop listeners
nodeElement.addEventListener('dragover', dragOverHandler);
nodeElement.addEventListener('dragleave', dragLeaveHandler);
nodeElement.addEventListener('drop', dropHandler);
this.eventListeners.push(
{ element: nodeElement, type: 'dragover', handler: dragOverHandler },
{ element: nodeElement, type: 'dragleave', handler: dragLeaveHandler },
{ element: nodeElement, type: 'drop', handler: dropHandler }
);
nodeElement.onDrop = (event) => {
if (this.hasImageFiles(event.dataTransfer)) {
event.preventDefault();
event.stopPropagation();
this.setDragState(false);
this.handleDrop(event);
return true;
}
return originalOnDrop?.call(nodeElement, event);
};
}
/**
* Setup image preview functionality
*/
setupPreview() {
// Create preview container (initially hidden)
// For ComfyUI nodes, we'll create a preview that appears when hovering or when an image is loaded
// The preview will be positioned relative to the node
this.previewElement = document.createElement('div');
this.previewElement.style.cssText = `
position: absolute;
top: -10px;
right: -10px;
position: fixed;
top: 10px;
right: 10px;
width: ${CONFIG.MAX_PREVIEW_SIZE}px;
max-height: ${CONFIG.MAX_PREVIEW_SIZE}px;
border: 2px solid #4CAF50;
border-radius: 4px;
border-radius: 8px;
background: #fff;
display: none;
z-index: 1000;
z-index: 9999;
overflow: hidden;
box-shadow: 0 4px 12px rgba(0,0,0,0.15);
pointer-events: none;
`;
if (this.node.element) {
this.node.element.style.position = 'relative';
this.node.element.appendChild(this.previewElement);
}
// Add to document body for better positioning
document.body.appendChild(this.previewElement);
// Store reference for cleanup
this.node._imageLoaderPreview = this.previewElement;
}
/**
* Handle clipboard paste events
*/
handlePaste(event) {
// Only handle if this node is focused or selected
if (!this.isNodeActive()) {
return;
}
const items = (event.clipboardData || event.originalEvent?.clipboardData)?.items;
if (!items) return;
@@ -277,6 +278,18 @@ class ImageLoaderNodeHandler {
}
}
/**
* Check if this node is currently active/selected
*/
isNodeActive() {
// Check if the node is selected in the graph
if (this.node.graph && this.node.graph.canvas) {
return this.node.graph.canvas.selected_nodes &&
this.node.graph.canvas.selected_nodes[this.node.id];
}
return false;
}
/**
* Handle drag and drop events
*/
@@ -328,23 +341,54 @@ class ImageLoaderNodeHandler {
}
/**
* Clear inputs other than the specified active one
* Clear inputs other than the specified active one based on precedence rules
*
* Precedence order:
* 1. Clipboard paste (highest) - clears all others
* 2. File path - clears image upload and base64
* 3. Base64 - clears image upload and filepath
* 4. Image upload (lowest) - clears filepath and base64
*/
clearOtherInputs(activeInput) {
const inputs = ['filepath', 'base64', 'paste'];
// Define what each input type should clear
const clearingRules = {
'paste': ['image', 'filepath', 'base64'], // Paste clears everything else
'filepath': ['image', 'base64'], // Filepath clears image upload and base64
'base64': ['image', 'filepath'], // Base64 clears image upload and filepath
'image': ['filepath', 'base64'] // Image upload clears filepath and base64
};
inputs.forEach(input => {
if (input !== activeInput) {
const widget = input === 'paste' ? this.widgets.pasted : this.widgets[input];
if (widget && widget.value !== '') {
widget.value = '';
const inputsToClear = clearingRules[activeInput] || [];
inputsToClear.forEach(inputName => {
const widget = inputName === 'paste' ? this.widgets.pasted : this.widgets[inputName];
if (widget && widget.value !== '') {
const oldValue = widget.value;
widget.value = '';
// Trigger change event to update the node
if (widget.callback) {
widget.callback('');
}
// Log the clearing action for debugging
console.info(`ImageLoader: Cleared ${inputName} (was: "${oldValue.substring(0, 50)}${oldValue.length > 50 ? '...' : ''}")`);
}
});
// Hide preview if clearing
// Update preview for new active input
this.hidePreview();
if (activeInput !== 'paste') {
this.hidePreview();
// For non-paste inputs, update preview after a short delay
setTimeout(() => {
if (activeInput === 'image' && this.widgets.image?.value) {
this.updatePreview('image', this.widgets.image.value);
} else if (activeInput === 'filepath' && this.widgets.filepath?.value) {
this.updatePreview('filepath', this.widgets.filepath.value);
} else if (activeInput === 'base64' && this.widgets.base64?.value) {
this.updatePreview('base64', this.widgets.base64.value);
}
}, 100);
}
}
@@ -382,24 +426,78 @@ class ImageLoaderNodeHandler {
let imageUrl = '';
if (source === 'filepath') {
// For file paths, we can't show preview without loading the file
this.hidePreview();
return;
if (source === 'filepath' || source === 'image') {
// For file paths, try to create a preview URL
if (value && value.trim()) {
// Try to load the image for preview
this.loadImagePreview(value);
return;
}
} else if (source === 'base64' || source === 'paste') {
imageUrl = value.startsWith('data:') ? value : `data:image/png;base64,${value}`;
}
if (imageUrl) {
this.previewElement.innerHTML = `
<img src="${imageUrl}"
style="width: 100%; height: auto; max-height: ${CONFIG.MAX_PREVIEW_SIZE}px; object-fit: contain;"
alt="Preview" />
`;
this.previewElement.style.display = 'block';
this.showPreview(imageUrl);
} else {
this.hidePreview();
}
}
/**
* Load image preview from file path
*/
async loadImagePreview(filePath) {
try {
// For ComfyUI, we can try to access the image through the input directory
const imageUrl = `/view?filename=${encodeURIComponent(filePath)}&type=input`;
// Test if the image loads
const img = new Image();
img.onload = () => {
this.showPreview(imageUrl);
};
img.onerror = () => {
// If direct access fails, show a placeholder
this.showPreviewPlaceholder(filePath);
};
img.src = imageUrl;
} catch (error) {
this.showPreviewPlaceholder(filePath);
}
}
/**
* Show image preview
*/
showPreview(imageUrl) {
if (!this.previewElement) return;
this.previewElement.innerHTML = `
<img src="${imageUrl}"
style="width: 100%; height: auto; max-height: ${CONFIG.MAX_PREVIEW_SIZE}px; object-fit: contain; display: block;"
alt="Preview"
onerror="this.parentElement.innerHTML='<div style=\\'padding: 10px; text-align: center; color: #666;\\'>Preview not available</div>'" />
`;
this.previewElement.style.display = 'block';
}
/**
* Show preview placeholder for file paths
*/
showPreviewPlaceholder(fileName) {
if (!this.previewElement) return;
const baseName = fileName.split(/[\\/]/).pop() || fileName;
this.previewElement.innerHTML = `
<div style="padding: 10px; text-align: center; color: #666; font-size: 12px; background: #f5f5f5; border-radius: 4px;">
📁 ${baseName}<br>
<small>File selected</small>
</div>
`;
this.previewElement.style.display = 'block';
}
/**
* Hide image preview
*/
@@ -416,28 +514,44 @@ class ImageLoaderNodeHandler {
const indicator = document.createElement('div');
indicator.textContent = '📋 Image Pasted!';
indicator.style.cssText = `
position: absolute;
position: fixed;
top: 50%;
left: 50%;
transform: translate(-50%, -50%);
background: #4CAF50;
color: white;
padding: 8px 16px;
border-radius: 4px;
font-size: 12px;
z-index: 1001;
padding: 12px 24px;
border-radius: 8px;
font-size: 14px;
font-weight: bold;
z-index: 10000;
pointer-events: none;
box-shadow: 0 4px 12px rgba(0,0,0,0.3);
animation: fadeInOut 2s ease-in-out;
`;
if (this.node.element) {
this.node.element.appendChild(indicator);
setTimeout(() => {
if (indicator.parentNode) {
indicator.parentNode.removeChild(indicator);
// Add animation keyframes if not already present
if (!document.getElementById('pasteIndicatorStyles')) {
const style = document.createElement('style');
style.id = 'pasteIndicatorStyles';
style.textContent = `
@keyframes fadeInOut {
0% { opacity: 0; transform: translate(-50%, -50%) scale(0.8); }
20% { opacity: 1; transform: translate(-50%, -50%) scale(1); }
80% { opacity: 1; transform: translate(-50%, -50%) scale(1); }
100% { opacity: 0; transform: translate(-50%, -50%) scale(0.8); }
}
}, CONFIG.PASTE_INDICATOR_DURATION);
`;
document.head.appendChild(style);
}
document.body.appendChild(indicator);
setTimeout(() => {
if (indicator.parentNode) {
indicator.parentNode.removeChild(indicator);
}
}, CONFIG.PASTE_INDICATOR_DURATION);
}
/**
@@ -446,13 +560,34 @@ class ImageLoaderNodeHandler {
setDragState(isDragging) {
this.isDragging = isDragging;
if (this.node.element) {
if (this.node) {
if (isDragging) {
this.node.element.style.outline = '2px dashed #4CAF50';
this.node.element.style.backgroundColor = 'rgba(76, 175, 80, 0.1)';
// Store original colors
this.originalColors = {
bgcolor: this.node.bgcolor,
color: this.node.color
};
// Set drag state colors
this.node.bgcolor = "rgba(76, 175, 80, 0.2)";
this.node.color = "#4CAF50";
// Force redraw
if (this.node.graph && this.node.graph.canvas) {
this.node.graph.canvas.draw(true, true);
}
} else {
this.node.element.style.outline = '';
this.node.element.style.backgroundColor = '';
// Restore original colors
if (this.originalColors) {
this.node.bgcolor = this.originalColors.bgcolor;
this.node.color = this.originalColors.color;
this.originalColors = null;
}
// Force redraw
if (this.node.graph && this.node.graph.canvas) {
this.node.graph.canvas.draw(true, true);
}
}
}
}
@@ -486,13 +621,18 @@ class ImageLoaderNodeHandler {
});
this.eventListeners = [];
// Remove preview element
// Remove preview element from document body
if (this.previewElement?.parentNode) {
this.previewElement.parentNode.removeChild(this.previewElement);
}
// Clear drag state
this.setDragState(false);
// Clear node reference
if (this.node._imageLoaderPreview) {
delete this.node._imageLoaderPreview;
}
}
}
-145
View File
@@ -1,145 +0,0 @@
"""
Basic tests for ComfyUI Universal Image Loader
These tests can be run independently or as part of a larger test suite.
"""
import unittest
import torch
import numpy as np
import base64
import io
from PIL import Image
# Import the node (adjust path as needed)
try:
from imgloader_node import ImageLoader
except ImportError:
# If running from different directory
import sys
import os
sys.path.append(os.path.dirname(os.path.abspath(__file__)))
from imgloader_node import ImageLoader
class TestImageLoader(unittest.TestCase):
"""Test cases for the ImageLoader node."""
def setUp(self):
"""Set up test fixtures."""
self.node = ImageLoader()
# Create a simple test image
self.test_image = Image.new('RGB', (100, 100), color='red')
# Convert to base64
buffer = io.BytesIO()
self.test_image.save(buffer, format='PNG')
self.test_base64 = base64.b64encode(buffer.getvalue()).decode('utf-8')
self.test_data_url = f"data:image/png;base64,{self.test_base64}"
def test_input_types(self):
"""Test that input types are properly defined."""
input_types = ImageLoader.INPUT_TYPES()
self.assertIn('optional', input_types)
self.assertIn('filepath', input_types['optional'])
self.assertIn('base64', input_types['optional'])
self.assertIn('pasted_base64', input_types['optional'])
def test_base64_loading(self):
"""Test loading from base64 string."""
image_tensor, mask_tensor = self.node.load_image(base64=self.test_base64)
# Check tensor properties
self.assertIsInstance(image_tensor, torch.Tensor)
self.assertIsInstance(mask_tensor, torch.Tensor)
self.assertEqual(len(image_tensor.shape), 4) # NHWC format
self.assertEqual(len(mask_tensor.shape), 3) # NHW format
# Check tensor values are in valid range
self.assertTrue(torch.all(image_tensor >= 0))
self.assertTrue(torch.all(image_tensor <= 1))
self.assertTrue(torch.all(mask_tensor >= 0))
self.assertTrue(torch.all(mask_tensor <= 1))
def test_data_url_loading(self):
"""Test loading from data URL format."""
image_tensor, mask_tensor = self.node.load_image(base64=self.test_data_url)
# Should successfully load
self.assertEqual(image_tensor.shape[0], 1) # Batch size 1
self.assertEqual(image_tensor.shape[3], 3) # RGB channels
def test_pasted_base64_precedence(self):
"""Test that pasted base64 takes precedence over other inputs."""
# Provide all three inputs
image_tensor, mask_tensor = self.node.load_image(
filepath="nonexistent.png",
base64="invalid_base64",
pasted_base64=self.test_data_url
)
# Should use pasted_base64 and succeed
self.assertEqual(image_tensor.shape[0], 1)
self.assertEqual(image_tensor.shape[3], 3)
def test_empty_inputs(self):
"""Test behavior with empty inputs."""
image_tensor, mask_tensor = self.node.load_image()
# Should return empty tensors
self.assertEqual(image_tensor.shape, (1, 1, 1, 3))
self.assertEqual(mask_tensor.shape, (1, 1, 1))
def test_invalid_base64(self):
"""Test handling of invalid base64."""
image_tensor, mask_tensor = self.node.load_image(base64="invalid_base64")
# Should fallback to empty tensors
self.assertEqual(image_tensor.shape, (1, 1, 1, 3))
self.assertEqual(mask_tensor.shape, (1, 1, 1))
def test_is_changed_method(self):
"""Test the IS_CHANGED method."""
# Same inputs should return same hash
hash1 = ImageLoader.IS_CHANGED(filepath="test.png", base64="")
hash2 = ImageLoader.IS_CHANGED(filepath="test.png", base64="")
self.assertEqual(hash1, hash2)
# Different inputs should return different hash
hash3 = ImageLoader.IS_CHANGED(filepath="test2.png", base64="")
self.assertNotEqual(hash1, hash3)
def test_alpha_channel_handling(self):
"""Test proper alpha channel extraction."""
# Create RGBA test image
rgba_image = Image.new('RGBA', (50, 50), color=(255, 0, 0, 128))
buffer = io.BytesIO()
rgba_image.save(buffer, format='PNG')
rgba_base64 = base64.b64encode(buffer.getvalue()).decode('utf-8')
image_tensor, mask_tensor = self.node.load_image(base64=rgba_base64)
# Check that alpha was extracted
self.assertEqual(mask_tensor.shape, (1, 50, 50))
# Alpha should be approximately 0.5 (128/255)
self.assertTrue(torch.allclose(mask_tensor, torch.tensor(128/255), atol=0.01))
class TestImageLoaderIntegration(unittest.TestCase):
"""Integration tests for the ImageLoader node."""
def test_node_registration(self):
"""Test that node registration exports are correct."""
from __init__ import NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS, WEB_DIRECTORY
self.assertIn("ImageLoader", NODE_CLASS_MAPPINGS)
self.assertEqual(NODE_CLASS_MAPPINGS["ImageLoader"], ImageLoader)
self.assertIn("ImageLoader", NODE_DISPLAY_NAME_MAPPINGS)
self.assertEqual(WEB_DIRECTORY, "./js")
if __name__ == '__main__':
# Run tests
unittest.main(verbosity=2)