From 9a599a4dfcd775993c16faaa4485bfaefd4a09cf Mon Sep 17 00:00:00 2001 From: negaga53 Date: Mon, 7 Jul 2025 23:11:12 +0200 Subject: [PATCH] Stable --- imgloader_node.py | 152 +++++++++++++++++---- js/imgloader.js | 342 ++++++++++++++++++++++++++++++++-------------- test_imgloader.py | 145 -------------------- 3 files changed, 367 insertions(+), 272 deletions(-) delete mode 100644 test_imgloader.py diff --git a/imgloader_node.py b/imgloader_node.py index d9c8fae..372b260 100644 --- a/imgloader_node.py +++ b/imgloader_node.py @@ -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 diff --git a/js/imgloader.js b/js/imgloader.js index c968a27..a49de81 100644 --- a/js/imgloader.js +++ b/js/imgloader.js @@ -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 = ` - 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 = ` + Preview + `; + 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 = ` +
+ 📁 ${baseName}
+ File selected +
+ `; + 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; + } } } diff --git a/test_imgloader.py b/test_imgloader.py deleted file mode 100644 index b07899c..0000000 --- a/test_imgloader.py +++ /dev/null @@ -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)