From de1fe4656393705ec4d493a5c4d2a6be8ce5a375 Mon Sep 17 00:00:00 2001 From: Vito Sansevero Date: Sat, 24 May 2025 12:26:31 -0700 Subject: [PATCH] feat(utils): add utility modules for diagnostics, hashing, image monitoring, metadata extraction, prompt tracking, and validation --- utils/__init__.py | 6 + utils/diagnostics.py | 314 ++++++++++++++++++++++++++++++++++++ utils/hashing.py | 73 +++++++++ utils/image_monitor.py | 229 ++++++++++++++++++++++++++ utils/metadata_extractor.py | 264 ++++++++++++++++++++++++++++++ utils/prompt_tracker.py | 246 ++++++++++++++++++++++++++++ utils/validators.py | 221 +++++++++++++++++++++++++ 7 files changed, 1353 insertions(+) create mode 100644 utils/__init__.py create mode 100644 utils/diagnostics.py create mode 100644 utils/hashing.py create mode 100644 utils/image_monitor.py create mode 100644 utils/metadata_extractor.py create mode 100644 utils/prompt_tracker.py create mode 100644 utils/validators.py diff --git a/utils/__init__.py b/utils/__init__.py new file mode 100644 index 0000000..0d281ca --- /dev/null +++ b/utils/__init__.py @@ -0,0 +1,6 @@ +"""Utility functions for KikoTextEncode.""" + +from .hashing import generate_prompt_hash +from .validators import validate_prompt_text, validate_rating, validate_tags + +__all__ = ["generate_prompt_hash", "validate_prompt_text", "validate_rating", "validate_tags"] \ No newline at end of file diff --git a/utils/diagnostics.py b/utils/diagnostics.py new file mode 100644 index 0000000..50d1a64 --- /dev/null +++ b/utils/diagnostics.py @@ -0,0 +1,314 @@ +""" +Diagnostic utilities for troubleshooting the PromptManager gallery system. +""" + +import os +import sqlite3 +from pathlib import Path +from typing import Dict, Any, List + + +class GalleryDiagnostics: + """Diagnostics for the gallery system.""" + + def __init__(self, db_path: str = "prompts.db"): + self.db_path = db_path + + def run_full_diagnostic(self) -> Dict[str, Any]: + """Run a complete diagnostic check.""" + print("\n" + "="*60) + print("๐Ÿ” PROMPTMANAGER GALLERY DIAGNOSTICS") + print("="*60) + + results = { + 'database': self.check_database(), + 'images_table': self.check_images_table(), + 'file_system': self.check_file_system(), + 'comfyui_output': self.check_comfyui_output(), + 'dependencies': self.check_dependencies() + } + + print("\n" + "="*60) + print("๐Ÿ“‹ DIAGNOSTIC SUMMARY") + print("="*60) + + for category, result in results.items(): + status = "โœ… PASS" if result['status'] == 'ok' else "โŒ FAIL" + print(f"{category.upper():<20} {status}") + if result['status'] != 'ok': + print(f" Issue: {result['message']}") + + print("\n" + "="*60) + return results + + def check_database(self) -> Dict[str, Any]: + """Check database connection and structure.""" + print("\n๐Ÿ—„๏ธ Checking Database...") + + try: + if not os.path.exists(self.db_path): + return { + 'status': 'error', + 'message': f'Database file not found: {self.db_path}' + } + + with sqlite3.connect(self.db_path) as conn: + conn.row_factory = sqlite3.Row + + # Check prompts table + cursor = conn.execute("SELECT COUNT(*) as count FROM prompts") + prompt_count = cursor.fetchone()['count'] + print(f" ๐Ÿ“ Prompts in database: {prompt_count}") + + # Check if generated_images table exists + cursor = conn.execute(""" + SELECT name FROM sqlite_master + WHERE type='table' AND name='generated_images' + """) + + has_images_table = cursor.fetchone() is not None + print(f" ๐Ÿ–ผ๏ธ Images table exists: {has_images_table}") + + return { + 'status': 'ok', + 'prompt_count': prompt_count, + 'has_images_table': has_images_table + } + + except Exception as e: + return { + 'status': 'error', + 'message': f'Database error: {str(e)}' + } + + def check_images_table(self) -> Dict[str, Any]: + """Check the generated_images table specifically.""" + print("\n๐Ÿ–ผ๏ธ Checking Images Table...") + + try: + with sqlite3.connect(self.db_path) as conn: + conn.row_factory = sqlite3.Row + + # Check if table exists + cursor = conn.execute(""" + SELECT name FROM sqlite_master + WHERE type='table' AND name='generated_images' + """) + + if not cursor.fetchone(): + return { + 'status': 'error', + 'message': 'generated_images table does not exist - run the updated code to create it' + } + + # Check image records + cursor = conn.execute("SELECT COUNT(*) as count FROM generated_images") + image_count = cursor.fetchone()['count'] + print(f" ๐Ÿ“Š Images in database: {image_count}") + + # Get recent images + cursor = conn.execute(""" + SELECT gi.*, p.text + FROM generated_images gi + LEFT JOIN prompts p ON gi.prompt_id = p.id + ORDER BY gi.generation_time DESC + LIMIT 5 + """) + recent_images = [dict(row) for row in cursor.fetchall()] + + print(f" ๐Ÿ•’ Recent images: {len(recent_images)}") + for img in recent_images: + print(f" - {img['filename']} -> Prompt {img['prompt_id']}") + + return { + 'status': 'ok', + 'image_count': image_count, + 'recent_images': recent_images + } + + except Exception as e: + return { + 'status': 'error', + 'message': f'Images table error: {str(e)}' + } + + def check_file_system(self) -> Dict[str, Any]: + """Check file system and permissions.""" + print("\n๐Ÿ“ Checking File System...") + + try: + # Check current directory + current_dir = os.getcwd() + print(f" ๐Ÿ“‚ Current directory: {current_dir}") + + # Check if we can write to current directory + test_file = "test_write.tmp" + try: + with open(test_file, 'w') as f: + f.write("test") + os.remove(test_file) + can_write = True + except: + can_write = False + + print(f" โœ๏ธ Can write to directory: {can_write}") + + return { + 'status': 'ok', + 'current_dir': current_dir, + 'can_write': can_write + } + + except Exception as e: + return { + 'status': 'error', + 'message': f'File system error: {str(e)}' + } + + def check_comfyui_output(self) -> Dict[str, Any]: + """Check ComfyUI output directories.""" + print("\n๐ŸŽจ Checking ComfyUI Output...") + + output_dirs = [] + + # Try to detect ComfyUI output directories + potential_dirs = [ + "output", + "../output", + "../../output", + "ComfyUI/output", + "../ComfyUI/output", + "../../ComfyUI/output" + ] + + for dir_path in potential_dirs: + abs_path = os.path.abspath(dir_path) + if os.path.exists(abs_path): + output_dirs.append(abs_path) + print(f" ๐Ÿ“ Found output dir: {abs_path}") + + # Count images in this directory + try: + image_files = [] + for ext in ['.png', '.jpg', '.jpeg', '.webp']: + image_files.extend(Path(abs_path).rglob(f'*{ext}')) + print(f" ๐Ÿ–ผ๏ธ Images found: {len(image_files)}") + except Exception as e: + print(f" โŒ Error scanning: {e}") + + # Try ComfyUI's folder_paths + try: + import folder_paths + comfyui_output = folder_paths.get_output_directory() + if comfyui_output and comfyui_output not in output_dirs: + output_dirs.append(comfyui_output) + print(f" ๐Ÿ“ ComfyUI output dir: {comfyui_output}") + except ImportError: + print(" โš ๏ธ ComfyUI folder_paths not available") + + return { + 'status': 'ok' if output_dirs else 'warning', + 'message': 'No output directories found' if not output_dirs else None, + 'output_dirs': output_dirs + } + + def check_dependencies(self) -> Dict[str, Any]: + """Check required dependencies.""" + print("\n๐Ÿ“ฆ Checking Dependencies...") + + dependencies = { + 'watchdog': False, + 'PIL': False, + 'sqlite3': False + } + + # Check watchdog + try: + import watchdog + dependencies['watchdog'] = True + print(f" โœ… watchdog: {watchdog.__version__}") + except ImportError: + print(" โŒ watchdog: NOT INSTALLED") + + # Check PIL + try: + from PIL import Image + dependencies['PIL'] = True + print(f" โœ… PIL (Pillow): Available") + except ImportError: + print(" โŒ PIL (Pillow): NOT AVAILABLE") + + # Check sqlite3 + try: + import sqlite3 + dependencies['sqlite3'] = True + print(f" โœ… sqlite3: {sqlite3.sqlite_version}") + except ImportError: + print(" โŒ sqlite3: NOT AVAILABLE") + + all_deps_ok = all(dependencies.values()) + + return { + 'status': 'ok' if all_deps_ok else 'error', + 'message': 'Missing dependencies' if not all_deps_ok else None, + 'dependencies': dependencies + } + + def create_test_image_link(self, prompt_id: int, test_image_path: str = None) -> Dict[str, Any]: + """Create a test image link to verify the system works.""" + print(f"\n๐Ÿงช Creating test image link for prompt {prompt_id}...") + + try: + # Import database operations + import sys + import os + sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__)))) + from database.operations import PromptDatabase + + db = PromptDatabase() + + # Create a fake image path if none provided + if not test_image_path: + test_image_path = "/fake/test/image.png" + + # Create test metadata + test_metadata = { + 'file_info': { + 'size': 1024000, + 'dimensions': [512, 512], + 'format': 'PNG' + }, + 'workflow': {'test': True}, + 'prompt': {'test_prompt': 'This is a test'} + } + + # Link the test image + image_id = db.link_image_to_prompt( + prompt_id=str(prompt_id), + image_path=test_image_path, + metadata=test_metadata + ) + + print(f" โœ… Test image linked with ID: {image_id}") + + return { + 'status': 'ok', + 'image_id': image_id, + 'message': f'Test image linked successfully with ID {image_id}' + } + + except Exception as e: + return { + 'status': 'error', + 'message': f'Failed to create test link: {str(e)}' + } + + +def run_diagnostics(): + """Run diagnostics from command line.""" + diagnostics = GalleryDiagnostics() + return diagnostics.run_full_diagnostic() + + +if __name__ == "__main__": + run_diagnostics() \ No newline at end of file diff --git a/utils/hashing.py b/utils/hashing.py new file mode 100644 index 0000000..881b0e3 --- /dev/null +++ b/utils/hashing.py @@ -0,0 +1,73 @@ +""" +Hashing utilities for KikoTextEncode prompt deduplication. +""" + +import hashlib + + +def generate_prompt_hash(text: str) -> str: + """ + Generate a SHA256 hash for prompt text to enable deduplication. + + Args: + text: The prompt text to hash + + Returns: + str: SHA256 hexdigest of the text + + Example: + >>> generate_prompt_hash("beautiful landscape") + 'a1b2c3d4e5f6...' + """ + if not isinstance(text, str): + raise TypeError("Text must be a string") + + # Normalize the text by stripping whitespace and converting to lowercase + # for consistent hashing regardless of minor formatting differences + normalized_text = text.strip().lower() + + return hashlib.sha256(normalized_text.encode('utf-8')).hexdigest() + + +def generate_content_hash(content: dict) -> str: + """ + Generate a hash for prompt content including metadata. + + Args: + content: Dictionary containing prompt data + + Returns: + str: SHA256 hexdigest of the content + """ + import json + + # Create a normalized representation of the content + normalized = { + 'text': content.get('text', '').strip().lower(), + 'category': content.get('category', '').strip().lower() if content.get('category') else '', + 'tags': sorted([tag.strip().lower() for tag in content.get('tags', []) if tag.strip()]), + 'workflow_name': content.get('workflow_name', '').strip().lower() if content.get('workflow_name') else '' + } + + # Convert to JSON string for consistent hashing + content_str = json.dumps(normalized, sort_keys=True) + + return hashlib.sha256(content_str.encode('utf-8')).hexdigest() + + +def is_duplicate_prompt(text1: str, text2: str, threshold: float = 0.95) -> bool: + """ + Check if two prompts are likely duplicates using hash comparison. + + Args: + text1: First prompt text + text2: Second prompt text + threshold: Similarity threshold (not used for exact hash matching) + + Returns: + bool: True if prompts are likely duplicates + """ + hash1 = generate_prompt_hash(text1) + hash2 = generate_prompt_hash(text2) + + return hash1 == hash2 \ No newline at end of file diff --git a/utils/image_monitor.py b/utils/image_monitor.py new file mode 100644 index 0000000..5b8f698 --- /dev/null +++ b/utils/image_monitor.py @@ -0,0 +1,229 @@ +""" +Image monitoring system for ComfyUI generated images. +Automatically detects new images and links them to prompts. +""" + +import os +import time +import threading +import json +from pathlib import Path +from typing import Optional, Dict, Any, Callable +from watchdog.observers import Observer +from watchdog.events import FileSystemEventHandler + +from .metadata_extractor import ComfyUIMetadataExtractor + + +class ImageGenerationHandler(FileSystemEventHandler): + """Handler for new image file creation events.""" + + def __init__(self, db_manager, prompt_tracker): + """ + Initialize the image handler. + + Args: + db_manager: Database manager instance + prompt_tracker: Prompt tracking instance + """ + self.db_manager = db_manager + self.prompt_tracker = prompt_tracker + self.metadata_extractor = ComfyUIMetadataExtractor() + self.processing_delay = 2.0 # Wait 2 seconds before processing + + def on_created(self, event): + """Handle file creation events.""" + if not event.is_directory and self.is_image_file(event.src_path): + print(f"[PromptManager] New image detected: {event.src_path}") + # Small delay to ensure file is fully written + threading.Timer( + self.processing_delay, + self.process_new_image, + args=[event.src_path] + ).start() + + def is_image_file(self, filepath: str) -> bool: + """Check if file is a supported image format.""" + return filepath.lower().endswith(('.png', '.jpg', '.jpeg', '.webp', '.gif')) + + def process_new_image(self, image_path: str): + """Process a newly created image file.""" + try: + print(f"[PromptManager] Processing image: {image_path}") + + if not os.path.exists(image_path): + print(f"[PromptManager] Image file no longer exists: {image_path}") + return + + # Get current prompt context first + current_prompt = self.prompt_tracker.get_current_prompt() + print(f"[PromptManager] Current prompt context: {current_prompt['id'] if current_prompt else 'None'}") + + if not current_prompt: + print(f"[PromptManager] No active prompt context for image: {image_path}") + return + + # Extract ComfyUI metadata + try: + metadata = self.metadata_extractor.extract_metadata(image_path) + print(f"[PromptManager] Extracted metadata: {bool(metadata)}") + except Exception as meta_error: + print(f"[PromptManager] Metadata extraction failed: {meta_error}") + metadata = None + + if metadata: + print(f"[PromptManager] Linking image with full metadata to prompt {current_prompt['id']}") + self.link_image_to_prompt(image_path, current_prompt, metadata) + else: + print(f"[PromptManager] Linking image with basic info to prompt {current_prompt['id']}") + # Link with basic file info even without metadata + basic_metadata = self.get_basic_file_info(image_path) + self.link_image_to_prompt(image_path, current_prompt, {'file_info': basic_metadata}) + + except Exception as e: + print(f"[PromptManager] Error processing image {image_path}: {e}") + import traceback + traceback.print_exc() + + def get_basic_file_info(self, image_path: str) -> Dict[str, Any]: + """Get basic file information when metadata extraction fails.""" + try: + from PIL import Image + + stat = os.stat(image_path) + file_info = { + 'size': stat.st_size, + 'format': None, + 'dimensions': None + } + + # Try to get image dimensions + try: + with Image.open(image_path) as img: + file_info['dimensions'] = list(img.size) + file_info['format'] = img.format + except Exception: + pass + + return file_info + except Exception as e: + print(f"[PromptManager] Error getting file info: {e}") + return {} + + def link_image_to_prompt(self, image_path: str, prompt_context: Dict, metadata: Dict): + """Link an image to a prompt in the database.""" + try: + image_id = self.db_manager.link_image_to_prompt( + prompt_id=prompt_context['id'], + image_path=image_path, + metadata=metadata + ) + print(f"[PromptManager] Successfully linked image {image_id} to prompt {prompt_context['id']}") + except Exception as e: + print(f"[PromptManager] Failed to link image to prompt: {e}") + + +class ImageMonitor: + """Main image monitoring system.""" + + def __init__(self, db_manager, prompt_tracker): + """ + Initialize the image monitor. + + Args: + db_manager: Database manager instance + prompt_tracker: Prompt tracking instance + """ + self.db_manager = db_manager + self.prompt_tracker = prompt_tracker + self.observer = None + self.handler = None + self.monitored_directories = [] + + def start_monitoring(self, output_directories: Optional[list] = None): + """ + Start monitoring ComfyUI output directories. + + Args: + output_directories: List of directories to monitor. If None, auto-detect. + """ + if self.observer: + print("[PromptManager] Image monitoring already running") + return + + # Auto-detect ComfyUI output directory if none provided + if not output_directories: + output_directories = self.detect_comfyui_output_dirs() + + if not output_directories: + print("[PromptManager] No output directories found to monitor") + return + + # Create event handler + self.handler = ImageGenerationHandler(self.db_manager, self.prompt_tracker) + + # Start observer + self.observer = Observer() + + for output_dir in output_directories: + if os.path.exists(output_dir): + self.observer.schedule(self.handler, output_dir, recursive=True) + self.monitored_directories.append(output_dir) + print(f"[PromptManager] Monitoring directory: {output_dir}") + else: + print(f"[PromptManager] Directory does not exist: {output_dir}") + + if self.monitored_directories: + self.observer.start() + print(f"[PromptManager] Image monitoring started for {len(self.monitored_directories)} directories") + else: + print("[PromptManager] No valid directories to monitor") + + def stop_monitoring(self): + """Stop the image monitoring system.""" + if self.observer: + self.observer.stop() + self.observer.join() + self.observer = None + self.handler = None + self.monitored_directories = [] + print("[PromptManager] Image monitoring stopped") + + def detect_comfyui_output_dirs(self) -> list: + """Auto-detect ComfyUI output directories.""" + potential_dirs = [] + + try: + # Try to import ComfyUI's folder_paths + import folder_paths + output_dir = folder_paths.get_output_directory() + if output_dir and os.path.exists(output_dir): + potential_dirs.append(output_dir) + print(f"[PromptManager] Detected ComfyUI output directory: {output_dir}") + except ImportError: + print("[PromptManager] ComfyUI folder_paths not available, using fallback detection") + + # Fallback: Look for common ComfyUI directory structures + fallback_paths = [ + "output", + "../output", + "../../output", + "ComfyUI/output", + "../ComfyUI/output" + ] + + for path in fallback_paths: + abs_path = os.path.abspath(path) + if os.path.exists(abs_path) and abs_path not in potential_dirs: + potential_dirs.append(abs_path) + print(f"[PromptManager] Found output directory: {abs_path}") + + return potential_dirs + + def get_status(self) -> Dict[str, Any]: + """Get monitoring status information.""" + return { + 'running': self.observer is not None, + 'monitored_directories': self.monitored_directories, + 'handler_active': self.handler is not None + } \ No newline at end of file diff --git a/utils/metadata_extractor.py b/utils/metadata_extractor.py new file mode 100644 index 0000000..c01bf2c --- /dev/null +++ b/utils/metadata_extractor.py @@ -0,0 +1,264 @@ +""" +ComfyUI metadata extraction utilities. +Extracts workflow and prompt information from generated images. +""" + +import os +import json +from typing import Optional, Dict, Any +from PIL import Image +from PIL.PngImagePlugin import PngInfo + + +class ComfyUIMetadataExtractor: + """Extracts ComfyUI metadata from generated images.""" + + def __init__(self): + """Initialize the metadata extractor.""" + pass + + def extract_metadata(self, image_path: str) -> Optional[Dict[str, Any]]: + """ + Extract ComfyUI workflow and prompt metadata from an image. + + Args: + image_path: Path to the image file + + Returns: + Dictionary containing extracted metadata or None if extraction fails + """ + try: + with Image.open(image_path) as image: + metadata = {} + + # Add basic file information + metadata['file_info'] = self.get_file_info(image_path, image) + + # Extract ComfyUI-specific metadata from PNG text chunks + if hasattr(image, 'text') and image.text: + # Look for ComfyUI workflow data + if 'workflow' in image.text: + try: + workflow_data = json.loads(image.text['workflow']) + metadata['workflow'] = workflow_data + + # Extract text encoder nodes from workflow + text_encoder_nodes = self.find_text_encoder_nodes(workflow_data) + if text_encoder_nodes: + metadata['text_encoder_nodes'] = text_encoder_nodes + + except json.JSONDecodeError as e: + print(f"[PromptManager] Failed to parse workflow JSON: {e}") + + # Look for prompt data + if 'prompt' in image.text: + try: + prompt_data = json.loads(image.text['prompt']) + metadata['prompt'] = prompt_data + except json.JSONDecodeError as e: + print(f"[PromptManager] Failed to parse prompt JSON: {e}") + + # Extract other common metadata fields + metadata_fields = [ + 'parameters', 'model', 'sampler', 'steps', 'cfg_scale', + 'seed', 'scheduler', 'positive', 'negative' + ] + + for field in metadata_fields: + if field in image.text: + try: + # Try to parse as JSON first + metadata[field] = json.loads(image.text[field]) + except (json.JSONDecodeError, TypeError): + # Store as string if not valid JSON + metadata[field] = image.text[field] + + return metadata if any(key != 'file_info' for key in metadata.keys()) else None + + except Exception as e: + print(f"[PromptManager] Error extracting metadata from {image_path}: {e}") + return None + + def get_file_info(self, image_path: str, image: Image.Image) -> Dict[str, Any]: + """ + Get basic file information. + + Args: + image_path: Path to the image file + image: PIL Image object + + Returns: + Dictionary containing file information + """ + try: + stat = os.stat(image_path) + return { + 'size': stat.st_size, + 'dimensions': list(image.size), + 'format': image.format, + 'mode': image.mode, + 'created_time': stat.st_ctime, + 'modified_time': stat.st_mtime + } + except Exception as e: + print(f"[PromptManager] Error getting file info: {e}") + return {} + + def find_text_encoder_nodes(self, workflow_data: Dict) -> list: + """ + Find text encoder nodes in the workflow data. + + Args: + workflow_data: ComfyUI workflow data + + Returns: + List of text encoder node data + """ + text_encoder_nodes = [] + + if not isinstance(workflow_data, dict): + return text_encoder_nodes + + # Check different possible workflow structures + nodes_data = None + + # Try different keys where nodes might be stored + if 'nodes' in workflow_data: + nodes_data = workflow_data['nodes'] + elif 'workflow' in workflow_data and 'nodes' in workflow_data['workflow']: + nodes_data = workflow_data['workflow']['nodes'] + elif isinstance(workflow_data, dict): + # Sometimes the workflow data is just a flat dict of node IDs + nodes_data = workflow_data + + if not nodes_data: + return text_encoder_nodes + + # Handle different node data structures + if isinstance(nodes_data, list): + # Nodes as a list + for node in nodes_data: + if self.is_text_encoder_node(node): + text_encoder_nodes.append(node) + elif isinstance(nodes_data, dict): + # Nodes as a dictionary (node_id -> node_data) + for node_id, node_data in nodes_data.items(): + if self.is_text_encoder_node(node_data): + node_data['node_id'] = node_id + text_encoder_nodes.append(node_data) + + return text_encoder_nodes + + def is_text_encoder_node(self, node_data: Any) -> bool: + """ + Check if a node is a text encoder node. + + Args: + node_data: Node data to check + + Returns: + True if the node is a text encoder + """ + if not isinstance(node_data, dict): + return False + + # Check for common text encoder node types + text_encoder_types = [ + 'CLIPTextEncode', + 'CLIPTextEncodeSDXL', + 'CLIPTextEncodeSDXLRefiner', + 'PromptManager', # Our custom node + 'BNK_CLIPTextEncoder', + 'Text Encoder', + 'CLIP Text Encode' + ] + + # Check node type/class_type + node_type = node_data.get('type') or node_data.get('class_type') or '' + + for encoder_type in text_encoder_types: + if encoder_type.lower() in node_type.lower(): + return True + + # Check node title/name for text encoding keywords + node_title = (node_data.get('title') or node_data.get('name') or '').lower() + text_keywords = ['text', 'prompt', 'encode', 'clip'] + + if any(keyword in node_title for keyword in text_keywords): + return True + + return False + + def extract_prompt_text_from_workflow(self, workflow_data: Dict) -> Optional[str]: + """ + Extract the actual prompt text from workflow data. + + Args: + workflow_data: ComfyUI workflow data + + Returns: + Extracted prompt text or None + """ + text_encoder_nodes = self.find_text_encoder_nodes(workflow_data) + + for node in text_encoder_nodes: + # Try different ways to get the prompt text + inputs = node.get('inputs', {}) + + # Common input field names for prompt text + text_fields = ['text', 'prompt', 'positive', 'conditioning'] + + for field in text_fields: + if field in inputs and inputs[field]: + if isinstance(inputs[field], str): + return inputs[field] + elif isinstance(inputs[field], list) and inputs[field]: + return str(inputs[field][0]) + + return None + + def get_generation_parameters(self, metadata: Dict[str, Any]) -> Dict[str, Any]: + """ + Extract generation parameters from metadata. + + Args: + metadata: Full metadata dictionary + + Returns: + Dictionary of generation parameters + """ + parameters = {} + + # Common generation parameters to extract + param_fields = [ + 'steps', 'cfg_scale', 'sampler', 'scheduler', 'seed', + 'model', 'width', 'height', 'batch_size' + ] + + for field in param_fields: + if field in metadata: + parameters[field] = metadata[field] + + # Extract from workflow if available + if 'workflow' in metadata: + workflow_params = self.extract_params_from_workflow(metadata['workflow']) + parameters.update(workflow_params) + + return parameters + + def extract_params_from_workflow(self, workflow_data: Dict) -> Dict[str, Any]: + """ + Extract generation parameters from workflow data. + + Args: + workflow_data: ComfyUI workflow data + + Returns: + Dictionary of extracted parameters + """ + parameters = {} + + # This would need to be customized based on your specific workflow structure + # For now, return empty dict - can be expanded based on specific needs + + return parameters \ No newline at end of file diff --git a/utils/prompt_tracker.py b/utils/prompt_tracker.py new file mode 100644 index 0000000..a1bd828 --- /dev/null +++ b/utils/prompt_tracker.py @@ -0,0 +1,246 @@ +""" +Prompt tracking system for linking generated images to prompts. +Tracks active prompt executions to associate them with generated images. +""" + +import threading +import time +import uuid +import hashlib +from typing import Optional, Dict, Any +from datetime import datetime, timezone + + +class PromptTracker: + """Thread-safe tracking of current prompt executions.""" + + def __init__(self, db_manager): + """ + Initialize the prompt tracker. + + Args: + db_manager: Database manager instance + """ + self.db_manager = db_manager + self._local = threading.local() + self.active_prompts = {} # Global tracking for multiple threads + self.lock = threading.Lock() + self.cleanup_interval = 300 # 5 minutes + self.prompt_timeout = 120 # 2 minutes + + # Start cleanup thread + self.cleanup_thread = threading.Thread(target=self._cleanup_expired_prompts, daemon=True) + self.cleanup_thread.start() + + def set_current_prompt(self, prompt_text: str, additional_data: Optional[Dict[str, Any]] = None) -> str: + """ + Set the current prompt for this thread and store in database. + + Args: + prompt_text: The prompt text being executed + additional_data: Additional prompt metadata + + Returns: + Execution ID for this prompt + """ + execution_id = self.generate_execution_id() + + # Save prompt to database first + prompt_hash = hashlib.sha256(prompt_text.encode('utf-8')).hexdigest() + + try: + # Check if prompt already exists + existing_prompt = self.db_manager.get_prompt_by_hash(prompt_hash) + + if existing_prompt: + prompt_id = existing_prompt['id'] + print(f"[PromptManager] Using existing prompt ID: {prompt_id}") + else: + # Save new prompt + prompt_id = self.db_manager.save_prompt( + text=prompt_text, + prompt_hash=prompt_hash, + category=additional_data.get('category') if additional_data else None, + tags=additional_data.get('tags') if additional_data else None, + rating=additional_data.get('rating') if additional_data else None, + notes=additional_data.get('notes') if additional_data else None + ) + print(f"[PromptManager] Created new prompt ID: {prompt_id}") + + except Exception as e: + print(f"[PromptManager] Error saving prompt: {e}") + # Generate a temporary ID for tracking + prompt_id = f"temp_{int(time.time())}" + + # Create execution context + execution_context = { + 'id': prompt_id, + 'execution_id': execution_id, + 'text': prompt_text, + 'timestamp': time.time(), + 'thread_id': threading.current_thread().ident, + 'additional_data': additional_data or {} + } + + # Store in thread-local storage + self._local.current_prompt = execution_context + + # Store in global tracking with thread safety + with self.lock: + self.active_prompts[execution_id] = execution_context + + print(f"[PromptManager] Set current prompt: {execution_id} -> {prompt_text[:50]}...") + return execution_id + + def get_current_prompt(self) -> Optional[Dict[str, Any]]: + """ + Get the current prompt context for this thread. + + Returns: + Current prompt context or None + """ + current = getattr(self._local, 'current_prompt', None) + + if current: + # Check if prompt hasn't expired + if time.time() - current['timestamp'] < self.prompt_timeout: + return current + else: + print(f"[PromptManager] Current prompt expired: {current['execution_id']}") + self.clear_current_prompt() + + # Fallback: try to find recent prompt from global tracking + return self._find_recent_prompt() + + def _find_recent_prompt(self) -> Optional[Dict[str, Any]]: + """Find the most recent prompt that's still valid.""" + with self.lock: + current_time = time.time() + recent_prompts = [ + prompt for prompt in self.active_prompts.values() + if current_time - prompt['timestamp'] < self.prompt_timeout + ] + + if recent_prompts: + # Return the most recent one + return max(recent_prompts, key=lambda p: p['timestamp']) + + return None + + def clear_current_prompt(self): + """Clear the current prompt context for this thread.""" + current = getattr(self._local, 'current_prompt', None) + if current: + execution_id = current['execution_id'] + print(f"[PromptManager] Clearing current prompt: {execution_id}") + + # Clear from thread-local storage + self._local.current_prompt = None + + # Remove from global tracking + with self.lock: + self.active_prompts.pop(execution_id, None) + + def extend_prompt_timeout(self, execution_id: str, additional_seconds: int = 60): + """ + Extend the timeout for a specific prompt execution. + + Args: + execution_id: The execution ID to extend + additional_seconds: Additional seconds to add to timeout + """ + with self.lock: + if execution_id in self.active_prompts: + self.active_prompts[execution_id]['timestamp'] = time.time() + print(f"[PromptManager] Extended timeout for prompt: {execution_id}") + + def generate_execution_id(self) -> str: + """Generate a unique execution ID.""" + return f"exec_{uuid.uuid4().hex[:8]}_{int(time.time())}" + + def _cleanup_expired_prompts(self): + """Background thread to clean up expired prompts.""" + while True: + try: + current_time = time.time() + expired_ids = [] + + with self.lock: + for exec_id, prompt_data in self.active_prompts.items(): + if current_time - prompt_data['timestamp'] > self.prompt_timeout: + expired_ids.append(exec_id) + + for exec_id in expired_ids: + self.active_prompts.pop(exec_id, None) + + if expired_ids: + print(f"[PromptManager] Cleaned up {len(expired_ids)} expired prompts") + + time.sleep(self.cleanup_interval) + + except Exception as e: + print(f"[PromptManager] Error in cleanup thread: {e}") + time.sleep(60) # Wait a minute before retrying + + def get_active_prompts(self) -> Dict[str, Dict[str, Any]]: + """ + Get all currently active prompts. + + Returns: + Dictionary of active prompts + """ + with self.lock: + return self.active_prompts.copy() + + def get_status(self) -> Dict[str, Any]: + """ + Get tracker status information. + + Returns: + Status information dictionary + """ + with self.lock: + active_count = len(self.active_prompts) + + current_prompt = self.get_current_prompt() + + return { + 'active_prompts_count': active_count, + 'current_prompt_id': current_prompt['id'] if current_prompt else None, + 'current_execution_id': current_prompt['execution_id'] if current_prompt else None, + 'thread_id': threading.current_thread().ident, + 'prompt_timeout': self.prompt_timeout, + 'cleanup_interval': self.cleanup_interval + } + + +class PromptExecutionContext: + """Context manager for prompt executions.""" + + def __init__(self, prompt_tracker: PromptTracker, prompt_text: str, **kwargs): + """ + Initialize execution context. + + Args: + prompt_tracker: PromptTracker instance + prompt_text: The prompt text + **kwargs: Additional prompt metadata + """ + self.prompt_tracker = prompt_tracker + self.prompt_text = prompt_text + self.additional_data = kwargs + self.execution_id = None + + def __enter__(self): + """Enter the execution context.""" + self.execution_id = self.prompt_tracker.set_current_prompt( + self.prompt_text, + self.additional_data + ) + return self.execution_id + + def __exit__(self, exc_type, exc_val, exc_tb): + """Exit the execution context.""" + # Don't clear immediately - let the timeout handle it + # This allows images to be generated after the prompt execution completes + pass \ No newline at end of file diff --git a/utils/validators.py b/utils/validators.py new file mode 100644 index 0000000..9cc9fd1 --- /dev/null +++ b/utils/validators.py @@ -0,0 +1,221 @@ +""" +Input validation utilities for KikoTextEncode. +""" + +import re +from typing import List, Optional, Union + + +def validate_prompt_text(text: str) -> bool: + """ + Validate prompt text input. + + Args: + text: The prompt text to validate + + Returns: + bool: True if valid, False otherwise + + Raises: + ValueError: If text is invalid with descriptive message + """ + if not isinstance(text, str): + raise ValueError("Prompt text must be a string") + + if not text or not text.strip(): + raise ValueError("Prompt text cannot be empty") + + if len(text.strip()) > 10000: # Reasonable limit for prompt length + raise ValueError("Prompt text is too long (maximum 10,000 characters)") + + return True + + +def validate_rating(rating: Optional[int]) -> bool: + """ + Validate rating input. + + Args: + rating: The rating to validate (1-5 or None) + + Returns: + bool: True if valid, False otherwise + + Raises: + ValueError: If rating is invalid + """ + if rating is None: + return True + + if not isinstance(rating, int): + raise ValueError("Rating must be an integer") + + if rating < 1 or rating > 5: + raise ValueError("Rating must be between 1 and 5") + + return True + + +def validate_tags(tags: Union[str, List[str], None]) -> bool: + """ + Validate tags input. + + Args: + tags: Tags as string, list, or None + + Returns: + bool: True if valid, False otherwise + + Raises: + ValueError: If tags are invalid + """ + if tags is None: + return True + + if isinstance(tags, str): + # Parse comma-separated tags + tag_list = [tag.strip() for tag in tags.split(',') if tag.strip()] + tags = tag_list + + if not isinstance(tags, list): + raise ValueError("Tags must be a string, list, or None") + + for tag in tags: + if not isinstance(tag, str): + raise ValueError("All tags must be strings") + + if not tag.strip(): + raise ValueError("Tags cannot be empty") + + if len(tag.strip()) > 50: + raise ValueError("Individual tags cannot exceed 50 characters") + + # Check for invalid characters (optional - you can adjust this) + if not re.match(r'^[a-zA-Z0-9\s\-_]+$', tag.strip()): + raise ValueError(f"Tag '{tag}' contains invalid characters") + + if len(tags) > 20: # Reasonable limit + raise ValueError("Maximum 20 tags allowed") + + return True + + +def validate_category(category: Optional[str]) -> bool: + """ + Validate category input. + + Args: + category: The category to validate + + Returns: + bool: True if valid, False otherwise + + Raises: + ValueError: If category is invalid + """ + if category is None: + return True + + if not isinstance(category, str): + raise ValueError("Category must be a string") + + category = category.strip() + if not category: + return True # Empty category is valid (same as None) + + if len(category) > 100: + raise ValueError("Category cannot exceed 100 characters") + + # Check for invalid characters (adjust as needed) + if not re.match(r'^[a-zA-Z0-9\s\-_]+$', category): + raise ValueError("Category contains invalid characters") + + return True + + +def validate_workflow_name(workflow_name: Optional[str]) -> bool: + """ + Validate workflow name input. + + Args: + workflow_name: The workflow name to validate + + Returns: + bool: True if valid, False otherwise + + Raises: + ValueError: If workflow name is invalid + """ + if workflow_name is None: + return True + + if not isinstance(workflow_name, str): + raise ValueError("Workflow name must be a string") + + workflow_name = workflow_name.strip() + if not workflow_name: + return True # Empty workflow name is valid + + if len(workflow_name) > 200: + raise ValueError("Workflow name cannot exceed 200 characters") + + return True + + +def sanitize_input(text: str) -> str: + """ + Sanitize text input by removing potentially harmful content. + + Args: + text: The text to sanitize + + Returns: + str: Sanitized text + """ + if not isinstance(text, str): + return "" + + # Remove null bytes and other control characters + sanitized = text.replace('\x00', '').replace('\r\n', '\n').replace('\r', '\n') + + # Strip excessive whitespace but preserve single newlines + lines = sanitized.split('\n') + sanitized_lines = [line.strip() for line in lines] + + # Remove excessive empty lines (keep max 2 consecutive) + result_lines = [] + empty_count = 0 + + for line in sanitized_lines: + if not line: + empty_count += 1 + if empty_count <= 2: + result_lines.append(line) + else: + empty_count = 0 + result_lines.append(line) + + return '\n'.join(result_lines).strip() + + +def parse_tags_string(tags_string: str) -> List[str]: + """ + Parse a comma-separated tags string into a clean list. + + Args: + tags_string: Comma-separated tags string + + Returns: + List[str]: Cleaned list of unique tags + """ + if not tags_string or not isinstance(tags_string, str): + return [] + + # Split by comma and clean each tag + tags = [] + for tag in tags_string.split(','): + clean_tag = sanitize_input(tag).strip() + if clean_tag and clean_tag not in tags: # Avoid duplicates + tags.append(clean_tag) + + return tags[:20] # Limit to 20 tags \ No newline at end of file