From 2737a4b1efb3539829fca3cb66d64c90566db096 Mon Sep 17 00:00:00 2001 From: Vito Sansevero Date: Sat, 7 Feb 2026 08:30:09 -0800 Subject: [PATCH] style: apply Black formatter (line-length=88) to all Python files Automated formatting pass across 30 files to establish consistent code style enforced by CI. No logic changes. --- __init__.py | 7 +- database/__init__.py | 2 +- database/models.py | 152 +++--- database/operations.py | 615 ++++++++++++++---------- prompt_manager.py | 13 +- prompt_manager_base.py | 1 + prompt_manager_text.py | 12 +- py/__init__.py | 2 +- py/api/__init__.py | 198 +++++--- py/api/admin.py | 732 +++++++++++++++++----------- py/api/autotag_routes.py | 283 +++++------ py/api/images.py | 898 ++++++++++++++++++++--------------- py/api/logging_routes.py | 155 +++--- py/api/prompts.py | 345 +++++++++----- py/autotag.py | 203 ++++---- py/config.py | 296 ++++++------ restart_gallery.py | 15 +- tests/__init__.py | 2 +- tests/test_api.py | 36 +- tests/test_basic.py | 129 +++-- tests/test_database.py | 42 +- utils/__init__.py | 7 +- utils/comfyui_integration.py | 196 ++++---- utils/diagnostics.py | 298 ++++++------ utils/hashing.py | 66 +-- utils/image_monitor.py | 193 ++++---- utils/logging_config.py | 294 ++++++------ utils/metadata_extractor.py | 236 ++++----- utils/prompt_tracker.py | 255 +++++----- utils/validators.py | 126 ++--- 30 files changed, 3278 insertions(+), 2531 deletions(-) diff --git a/__init__.py b/__init__.py index 159140a..86d9d6e 100644 --- a/__init__.py +++ b/__init__.py @@ -1,6 +1,6 @@ """ -ComfyUI_PromptManager: A comprehensive ComfyUI custom node that extends text encoding -with persistent prompt storage, advanced search capabilities, and an automatic image +ComfyUI_PromptManager: A comprehensive ComfyUI custom node that extends text encoding +with persistent prompt storage, advanced search capabilities, and an automatic image gallery system. This module provides two main node types: @@ -13,7 +13,7 @@ generated images to their source prompts. Features: - SQLite-based persistent prompt storage with deduplication -- Advanced search and categorization system +- Advanced search and categorization system - Real-time image gallery with metadata extraction - Web-based admin dashboard for prompt management - Comprehensive logging and diagnostics @@ -105,6 +105,7 @@ try: except Exception as e: try: from .utils.logging_config import get_logger + _init_logger = get_logger("prompt_manager.init") _init_logger.error(f"Failed to start image monitoring: {e}") except Exception: diff --git a/database/__init__.py b/database/__init__.py index f27ad19..e85525b 100644 --- a/database/__init__.py +++ b/database/__init__.py @@ -8,4 +8,4 @@ search capabilities, metadata management, and image tracking functionality. from .operations import PromptDatabase from .models import PromptModel -__all__ = ["PromptDatabase", "PromptModel"] \ No newline at end of file +__all__ = ["PromptDatabase", "PromptModel"] diff --git a/database/models.py b/database/models.py index 299380d..0a028d6 100644 --- a/database/models.py +++ b/database/models.py @@ -12,6 +12,7 @@ try: from ..utils.logging_config import get_logger except ImportError: import sys + current_dir = os.path.dirname(os.path.dirname(os.path.abspath(__file__))) sys.path.insert(0, current_dir) from utils.logging_config import get_logger @@ -19,28 +20,28 @@ except ImportError: class PromptModel: """Database model for prompt storage and schema management.""" - + def __init__(self, db_path: str = "prompts.db"): """ Initialize the database model. - + Args: db_path: Path to the SQLite database file """ - self.logger = get_logger('prompt_manager.database.models') + self.logger = get_logger("prompt_manager.database.models") self.logger.debug(f"Initializing database model with path: {db_path}") self.db_path = db_path self._conn: Optional[sqlite3.Connection] = None self._conn_lock = threading.Lock() self._ensure_database_exists() - + def _ensure_database_exists(self) -> None: """ Create database and tables if they don't exist. - + Sets up the database schema including tables for prompts and generated images, creates necessary indexes, and applies any pending migrations. - + Raises: Exception: If database creation fails """ @@ -55,14 +56,14 @@ class PromptModel: except Exception as e: self.logger.error(f"Error creating database: {e}") raise - + def _create_tables(self, conn: sqlite3.Connection) -> None: """ Create the prompts and generated_images tables with all required columns. - + Args: conn: Active database connection - + Creates: - prompts table: Stores prompt text and metadata - generated_images table: Links generated images to their source prompts @@ -80,7 +81,7 @@ class PromptModel: hash TEXT UNIQUE ) """) - + # Create images table for gallery functionality conn.execute(""" CREATE TABLE IF NOT EXISTS generated_images ( @@ -129,14 +130,14 @@ class PromptModel: # Migrate JSON tags to normalized junction tables self._migrate_json_tags_to_junction(conn) - + def _create_indexes(self, conn: sqlite3.Connection) -> None: """ Create indexes for better query performance. - + Args: conn: Active database connection - + Creates indexes on: - Text content for search operations - Categories and tags for filtering @@ -155,10 +156,10 @@ class PromptModel: "CREATE INDEX IF NOT EXISTS idx_prompt_tags_tag ON prompt_tags(tag_id)", "CREATE INDEX IF NOT EXISTS idx_tags_name ON tags(name)", ] - + for index_sql in indexes: conn.execute(index_sql) - + def get_connection(self) -> sqlite3.Connection: """ Get the persistent database connection. @@ -173,22 +174,20 @@ class PromptModel: """ with self._conn_lock: if self._conn is None: - self._conn = sqlite3.connect( - self.db_path, check_same_thread=False - ) + self._conn = sqlite3.connect(self.db_path, check_same_thread=False) self._conn.row_factory = sqlite3.Row self._conn.execute("PRAGMA journal_mode = WAL") self._conn.execute("PRAGMA foreign_keys = ON") self._conn.execute("PRAGMA busy_timeout = 5000") return self._conn - + def _migrate_workflow_name_removal(self, conn: sqlite3.Connection) -> None: """ Remove workflow_name column if it exists in existing database. - + Args: conn: Active database connection - + This migration handles legacy schema updates by removing the deprecated workflow_name column while preserving all other data. """ @@ -196,10 +195,10 @@ class PromptModel: # Check if workflow_name column exists cursor = conn.execute("PRAGMA table_info(prompts)") columns = [column[1] for column in cursor.fetchall()] - - if 'workflow_name' in columns: + + if "workflow_name" in columns: self.logger.info("Migrating database: removing workflow_name column") - + # Create new table without workflow_name conn.execute(""" CREATE TABLE prompts_new ( @@ -214,31 +213,31 @@ class PromptModel: hash TEXT UNIQUE ) """) - + # Copy data from old table to new table conn.execute(""" INSERT INTO prompts_new (id, text, created_at, updated_at, category, tags, rating, notes, hash) SELECT id, text, created_at, updated_at, category, tags, rating, notes, hash FROM prompts """) - + # Drop old table and rename new one conn.execute("DROP TABLE prompts") conn.execute("ALTER TABLE prompts_new RENAME TO prompts") - + self.logger.info("Database migration completed") - + except Exception as e: self.logger.error(f"Migration error: {e}") # If migration fails, the table creation will handle it - + def _migrate_foreign_key_types(self, conn: sqlite3.Connection) -> None: """ Fix foreign key data type mismatch in generated_images table. - + Args: conn: Active database connection - + Converts prompt_id from TEXT to INTEGER type to match the prompts table's primary key type, ensuring referential integrity. """ @@ -246,10 +245,12 @@ class PromptModel: # Check if generated_images table exists and has TEXT prompt_id cursor = conn.execute("PRAGMA table_info(generated_images)") columns = {column[1]: column[2] for column in cursor.fetchall()} - - if 'prompt_id' in columns and columns['prompt_id'] == 'TEXT': - self.logger.info("Migrating foreign key types: prompt_id TEXT -> INTEGER") - + + if "prompt_id" in columns and columns["prompt_id"] == "TEXT": + self.logger.info( + "Migrating foreign key types: prompt_id TEXT -> INTEGER" + ) + # Create new table with correct types conn.execute(""" CREATE TABLE generated_images_new ( @@ -268,7 +269,7 @@ class PromptModel: FOREIGN KEY (prompt_id) REFERENCES prompts(id) ON DELETE CASCADE ) """) - + # Copy data, converting prompt_id from TEXT to INTEGER conn.execute(""" INSERT INTO generated_images_new @@ -280,17 +281,19 @@ class PromptModel: WHERE prompt_id != '' AND prompt_id IS NOT NULL AND CAST(prompt_id AS INTEGER) IN (SELECT id FROM prompts) """) - + # Drop old table and rename new one conn.execute("DROP TABLE generated_images") - conn.execute("ALTER TABLE generated_images_new RENAME TO generated_images") - + conn.execute( + "ALTER TABLE generated_images_new RENAME TO generated_images" + ) + self.logger.info("Foreign key migration completed") - + except Exception as e: self.logger.error(f"Foreign key migration error: {e}") # If migration fails, continue with existing schema - + def _migrate_add_unique_constraint(self, conn: sqlite3.Connection) -> None: """ Add UNIQUE constraint on (prompt_id, filename) to prevent duplicate image entries. @@ -314,14 +317,16 @@ class PromptModel: if idx[2] == 1: # unique flag cursor = conn.execute(f"PRAGMA index_info({idx[1]})") columns = [col[2] for col in cursor.fetchall()] - if 'prompt_id' in columns and 'filename' in columns: + if "prompt_id" in columns and "filename" in columns: has_unique_constraint = True break if has_unique_constraint: return # Already migrated - self.logger.info("Migrating database: adding UNIQUE constraint on (prompt_id, filename)") + self.logger.info( + "Migrating database: adding UNIQUE constraint on (prompt_id, filename)" + ) # First, remove duplicates keeping only the most recent (highest id) conn.execute(""" @@ -334,7 +339,9 @@ class PromptModel: duplicates_removed = conn.total_changes if duplicates_removed > 0: - self.logger.info(f"Removed {duplicates_removed} duplicate image entries") + self.logger.info( + f"Removed {duplicates_removed} duplicate image entries" + ) # Create new table with UNIQUE constraint conn.execute(""" @@ -371,9 +378,15 @@ class PromptModel: conn.execute("ALTER TABLE generated_images_new RENAME TO generated_images") # Recreate indexes - conn.execute("CREATE INDEX IF NOT EXISTS idx_prompt_images ON generated_images(prompt_id)") - conn.execute("CREATE INDEX IF NOT EXISTS idx_image_path ON generated_images(image_path)") - conn.execute("CREATE INDEX IF NOT EXISTS idx_generation_time ON generated_images(generation_time)") + conn.execute( + "CREATE INDEX IF NOT EXISTS idx_prompt_images ON generated_images(prompt_id)" + ) + conn.execute( + "CREATE INDEX IF NOT EXISTS idx_image_path ON generated_images(image_path)" + ) + conn.execute( + "CREATE INDEX IF NOT EXISTS idx_generation_time ON generated_images(generation_time)" + ) self.logger.info("UNIQUE constraint migration completed successfully") @@ -436,11 +449,11 @@ class PromptModel: """ # Future migrations can be added here pass - + def vacuum_database(self) -> None: """ Optimize database by running VACUUM command. - + Reclaims unused space and defragments the database file, improving query performance and reducing file size. """ @@ -450,57 +463,60 @@ class PromptModel: conn.commit() except Exception as e: self.logger.error(f"Error vacuuming database: {e}") - + def get_database_info(self) -> dict: """ Get information about the database. - + Returns: dict: Database statistics and information """ try: with self.get_connection() as conn: cursor = conn.execute("SELECT COUNT(*) as total_prompts FROM prompts") - total_prompts = cursor.fetchone()['total_prompts'] - + total_prompts = cursor.fetchone()["total_prompts"] + cursor = conn.execute( "SELECT COUNT(DISTINCT category) as unique_categories FROM prompts WHERE category IS NOT NULL" ) - unique_categories = cursor.fetchone()['unique_categories'] - + unique_categories = cursor.fetchone()["unique_categories"] + cursor = conn.execute( "SELECT AVG(rating) as avg_rating FROM prompts WHERE rating IS NOT NULL" ) - avg_rating = cursor.fetchone()['avg_rating'] - + avg_rating = cursor.fetchone()["avg_rating"] + # Get database file size - db_size = os.path.getsize(self.db_path) if os.path.exists(self.db_path) else 0 - + db_size = ( + os.path.getsize(self.db_path) if os.path.exists(self.db_path) else 0 + ) + return { - 'total_prompts': total_prompts, - 'unique_categories': unique_categories, - 'average_rating': round(avg_rating, 2) if avg_rating else None, - 'database_size_bytes': db_size, - 'database_path': os.path.abspath(self.db_path) + "total_prompts": total_prompts, + "unique_categories": unique_categories, + "average_rating": round(avg_rating, 2) if avg_rating else None, + "database_size_bytes": db_size, + "database_path": os.path.abspath(self.db_path), } except Exception as e: self.logger.error(f"Error getting database info: {e}") return {} - + def backup_database(self, backup_path: str) -> bool: """ Create a backup of the database. - + Args: backup_path: Path where the backup should be saved - + Returns: bool: True if backup was successful, False otherwise """ try: import shutil + shutil.copy2(self.db_path, backup_path) return True except Exception as e: self.logger.error(f"Error creating database backup: {e}") - return False \ No newline at end of file + return False diff --git a/database/operations.py b/database/operations.py index 440eaee..2e04d9a 100644 --- a/database/operations.py +++ b/database/operations.py @@ -15,6 +15,7 @@ try: from ..utils.logging_config import get_logger except ImportError: import sys + current_dir = os.path.dirname(os.path.dirname(os.path.abspath(__file__))) sys.path.insert(0, current_dir) from utils.logging_config import get_logger @@ -34,15 +35,15 @@ class PromptDatabase: def __init__(self, db_path: str = "prompts.db"): """ Initialize the database operations. - + Args: db_path: Path to the SQLite database file """ - self.logger = get_logger('prompt_manager.database') + self.logger = get_logger("prompt_manager.database") self.logger.debug(f"Initializing database operations with path: {db_path}") self.model = PromptModel(db_path) self.logger.debug("Database operations initialized successfully") - + def save_prompt( self, text: str, @@ -50,11 +51,11 @@ class PromptDatabase: tags: Optional[List[str]] = None, rating: Optional[int] = None, notes: Optional[str] = None, - prompt_hash: Optional[str] = None + prompt_hash: Optional[str] = None, ) -> int: """ Save a new prompt to the database. - + Args: text: The prompt text category: Optional category @@ -62,21 +63,23 @@ class PromptDatabase: rating: Rating 1-5 notes: Optional notes prompt_hash: SHA256 hash of the prompt - + Returns: int: The ID of the saved prompt - + Raises: ValueError: If required parameters are invalid sqlite3.Error: If database operation fails """ if not text or not text.strip(): raise ValueError("Prompt text cannot be empty") - + if rating is not None and (rating < 1 or rating > 5): raise ValueError("Rating must be between 1 and 5") - - self.logger.debug(f"Saving prompt: text_length={len(text)}, category={category}, tags={tags}, rating={rating}") + + self.logger.debug( + f"Saving prompt: text_length={len(text)}, category={category}, tags={tags}, rating={rating}" + ) with self.model.get_connection() as conn: cursor = conn.execute( @@ -92,8 +95,8 @@ class PromptDatabase: notes, prompt_hash, datetime.datetime.now(datetime.timezone.utc).isoformat(), - datetime.datetime.now(datetime.timezone.utc).isoformat() - ) + datetime.datetime.now(datetime.timezone.utc).isoformat(), + ), ) prompt_id = cursor.lastrowid if tags: @@ -101,20 +104,21 @@ class PromptDatabase: conn.commit() self.logger.debug(f"Successfully saved prompt with ID: {prompt_id}") return prompt_id - + def get_prompt_by_id(self, prompt_id: int) -> Optional[Dict[str, Any]]: """ Get a prompt by its ID. - + Args: prompt_id: The prompt ID - + Returns: Dict containing prompt data or None if not found """ with self.model.get_connection() as conn: cursor = conn.execute( - f"SELECT prompts.*, {TAG_SUBQUERY} FROM prompts WHERE id = ?", (prompt_id,) + f"SELECT prompts.*, {TAG_SUBQUERY} FROM prompts WHERE id = ?", + (prompt_id,), ) row = cursor.fetchone() return self._row_to_dict(row) if row else None @@ -131,11 +135,12 @@ class PromptDatabase: """ with self.model.get_connection() as conn: cursor = conn.execute( - f"SELECT prompts.*, {TAG_SUBQUERY} FROM prompts WHERE hash = ?", (prompt_hash,) + f"SELECT prompts.*, {TAG_SUBQUERY} FROM prompts WHERE hash = ?", + (prompt_hash,), ) row = cursor.fetchone() return self._row_to_dict(row) if row else None - + def search_prompts( self, text: Optional[str] = None, @@ -146,11 +151,11 @@ class PromptDatabase: date_from: Optional[str] = None, date_to: Optional[str] = None, limit: int = 100, - offset: int = 0 + offset: int = 0, ) -> List[Dict[str, Any]]: """ Search prompts with various filters. - + Args: text: Text to search for in prompt content category: Filter by category @@ -161,7 +166,7 @@ class PromptDatabase: date_to: End date filter (ISO format) limit: Maximum number of results offset: Number of results to skip - + Returns: List of dictionaries containing prompt data """ @@ -220,11 +225,11 @@ class PromptDatabase: def get_recent_prompts(self, limit: int = 10, offset: int = 0) -> Dict[str, Any]: """ Get the most recent prompts with pagination support. - + Args: limit: Maximum number of prompts to return offset: Number of prompts to skip (for pagination) - + Returns: Dictionary containing prompt data and pagination info """ @@ -233,11 +238,11 @@ class PromptDatabase: cursor = conn.execute("SELECT COUNT(*) FROM prompts") row = cursor.fetchone() total_count = (row[0] if row else 0) or 0 - + # Get paginated results cursor = conn.execute( f"SELECT prompts.*, {TAG_SUBQUERY} FROM prompts ORDER BY created_at DESC LIMIT ? OFFSET ?", - (limit, offset) + (limit, offset), ) rows = cursor.fetchall() prompts = [self._row_to_dict(row) for row in rows] @@ -247,41 +252,43 @@ class PromptDatabase: self._attach_preview_images(conn, prompts, prompt_ids) return { - 'prompts': prompts, - 'total': total_count, - 'limit': limit, - 'offset': offset, - 'has_more': (offset + limit) < total_count, - 'page': (offset // limit) + 1, - 'total_pages': (total_count + limit - 1) // limit # Ceiling division + "prompts": prompts, + "total": total_count, + "limit": limit, + "offset": offset, + "has_more": (offset + limit) < total_count, + "page": (offset // limit) + 1, + "total_pages": (total_count + limit - 1) // limit, # Ceiling division } - - def get_prompts_by_category(self, category: str, limit: int = 100) -> List[Dict[str, Any]]: + + def get_prompts_by_category( + self, category: str, limit: int = 100 + ) -> List[Dict[str, Any]]: """ Get all prompts in a specific category. - + Args: category: The category name limit: Maximum number of prompts to return - + Returns: List of dictionaries containing prompt data """ with self.model.get_connection() as conn: cursor = conn.execute( f"SELECT prompts.*, {TAG_SUBQUERY} FROM prompts WHERE category = ? ORDER BY created_at DESC LIMIT ?", - (category, limit) + (category, limit), ) rows = cursor.fetchall() return [self._row_to_dict(row) for row in rows] - + def get_top_rated_prompts(self, limit: int = 10) -> List[Dict[str, Any]]: """ Get the highest rated prompts. - + Args: limit: Maximum number of prompts to return - + Returns: List of dictionaries containing prompt data """ @@ -293,29 +300,29 @@ class PromptDatabase: ORDER BY rating DESC, created_at DESC LIMIT ? """, - (limit,) + (limit,), ) rows = cursor.fetchall() return [self._row_to_dict(row) for row in rows] - + def update_prompt_metadata( self, prompt_id: int, category: Optional[str] = None, tags: Optional[List[str]] = None, rating: Optional[int] = None, - notes: Optional[str] = None + notes: Optional[str] = None, ) -> bool: """ Update metadata for an existing prompt. - + Args: prompt_id: The prompt ID category: New category tags: New tags list rating: New rating notes: New notes - + Returns: bool: True if update was successful, False otherwise """ @@ -350,33 +357,38 @@ class PromptDatabase: self._sync_prompt_tags(conn, prompt_id, tags) conn.execute( "UPDATE prompts SET updated_at = ? WHERE id = ?", - (datetime.datetime.now(datetime.timezone.utc).isoformat(), prompt_id), + ( + datetime.datetime.now(datetime.timezone.utc).isoformat(), + prompt_id, + ), ) conn.commit() return True - + def delete_prompt(self, prompt_id: int) -> bool: """ Delete a prompt by its ID. - + Args: prompt_id: The prompt ID to delete - + Returns: bool: True if deletion was successful, False otherwise """ with self.model.get_connection() as conn: # First delete related images to avoid foreign key constraint - conn.execute("DELETE FROM generated_images WHERE prompt_id = ?", (prompt_id,)) + conn.execute( + "DELETE FROM generated_images WHERE prompt_id = ?", (prompt_id,) + ) # Then delete the prompt cursor = conn.execute("DELETE FROM prompts WHERE id = ?", (prompt_id,)) conn.commit() return cursor.rowcount > 0 - + def get_all_categories(self) -> List[str]: """ Get all unique categories from the database. - + Returns: List of category names """ @@ -384,8 +396,8 @@ class PromptDatabase: cursor = conn.execute( "SELECT DISTINCT TRIM(category) as category FROM prompts WHERE category IS NOT NULL AND TRIM(category) != '' ORDER BY category" ) - return [row['category'] for row in cursor.fetchall()] - + return [row["category"] for row in cursor.fetchall()] + def get_all_tags(self) -> List[str]: """ Get all unique tags that are in use (linked to at least one prompt). @@ -407,7 +419,7 @@ class PromptDatabase: limit: int = 50, offset: int = 0, search: Optional[str] = None, - sort: str = "alpha_asc" + sort: str = "alpha_asc", ) -> Dict[str, Any]: """ Get all unique tags with their usage counts via junction table. @@ -454,7 +466,9 @@ class PromptDatabase: " LIMIT ? OFFSET ?" ) cursor = conn.execute(data_sql, params + [limit, offset]) - tags = [{"name": row["tag"], "count": row["count"]} for row in cursor.fetchall()] + tags = [ + {"name": row["tag"], "count": row["count"]} for row in cursor.fetchall() + ] return { "tags": tags, @@ -465,11 +479,7 @@ class PromptDatabase: } def get_prompts_by_tags( - self, - tags: List[str], - mode: str = "and", - limit: int = 20, - offset: int = 0 + self, tags: List[str], mode: str = "and", limit: int = 20, offset: int = 0 ) -> Dict[str, Any]: """ Get prompts that match the given tags with AND/OR filtering. @@ -487,7 +497,13 @@ class PromptDatabase: Dict with prompts list (including preview images), total count, pagination """ if not tags: - return {'prompts': [], 'total': 0, 'limit': limit, 'offset': offset, 'has_more': False} + return { + "prompts": [], + "total": 0, + "limit": limit, + "offset": offset, + "has_more": False, + } with self.model.get_connection() as conn: placeholders = ",".join(["?"] * len(tags)) @@ -564,22 +580,22 @@ class PromptDatabase: cursor = conn.execute("SELECT id FROM tags WHERE name = ?", (old_name,)) old_tag = cursor.fetchone() if not old_tag: - return {'success': True, 'affected_count': 0, 'skipped_count': 0} + return {"success": True, "affected_count": 0, "skipped_count": 0} - old_tag_id = old_tag['id'] + old_tag_id = old_tag["id"] # Count affected prompts before the operation cursor = conn.execute( "SELECT COUNT(*) as c FROM prompt_tags WHERE tag_id = ?", (old_tag_id,) ) - affected = cursor.fetchone()['c'] + affected = cursor.fetchone()["c"] cursor = conn.execute("SELECT id FROM tags WHERE name = ?", (new_name,)) existing_new = cursor.fetchone() if existing_new: # Target tag exists — merge: move links, handle conflicts, delete old - new_tag_id = existing_new['id'] + new_tag_id = existing_new["id"] conn.execute( "UPDATE OR IGNORE prompt_tags SET tag_id = ? WHERE tag_id = ?", (new_tag_id, old_tag_id), @@ -589,12 +605,16 @@ class PromptDatabase: conn.execute("DELETE FROM tags WHERE id = ?", (old_tag_id,)) else: # Simple rename - conn.execute("UPDATE tags SET name = ? WHERE id = ?", (new_name, old_tag_id)) + conn.execute( + "UPDATE tags SET name = ? WHERE id = ?", (new_name, old_tag_id) + ) conn.commit() - self.logger.info(f"Renamed tag '{old_name}' -> '{new_name}' in {affected} prompts") - return {'success': True, 'affected_count': affected, 'skipped_count': 0} + self.logger.info( + f"Renamed tag '{old_name}' -> '{new_name}' in {affected} prompts" + ) + return {"success": True, "affected_count": affected, "skipped_count": 0} def delete_tag_all_prompts(self, tag_name: str) -> Dict[str, Any]: """ @@ -617,20 +637,20 @@ class PromptDatabase: cursor = conn.execute("SELECT id FROM tags WHERE name = ?", (tag_name,)) tag_row = cursor.fetchone() if not tag_row: - return {'success': True, 'affected_count': 0, 'skipped_count': 0} + return {"success": True, "affected_count": 0, "skipped_count": 0} - tag_id = tag_row['id'] + tag_id = tag_row["id"] cursor = conn.execute( "SELECT COUNT(*) as c FROM prompt_tags WHERE tag_id = ?", (tag_id,) ) - affected = cursor.fetchone()['c'] + affected = cursor.fetchone()["c"] # CASCADE delete handles prompt_tags entries conn.execute("DELETE FROM tags WHERE id = ?", (tag_id,)) conn.commit() self.logger.info(f"Deleted tag '{tag_name}' from {affected} prompts") - return {'success': True, 'affected_count': affected, 'skipped_count': 0} + return {"success": True, "affected_count": affected, "skipped_count": 0} def merge_tags(self, source_tags: List[str], target_tag: str) -> Dict[str, Any]: """ @@ -659,19 +679,19 @@ class PromptDatabase: # Ensure target tag exists conn.execute("INSERT OR IGNORE INTO tags (name) VALUES (?)", (target_tag,)) cursor = conn.execute("SELECT id FROM tags WHERE name = ?", (target_tag,)) - target_id = cursor.fetchone()['id'] + target_id = cursor.fetchone()["id"] for src_tag in source_tags: cursor = conn.execute("SELECT id FROM tags WHERE name = ?", (src_tag,)) src_row = cursor.fetchone() if not src_row: continue - src_id = src_row['id'] + src_id = src_row["id"] cursor = conn.execute( "SELECT COUNT(*) as c FROM prompt_tags WHERE tag_id = ?", (src_id,) ) - src_count = cursor.fetchone()['c'] + src_count = cursor.fetchone()["c"] if src_count > 0: # Move links from source to target (ignore conflicts) @@ -689,8 +709,15 @@ class PromptDatabase: conn.commit() - self.logger.info(f"Merged {tags_merged} tags into '{target_tag}', affected {affected} prompts") - return {'success': True, 'affected_count': affected, 'tags_merged': tags_merged, 'skipped_count': 0} + self.logger.info( + f"Merged {tags_merged} tags into '{target_tag}', affected {affected} prompts" + ) + return { + "success": True, + "affected_count": affected, + "tags_merged": tags_merged, + "skipped_count": 0, + } def get_untagged_prompts_count(self) -> int: """ @@ -704,7 +731,7 @@ class PromptDatabase: "SELECT COUNT(*) as total FROM prompts " "WHERE NOT EXISTS (SELECT 1 FROM prompt_tags WHERE prompt_id = prompts.id)" ) - return cursor.fetchone()['total'] + return cursor.fetchone()["total"] def get_untagged_prompts(self, limit: int = 20, offset: int = 0) -> Dict[str, Any]: """ @@ -720,10 +747,14 @@ class PromptDatabase: Dict with prompts list, total count, and pagination """ with self.model.get_connection() as conn: - where = "NOT EXISTS (SELECT 1 FROM prompt_tags WHERE prompt_id = prompts.id)" + where = ( + "NOT EXISTS (SELECT 1 FROM prompt_tags WHERE prompt_id = prompts.id)" + ) - cursor = conn.execute(f"SELECT COUNT(*) as total FROM prompts WHERE {where}") - total = cursor.fetchone()['total'] + cursor = conn.execute( + f"SELECT COUNT(*) as total FROM prompts WHERE {where}" + ) + total = cursor.fetchone()["total"] cursor = conn.execute( f"SELECT prompts.*, {TAG_SUBQUERY} FROM prompts WHERE {where} " @@ -738,11 +769,11 @@ class PromptDatabase: self._attach_preview_images(conn, prompts, prompt_ids) return { - 'prompts': prompts, - 'total': total, - 'limit': limit, - 'offset': offset, - 'has_more': (offset + limit) < total, + "prompts": prompts, + "total": total, + "limit": limit, + "offset": offset, + "has_more": (offset + limit) < total, } def _attach_preview_images( @@ -776,14 +807,18 @@ class PromptDatabase: f"WHERE prompt_id IN ({id_placeholders}) GROUP BY prompt_id", prompt_ids, ) - counts_by_prompt = {row["prompt_id"]: row["cnt"] for row in cnt_cursor.fetchall()} + counts_by_prompt = { + row["prompt_id"]: row["cnt"] for row in cnt_cursor.fetchall() + } for prompt in prompts: pid = prompt["id"] prompt["images"] = images_by_prompt.get(pid, []) prompt["image_count"] = counts_by_prompt.get(pid, 0) - def _ensure_tags(self, conn: sqlite3.Connection, tag_names: List[str]) -> Dict[str, int]: + def _ensure_tags( + self, conn: sqlite3.Connection, tag_names: List[str] + ) -> Dict[str, int]: """Ensure tag names exist in tags table, return name->id mapping.""" if not tag_names: return {} @@ -794,9 +829,11 @@ class PromptDatabase: f"SELECT id, name FROM tags WHERE name IN ({placeholders})", list(tag_names), ) - return {row['name']: row['id'] for row in cursor.fetchall()} + return {row["name"]: row["id"] for row in cursor.fetchall()} - def _sync_prompt_tags(self, conn: sqlite3.Connection, prompt_id: int, tags: List[str]) -> None: + def _sync_prompt_tags( + self, conn: sqlite3.Connection, prompt_id: int, tags: List[str] + ) -> None: """Replace all junction table entries for a prompt.""" conn.execute("DELETE FROM prompt_tags WHERE prompt_id = ?", (prompt_id,)) if tags: @@ -817,7 +854,7 @@ class PromptDatabase: "WHERE pt.prompt_id = ?", (prompt_id,), ) - return [row['name'] for row in cursor.fetchall()] + return [row["name"] for row in cursor.fetchall()] def _row_to_dict(self, row: sqlite3.Row) -> Dict[str, Any]: """ @@ -832,71 +869,72 @@ class PromptDatabase: data = dict(row) # Parse tags from junction table subquery result (preferred) - if '_tag_list' in data and data['_tag_list']: - data['tags'] = data['_tag_list'].split('|||') - del data['_tag_list'] - elif '_tag_list' in data: + if "_tag_list" in data and data["_tag_list"]: + data["tags"] = data["_tag_list"].split("|||") + del data["_tag_list"] + elif "_tag_list" in data: # TAG_SUBQUERY present but NULL (no tags) - data['tags'] = [] - del data['_tag_list'] - elif data.get('tags'): + data["tags"] = [] + del data["_tag_list"] + elif data.get("tags"): # Fallback: parse legacy JSON column try: - parsed = json.loads(data['tags']) + parsed = json.loads(data["tags"]) if isinstance(parsed, list): - data['tags'] = parsed + data["tags"] = parsed elif isinstance(parsed, str): - data['tags'] = [t.strip() for t in parsed.split(',') if t.strip()] + data["tags"] = [t.strip() for t in parsed.split(",") if t.strip()] else: - data['tags'] = [] + data["tags"] = [] except (json.JSONDecodeError, TypeError): - data['tags'] = [] + data["tags"] = [] else: - data['tags'] = [] + data["tags"] = [] return data - + def export_prompts(self, file_path: str, format: str = "json") -> bool: """ Export all prompts to a file. - + Args: file_path: Path to save the export file format: Export format ("json" or "csv") - + Returns: bool: True if export was successful, False otherwise """ try: prompts = self.search_prompts(limit=10000) # Get all prompts - + if format.lower() == "json": - with open(file_path, 'w', encoding='utf-8') as f: + with open(file_path, "w", encoding="utf-8") as f: json.dump(prompts, f, indent=2, ensure_ascii=False) elif format.lower() == "csv": import csv + if prompts: - with open(file_path, 'w', newline='', encoding='utf-8') as f: + with open(file_path, "w", newline="", encoding="utf-8") as f: writer = csv.DictWriter(f, fieldnames=prompts[0].keys()) writer.writeheader() for prompt in prompts: # Convert lists to strings for CSV row = prompt.copy() - if isinstance(row.get('tags'), list): - row['tags'] = ', '.join(row['tags']) + if isinstance(row.get("tags"), list): + row["tags"] = ", ".join(row["tags"]) writer.writerow(row) else: raise ValueError(f"Unsupported export format: {format}") - + return True except Exception as e: self.logger.error(f"Error exporting prompts: {e}") return False - + def find_duplicates(self) -> List[Dict[str, Any]]: """ Find duplicate prompts based on text content without removing them. - + Returns: List of duplicate groups, each containing: - text: The duplicate text content @@ -916,21 +954,23 @@ class PromptDatabase: GROUP BY LOWER(TRIM(text)) HAVING COUNT(*) > 1 """) - + duplicate_groups = cursor.fetchall() - self.logger.debug(f"Found {len(duplicate_groups)} groups of duplicate prompts") - + self.logger.debug( + f"Found {len(duplicate_groups)} groups of duplicate prompts" + ) + result = [] - + for group in duplicate_groups: - ids = group['ids'].split(',') - created_dates = group['created_dates'].split(',') - + ids = group["ids"].split(",") + created_dates = group["created_dates"].split(",") + # Sort IDs by created_at date id_date_pairs = list(zip(ids, created_dates)) id_date_pairs.sort(key=lambda x: x[1]) # Sort by date ids = [pair[0] for pair in id_date_pairs] - + # Get full details for all prompts in this duplicate group prompts = [] for prompt_id in ids: @@ -942,27 +982,33 @@ class PromptDatabase: prompt_data = cursor.fetchone() if prompt_data: prompt_dict = dict(prompt_data) - prompt_dict['tags'] = self._get_prompt_tags(conn, int(prompt_id)) + prompt_dict["tags"] = self._get_prompt_tags( + conn, int(prompt_id) + ) prompts.append(prompt_dict) - + if prompts: - result.append({ - 'text': prompts[0]['text'], # Use the actual text (not normalized) - 'prompts': prompts - }) - + result.append( + { + "text": prompts[0][ + "text" + ], # Use the actual text (not normalized) + "prompts": prompts, + } + ) + self.logger.info(f"Found {len(result)} groups with duplicates") return result - + except Exception as e: self.logger.error(f"Error finding duplicates: {e}") return [] - + def cleanup_duplicates(self) -> int: """ Remove duplicate prompts based on text content, preserving all image links. Merges metadata and transfers images to the retained prompt. - + Returns: int: Number of duplicates removed """ @@ -980,78 +1026,91 @@ class PromptDatabase: GROUP BY LOWER(TRIM(text)) HAVING COUNT(*) > 1 """) - + duplicates = cursor.fetchall() - self.logger.debug(f"Found {len(duplicates)} groups of duplicate prompts") + self.logger.debug( + f"Found {len(duplicates)} groups of duplicate prompts" + ) total_removed = 0 total_images_transferred = 0 - + for duplicate in duplicates: - ids = duplicate['ids'].split(',') - created_dates = duplicate['created_dates'].split(',') - + ids = duplicate["ids"].split(",") + created_dates = duplicate["created_dates"].split(",") + # Sort IDs by created_at date to keep the oldest id_date_pairs = list(zip(ids, created_dates)) id_date_pairs.sort(key=lambda x: x[1]) # Sort by date sorted_ids = [int(pair[0]) for pair in id_date_pairs] - + # Keep the oldest one (first), merge and delete the rest primary_id = sorted_ids[0] # Keep the oldest duplicate_ids = sorted_ids[1:] - - self.logger.debug(f"Merging duplicates: keeping {primary_id}, removing {duplicate_ids}") - + + self.logger.debug( + f"Merging duplicates: keeping {primary_id}, removing {duplicate_ids}" + ) + # Get primary prompt details - cursor = conn.execute("SELECT * FROM prompts WHERE id = ?", (primary_id,)) + cursor = conn.execute( + "SELECT * FROM prompts WHERE id = ?", (primary_id,) + ) primary_prompt = cursor.fetchone() if not primary_prompt: continue - + # Collect and merge metadata from all duplicates - merged_metadata = self._merge_duplicate_metadata(conn, primary_id, duplicate_ids) - + merged_metadata = self._merge_duplicate_metadata( + conn, primary_id, duplicate_ids + ) + # Transfer all images from duplicates to primary prompt - images_transferred = self._transfer_images_to_primary(conn, primary_id, duplicate_ids) + images_transferred = self._transfer_images_to_primary( + conn, primary_id, duplicate_ids + ) total_images_transferred += images_transferred - + # Update primary prompt with merged metadata if merged_metadata: - self._update_primary_with_merged_metadata(conn, primary_id, merged_metadata) - + self._update_primary_with_merged_metadata( + conn, primary_id, merged_metadata + ) + # Delete duplicate prompts (images already transferred) for duplicate_id in duplicate_ids: - conn.execute("DELETE FROM prompts WHERE id = ?", (duplicate_id,)) + conn.execute( + "DELETE FROM prompts WHERE id = ?", (duplicate_id,) + ) total_removed += 1 self.logger.debug(f"Removed duplicate prompt {duplicate_id}") - + conn.commit() - + if total_removed > 0: - self.logger.info(f"Removed {total_removed} duplicate prompts, transferred {total_images_transferred} images") + self.logger.info( + f"Removed {total_removed} duplicate prompts, transferred {total_images_transferred} images" + ) else: self.logger.info("No duplicate prompts found") - + return total_removed - + except Exception as e: self.logger.error(f"Error cleaning up duplicates: {e}", exc_info=True) return 0 # Gallery-related methods def link_image_to_prompt( - self, - prompt_id: str, - image_path: str, - metadata: Optional[Dict[str, Any]] = None + self, prompt_id: str, image_path: str, metadata: Optional[Dict[str, Any]] = None ) -> int: """ Link a generated image to a prompt. - + Args: prompt_id: ID of the prompt that generated this image image_path: Full path to the image file metadata: Optional metadata about the image - + Returns: int: The ID of the created image record """ @@ -1064,23 +1123,31 @@ class PromptDatabase: prompt_id_int = prompt_id else: # Handle temporary IDs like "temp_123456" - skip linking - if isinstance(prompt_id, str) and prompt_id.startswith('temp_'): - self.logger.debug(f"Skipping image linking for temporary prompt ID: {prompt_id}") + if isinstance(prompt_id, str) and prompt_id.startswith("temp_"): + self.logger.debug( + f"Skipping image linking for temporary prompt ID: {prompt_id}" + ) return 0 else: - self.logger.warning(f"Invalid prompt_id format: {prompt_id}, skipping image linking") + self.logger.warning( + f"Invalid prompt_id format: {prompt_id}, skipping image linking" + ) return 0 - + # Verify the prompt exists in the database with self.model.get_connection() as conn: - cursor = conn.execute("SELECT id FROM prompts WHERE id = ?", (prompt_id_int,)) + cursor = conn.execute( + "SELECT id FROM prompts WHERE id = ?", (prompt_id_int,) + ) if not cursor.fetchone(): - self.logger.warning(f"Prompt ID {prompt_id_int} not found in database, skipping image linking") + self.logger.warning( + f"Prompt ID {prompt_id_int} not found in database, skipping image linking" + ) return 0 - + # Proceed with linking filename = os.path.basename(image_path) - file_info = metadata.get('file_info', {}) if metadata else {} + file_info = metadata.get("file_info", {}) if metadata else {} # Use INSERT OR IGNORE to skip duplicates (same prompt_id + filename) cursor = conn.execute( @@ -1094,23 +1161,33 @@ class PromptDatabase: prompt_id_int, image_path, filename, - file_info.get('size'), - file_info.get('dimensions', [None, None])[0] if file_info.get('dimensions') else None, - file_info.get('dimensions', [None, None])[1] if file_info.get('dimensions') else None, - file_info.get('format'), - json.dumps(metadata.get('workflow', {}) if metadata else {}), - json.dumps(metadata.get('prompt', {}) if metadata else {}), - json.dumps(metadata.get('parameters', {}) if metadata else {}) - ) + file_info.get("size"), + ( + file_info.get("dimensions", [None, None])[0] + if file_info.get("dimensions") + else None + ), + ( + file_info.get("dimensions", [None, None])[1] + if file_info.get("dimensions") + else None + ), + file_info.get("format"), + json.dumps(metadata.get("workflow", {}) if metadata else {}), + json.dumps(metadata.get("prompt", {}) if metadata else {}), + json.dumps(metadata.get("parameters", {}) if metadata else {}), + ), ) conn.commit() if cursor.lastrowid == 0: - self.logger.debug(f"Image {filename} already linked to prompt {prompt_id_int}") + self.logger.debug( + f"Image {filename} already linked to prompt {prompt_id_int}" + ) return 0 return cursor.lastrowid - + except Exception as e: self.logger.error(f"Error linking image to prompt {prompt_id}: {e}") return 0 @@ -1132,17 +1209,17 @@ class PromptDatabase: WHERE prompt_id = ? ORDER BY generation_time DESC """, - (prompt_id,) + (prompt_id,), ) return [self._image_row_to_dict(row) for row in cursor.fetchall()] def get_recent_images(self, limit: int = 50) -> List[Dict[str, Any]]: """ Get recently generated images across all prompts. - + Args: limit: Maximum number of images to return - + Returns: List of image records with prompt text """ @@ -1155,7 +1232,7 @@ class PromptDatabase: ORDER BY gi.generation_time DESC LIMIT ? """, - (limit,) + (limit,), ) return [self._image_row_to_dict(row) for row in cursor.fetchall()] @@ -1194,20 +1271,20 @@ class PromptDatabase: result = [] for row in rows: data = self._image_row_to_dict(row) - if row['_prompt_tags_list']: - data['prompt_tags'] = row['_prompt_tags_list'].split('|||') + if row["_prompt_tags_list"]: + data["prompt_tags"] = row["_prompt_tags_list"].split("|||") else: - data['prompt_tags'] = [] + data["prompt_tags"] = [] result.append(data) return result def search_images_by_prompt(self, search_term: str) -> List[Dict[str, Any]]: """ Search images by prompt text. - + Args: search_term: Text to search for in prompt content - + Returns: List of image records with prompt text """ @@ -1220,24 +1297,23 @@ class PromptDatabase: WHERE p.text LIKE ? ORDER BY gi.generation_time DESC """, - (f"%{search_term}%",) + (f"%{search_term}%",), ) return [self._image_row_to_dict(row) for row in cursor.fetchall()] def get_image_by_id(self, image_id: int) -> Optional[Dict[str, Any]]: """ Get an image record by its ID. - + Args: image_id: The image ID - + Returns: Image record or None if not found """ with self.model.get_connection() as conn: cursor = conn.execute( - "SELECT * FROM generated_images WHERE id = ?", - (image_id,) + "SELECT * FROM generated_images WHERE id = ?", (image_id,) ) row = cursor.fetchone() return self._image_row_to_dict(row) if row else None @@ -1245,17 +1321,16 @@ class PromptDatabase: def delete_image(self, image_id: int) -> bool: """ Delete an image record by its ID. - + Args: image_id: The image ID to delete - + Returns: bool: True if deletion was successful """ with self.model.get_connection() as conn: cursor = conn.execute( - "DELETE FROM generated_images WHERE id = ?", - (image_id,) + "DELETE FROM generated_images WHERE id = ?", (image_id,) ) conn.commit() return cursor.rowcount > 0 @@ -1263,32 +1338,34 @@ class PromptDatabase: def cleanup_missing_images(self) -> int: """ Remove image records where the actual file no longer exists. - + Returns: int: Number of orphaned records removed """ removed_count = 0 - + with self.model.get_connection() as conn: cursor = conn.execute("SELECT id, image_path FROM generated_images") images = cursor.fetchall() - + for image in images: - if not os.path.exists(image['image_path']): - conn.execute("DELETE FROM generated_images WHERE id = ?", (image['id'],)) + if not os.path.exists(image["image_path"]): + conn.execute( + "DELETE FROM generated_images WHERE id = ?", (image["id"],) + ) removed_count += 1 - + conn.commit() - + return removed_count def _clean_nan_values(self, obj: Any) -> Any: """ Recursively clean NaN values from nested data structures. - + Args: obj: The object to clean (dict, list, or scalar) - + Returns: Cleaned object with NaN values replaced by None """ @@ -1296,78 +1373,86 @@ class PromptDatabase: return {key: self._clean_nan_values(value) for key, value in obj.items()} elif isinstance(obj, list): return [self._clean_nan_values(item) for item in obj] - elif isinstance(obj, float) and str(obj) == 'nan': + elif isinstance(obj, float) and str(obj) == "nan": return None else: return obj - def _merge_duplicate_metadata(self, conn: sqlite3.Connection, primary_id: int, duplicate_ids: List[int]) -> Dict[str, Any]: + def _merge_duplicate_metadata( + self, conn: sqlite3.Connection, primary_id: int, duplicate_ids: List[int] + ) -> Dict[str, Any]: """ Merge metadata from duplicate prompts, prioritizing non-empty values. - + Args: conn: Database connection primary_id: ID of the prompt to keep duplicate_ids: List of duplicate prompt IDs - + Returns: Dict containing merged metadata """ try: cursor = conn.execute( - "SELECT category, rating, notes FROM prompts WHERE id = ?", (primary_id,) + "SELECT category, rating, notes FROM prompts WHERE id = ?", + (primary_id,), ) primary_data = cursor.fetchone() if not primary_data: return {} merged = { - 'category': primary_data['category'], - 'tags': self._get_prompt_tags(conn, primary_id), - 'rating': primary_data['rating'], - 'notes': primary_data['notes'] or '', + "category": primary_data["category"], + "tags": self._get_prompt_tags(conn, primary_id), + "rating": primary_data["rating"], + "notes": primary_data["notes"] or "", } for dup_id in duplicate_ids: cursor = conn.execute( - "SELECT category, rating, notes FROM prompts WHERE id = ?", (dup_id,) + "SELECT category, rating, notes FROM prompts WHERE id = ?", + (dup_id,), ) dup_data = cursor.fetchone() if not dup_data: continue - if not merged['category'] and dup_data['category']: - merged['category'] = dup_data['category'] + if not merged["category"] and dup_data["category"]: + merged["category"] = dup_data["category"] dup_tags = self._get_prompt_tags(conn, dup_id) for tag in dup_tags: - if tag not in merged['tags']: - merged['tags'].append(tag) + if tag not in merged["tags"]: + merged["tags"].append(tag) - if dup_data['rating'] and (not merged['rating'] or dup_data['rating'] > merged['rating']): - merged['rating'] = dup_data['rating'] + if dup_data["rating"] and ( + not merged["rating"] or dup_data["rating"] > merged["rating"] + ): + merged["rating"] = dup_data["rating"] - if dup_data['notes'] and dup_data['notes'].strip(): - if merged['notes']: - merged['notes'] += f" | {dup_data['notes']}" + if dup_data["notes"] and dup_data["notes"].strip(): + if merged["notes"]: + merged["notes"] += f" | {dup_data['notes']}" else: - merged['notes'] = dup_data['notes'] + merged["notes"] = dup_data["notes"] return merged except Exception as e: self.logger.error(f"Error merging metadata: {e}", exc_info=True) return {} - - def _transfer_images_to_primary(self, conn: sqlite3.Connection, primary_id: int, duplicate_ids: List[int]) -> int: + + def _transfer_images_to_primary( + self, conn: sqlite3.Connection, primary_id: int, duplicate_ids: List[int] + ) -> int: """ Transfer all images from duplicate prompts to the primary prompt. - + Args: conn: Database connection primary_id: ID of the prompt to keep duplicate_ids: List of duplicate prompt IDs - + Returns: Number of images transferred """ @@ -1377,23 +1462,27 @@ class PromptDatabase: # Update all images to point to primary prompt cursor = conn.execute( "UPDATE generated_images SET prompt_id = ? WHERE prompt_id = ?", - (primary_id, dup_id) + (primary_id, dup_id), ) transferred_count += cursor.rowcount - + if cursor.rowcount > 0: - self.logger.debug(f"Transferred {cursor.rowcount} images from prompt {dup_id} to {primary_id}") - + self.logger.debug( + f"Transferred {cursor.rowcount} images from prompt {dup_id} to {primary_id}" + ) + return transferred_count - + except Exception as e: self.logger.error(f"Error transferring images: {e}", exc_info=True) return 0 - - def _update_primary_with_merged_metadata(self, conn: sqlite3.Connection, primary_id: int, merged_metadata: Dict[str, Any]) -> None: + + def _update_primary_with_merged_metadata( + self, conn: sqlite3.Connection, primary_id: int, merged_metadata: Dict[str, Any] + ) -> None: """ Update the primary prompt with merged metadata. - + Args: conn: Database connection primary_id: ID of the prompt to update @@ -1407,19 +1496,23 @@ class PromptDatabase: WHERE id = ? """, ( - merged_metadata.get('category'), - merged_metadata.get('rating'), - merged_metadata.get('notes'), + merged_metadata.get("category"), + merged_metadata.get("rating"), + merged_metadata.get("notes"), datetime.datetime.now(datetime.timezone.utc).isoformat(), primary_id, ), ) - self._sync_prompt_tags(conn, primary_id, merged_metadata.get('tags', [])) - self.logger.debug(f"Updated primary prompt {primary_id} with merged metadata") + self._sync_prompt_tags(conn, primary_id, merged_metadata.get("tags", [])) + self.logger.debug( + f"Updated primary prompt {primary_id} with merged metadata" + ) except Exception as e: - self.logger.error(f"Error updating primary prompt metadata: {e}", exc_info=True) - + self.logger.error( + f"Error updating primary prompt metadata: {e}", exc_info=True + ) + # ------------------------------------------------------------------ # Methods extracted from api.py raw SQL (Fix 3.1) # ------------------------------------------------------------------ @@ -1481,7 +1574,11 @@ class PromptDatabase: with self.model.get_connection() as conn: cursor = conn.execute( "UPDATE prompts SET text = ?, updated_at = ? WHERE id = ?", - (new_text, datetime.datetime.now(datetime.timezone.utc).isoformat(), prompt_id), + ( + new_text, + datetime.datetime.now(datetime.timezone.utc).isoformat(), + prompt_id, + ), ) conn.commit() return cursor.rowcount > 0 @@ -1495,7 +1592,11 @@ class PromptDatabase: with self.model.get_connection() as conn: cursor = conn.execute( "UPDATE prompts SET rating = ?, updated_at = ? WHERE id = ?", - (rating, datetime.datetime.now(datetime.timezone.utc).isoformat(), prompt_id), + ( + rating, + datetime.datetime.now(datetime.timezone.utc).isoformat(), + prompt_id, + ), ) conn.commit() return cursor.rowcount > 0 @@ -1590,7 +1691,7 @@ class PromptDatabase: FROM generated_images gi JOIN prompts p ON gi.prompt_id = p.id WHERE gi.image_path = ? OR gi.image_path LIKE ?""", - (image_path, f'%{os.path.basename(image_path)}'), + (image_path, f"%{os.path.basename(image_path)}"), ) row = cursor.fetchone() if not row: @@ -1694,17 +1795,17 @@ class PromptDatabase: def _image_row_to_dict(self, row: sqlite3.Row) -> Dict[str, Any]: """ Convert an image database row to a dictionary with parsed JSON fields. - + Args: row: SQLite row object - + Returns: Dictionary representation of the row """ data = dict(row) - + # Parse JSON fields - for field in ['workflow_data', 'prompt_metadata', 'parameters']: + for field in ["workflow_data", "prompt_metadata", "parameters"]: if data.get(field): try: parsed_data = json.loads(data[field]) @@ -1713,5 +1814,5 @@ class PromptDatabase: data[field] = {} else: data[field] = {} - - return data \ No newline at end of file + + return data diff --git a/prompt_manager.py b/prompt_manager.py index c3ad1db..5259d71 100644 --- a/prompt_manager.py +++ b/prompt_manager.py @@ -228,8 +228,17 @@ class PromptManager(PromptManagerBase, ComfyNodeABC): return (conditioning, encoding_text) @classmethod - def IS_CHANGED(cls, clip, text="", category="", tags="", search_text="", - prepend_text="", append_text="", **kwargs): + def IS_CHANGED( + cls, + clip, + text="", + category="", + tags="", + search_text="", + prepend_text="", + append_text="", + **kwargs, + ): """ ComfyUI method to determine if node needs re-execution. diff --git a/prompt_manager_base.py b/prompt_manager_base.py index db46a66..75106f8 100644 --- a/prompt_manager_base.py +++ b/prompt_manager_base.py @@ -144,6 +144,7 @@ class PromptManagerBase: try: from .py.config import PromptManagerConfig + max_results = PromptManagerConfig.MAX_SEARCH_RESULTS except Exception: max_results = 100 diff --git a/prompt_manager_text.py b/prompt_manager_text.py index 7932bfe..7c41051 100644 --- a/prompt_manager_text.py +++ b/prompt_manager_text.py @@ -201,8 +201,16 @@ class PromptManagerText(PromptManagerBase, ComfyNodeABC): return (final_text,) @classmethod - def IS_CHANGED(cls, text="", category="", tags="", search_text="", - prepend_text="", append_text="", **kwargs): + def IS_CHANGED( + cls, + text="", + category="", + tags="", + search_text="", + prepend_text="", + append_text="", + **kwargs, + ): """ ComfyUI method to determine if node needs re-execution. diff --git a/py/__init__.py b/py/__init__.py index d8a0943..b2287f8 100644 --- a/py/__init__.py +++ b/py/__init__.py @@ -3,4 +3,4 @@ ComfyUI_PromptManager Python API modules. This package contains the core API components for the web interface and configuration management, including REST endpoints, configuration handling, and server integration. -""" \ No newline at end of file +""" diff --git a/py/api/__init__.py b/py/api/__init__.py index c3c30af..2495df7 100644 --- a/py/api/__init__.py +++ b/py/api/__init__.py @@ -40,20 +40,20 @@ except ImportError: def _get_project_root(): """Get the project root directory (3 levels up from py/api/__init__.py).""" - return os.path.dirname( - os.path.dirname(os.path.dirname(os.path.abspath(__file__))) - ) + return os.path.dirname(os.path.dirname(os.path.dirname(os.path.abspath(__file__)))) # ── Gzip compression middleware ──────────────────────────────────────── _GZIP_MIN_SIZE = 1024 # Only compress bodies larger than 1 KB -_GZIP_TYPES = frozenset(( - 'application/json', - 'text/html', - 'text/css', - 'application/javascript', - 'text/plain', -)) +_GZIP_TYPES = frozenset( + ( + "application/json", + "text/html", + "text/css", + "application/javascript", + "text/plain", + ) +) _gzip_registered = False @@ -62,24 +62,24 @@ async def _gzip_middleware(request, handler): """Compress PromptManager responses when the client accepts gzip.""" response = await handler(request) - if not request.path.startswith('/prompt_manager/'): + if not request.path.startswith("/prompt_manager/"): return response # Only handle regular Response objects (not StreamResponse/WebSocket) if not isinstance(response, web.Response): return response - if 'gzip' not in request.headers.get('Accept-Encoding', ''): + if "gzip" not in request.headers.get("Accept-Encoding", ""): return response - if 'Content-Encoding' in response.headers: + if "Content-Encoding" in response.headers: return response body = response.body if body is None or len(body) < _GZIP_MIN_SIZE: return response - content_type = response.content_type or '' + content_type = response.content_type or "" if not any(ct in content_type for ct in _GZIP_TYPES): return response @@ -88,8 +88,8 @@ async def _gzip_middleware(request, handler): return response response.body = compressed - response.headers['Content-Encoding'] = 'gzip' - response.headers['Vary'] = 'Accept-Encoding' + response.headers["Content-Encoding"] = "gzip" + response.headers["Vary"] = "Accept-Encoding" return response @@ -117,7 +117,7 @@ class PromptManagerAPI( def __init__(self): """Initialize the PromptManager API with database connection and cleanup.""" - self.logger = get_logger('prompt_manager.api') + self.logger = get_logger("prompt_manager.api") self.logger.info("Initializing PromptManager API") self.db = PromptDatabase() @@ -131,7 +131,9 @@ class PromptManagerAPI( try: removed = self.db.cleanup_duplicates() if removed > 0: - self.logger.info(f"Startup cleanup: removed {removed} duplicate prompts") + self.logger.info( + f"Startup cleanup: removed {removed} duplicate prompts" + ) except Exception as e: self.logger.error(f"Startup cleanup failed: {e}") @@ -207,7 +209,9 @@ class PromptManagerAPI( async def serve_gallery_ui(request): try: html_path = os.path.join( - _get_project_root(), "web", "metadata.html", + _get_project_root(), + "web", + "metadata.html", ) if html_path not in self._html_cache: @@ -222,7 +226,9 @@ class PromptManagerAPI( ) return web.Response( - text=self._html_cache[html_path], content_type="text/html", charset="utf-8" + text=self._html_cache[html_path], + content_type="text/html", + charset="utf-8", ) except Exception as e: @@ -236,7 +242,9 @@ class PromptManagerAPI( async def serve_admin_ui(request): try: html_path = os.path.join( - _get_project_root(), "web", "admin.html", + _get_project_root(), + "web", + "admin.html", ) if html_path not in self._html_cache: @@ -251,7 +259,9 @@ class PromptManagerAPI( ) return web.Response( - text=self._html_cache[html_path], content_type="text/html", charset="utf-8" + text=self._html_cache[html_path], + content_type="text/html", + charset="utf-8", ) except Exception as e: @@ -265,7 +275,9 @@ class PromptManagerAPI( async def serve_gallery_admin_ui(request): try: html_path = os.path.join( - _get_project_root(), "web", "gallery.html", + _get_project_root(), + "web", + "gallery.html", ) if html_path not in self._html_cache: @@ -280,7 +292,9 @@ class PromptManagerAPI( ) return web.Response( - text=self._html_cache[html_path], content_type="text/html", charset="utf-8" + text=self._html_cache[html_path], + content_type="text/html", + charset="utf-8", ) except Exception as e: @@ -362,6 +376,7 @@ class PromptManagerAPI( if not _gzip_registered: try: from ..config import server_instance + server_instance.app.middlewares.append(_gzip_middleware) _gzip_registered = True self.logger.info("Gzip compression middleware registered") @@ -380,30 +395,36 @@ class PromptManagerAPI( output_path = Path(output_dir) if output_dir else None for prompt in prompts: - for image in prompt.get('images', []): - image_path_str = image.get('image_path', '') + for image in prompt.get("images", []): + image_path_str = image.get("image_path", "") if not image_path_str: continue img_path = Path(image_path_str) # Set fallback url via image ID - if image.get('id'): - image['url'] = f"/prompt_manager/images/{image['id']}/file" + if image.get("id"): + image["url"] = f"/prompt_manager/images/{image['id']}/file" # Try to compute relative path and thumbnail URL if output_path: try: rel_path = img_path.resolve().relative_to(output_path.resolve()) - image['relative_path'] = str(rel_path) - image['url'] = f"/prompt_manager/images/serve/{url_quote(rel_path.as_posix(), safe='/')}" + image["relative_path"] = str(rel_path) + image["url"] = ( + f"/prompt_manager/images/serve/{url_quote(rel_path.as_posix(), safe='/')}" + ) # Check for thumbnail - rel_no_ext = rel_path.with_suffix('') - thumb_rel = f"thumbnails/{rel_no_ext.as_posix()}_thumb{rel_path.suffix}" + rel_no_ext = rel_path.with_suffix("") + thumb_rel = ( + f"thumbnails/{rel_no_ext.as_posix()}_thumb{rel_path.suffix}" + ) thumb_abs = output_path / thumb_rel if thumb_abs.exists(): - image['thumbnail_url'] = f"/prompt_manager/images/serve/{url_quote(thumb_rel, safe='/')}" + image["thumbnail_url"] = ( + f"/prompt_manager/images/serve/{url_quote(thumb_rel, safe='/')}" + ) except (ValueError, RuntimeError): pass @@ -415,7 +436,7 @@ class PromptManagerAPI( return {key: self._clean_nan_recursive(value) for key, value in obj.items()} elif isinstance(obj, list): return [self._clean_nan_recursive(item) for item in obj] - elif isinstance(obj, float) and str(obj) == 'nan': + elif isinstance(obj, float) and str(obj) == "nan": return None else: return obj @@ -442,11 +463,13 @@ class PromptManagerAPI( for i in range(max_depth): # Check if current directory contains ComfyUI markers - comfyui_markers = ['main.py', 'nodes.py', 'server.py'] + comfyui_markers = ["main.py", "nodes.py", "server.py"] if any((current_dir / marker).exists() for marker in comfyui_markers): output_dir = current_dir / "output" if output_dir.exists() and output_dir.is_dir(): - self.logger.debug(f"Found ComfyUI output directory via upward search: {output_dir}") + self.logger.debug( + f"Found ComfyUI output directory via upward search: {output_dir}" + ) self._cached_output_dir = str(output_dir) return self._cached_output_dir @@ -500,7 +523,7 @@ class PromptManagerAPI( try: with Image.open(image_path) as img: metadata = {} - if hasattr(img, 'text'): + if hasattr(img, "text"): for key, value in img.text.items(): metadata[key] = value return metadata @@ -511,11 +534,11 @@ class PromptManagerAPI( def _parse_comfyui_prompt(self, metadata): """Parse ComfyUI and A1111 prompt data from metadata.""" result = { - 'prompt': None, - 'workflow': None, - 'parameters': {}, - 'positive_prompt': None, - 'negative_prompt': None + "prompt": None, + "workflow": None, + "parameters": {}, + "positive_prompt": None, + "negative_prompt": None, } # Check for A1111 style parameters first (like parse-metadata.py) @@ -523,40 +546,48 @@ class PromptManagerAPI( params = metadata["parameters"] lines = params.splitlines() if lines: - result['positive_prompt'] = lines[0].strip() + result["positive_prompt"] = lines[0].strip() for line in lines: if line.lower().startswith("negative prompt:"): - result['negative_prompt'] = line.split(":", 1)[1].strip() + result["negative_prompt"] = line.split(":", 1)[1].strip() break # Store raw parameters too - result['parameters']['parameters'] = params + result["parameters"]["parameters"] = params # If no A1111 format found, proceed with ComfyUI parsing - if result['positive_prompt'] is None: + if result["positive_prompt"] is None: # Check for direct prompt field - if 'prompt' in metadata: + if "prompt" in metadata: try: - prompt_data = json.loads(metadata['prompt']) - result['prompt'] = prompt_data + prompt_data = json.loads(metadata["prompt"]) + result["prompt"] = prompt_data except json.JSONDecodeError: - result['prompt'] = metadata['prompt'] + result["prompt"] = metadata["prompt"] # Check for workflow - if 'workflow' in metadata: + if "workflow" in metadata: try: - workflow_data = json.loads(metadata['workflow']) - result['workflow'] = workflow_data + workflow_data = json.loads(metadata["workflow"]) + result["workflow"] = workflow_data except json.JSONDecodeError: - result['workflow'] = metadata['workflow'] + result["workflow"] = metadata["workflow"] # Check for other common ComfyUI fields - common_fields = ['positive', 'negative', 'steps', 'cfg', 'sampler', 'scheduler', 'seed'] + common_fields = [ + "positive", + "negative", + "steps", + "cfg", + "sampler", + "scheduler", + "seed", + ] for field in common_fields: if field in metadata: try: - result['parameters'][field] = json.loads(metadata[field]) + result["parameters"][field] = json.loads(metadata[field]) except json.JSONDecodeError: - result['parameters'][field] = metadata[field] + result["parameters"][field] = metadata[field] return result @@ -567,40 +598,46 @@ class PromptManagerAPI( if isinstance(value, str): return value elif isinstance(value, list): - return ' '.join(str(item) for item in value if item) + return " ".join(str(item) for item in value if item) elif value is not None: return str(value) return None # First check if we already extracted a positive prompt (A1111 format) - if parsed_data.get('positive_prompt'): - return parsed_data['positive_prompt'] + if parsed_data.get("positive_prompt"): + return parsed_data["positive_prompt"] # Check if prompt is already a string - if isinstance(parsed_data.get('prompt'), str): - return parsed_data['prompt'] + if isinstance(parsed_data.get("prompt"), str): + return parsed_data["prompt"] # Check if prompt is a simple value that can be converted - if parsed_data.get('prompt') and not isinstance(parsed_data.get('prompt'), dict): - return safe_to_string(parsed_data['prompt']) + if parsed_data.get("prompt") and not isinstance( + parsed_data.get("prompt"), dict + ): + return safe_to_string(parsed_data["prompt"]) - prompt_data = parsed_data.get('prompt') + prompt_data = parsed_data.get("prompt") if isinstance(prompt_data, dict): # Use the enhanced logic from parse-metadata.py - positive_prompt = self._extract_positive_prompt_from_comfyui_data(prompt_data) + positive_prompt = self._extract_positive_prompt_from_comfyui_data( + prompt_data + ) if positive_prompt: return positive_prompt # Check workflow data if available - workflow_data = parsed_data.get('workflow') + workflow_data = parsed_data.get("workflow") if isinstance(workflow_data, dict): - positive_prompt = self._extract_positive_prompt_from_comfyui_data(workflow_data) + positive_prompt = self._extract_positive_prompt_from_comfyui_data( + workflow_data + ) if positive_prompt: return positive_prompt # Check parameters for positive prompt - if parsed_data.get('parameters', {}).get('positive'): - return safe_to_string(parsed_data['parameters']['positive']) + if parsed_data.get("parameters", {}).get("positive"): + return safe_to_string(parsed_data["parameters"]["positive"]) return None @@ -652,12 +689,21 @@ class PromptManagerAPI( # Strategy 2: For text encoder nodes, check widgets_values class_type = node.get("class_type", node.get("type", "")) text_encoder_types = [ - 'CLIPTextEncode', 'CLIPTextEncodeSDXL', 'CLIPTextEncodeSDXLRefiner', - 'CLIPTextEncodeFlux', 'PromptManager', 'PromptManagerText', - 'BNK_CLIPTextEncoder', 'Text Encoder', 'CLIP Text Encode' + "CLIPTextEncode", + "CLIPTextEncodeSDXL", + "CLIPTextEncodeSDXLRefiner", + "CLIPTextEncodeFlux", + "PromptManager", + "PromptManagerText", + "BNK_CLIPTextEncoder", + "Text Encoder", + "CLIP Text Encode", ] - if any(encoder_type.lower() in class_type.lower() for encoder_type in text_encoder_types): + if any( + encoder_type.lower() in class_type.lower() + for encoder_type in text_encoder_types + ): widgets_values = node.get("widgets_values", []) if widgets_values and len(widgets_values) > 0: if isinstance(widgets_values[0], str) and widgets_values[0].strip(): @@ -721,8 +767,8 @@ class PromptManagerAPI( text_val = self._find_text_in_node(node) if text_val: # Try to determine if this is positive or negative - node_title = node.get('title', '').lower() - if 'neg' not in node_title and 'negative' not in node_title: + node_title = node.get("title", "").lower() + if "neg" not in node_title and "negative" not in node_title: # Prioritize non-negative prompts text_nodes.insert(0, text_val) else: diff --git a/py/api/admin.py b/py/api/admin.py index c64c1d7..09de150 100644 --- a/py/api/admin.py +++ b/py/api/admin.py @@ -82,7 +82,10 @@ class AdminRoutesMixin: except Exception as e: self.logger.error(f"Scan duplicates error: {e}") return web.json_response( - {"success": False, "error": f"Failed to scan duplicate images: {str(e)}"}, + { + "success": False, + "error": f"Failed to scan duplicate images: {str(e)}", + }, status=500, ) @@ -118,8 +121,16 @@ class AdminRoutesMixin: output_path = Path(output_dir) - image_extensions = ['.png', '.jpg', '.jpeg', '.gif', '.webp', '.bmp', '.tiff'] - video_extensions = ['.mp4', '.avi', '.mov', '.mkv', '.webm', '.gif'] + image_extensions = [ + ".png", + ".jpg", + ".jpeg", + ".gif", + ".webp", + ".bmp", + ".tiff", + ] + video_extensions = [".mp4", ".avi", ".mov", ".mkv", ".webm", ".gif"] all_extensions = image_extensions + video_extensions media_files = [] @@ -127,7 +138,7 @@ class AdminRoutesMixin: for ext in all_extensions: for pattern in [f"*{ext.lower()}", f"*{ext.upper()}"]: for media_path in output_path.rglob(pattern): - if 'thumbnails' not in media_path.parts: + if "thumbnails" not in media_path.parts: normalized_path = str(media_path).lower() if normalized_path not in seen_paths: seen_paths.add(normalized_path) @@ -149,41 +160,44 @@ class AdminRoutesMixin: rel_path = media_path.relative_to(output_path) extension = media_path.suffix.lower() is_video = extension in [ext.lower() for ext in video_extensions] - media_type = 'video' if is_video else 'image' + media_type = "video" if is_video else "image" # Check if thumbnail exists thumbnail_url = None thumbnails_dir = output_path / "thumbnails" if thumbnails_dir.exists(): - thumbnail_ext = '.jpg' if is_video else extension - rel_path_no_ext = rel_path.with_suffix('') + thumbnail_ext = ".jpg" if is_video else extension + rel_path_no_ext = rel_path.with_suffix("") thumbnail_rel_path = f"thumbnails/{rel_path_no_ext.as_posix()}_thumb{thumbnail_ext}" thumbnail_abs_path = output_path / thumbnail_rel_path if thumbnail_abs_path.exists(): from urllib.parse import quote + thumbnail_url = f'/prompt_manager/images/serve/{quote(thumbnail_rel_path, safe="/")}' image_info = { - 'id': str(hash(str(media_path))), - 'filename': media_path.name, - 'path': str(media_path), - 'relative_path': str(rel_path), - 'url': f'/prompt_manager/images/serve/{rel_path.as_posix()}', - 'thumbnail_url': thumbnail_url, - 'size': stat.st_size, - 'modified_time': stat.st_mtime, - 'extension': extension, - 'media_type': media_type, - 'is_video': is_video, - 'hash': file_hash + "id": str(hash(str(media_path))), + "filename": media_path.name, + "path": str(media_path), + "relative_path": str(rel_path), + "url": f"/prompt_manager/images/serve/{rel_path.as_posix()}", + "thumbnail_url": thumbnail_url, + "size": stat.st_size, + "modified_time": stat.st_mtime, + "extension": extension, + "media_type": media_type, + "is_video": is_video, + "hash": file_hash, } file_hashes[file_hash].append(image_info) processed += 1 if processed % 100 == 0: - self.logger.info(f"Processed {processed}/{len(media_files)} files for duplicate detection") + self.logger.info( + f"Processed {processed}/{len(media_files)} files for duplicate detection" + ) except Exception as e: self.logger.error(f"Error processing file {media_path}: {e}") @@ -193,12 +207,10 @@ class AdminRoutesMixin: duplicates = [] for file_hash, images in file_hashes.items(): if len(images) > 1: - images.sort(key=lambda x: x['modified_time']) - duplicates.append({ - 'hash': file_hash, - 'images': images, - 'count': len(images) - }) + images.sort(key=lambda x: x["modified_time"]) + duplicates.append( + {"hash": file_hash, "images": images, "count": len(images)} + ) self.logger.info(f"Found {len(duplicates)} groups of duplicate images") return duplicates @@ -219,7 +231,7 @@ class AdminRoutesMixin: """Delete duplicate image files from disk.""" try: data = await request.json() - image_paths = data.get('image_paths', []) + image_paths = data.get("image_paths", []) if not image_paths: return web.json_response( @@ -236,7 +248,9 @@ class AdminRoutesMixin: # Ensure the path is within the output directory for security output_dir = self._find_comfyui_output_dir() if not output_dir: - failed_files.append(f"{image_path} (output directory not found)") + failed_files.append( + f"{image_path} (output directory not found)" + ) failed_count += 1 continue @@ -247,7 +261,9 @@ class AdminRoutesMixin: try: file_path.resolve().relative_to(output_path.resolve()) except ValueError: - self.logger.warning(f"Attempted to delete file outside output directory: {image_path}") + self.logger.warning( + f"Attempted to delete file outside output directory: {image_path}" + ) failed_files.append(f"{image_path} (outside output directory)") failed_count += 1 continue @@ -260,13 +276,21 @@ class AdminRoutesMixin: # Also try to delete associated thumbnail if it exists try: rel_path = file_path.relative_to(output_path) - rel_path_no_ext = rel_path.with_suffix('') - thumbnail_path = output_path / "thumbnails" / f"{rel_path_no_ext.as_posix()}_thumb{file_path.suffix}" + rel_path_no_ext = rel_path.with_suffix("") + thumbnail_path = ( + output_path + / "thumbnails" + / f"{rel_path_no_ext.as_posix()}_thumb{file_path.suffix}" + ) if thumbnail_path.exists(): os.remove(thumbnail_path) - self.logger.debug(f"Deleted associated thumbnail: {thumbnail_path}") + self.logger.debug( + f"Deleted associated thumbnail: {thumbnail_path}" + ) except Exception as e: - self.logger.warning(f"Could not delete thumbnail for {image_path}: {e}") + self.logger.warning( + f"Could not delete thumbnail for {image_path}: {e}" + ) else: failed_files.append(f"{image_path} (file not found)") failed_count += 1 @@ -280,7 +304,7 @@ class AdminRoutesMixin: "success": True, "deleted_count": deleted_count, "failed_count": failed_count, - "message": f"Deleted {deleted_count} files successfully" + "message": f"Deleted {deleted_count} files successfully", } if failed_count > 0: @@ -292,7 +316,10 @@ class AdminRoutesMixin: except Exception as e: self.logger.error(f"Delete duplicate images error: {e}") return web.json_response( - {"success": False, "error": f"Failed to delete duplicate images: {str(e)}"}, + { + "success": False, + "error": f"Failed to delete duplicate images: {str(e)}", + }, status=500, ) @@ -319,29 +346,40 @@ class AdminRoutesMixin: monitored_dirs = [] try: import sys + monitor_module = None for mod_name in list(sys.modules.keys()): - if 'image_monitor' in mod_name and hasattr(sys.modules[mod_name], '_monitor_instance'): + if "image_monitor" in mod_name and hasattr( + sys.modules[mod_name], "_monitor_instance" + ): monitor_module = sys.modules[mod_name] break if monitor_module and monitor_module._monitor_instance is not None: - monitored_dirs = getattr(monitor_module._monitor_instance, 'monitored_directories', []) + monitored_dirs = getattr( + monitor_module._monitor_instance, "monitored_directories", [] + ) elif GalleryConfig.MONITORING_DIRECTORIES: monitored_dirs = GalleryConfig.MONITORING_DIRECTORIES except Exception: if GalleryConfig.MONITORING_DIRECTORIES: monitored_dirs = GalleryConfig.MONITORING_DIRECTORIES - return web.json_response({ - "success": True, - "settings": { - "result_timeout": PromptManagerConfig.RESULT_TIMEOUT, - "webui_display_mode": PromptManagerConfig.WEBUI_DISPLAY_MODE, - "gallery_root_path": GalleryConfig.MONITORING_DIRECTORIES[0] if GalleryConfig.MONITORING_DIRECTORIES else "", - "monitored_directories": monitored_dirs + return web.json_response( + { + "success": True, + "settings": { + "result_timeout": PromptManagerConfig.RESULT_TIMEOUT, + "webui_display_mode": PromptManagerConfig.WEBUI_DISPLAY_MODE, + "gallery_root_path": ( + GalleryConfig.MONITORING_DIRECTORIES[0] + if GalleryConfig.MONITORING_DIRECTORIES + else "" + ), + "monitored_directories": monitored_dirs, + }, } - }) + ) except Exception as e: return web.json_response( {"success": False, "error": f"Failed to get settings: {str(e)}"}, @@ -357,66 +395,91 @@ class AdminRoutesMixin: restart_required = False # Update in-memory config - if 'result_timeout' in data: - PromptManagerConfig.RESULT_TIMEOUT = data['result_timeout'] - if 'webui_display_mode' in data: - PromptManagerConfig.WEBUI_DISPLAY_MODE = data['webui_display_mode'] + if "result_timeout" in data: + PromptManagerConfig.RESULT_TIMEOUT = data["result_timeout"] + if "webui_display_mode" in data: + PromptManagerConfig.WEBUI_DISPLAY_MODE = data["webui_display_mode"] # Handle gallery root path - if 'gallery_root_path' in data: - new_path = data['gallery_root_path'].strip() - old_path = GalleryConfig.MONITORING_DIRECTORIES[0] if GalleryConfig.MONITORING_DIRECTORIES else "" + if "gallery_root_path" in data: + new_path = data["gallery_root_path"].strip() + old_path = ( + GalleryConfig.MONITORING_DIRECTORIES[0] + if GalleryConfig.MONITORING_DIRECTORIES + else "" + ) if new_path != old_path: if new_path: from pathlib import Path as _Path + resolved = _Path(new_path).resolve() if not resolved.is_dir(): - return web.json_response({ - 'success': False, - 'error': f'Gallery path does not exist or is not a directory: {new_path}' - }, status=400) - blocked = ['/etc', '/usr', '/bin', '/sbin', '/boot', '/proc', '/sys', '/dev', - '/var/log', '/root', 'C:\\Windows', 'C:\\Program Files'] + return web.json_response( + { + "success": False, + "error": f"Gallery path does not exist or is not a directory: {new_path}", + }, + status=400, + ) + blocked = [ + "/etc", + "/usr", + "/bin", + "/sbin", + "/boot", + "/proc", + "/sys", + "/dev", + "/var/log", + "/root", + "C:\\Windows", + "C:\\Program Files", + ] for b in blocked: if str(resolved).startswith(b): - return web.json_response({ - 'success': False, - 'error': 'Gallery path cannot point to a system directory' - }, status=400) + return web.json_response( + { + "success": False, + "error": "Gallery path cannot point to a system directory", + }, + status=400, + ) GalleryConfig.MONITORING_DIRECTORIES = [new_path] else: GalleryConfig.MONITORING_DIRECTORIES = [] restart_required = True # Save to config file for persistence - config_dir = os.path.dirname(os.path.dirname(os.path.dirname(os.path.abspath(__file__)))) - config_file = os.path.join(config_dir, 'config.json') + config_dir = os.path.dirname( + os.path.dirname(os.path.dirname(os.path.abspath(__file__))) + ) + config_file = os.path.join(config_dir, "config.json") config_data = { - 'web_ui': { - 'result_timeout': PromptManagerConfig.RESULT_TIMEOUT, - 'webui_display_mode': PromptManagerConfig.WEBUI_DISPLAY_MODE + "web_ui": { + "result_timeout": PromptManagerConfig.RESULT_TIMEOUT, + "webui_display_mode": PromptManagerConfig.WEBUI_DISPLAY_MODE, + }, + "gallery": { + "monitoring": {"directories": GalleryConfig.MONITORING_DIRECTORIES} }, - 'gallery': { - 'monitoring': { - 'directories': GalleryConfig.MONITORING_DIRECTORIES - } - } } try: - with open(config_file, 'w') as f: + with open(config_file, "w") as f: json.dump(config_data, f, indent=2) self.logger.info(f"Settings saved to {config_file}") except Exception as save_err: self.logger.warning(f"Could not save config file: {save_err}") - return web.json_response({ - "success": True, - "message": "Settings saved successfully", - "restart_required": restart_required - }) + return web.json_response( + { + "success": True, + "message": "Settings saved successfully", + "restart_required": restart_required, + } + ) except Exception as e: return web.json_response( {"success": False, "error": f"Failed to save settings: {str(e)}"}, @@ -435,53 +498,59 @@ class AdminRoutesMixin: with sqlite3.connect(db_path) as conn: conn.row_factory = sqlite3.Row cursor = conn.execute("SELECT COUNT(*) as count FROM prompts") - prompt_count = cursor.fetchone()['count'] + prompt_count = cursor.fetchone()["count"] - cursor = conn.execute("SELECT name FROM sqlite_master WHERE type='table' AND name='generated_images'") + cursor = conn.execute( + "SELECT name FROM sqlite_master WHERE type='table' AND name='generated_images'" + ) has_images_table = cursor.fetchone() is not None if has_images_table: - cursor = conn.execute("SELECT COUNT(*) as count FROM generated_images") - image_count = cursor.fetchone()['count'] + cursor = conn.execute( + "SELECT COUNT(*) as count FROM generated_images" + ) + image_count = cursor.fetchone()["count"] else: image_count = 0 - results['database'] = { - 'status': 'ok', - 'prompt_count': prompt_count, - 'has_images_table': has_images_table, - 'image_count': image_count + results["database"] = { + "status": "ok", + "prompt_count": prompt_count, + "has_images_table": has_images_table, + "image_count": image_count, } else: - results['database'] = { - 'status': 'error', - 'message': f'Database file not found: {db_path}' + results["database"] = { + "status": "error", + "message": f"Database file not found: {db_path}", } except Exception as e: - results['database'] = { - 'status': 'error', - 'message': f'Database error: {str(e)}' + results["database"] = { + "status": "error", + "message": f"Database error: {str(e)}", } # Check dependencies dependencies = {} try: import watchdog - dependencies['watchdog'] = True + + dependencies["watchdog"] = True except ImportError: - dependencies['watchdog'] = False + dependencies["watchdog"] = False try: from PIL import Image - dependencies['PIL'] = True + + dependencies["PIL"] = True except ImportError: - dependencies['PIL'] = False + dependencies["PIL"] = False - dependencies['sqlite3'] = True # Always available in Python + dependencies["sqlite3"] = True # Always available in Python - results['dependencies'] = { - 'status': 'ok' if all(dependencies.values()) else 'error', - 'dependencies': dependencies + results["dependencies"] = { + "status": "ok" if all(dependencies.values()) else "error", + "dependencies": dependencies, } # Check output directories @@ -493,64 +562,60 @@ class AdminRoutesMixin: if os.path.exists(abs_path): output_dirs.append(abs_path) - results['comfyui_output'] = { - 'status': 'ok' if output_dirs else 'warning', - 'output_dirs': output_dirs + results["comfyui_output"] = { + "status": "ok" if output_dirs else "warning", + "output_dirs": output_dirs, } # Check image monitor status try: from ...utils.image_monitor import _monitor_instance + if _monitor_instance is not None: monitor_status = _monitor_instance.get_status() - results['image_monitor'] = { - 'status': 'ok' if monitor_status.get('observer_alive') else 'error', - **monitor_status + results["image_monitor"] = { + "status": ( + "ok" if monitor_status.get("observer_alive") else "error" + ), + **monitor_status, } else: - results['image_monitor'] = { - 'status': 'error', - 'message': 'Image monitor not initialized' + results["image_monitor"] = { + "status": "error", + "message": "Image monitor not initialized", } except Exception as e: - results['image_monitor'] = { - 'status': 'error', - 'message': f'Failed to get monitor status: {str(e)}' + results["image_monitor"] = { + "status": "error", + "message": f"Failed to get monitor status: {str(e)}", } - return web.json_response({ - 'success': True, - 'diagnostics': results - }) + return web.json_response({"success": True, "diagnostics": results}) except Exception as e: self.logger.error(f"Diagnostics error: {e}", exc_info=True) - return web.json_response({ - 'success': False, - 'error': str(e) - }, status=500) + return web.json_response({"success": False, "error": str(e)}, status=500) async def test_image_link(self, request): """Test creating an image link.""" try: data = await request.json() - prompt_id = data.get('prompt_id') - test_image_path = data.get('image_path', '/test/fake/image.png') + prompt_id = data.get("prompt_id") + test_image_path = data.get("image_path", "/test/fake/image.png") if not prompt_id: - return web.json_response({ - 'success': False, - 'error': 'prompt_id is required' - }, status=400) + return web.json_response( + {"success": False, "error": "prompt_id is required"}, status=400 + ) test_metadata = { - 'file_info': { - 'size': 1024000, - 'dimensions': [512, 512], - 'format': 'PNG' + "file_info": { + "size": 1024000, + "dimensions": [512, 512], + "format": "PNG", }, - 'workflow': {'test': True}, - 'prompt': {'test_prompt': 'This is a test image'} + "workflow": {"test": True}, + "prompt": {"test_prompt": "This is a test image"}, } try: @@ -558,125 +623,176 @@ class AdminRoutesMixin: self.db.link_image_to_prompt, prompt_id=str(prompt_id), image_path=test_image_path, - metadata=test_metadata + metadata=test_metadata, ) - return web.json_response({ - 'success': True, - 'result': { - 'status': 'ok', - 'image_id': image_id, - 'message': f'Test image linked successfully with ID {image_id}' + return web.json_response( + { + "success": True, + "result": { + "status": "ok", + "image_id": image_id, + "message": f"Test image linked successfully with ID {image_id}", + }, } - }) + ) except Exception as e: - return web.json_response({ - 'success': False, - 'result': { - 'status': 'error', - 'message': f'Failed to create test link: {str(e)}' + return web.json_response( + { + "success": False, + "result": { + "status": "error", + "message": f"Failed to create test link: {str(e)}", + }, } - }) + ) except Exception as e: self.logger.error(f"Test link error: {e}", exc_info=True) - return web.json_response({ - 'success': False, - 'error': str(e) - }, status=500) + return web.json_response({"success": False, "error": str(e)}, status=500) async def run_maintenance(self, request): """Perform comprehensive database maintenance and optimization.""" try: - data = await request.json() if request.content_type == 'application/json' else {} - operations = data.get('operations', ['cleanup_duplicates', 'vacuum', 'cleanup_orphaned_images']) + data = ( + await request.json() + if request.content_type == "application/json" + else {} + ) + operations = data.get( + "operations", + ["cleanup_duplicates", "vacuum", "cleanup_orphaned_images"], + ) results = {} def _run_maintenance(): - if 'cleanup_duplicates' in operations: + if "cleanup_duplicates" in operations: try: duplicates_removed = self.db.cleanup_duplicates() - results['cleanup_duplicates'] = { - 'success': True, 'removed_count': duplicates_removed, - 'message': f'Removed {duplicates_removed} duplicate prompts' + results["cleanup_duplicates"] = { + "success": True, + "removed_count": duplicates_removed, + "message": f"Removed {duplicates_removed} duplicate prompts", } except Exception as e: - results['cleanup_duplicates'] = {'success': False, 'error': str(e), 'message': 'Failed to cleanup duplicates'} + results["cleanup_duplicates"] = { + "success": False, + "error": str(e), + "message": "Failed to cleanup duplicates", + } - if 'vacuum' in operations: + if "vacuum" in operations: try: self.db.model.vacuum_database() - results['vacuum'] = {'success': True, 'message': 'Database vacuum completed successfully'} + results["vacuum"] = { + "success": True, + "message": "Database vacuum completed successfully", + } except Exception as e: - results['vacuum'] = {'success': False, 'error': str(e), 'message': 'Failed to vacuum database'} + results["vacuum"] = { + "success": False, + "error": str(e), + "message": "Failed to vacuum database", + } - if 'cleanup_orphaned_images' in operations: + if "cleanup_orphaned_images" in operations: try: orphaned_removed = self.db.cleanup_missing_images() - results['cleanup_orphaned_images'] = { - 'success': True, 'removed_count': orphaned_removed, - 'message': f'Removed {orphaned_removed} orphaned image records' + results["cleanup_orphaned_images"] = { + "success": True, + "removed_count": orphaned_removed, + "message": f"Removed {orphaned_removed} orphaned image records", } except Exception as e: - results['cleanup_orphaned_images'] = {'success': False, 'error': str(e), 'message': 'Failed to cleanup orphaned images'} + results["cleanup_orphaned_images"] = { + "success": False, + "error": str(e), + "message": "Failed to cleanup orphaned images", + } - if 'check_hash_duplicates' in operations: + if "check_hash_duplicates" in operations: try: hash_duplicates = self.db.check_hash_duplicates() - results['check_hash_duplicates'] = { - 'success': True, 'duplicate_hashes': len(hash_duplicates), - 'message': f'Found {len(hash_duplicates)} duplicate hash groups' + results["check_hash_duplicates"] = { + "success": True, + "duplicate_hashes": len(hash_duplicates), + "message": f"Found {len(hash_duplicates)} duplicate hash groups", } except Exception as e: - results['check_hash_duplicates'] = {'success': False, 'error': str(e), 'message': 'Failed to check hash duplicates'} + results["check_hash_duplicates"] = { + "success": False, + "error": str(e), + "message": "Failed to check hash duplicates", + } - if 'statistics' in operations: + if "statistics" in operations: try: db_info = self.db.model.get_database_info() - results['statistics'] = {'success': True, 'info': db_info, 'message': 'Database statistics retrieved'} + results["statistics"] = { + "success": True, + "info": db_info, + "message": "Database statistics retrieved", + } except Exception as e: - results['statistics'] = {'success': False, 'error': str(e), 'message': 'Failed to get database statistics'} + results["statistics"] = { + "success": False, + "error": str(e), + "message": "Failed to get database statistics", + } - if 'prune_orphaned_prompts' in operations: + if "prune_orphaned_prompts" in operations: try: removed_count = self.db.prune_orphaned_prompts() - results['prune_orphaned_prompts'] = { - 'success': True, 'removed_count': removed_count, - 'message': f'Removed {removed_count} orphaned prompts (prompts with no linked images, excluding protected prompts)' + results["prune_orphaned_prompts"] = { + "success": True, + "removed_count": removed_count, + "message": f"Removed {removed_count} orphaned prompts (prompts with no linked images, excluding protected prompts)", } except Exception as e: - results['prune_orphaned_prompts'] = {'success': False, 'error': str(e), 'message': 'Failed to prune orphaned prompts'} + results["prune_orphaned_prompts"] = { + "success": False, + "error": str(e), + "message": "Failed to prune orphaned prompts", + } - if 'check_consistency' in operations: + if "check_consistency" in operations: try: consistency_issues = self.db.check_consistency() - results['check_consistency'] = { - 'success': True, 'issues_found': len(consistency_issues), - 'issues': consistency_issues[:10], - 'message': f'Found {len(consistency_issues)} consistency issues' + results["check_consistency"] = { + "success": True, + "issues_found": len(consistency_issues), + "issues": consistency_issues[:10], + "message": f"Found {len(consistency_issues)} consistency issues", } except Exception as e: - results['check_consistency'] = {'success': False, 'error': str(e), 'message': 'Failed to check database consistency'} + results["check_consistency"] = { + "success": False, + "error": str(e), + "message": "Failed to check database consistency", + } await self._run_in_executor(_run_maintenance) - all_successful = all(result.get('success', False) for result in results.values()) + all_successful = all( + result.get("success", False) for result in results.values() + ) - return web.json_response({ - 'success': True, - 'operations_completed': len(results), - 'all_successful': all_successful, - 'results': results, - 'message': f'Maintenance completed: {len(results)} operations processed' - }) + return web.json_response( + { + "success": True, + "operations_completed": len(results), + "all_successful": all_successful, + "results": results, + "message": f"Maintenance completed: {len(results)} operations processed", + } + ) except Exception as e: self.logger.error(f"Maintenance error: {e}", exc_info=True) - return web.json_response({ - 'success': False, - 'error': f'Maintenance failed: {str(e)}' - }, status=500) + return web.json_response( + {"success": False, "error": f"Maintenance failed: {str(e)}"}, status=500 + ) async def backup_database(self, request): """Backup the entire prompts.db database file.""" @@ -684,17 +800,16 @@ class AdminRoutesMixin: db_path = "prompts.db" if not os.path.exists(db_path): - return web.json_response({ - 'success': False, - 'error': 'Database file not found' - }, status=404) + return web.json_response( + {"success": False, "error": "Database file not found"}, status=404 + ) - with tempfile.NamedTemporaryFile(delete=False, suffix='.db') as temp_file: + with tempfile.NamedTemporaryFile(delete=False, suffix=".db") as temp_file: temp_path = temp_file.name shutil.copy2(db_path, temp_path) - with open(temp_path, 'rb') as f: + with open(temp_path, "rb") as f: file_data = f.read() os.unlink(temp_path) @@ -704,19 +819,19 @@ class AdminRoutesMixin: return web.Response( body=file_data, - content_type='application/octet-stream', + content_type="application/octet-stream", headers={ - 'Content-Disposition': f'attachment; filename="{filename}"', - 'Content-Length': str(len(file_data)) - } + "Content-Disposition": f'attachment; filename="{filename}"', + "Content-Length": str(len(file_data)), + }, ) except Exception as e: self.logger.error(f"Backup error: {e}", exc_info=True) - return web.json_response({ - 'success': False, - 'error': f'Failed to backup database: {str(e)}' - }, status=500) + return web.json_response( + {"success": False, "error": f"Failed to backup database: {str(e)}"}, + status=500, + ) async def restore_database(self, request): """Restore the prompts.db database from uploaded file.""" @@ -724,28 +839,33 @@ class AdminRoutesMixin: reader = await request.multipart() field = await reader.next() - if not field or field.name != 'database_file': - return web.json_response({ - 'success': False, - 'error': 'No database file uploaded. Expected field name: database_file' - }, status=400) + if not field or field.name != "database_file": + return web.json_response( + { + "success": False, + "error": "No database file uploaded. Expected field name: database_file", + }, + status=400, + ) MAX_RESTORE_SIZE = 100 * 1024 * 1024 # 100MB file_data = await field.read() if not file_data: - return web.json_response({ - 'success': False, - 'error': 'Uploaded file is empty' - }, status=400) + return web.json_response( + {"success": False, "error": "Uploaded file is empty"}, status=400 + ) if len(file_data) > MAX_RESTORE_SIZE: - return web.json_response({ - 'success': False, - 'error': f'File too large. Maximum size is {MAX_RESTORE_SIZE // (1024*1024)}MB' - }, status=400) + return web.json_response( + { + "success": False, + "error": f"File too large. Maximum size is {MAX_RESTORE_SIZE // (1024*1024)}MB", + }, + status=400, + ) - with tempfile.NamedTemporaryFile(delete=False, suffix='.db') as temp_file: + with tempfile.NamedTemporaryFile(delete=False, suffix=".db") as temp_file: temp_path = temp_file.name temp_file.write(file_data) @@ -753,20 +873,22 @@ class AdminRoutesMixin: with sqlite3.connect(temp_path) as conn: conn.row_factory = sqlite3.Row - cursor = conn.execute("SELECT name FROM sqlite_master WHERE type='table' AND name='prompts'") + cursor = conn.execute( + "SELECT name FROM sqlite_master WHERE type='table' AND name='prompts'" + ) if not cursor.fetchone(): raise ValueError("Database does not contain a 'prompts' table") cursor = conn.execute("PRAGMA table_info(prompts)") - columns = [row['name'] for row in cursor.fetchall()] - required_columns = ['id', 'text', 'created_at'] + columns = [row["name"] for row in cursor.fetchall()] + required_columns = ["id", "text", "created_at"] for col in required_columns: if col not in columns: raise ValueError(f"Database missing required column: {col}") cursor = conn.execute("SELECT COUNT(*) as count FROM prompts") - prompt_count = cursor.fetchone()['count'] + prompt_count = cursor.fetchone()["count"] db_path = "prompts.db" backup_path = f"{db_path}.backup_{datetime.datetime.now().strftime('%Y%m%d_%H%M%S')}" @@ -782,37 +904,46 @@ class AdminRoutesMixin: from ...database.operations import PromptDatabase except ImportError: import sys - sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.dirname(os.path.abspath(__file__))))) + + sys.path.insert( + 0, + os.path.dirname( + os.path.dirname(os.path.dirname(os.path.abspath(__file__))) + ), + ) from database.operations import PromptDatabase self.db = PromptDatabase() - return web.json_response({ - 'success': True, - 'message': f'Database restored successfully. Found {prompt_count} prompts.', - 'prompt_count': prompt_count, - 'backup_created': backup_path if os.path.exists(db_path) else None - }) + return web.json_response( + { + "success": True, + "message": f"Database restored successfully. Found {prompt_count} prompts.", + "prompt_count": prompt_count, + "backup_created": ( + backup_path if os.path.exists(db_path) else None + ), + } + ) except sqlite3.Error as e: - return web.json_response({ - 'success': False, - 'error': f'Invalid SQLite database: {str(e)}' - }, status=400) + return web.json_response( + {"success": False, "error": f"Invalid SQLite database: {str(e)}"}, + status=400, + ) except ValueError as e: - return web.json_response({ - 'success': False, - 'error': str(e) - }, status=400) + return web.json_response( + {"success": False, "error": str(e)}, status=400 + ) finally: if os.path.exists(temp_path): os.unlink(temp_path) except Exception as e: self.logger.error(f"Restore error: {e}", exc_info=True) - return web.json_response({ - 'success': False, - 'error': f'Failed to restore database: {str(e)}' - }, status=500) + return web.json_response( + {"success": False, "error": f"Failed to restore database: {str(e)}"}, + status=500, + ) async def scan_images(self, request): """Scan ComfyUI output images for prompt metadata and add them to the database.""" @@ -850,7 +981,10 @@ class AdminRoutesMixin: from ...utils.hashing import generate_prompt_hash except ImportError: import sys - current_dir = os.path.dirname(os.path.dirname(os.path.dirname(os.path.abspath(__file__)))) + + current_dir = os.path.dirname( + os.path.dirname(os.path.dirname(os.path.abspath(__file__))) + ) sys.path.insert(0, current_dir) from utils.hashing import generate_prompt_hash @@ -862,66 +996,106 @@ class AdminRoutesMixin: processed_count += 1 if metadata: - self.logger.debug(f"Found metadata in {os.path.basename(png_file)}: {list(metadata.keys())}") + self.logger.debug( + f"Found metadata in {os.path.basename(png_file)}: {list(metadata.keys())}" + ) parsed_data = self._parse_comfyui_prompt(metadata) - self.logger.debug(f"Parsed data keys: {list(parsed_data.keys())}, has prompt: {bool(parsed_data.get('prompt'))}, has parameters: {bool(parsed_data.get('parameters'))}") + self.logger.debug( + f"Parsed data keys: {list(parsed_data.keys())}, has prompt: {bool(parsed_data.get('prompt'))}, has parameters: {bool(parsed_data.get('parameters'))}" + ) - if parsed_data.get('prompt') or parsed_data.get('parameters'): + if parsed_data.get("prompt") or parsed_data.get( + "parameters" + ): found_count += 1 prompt_text = self._extract_readable_prompt(parsed_data) if prompt_text: - self.logger.debug(f"Found prompt in {os.path.basename(png_file)} (type: {type(prompt_text)}): {str(prompt_text)[:100]}...") + self.logger.debug( + f"Found prompt in {os.path.basename(png_file)} (type: {type(prompt_text)}): {str(prompt_text)[:100]}..." + ) else: - self.logger.debug(f"No readable prompt found in {os.path.basename(png_file)}, parsed_data keys: {list(parsed_data.keys())}") + self.logger.debug( + f"No readable prompt found in {os.path.basename(png_file)}, parsed_data keys: {list(parsed_data.keys())}" + ) if prompt_text and not isinstance(prompt_text, str): - self.logger.debug(f"Converting prompt_text from {type(prompt_text)} to string") + self.logger.debug( + f"Converting prompt_text from {type(prompt_text)} to string" + ) prompt_text = str(prompt_text) if prompt_text and prompt_text.strip(): try: - prompt_hash = generate_prompt_hash(prompt_text.strip()) - self.logger.debug(f"Generated hash for prompt: {prompt_hash[:16]}...") + prompt_hash = generate_prompt_hash( + prompt_text.strip() + ) + self.logger.debug( + f"Generated hash for prompt: {prompt_hash[:16]}..." + ) existing = await self._run_in_executor( self.db.get_prompt_by_hash, prompt_hash ) if existing: - self.logger.debug(f"Found existing prompt ID {existing['id']} for image {os.path.basename(png_file)}") + self.logger.debug( + f"Found existing prompt ID {existing['id']} for image {os.path.basename(png_file)}" + ) try: await self._run_in_executor( - self.db.link_image_to_prompt, existing['id'], str(png_file) + self.db.link_image_to_prompt, + existing["id"], + str(png_file), ) linked_count += 1 - self.logger.debug(f"Linked image {os.path.basename(png_file)} to existing prompt {existing['id']}") + self.logger.debug( + f"Linked image {os.path.basename(png_file)} to existing prompt {existing['id']}" + ) except Exception as e: - self.logger.error(f"Failed to link image {png_file} to existing prompt: {e}") + self.logger.error( + f"Failed to link image {png_file} to existing prompt: {e}" + ) else: - self.logger.debug(f"Saving new prompt from {os.path.basename(png_file)}") + self.logger.debug( + f"Saving new prompt from {os.path.basename(png_file)}" + ) prompt_id = await self._run_in_executor( self.db.save_prompt, - prompt_text.strip(), 'scanned', ['auto-scanned'], - f'Auto-scanned from {os.path.basename(png_file)}', - prompt_hash + prompt_text.strip(), + "scanned", + ["auto-scanned"], + f"Auto-scanned from {os.path.basename(png_file)}", + prompt_hash, ) if prompt_id: added_count += 1 - self.logger.info(f"Successfully saved new prompt with ID {prompt_id} from {os.path.basename(png_file)}") + self.logger.info( + f"Successfully saved new prompt with ID {prompt_id} from {os.path.basename(png_file)}" + ) try: await self._run_in_executor( - self.db.link_image_to_prompt, prompt_id, str(png_file) + self.db.link_image_to_prompt, + prompt_id, + str(png_file), + ) + self.logger.debug( + f"Linked image {os.path.basename(png_file)} to new prompt {prompt_id}" ) - self.logger.debug(f"Linked image {os.path.basename(png_file)} to new prompt {prompt_id}") except Exception as e: - self.logger.error(f"Failed to link image {png_file} to new prompt: {e}") + self.logger.error( + f"Failed to link image {png_file} to new prompt: {e}" + ) else: - self.logger.error(f"Failed to save prompt from {os.path.basename(png_file)} - no ID returned") + self.logger.error( + f"Failed to save prompt from {os.path.basename(png_file)} - no ID returned" + ) except Exception as e: - self.logger.error(f"Failed to save prompt from {png_file}: {e}") + self.logger.error( + f"Failed to save prompt from {png_file}: {e}" + ) # Update progress every 10 files if i % 10 == 0 or i == total_files - 1: @@ -936,7 +1110,9 @@ class AdminRoutesMixin: self.logger.error(f"Error processing {png_file}: {e}") continue - self.logger.info(f"Scan completed: processed={processed_count}, found={found_count}, new_prompts_added={added_count}, images_linked_to_existing={linked_count}") + self.logger.info( + f"Scan completed: processed={processed_count}, found={found_count}, new_prompts_added={added_count}, images_linked_to_existing={linked_count}" + ) yield f"data: {json.dumps({'type': 'complete', 'processed': processed_count, 'found': found_count, 'added': added_count, 'linked': linked_count})}\n\n" except Exception as e: @@ -946,18 +1122,18 @@ class AdminRoutesMixin: response = web.StreamResponse( status=200, - reason='OK', + reason="OK", headers={ - 'Content-Type': 'text/event-stream', - 'Cache-Control': 'no-cache', - 'Connection': 'keep-alive' - } + "Content-Type": "text/event-stream", + "Cache-Control": "no-cache", + "Connection": "keep-alive", + }, ) await response.prepare(request) async for chunk in stream_response(): - await response.write(chunk.encode('utf-8')) + await response.write(chunk.encode("utf-8")) await response.write_eof() return response diff --git a/py/api/autotag_routes.py b/py/api/autotag_routes.py index bfac0c9..2864570 100644 --- a/py/api/autotag_routes.py +++ b/py/api/autotag_routes.py @@ -47,24 +47,23 @@ class AutotagRoutesMixin: service = get_autotag_service() models_status = service.get_models_status() - return web.json_response({ - 'success': True, - 'models': models_status, - 'default_prompt': service.default_prompt, - 'model_loaded': service.is_model_loaded(), - 'loaded_model_type': service.get_loaded_model_type() - }) + return web.json_response( + { + "success": True, + "models": models_status, + "default_prompt": service.default_prompt, + "model_loaded": service.is_model_loaded(), + "loaded_model_type": service.get_loaded_model_type(), + } + ) except Exception as e: self.logger.error(f"Get autotag models error: {e}") - return web.json_response({ - 'success': False, - 'error': str(e) - }, status=500) + return web.json_response({"success": False, "error": str(e)}, status=500) async def download_autotag_model(self, request): """Download an AutoTag model with streaming progress.""" - model_type = request.match_info.get('model_type') + model_type = request.match_info.get("model_type") async def stream_response(): try: @@ -78,15 +77,14 @@ class AutotagRoutesMixin: yield f"data: {json.dumps({'type': 'progress', 'progress': 0, 'status': 'Starting download...'})}\n\n" - progress_data = {'last_progress': 0} + progress_data = {"last_progress": 0} def progress_callback(status: str, progress: float): - progress_data['last_progress'] = progress + progress_data["last_progress"] = progress loop = asyncio.get_event_loop() success = await loop.run_in_executor( - None, - lambda: service.download_model(model_type, progress_callback) + None, lambda: service.download_model(model_type, progress_callback) ) if success: @@ -100,28 +98,28 @@ class AutotagRoutesMixin: response = web.StreamResponse( status=200, - reason='OK', + reason="OK", headers={ - 'Content-Type': 'text/event-stream', - 'Cache-Control': 'no-cache', - 'Connection': 'keep-alive' - } + "Content-Type": "text/event-stream", + "Cache-Control": "no-cache", + "Connection": "keep-alive", + }, ) await response.prepare(request) async for chunk in stream_response(): - await response.write(chunk.encode('utf-8')) + await response.write(chunk.encode("utf-8")) await response.write_eof() return response async def start_autotag(self, request): """Start batch auto-tagging with streaming progress.""" - model_type = request.query.get('model_type', 'gguf') - custom_prompt = request.query.get('prompt', '') - skip_tagged = request.query.get('skip_tagged', 'true').lower() == 'true' - keep_in_memory = request.query.get('keep_in_memory', 'true').lower() == 'true' + model_type = request.query.get("model_type", "gguf") + custom_prompt = request.query.get("prompt", "") + skip_tagged = request.query.get("skip_tagged", "true").lower() == "true" + keep_in_memory = request.query.get("keep_in_memory", "true").lower() == "true" use_gpu = True async def stream_response(): @@ -132,7 +130,7 @@ class AutotagRoutesMixin: service = get_autotag_service() status = service.get_models_status() - if not status.get(model_type, {}).get('downloaded'): + if not status.get(model_type, {}).get("downloaded"): yield f"data: {json.dumps({'type': 'error', 'message': f'Model {model_type} not downloaded'})}\n\n" return @@ -141,8 +139,7 @@ class AutotagRoutesMixin: loop = asyncio.get_event_loop() try: await loop.run_in_executor( - None, - lambda: service.load_model(model_type, use_gpu) + None, lambda: service.load_model(model_type, use_gpu) ) except Exception as e: yield f"data: {json.dumps({'type': 'error', 'message': f'Failed to load model: {str(e)}'})}\n\n" @@ -172,8 +169,8 @@ class AutotagRoutesMixin: last_update_time = _time.monotonic() for i, image_data in enumerate(images): - image_path = image_data.get('image_path') - prompt_id = image_data.get('prompt_id') + image_path = image_data.get("image_path") + prompt_id = image_data.get("prompt_id") if not image_path or not prompt_id: skipped += 1 @@ -200,10 +197,12 @@ class AutotagRoutesMixin: continue if skip_tagged: - prompt_tags = image_data.get('prompt_tags', []) + prompt_tags = image_data.get("prompt_tags", []) if isinstance(prompt_tags, str): - prompt_tags = [t.strip() for t in prompt_tags.split(',') if t.strip()] - real_tags = [t for t in prompt_tags if t != 'auto-scanned'] + prompt_tags = [ + t.strip() for t in prompt_tags.split(",") if t.strip() + ] + real_tags = [t for t in prompt_tags if t != "auto-scanned"] if real_tags: tagged_prompt_ids.add(prompt_id) skipped += 1 @@ -217,18 +216,23 @@ class AutotagRoutesMixin: try: tags = await loop.run_in_executor( - None, - lambda p=str(image_path): service.generate_tags(p) + None, lambda p=str(image_path): service.generate_tags(p) ) processed += 1 if tags: - existing_prompt = await self._run_in_executor(self.db.get_prompt_by_id, prompt_id) + existing_prompt = await self._run_in_executor( + self.db.get_prompt_by_id, prompt_id + ) if existing_prompt: - existing_tags = existing_prompt.get('tags', []) + existing_tags = existing_prompt.get("tags", []) if isinstance(existing_tags, str): - existing_tags = [t.strip() for t in existing_tags.split(',') if t.strip()] + existing_tags = [ + t.strip() + for t in existing_tags.split(",") + if t.strip() + ] new_tags = [t for t in tags if t not in existing_tags] if new_tags: @@ -236,7 +240,7 @@ class AutotagRoutesMixin: await self._run_in_executor( self.db.update_prompt_metadata, prompt_id, - tags=all_tags + tags=all_tags, ) tagged += 1 else: @@ -260,32 +264,33 @@ class AutotagRoutesMixin: if not keep_in_memory: service.unload_model() - model_status = 'Model unloaded' + model_status = "Model unloaded" else: - model_status = 'Model kept in memory' + model_status = "Model kept in memory" yield f"data: {json.dumps({'type': 'complete', 'progress': 100, 'processed': processed, 'tagged': tagged, 'skipped': skipped, 'errors': errors, 'status': 'Complete', 'model_status': model_status, 'model_loaded': keep_in_memory})}\n\n" except Exception as e: self.logger.error(f"AutoTag error: {e}") import traceback + self.logger.error(f"AutoTag traceback: {traceback.format_exc()}") yield f"data: {json.dumps({'type': 'error', 'message': str(e)})}\n\n" response = web.StreamResponse( status=200, - reason='OK', + reason="OK", headers={ - 'Content-Type': 'text/event-stream', - 'Cache-Control': 'no-cache', - 'Connection': 'keep-alive' - } + "Content-Type": "text/event-stream", + "Cache-Control": "no-cache", + "Connection": "keep-alive", + }, ) await response.prepare(request) async for chunk in stream_response(): - await response.write(chunk.encode('utf-8')) + await response.write(chunk.encode("utf-8")) await response.write_eof() return response @@ -294,26 +299,27 @@ class AutotagRoutesMixin: """Generate tags for a single image.""" try: data = await request.json() - image_path = data.get('image_path') - model_type = data.get('model_type', 'gguf') - custom_prompt = data.get('prompt') - use_gpu = data.get('use_gpu', True) + image_path = data.get("image_path") + model_type = data.get("model_type", "gguf") + custom_prompt = data.get("prompt") + use_gpu = data.get("use_gpu", True) if not image_path: - return web.json_response({ - 'success': False, - 'error': 'image_path is required' - }, status=400) + return web.json_response( + {"success": False, "error": "image_path is required"}, status=400 + ) from ..autotag import get_autotag_service service = get_autotag_service() - if not service.is_model_loaded() or service.get_loaded_model_type() != model_type: + if ( + not service.is_model_loaded() + or service.get_loaded_model_type() != model_type + ): loop = asyncio.get_event_loop() await loop.run_in_executor( - None, - lambda: service.load_model(model_type, use_gpu) + None, lambda: service.load_model(model_type, use_gpu) ) if custom_prompt: @@ -321,81 +327,74 @@ class AutotagRoutesMixin: loop = asyncio.get_event_loop() tags = await loop.run_in_executor( - None, - lambda: service.generate_tags(image_path) + None, lambda: service.generate_tags(image_path) ) prompt_id = None try: - prompt_id = await self._run_in_executor(self.db.get_prompt_id_for_image, image_path) + prompt_id = await self._run_in_executor( + self.db.get_prompt_id_for_image, image_path + ) except Exception as e: self.logger.warning(f"Could not find linked prompt: {e}") - return web.json_response({ - 'success': True, - 'tags': tags, - 'prompt_id': prompt_id, - 'image_path': image_path - }) + return web.json_response( + { + "success": True, + "tags": tags, + "prompt_id": prompt_id, + "image_path": image_path, + } + ) except Exception as e: self.logger.error(f"AutoTag single error: {e}") - return web.json_response({ - 'success': False, - 'error': str(e) - }, status=500) + return web.json_response({"success": False, "error": str(e)}, status=500) async def apply_autotag(self, request): """Apply selected tags to a prompt.""" try: data = await request.json() - prompt_id = data.get('prompt_id') - tags = data.get('tags', []) + prompt_id = data.get("prompt_id") + tags = data.get("tags", []) if not prompt_id: - return web.json_response({ - 'success': False, - 'error': 'prompt_id is required' - }, status=400) + return web.json_response( + {"success": False, "error": "prompt_id is required"}, status=400 + ) if not tags: - return web.json_response({ - 'success': True, - 'message': 'No tags to apply' - }) + return web.json_response( + {"success": True, "message": "No tags to apply"} + ) prompt = await self._run_in_executor(self.db.get_prompt_by_id, prompt_id) if not prompt: - return web.json_response({ - 'success': False, - 'error': f'Prompt {prompt_id} not found' - }, status=404) + return web.json_response( + {"success": False, "error": f"Prompt {prompt_id} not found"}, + status=404, + ) - existing_tags = prompt.get('tags', []) + existing_tags = prompt.get("tags", []) if isinstance(existing_tags, str): - existing_tags = [t.strip() for t in existing_tags.split(',') if t.strip()] + existing_tags = [ + t.strip() for t in existing_tags.split(",") if t.strip() + ] new_tags = [t for t in tags if t not in existing_tags] all_tags = existing_tags + new_tags await self._run_in_executor( - self.db.update_prompt_metadata, - prompt_id, - tags=all_tags + self.db.update_prompt_metadata, prompt_id, tags=all_tags ) - return web.json_response({ - 'success': True, - 'added_tags': new_tags, - 'total_tags': len(all_tags) - }) + return web.json_response( + {"success": True, "added_tags": new_tags, "total_tags": len(all_tags)} + ) except Exception as e: self.logger.error(f"Apply autotag error: {e}") - return web.json_response({ - 'success': False, - 'error': str(e) - }, status=500) + return web.json_response({"success": False, "error": str(e)}, status=500) async def unload_autotag_model(self, request): """Manually unload the AutoTag model from memory.""" @@ -405,39 +404,45 @@ class AutotagRoutesMixin: service = get_autotag_service() if not service.is_model_loaded(): - return web.json_response({ - 'success': True, - 'message': 'No model was loaded' - }) + return web.json_response( + {"success": True, "message": "No model was loaded"} + ) model_type = service.get_loaded_model_type() service.unload_model() - return web.json_response({ - 'success': True, - 'message': f'{model_type.upper()} model unloaded successfully', - 'model_loaded': False - }) + return web.json_response( + { + "success": True, + "message": f"{model_type.upper()} model unloaded successfully", + "model_loaded": False, + } + ) except Exception as e: self.logger.error(f"Unload autotag model error: {e}") - return web.json_response({ - 'success': False, - 'error': str(e) - }, status=500) + return web.json_response({"success": False, "error": str(e)}, status=500) async def scan_output_dir(self, request): """Scan ComfyUI output directory for images.""" try: output_dir = self._find_comfyui_output_dir() if not output_dir: - return web.json_response({ - 'success': False, - 'error': 'ComfyUI output directory not found' - }, status=404) + return web.json_response( + {"success": False, "error": "ComfyUI output directory not found"}, + status=404, + ) output_path = Path(output_dir) - image_extensions = ['.png', '.jpg', '.jpeg', '.gif', '.webp', '.bmp', '.tiff'] + image_extensions = [ + ".png", + ".jpg", + ".jpeg", + ".gif", + ".webp", + ".bmp", + ".tiff", + ] images = [] seen_paths = set() @@ -445,7 +450,7 @@ class AutotagRoutesMixin: for ext in image_extensions: for pattern in [f"*{ext.lower()}", f"*{ext.upper()}"]: for image_path in output_path.rglob(pattern): - if 'thumbnails' not in image_path.parts: + if "thumbnails" not in image_path.parts: normalized_path = str(image_path).lower() if normalized_path not in seen_paths: seen_paths.add(normalized_path) @@ -455,35 +460,37 @@ class AutotagRoutesMixin: thumbnail_url = None thumbnails_dir = output_path / "thumbnails" if thumbnails_dir.exists(): - rel_path_no_ext = rel_path.with_suffix('') + rel_path_no_ext = rel_path.with_suffix("") thumbnail_rel_path = f"thumbnails/{rel_path_no_ext.as_posix()}_thumb{image_path.suffix}" - thumbnail_abs_path = thumbnails_dir / f"{rel_path_no_ext.as_posix()}_thumb{image_path.suffix}" + thumbnail_abs_path = ( + thumbnails_dir + / f"{rel_path_no_ext.as_posix()}_thumb{image_path.suffix}" + ) if thumbnail_abs_path.exists(): from urllib.parse import quote + thumbnail_url = f'/prompt_manager/images/serve/{quote(thumbnail_rel_path, safe="/")}' from urllib.parse import quote as url_quote - images.append({ - 'filename': image_path.name, - 'path': str(image_path), - 'relative_path': str(rel_path), - 'url': f'/prompt_manager/images/serve/{url_quote(rel_path.as_posix(), safe="/")}', - 'thumbnail_url': thumbnail_url - }) - images.sort(key=lambda x: x['filename']) + images.append( + { + "filename": image_path.name, + "path": str(image_path), + "relative_path": str(rel_path), + "url": f'/prompt_manager/images/serve/{url_quote(rel_path.as_posix(), safe="/")}', + "thumbnail_url": thumbnail_url, + } + ) + + images.sort(key=lambda x: x["filename"]) self.logger.info(f"Found {len(images)} images in output directory") - return web.json_response({ - 'success': True, - 'images': images, - 'count': len(images) - }) + return web.json_response( + {"success": True, "images": images, "count": len(images)} + ) except Exception as e: self.logger.error(f"Scan output dir error: {e}") - return web.json_response({ - 'success': False, - 'error': str(e) - }, status=500) + return web.json_response({"success": False, "error": str(e)}, status=500) diff --git a/py/api/images.py b/py/api/images.py index c34a223..599acaf 100644 --- a/py/api/images.py +++ b/py/api/images.py @@ -79,19 +79,16 @@ class ImageRoutesMixin: # Additional fallback: convert to JSON string and clean NaN values manually try: - response_data = { - 'success': True, - 'images': cleaned_images - } + response_data = {"success": True, "images": cleaned_images} # Convert to JSON string json_str = json.dumps(response_data, default=str) # Clean any remaining NaN values with regex - json_str = re.sub(r':\s*NaN', ': null', json_str) - json_str = re.sub(r'\[\s*NaN\s*\]', '[null]', json_str) - json_str = re.sub(r',\s*NaN\s*,', ', null,', json_str) - json_str = re.sub(r',\s*NaN\s*\]', ', null]', json_str) - json_str = re.sub(r'\[\s*NaN\s*,', '[null,', json_str) + json_str = re.sub(r":\s*NaN", ": null", json_str) + json_str = re.sub(r"\[\s*NaN\s*\]", "[null]", json_str) + json_str = re.sub(r",\s*NaN\s*,", ", null,", json_str) + json_str = re.sub(r",\s*NaN\s*\]", ", null]", json_str) + json_str = re.sub(r"\[\s*NaN\s*,", "[null,", json_str) # Parse back to verify it's valid JSON cleaned_data = json.loads(json_str) @@ -100,80 +97,57 @@ class ImageRoutesMixin: except Exception as json_error: self.logger.error(f"JSON cleaning error: {json_error}") # Fallback to original response - return web.json_response({ - 'success': True, - 'images': cleaned_images - }) + return web.json_response({"success": True, "images": cleaned_images}) except Exception as e: self.logger.error(f"Get prompt images error: {e}") - return web.json_response({ - 'success': False, - 'error': str(e) - }, status=500) + return web.json_response({"success": False, "error": str(e)}, status=500) async def get_recent_images(self, request): """Get recently generated images.""" try: - limit = int(request.query.get('limit', 50)) + limit = int(request.query.get("limit", 50)) images = await self._run_in_executor(self.db.get_recent_images, limit) - return web.json_response({ - 'success': True, - 'images': images - }) + return web.json_response({"success": True, "images": images}) except Exception as e: self.logger.error(f"Get recent images error: {e}") - return web.json_response({ - 'success': False, - 'error': str(e) - }, status=500) + return web.json_response({"success": False, "error": str(e)}, status=500) async def get_all_images(self, request): """Get all generated images with linked prompts.""" try: images = await self._run_in_executor(self.db.get_all_images) - return web.json_response({ - 'success': True, - 'images': images, - 'count': len(images) - }) + return web.json_response( + {"success": True, "images": images, "count": len(images)} + ) except Exception as e: self.logger.error(f"Get all images error: {e}") - return web.json_response({ - 'success': False, - 'error': str(e) - }, status=500) + return web.json_response({"success": False, "error": str(e)}, status=500) async def search_images(self, request): """Search images by prompt text.""" try: - query = request.query.get('q', '') + query = request.query.get("q", "") if not query: - return web.json_response({ - 'success': False, - 'error': 'Search query required' - }, status=400) + return web.json_response( + {"success": False, "error": "Search query required"}, status=400 + ) images = await self._run_in_executor(self.db.search_images_by_prompt, query) - return web.json_response({ - 'success': True, - 'images': images, - 'query': query - }) + return web.json_response( + {"success": True, "images": images, "query": query} + ) except Exception as e: self.logger.error(f"Search images error: {e}") - return web.json_response({ - 'success': False, - 'error': str(e) - }, status=500) + return web.json_response({"success": False, "error": str(e)}, status=500) def _scan_gallery_files_sync(self, output_path): """Scan output directory for media files (blocking I/O, run in executor).""" - image_extensions = ['.png', '.jpg', '.jpeg', '.webp', '.gif'] - video_extensions = ['.mp4', '.webm', '.avi', '.mov', '.mkv', '.m4v', '.wmv'] + image_extensions = [".png", ".jpg", ".jpeg", ".webp", ".gif"] + video_extensions = [".mp4", ".webm", ".avi", ".mov", ".mkv", ".m4v", ".wmv"] media_extensions = image_extensions + video_extensions all_images = [] @@ -181,7 +155,7 @@ class ImageRoutesMixin: for ext in media_extensions: for pattern in [f"*{ext}", f"*{ext.upper()}"]: for media_path in output_path.rglob(pattern): - if 'thumbnails' not in media_path.parts: + if "thumbnails" not in media_path.parts: normalized_path = str(media_path).lower() if normalized_path not in seen_paths: seen_paths.add(normalized_path) @@ -194,8 +168,10 @@ class ImageRoutesMixin: async def _get_gallery_files(self, output_path): """Get gallery files with TTL cache. Invalidated by image monitor.""" now = _time.monotonic() - if (self._gallery_cache is not None - and (now - self._gallery_cache_time) < self._gallery_cache_ttl): + if ( + self._gallery_cache is not None + and (now - self._gallery_cache_time) < self._gallery_cache_ttl + ): return self._gallery_cache all_images = await self._run_in_executor( @@ -213,25 +189,27 @@ class ImageRoutesMixin: # Find ComfyUI output directory output_dir = self._find_comfyui_output_dir() if not output_dir: - return web.json_response({ - 'success': False, - 'error': 'ComfyUI output directory not found', - 'images': [] - }) + return web.json_response( + { + "success": False, + "error": "ComfyUI output directory not found", + "images": [], + } + ) # Get pagination parameters - limit = int(request.query.get('limit', 100)) - offset = int(request.query.get('offset', 0)) + limit = int(request.query.get("limit", 100)) + offset = int(request.query.get("offset", 0)) output_path = Path(output_dir) - thumbnails_dir = output_path / 'thumbnails' - video_extensions = ['.mp4', '.webm', '.avi', '.mov', '.mkv', '.m4v', '.wmv'] + thumbnails_dir = output_path / "thumbnails" + video_extensions = [".mp4", ".webm", ".avi", ".mov", ".mkv", ".m4v", ".wmv"] # Use cached file listing (Fix 2.4) all_images = await self._get_gallery_files(output_path) # Apply pagination - paginated_images = all_images[offset:offset + limit] + paginated_images = all_images[offset : offset + limit] # Format media data in executor (stat calls are blocking) def _format_page(): @@ -242,30 +220,32 @@ class ImageRoutesMixin: rel_path = media_path.relative_to(output_path) extension = media_path.suffix.lower() is_video = extension in video_extensions - media_type = 'video' if is_video else 'image' + media_type = "video" if is_video else "image" thumbnail_url = None if thumbnails_dir.exists(): - thumbnail_ext = '.jpg' if is_video else extension - rel_path_no_ext = rel_path.with_suffix('') + thumbnail_ext = ".jpg" if is_video else extension + rel_path_no_ext = rel_path.with_suffix("") thumbnail_rel_path = f"thumbnails/{rel_path_no_ext.as_posix()}_thumb{thumbnail_ext}" thumbnail_abs_path = output_path / thumbnail_rel_path if thumbnail_abs_path.exists(): thumbnail_url = f'/prompt_manager/images/serve/{quote(thumbnail_rel_path, safe="/")}' - images.append({ - 'id': str(hash(str(media_path))), - 'filename': media_path.name, - 'path': str(media_path), - 'relative_path': str(rel_path), - 'url': f'/prompt_manager/images/serve/{rel_path.as_posix()}', - 'thumbnail_url': thumbnail_url, - 'size': stat.st_size, - 'modified_time': stat.st_mtime, - 'extension': extension, - 'media_type': media_type, - 'is_video': is_video - }) + images.append( + { + "id": str(hash(str(media_path))), + "filename": media_path.name, + "path": str(media_path), + "relative_path": str(rel_path), + "url": f"/prompt_manager/images/serve/{rel_path.as_posix()}", + "thumbnail_url": thumbnail_url, + "size": stat.st_size, + "modified_time": stat.st_mtime, + "extension": extension, + "media_type": media_type, + "is_video": is_video, + } + ) except Exception as e: self.logger.error(f"Error processing media {media_path}: {e}") continue @@ -273,22 +253,22 @@ class ImageRoutesMixin: images = await self._run_in_executor(_format_page) - return web.json_response({ - 'success': True, - 'images': images, - 'total': len(all_images), - 'offset': offset, - 'limit': limit, - 'has_more': offset + limit < len(all_images) - }) + return web.json_response( + { + "success": True, + "images": images, + "total": len(all_images), + "offset": offset, + "limit": limit, + "has_more": offset + limit < len(all_images), + } + ) except Exception as e: self.logger.error(f"Get output images error: {e}") - return web.json_response({ - 'success': False, - 'error': str(e), - 'images': [] - }, status=500) + return web.json_response( + {"success": False, "error": str(e), "images": []}, status=500 + ) async def serve_image(self, request): """Serve the actual image file using streamed FileResponse.""" @@ -297,29 +277,37 @@ class ImageRoutesMixin: image = await self._run_in_executor(self.db.get_image_by_id, image_id) if not image: - return web.json_response({'success': False, 'error': 'Image not found'}, status=404) + return web.json_response( + {"success": False, "error": "Image not found"}, status=404 + ) - image_path = Path(image['image_path']).resolve() + image_path = Path(image["image_path"]).resolve() # Validate path is within the ComfyUI output directory output_dir = self._find_comfyui_output_dir() if output_dir: output_path = Path(output_dir).resolve() if not image_path.is_relative_to(output_path): - return web.json_response({'success': False, 'error': 'Access denied'}, status=403) + return web.json_response( + {"success": False, "error": "Access denied"}, status=403 + ) if not image_path.exists(): - return web.json_response({'success': False, 'error': 'Image file not found'}, status=404) + return web.json_response( + {"success": False, "error": "Image file not found"}, status=404 + ) response = web.FileResponse(image_path) - response.headers['Cache-Control'] = 'public, max-age=3600' + response.headers["Cache-Control"] = "public, max-age=3600" return response except ValueError: - return web.json_response({'success': False, 'error': 'Invalid image ID'}, status=400) + return web.json_response( + {"success": False, "error": "Invalid image ID"}, status=400 + ) except Exception as e: self.logger.error(f"Serve image error: {e}") - return web.json_response({'success': False, 'error': str(e)}, status=500) + return web.json_response({"success": False, "error": str(e)}, status=500) async def serve_output_image(self, request): """Serve image file directly from ComfyUI output folder using streamed FileResponse.""" @@ -329,7 +317,10 @@ class ImageRoutesMixin: # Find ComfyUI output directory output_dir = self._find_comfyui_output_dir() if not output_dir: - return web.json_response({'success': False, 'error': 'ComfyUI output directory not found'}, status=404) + return web.json_response( + {"success": False, "error": "ComfyUI output directory not found"}, + status=404, + ) # Construct full image path image_path = Path(output_dir) / filepath @@ -339,51 +330,55 @@ class ImageRoutesMixin: image_path = image_path.resolve() output_path = Path(output_dir).resolve() if not image_path.is_relative_to(output_path): - return web.json_response({'success': False, 'error': 'Access denied'}, status=403) + return web.json_response( + {"success": False, "error": "Access denied"}, status=403 + ) except Exception: - return web.json_response({'success': False, 'error': 'Invalid file path'}, status=400) + return web.json_response( + {"success": False, "error": "Invalid file path"}, status=400 + ) if not image_path.exists(): - return web.json_response({'success': False, 'error': 'Image file not found'}, status=404) + return web.json_response( + {"success": False, "error": "Image file not found"}, status=404 + ) response = web.FileResponse(image_path) - response.headers['Cache-Control'] = 'public, max-age=3600' + response.headers["Cache-Control"] = "public, max-age=3600" return response except Exception as e: self.logger.error(f"Serve output image error: {e}") - return web.json_response({'success': False, 'error': str(e)}, status=500) + return web.json_response({"success": False, "error": str(e)}, status=500) async def generate_thumbnails(self, request): """Generate thumbnails for all images and videos in the ComfyUI output directory.""" try: # Get request parameters data = await request.json() - quality = data.get('quality', 'medium') + quality = data.get("quality", "medium") # Map quality to size - size_map = { - 'low': (150, 150), - 'medium': (300, 300), - 'high': (600, 600) - } + size_map = {"low": (150, 150), "medium": (300, 300), "high": (600, 600)} thumbnail_size = size_map.get(quality, (300, 300)) # Find ComfyUI output directory output_dir = self._find_comfyui_output_dir() if not output_dir: - return web.json_response({ - 'success': False, - 'error': 'ComfyUI output directory not found' - }, status=404) + return web.json_response( + {"success": False, "error": "ComfyUI output directory not found"}, + status=404, + ) output_path = Path(output_dir) - thumbnails_dir = output_path / 'thumbnails' + thumbnails_dir = output_path / "thumbnails" # Run the entire thumbnail generation in executor (heavy PIL I/O) result = await self._run_in_executor( self._generate_thumbnails_sync, - output_path, thumbnails_dir, thumbnail_size + output_path, + thumbnails_dir, + thumbnail_size, ) # Invalidate gallery cache since thumbnails changed @@ -392,16 +387,16 @@ class ImageRoutesMixin: return web.json_response(result) except ImportError: - return web.json_response({ - 'success': False, - 'error': 'PIL (Pillow) library not available. Install with: pip install Pillow' - }, status=500) + return web.json_response( + { + "success": False, + "error": "PIL (Pillow) library not available. Install with: pip install Pillow", + }, + status=500, + ) except Exception as e: self.logger.error(f"Generate thumbnails error: {e}") - return web.json_response({ - 'success': False, - 'error': str(e) - }, status=500) + return web.json_response({"success": False, "error": str(e)}, status=500) def _generate_thumbnails_sync(self, output_path, thumbnails_dir, thumbnail_size): """Blocking thumbnail generation loop (run in executor).""" @@ -410,12 +405,12 @@ class ImageRoutesMixin: thumbnails_dir.mkdir(exist_ok=True) self.logger.info("Scanning for media files to generate thumbnails...") - image_extensions = ['.png', '.jpg', '.jpeg', '.webp', '.gif'] - video_extensions = ['.mp4', '.webm', '.avi', '.mov', '.mkv', '.m4v', '.wmv'] + image_extensions = [".png", ".jpg", ".jpeg", ".webp", ".gif"] + video_extensions = [".mp4", ".webm", ".avi", ".mov", ".mkv", ".m4v", ".wmv"] media_extensions = image_extensions + video_extensions media_files = [] for root, dirs, files in os.walk(output_path): - if 'thumbnails' in Path(root).parts: + if "thumbnails" in Path(root).parts: continue for file in files: if any(file.lower().endswith(ext) for ext in media_extensions): @@ -426,11 +421,11 @@ class ImageRoutesMixin: if total_images == 0: return { - 'success': True, - 'count': 0, - 'total_images': 0, - 'message': 'No media files found to process', - 'errors': [] + "success": True, + "count": 0, + "total_images": 0, + "message": "No media files found to process", + "errors": [], } generated_count = 0 @@ -441,64 +436,88 @@ class ImageRoutesMixin: for i, media_file in enumerate(media_files): try: rel_path = media_file.relative_to(output_path) - is_video = any(media_file.name.lower().endswith(ext) for ext in video_extensions) + is_video = any( + media_file.name.lower().endswith(ext) for ext in video_extensions + ) - rel_path_no_ext = rel_path.with_suffix('') + rel_path_no_ext = rel_path.with_suffix("") if is_video: - thumbnail_path = thumbnails_dir / f"{rel_path_no_ext.as_posix()}_thumb.jpg" + thumbnail_path = ( + thumbnails_dir / f"{rel_path_no_ext.as_posix()}_thumb.jpg" + ) else: - thumbnail_path = thumbnails_dir / f"{rel_path_no_ext.as_posix()}_thumb{rel_path.suffix}" + thumbnail_path = ( + thumbnails_dir + / f"{rel_path_no_ext.as_posix()}_thumb{rel_path.suffix}" + ) - if (thumbnail_path.exists() and - thumbnail_path.stat().st_mtime > media_file.stat().st_mtime): + if ( + thumbnail_path.exists() + and thumbnail_path.stat().st_mtime > media_file.stat().st_mtime + ): skipped_count += 1 continue thumbnail_path.parent.mkdir(parents=True, exist_ok=True) if is_video: - if self._generate_video_thumbnail(media_file, thumbnail_path, thumbnail_size): + if self._generate_video_thumbnail( + media_file, thumbnail_path, thumbnail_size + ): generated_count += 1 else: - errors.append(f"Failed to generate video thumbnail for {media_file.name}") + errors.append( + f"Failed to generate video thumbnail for {media_file.name}" + ) else: with Image.open(media_file) as img: - if img.mode in ('RGBA', 'LA', 'P'): - img = img.convert('RGB') + if img.mode in ("RGBA", "LA", "P"): + img = img.convert("RGB") img.thumbnail(thumbnail_size, Image.Resampling.LANCZOS) - save_kwargs = {'quality': 85, 'optimize': True} - if thumbnail_path.suffix.lower() == '.png': - save_kwargs = {'optimize': True} + save_kwargs = {"quality": 85, "optimize": True} + if thumbnail_path.suffix.lower() == ".png": + save_kwargs = {"optimize": True} img.save(thumbnail_path, **save_kwargs) generated_count += 1 - if generated_count % 100 == 0 or (i + 1) % max(1, total_images // 10) == 0: + if ( + generated_count % 100 == 0 + or (i + 1) % max(1, total_images // 10) == 0 + ): progress = ((i + 1) / total_images) * 100 elapsed = time.time() - start_time rate = (i + 1) / elapsed if elapsed > 0 else 0 eta = ((total_images - i - 1) / rate) if rate > 0 else 0 - self.logger.info(f"Thumbnail progress: {i+1}/{total_images} ({progress:.1f}%) - " - f"Generated: {generated_count}, Skipped: {skipped_count}, " - f"Rate: {rate:.1f} img/s, ETA: {eta:.0f}s") + self.logger.info( + f"Thumbnail progress: {i+1}/{total_images} ({progress:.1f}%) - " + f"Generated: {generated_count}, Skipped: {skipped_count}, " + f"Rate: {rate:.1f} img/s, ETA: {eta:.0f}s" + ) except Exception as e: - error_msg = f"Failed to generate thumbnail for {media_file.name}: {str(e)}" + error_msg = ( + f"Failed to generate thumbnail for {media_file.name}: {str(e)}" + ) errors.append(error_msg) self.logger.warning(error_msg) elapsed_time = time.time() - start_time - self.logger.info(f"Thumbnail generation completed: {generated_count} generated, " - f"{skipped_count} skipped, {len(errors)} errors in {elapsed_time:.1f}s") + self.logger.info( + f"Thumbnail generation completed: {generated_count} generated, " + f"{skipped_count} skipped, {len(errors)} errors in {elapsed_time:.1f}s" + ) return { - 'success': True, - 'count': generated_count, - 'skipped': skipped_count, - 'total_images': total_images, - 'errors': errors, - 'thumbnails_path': str(thumbnails_dir), - 'elapsed_time': round(elapsed_time, 2), - 'processing_rate': round((total_images / elapsed_time) if elapsed_time > 0 else 0, 2) + "success": True, + "count": generated_count, + "skipped": skipped_count, + "total_images": total_images, + "errors": errors, + "thumbnails_path": str(thumbnails_dir), + "elapsed_time": round(elapsed_time, 2), + "processing_rate": round( + (total_images / elapsed_time) if elapsed_time > 0 else 0, 2 + ), } async def generate_thumbnails_with_progress(self, request): @@ -507,25 +526,21 @@ class ImageRoutesMixin: import time # Parse query parameters - quality = request.query.get('quality', 'medium') + quality = request.query.get("quality", "medium") # Map quality to size - size_map = { - 'low': (150, 150), - 'medium': (300, 300), - 'high': (600, 600) - } + size_map = {"low": (150, 150), "medium": (300, 300), "high": (600, 600)} thumbnail_size = size_map.get(quality, (300, 300)) # Set up SSE response response = web.StreamResponse( status=200, headers={ - 'Content-Type': 'text/event-stream', - 'Cache-Control': 'no-cache', - 'Connection': 'keep-alive', - 'Access-Control-Allow-Origin': '*', - } + "Content-Type": "text/event-stream", + "Cache-Control": "no-cache", + "Connection": "keep-alive", + "Access-Control-Allow-Origin": "*", + }, ) await response.prepare(request) @@ -533,7 +548,7 @@ class ImageRoutesMixin: """Send SSE event to client.""" try: message = f"event: {event_type}\ndata: {json.dumps(data)}\n\n" - await response.write(message.encode('utf-8')) + await response.write(message.encode("utf-8")) await asyncio.sleep(0.01) except Exception as e: self.logger.warning(f"Failed to send SSE message: {e}") @@ -542,71 +557,101 @@ class ImageRoutesMixin: # Find ComfyUI output directory output_dir = self._find_comfyui_output_dir() if not output_dir: - await send_progress('error', { - 'error': 'ComfyUI output directory not found' - }) + await send_progress( + "error", {"error": "ComfyUI output directory not found"} + ) return response output_path = Path(output_dir) - thumbnails_dir = output_path / 'thumbnails' + thumbnails_dir = output_path / "thumbnails" thumbnails_dir.mkdir(exist_ok=True) # Send scanning event - await send_progress('status', { - 'phase': 'scanning', - 'message': f'Scanning {output_path} for images and videos to process...' - }) + await send_progress( + "status", + { + "phase": "scanning", + "message": f"Scanning {output_path} for images and videos to process...", + }, + ) - self.logger.info(f"Starting thumbnail generation scan in: {output_path}") + self.logger.info( + f"Starting thumbnail generation scan in: {output_path}" + ) # Find all media files (images and videos) - image_extensions = ['.png', '.jpg', '.jpeg', '.webp', '.gif'] - video_extensions = ['.mp4', '.webm', '.avi', '.mov', '.mkv', '.m4v', '.wmv'] + image_extensions = [".png", ".jpg", ".jpeg", ".webp", ".gif"] + video_extensions = [ + ".mp4", + ".webm", + ".avi", + ".mov", + ".mkv", + ".m4v", + ".wmv", + ] media_extensions = image_extensions + video_extensions media_files = [] scanned_dirs = 0 for root, dirs, files in os.walk(output_path): - if 'thumbnails' in Path(root).parts: + if "thumbnails" in Path(root).parts: continue scanned_dirs += 1 if scanned_dirs % 5 == 0: - await send_progress('status', { - 'phase': 'scanning', - 'message': f'Scanning directories... ({scanned_dirs} checked, {len(media_files)} files found)' - }) + await send_progress( + "status", + { + "phase": "scanning", + "message": f"Scanning directories... ({scanned_dirs} checked, {len(media_files)} files found)", + }, + ) for file in files: if any(file.lower().endswith(ext) for ext in media_extensions): media_files.append(Path(root) / file) - self.logger.info(f"Scan complete: Found {len(media_files)} media files in {scanned_dirs} directories") + self.logger.info( + f"Scan complete: Found {len(media_files)} media files in {scanned_dirs} directories" + ) total_images = len(media_files) # Count images vs videos for more detail - image_count = sum(1 for f in media_files if any(f.name.lower().endswith(ext) for ext in image_extensions)) + image_count = sum( + 1 + for f in media_files + if any(f.name.lower().endswith(ext) for ext in image_extensions) + ) video_count = total_images - image_count - await send_progress('start', { - 'total_images': total_images, - 'phase': 'processing', - 'message': f'Found {image_count} images and {video_count} videos to process', - 'image_count': image_count, - 'video_count': video_count - }) + await send_progress( + "start", + { + "total_images": total_images, + "phase": "processing", + "message": f"Found {image_count} images and {video_count} videos to process", + "image_count": image_count, + "video_count": video_count, + }, + ) - self.logger.info(f"Starting thumbnail generation for {image_count} images and {video_count} videos") + self.logger.info( + f"Starting thumbnail generation for {image_count} images and {video_count} videos" + ) if total_images == 0: - await send_progress('complete', { - 'count': 0, - 'skipped': 0, - 'total_images': 0, - 'elapsed_time': 0, - 'message': 'No media files found to process' - }) + await send_progress( + "complete", + { + "count": 0, + "skipped": 0, + "total_images": 0, + "elapsed_time": 0, + "message": "No media files found to process", + }, + ) return response generated_count = 0 @@ -619,80 +664,117 @@ class ImageRoutesMixin: if not media_file.exists() or not media_file.is_file(): continue - is_video = any(media_file.name.lower().endswith(ext) for ext in video_extensions) + is_video = any( + media_file.name.lower().endswith(ext) + for ext in video_extensions + ) rel_path = media_file.relative_to(output_path) - rel_path_no_ext = rel_path.with_suffix('') + rel_path_no_ext = rel_path.with_suffix("") if is_video: - thumbnail_path = thumbnails_dir / f"{rel_path_no_ext.as_posix()}_thumb.jpg" + thumbnail_path = ( + thumbnails_dir + / f"{rel_path_no_ext.as_posix()}_thumb.jpg" + ) else: - thumbnail_path = thumbnails_dir / f"{rel_path_no_ext.as_posix()}_thumb{rel_path.suffix}" + thumbnail_path = ( + thumbnails_dir + / f"{rel_path_no_ext.as_posix()}_thumb{rel_path.suffix}" + ) # Ensure thumbnail path is within our thumbnails directory try: thumbnail_path = thumbnail_path.resolve() thumbnails_dir_resolved = thumbnails_dir.resolve() - if not str(thumbnail_path).startswith(str(thumbnails_dir_resolved)): - self.logger.warning(f"Skipping thumbnail outside safe directory: {thumbnail_path}") + if not str(thumbnail_path).startswith( + str(thumbnails_dir_resolved) + ): + self.logger.warning( + f"Skipping thumbnail outside safe directory: {thumbnail_path}" + ) continue except Exception as e: - self.logger.warning(f"Path validation failed for {rel_path}: {e}") + self.logger.warning( + f"Path validation failed for {rel_path}: {e}" + ) continue # Skip if thumbnail already exists and is newer than original - if (thumbnail_path.exists() and - thumbnail_path.stat().st_mtime > media_file.stat().st_mtime): + if ( + thumbnail_path.exists() + and thumbnail_path.stat().st_mtime + > media_file.stat().st_mtime + ): skipped_count += 1 else: thumbnail_path.parent.mkdir(parents=True, exist_ok=True) if is_video: - if self._generate_video_thumbnail(media_file, thumbnail_path, thumbnail_size): + if self._generate_video_thumbnail( + media_file, thumbnail_path, thumbnail_size + ): generated_count += 1 else: - errors.append(f"Failed to generate video thumbnail for {media_file.name}") + errors.append( + f"Failed to generate video thumbnail for {media_file.name}" + ) else: with Image.open(media_file) as img: - if img.mode in ('RGBA', 'LA', 'P'): - img = img.convert('RGB') - img.thumbnail(thumbnail_size, Image.Resampling.LANCZOS) - save_kwargs = {'quality': 85, 'optimize': True} - if thumbnail_path.suffix.lower() == '.png': - save_kwargs = {'optimize': True} + if img.mode in ("RGBA", "LA", "P"): + img = img.convert("RGB") + img.thumbnail( + thumbnail_size, Image.Resampling.LANCZOS + ) + save_kwargs = {"quality": 85, "optimize": True} + if thumbnail_path.suffix.lower() == ".png": + save_kwargs = {"optimize": True} img.save(thumbnail_path, **save_kwargs) generated_count += 1 # Send progress update - if (i + 1) % 5 == 0 or (i + 1) % max(1, total_images // 100) == 0 or i == total_images - 1: + if ( + (i + 1) % 5 == 0 + or (i + 1) % max(1, total_images // 100) == 0 + or i == total_images - 1 + ): elapsed = time.time() - start_time progress_percent = ((i + 1) / total_images) * 100 rate = (i + 1) / elapsed if elapsed > 0 else 0 eta = ((total_images - i - 1) / rate) if rate > 0 else 0 file_info = { - 'name': media_file.name, - 'dir': media_file.parent.name, - 'type': 'video' if is_video else 'image', - 'action': 'skipped' if thumbnail_path.exists() else 'generating' + "name": media_file.name, + "dir": media_file.parent.name, + "type": "video" if is_video else "image", + "action": ( + "skipped" + if thumbnail_path.exists() + else "generating" + ), } - await send_progress('progress', { - 'processed': i + 1, - 'total_images': total_images, - 'generated': generated_count, - 'skipped': skipped_count, - 'percentage': round(progress_percent, 1), - 'rate': round(rate, 1), - 'eta': round(eta, 0), - 'elapsed': round(elapsed, 1), - 'current_file': f"{file_info['dir']}/{file_info['name']}", - 'file_type': file_info['type'], - 'action': file_info['action'] - }) + await send_progress( + "progress", + { + "processed": i + 1, + "total_images": total_images, + "generated": generated_count, + "skipped": skipped_count, + "percentage": round(progress_percent, 1), + "rate": round(rate, 1), + "eta": round(eta, 0), + "elapsed": round(elapsed, 1), + "current_file": f"{file_info['dir']}/{file_info['name']}", + "file_type": file_info["type"], + "action": file_info["action"], + }, + ) if (i + 1) % 50 == 0: - self.logger.info(f"Thumbnail progress: {i+1}/{total_images} ({progress_percent:.1f}%) - Generated: {generated_count}, Skipped: {skipped_count}") + self.logger.info( + f"Thumbnail progress: {i+1}/{total_images} ({progress_percent:.1f}%) - Generated: {generated_count}, Skipped: {skipped_count}" + ) except Exception as e: error_msg = f"Failed to generate thumbnail for {media_file.name}: {str(e)}" @@ -700,49 +782,62 @@ class ImageRoutesMixin: self.logger.warning(error_msg) if len(errors) <= 5: - await send_progress('status', { - 'phase': 'processing', - 'message': f'Error processing {media_file.name}: {str(e)}' - }) + await send_progress( + "status", + { + "phase": "processing", + "message": f"Error processing {media_file.name}: {str(e)}", + }, + ) elapsed_time = time.time() - start_time - completion_message = f'Successfully generated {generated_count} new thumbnails, skipped {skipped_count} existing' + completion_message = f"Successfully generated {generated_count} new thumbnails, skipped {skipped_count} existing" if errors: - completion_message += f' ({len(errors)} errors occurred)' + completion_message += f" ({len(errors)} errors occurred)" - await send_progress('complete', { - 'count': generated_count, - 'skipped': skipped_count, - 'total_images': total_images, - 'errors': errors[:10], - 'error_count': len(errors), - 'elapsed_time': round(elapsed_time, 2), - 'processing_rate': round((total_images / elapsed_time) if elapsed_time > 0 else 0, 2), - 'message': completion_message - }) + await send_progress( + "complete", + { + "count": generated_count, + "skipped": skipped_count, + "total_images": total_images, + "errors": errors[:10], + "error_count": len(errors), + "elapsed_time": round(elapsed_time, 2), + "processing_rate": round( + (total_images / elapsed_time) if elapsed_time > 0 else 0, 2 + ), + "message": completion_message, + }, + ) - self.logger.info(f"Thumbnail generation completed: {generated_count} generated, {skipped_count} skipped, {len(errors)} errors in {elapsed_time:.2f}s") + self.logger.info( + f"Thumbnail generation completed: {generated_count} generated, {skipped_count} skipped, {len(errors)} errors in {elapsed_time:.2f}s" + ) except Exception as e: - await send_progress('error', { - 'error': str(e), - 'message': f'Thumbnail generation failed: {str(e)}' - }) + await send_progress( + "error", + { + "error": str(e), + "message": f"Thumbnail generation failed: {str(e)}", + }, + ) return response except ImportError: - return web.json_response({ - 'success': False, - 'error': 'PIL (Pillow) library not available. Install with: pip install Pillow' - }, status=500) + return web.json_response( + { + "success": False, + "error": "PIL (Pillow) library not available. Install with: pip install Pillow", + }, + status=500, + ) except Exception as e: self.logger.error(f"Generate thumbnails with progress error: {e}") - return web.json_response({ - 'success': False, - 'error': str(e) - }, status=500) + return web.json_response({"success": False, "error": str(e)}, status=500) def _generate_video_thumbnail(self, video_path, thumbnail_path, thumbnail_size): """Generate thumbnail from video file. Returns True if successful.""" @@ -769,8 +864,10 @@ class ImageRoutesMixin: img = Image.fromarray(frame_rgb) img.thumbnail(thumbnail_size, Image.Resampling.LANCZOS) - img.save(thumbnail_path, 'JPEG', quality=85, optimize=True) - self.logger.debug(f"Generated video thumbnail using OpenCV: {thumbnail_path}") + img.save(thumbnail_path, "JPEG", quality=85, optimize=True) + self.logger.debug( + f"Generated video thumbnail using OpenCV: {thumbnail_path}" + ) return True except ImportError: @@ -781,20 +878,29 @@ class ImageRoutesMixin: import subprocess cmd = [ - 'ffmpeg', '-i', str(video_path), - '-ss', '00:00:01', - '-vframes', '1', - '-s', f"{thumbnail_size[0]}x{thumbnail_size[1]}", - '-y', - str(thumbnail_path) + "ffmpeg", + "-i", + str(video_path), + "-ss", + "00:00:01", + "-vframes", + "1", + "-s", + f"{thumbnail_size[0]}x{thumbnail_size[1]}", + "-y", + str(thumbnail_path), ] result = subprocess.run(cmd, capture_output=True, text=True, timeout=30) if result.returncode == 0: - self.logger.debug(f"Generated video thumbnail using ffmpeg: {thumbnail_path}") + self.logger.debug( + f"Generated video thumbnail using ffmpeg: {thumbnail_path}" + ) return True else: - self.logger.warning(f"ffmpeg failed for {video_path}: {result.stderr}") + self.logger.warning( + f"ffmpeg failed for {video_path}: {result.stderr}" + ) except (ImportError, subprocess.TimeoutExpired, FileNotFoundError): pass @@ -803,16 +909,16 @@ class ImageRoutesMixin: try: from PIL import ImageDraw, ImageFont - img = Image.new('RGB', thumbnail_size, color=(50, 50, 50)) + img = Image.new("RGB", thumbnail_size, color=(50, 50, 50)) draw = ImageDraw.Draw(img) center_x, center_y = thumbnail_size[0] // 2, thumbnail_size[1] // 2 triangle_size = min(thumbnail_size) // 4 points = [ - (center_x - triangle_size//2, center_y - triangle_size//2), - (center_x - triangle_size//2, center_y + triangle_size//2), - (center_x + triangle_size//2, center_y) + (center_x - triangle_size // 2, center_y - triangle_size // 2), + (center_x - triangle_size // 2, center_y + triangle_size // 2), + (center_x + triangle_size // 2, center_y), ] draw.polygon(points, fill=(255, 255, 255)) @@ -822,22 +928,33 @@ class ImageRoutesMixin: bbox = draw.textbbox((0, 0), text, font=font) text_width = bbox[2] - bbox[0] draw.text( - (center_x - text_width//2, center_y + triangle_size//2 + 10), - text, fill=(255, 255, 255), font=font + ( + center_x - text_width // 2, + center_y + triangle_size // 2 + 10, + ), + text, + fill=(255, 255, 255), + font=font, ) except (OSError, AttributeError): pass - img.save(thumbnail_path, 'JPEG', quality=85) - self.logger.debug(f"Generated placeholder video thumbnail: {thumbnail_path}") + img.save(thumbnail_path, "JPEG", quality=85) + self.logger.debug( + f"Generated placeholder video thumbnail: {thumbnail_path}" + ) return True except Exception as e: - self.logger.warning(f"Failed to create placeholder thumbnail for {video_path}: {e}") + self.logger.warning( + f"Failed to create placeholder thumbnail for {video_path}: {e}" + ) return False except Exception as e: - self.logger.error(f"Video thumbnail generation failed for {video_path}: {e}") + self.logger.error( + f"Video thumbnail generation failed for {video_path}: {e}" + ) return False async def clear_thumbnails(self, request): @@ -848,40 +965,50 @@ class ImageRoutesMixin: # Find ComfyUI output directory output_dir = self._find_comfyui_output_dir() if not output_dir: - return web.json_response({ - 'success': False, - 'error': 'ComfyUI output directory not found' - }, status=404) + return web.json_response( + {"success": False, "error": "ComfyUI output directory not found"}, + status=404, + ) output_path = Path(output_dir) - thumbnails_dir = output_path / 'thumbnails' + thumbnails_dir = output_path / "thumbnails" if not thumbnails_dir.exists(): - return web.json_response({ - 'success': True, - 'message': 'No thumbnails directory found - nothing to clear', - 'cleared_files': 0 - }) + return web.json_response( + { + "success": True, + "message": "No thumbnails directory found - nothing to clear", + "cleared_files": 0, + } + ) # Verify this is actually our thumbnails directory try: thumbnails_dir_resolved = thumbnails_dir.resolve() output_path_resolved = output_path.resolve() - if (not str(thumbnails_dir_resolved).startswith(str(output_path_resolved)) or - thumbnails_dir.name != 'thumbnails'): - self.logger.error(f"Safety check failed: thumbnails directory path invalid: {thumbnails_dir}") - return web.json_response({ - 'success': False, - 'error': 'Safety check failed: invalid thumbnails directory path' - }, status=400) + if ( + not str(thumbnails_dir_resolved).startswith( + str(output_path_resolved) + ) + or thumbnails_dir.name != "thumbnails" + ): + self.logger.error( + f"Safety check failed: thumbnails directory path invalid: {thumbnails_dir}" + ) + return web.json_response( + { + "success": False, + "error": "Safety check failed: invalid thumbnails directory path", + }, + status=400, + ) except Exception as e: self.logger.error(f"Path validation failed: {e}") - return web.json_response({ - 'success': False, - 'error': 'Path validation failed' - }, status=500) + return web.json_response( + {"success": False, "error": "Path validation failed"}, status=500 + ) # Count files before deletion cleared_count = 0 @@ -891,7 +1018,10 @@ class ImageRoutesMixin: for file in files: file_path = Path(root) / file - if '_thumb' in file.lower() and any(file.lower().endswith(ext) for ext in ['.png', '.jpg', '.jpeg', '.webp', '.gif']): + if "_thumb" in file.lower() and any( + file.lower().endswith(ext) + for ext in [".png", ".jpg", ".jpeg", ".webp", ".gif"] + ): try: file_size = file_path.stat().st_size file_path.unlink() @@ -899,7 +1029,9 @@ class ImageRoutesMixin: cleared_size += file_size self.logger.debug(f"Cleared thumbnail: {file_path}") except Exception as e: - self.logger.warning(f"Failed to delete thumbnail {file_path}: {e}") + self.logger.warning( + f"Failed to delete thumbnail {file_path}: {e}" + ) # Remove empty directories within thumbnails folder try: @@ -913,76 +1045,79 @@ class ImageRoutesMixin: self.logger.debug(f"Directory cleanup info: {e}") def format_size(bytes_size): - for unit in ['B', 'KB', 'MB', 'GB']: + for unit in ["B", "KB", "MB", "GB"]: if bytes_size < 1024.0: return f"{bytes_size:.1f} {unit}" bytes_size /= 1024.0 return f"{bytes_size:.1f} TB" - self.logger.info(f"Thumbnail cleanup: cleared {cleared_count} files ({format_size(cleared_size)})") + self.logger.info( + f"Thumbnail cleanup: cleared {cleared_count} files ({format_size(cleared_size)})" + ) - return web.json_response({ - 'success': True, - 'cleared_files': cleared_count, - 'cleared_size': cleared_size, - 'cleared_size_formatted': format_size(cleared_size), - 'message': f'Cleared {cleared_count} thumbnail files ({format_size(cleared_size)})' - }) + return web.json_response( + { + "success": True, + "cleared_files": cleared_count, + "cleared_size": cleared_size, + "cleared_size_formatted": format_size(cleared_size), + "message": f"Cleared {cleared_count} thumbnail files ({format_size(cleared_size)})", + } + ) except Exception as e: self.logger.error(f"Clear thumbnails error: {e}") - return web.json_response({ - 'success': False, - 'error': str(e) - }, status=500) + return web.json_response({"success": False, "error": str(e)}, status=500) async def link_image_to_prompt(self, request): """Link a generated image to a prompt.""" try: data = await request.json() - prompt_id = data.get('prompt_id') - image_path = data.get('image_path') - metadata = data.get('metadata', {}) + prompt_id = data.get("prompt_id") + image_path = data.get("image_path") + metadata = data.get("metadata", {}) if not prompt_id or not image_path: - return web.json_response({ - 'success': False, - 'error': 'prompt_id and image_path are required' - }, status=400) + return web.json_response( + { + "success": False, + "error": "prompt_id and image_path are required", + }, + status=400, + ) if not os.path.exists(image_path): - return web.json_response({ - 'success': False, - 'error': 'Image file not found' - }, status=404) + return web.json_response( + {"success": False, "error": "Image file not found"}, status=404 + ) - image_id = await self._run_in_executor(self.db.link_image_to_prompt, prompt_id, image_path, metadata) + image_id = await self._run_in_executor( + self.db.link_image_to_prompt, prompt_id, image_path, metadata + ) - return web.json_response({ - 'success': True, - 'image_id': image_id, - 'message': 'Image linked successfully' - }) + return web.json_response( + { + "success": True, + "image_id": image_id, + "message": "Image linked successfully", + } + ) except Exception as e: self.logger.error(f"Link image error: {e}") - return web.json_response({ - 'success': False, - 'error': str(e) - }, status=500) + return web.json_response({"success": False, "error": str(e)}, status=500) async def get_image_prompt(self, request): """Get prompt information for a specific image path.""" try: # Get the image path from URL - raw_image_path = request.match_info.get('image_path', '') + raw_image_path = request.match_info.get("image_path", "") image_path = urllib.parse.unquote(raw_image_path) if not image_path: - return web.json_response({ - 'success': False, - 'error': 'Image path is required' - }, status=400) + return web.json_response( + {"success": False, "error": "Image path is required"}, status=400 + ) # Convert relative path to absolute if needed if not os.path.isabs(image_path): @@ -992,31 +1127,31 @@ class ImageRoutesMixin: # Look up the image in generated_images table try: - prompt_data = await self._run_in_executor(self.db.get_image_prompt_info, image_path) + prompt_data = await self._run_in_executor( + self.db.get_image_prompt_info, image_path + ) if prompt_data: - prompt_data['image_path'] = image_path + prompt_data["image_path"] = image_path if prompt_data: - return web.json_response({'success': True, 'prompt': prompt_data}) + return web.json_response({"success": True, "prompt": prompt_data}) else: - return web.json_response({ - 'success': False, - 'error': 'No prompt found for this image', - 'image_path': image_path - }) + return web.json_response( + { + "success": False, + "error": "No prompt found for this image", + "image_path": image_path, + } + ) except Exception as db_error: self.logger.error(f"Database error in get_image_prompt: {db_error}") - return web.json_response({ - 'success': False, - 'error': 'Database error occurred' - }, status=500) + return web.json_response( + {"success": False, "error": "Database error occurred"}, status=500 + ) except Exception as e: self.logger.error(f"Get image prompt error: {e}") - return web.json_response({ - 'success': False, - 'error': str(e) - }, status=500) + return web.json_response({"success": False, "error": str(e)}, status=500) async def delete_image(self, request): """Delete an image record.""" @@ -1025,21 +1160,18 @@ class ImageRoutesMixin: success = await self._run_in_executor(self.db.delete_image, image_id) if success: - return web.json_response({ - 'success': True, - 'message': 'Image deleted successfully' - }) + return web.json_response( + {"success": True, "message": "Image deleted successfully"} + ) else: - return web.json_response({ - 'success': False, - 'error': 'Image not found' - }, status=404) + return web.json_response( + {"success": False, "error": "Image not found"}, status=404 + ) except ValueError: - return web.json_response({'success': False, 'error': 'Invalid image ID'}, status=400) + return web.json_response( + {"success": False, "error": "Invalid image ID"}, status=400 + ) except Exception as e: self.logger.error(f"Delete image error: {e}") - return web.json_response({ - 'success': False, - 'error': str(e) - }, status=500) + return web.json_response({"success": False, "error": str(e)}, status=500) diff --git a/py/api/logging_routes.py b/py/api/logging_routes.py index 829dd80..9121a86 100644 --- a/py/api/logging_routes.py +++ b/py/api/logging_routes.py @@ -43,7 +43,10 @@ class LoggingRoutesMixin: from ...utils.logging_config import get_logger_manager except ImportError: import sys - current_dir = os.path.dirname(os.path.dirname(os.path.dirname(os.path.abspath(__file__)))) + + current_dir = os.path.dirname( + os.path.dirname(os.path.dirname(os.path.abspath(__file__))) + ) sys.path.insert(0, current_dir) from utils.logging_config import get_logger_manager return get_logger_manager() @@ -53,8 +56,8 @@ class LoggingRoutesMixin: try: logger_manager = self._get_logger_manager() - limit = int(request.query.get('limit', 100)) - level = request.query.get('level', None) + limit = int(request.query.get("limit", 100)) + level = request.query.get("level", None) if limit > 1000: limit = 1000 @@ -63,21 +66,21 @@ class LoggingRoutesMixin: logs = logger_manager.get_recent_logs(limit=limit, level=level) - return web.json_response({ - 'success': True, - 'logs': logs, - 'count': len(logs), - 'level_filter': level, - 'limit': limit - }) + return web.json_response( + { + "success": True, + "logs": logs, + "count": len(logs), + "level_filter": level, + "limit": limit, + } + ) except Exception as e: self.logger.error(f"Get logs error: {e}") - return web.json_response({ - 'success': False, - 'error': str(e), - 'logs': [] - }, status=500) + return web.json_response( + {"success": False, "error": str(e), "logs": []}, status=500 + ) async def get_log_files(self, request): """Get information about available log files.""" @@ -85,58 +88,49 @@ class LoggingRoutesMixin: logger_manager = self._get_logger_manager() log_files = logger_manager.get_log_files() - return web.json_response({ - 'success': True, - 'files': log_files, - 'count': len(log_files) - }) + return web.json_response( + {"success": True, "files": log_files, "count": len(log_files)} + ) except Exception as e: self.logger.error(f"Get log files error: {e}") - return web.json_response({ - 'success': False, - 'error': str(e), - 'files': [] - }, status=500) + return web.json_response( + {"success": False, "error": str(e), "files": []}, status=500 + ) async def download_log_file(self, request): """Download a specific log file.""" try: - filename = request.match_info['filename'] + filename = request.match_info["filename"] logger_manager = self._get_logger_manager() - if not filename or '..' in filename or '/' in filename or '\\' in filename: - return web.json_response({ - 'success': False, - 'error': 'Invalid filename' - }, status=400) + if not filename or ".." in filename or "/" in filename or "\\" in filename: + return web.json_response( + {"success": False, "error": "Invalid filename"}, status=400 + ) log_file_path = logger_manager.log_dir / filename if not log_file_path.exists(): - return web.json_response({ - 'success': False, - 'error': 'Log file not found' - }, status=404) + return web.json_response( + {"success": False, "error": "Log file not found"}, status=404 + ) - with open(log_file_path, 'rb') as f: + with open(log_file_path, "rb") as f: file_content = f.read() return web.Response( body=file_content, - content_type='text/plain', + content_type="text/plain", headers={ - 'Content-Disposition': f'attachment; filename="{filename}"', - 'Content-Length': str(len(file_content)) - } + "Content-Disposition": f'attachment; filename="{filename}"', + "Content-Length": str(len(file_content)), + }, ) except Exception as e: self.logger.error(f"Download log file error: {e}") - return web.json_response({ - 'success': False, - 'error': str(e) - }, status=500) + return web.json_response({"success": False, "error": str(e)}, status=500) async def truncate_logs(self, request): """Truncate all log files.""" @@ -144,18 +138,17 @@ class LoggingRoutesMixin: logger_manager = self._get_logger_manager() results = logger_manager.truncate_logs() - return web.json_response({ - 'success': True, - 'message': f"Truncated {len(results['truncated'])} log files", - 'results': results - }) + return web.json_response( + { + "success": True, + "message": f"Truncated {len(results['truncated'])} log files", + "results": results, + } + ) except Exception as e: self.logger.error(f"Truncate logs error: {e}") - return web.json_response({ - 'success': False, - 'error': str(e) - }, status=500) + return web.json_response({"success": False, "error": str(e)}, status=500) async def get_log_config(self, request): """Get current logging configuration.""" @@ -163,17 +156,11 @@ class LoggingRoutesMixin: logger_manager = self._get_logger_manager() config = logger_manager.get_config() - return web.json_response({ - 'success': True, - 'config': config - }) + return web.json_response({"success": True, "config": config}) except Exception as e: self.logger.error(f"Get log config error: {e}") - return web.json_response({ - 'success': False, - 'error': str(e) - }, status=500) + return web.json_response({"success": False, "error": str(e)}, status=500) async def update_log_config(self, request): """Update logging configuration.""" @@ -181,29 +168,31 @@ class LoggingRoutesMixin: data = await request.json() logger_manager = self._get_logger_manager() - if 'level' in data: - valid_levels = ['DEBUG', 'INFO', 'WARNING', 'ERROR', 'CRITICAL'] - if data['level'].upper() not in valid_levels: - return web.json_response({ - 'success': False, - 'error': f'Invalid log level. Must be one of: {valid_levels}' - }, status=400) - data['level'] = data['level'].upper() + if "level" in data: + valid_levels = ["DEBUG", "INFO", "WARNING", "ERROR", "CRITICAL"] + if data["level"].upper() not in valid_levels: + return web.json_response( + { + "success": False, + "error": f"Invalid log level. Must be one of: {valid_levels}", + }, + status=400, + ) + data["level"] = data["level"].upper() logger_manager.update_config(data) - return web.json_response({ - 'success': True, - 'message': 'Logging configuration updated', - 'config': logger_manager.get_config() - }) + return web.json_response( + { + "success": True, + "message": "Logging configuration updated", + "config": logger_manager.get_config(), + } + ) except Exception as e: self.logger.error(f"Update log config error: {e}") - return web.json_response({ - 'success': False, - 'error': str(e) - }, status=500) + return web.json_response({"success": False, "error": str(e)}, status=500) async def get_log_stats(self, request): """Get logging statistics.""" @@ -211,14 +200,8 @@ class LoggingRoutesMixin: logger_manager = self._get_logger_manager() stats = logger_manager.get_log_stats() - return web.json_response({ - 'success': True, - 'stats': stats - }) + return web.json_response({"success": True, "stats": stats}) except Exception as e: self.logger.error(f"Get log stats error: {e}") - return web.json_response({ - 'success': False, - 'error': str(e) - }, status=500) + return web.json_response({"success": False, "error": str(e)}, status=500) diff --git a/py/api/prompts.py b/py/api/prompts.py index 34cd31d..ed7b337 100644 --- a/py/api/prompts.py +++ b/py/api/prompts.py @@ -153,22 +153,26 @@ class PromptRoutesMixin: elif limit < 1: limit = 1 - results = await self._run_in_executor(self.db.get_recent_prompts, limit=limit, offset=offset) - self._enrich_prompt_images(results['prompts']) + results = await self._run_in_executor( + self.db.get_recent_prompts, limit=limit, offset=offset + ) + self._enrich_prompt_images(results["prompts"]) - return web.json_response({ - "success": True, - "results": results['prompts'], - "pagination": { - "total": results['total'], - "limit": results['limit'], - "offset": results['offset'], - "page": results['page'], - "total_pages": results['total_pages'], - "has_more": results['has_more'], - "count": len(results['prompts']) + return web.json_response( + { + "success": True, + "results": results["prompts"], + "pagination": { + "total": results["total"], + "limit": results["limit"], + "offset": results["offset"], + "page": results["page"], + "total_pages": results["total_pages"], + "has_more": results["has_more"], + "count": len(results["prompts"]), + }, } - }) + ) except Exception as e: self.logger.error(f"Recent prompts error: {e}", exc_info=True) @@ -177,7 +181,7 @@ class PromptRoutesMixin: "success": False, "error": f"Failed to get recent prompts: {str(e)}", "results": [], - "pagination": {"total": 0, "page": 1, "total_pages": 0} + "pagination": {"total": 0, "page": 1, "total_pages": 0}, }, status=500, ) @@ -190,7 +194,11 @@ class PromptRoutesMixin: except Exception as e: self.logger.error(f"Categories error: {e}") return web.json_response( - {"success": False, "error": f"Failed to get categories: {str(e)}", "categories": []}, + { + "success": False, + "error": f"Failed to get categories: {str(e)}", + "categories": [], + }, status=500, ) @@ -202,7 +210,11 @@ class PromptRoutesMixin: except Exception as e: self.logger.error(f"Tags error: {e}") return web.json_response( - {"success": False, "error": f"Failed to get tags: {str(e)}", "tags": []}, + { + "success": False, + "error": f"Failed to get tags: {str(e)}", + "tags": [], + }, status=500, ) @@ -214,25 +226,32 @@ class PromptRoutesMixin: offset = int(request.query.get("offset", 0)) except (ValueError, TypeError): return web.json_response( - {"success": False, "error": "Invalid limit or offset parameter"}, status=400 + {"success": False, "error": "Invalid limit or offset parameter"}, + status=400, ) search = request.query.get("search", "").strip() or None sort = request.query.get("sort", "alpha_asc") - result = await self._run_in_executor(self.db.get_tags_with_counts, limit, offset, search, sort) - untagged_count = await self._run_in_executor(self.db.get_untagged_prompts_count) + result = await self._run_in_executor( + self.db.get_tags_with_counts, limit, offset, search, sort + ) + untagged_count = await self._run_in_executor( + self.db.get_untagged_prompts_count + ) - return web.json_response({ - "success": True, - "tags": result['tags'], - "untagged_count": untagged_count, - "pagination": { - "total": result['total'], - "limit": result['limit'], - "offset": result['offset'], - "has_more": result['has_more'] + return web.json_response( + { + "success": True, + "tags": result["tags"], + "untagged_count": untagged_count, + "pagination": { + "total": result["total"], + "limit": result["limit"], + "offset": result["offset"], + "has_more": result["has_more"], + }, } - }) + ) except Exception as e: self.logger.error(f"Tags stats error: {e}", exc_info=True) return web.json_response({"success": False, "error": str(e)}, status=500) @@ -241,6 +260,7 @@ class PromptRoutesMixin: """Get prompts for a single tag.""" try: from urllib.parse import unquote + tag_name = unquote(request.match_info.get("tag_name", "")) if not tag_name: return web.json_response( @@ -252,23 +272,28 @@ class PromptRoutesMixin: offset = int(request.query.get("offset", 0)) except (ValueError, TypeError): return web.json_response( - {"success": False, "error": "Invalid limit or offset parameter"}, status=400 + {"success": False, "error": "Invalid limit or offset parameter"}, + status=400, ) - result = await self._run_in_executor(self.db.get_prompts_by_tags, [tag_name], 'and', limit, offset) - self._enrich_prompt_images(result['prompts']) + result = await self._run_in_executor( + self.db.get_prompts_by_tags, [tag_name], "and", limit, offset + ) + self._enrich_prompt_images(result["prompts"]) - return web.json_response({ - "success": True, - "tag": tag_name, - "prompts": result['prompts'], - "pagination": { - "total": result['total'], - "limit": result['limit'], - "offset": result['offset'], - "has_more": result['has_more'] + return web.json_response( + { + "success": True, + "tag": tag_name, + "prompts": result["prompts"], + "pagination": { + "total": result["total"], + "limit": result["limit"], + "offset": result["offset"], + "has_more": result["has_more"], + }, } - }) + ) except Exception as e: self.logger.error(f"Tag prompts error: {e}", exc_info=True) return web.json_response({"success": False, "error": str(e)}, status=500) @@ -284,22 +309,30 @@ class PromptRoutesMixin: offset = int(request.query.get("offset", 0)) except (ValueError, TypeError): return web.json_response( - {"success": False, "error": "Invalid limit or offset parameter"}, status=400 + { + "success": False, + "error": "Invalid limit or offset parameter", + }, + status=400, ) - result = await self._run_in_executor(self.db.get_untagged_prompts, limit, offset) - self._enrich_prompt_images(result['prompts']) - return web.json_response({ - "success": True, - "tags": [], - "mode": "untagged", - "prompts": result['prompts'], - "pagination": { - "total": result['total'], - "limit": result['limit'], - "offset": result['offset'], - "has_more": result['has_more'] + result = await self._run_in_executor( + self.db.get_untagged_prompts, limit, offset + ) + self._enrich_prompt_images(result["prompts"]) + return web.json_response( + { + "success": True, + "tags": [], + "mode": "untagged", + "prompts": result["prompts"], + "pagination": { + "total": result["total"], + "limit": result["limit"], + "offset": result["offset"], + "has_more": result["has_more"], + }, } - }) + ) tags_str = request.query.get("tags", "").strip() if not tags_str: @@ -317,24 +350,29 @@ class PromptRoutesMixin: offset = int(request.query.get("offset", 0)) except (ValueError, TypeError): return web.json_response( - {"success": False, "error": "Invalid limit or offset parameter"}, status=400 + {"success": False, "error": "Invalid limit or offset parameter"}, + status=400, ) - result = await self._run_in_executor(self.db.get_prompts_by_tags, tags_list, mode, limit, offset) - self._enrich_prompt_images(result['prompts']) + result = await self._run_in_executor( + self.db.get_prompts_by_tags, tags_list, mode, limit, offset + ) + self._enrich_prompt_images(result["prompts"]) - return web.json_response({ - "success": True, - "tags": tags_list, - "mode": mode, - "prompts": result['prompts'], - "pagination": { - "total": result['total'], - "limit": result['limit'], - "offset": result['offset'], - "has_more": result['has_more'] + return web.json_response( + { + "success": True, + "tags": tags_list, + "mode": mode, + "prompts": result["prompts"], + "pagination": { + "total": result["total"], + "limit": result["limit"], + "offset": result["offset"], + "has_more": result["has_more"], + }, } - }) + ) except Exception as e: self.logger.error(f"Tags filter error: {e}", exc_info=True) return web.json_response({"success": False, "error": str(e)}, status=500) @@ -343,6 +381,7 @@ class PromptRoutesMixin: """Rename a tag across all prompts.""" try: from urllib.parse import unquote + tag_name = unquote(request.match_info.get("tag_name", "")) if not tag_name: return web.json_response( @@ -361,16 +400,20 @@ class PromptRoutesMixin: {"success": False, "error": "New tag name required"}, status=400 ) - result = await self._run_in_executor(self.db.rename_tag_all_prompts, tag_name, new_name) + result = await self._run_in_executor( + self.db.rename_tag_all_prompts, tag_name, new_name + ) resp = { "success": True, "old_name": tag_name, "new_name": new_name, - "affected_count": result['affected_count'] + "affected_count": result["affected_count"], } - if result.get('skipped_count', 0) > 0: - resp['skipped_count'] = result['skipped_count'] - resp['warning'] = f"{result['skipped_count']} prompt(s) had corrupted tag data and were skipped" + if result.get("skipped_count", 0) > 0: + resp["skipped_count"] = result["skipped_count"] + resp["warning"] = ( + f"{result['skipped_count']} prompt(s) had corrupted tag data and were skipped" + ) return web.json_response(resp) except Exception as e: self.logger.error(f"Rename tag error: {e}", exc_info=True) @@ -380,21 +423,26 @@ class PromptRoutesMixin: """Delete a tag from all prompts.""" try: from urllib.parse import unquote + tag_name = unquote(request.match_info.get("tag_name", "")) if not tag_name: return web.json_response( {"success": False, "error": "Tag name required"}, status=400 ) - result = await self._run_in_executor(self.db.delete_tag_all_prompts, tag_name) + result = await self._run_in_executor( + self.db.delete_tag_all_prompts, tag_name + ) resp = { "success": True, "tag_name": tag_name, - "affected_count": result['affected_count'] + "affected_count": result["affected_count"], } - if result.get('skipped_count', 0) > 0: - resp['skipped_count'] = result['skipped_count'] - resp['warning'] = f"{result['skipped_count']} prompt(s) had corrupted tag data and were skipped" + if result.get("skipped_count", 0) > 0: + resp["skipped_count"] = result["skipped_count"] + resp["warning"] = ( + f"{result['skipped_count']} prompt(s) had corrupted tag data and were skipped" + ) return web.json_response(resp) except Exception as e: self.logger.error(f"Delete tag error: {e}", exc_info=True) @@ -421,16 +469,20 @@ class PromptRoutesMixin: {"success": False, "error": "Target tag required"}, status=400 ) - result = await self._run_in_executor(self.db.merge_tags, source_tags, target_tag) + result = await self._run_in_executor( + self.db.merge_tags, source_tags, target_tag + ) resp = { "success": True, "target_tag": target_tag, - "affected_count": result['affected_count'], - "tags_merged": result['tags_merged'] + "affected_count": result["affected_count"], + "tags_merged": result["tags_merged"], } - if result.get('skipped_count', 0) > 0: - resp['skipped_count'] = result['skipped_count'] - resp['warning'] = f"{result['skipped_count']} prompt(s) had corrupted tag data and were skipped" + if result.get("skipped_count", 0) > 0: + resp["skipped_count"] = result["skipped_count"] + resp["warning"] = ( + f"{result['skipped_count']} prompt(s) had corrupted tag data and were skipped" + ) return web.json_response(resp) except Exception as e: self.logger.error(f"Merge tags error: {e}", exc_info=True) @@ -440,9 +492,13 @@ class PromptRoutesMixin: """Save a new prompt with metadata and duplicate detection.""" try: from utils.validators import ( - validate_prompt_text, validate_rating, validate_tags, - validate_category, sanitize_input, + validate_prompt_text, + validate_rating, + validate_tags, + validate_category, + sanitize_input, ) + data = await request.json() text = data.get("text", "").strip() @@ -469,25 +525,30 @@ class PromptRoutesMixin: text = sanitize_input(text) from utils.hashing import generate_prompt_hash + prompt_hash = generate_prompt_hash(text) - existing = await self._run_in_executor(self.db.get_prompt_by_hash, prompt_hash) + existing = await self._run_in_executor( + self.db.get_prompt_by_hash, prompt_hash + ) if existing: if any([category, tags, rating, notes]): await self._run_in_executor( self.db.update_prompt_metadata, - prompt_id=existing['id'], + prompt_id=existing["id"], category=category, tags=tags, rating=rating, - notes=notes + notes=notes, ) - return web.json_response({ - "success": True, - "prompt_id": existing['id'], - "message": "Prompt already exists, metadata updated", - "is_duplicate": True - }) + return web.json_response( + { + "success": True, + "prompt_id": existing["id"], + "message": "Prompt already exists, metadata updated", + "is_duplicate": True, + } + ) prompt_id = await self._run_in_executor( self.db.save_prompt, @@ -499,11 +560,13 @@ class PromptRoutesMixin: prompt_hash=prompt_hash, ) - return web.json_response({ - "success": True, - "prompt_id": prompt_id, - "message": "Prompt saved successfully", - }) + return web.json_response( + { + "success": True, + "prompt_id": prompt_id, + "message": "Prompt saved successfully", + } + ) except Exception as e: self.logger.error(f"Save error: {e}", exc_info=True) @@ -524,7 +587,10 @@ class PromptRoutesMixin: ) else: return web.json_response( - {"success": False, "error": "Prompt not found or could not be deleted"}, + { + "success": False, + "error": "Prompt not found or could not be deleted", + }, status=404, ) @@ -562,7 +628,9 @@ class PromptRoutesMixin: new_text = sanitize_input(new_text) - updated = await self._run_in_executor(self.db.update_prompt_text, prompt_id, new_text) + updated = await self._run_in_executor( + self.db.update_prompt_text, prompt_id, new_text + ) if updated: return web.json_response( {"success": True, "message": "Prompt updated successfully"} @@ -599,7 +667,9 @@ class PromptRoutesMixin: {"success": False, "error": str(ve)}, status=400 ) - updated = await self._run_in_executor(self.db.update_prompt_rating, prompt_id, rating) + updated = await self._run_in_executor( + self.db.update_prompt_rating, prompt_id, rating + ) if updated: return web.json_response( {"success": True, "message": "Rating updated successfully"} @@ -644,7 +714,9 @@ class PromptRoutesMixin: if new_tag not in current_tags: current_tags.append(new_tag) - await self._run_in_executor(self.db.set_prompt_tags, prompt_id, current_tags) + await self._run_in_executor( + self.db.set_prompt_tags, prompt_id, current_tags + ) return web.json_response( {"success": True, "message": "Tag added successfully"} @@ -674,7 +746,8 @@ class PromptRoutesMixin: if not new_tags or not isinstance(new_tags, list): return web.json_response( - {"success": False, "error": "Tags must be a non-empty list"}, status=400 + {"success": False, "error": "Tags must be a non-empty list"}, + status=400, ) prompt = await self._run_in_executor(self.db.get_prompt_by_id, prompt_id) @@ -695,7 +768,9 @@ class PromptRoutesMixin: tags_added += 1 if tags_added > 0: - await self._run_in_executor(self.db.set_prompt_tags, prompt_id, current_tags) + await self._run_in_executor( + self.db.set_prompt_tags, prompt_id, current_tags + ) message = f"{tags_added} tag(s) added successfully" if tags_added == 0: @@ -734,7 +809,9 @@ class PromptRoutesMixin: if tag_to_remove in current_tags: current_tags.remove(tag_to_remove) - await self._run_in_executor(self.db.set_prompt_tags, prompt_id, current_tags) + await self._run_in_executor( + self.db.set_prompt_tags, prompt_id, current_tags + ) return web.json_response( {"success": True, "message": "Tag removed successfully"} @@ -762,13 +839,17 @@ class PromptRoutesMixin: {"success": False, "error": "No prompt IDs provided"}, status=400 ) - deleted_count = await self._run_in_executor(self.db.bulk_delete_prompts, prompt_ids) + deleted_count = await self._run_in_executor( + self.db.bulk_delete_prompts, prompt_ids + ) - return web.json_response({ - "success": True, - "message": f"Deleted {deleted_count} prompts", - "deleted_count": deleted_count, - }) + return web.json_response( + { + "success": True, + "message": f"Deleted {deleted_count} prompts", + "deleted_count": deleted_count, + } + ) except Exception as e: self.logger.error(f"Bulk delete error: {e}") @@ -790,13 +871,17 @@ class PromptRoutesMixin: status=400, ) - updated_count = await self._run_in_executor(self.db.bulk_add_tags, prompt_ids, new_tags) + updated_count = await self._run_in_executor( + self.db.bulk_add_tags, prompt_ids, new_tags + ) - return web.json_response({ - "success": True, - "message": f"Added tags to {updated_count} prompts", - "updated_count": updated_count, - }) + return web.json_response( + { + "success": True, + "message": f"Added tags to {updated_count} prompts", + "updated_count": updated_count, + } + ) except Exception as e: self.logger.error(f"Bulk add tags error: {e}") @@ -816,13 +901,17 @@ class PromptRoutesMixin: {"success": False, "error": "No prompt IDs provided"}, status=400 ) - updated_count = await self._run_in_executor(self.db.bulk_set_category, prompt_ids, category) + updated_count = await self._run_in_executor( + self.db.bulk_set_category, prompt_ids, category + ) - return web.json_response({ - "success": True, - "message": f"Set category for {updated_count} prompts", - "updated_count": updated_count, - }) + return web.json_response( + { + "success": True, + "message": f"Set category for {updated_count} prompts", + "updated_count": updated_count, + } + ) except Exception as e: self.logger.error(f"Bulk set category error: {e}") diff --git a/py/autotag.py b/py/autotag.py index d74a0f9..32302e7 100644 --- a/py/autotag.py +++ b/py/autotag.py @@ -16,13 +16,16 @@ try: from ..utils.logging_config import get_logger except ImportError: import logging + def get_logger(name: str) -> logging.Logger: logger = logging.getLogger(name) if not logger.handlers: handler = logging.StreamHandler() - handler.setFormatter(logging.Formatter( - '%(asctime)s - %(name)s - %(levelname)s - %(message)s' - )) + handler.setFormatter( + logging.Formatter( + "%(asctime)s - %(name)s - %(levelname)s - %(message)s" + ) + ) logger.addHandler(handler) logger.setLevel(logging.DEBUG) return logger @@ -46,7 +49,7 @@ MODELS = { "size": "~16GB", "repo": "fancyfeast/llama-joycaption-beta-one-hf-llava", "subdir": "llama-joycaption-beta-one-hf-llava", - } + }, } # Default prompts @@ -84,7 +87,7 @@ class AutoTagService: models_dir: Directory for storing models. If None, uses ComfyUI's folder_paths.models_dir / "LLM" path. """ - self.logger = get_logger('autotag.service') + self.logger = get_logger("autotag.service") # Determine models directory if models_dir: @@ -93,11 +96,14 @@ class AutoTagService: # Use ComfyUI's folder_paths system try: import folder_paths + self.models_dir = Path(folder_paths.models_dir) / "LLM" except ImportError: # Fallback for standalone usage (not running in ComfyUI) self.logger.warning("folder_paths not available, using fallback path") - self.models_dir = Path(__file__).parent.parent.parent.parent / "models" / "LLM" + self.models_dir = ( + Path(__file__).parent.parent.parent.parent / "models" / "LLM" + ) self.logger.info(f"AutoTag service initialized. Models dir: {self.models_dir}") @@ -148,29 +154,27 @@ class AutoTagService: for model_type, config in MODELS.items(): model_status = { - 'name': config['name'], - 'description': config['description'], - 'size': config['size'], - 'downloaded': False, - 'model_path': None + "name": config["name"], + "description": config["description"], + "size": config["size"], + "downloaded": False, + "model_path": None, } - if model_type == 'gguf': + if model_type == "gguf": model_exists, mmproj_exists = self._check_gguf_models() - model_status['model_exists'] = model_exists - model_status['mmproj_exists'] = mmproj_exists - model_status['downloaded'] = model_exists and mmproj_exists - if model_status['downloaded']: - model_status['model_path'] = str( - self.models_dir / config['subdir'] / config['filename'] + model_status["model_exists"] = model_exists + model_status["mmproj_exists"] = mmproj_exists + model_status["downloaded"] = model_exists and mmproj_exists + if model_status["downloaded"]: + model_status["model_path"] = str( + self.models_dir / config["subdir"] / config["filename"] ) else: # hf - model_status['downloaded'] = self._check_hf_model() - if model_status['downloaded']: + model_status["downloaded"] = self._check_hf_model() + if model_status["downloaded"]: # Get the actual path (local or cache) - model_status['model_path'] = str( - self._get_hf_model_path() - ) + model_status["model_path"] = str(self._get_hf_model_path()) status[model_type] = model_status @@ -182,10 +186,10 @@ class AutoTagService: Returns: Tuple of (model_exists, mmproj_exists) """ - config = MODELS['gguf'] - gguf_dir = self.models_dir / config['subdir'] - model_path = gguf_dir / config['filename'] - mmproj_path = gguf_dir / config['mmproj_filename'] + config = MODELS["gguf"] + gguf_dir = self.models_dir / config["subdir"] + model_path = gguf_dir / config["filename"] + mmproj_path = gguf_dir / config["mmproj_filename"] return model_path.exists(), mmproj_path.exists() def _check_hf_model(self) -> bool: @@ -196,17 +200,17 @@ class AutoTagService: Returns: True if model directory contains config.json """ - config = MODELS['hf'] + config = MODELS["hf"] # Check local directory first - model_dir = self.models_dir / config['subdir'] + model_dir = self.models_dir / config["subdir"] self.logger.debug(f"Checking local HF model path: {model_dir}") if (model_dir / "config.json").exists(): self.logger.debug("Found HF model in local directory") return True # Check HuggingFace cache as fallback - hf_cache_path = self._get_hf_cache_path(config['repo']) + hf_cache_path = self._get_hf_cache_path(config["repo"]) if hf_cache_path: self.logger.debug(f"Found HF model in cache: {hf_cache_path}") return True @@ -251,15 +255,15 @@ class AutoTagService: Returns: Path to the model directory, or None if not found """ - config = MODELS['hf'] + config = MODELS["hf"] # Check local directory first - model_dir = self.models_dir / config['subdir'] + model_dir = self.models_dir / config["subdir"] if (model_dir / "config.json").exists(): return model_dir # Check HuggingFace cache as fallback - cache_path = self._get_hf_cache_path(config['repo']) + cache_path = self._get_hf_cache_path(config["repo"]) if cache_path: return cache_path @@ -268,7 +272,7 @@ class AutoTagService: def download_model( self, model_type: str, - progress_callback: Optional[Callable[[str, float], None]] = None + progress_callback: Optional[Callable[[str, float], None]] = None, ) -> bool: """Download a model with optional progress updates. @@ -283,7 +287,9 @@ class AutoTagService: ValueError: If model_type is not valid """ if model_type not in MODELS: - raise ValueError(f"Invalid model type: {model_type}. Must be 'gguf' or 'hf'") + raise ValueError( + f"Invalid model type: {model_type}. Must be 'gguf' or 'hf'" + ) try: from huggingface_hub import hf_hub_download, snapshot_download @@ -294,7 +300,7 @@ class AutoTagService: return False try: - if model_type == 'gguf': + if model_type == "gguf": return self._download_gguf_models(progress_callback) else: return self._download_hf_model(progress_callback) @@ -305,18 +311,17 @@ class AutoTagService: return False def _download_gguf_models( - self, - progress_callback: Optional[Callable[[str, float], None]] = None + self, progress_callback: Optional[Callable[[str, float], None]] = None ) -> bool: """Download GGUF model and mmproj files.""" from huggingface_hub import hf_hub_download - config = MODELS['gguf'] - gguf_dir = self.models_dir / config['subdir'] + config = MODELS["gguf"] + gguf_dir = self.models_dir / config["subdir"] gguf_dir.mkdir(parents=True, exist_ok=True) - model_path = gguf_dir / config['filename'] - mmproj_path = gguf_dir / config['mmproj_filename'] + model_path = gguf_dir / config["filename"] + mmproj_path = gguf_dir / config["mmproj_filename"] # Download main model if not model_path.exists(): @@ -325,10 +330,10 @@ class AutoTagService: self.logger.info(f"Downloading GGUF model: {config['filename']}") hf_hub_download( - repo_id=config['repo'], - filename=config['filename'], + repo_id=config["repo"], + filename=config["filename"], local_dir=str(gguf_dir), - local_dir_use_symlinks=False + local_dir_use_symlinks=False, ) self.logger.info("GGUF model downloaded") @@ -342,10 +347,10 @@ class AutoTagService: self.logger.info(f"Downloading mmproj: {config['mmproj_filename']}") hf_hub_download( - repo_id=config['mmproj_repo'], - filename=config['mmproj_filename'], + repo_id=config["mmproj_repo"], + filename=config["mmproj_filename"], local_dir=str(gguf_dir), - local_dir_use_symlinks=False + local_dir_use_symlinks=False, ) self.logger.info("mmproj downloaded") @@ -355,14 +360,13 @@ class AutoTagService: return True def _download_hf_model( - self, - progress_callback: Optional[Callable[[str, float], None]] = None + self, progress_callback: Optional[Callable[[str, float], None]] = None ) -> bool: """Download HuggingFace model.""" from huggingface_hub import snapshot_download - config = MODELS['hf'] - model_dir = self.models_dir / config['subdir'] + config = MODELS["hf"] + model_dir = self.models_dir / config["subdir"] if not self._check_hf_model(): if progress_callback: @@ -370,9 +374,9 @@ class AutoTagService: self.logger.info(f"Downloading HF model: {config['repo']}") snapshot_download( - repo_id=config['repo'], + repo_id=config["repo"], local_dir=str(model_dir), - local_dir_use_symlinks=False + local_dir_use_symlinks=False, ) self.logger.info("HF model downloaded") @@ -403,11 +407,11 @@ class AutoTagService: self.unload_model() status = self.get_models_status() - if not status[model_type]['downloaded']: + if not status[model_type]["downloaded"]: raise RuntimeError(f"Model {model_type} not downloaded") try: - if model_type == 'gguf': + if model_type == "gguf": self._tagger = self._load_gguf_tagger(use_gpu) else: self._tagger = self._load_hf_tagger() @@ -427,10 +431,10 @@ class AutoTagService: from llama_cpp import Llama from llama_cpp.llama_chat_format import Llava15ChatHandler - config = MODELS['gguf'] - gguf_dir = self.models_dir / config['subdir'] - model_path = gguf_dir / config['filename'] - mmproj_path = gguf_dir / config['mmproj_filename'] + config = MODELS["gguf"] + gguf_dir = self.models_dir / config["subdir"] + model_path = gguf_dir / config["filename"] + mmproj_path = gguf_dir / config["mmproj_filename"] self.logger.info("Loading GGUF model...") n_gpu_layers = -1 if use_gpu else 0 @@ -447,17 +451,23 @@ class AutoTagService: ) self.logger.info("GGUF model loaded") - return ('gguf', tagger) + return ("gguf", tagger) def _load_hf_tagger(self, quantization: str = "8bit"): """Load HuggingFace-based tagger.""" import torch - from transformers import AutoProcessor, LlavaForConditionalGeneration, BitsAndBytesConfig + from transformers import ( + AutoProcessor, + LlavaForConditionalGeneration, + BitsAndBytesConfig, + ) # Get the actual model path (local or cache) model_path = self._get_hf_model_path() if model_path is None: - raise RuntimeError("HuggingFace model not found in local directory or cache") + raise RuntimeError( + "HuggingFace model not found in local directory or cache" + ) self.logger.info(f"Loading HuggingFace model from {model_path}...") device = "cuda" if torch.cuda.is_available() else "cpu" @@ -476,20 +486,18 @@ class AutoTagService: str(model_path), torch_dtype=torch.float16, quantization_config=qnt_config, - **model_kwargs + **model_kwargs, ) else: model = LlavaForConditionalGeneration.from_pretrained( - str(model_path), - torch_dtype=torch.bfloat16, - **model_kwargs + str(model_path), torch_dtype=torch.bfloat16, **model_kwargs ) model.eval() self.logger.info("HuggingFace model loaded") # Track the compute dtype for pixel_values conversion compute_dtype = torch.float16 if quantization == "8bit" else torch.bfloat16 - return ('hf', (model, processor, device, compute_dtype)) + return ("hf", (model, processor, device, compute_dtype)) def unload_model(self): """Unload the current model and free memory.""" @@ -503,6 +511,7 @@ class AutoTagService: # Try to clear CUDA cache if available try: import torch + if torch.cuda.is_available(): torch.cuda.empty_cache() except ImportError: @@ -518,11 +527,7 @@ class AutoTagService: """Get the type of currently loaded model.""" return self._current_model_type - def generate_tags( - self, - image_path: str, - prompt: Optional[str] = None - ) -> List[str]: + def generate_tags(self, image_path: str, prompt: Optional[str] = None) -> List[str]: """Generate tags for an image. Args: @@ -546,17 +551,19 @@ class AutoTagService: # Load image image = Image.open(image_path) - if image.mode != 'RGB': - image = image.convert('RGB') + if image.mode != "RGB": + image = image.convert("RGB") # Generate based on model type model_type, tagger_obj = self._tagger - if model_type == 'gguf': + if model_type == "gguf": raw_tags = self._generate_gguf(tagger_obj, image, use_prompt) else: model, processor, device, compute_dtype = tagger_obj - raw_tags = self._generate_hf(model, processor, device, compute_dtype, image, use_prompt) + raw_tags = self._generate_hf( + model, processor, device, compute_dtype, image, use_prompt + ) # Parse tags from response tags = self._parse_tags(raw_tags) @@ -572,9 +579,9 @@ class AutoTagService: # Encode to base64 buffer = io.BytesIO() - image.save(buffer, format='PNG') + image.save(buffer, format="PNG") buffer.seek(0) - img_base64 = base64.b64encode(buffer.read()).decode('utf-8') + img_base64 = base64.b64encode(buffer.read()).decode("utf-8") data_uri = f"data:image/png;base64,{img_base64}" # Create message @@ -584,9 +591,9 @@ class AutoTagService: "role": "user", "content": [ {"type": "text", "text": prompt}, - {"type": "image_url", "image_url": {"url": data_uri}} - ] - } + {"type": "image_url", "image_url": {"url": data_uri}}, + ], + }, ] # Generate @@ -608,7 +615,7 @@ class AutoTagService: device: str, compute_dtype, image: Image.Image, - prompt: str + prompt: str, ) -> str: """Generate tags using HuggingFace model.""" import torch @@ -624,14 +631,14 @@ class AutoTagService: convo, tokenize=False, add_generation_prompt=True ) - inputs = processor( - text=[convo_string], images=[image], return_tensors="pt" - ).to(device) + inputs = processor(text=[convo_string], images=[image], return_tensors="pt").to( + device + ) # Convert pixel_values to the model's compute dtype # (float16 for 8-bit quantized, bfloat16 for non-quantized) - if 'pixel_values' in inputs and inputs['pixel_values'] is not None: - inputs['pixel_values'] = inputs['pixel_values'].to(compute_dtype) + if "pixel_values" in inputs and inputs["pixel_values"] is not None: + inputs["pixel_values"] = inputs["pixel_values"].to(compute_dtype) with torch.inference_mode(), torch.cuda.amp.autocast(enabled=True): generate_ids = model.generate( @@ -643,8 +650,10 @@ class AutoTagService: use_cache=True, )[0] - generate_ids = generate_ids[inputs['input_ids'].shape[1]:] - return processor.tokenizer.decode(generate_ids, skip_special_tokens=True).strip() + generate_ids = generate_ids[inputs["input_ids"].shape[1] :] + return processor.tokenizer.decode( + generate_ids, skip_special_tokens=True + ).strip() def _parse_tags(self, raw_output: str) -> List[str]: """Parse raw model output into a list of clean tags. @@ -657,22 +666,22 @@ class AutoTagService: """ # Patterns to exclude exclude_prefixes = ( - 'copyright:', - 'meta:', - 'photo_', - 'photo:', + "copyright:", + "meta:", + "photo_", + "photo:", ) # Split by common delimiters tags = [] # Handle comma-separated tags - for part in raw_output.split(','): + for part in raw_output.split(","): tag = part.strip().lower() # Remove any quotes or extra characters - tag = tag.strip('"\'') + tag = tag.strip("\"'") # Replace spaces with underscores (Danbooru style) - tag = tag.replace(' ', '_') + tag = tag.replace(" ", "_") # Remove empty tags if tag and len(tag) > 1: # Filter out unwanted tag patterns diff --git a/py/config.py b/py/config.py index 10ece56..bea4b0d 100644 --- a/py/config.py +++ b/py/config.py @@ -23,6 +23,7 @@ extension_name = "PromptManager" # Get server instance and routes (same pattern as ComfyUI_Assets) from server import PromptServer + server_instance = PromptServer.instance routes = server_instance.routes @@ -37,24 +38,25 @@ try: from ..utils.logging_config import get_logger except ImportError: import sys + current_dir = os.path.dirname(os.path.dirname(os.path.abspath(__file__))) sys.path.insert(0, current_dir) from utils.logging_config import get_logger # Initialize logger for config operations -config_logger = get_logger('prompt_manager.config') +config_logger = get_logger("prompt_manager.config") class GalleryConfig: """Configuration class for the gallery monitoring and image processing system. - + This class manages all settings related to automatic image monitoring, prompt tracking, database cleanup, web interface display, and performance optimization for the gallery functionality. - + All configuration values are class attributes that can be modified at runtime or loaded from external configuration files. - + Attributes: MONITORING_ENABLED (bool): Enable/disable automatic image monitoring MONITORING_DIRECTORIES (List[str]): Directories to monitor for new images @@ -71,86 +73,88 @@ class GalleryConfig: MAX_CONCURRENT_PROCESSING (int): Maximum concurrent image processing tasks METADATA_EXTRACTION_TIMEOUT (int): Timeout for metadata extraction operations """ - + # Image monitoring settings MONITORING_ENABLED = True MONITORING_DIRECTORIES = [] # Auto-detect if empty - SUPPORTED_EXTENSIONS = ['.png', '.jpg', '.jpeg', '.webp', '.gif'] + SUPPORTED_EXTENSIONS = [".png", ".jpg", ".jpeg", ".webp", ".gif"] PROCESSING_DELAY = 2.0 # Seconds to wait before processing new files - + # Prompt tracking settings - PROMPT_TIMEOUT = 600 # Seconds to keep prompt context active (10 min for long generations) + PROMPT_TIMEOUT = ( + 600 # Seconds to keep prompt context active (10 min for long generations) + ) CLEANUP_INTERVAL = 300 # Seconds between cleanup of expired prompts - + # Database settings AUTO_CLEANUP_MISSING_FILES = True MAX_IMAGE_AGE_DAYS = 365 # Clean up images older than this - + # Web interface settings IMAGES_PER_PAGE = 20 THUMBNAIL_SIZE = 256 ENABLE_SEARCH = True ENABLE_METADATA_VIEW = True - + # Performance settings MAX_CONCURRENT_PROCESSING = 3 METADATA_EXTRACTION_TIMEOUT = 10 # Seconds - + @classmethod def get_config(cls) -> Dict[str, Any]: """Get the complete gallery configuration as a structured dictionary. - + Returns: Dict[str, Any]: A nested dictionary containing all gallery configuration sections: monitoring, tracking, database, web_interface, and performance. Each section contains the relevant configuration parameters as key-value pairs. - + Example: config = GalleryConfig.get_config() monitoring_enabled = config['monitoring']['enabled'] images_per_page = config['web_interface']['images_per_page'] """ return { - 'monitoring': { - 'enabled': cls.MONITORING_ENABLED, - 'directories': cls.MONITORING_DIRECTORIES, - 'extensions': cls.SUPPORTED_EXTENSIONS, - 'processing_delay': cls.PROCESSING_DELAY + "monitoring": { + "enabled": cls.MONITORING_ENABLED, + "directories": cls.MONITORING_DIRECTORIES, + "extensions": cls.SUPPORTED_EXTENSIONS, + "processing_delay": cls.PROCESSING_DELAY, }, - 'tracking': { - 'prompt_timeout': cls.PROMPT_TIMEOUT, - 'cleanup_interval': cls.CLEANUP_INTERVAL + "tracking": { + "prompt_timeout": cls.PROMPT_TIMEOUT, + "cleanup_interval": cls.CLEANUP_INTERVAL, }, - 'database': { - 'auto_cleanup': cls.AUTO_CLEANUP_MISSING_FILES, - 'max_image_age_days': cls.MAX_IMAGE_AGE_DAYS + "database": { + "auto_cleanup": cls.AUTO_CLEANUP_MISSING_FILES, + "max_image_age_days": cls.MAX_IMAGE_AGE_DAYS, }, - 'web_interface': { - 'images_per_page': cls.IMAGES_PER_PAGE, - 'thumbnail_size': cls.THUMBNAIL_SIZE, - 'enable_search': cls.ENABLE_SEARCH, - 'enable_metadata_view': cls.ENABLE_METADATA_VIEW + "web_interface": { + "images_per_page": cls.IMAGES_PER_PAGE, + "thumbnail_size": cls.THUMBNAIL_SIZE, + "enable_search": cls.ENABLE_SEARCH, + "enable_metadata_view": cls.ENABLE_METADATA_VIEW, + }, + "performance": { + "max_concurrent_processing": cls.MAX_CONCURRENT_PROCESSING, + "metadata_extraction_timeout": cls.METADATA_EXTRACTION_TIMEOUT, }, - 'performance': { - 'max_concurrent_processing': cls.MAX_CONCURRENT_PROCESSING, - 'metadata_extraction_timeout': cls.METADATA_EXTRACTION_TIMEOUT - } } - + @classmethod def update_config(cls, new_config: Dict[str, Any]): """Update gallery configuration attributes from a dictionary. - + Takes a nested dictionary with gallery configuration sections and updates the corresponding class attributes. Only updates attributes that are present in the input dictionary, leaving others unchanged. - + Args: new_config (Dict[str, Any]): Nested dictionary containing gallery configuration updates. Should follow the same structure as returned by get_config(). Valid top-level keys are: 'monitoring', 'tracking', 'database', 'web_interface', 'performance'. - + Example: gallery_settings = { 'monitoring': {'enabled': False}, @@ -158,58 +162,58 @@ class GalleryConfig: } GalleryConfig.update_config(gallery_settings) """ - monitoring = new_config.get('monitoring', {}) - if 'enabled' in monitoring: - cls.MONITORING_ENABLED = monitoring['enabled'] - if 'directories' in monitoring: - cls.MONITORING_DIRECTORIES = monitoring['directories'] - if 'extensions' in monitoring: - cls.SUPPORTED_EXTENSIONS = monitoring['extensions'] - if 'processing_delay' in monitoring: - cls.PROCESSING_DELAY = monitoring['processing_delay'] - - tracking = new_config.get('tracking', {}) - if 'prompt_timeout' in tracking: - cls.PROMPT_TIMEOUT = tracking['prompt_timeout'] - if 'cleanup_interval' in tracking: - cls.CLEANUP_INTERVAL = tracking['cleanup_interval'] - - database = new_config.get('database', {}) - if 'auto_cleanup' in database: - cls.AUTO_CLEANUP_MISSING_FILES = database['auto_cleanup'] - if 'max_image_age_days' in database: - cls.MAX_IMAGE_AGE_DAYS = database['max_image_age_days'] - - web_interface = new_config.get('web_interface', {}) - if 'images_per_page' in web_interface: - cls.IMAGES_PER_PAGE = web_interface['images_per_page'] - if 'thumbnail_size' in web_interface: - cls.THUMBNAIL_SIZE = web_interface['thumbnail_size'] - if 'enable_search' in web_interface: - cls.ENABLE_SEARCH = web_interface['enable_search'] - if 'enable_metadata_view' in web_interface: - cls.ENABLE_METADATA_VIEW = web_interface['enable_metadata_view'] - - performance = new_config.get('performance', {}) - if 'max_concurrent_processing' in performance: - cls.MAX_CONCURRENT_PROCESSING = performance['max_concurrent_processing'] - if 'metadata_extraction_timeout' in performance: - cls.METADATA_EXTRACTION_TIMEOUT = performance['metadata_extraction_timeout'] + monitoring = new_config.get("monitoring", {}) + if "enabled" in monitoring: + cls.MONITORING_ENABLED = monitoring["enabled"] + if "directories" in monitoring: + cls.MONITORING_DIRECTORIES = monitoring["directories"] + if "extensions" in monitoring: + cls.SUPPORTED_EXTENSIONS = monitoring["extensions"] + if "processing_delay" in monitoring: + cls.PROCESSING_DELAY = monitoring["processing_delay"] + + tracking = new_config.get("tracking", {}) + if "prompt_timeout" in tracking: + cls.PROMPT_TIMEOUT = tracking["prompt_timeout"] + if "cleanup_interval" in tracking: + cls.CLEANUP_INTERVAL = tracking["cleanup_interval"] + + database = new_config.get("database", {}) + if "auto_cleanup" in database: + cls.AUTO_CLEANUP_MISSING_FILES = database["auto_cleanup"] + if "max_image_age_days" in database: + cls.MAX_IMAGE_AGE_DAYS = database["max_image_age_days"] + + web_interface = new_config.get("web_interface", {}) + if "images_per_page" in web_interface: + cls.IMAGES_PER_PAGE = web_interface["images_per_page"] + if "thumbnail_size" in web_interface: + cls.THUMBNAIL_SIZE = web_interface["thumbnail_size"] + if "enable_search" in web_interface: + cls.ENABLE_SEARCH = web_interface["enable_search"] + if "enable_metadata_view" in web_interface: + cls.ENABLE_METADATA_VIEW = web_interface["enable_metadata_view"] + + performance = new_config.get("performance", {}) + if "max_concurrent_processing" in performance: + cls.MAX_CONCURRENT_PROCESSING = performance["max_concurrent_processing"] + if "metadata_extraction_timeout" in performance: + cls.METADATA_EXTRACTION_TIMEOUT = performance["metadata_extraction_timeout"] class PromptManagerConfig: """Main configuration class for PromptManager core functionality. - + This class manages configuration for database operations, web UI behavior, performance settings, and integrates gallery configuration. It provides methods for loading and saving configuration from/to JSON files. - + The configuration is organized into logical sections: - Database: Settings for SQLite operations and data management - Web UI: User interface behavior and display options - Performance: Optimization and resource management settings - Gallery: Embedded gallery configuration (via GalleryConfig) - + Attributes: DEFAULT_DB_PATH (str): Default path for the SQLite database file ENABLE_DUPLICATE_DETECTION (bool): Enable automatic duplicate detection @@ -221,82 +225,82 @@ class PromptManagerConfig: ENABLE_FUZZY_SEARCH (bool): Enable fuzzy search capabilities AUTO_BACKUP_INTERVAL (int): Hours between automatic database backups """ - + # Database settings DEFAULT_DB_PATH = "prompts.db" ENABLE_DUPLICATE_DETECTION = True ENABLE_AUTO_SAVE = True - + # Web UI settings RESULT_TIMEOUT = 5 # Seconds to auto-hide results in ComfyUI node SHOW_TEST_BUTTON = False # Show API test button in node UI - WEBUI_DISPLAY_MODE = 'newtab' # 'popup' or 'newtab' - + WEBUI_DISPLAY_MODE = "newtab" # 'popup' or 'newtab' + # Performance settings MAX_SEARCH_RESULTS = 100 ENABLE_FUZZY_SEARCH = False # Requires fuzzywuzzy AUTO_BACKUP_INTERVAL = 24 # Hours - + @classmethod def get_config(cls) -> Dict[str, Any]: """Get the complete PromptManager configuration as a structured dictionary. - + Returns: Dict[str, Any]: A nested dictionary containing all configuration sections: - database: Database-related settings - web_ui: Web interface configuration - performance: Performance and optimization settings - gallery: Complete gallery configuration (from GalleryConfig) - + Example: config = PromptManagerConfig.get_config() db_path = config['database']['default_path'] max_results = config['performance']['max_search_results'] """ return { - 'database': { - 'default_path': cls.DEFAULT_DB_PATH, - 'enable_duplicate_detection': cls.ENABLE_DUPLICATE_DETECTION, - 'enable_auto_save': cls.ENABLE_AUTO_SAVE + "database": { + "default_path": cls.DEFAULT_DB_PATH, + "enable_duplicate_detection": cls.ENABLE_DUPLICATE_DETECTION, + "enable_auto_save": cls.ENABLE_AUTO_SAVE, }, - 'web_ui': { - 'result_timeout': cls.RESULT_TIMEOUT, - 'show_test_button': cls.SHOW_TEST_BUTTON, - 'webui_display_mode': cls.WEBUI_DISPLAY_MODE + "web_ui": { + "result_timeout": cls.RESULT_TIMEOUT, + "show_test_button": cls.SHOW_TEST_BUTTON, + "webui_display_mode": cls.WEBUI_DISPLAY_MODE, }, - 'performance': { - 'max_search_results': cls.MAX_SEARCH_RESULTS, - 'enable_fuzzy_search': cls.ENABLE_FUZZY_SEARCH, - 'auto_backup_interval': cls.AUTO_BACKUP_INTERVAL + "performance": { + "max_search_results": cls.MAX_SEARCH_RESULTS, + "enable_fuzzy_search": cls.ENABLE_FUZZY_SEARCH, + "auto_backup_interval": cls.AUTO_BACKUP_INTERVAL, }, - 'gallery': GalleryConfig.get_config() + "gallery": GalleryConfig.get_config(), } - + @classmethod def load_from_file(cls, config_path: str): """Load configuration settings from a JSON file. - + Reads configuration from the specified JSON file and updates the current configuration attributes. If the file doesn't exist or contains invalid JSON, logs an appropriate message and continues with default values. - + Args: config_path (str): Path to the JSON configuration file to load. Can be relative or absolute path. - + Raises: The method handles all exceptions internally and logs errors rather than propagating them, ensuring the system continues with defaults. - + Example: PromptManagerConfig.load_from_file('custom_config.json') PromptManagerConfig.load_from_file('/path/to/config.json') """ import json - + if os.path.exists(config_path): try: - with open(config_path, 'r') as f: + with open(config_path, "r") as f: config = json.load(f) cls.update_config(config) config_logger.info(f"Loaded configuration from {config_path}") @@ -304,52 +308,52 @@ class PromptManagerConfig: config_logger.error(f"Error loading config from {config_path}: {e}") else: config_logger.info(f"Config file not found: {config_path}, using defaults") - + @classmethod def save_to_file(cls, config_path: str): """Save the current configuration to a JSON file. - + Serializes the complete configuration (including gallery settings) to a JSON file. Creates the directory structure if it doesn't exist. - + Args: config_path (str): Path where the JSON configuration file should be saved. Parent directories will be created if they don't exist. - + Raises: The method handles all exceptions internally and logs errors rather than propagating them. - + Example: PromptManagerConfig.save_to_file('backup_config.json') PromptManagerConfig.save_to_file('/etc/comfyui/prompt_manager.json') """ import json - + try: config = cls.get_config() os.makedirs(os.path.dirname(config_path), exist_ok=True) - - with open(config_path, 'w') as f: + + with open(config_path, "w") as f: json.dump(config, f, indent=2) - + config_logger.info(f"Saved configuration to {config_path}") except Exception as e: config_logger.error(f"Error saving config to {config_path}: {e}") - + @classmethod def update_config(cls, new_config: Dict[str, Any]): """Update configuration attributes from a dictionary. - + Takes a nested dictionary with configuration sections and updates the corresponding class attributes. Only updates attributes that are present in the input dictionary, leaving others unchanged. - + Args: new_config (Dict[str, Any]): Nested dictionary containing configuration updates. Should follow the same structure as returned by get_config(). Valid top-level keys are: 'database', 'web_ui', 'performance', 'gallery'. - + Example: new_settings = { 'database': {'default_path': 'custom.db'}, @@ -357,39 +361,39 @@ class PromptManagerConfig: } PromptManagerConfig.update_config(new_settings) """ - database = new_config.get('database', {}) - if 'default_path' in database: - cls.DEFAULT_DB_PATH = database['default_path'] - if 'enable_duplicate_detection' in database: - cls.ENABLE_DUPLICATE_DETECTION = database['enable_duplicate_detection'] - if 'enable_auto_save' in database: - cls.ENABLE_AUTO_SAVE = database['enable_auto_save'] - - web_ui = new_config.get('web_ui', {}) - if 'result_timeout' in web_ui: - cls.RESULT_TIMEOUT = web_ui['result_timeout'] - if 'show_test_button' in web_ui: - cls.SHOW_TEST_BUTTON = web_ui['show_test_button'] - if 'webui_display_mode' in web_ui: - cls.WEBUI_DISPLAY_MODE = web_ui['webui_display_mode'] - - performance = new_config.get('performance', {}) - if 'max_search_results' in performance: - cls.MAX_SEARCH_RESULTS = performance['max_search_results'] - if 'enable_fuzzy_search' in performance: - cls.ENABLE_FUZZY_SEARCH = performance['enable_fuzzy_search'] - if 'auto_backup_interval' in performance: - cls.AUTO_BACKUP_INTERVAL = performance['auto_backup_interval'] - + database = new_config.get("database", {}) + if "default_path" in database: + cls.DEFAULT_DB_PATH = database["default_path"] + if "enable_duplicate_detection" in database: + cls.ENABLE_DUPLICATE_DETECTION = database["enable_duplicate_detection"] + if "enable_auto_save" in database: + cls.ENABLE_AUTO_SAVE = database["enable_auto_save"] + + web_ui = new_config.get("web_ui", {}) + if "result_timeout" in web_ui: + cls.RESULT_TIMEOUT = web_ui["result_timeout"] + if "show_test_button" in web_ui: + cls.SHOW_TEST_BUTTON = web_ui["show_test_button"] + if "webui_display_mode" in web_ui: + cls.WEBUI_DISPLAY_MODE = web_ui["webui_display_mode"] + + performance = new_config.get("performance", {}) + if "max_search_results" in performance: + cls.MAX_SEARCH_RESULTS = performance["max_search_results"] + if "enable_fuzzy_search" in performance: + cls.ENABLE_FUZZY_SEARCH = performance["enable_fuzzy_search"] + if "auto_backup_interval" in performance: + cls.AUTO_BACKUP_INTERVAL = performance["auto_backup_interval"] + # Update gallery config - if 'gallery' in new_config: - GalleryConfig.update_config(new_config['gallery']) + if "gallery" in new_config: + GalleryConfig.update_config(new_config["gallery"]) # Load configuration on import try: config_dir = os.path.dirname(os.path.dirname(os.path.abspath(__file__))) - config_file = os.path.join(config_dir, 'config.json') + config_file = os.path.join(config_dir, "config.json") PromptManagerConfig.load_from_file(config_file) except Exception as e: - config_logger.error(f"Error during config initialization: {e}") \ No newline at end of file + config_logger.error(f"Error during config initialization: {e}") diff --git a/restart_gallery.py b/restart_gallery.py index f734263..7403594 100644 --- a/restart_gallery.py +++ b/restart_gallery.py @@ -3,7 +3,7 @@ ComfyUI_PromptManager Gallery System Restart Script. This script reinitializes and restarts the automatic image gallery monitoring system -for ComfyUI_PromptManager. Use this script after code modifications or when the +for ComfyUI_PromptManager. Use this script after code modifications or when the gallery system needs to be reloaded. The script performs the following operations: @@ -41,23 +41,24 @@ try: monitor = get_image_monitor(db, tracker) print("[SUCCESS] Components initialized successfully") - + # Test image monitoring directories status = monitor.get_status() print(f"[INFO] Monitoring status: {status}") - + # Start monitoring monitor.start_monitoring() print("[SUCCESS] Image monitoring restarted") - + print("\n[READY] Gallery system restart completed!") print("Generate some images now to test the automatic linking.") - + except ImportError as e: print(f"[ERROR] Import error: {e}") print("Make sure you're running this from the PromptManager directory") - + except Exception as e: print(f"[ERROR] Error: {e}") import traceback - traceback.print_exc() \ No newline at end of file + + traceback.print_exc() diff --git a/tests/__init__.py b/tests/__init__.py index 1d61aae..c1e6fb7 100644 --- a/tests/__init__.py +++ b/tests/__init__.py @@ -3,4 +3,4 @@ Test suite for ComfyUI_PromptManager. This package contains unit tests and integration tests for the PromptManager custom nodes, database operations, utility functions, and web interface components. -""" \ No newline at end of file +""" diff --git a/tests/test_api.py b/tests/test_api.py index 07664cb..54ba8f7 100644 --- a/tests/test_api.py +++ b/tests/test_api.py @@ -38,7 +38,11 @@ class APITestCase(AioHTTPTestCase): return app async def tearDownAsync(self): - for path in (self._temp_db.name, self._temp_db.name + "-wal", self._temp_db.name + "-shm"): + for path in ( + self._temp_db.name, + self._temp_db.name + "-wal", + self._temp_db.name + "-shm", + ): if os.path.exists(path): os.unlink(path) @@ -86,7 +90,9 @@ class TestRecentPrompts(APITestCase): async def test_recent_pagination(self): for i in range(15): self._save_prompt(f"Prompt {i:02d}") - resp = await self.client.request("GET", "/prompt_manager/recent?limit=5&offset=0") + resp = await self.client.request( + "GET", "/prompt_manager/recent?limit=5&offset=0" + ) data = await resp.json() self.assertEqual(len(data["results"]), 5) self.assertEqual(data["pagination"]["total"], 15) @@ -95,7 +101,9 @@ class TestRecentPrompts(APITestCase): async def test_recent_last_page(self): for i in range(7): self._save_prompt(f"Prompt {i}") - resp = await self.client.request("GET", "/prompt_manager/recent?limit=5&offset=5") + resp = await self.client.request( + "GET", "/prompt_manager/recent?limit=5&offset=5" + ) data = await resp.json() self.assertEqual(len(data["results"]), 2) self.assertFalse(data["pagination"]["has_more"]) @@ -116,7 +124,9 @@ class TestSearch(APITestCase): async def test_search_by_category(self): self._save_prompt("Nature scene", category="nature") self._save_prompt("Urban view", category="urban") - resp = await self.client.request("GET", "/prompt_manager/search?category=nature") + resp = await self.client.request( + "GET", "/prompt_manager/search?category=nature" + ) data = await resp.json() self.assertEqual(len(data["results"]), 1) @@ -128,7 +138,9 @@ class TestSearch(APITestCase): self.assertEqual(len(data["results"]), 1) async def test_search_empty_result(self): - resp = await self.client.request("GET", "/prompt_manager/search?text=nonexistent_xyz") + resp = await self.client.request( + "GET", "/prompt_manager/search?text=nonexistent_xyz" + ) data = await resp.json() self.assertTrue(data["success"]) self.assertEqual(len(data["results"]), 0) @@ -140,7 +152,11 @@ class TestSaveAndDelete(APITestCase): resp = await self.client.request( "POST", "/prompt_manager/save", - json={"text": "New prompt via API", "category": "test", "tags": ["api", "test"]}, + json={ + "text": "New prompt via API", + "category": "test", + "tags": ["api", "test"], + }, ) self.assertEqual(resp.status, 200) data = await resp.json() @@ -148,7 +164,9 @@ class TestSaveAndDelete(APITestCase): self.assertIn("prompt_id", data) async def test_save_empty_text(self): - resp = await self.client.request("POST", "/prompt_manager/save", json={"text": ""}) + resp = await self.client.request( + "POST", "/prompt_manager/save", json={"text": ""} + ) self.assertEqual(resp.status, 400) data = await resp.json() self.assertFalse(data["success"]) @@ -307,7 +325,9 @@ class TestResponseEnvelope(APITestCase): resp = await self.client.request(method, path) data = await resp.json() self.assertIn("success", data, f"Missing 'success' key in {method} {path}") - self.assertTrue(data["success"], f"Expected success=True for {method} {path}") + self.assertTrue( + data["success"], f"Expected success=True for {method} {path}" + ) async def test_error_responses_have_success_false(self): resp = await self.client.request("DELETE", "/prompt_manager/delete/99999") diff --git a/tests/test_basic.py b/tests/test_basic.py index 89c3159..ca70b02 100644 --- a/tests/test_basic.py +++ b/tests/test_basic.py @@ -9,6 +9,7 @@ from unittest.mock import Mock, patch # Add the parent directory to the path for imports import sys + sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__)))) from database.operations import PromptDatabase @@ -18,61 +19,61 @@ from utils.validators import validate_prompt_text, validate_rating, validate_tag class TestBasicFunctionality(unittest.TestCase): """Test basic functionality of PromptManager components.""" - + def setUp(self): """Set up test fixtures.""" # Use a temporary database for testing - self.temp_db = tempfile.NamedTemporaryFile(delete=False, suffix='.db') + self.temp_db = tempfile.NamedTemporaryFile(delete=False, suffix=".db") self.temp_db.close() self.db = PromptDatabase(self.temp_db.name) - + def tearDown(self): """Clean up test fixtures.""" # Remove temporary database if os.path.exists(self.temp_db.name): os.unlink(self.temp_db.name) - + def test_prompt_hash_generation(self): """Test prompt hash generation.""" text1 = "A beautiful landscape with mountains" text2 = "A BEAUTIFUL LANDSCAPE WITH MOUNTAINS" # Different case text3 = " A beautiful landscape with mountains " # Extra whitespace text4 = "A different prompt" - + hash1 = generate_prompt_hash(text1) hash2 = generate_prompt_hash(text2) hash3 = generate_prompt_hash(text3) hash4 = generate_prompt_hash(text4) - + # Same content should produce same hash (case and whitespace insensitive) self.assertEqual(hash1, hash2) self.assertEqual(hash1, hash3) - + # Different content should produce different hash self.assertNotEqual(hash1, hash4) - + # Hash should be 64 characters (SHA256 hex) self.assertEqual(len(hash1), 64) - + def test_prompt_validation(self): """Test prompt text validation.""" # Valid prompts self.assertTrue(validate_prompt_text("Valid prompt text")) self.assertTrue(validate_prompt_text("Another valid prompt")) - + # Invalid prompts with self.assertRaises(ValueError): validate_prompt_text("") # Empty - + with self.assertRaises(ValueError): validate_prompt_text(" ") # Whitespace only - + with self.assertRaises(ValueError): validate_prompt_text("x" * 10001) # Too long - + with self.assertRaises(ValueError): validate_prompt_text(123) # Not a string - + def test_rating_validation(self): """Test rating validation.""" # Valid ratings @@ -80,17 +81,17 @@ class TestBasicFunctionality(unittest.TestCase): self.assertTrue(validate_rating(1)) self.assertTrue(validate_rating(3)) self.assertTrue(validate_rating(5)) - + # Invalid ratings with self.assertRaises(ValueError): validate_rating(0) # Too low - + with self.assertRaises(ValueError): validate_rating(6) # Too high - + with self.assertRaises(ValueError): validate_rating("3") # Not an integer - + def test_tags_validation(self): """Test tags validation.""" # Valid tags @@ -98,17 +99,17 @@ class TestBasicFunctionality(unittest.TestCase): self.assertTrue(validate_tags([])) self.assertTrue(validate_tags(["tag1", "tag2"])) self.assertTrue(validate_tags("tag1, tag2, tag3")) - + # Invalid tags with self.assertRaises(ValueError): validate_tags([""]) # Empty tag - + with self.assertRaises(ValueError): validate_tags(["x" * 51]) # Tag too long - + with self.assertRaises(ValueError): validate_tags(["tag"] * 21) # Too many tags - + def test_database_save_and_retrieve(self): """Test basic database save and retrieve operations.""" # Save a prompt @@ -118,38 +119,36 @@ class TestBasicFunctionality(unittest.TestCase): tags=["test", "database"], rating=4, notes="Test notes", - prompt_hash=generate_prompt_hash("Test prompt for database") + prompt_hash=generate_prompt_hash("Test prompt for database"), ) - + self.assertIsInstance(prompt_id, int) self.assertGreater(prompt_id, 0) - + # Retrieve the prompt retrieved = self.db.get_prompt_by_id(prompt_id) self.assertIsNotNone(retrieved) - self.assertEqual(retrieved['text'], "Test prompt for database") - self.assertEqual(retrieved['category'], "test") - self.assertEqual(retrieved['tags'], ["test", "database"]) - self.assertEqual(retrieved['rating'], 4) - self.assertEqual(retrieved['notes'], "Test notes") - + self.assertEqual(retrieved["text"], "Test prompt for database") + self.assertEqual(retrieved["category"], "test") + self.assertEqual(retrieved["tags"], ["test", "database"]) + self.assertEqual(retrieved["rating"], 4) + self.assertEqual(retrieved["notes"], "Test notes") + def test_duplicate_detection(self): """Test duplicate prompt detection.""" text = "Duplicate test prompt" hash_val = generate_prompt_hash(text) - + # Save original prompt prompt_id1 = self.db.save_prompt( - text=text, - category="test", - prompt_hash=hash_val + text=text, category="test", prompt_hash=hash_val ) - + # Try to save the same prompt again existing = self.db.get_prompt_by_hash(hash_val) self.assertIsNotNone(existing) - self.assertEqual(existing['id'], prompt_id1) - + self.assertEqual(existing["id"], prompt_id1) + def test_search_functionality(self): """Test prompt search functionality.""" # Save multiple prompts @@ -158,67 +157,67 @@ class TestBasicFunctionality(unittest.TestCase): "text": "Beautiful landscape with mountains", "category": "landscape", "tags": ["nature", "mountains"], - "rating": 5 + "rating": 5, }, { "text": "Portrait of a woman", "category": "portrait", "tags": ["people", "woman"], - "rating": 4 + "rating": 4, }, { "text": "Abstract art piece", "category": "abstract", "tags": ["art", "abstract"], - "rating": 3 - } + "rating": 3, + }, ] - + for prompt_data in prompts: self.db.save_prompt( text=prompt_data["text"], category=prompt_data["category"], tags=prompt_data["tags"], rating=prompt_data["rating"], - prompt_hash=generate_prompt_hash(prompt_data["text"]) + prompt_hash=generate_prompt_hash(prompt_data["text"]), ) - + # Test text search results = self.db.search_prompts(text="landscape") self.assertEqual(len(results), 1) - self.assertIn("landscape", results[0]['text']) - + self.assertIn("landscape", results[0]["text"]) + # Test category search results = self.db.search_prompts(category="portrait") self.assertEqual(len(results), 1) - self.assertEqual(results[0]['category'], "portrait") - + self.assertEqual(results[0]["category"], "portrait") + # Test rating search results = self.db.search_prompts(rating_min=4) self.assertEqual(len(results), 2) # Rating 4 and 5 - + # Test tag search results = self.db.search_prompts(tags=["nature"]) self.assertEqual(len(results), 1) - self.assertIn("nature", results[0]['tags']) + self.assertIn("nature", results[0]["tags"]) class TestNodeIntegration(unittest.TestCase): """Test the actual node integration (mocked).""" - + def setUp(self): """Set up test fixtures.""" # Use a temporary database for testing - self.temp_db = tempfile.NamedTemporaryFile(delete=False, suffix='.db') + self.temp_db = tempfile.NamedTemporaryFile(delete=False, suffix=".db") self.temp_db.close() - + def tearDown(self): """Clean up test fixtures.""" # Remove temporary database if os.path.exists(self.temp_db.name): os.unlink(self.temp_db.name) - - @patch('prompt_manager_base.PromptDatabase') + + @patch("prompt_manager_base.PromptDatabase") def test_node_encode_function(self, mock_db_class): """Test the node's encode function with mocked dependencies.""" # Mock the database @@ -238,11 +237,7 @@ class TestNodeIntegration(unittest.TestCase): node = PromptManager() # Test encoding - result = node.encode_prompt( - clip=mock_clip, - text="Test prompt", - search_text="" - ) + result = node.encode_prompt(clip=mock_clip, text="Test prompt", search_text="") # Verify CLIP was called correctly mock_clip.tokenize.assert_called_once_with("Test prompt") @@ -250,20 +245,20 @@ class TestNodeIntegration(unittest.TestCase): # Verify result tuple (conditioning, prompt_text) self.assertEqual(result[1], "Test prompt") - + def test_node_input_types(self): """Test the node's input type definitions.""" from prompt_manager import PromptManager - + input_types = PromptManager.INPUT_TYPES() - + # Check required inputs self.assertIn("text", input_types["required"]) self.assertIn("clip", input_types["required"]) - + # Check optional inputs self.assertIn("search_text", input_types["optional"]) -if __name__ == '__main__': - unittest.main() \ No newline at end of file +if __name__ == "__main__": + unittest.main() diff --git a/tests/test_database.py b/tests/test_database.py index 666cc15..80087cd 100644 --- a/tests/test_database.py +++ b/tests/test_database.py @@ -34,7 +34,9 @@ class DatabaseTestCase(unittest.TestCase): if os.path.exists(f): os.unlink(f) - def _save(self, text="Test prompt", category=None, tags=None, rating=None, notes=None): + def _save( + self, text="Test prompt", category=None, tags=None, rating=None, notes=None + ): """Helper to save a prompt and return its ID.""" return self.db.save_prompt( text=text, @@ -50,7 +52,9 @@ class TestPromptCRUD(DatabaseTestCase): """Test basic create, read, update, delete operations.""" def test_save_and_retrieve(self): - pid = self._save("A beautiful sunset", category="nature", tags=["sunset", "sky"], rating=5) + pid = self._save( + "A beautiful sunset", category="nature", tags=["sunset", "sky"], rating=5 + ) prompt = self.db.get_prompt_by_id(pid) self.assertEqual(prompt["text"], "A beautiful sunset") self.assertEqual(prompt["category"], "nature") @@ -84,7 +88,9 @@ class TestPromptCRUD(DatabaseTestCase): def test_update_metadata(self): pid = self._save("Updatable prompt", category="old", tags=["old_tag"], rating=2) - self.db.update_prompt_metadata(pid, category="new", tags=["new_tag"], rating=5, notes="updated") + self.db.update_prompt_metadata( + pid, category="new", tags=["new_tag"], rating=5, notes="updated" + ) prompt = self.db.get_prompt_by_id(pid) self.assertEqual(prompt["category"], "new") self.assertIn("new_tag", prompt["tags"]) @@ -204,10 +210,27 @@ class TestSearch(DatabaseTestCase): def setUp(self): super().setUp() - self._save("Beautiful mountain landscape", category="nature", tags=["mountain", "landscape"], rating=5) - self._save("City skyline at night", category="urban", tags=["city", "night"], rating=4) - self._save("Portrait of an artist", category="portrait", tags=["person", "art"], rating=3) - self._save("Abstract geometric shapes", category="abstract", tags=["art", "geometric"], rating=2) + self._save( + "Beautiful mountain landscape", + category="nature", + tags=["mountain", "landscape"], + rating=5, + ) + self._save( + "City skyline at night", category="urban", tags=["city", "night"], rating=4 + ) + self._save( + "Portrait of an artist", + category="portrait", + tags=["person", "art"], + rating=3, + ) + self._save( + "Abstract geometric shapes", + category="abstract", + tags=["art", "geometric"], + rating=2, + ) def test_search_by_text(self): results = self.db.search_prompts(text="mountain") @@ -428,7 +451,10 @@ class TestEdgeCases(DatabaseTestCase): self.assertEqual(prompt["text"], long_text.strip()) def test_tag_with_special_characters(self): - pid = self._save("Special tags", tags=["tag-with-dash", "tag_with_underscore", "tag.with.dots"]) + pid = self._save( + "Special tags", + tags=["tag-with-dash", "tag_with_underscore", "tag.with.dots"], + ) prompt = self.db.get_prompt_by_id(pid) self.assertEqual(len(prompt["tags"]), 3) diff --git a/utils/__init__.py b/utils/__init__.py index e02dbf2..71f096c 100644 --- a/utils/__init__.py +++ b/utils/__init__.py @@ -8,4 +8,9 @@ image monitoring, metadata extraction, logging, and system diagnostics. 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 +__all__ = [ + "generate_prompt_hash", + "validate_prompt_text", + "validate_rating", + "validate_tags", +] diff --git a/utils/comfyui_integration.py b/utils/comfyui_integration.py index 5aa334a..2468161 100644 --- a/utils/comfyui_integration.py +++ b/utils/comfyui_integration.py @@ -19,7 +19,7 @@ The integration works by: Typical usage: from utils.comfyui_integration import get_comfyui_integration - + integration = get_comfyui_integration() integration.register_prompt(node_id, prompt_text, metadata) # Generated images will now include this prompt in their metadata @@ -41,34 +41,35 @@ try: except ImportError: import sys import os + sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__)))) from utils.logging_config import get_logger class ComfyUIMetadataIntegration: """Integrates PromptManager with ComfyUI's standard metadata system. - + This singleton class manages the integration between PromptManager custom nodes and ComfyUI's standard metadata system. It ensures that prompts generated by PromptManager appear in the standard ComfyUI metadata format for compatibility with third-party tools and parsers. - + The integration uses thread-local storage combined with global tracking to handle prompt context across different execution threads, which is necessary because ComfyUI executions may span multiple threads. - + Key responsibilities: - Register prompts from PromptManager nodes during execution - Patch SaveImage to include PromptManager prompts in metadata - Manage prompt lifecycle and cleanup """ - + _instance = None _lock = threading.Lock() - + def __new__(cls): """Ensure singleton pattern with thread safety. - + Returns: The single instance of ComfyUIMetadataIntegration """ @@ -77,109 +78,118 @@ class ComfyUIMetadataIntegration: if cls._instance is None: cls._instance = super().__new__(cls) return cls._instance - + def __init__(self): """Initialize the ComfyUI integration system. - + Sets up prompt tracking data structures and attempts to patch the SaveImage node. Uses _initialized flag to prevent duplicate initialization. """ - if hasattr(self, '_initialized'): + if hasattr(self, "_initialized"): return - - self.logger = get_logger('prompt_manager.comfyui_integration') + + self.logger = get_logger("prompt_manager.comfyui_integration") self._current_prompts = {} self._thread_local = threading.local() self._saveimage_patched = False self._initialized = True - + # Try to patch SaveImage node on initialization self._patch_saveimage_node() - + def register_prompt(self, node_id: str, prompt_text: str, metadata: Dict[str, Any]): """ Register a prompt from PromptManager for inclusion in ComfyUI metadata. - + This method is called by PromptManager nodes during execution to register their prompt text for later inclusion in image metadata. The prompt is stored in both thread-local and global storage to handle cross-thread access scenarios. - + Args: node_id: Unique identifier for this prompt (typically the node ID) prompt_text: The actual prompt text that was encoded by PromptManager metadata: Additional metadata from PromptManager (category, tags, etc.) """ thread_id = threading.current_thread().ident - + # Store in thread-local storage - if not hasattr(self._thread_local, 'prompts'): + if not hasattr(self._thread_local, "prompts"): self._thread_local.prompts = {} - + self._thread_local.prompts[node_id] = { - 'text': prompt_text, - 'metadata': metadata, - 'timestamp': time.time(), - 'thread_id': thread_id + "text": prompt_text, + "metadata": metadata, + "timestamp": time.time(), + "thread_id": thread_id, } - + # Also store globally for cross-thread access with self._lock: self._current_prompts[f"{thread_id}_{node_id}"] = { - 'text': prompt_text, - 'metadata': metadata, - 'timestamp': time.time(), - 'thread_id': thread_id + "text": prompt_text, + "metadata": metadata, + "timestamp": time.time(), + "thread_id": thread_id, } - - self.logger.debug(f"Registered prompt for node {node_id}: {prompt_text[:50]}...") - + + self.logger.debug( + f"Registered prompt for node {node_id}: {prompt_text[:50]}..." + ) + def get_current_prompt_text(self, node_id: str = None) -> Optional[str]: """ Get the current prompt text for metadata inclusion. - + Retrieves the most appropriate prompt text for inclusion in image metadata. Uses a fallback strategy: thread-local storage first, then global storage, with recency checks to avoid stale prompts. - + Args: node_id: Optional specific node ID to retrieve prompt for - + Returns: The prompt text string if available, None if no suitable prompt found """ # First try thread-local storage - if hasattr(self._thread_local, 'prompts'): + if hasattr(self._thread_local, "prompts"): if node_id and node_id in self._thread_local.prompts: - return self._thread_local.prompts[node_id]['text'] + return self._thread_local.prompts[node_id]["text"] elif self._thread_local.prompts: # Return the most recent prompt from this thread - latest = max(self._thread_local.prompts.values(), key=lambda x: x['timestamp']) - return latest['text'] - + latest = max( + self._thread_local.prompts.values(), key=lambda x: x["timestamp"] + ) + return latest["text"] + # Fallback to global storage thread_id = threading.current_thread().ident with self._lock: # Look for prompts from current thread - thread_prompts = {k: v for k, v in self._current_prompts.items() - if v['thread_id'] == thread_id} - + thread_prompts = { + k: v + for k, v in self._current_prompts.items() + if v["thread_id"] == thread_id + } + if thread_prompts: - latest = max(thread_prompts.values(), key=lambda x: x['timestamp']) - return latest['text'] - + latest = max(thread_prompts.values(), key=lambda x: x["timestamp"]) + return latest["text"] + # Last resort: return the most recent prompt from any thread if self._current_prompts: - latest = max(self._current_prompts.values(), key=lambda x: x['timestamp']) + latest = max( + self._current_prompts.values(), key=lambda x: x["timestamp"] + ) # Only return if it's recent (within last 5 minutes) - if time.time() - latest['timestamp'] < 300: - return latest['text'] - + if time.time() - latest["timestamp"] < 300: + return latest["text"] + return None - + def _patch_saveimage_node(self): """ Patch ComfyUI's SaveImage node to include PromptManager prompts in metadata. - + This method modifies ComfyUI's SaveImage.save_images method to automatically include PromptManager prompts in the image metadata. The patching: @@ -195,80 +205,97 @@ class ComfyUIMetadataIntegration: """ try: import nodes - - if not hasattr(nodes, 'SaveImage'): + + if not hasattr(nodes, "SaveImage"): self.logger.warning("SaveImage node not found in ComfyUI nodes") return - + # Store original save_images method original_save_images = nodes.SaveImage.save_images integration = self # Capture self reference - - def patched_save_images(self_node, images, filename_prefix="ComfyUI", prompt=None, extra_pnginfo=None): + + def patched_save_images( + self_node, + images, + filename_prefix="ComfyUI", + prompt=None, + extra_pnginfo=None, + ): """Patched save_images method that includes PromptManager prompts.""" - + # Get current prompt text from PromptManager current_prompt_text = integration.get_current_prompt_text() - + if current_prompt_text: - integration.logger.debug(f"Including PromptManager prompt in SaveImage metadata: {current_prompt_text[:50]}...") - + integration.logger.debug( + f"Including PromptManager prompt in SaveImage metadata: {current_prompt_text[:50]}..." + ) + # If no prompt provided, create one with our text if prompt is None: prompt = {} - + # Ensure prompt has the standard structure ComfyUI expects if not isinstance(prompt, dict): prompt = {} - + # Find PromptManager nodes and ensure prompt text is captured # NOTE: We do NOT change class_type anymore - that was corrupting saved workflows # when users reload them. PromptManager stays as PromptManager. for node_id, node_data in prompt.items(): if isinstance(node_data, dict): - class_type = node_data.get('class_type', '') - if 'promptmanager' in class_type.lower(): + class_type = node_data.get("class_type", "") + if "promptmanager" in class_type.lower(): # Ensure the text input reflects the actual prompt used - if 'inputs' not in node_data: - node_data['inputs'] = {} - node_data['inputs']['text'] = current_prompt_text - integration.logger.debug(f"Updated PromptManager node {node_id} with prompt text") - + if "inputs" not in node_data: + node_data["inputs"] = {} + node_data["inputs"]["text"] = current_prompt_text + integration.logger.debug( + f"Updated PromptManager node {node_id} with prompt text" + ) + # Call original method with potentially modified prompt - return original_save_images(self_node, images, filename_prefix, prompt, extra_pnginfo) - + return original_save_images( + self_node, images, filename_prefix, prompt, extra_pnginfo + ) + # Apply the patch nodes.SaveImage.save_images = patched_save_images self._saveimage_patched = True - self.logger.info("Successfully patched SaveImage node for PromptManager integration") - + self.logger.info( + "Successfully patched SaveImage node for PromptManager integration" + ) + except Exception as e: self.logger.error(f"Failed to patch SaveImage node: {e}") - self.logger.warning("PromptManager prompts may not appear in standard ComfyUI metadata") - + self.logger.warning( + "PromptManager prompts may not appear in standard ComfyUI metadata" + ) + def cleanup_old_prompts(self, max_age_seconds: int = 600): """ Clean up old prompt registrations. - + Removes prompt registrations that are older than the specified age to prevent memory leaks and ensure that only recent, relevant prompts are used for metadata inclusion. - + Args: max_age_seconds: Maximum age in seconds before a prompt registration is considered stale and removed (default: 600 seconds/10 minutes) """ current_time = time.time() - + with self._lock: old_keys = [ - key for key, prompt_data in self._current_prompts.items() - if current_time - prompt_data['timestamp'] > max_age_seconds + key + for key, prompt_data in self._current_prompts.items() + if current_time - prompt_data["timestamp"] > max_age_seconds ] - + for key in old_keys: del self._current_prompts[key] - + if old_keys: self.logger.debug(f"Cleaned up {len(old_keys)} old prompt registrations") @@ -276,13 +303,14 @@ class ComfyUIMetadataIntegration: # Global instance _integration_instance = None + def get_comfyui_integration() -> ComfyUIMetadataIntegration: """Get the global ComfyUI integration instance. - + Returns: The singleton ComfyUIMetadataIntegration instance, creating it if necessary """ global _integration_instance if _integration_instance is None: _integration_instance = ComfyUIMetadataIntegration() - return _integration_instance \ No newline at end of file + return _integration_instance diff --git a/utils/diagnostics.py b/utils/diagnostics.py index d6bccb4..60bb7c6 100644 --- a/utils/diagnostics.py +++ b/utils/diagnostics.py @@ -15,10 +15,10 @@ The diagnostic system checks: Typical usage: from utils.diagnostics import GalleryDiagnostics, run_diagnostics - + # Run full diagnostic suite results = run_diagnostics() - + # Or create custom diagnostic instance diagnostics = GalleryDiagnostics("custom_db.db") database_status = diagnostics.check_database() @@ -40,39 +40,39 @@ from .logging_config import get_logger class GalleryDiagnostics: """Diagnostics for the gallery system. - + This class provides comprehensive diagnostic capabilities for the PromptManager gallery system. It systematically checks all components and dependencies to identify potential issues and provide actionable feedback. - + The diagnostics cover: - Database structure and connectivity - Image tracking table integrity - File system access and permissions - ComfyUI integration points - Python dependency availability - + Each diagnostic method returns a standardized result dictionary with: - status: 'ok', 'warning', or 'error' - message: Descriptive message (for warnings/errors) - Additional data specific to the diagnostic """ - + def __init__(self, db_path: str = "prompts.db"): """Initialize the diagnostics system. - + Args: db_path: Path to the SQLite database file to diagnose """ self.db_path = db_path - self.logger = get_logger('prompt_manager.diagnostics') - + self.logger = get_logger("prompt_manager.diagnostics") + def run_full_diagnostic(self) -> Dict[str, Any]: """Run a complete diagnostic check. - + Executes all diagnostic checks in sequence and provides a comprehensive report of system status. Logs detailed information during the process. - + Returns: Dictionary mapping diagnostic categories to their results: - database: Database connectivity and structure check @@ -81,37 +81,37 @@ class GalleryDiagnostics: - comfyui_output: ComfyUI output directory detection - dependencies: Python dependency availability check """ - self.logger.info("\n" + "="*60) + self.logger.info("\n" + "=" * 60) self.logger.info("[DIAG] PROMPTMANAGER GALLERY DIAGNOSTICS") - self.logger.info("="*60) - + self.logger.info("=" * 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() + "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(), } - - self.logger.info("\n" + "="*60) + + self.logger.info("\n" + "=" * 60) self.logger.info("[SUMMARY] DIAGNOSTIC SUMMARY") - self.logger.info("="*60) - + self.logger.info("=" * 60) + for category, result in results.items(): - status = "[PASS] PASS" if result['status'] == 'ok' else "[FAIL] FAIL" + status = "[PASS] PASS" if result["status"] == "ok" else "[FAIL] FAIL" self.logger.info(f"{category.upper():<20} {status}") - if result['status'] != 'ok': + if result["status"] != "ok": self.logger.warning(f" Issue: {result['message']}") - - self.logger.info("\n" + "="*60) + + self.logger.info("\n" + "=" * 60) return results - + def check_database(self) -> Dict[str, Any]: """Check database connection and structure. - + Verifies that the database file exists, is accessible, and contains the expected prompt table structure. - + Returns: Dictionary containing: - status: 'ok' or 'error' @@ -120,49 +120,46 @@ class GalleryDiagnostics: - has_images_table: Whether generated_images table exists (if successful) """ self.logger.info("\n[DB] Checking Database...") - + try: if not os.path.exists(self.db_path): return { - 'status': 'error', - 'message': f'Database file not found: {self.db_path}' + "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'] + prompt_count = cursor.fetchone()["count"] self.logger.info(f" [NOTE] 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 self.logger.info(f" [IMG] Images table exists: {has_images_table}") - + return { - 'status': 'ok', - 'prompt_count': prompt_count, - 'has_images_table': has_images_table + "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)}' - } - + return {"status": "error", "message": f"Database error: {str(e)}"} + def check_images_table(self) -> Dict[str, Any]: """Check the generated_images table specifically. - + Examines the generated_images table structure and content to verify the gallery system can function properly. - + Returns: Dictionary containing: - status: 'ok' or 'error' @@ -171,28 +168,28 @@ class GalleryDiagnostics: - recent_images: List of recent image records (if successful) """ self.logger.info("\n[IMG] 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' + "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'] + image_count = cursor.fetchone()["count"] self.logger.info(f" [STATS] Images in database: {image_count}") - + # Get recent images cursor = conn.execute(""" SELECT gi.*, p.text @@ -202,29 +199,28 @@ class GalleryDiagnostics: LIMIT 5 """) recent_images = [dict(row) for row in cursor.fetchall()] - + self.logger.info(f" [TIME] Recent images: {len(recent_images)}") for img in recent_images: - self.logger.info(f" - {img['filename']} -> Prompt {img['prompt_id']}") - + self.logger.info( + f" - {img['filename']} -> Prompt {img['prompt_id']}" + ) + return { - 'status': 'ok', - 'image_count': image_count, - 'recent_images': recent_images + "status": "ok", + "image_count": image_count, + "recent_images": recent_images, } - + except Exception as e: - return { - 'status': 'error', - 'message': f'Images table error: {str(e)}' - } - + return {"status": "error", "message": f"Images table error: {str(e)}"} + def check_file_system(self) -> Dict[str, Any]: """Check file system and permissions. - + Verifies that the application has appropriate file system access for reading images and writing database files. - + Returns: Dictionary containing: - status: 'ok' or 'error' @@ -233,42 +229,35 @@ class GalleryDiagnostics: - can_write: Whether write access is available """ self.logger.info("\n[DIR] Checking File System...") - + try: # Check current directory current_dir = os.getcwd() self.logger.info(f" [FOLDER] 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: + with open(test_file, "w") as f: f.write("test") os.remove(test_file) can_write = True except (IOError, OSError): can_write = False - + self.logger.info(f" [EDIT] Can write to directory: {can_write}") - - return { - 'status': 'ok', - 'current_dir': current_dir, - 'can_write': 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)}' - } - + return {"status": "error", "message": f"File system error: {str(e)}"} + def check_comfyui_output(self) -> Dict[str, Any]: """Check ComfyUI output directories. - + Attempts to locate ComfyUI output directories where generated images would be stored. Checks both common locations and ComfyUI's configured paths. - + Returns: Dictionary containing: - status: 'ok' if directories found, 'warning' if none found @@ -276,56 +265,57 @@ class GalleryDiagnostics: - output_dirs: List of detected output directory paths """ self.logger.info("\n[STYLE] Checking ComfyUI Output...") - + output_dirs = [] - + # Try to detect ComfyUI output directories potential_dirs = [ "output", "../output", - "../../output", + "../../output", "ComfyUI/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) self.logger.info(f" [DIR] 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}')) + for ext in [".png", ".jpg", ".jpeg", ".webp"]: + image_files.extend(Path(abs_path).rglob(f"*{ext}")) self.logger.info(f" [IMG] Images found: {len(image_files)}") except Exception as e: self.logger.error(f" [FAIL] 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) self.logger.info(f" [DIR] ComfyUI output dir: {comfyui_output}") except ImportError: self.logger.warning(" [WARN] 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 + "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. - + Verifies that all required Python packages are available for the PromptManager system to function properly. - + Returns: Dictionary containing: - status: 'ok' if all dependencies available, 'error' if any missing @@ -333,55 +323,56 @@ class GalleryDiagnostics: - dependencies: Dictionary mapping package names to availability status """ self.logger.info("\n[PKG] Checking Dependencies...") - - dependencies = { - 'watchdog': False, - 'PIL': False, - 'sqlite3': False - } - + + dependencies = {"watchdog": False, "PIL": False, "sqlite3": False} + # Check watchdog try: import watchdog - dependencies['watchdog'] = True + + dependencies["watchdog"] = True self.logger.info(f" [PASS] watchdog: {watchdog.__version__}") except ImportError: self.logger.error(" [FAIL] watchdog: NOT INSTALLED") - + # Check PIL try: from PIL import Image - dependencies['PIL'] = True + + dependencies["PIL"] = True self.logger.info(f" [PASS] PIL (Pillow): Available") except ImportError: self.logger.error(" [FAIL] PIL (Pillow): NOT AVAILABLE") - + # Check sqlite3 try: import sqlite3 - dependencies['sqlite3'] = True + + dependencies["sqlite3"] = True self.logger.info(f" [PASS] sqlite3: {sqlite3.sqlite_version}") except ImportError: self.logger.error(" [FAIL] 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 + "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]: + + 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. - + Creates a test entry in the generated_images table to verify that the image linking functionality is working correctly. - + Args: prompt_id: ID of an existing prompt to link the test image to test_image_path: Optional path for the test image (uses fake path if None) - + Returns: Dictionary containing: - status: 'ok' if test link created successfully, 'error' otherwise @@ -389,59 +380,62 @@ class GalleryDiagnostics: - image_id: ID of created test image record (if successful) """ self.logger.info(f"\n[TEST] 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__)))) + + 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' + "file_info": { + "size": 1024000, + "dimensions": [512, 512], + "format": "PNG", }, - 'workflow': {'test': True}, - 'prompt': {'test_prompt': 'This is a test'} + "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 + metadata=test_metadata, ) - + self.logger.info(f" [PASS] Test image linked with ID: {image_id}") - + return { - 'status': 'ok', - 'image_id': image_id, - 'message': f'Test image linked successfully with ID {image_id}' + "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)}' + "status": "error", + "message": f"Failed to create test link: {str(e)}", } def run_diagnostics() -> Dict[str, Any]: """Run diagnostics from command line. - + Convenience function to create a GalleryDiagnostics instance and run the full diagnostic suite with default settings. - + Returns: Complete diagnostic results dictionary """ @@ -450,4 +444,4 @@ def run_diagnostics() -> Dict[str, Any]: if __name__ == "__main__": - run_diagnostics() \ No newline at end of file + run_diagnostics() diff --git a/utils/hashing.py b/utils/hashing.py index fc2253e..c37d05b 100644 --- a/utils/hashing.py +++ b/utils/hashing.py @@ -13,11 +13,11 @@ Key features: Typical usage: from utils.hashing import generate_prompt_hash, is_duplicate_prompt - + hash1 = generate_prompt_hash("Beautiful landscape") hash2 = generate_prompt_hash(" beautiful landscape ") # hash1 == hash2 (normalization makes them identical) - + if is_duplicate_prompt(text1, text2): print("Duplicate prompts detected") @@ -34,20 +34,20 @@ import hashlib def generate_prompt_hash(text: str) -> str: """ Generate a SHA256 hash for prompt text to enable deduplication. - + Normalizes the input text (strips whitespace, converts to lowercase) before hashing to ensure consistent results for functionally identical prompts with minor formatting differences. - + Args: text: The prompt text to hash - + Returns: SHA256 hexadecimal digest of the normalized text - + Raises: TypeError: If the input is not a string - + Example: >>> generate_prompt_hash("Beautiful Landscape") 'c3d5f8a9b2e1...' @@ -57,69 +57,79 @@ def generate_prompt_hash(text: str) -> str: """ 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() + + return hashlib.sha256(normalized_text.encode("utf-8")).hexdigest() def generate_content_hash(content: dict) -> str: """ Generate a hash for prompt content including metadata. - + Creates a comprehensive hash that includes not just the prompt text but also associated metadata like category, tags, and workflow name. This enables detection of prompts that are identical in all aspects. - + Args: content: Dictionary containing prompt data with optional keys: - text: The prompt text - category: Prompt category - tags: List of tags - workflow_name: Associated workflow name - + Returns: SHA256 hexadecimal digest of the normalized content structure - + Note: The hash is generated from a normalized JSON representation with sorted keys and normalized text fields to ensure consistency. """ 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 '' + "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() + + 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. - + Compares the normalized hashes of two prompt texts to determine if they are functionally identical. This is an exact match after normalization, not a similarity measure. - + Args: text1: First prompt text to compare text2: Second prompt text to compare threshold: Similarity threshold (not used - kept for API compatibility) - + Returns: True if the prompts have identical normalized hashes (i.e., are duplicates), False otherwise - + Note: The threshold parameter is not used in the current implementation as this performs exact hash matching rather than similarity scoring. @@ -127,5 +137,5 @@ def is_duplicate_prompt(text1: str, text2: str, threshold: float = 0.95) -> bool """ hash1 = generate_prompt_hash(text1) hash2 = generate_prompt_hash(text2) - - return hash1 == hash2 \ No newline at end of file + + return hash1 == hash2 diff --git a/utils/image_monitor.py b/utils/image_monitor.py index 32015d8..dd69c02 100644 --- a/utils/image_monitor.py +++ b/utils/image_monitor.py @@ -11,7 +11,7 @@ The main components are: Typical usage: from utils.image_monitor import ImageMonitor - + monitor = ImageMonitor(db_manager, prompt_tracker) monitor.start_monitoring(['/path/to/comfyui/output']) @@ -38,22 +38,22 @@ from .logging_config import get_logger class ImageGenerationHandler(FileSystemEventHandler): """Filesystem event handler for detecting new image generation. - + This handler extends watchdog's FileSystemEventHandler to specifically handle new image file creation events in ComfyUI output directories. When a new image is detected, it attempts to: 1. Extract ComfyUI metadata from the image 2. Associate the image with the currently active prompt 3. Store the relationship in the database - + The handler implements a small delay before processing to ensure files are completely written before attempting to read them. """ - + def __init__(self, db_manager, prompt_tracker): """ Initialize the image generation handler. - + Args: db_manager: Database manager instance for storing image-prompt relationships prompt_tracker: Prompt tracking instance for getting current active prompts @@ -61,24 +61,25 @@ class ImageGenerationHandler(FileSystemEventHandler): self.db_manager = db_manager self.prompt_tracker = prompt_tracker self.metadata_extractor = ComfyUIMetadataExtractor() - self.logger = get_logger('prompt_manager.image_monitor') + self.logger = get_logger("prompt_manager.image_monitor") # Read from GalleryConfig if available, otherwise use defaults try: from ..py.config import GalleryConfig + self.processing_delay = GalleryConfig.PROCESSING_DELAY self.supported_extensions = tuple(GalleryConfig.SUPPORTED_EXTENSIONS) except Exception: self.processing_delay = 2.0 - self.supported_extensions = ('.png', '.jpg', '.jpeg', '.webp', '.gif') - + self.supported_extensions = (".png", ".jpg", ".jpeg", ".webp", ".gif") + def on_created(self, event): """Handle filesystem creation events. - + This method is called by watchdog when a new file is created in a monitored directory. It filters for image files and schedules them for processing after a small delay to ensure the file is fully written. - + Args: event: FileSystemEvent object containing event details """ @@ -86,11 +87,9 @@ class ImageGenerationHandler(FileSystemEventHandler): self.logger.info(f"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] + 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 and not a thumbnail. @@ -102,20 +101,20 @@ class ImageGenerationHandler(FileSystemEventHandler): inside a thumbnails directory, False otherwise """ # Skip files in thumbnails directory - those are derivatives, not generated images - if '/thumbnails/' in filepath or '\\thumbnails\\' in filepath: + if "/thumbnails/" in filepath or "\\thumbnails\\" in filepath: return False return filepath.lower().endswith(self.supported_extensions) - + def process_new_image(self, image_path: str): """Process a newly created image file for gallery integration. - + This method handles the complete processing pipeline for a new image: 1. Verifies the file still exists 2. Gets the current prompt context from the tracker 3. Extracts ComfyUI metadata from the image 4. Links the image to the appropriate prompt in the database 5. Handles fallback scenarios when no active prompt is available - + Args: image_path: Full path to the newly created image file """ @@ -128,8 +127,10 @@ class ImageGenerationHandler(FileSystemEventHandler): # Get current prompt context first current_prompt = self.prompt_tracker.get_current_prompt() - self.logger.info(f"Current prompt context: {current_prompt['id'] if current_prompt else 'None'}") - + self.logger.info( + f"Current prompt context: {current_prompt['id'] if current_prompt else 'None'}" + ) + if not current_prompt: self.logger.debug(f"No active prompt context for image: {image_path}") # Fallback: try to link to the most recent prompt in database @@ -141,9 +142,11 @@ class ImageGenerationHandler(FileSystemEventHandler): return else: # Extend the timeout for this prompt since we're still getting images - if 'execution_id' in current_prompt: - self.prompt_tracker.extend_prompt_timeout(current_prompt['execution_id'], 300) # Add 5 more minutes - + if "execution_id" in current_prompt: + self.prompt_tracker.extend_prompt_timeout( + current_prompt["execution_id"], 300 + ) # Add 5 more minutes + # Extract ComfyUI metadata try: metadata = self.metadata_extractor.extract_metadata(image_path) @@ -151,30 +154,37 @@ class ImageGenerationHandler(FileSystemEventHandler): except Exception as meta_error: self.logger.warning(f"Metadata extraction failed: {meta_error}") metadata = None - + if metadata: - self.logger.debug(f"Linking image with full metadata to prompt {current_prompt['id']}") + self.logger.debug( + f"Linking image with full metadata to prompt {current_prompt['id']}" + ) self.link_image_to_prompt(image_path, current_prompt, metadata) else: - self.logger.debug(f"Linking image with basic info to prompt {current_prompt['id']}") + self.logger.debug( + f"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}) - + self.link_image_to_prompt( + image_path, current_prompt, {"file_info": basic_metadata} + ) + except Exception as e: self.logger.error(f"Error processing image {image_path}: {e}") import traceback + self.logger.error(traceback.format_exc()) - + def get_basic_file_info(self, image_path: str) -> Dict[str, Any]: """Get basic file information when metadata extraction fails. - + Provides fallback file information when ComfyUI metadata cannot be extracted from the image. Includes file size, format, and dimensions when possible. - + Args: image_path: Path to the image file - + Returns: Dictionary containing basic file information: - size: File size in bytes @@ -183,34 +193,30 @@ class ImageGenerationHandler(FileSystemEventHandler): """ try: from PIL import Image - + stat = os.stat(image_path) - file_info = { - 'size': stat.st_size, - 'format': None, - 'dimensions': None - } - + 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 + file_info["dimensions"] = list(img.size) + file_info["format"] = img.format except Exception: pass - + return file_info except Exception as e: self.logger.error(f"Error getting file info: {e}") return {} - + def _get_fallback_prompt(self) -> Optional[Dict[str, Any]]: """Get the most recent prompt from database as fallback. - + When no active prompt is available from the tracker, this method attempts to find the most recently created prompt in the database to use as a fallback for image linking. - + Returns: Dictionary containing prompt information with 'fallback' flag set to True, or None if no recent prompt is available @@ -220,21 +226,23 @@ class ImageGenerationHandler(FileSystemEventHandler): if recent_prompts: prompt = recent_prompts[0] return { - 'id': prompt['id'], - 'text': prompt['text'], - 'timestamp': prompt.get('created_at'), - 'fallback': True + "id": prompt["id"], + "text": prompt["text"], + "timestamp": prompt.get("created_at"), + "fallback": True, } except Exception as e: self.logger.error(f"Error getting fallback prompt: {e}") return None - - def link_image_to_prompt(self, image_path: str, prompt_context: Dict, metadata: Dict): + + def link_image_to_prompt( + self, image_path: str, prompt_context: Dict, metadata: Dict + ): """Link an image to a prompt in the database. - + Creates a database record associating the generated image with its source prompt, including any extracted metadata from the image file. - + Args: image_path: Full path to the image file prompt_context: Dictionary containing prompt information including ID and text @@ -242,33 +250,33 @@ class ImageGenerationHandler(FileSystemEventHandler): """ try: image_id = self.db_manager.link_image_to_prompt( - prompt_id=prompt_context['id'], - image_path=image_path, - metadata=metadata + prompt_id=prompt_context["id"], image_path=image_path, metadata=metadata + ) + fallback_note = " (fallback)" if prompt_context.get("fallback") else "" + self.logger.debug( + f"Successfully linked image {image_id} to prompt {prompt_context['id']}{fallback_note}" ) - fallback_note = " (fallback)" if prompt_context.get('fallback') else "" - self.logger.debug(f"Successfully linked image {image_id} to prompt {prompt_context['id']}{fallback_note}") except Exception as e: self.logger.error(f"Failed to link image to prompt: {e}") class ImageMonitor: """Main image monitoring system for ComfyUI gallery integration. - + This class manages the overall image monitoring system, including: - Setting up filesystem watchers for output directories - Auto-detecting ComfyUI output locations - Managing the lifecycle of monitoring operations - Providing status information - + The monitor uses watchdog to efficiently watch filesystem changes and can monitor multiple directories simultaneously with recursive subdirectory support. """ - + def __init__(self, db_manager, prompt_tracker): """ Initialize the image monitor. - + Args: db_manager: Database manager instance for storing image relationships prompt_tracker: Prompt tracking instance for getting active prompt context @@ -278,8 +286,8 @@ class ImageMonitor: self.observer = None self.handler = None self.monitored_directories = [] - self.logger = get_logger('prompt_manager.image_monitor') - + self.logger = get_logger("prompt_manager.image_monitor") + def start_monitoring(self, output_directories: Optional[list] = None): """ Start monitoring ComfyUI output directories for new images. @@ -299,6 +307,7 @@ class ImageMonitor: # Check if monitoring is enabled in config try: from ..py.config import GalleryConfig + if not GalleryConfig.MONITORING_ENABLED: self.logger.info("Image monitoring disabled in config") return @@ -309,26 +318,29 @@ class ImageMonitor: if not output_directories: try: from py.config import GalleryConfig + if GalleryConfig.MONITORING_DIRECTORIES: output_directories = GalleryConfig.MONITORING_DIRECTORIES - self.logger.info(f"Using configured monitoring directories: {output_directories}") + self.logger.info( + f"Using configured monitoring directories: {output_directories}" + ) except ImportError: pass # Auto-detect ComfyUI output directory if still none if not output_directories: output_directories = self.detect_comfyui_output_dirs() - + if not output_directories: self.logger.warning("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) @@ -339,13 +351,15 @@ class ImageMonitor: if self.monitored_directories: self.observer.start() - self.logger.info(f"Image monitoring started for {len(self.monitored_directories)} directories") + self.logger.info( + f"Image monitoring started for {len(self.monitored_directories)} directories" + ) else: self.logger.warning("No valid directories to monitor") - + def stop_monitoring(self): """Stop the image monitoring system. - + Cleanly shuts down the filesystem watcher and clears all monitoring state. This method should be called before program exit to ensure proper cleanup. """ @@ -356,47 +370,50 @@ class ImageMonitor: self.handler = None self.monitored_directories = [] self.logger.debug("Image monitoring stopped") - + def detect_comfyui_output_dirs(self) -> list: """Auto-detect ComfyUI output directories. - + Attempts to locate ComfyUI output directories using multiple strategies: 1. Import ComfyUI's folder_paths module to get the configured output directory 2. Search common relative paths where ComfyUI output directories are typically located 3. Verify that detected directories actually exist - + Returns: List of absolute paths to detected 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) self.logger.debug(f"Detected ComfyUI output directory: {output_dir}") except ImportError: - self.logger.debug("ComfyUI folder_paths not available, using fallback detection") - + self.logger.debug( + "ComfyUI folder_paths not available, using fallback detection" + ) + # Fallback: Look for common ComfyUI directory structures fallback_paths = [ "output", - "../output", + "../output", "../../output", "ComfyUI/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) self.logger.debug(f"Found output directory: {abs_path}") - + return potential_dirs - + def get_status(self) -> Dict[str, Any]: """Get monitoring status information. @@ -414,10 +431,10 @@ class ImageMonitor: except Exception: pass return { - 'running': self.observer is not None, - 'monitored_directories': self.monitored_directories, - 'handler_active': self.handler is not None, - 'observer_alive': observer_alive + "running": self.observer is not None, + "monitored_directories": self.monitored_directories, + "handler_active": self.handler is not None, + "observer_alive": observer_alive, } @@ -444,4 +461,4 @@ def get_image_monitor(db_manager, prompt_tracker) -> ImageMonitor: with _monitor_lock: if _monitor_instance is None: _monitor_instance = ImageMonitor(db_manager, prompt_tracker) - return _monitor_instance \ No newline at end of file + return _monitor_instance diff --git a/utils/logging_config.py b/utils/logging_config.py index 2d4535d..023bae7 100644 --- a/utils/logging_config.py +++ b/utils/logging_config.py @@ -23,7 +23,7 @@ from pathlib import Path class PromptManagerLogger: """ Centralized logging system for PromptManager. - + Features: - Configurable log levels - File rotation with size limits @@ -31,13 +31,13 @@ class PromptManagerLogger: - Thread-safe operations - Memory buffer for recent logs """ - + _instance = None _lock = threading.Lock() - + def __new__(cls): """Ensure singleton pattern with thread safety. - + Returns: The single instance of PromptManagerLogger """ @@ -46,145 +46,149 @@ class PromptManagerLogger: if cls._instance is None: cls._instance = super().__new__(cls) return cls._instance - + def __init__(self): """Initialize the logging system if not already initialized. - + Sets up log directory, configuration, memory buffer, and all handlers. Uses _initialized flag to prevent duplicate initialization. """ - if hasattr(self, '_initialized'): + if hasattr(self, "_initialized"): return - + self._initialized = True self.log_dir = Path(__file__).parent.parent / "logs" self.log_dir.mkdir(exist_ok=True) - + # Configuration self.config = { - 'level': 'INFO', - 'max_file_size': 10 * 1024 * 1024, # 10MB - 'backup_count': 5, - 'console_logging': True, - 'file_logging': True, - 'buffer_size': 1000 # Keep last 1000 log entries in memory + "level": "INFO", + "max_file_size": 10 * 1024 * 1024, # 10MB + "backup_count": 5, + "console_logging": True, + "file_logging": True, + "buffer_size": 1000, # Keep last 1000 log entries in memory } - + # Memory buffer for recent logs (for web viewer) - self._log_buffer = collections.deque(maxlen=self.config['buffer_size']) + self._log_buffer = collections.deque(maxlen=self.config["buffer_size"]) self._buffer_lock = threading.Lock() - + # Initialize loggers self._setup_loggers() - + def _setup_loggers(self): """Set up all loggers with appropriate handlers. - + Configures the main logger with console, file, and memory handlers. Sets up file rotation and safe encoding for cross-platform compatibility. """ # Main logger - self.logger = logging.getLogger('prompt_manager') - self.logger.setLevel(getattr(logging, self.config['level'])) - + self.logger = logging.getLogger("prompt_manager") + self.logger.setLevel(getattr(logging, self.config["level"])) + # Clear existing handlers for handler in self.logger.handlers[:]: self.logger.removeHandler(handler) - + # Custom formatter formatter = logging.Formatter( - '%(asctime)s - %(name)s - %(levelname)s - %(filename)s:%(lineno)d - %(message)s', - datefmt='%Y-%m-%d %H:%M:%S' + "%(asctime)s - %(name)s - %(levelname)s - %(filename)s:%(lineno)d - %(message)s", + datefmt="%Y-%m-%d %H:%M:%S", ) - + # Console handler - if self.config['console_logging']: + if self.config["console_logging"]: console_handler = logging.StreamHandler() console_handler.setFormatter(formatter) self.logger.addHandler(console_handler) - + # File handler with rotation and safe encoding for Windows - if self.config['file_logging']: + if self.config["file_logging"]: log_file = self.log_dir / "prompt_manager.log" file_handler = logging.handlers.RotatingFileHandler( log_file, - maxBytes=self.config['max_file_size'], - backupCount=self.config['backup_count'], - encoding='utf-8', - errors='replace' # Replace problematic characters instead of crashing + maxBytes=self.config["max_file_size"], + backupCount=self.config["backup_count"], + encoding="utf-8", + errors="replace", # Replace problematic characters instead of crashing ) file_handler.setFormatter(formatter) self.logger.addHandler(file_handler) - + # Memory handler for web viewer memory_handler = MemoryBufferHandler(self) memory_handler.setFormatter(formatter) self.logger.addHandler(memory_handler) - + # Component-specific loggers self._setup_component_loggers() - + def _setup_component_loggers(self): """Set up loggers for specific components. - + Creates child loggers for different PromptManager components. Child loggers inherit handlers from the main logger. """ components = [ - 'prompt_manager.database', - 'prompt_manager.api', - 'prompt_manager.image_monitor', - 'prompt_manager.prompt_tracker', - 'prompt_manager.web_ui' + "prompt_manager.database", + "prompt_manager.api", + "prompt_manager.image_monitor", + "prompt_manager.prompt_tracker", + "prompt_manager.web_ui", ] - + for component in components: comp_logger = logging.getLogger(component) - comp_logger.setLevel(getattr(logging, self.config['level'])) + comp_logger.setLevel(getattr(logging, self.config["level"])) # Child loggers inherit handlers from parent - - def get_logger(self, name: str = 'prompt_manager') -> logging.Logger: + + def get_logger(self, name: str = "prompt_manager") -> logging.Logger: """Get a logger instance for a specific component. - + Args: name: Logger name, typically in format 'prompt_manager.component' - + Returns: Configured logger instance """ return logging.getLogger(name) - + def add_to_buffer(self, record: logging.LogRecord, formatted_message: str): """Add a log entry to the memory buffer for web viewer. - + Args: record: The LogRecord object from the logging system formatted_message: The formatted log message string """ with self._buffer_lock: log_entry = { - 'timestamp': datetime.fromtimestamp(record.created).isoformat(), - 'level': record.levelname, - 'logger': record.name, - 'message': record.getMessage(), - 'formatted': formatted_message, - 'module': record.module, - 'filename': record.filename, - 'lineno': record.lineno, - 'thread': record.thread, - 'thread_name': record.threadName if hasattr(record, 'threadName') else '', - 'process': record.process + "timestamp": datetime.fromtimestamp(record.created).isoformat(), + "level": record.levelname, + "logger": record.name, + "message": record.getMessage(), + "formatted": formatted_message, + "module": record.module, + "filename": record.filename, + "lineno": record.lineno, + "thread": record.thread, + "thread_name": ( + record.threadName if hasattr(record, "threadName") else "" + ), + "process": record.process, } - + self._log_buffer.append(log_entry) - - def get_recent_logs(self, limit: int = 100, level: Optional[str] = None) -> List[Dict[str, Any]]: + + def get_recent_logs( + self, limit: int = 100, level: Optional[str] = None + ) -> List[Dict[str, Any]]: """Get recent log entries from memory buffer. - + Args: limit: Maximum number of log entries to return level: Optional log level filter (DEBUG, INFO, WARNING, ERROR, CRITICAL) - + Returns: List of log entry dictionaries, most recent first """ @@ -195,14 +199,16 @@ class PromptManagerLogger: if level: level_num = getattr(logging, level.upper(), None) if level_num: - logs = [log for log in logs if getattr(logging, log['level']) >= level_num] + logs = [ + log for log in logs if getattr(logging, log["level"]) >= level_num + ] # Return most recent first return list(reversed(logs[-limit:])) - + def get_log_files(self) -> List[Dict[str, Any]]: """Get information about available log files. - + Returns: List of dictionaries containing file information: - filename: Name of the log file @@ -212,133 +218,140 @@ class PromptManagerLogger: - is_main: True if this is the main log file """ log_files = [] - + for log_file in self.log_dir.glob("*.log*"): try: stat = log_file.stat() - log_files.append({ - 'filename': log_file.name, - 'path': str(log_file), - 'size': stat.st_size, - 'modified': datetime.fromtimestamp(stat.st_mtime).isoformat(), - 'is_main': log_file.name == 'prompt_manager.log' - }) + log_files.append( + { + "filename": log_file.name, + "path": str(log_file), + "size": stat.st_size, + "modified": datetime.fromtimestamp(stat.st_mtime).isoformat(), + "is_main": log_file.name == "prompt_manager.log", + } + ) except Exception as e: self.logger.error(f"Error getting info for log file {log_file}: {e}") - + # Sort by modification time, newest first - log_files.sort(key=lambda x: x['modified'], reverse=True) + log_files.sort(key=lambda x: x["modified"], reverse=True) return log_files - + def read_log_file(self, filename: str, lines: int = 100) -> List[str]: """Read lines from a specific log file. - + Args: filename: Name of the log file to read (must be in log directory) lines: Number of lines to read from the end of the file (0 for all) - + Returns: List of line strings from the log file - + Raises: FileNotFoundError: If the log file doesn't exist ValueError: If the file path is outside the log directory """ log_file = self.log_dir / filename - + if not log_file.exists() or not log_file.is_file(): raise FileNotFoundError(f"Log file not found: {filename}") - + # Security check - ensure file is in log directory if not str(log_file.resolve()).startswith(str(self.log_dir.resolve())): raise ValueError("Invalid log file path") - + try: # Use UTF-8 encoding with error handling for Windows compatibility - with open(log_file, 'r', encoding='utf-8', errors='replace') as f: + with open(log_file, "r", encoding="utf-8", errors="replace") as f: all_lines = f.readlines() return all_lines[-lines:] if lines > 0 else all_lines except Exception as e: self.logger.error(f"Error reading log file {filename}: {e}") raise - + def truncate_logs(self) -> Dict[str, Any]: """Truncate all log files. - + Clears the main log file and deletes rotated log files. Also clears the memory buffer. - + Returns: Dictionary with: - truncated: List of successfully truncated files - errors: List of error messages for failed operations """ - results = { - 'truncated': [], - 'errors': [] - } - + results = {"truncated": [], "errors": []} + for log_file in self.log_dir.glob("prompt_manager.log*"): try: - if log_file.name == 'prompt_manager.log': + if log_file.name == "prompt_manager.log": # For main log file, just clear it with safe encoding - with open(log_file, 'w', encoding='utf-8', errors='replace') as f: + with open(log_file, "w", encoding="utf-8", errors="replace") as f: f.write("") else: # For rotated files, delete them log_file.unlink() - - results['truncated'].append(log_file.name) + + results["truncated"].append(log_file.name) self.logger.info(f"Truncated log file: {log_file.name}") - + except Exception as e: error_msg = f"Failed to truncate {log_file.name}: {str(e)}" - results['errors'].append(error_msg) + results["errors"].append(error_msg) self.logger.error(error_msg) - + # Clear memory buffer with self._buffer_lock: self._log_buffer.clear() - + return results - + def update_config(self, new_config: Dict[str, Any]): """Update logging configuration. - + Args: new_config: Dictionary of configuration updates (level, console_logging, file_logging, etc.) """ self.config.update(new_config) - + # Reconfigure loggers if level changed - if 'level' in new_config: - level = getattr(logging, new_config['level'].upper(), logging.INFO) + if "level" in new_config: + level = getattr(logging, new_config["level"].upper(), logging.INFO) self.logger.setLevel(level) - + # Update all component loggers for logger_name in logging.Logger.manager.loggerDict: - if logger_name.startswith('prompt_manager'): + if logger_name.startswith("prompt_manager"): logger = logging.getLogger(logger_name) logger.setLevel(level) - + # Re-setup loggers if handlers changed - if any(key in new_config for key in ['console_logging', 'file_logging', 'max_file_size', 'backup_count']): + if any( + key in new_config + for key in [ + "console_logging", + "file_logging", + "max_file_size", + "backup_count", + ] + ): self._setup_loggers() - + self.logger.info(f"Updated logging configuration: {new_config}") - + def get_config(self) -> Dict[str, Any]: """Get current logging configuration. - + Returns: Copy of the current configuration dictionary """ return self.config.copy() - + def get_log_stats(self) -> Dict[str, Any]: """Get logging statistics. - + Returns: Dictionary containing: - buffer_count: Number of entries in memory buffer @@ -352,45 +365,45 @@ class PromptManagerLogger: buffer_count = len(self._log_buffer) level_counts = {} for entry in self._log_buffer: - level = entry['level'] + level = entry["level"] level_counts[level] = level_counts.get(level, 0) + 1 - + log_files = self.get_log_files() - total_log_size = sum(f['size'] for f in log_files) - + total_log_size = sum(f["size"] for f in log_files) + return { - 'buffer_count': buffer_count, - 'level_counts': level_counts, - 'log_files_count': len(log_files), - 'total_log_size': total_log_size, - 'log_directory': str(self.log_dir), - 'current_level': self.config['level'] + "buffer_count": buffer_count, + "level_counts": level_counts, + "log_files_count": len(log_files), + "total_log_size": total_log_size, + "log_directory": str(self.log_dir), + "current_level": self.config["level"], } class MemoryBufferHandler(logging.Handler): """Custom logging handler that stores logs in memory for web viewer. - + This handler extends the standard logging.Handler to capture log records and store them in the PromptManagerLogger's memory buffer. This enables the web interface to display recent log entries without reading from files. - + The handler is designed to be failure-safe - if any error occurs during log processing, it silently continues without disrupting the application. """ - + def __init__(self, logger_manager: PromptManagerLogger): """Initialize the memory buffer handler. - + Args: logger_manager: The PromptManagerLogger instance to store logs in """ super().__init__() self.logger_manager = logger_manager - + def emit(self, record: logging.LogRecord): """Handle a log record by adding it to the memory buffer. - + Args: record: The LogRecord to process and store """ @@ -405,9 +418,10 @@ class MemoryBufferHandler(logging.Handler): # Global logger instance _logger_manager = None + def get_logger_manager() -> PromptManagerLogger: """Get the global logger manager instance. - + Returns: The singleton PromptManagerLogger instance, creating it if necessary """ @@ -416,16 +430,18 @@ def get_logger_manager() -> PromptManagerLogger: _logger_manager = PromptManagerLogger() return _logger_manager -def get_logger(name: str = 'prompt_manager') -> logging.Logger: + +def get_logger(name: str = "prompt_manager") -> logging.Logger: """Convenience function to get a logger. - + Args: name: Logger name, defaults to 'prompt_manager' - + Returns: Configured logger instance for the specified name """ return get_logger_manager().get_logger(name) + # Initialize logging on import -get_logger_manager() \ No newline at end of file +get_logger_manager() diff --git a/utils/metadata_extractor.py b/utils/metadata_extractor.py index 2fecac7..d24f1f4 100644 --- a/utils/metadata_extractor.py +++ b/utils/metadata_extractor.py @@ -13,10 +13,10 @@ The extractor supports: Typical usage: from utils.metadata_extractor import ComfyUIMetadataExtractor - + extractor = ComfyUIMetadataExtractor() metadata = extractor.extract_metadata('/path/to/image.png') - + if metadata: workflow = metadata.get('workflow', {}) text_nodes = metadata.get('text_encoder_nodes', []) @@ -41,33 +41,33 @@ from .logging_config import get_logger class ComfyUIMetadataExtractor: """Extracts ComfyUI metadata from generated images. - + This class handles extraction of ComfyUI workflow and generation metadata from PNG images. It parses the text chunks embedded by ComfyUI and extracts structured information about workflows, prompts, and generation parameters. - + The extractor is designed to be robust and handle various ComfyUI workflow formats, including custom nodes and different metadata structures. """ - + def __init__(self): """Initialize the metadata extractor. - + Sets up logging and prepares the extractor for metadata parsing operations. """ - self.logger = get_logger('prompt_manager.metadata_extractor') - + self.logger = get_logger("prompt_manager.metadata_extractor") + def extract_metadata(self, image_path: str) -> Optional[Dict[str, Any]]: """ Extract ComfyUI workflow and prompt metadata from an image. - + This method opens a PNG image and extracts all available ComfyUI metadata from the embedded text chunks. It handles JSON parsing, workflow analysis, and parameter extraction. - + Args: image_path: Path to the PNG image file to analyze - + Returns: Dictionary containing extracted metadata with keys: - file_info: Basic file information (always present) @@ -76,7 +76,7 @@ class ComfyUIMetadataExtractor: - prompt: ComfyUI prompt data (if available) - Additional fields: parameters, model, sampler, steps, etc. (if present) Returns None if no ComfyUI metadata is found or extraction fails - + Raises: Exception: Re-raises any exception that occurs during extraction, with appropriate error logging @@ -84,40 +84,49 @@ class ComfyUIMetadataExtractor: try: with Image.open(image_path) as image: metadata = {} - + # Add basic file information - metadata['file_info'] = self.get_file_info(image_path, image) - + 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: + if hasattr(image, "text") and image.text: # Look for ComfyUI workflow data - if 'workflow' in image.text: + if "workflow" in image.text: try: - workflow_data = json.loads(image.text['workflow']) - metadata['workflow'] = workflow_data - + 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) + text_encoder_nodes = self.find_text_encoder_nodes( + workflow_data + ) if text_encoder_nodes: - metadata['text_encoder_nodes'] = text_encoder_nodes - + metadata["text_encoder_nodes"] = text_encoder_nodes + except json.JSONDecodeError as e: self.logger.warning(f"Failed to parse workflow JSON: {e}") - + # Look for prompt data - if 'prompt' in image.text: + if "prompt" in image.text: try: - prompt_data = json.loads(image.text['prompt']) - metadata['prompt'] = prompt_data + prompt_data = json.loads(image.text["prompt"]) + metadata["prompt"] = prompt_data except json.JSONDecodeError as e: self.logger.warning(f"Failed to parse prompt JSON: {e}") - + # Extract other common metadata fields metadata_fields = [ - 'parameters', 'model', 'sampler', 'steps', 'cfg_scale', - 'seed', 'scheduler', 'positive', 'negative' + "parameters", + "model", + "sampler", + "steps", + "cfg_scale", + "seed", + "scheduler", + "positive", + "negative", ] - + for field in metadata_fields: if field in image.text: try: @@ -126,24 +135,28 @@ class ComfyUIMetadataExtractor: 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 - + + return ( + metadata + if any(key != "file_info" for key in metadata.keys()) + else None + ) + except Exception as e: self.logger.error(f"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. - + Extracts fundamental file properties including size, dimensions, format, and filesystem timestamps. - + Args: image_path: Path to the image file image: Opened PIL Image object - + Returns: Dictionary containing: - size: File size in bytes @@ -156,53 +169,53 @@ class ComfyUIMetadataExtractor: 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 + "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: self.logger.error(f"Error getting file info: {e}") return {} - + def find_text_encoder_nodes(self, workflow_data: Dict) -> list: """ Find text encoder nodes in the workflow data. - + Searches through the workflow structure to identify nodes that perform text encoding operations. This includes standard ComfyUI nodes like CLIPTextEncode as well as custom nodes and extensions. - + Args: workflow_data: ComfyUI workflow data structure (can be in various formats) - + Returns: List of dictionaries representing text encoder nodes, each containing the node's configuration, inputs, and metadata. For dictionary-based workflows, adds 'node_id' field to each node. """ 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'] + 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 @@ -213,98 +226,98 @@ class ComfyUIMetadataExtractor: # 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 + 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. - + Determines whether a given node performs text encoding by examining its type, class, and title for known text encoding patterns. Supports various ComfyUI node types and custom extensions. - + Args: node_data: Node configuration dictionary to analyze - + Returns: True if the node is identified as a text encoder, False otherwise """ 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' + "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 '' - + 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'] - + 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. - + Searches through text encoder nodes to find and extract the actual prompt text that was used for generation. Handles various input field names and data structures. - + Args: workflow_data: ComfyUI workflow data structure - + Returns: The extracted prompt text string, or None if no prompt text is found """ 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', {}) - + inputs = node.get("inputs", {}) + # Common input field names for prompt text - text_fields = ['text', 'prompt', 'positive', 'conditioning'] - + 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. - + Parses the metadata to extract key generation parameters used for image creation, including sampling settings, model information, and other configuration values. - + Args: metadata: Full metadata dictionary from extract_metadata() - + Returns: Dictionary containing generation parameters such as: - steps: Number of sampling steps @@ -317,42 +330,49 @@ class ComfyUIMetadataExtractor: - batch_size: Number of images generated """ parameters = {} - + # Common generation parameters to extract param_fields = [ - 'steps', 'cfg_scale', 'sampler', 'scheduler', 'seed', - 'model', 'width', 'height', 'batch_size' + "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']) + 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. - + Analyzes the workflow structure to identify and extract generation parameters from various node types. This method can be extended to support specific workflow patterns and custom nodes. - + Args: workflow_data: ComfyUI workflow data structure - + Returns: Dictionary of extracted parameters. Currently returns empty dict but can be extended based on specific workflow analysis needs. """ 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 + + return parameters diff --git a/utils/prompt_tracker.py b/utils/prompt_tracker.py index 3ab1523..183acc5 100644 --- a/utils/prompt_tracker.py +++ b/utils/prompt_tracker.py @@ -14,11 +14,11 @@ Key features: Typical usage: from utils.prompt_tracker import PromptTracker - + tracker = PromptTracker(db_manager) execution_id = tracker.set_current_prompt("beautiful landscape") # Images generated after this point will be linked to this prompt - + Or using the context manager: with PromptExecutionContext(tracker, "beautiful landscape"): # Generate images here @@ -38,6 +38,7 @@ try: except ImportError: import sys import os + current_dir = os.path.dirname(os.path.dirname(os.path.abspath(__file__))) sys.path.insert(0, current_dir) from utils.logging_config import get_logger @@ -45,32 +46,32 @@ except ImportError: class PromptTracker: """Thread-safe tracking of current prompt executions. - + This class manages active prompt contexts across multiple threads, enabling the image monitoring system to correctly associate generated images with their source prompts. It uses both thread-local storage and global tracking to handle various execution scenarios. - + The tracker automatically cleans up expired prompts and provides fallback mechanisms when prompts are accessed from different threads (e.g., image monitoring running in a separate thread from prompt execution). - + Attributes: active_prompts: Global dictionary of active prompt contexts cleanup_interval: Seconds between cleanup operations (default: 300) prompt_timeout: Seconds before a prompt expires (default: 600) """ - + def __init__(self, db_manager): """ Initialize the prompt tracker. - + Args: db_manager: Database manager instance for prompt operations """ - self.logger = get_logger('prompt_manager.prompt_tracker') + self.logger = get_logger("prompt_manager.prompt_tracker") self.logger.debug("Initializing PromptTracker") - + self.db_manager = db_manager self._local = threading.local() self.active_prompts = {} # Global tracking for multiple threads @@ -79,51 +80,59 @@ class PromptTracker: # Read from GalleryConfig if available, otherwise use defaults try: from ..py.config import GalleryConfig + self.cleanup_interval = GalleryConfig.CLEANUP_INTERVAL self.prompt_timeout = GalleryConfig.PROMPT_TIMEOUT except Exception: - self.cleanup_interval = 300 # 5 minutes - self.prompt_timeout = 600 # 10 minutes - + self.cleanup_interval = 300 # 5 minutes + self.prompt_timeout = 600 # 10 minutes + # Start cleanup thread - self.cleanup_thread = threading.Thread(target=self._cleanup_expired_prompts, daemon=True) + self.cleanup_thread = threading.Thread( + target=self._cleanup_expired_prompts, daemon=True + ) self.cleanup_thread.start() self.logger.debug("PromptTracker initialization completed") - - def set_current_prompt(self, prompt_text: str, additional_data: Optional[Dict[str, Any]] = None) -> str: + + def set_current_prompt( + self, prompt_text: str, additional_data: Optional[Dict[str, Any]] = None + ) -> str: """ Set the current prompt for this thread for image tracking. - + This method establishes the prompt context that will be used to link any subsequently generated images. The prompt is stored in both thread-local storage and global tracking to handle cross-thread access scenarios. - + Note: Prompt saving is handled by PromptManager to avoid duplicates. - + Args: prompt_text: The prompt text being executed additional_data: Additional prompt metadata, should include prompt_id from PromptManager - + Returns: Unique execution ID for this prompt execution """ execution_id = self.generate_execution_id() - + # Get prompt_id from additional_data (passed from PromptManager) # PromptTracker should NOT save prompts to avoid duplicates - prompt_id = additional_data.get('prompt_id') if additional_data else None - + prompt_id = additional_data.get("prompt_id") if additional_data else None + if not prompt_id: # Fallback: try to find existing prompt using consistent hash calculation try: # Use consistent hash calculation (same as PromptManager) import hashlib + normalized_text = prompt_text.strip().lower() - prompt_hash = hashlib.sha256(normalized_text.encode('utf-8')).hexdigest() + prompt_hash = hashlib.sha256( + normalized_text.encode("utf-8") + ).hexdigest() existing_prompt = self.db_manager.get_prompt_by_hash(prompt_hash) - + if existing_prompt: - prompt_id = existing_prompt['id'] + prompt_id = existing_prompt["id"] self.logger.debug(f"Found existing prompt ID: {prompt_id}") else: # Generate a temporary ID for tracking (prompt should be saved by PromptManager) @@ -132,36 +141,38 @@ class PromptTracker: except Exception as e: self.logger.error(f"Error finding prompt: {e}") 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 {} + "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 - - self.logger.debug(f"Set current prompt: {execution_id} -> {prompt_text[:50]}... (thread: {threading.current_thread().ident})") + + self.logger.debug( + f"Set current prompt: {execution_id} -> {prompt_text[:50]}... (thread: {threading.current_thread().ident})" + ) self.logger.debug(f"Active prompts count: {len(self.active_prompts)}") return execution_id - + def get_current_prompt(self) -> Optional[Dict[str, Any]]: """ Get the current prompt context for this thread. - + Attempts to retrieve the current prompt context, first from thread-local storage, then from global tracking as a fallback. This enables image monitoring (which may run in a different thread) to access prompt context. - + Returns: Dictionary containing prompt context with keys: - id: Prompt ID in database @@ -172,94 +183,112 @@ class PromptTracker: - additional_data: Any additional metadata Returns None if no valid prompt context is found """ - current = getattr(self._local, 'current_prompt', None) - + current = getattr(self._local, "current_prompt", None) + if current: # Check if prompt hasn't expired - if time.time() - current['timestamp'] < self.prompt_timeout: + if time.time() - current["timestamp"] < self.prompt_timeout: return current else: self.logger.debug(f"Current prompt expired: {current['execution_id']}") self.clear_current_prompt() - + # Fallback: try to find recent prompt from global tracking # This is crucial for image monitoring which runs in different threads recent_prompt = self._find_recent_prompt() if recent_prompt: - self.logger.debug(f"Using recent prompt from global tracking: {recent_prompt['execution_id']}") + self.logger.debug( + f"Using recent prompt from global tracking: {recent_prompt['execution_id']}" + ) return recent_prompt - - self.logger.debug(f"No prompt context found (thread: {threading.current_thread().ident})") + + self.logger.debug( + f"No prompt context found (thread: {threading.current_thread().ident})" + ) return None - + def _find_recent_prompt(self) -> Optional[Dict[str, Any]]: """Find the most recent prompt that's still valid. - + Searches the global prompt tracking for the most recently set prompt that hasn't expired. This is used as a fallback when thread-local storage doesn't contain a current prompt. - + Returns: Most recent valid prompt context, or None if no valid prompts found """ with self.lock: current_time = time.time() - self.logger.debug(f"Searching for recent prompt among {len(self.active_prompts)} active prompts") - + self.logger.debug( + f"Searching for recent prompt among {len(self.active_prompts)} active prompts" + ) + recent_prompts = [] for exec_id, prompt in self.active_prompts.items(): - age_seconds = current_time - prompt['timestamp'] - self.logger.debug(f"Prompt {exec_id}: age={age_seconds:.1f}s, timeout={self.prompt_timeout}s") - + age_seconds = current_time - prompt["timestamp"] + self.logger.debug( + f"Prompt {exec_id}: age={age_seconds:.1f}s, timeout={self.prompt_timeout}s" + ) + if age_seconds < self.prompt_timeout: recent_prompts.append(prompt) else: self.logger.debug(f"Prompt {exec_id} is expired") - + if recent_prompts: # Return the most recent one - most_recent = max(recent_prompts, key=lambda p: p['timestamp']) + most_recent = max(recent_prompts, key=lambda p: p["timestamp"]) self.logger.debug(f"Found recent prompt: {most_recent['execution_id']}") return most_recent else: self.logger.debug(f"No recent prompts found") - + return None - + def clear_current_prompt(self): """Clear the current prompt context for this thread. - + Removes the prompt context from both thread-local storage and global tracking. This should be called when prompt execution is complete, though the system also handles automatic cleanup via timeouts. """ - current = getattr(self._local, 'current_prompt', None) + current = getattr(self._local, "current_prompt", None) if current: - execution_id = current['execution_id'] - self.logger.debug(f"Clearing current prompt: {execution_id} (thread: {threading.current_thread().ident})") - + execution_id = current["execution_id"] + self.logger.debug( + f"Clearing current prompt: {execution_id} (thread: {threading.current_thread().ident})" + ) + # Clear from thread-local storage self._local.current_prompt = None - + # Remove from global tracking with self.lock: removed = self.active_prompts.pop(execution_id, None) if removed: - self.logger.debug(f"Removed prompt from global tracking: {execution_id}") - self.logger.debug(f"Remaining active prompts: {len(self.active_prompts)}") + self.logger.debug( + f"Removed prompt from global tracking: {execution_id}" + ) + self.logger.debug( + f"Remaining active prompts: {len(self.active_prompts)}" + ) else: - self.logger.debug(f"Prompt {execution_id} was not in global tracking") + self.logger.debug( + f"Prompt {execution_id} was not in global tracking" + ) else: - self.logger.debug(f"No current prompt to clear (thread: {threading.current_thread().ident})") - + self.logger.debug( + f"No current prompt to clear (thread: {threading.current_thread().ident})" + ) + def extend_prompt_timeout(self, execution_id: str, additional_seconds: int = 60): """ Extend the timeout for a specific prompt execution. - + This is useful for long-running image generation processes where images may continue to be generated well after the initial prompt execution. Updates the timestamp to prevent the prompt from being cleaned up. - + Args: execution_id: The execution ID to extend additional_seconds: Additional seconds to add to timeout (currently unused, @@ -267,22 +296,22 @@ class PromptTracker: """ with self.lock: if execution_id in self.active_prompts: - self.active_prompts[execution_id]['timestamp'] = time.time() + self.active_prompts[execution_id]["timestamp"] = time.time() self.logger.debug(f"Extended timeout for prompt: {execution_id}") - + def generate_execution_id(self) -> str: """Generate a unique execution ID. - + Creates a unique identifier for this prompt execution using UUID and timestamp. - + Returns: Unique execution ID in format 'exec_{uuid8}_{timestamp}' """ return f"exec_{uuid.uuid4().hex[:8]}_{int(time.time())}" - + def _cleanup_expired_prompts(self): """Background thread to clean up expired prompts. - + Runs continuously as a daemon thread, periodically removing expired prompt contexts from the global tracking dictionary. This prevents memory leaks from accumulating old prompt data. @@ -291,59 +320,62 @@ class PromptTracker: 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: + 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: self.logger.debug(f"Cleaned up {len(expired_ids)} expired prompts") - + time.sleep(self.cleanup_interval) - + except Exception as e: self.logger.error(f"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 mapping execution IDs to prompt context dictionaries. This is a copy of the internal tracking dictionary. """ with self.lock: return self.active_prompts.copy() - + def clear_all_active_prompts(self) -> int: """ Clear all active prompts. - + Removes all prompt contexts from both global and thread-local storage. This is useful before batch operations or when resetting the tracking state. - + Returns: Number of prompts that were cleared from global tracking """ with self.lock: cleared_count = len(self.active_prompts) self.active_prompts.clear() - + # Clear thread-local storage as well self._local.current_prompt = None - + self.logger.debug(f"Cleared {cleared_count} active prompts") return cleared_count - + def get_status(self) -> Dict[str, Any]: """ Get tracker status information. - + Returns: Dictionary containing: - active_prompts_count: Number of currently active prompts @@ -355,28 +387,30 @@ class PromptTracker: """ 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 + "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. - + This context manager provides automatic prompt lifecycle management, ensuring prompt context is properly set and cleaned up. While the current implementation doesn't automatically clear prompts on exit (to allow for images generated after prompt execution), it provides a clean interface for prompt management. - + Example: with PromptExecutionContext(tracker, "beautiful landscape") as exec_id: # Generate images here @@ -384,11 +418,11 @@ class PromptExecutionContext: pass # Prompt context remains active for timeout period """ - + def __init__(self, prompt_tracker: PromptTracker, prompt_text: str, **kwargs): """ Initialize execution context. - + Args: prompt_tracker: PromptTracker instance to use for tracking prompt_text: The prompt text to track @@ -398,21 +432,20 @@ class PromptExecutionContext: self.prompt_text = prompt_text self.additional_data = kwargs self.execution_id = None - + def __enter__(self): """Enter the execution context. - + Sets the current prompt in the tracker and returns the execution ID. - + Returns: Unique execution ID for this prompt """ self.execution_id = self.prompt_tracker.set_current_prompt( - self.prompt_text, - self.additional_data + self.prompt_text, self.additional_data ) return self.execution_id - + def __exit__(self, exc_type, exc_val, exc_tb): """Exit the execution context. @@ -452,4 +485,4 @@ def get_prompt_tracker(db_manager) -> PromptTracker: with _tracker_lock: if _tracker_instance is None: _tracker_instance = PromptTracker(db_manager) - return _tracker_instance \ No newline at end of file + return _tracker_instance diff --git a/utils/validators.py b/utils/validators.py index 6796d2a..edf3fcc 100644 --- a/utils/validators.py +++ b/utils/validators.py @@ -20,7 +20,7 @@ All validation functions follow a consistent pattern: Typical usage: from utils.validators import validate_prompt_text, sanitize_input - + try: validate_prompt_text(user_input) clean_text = sanitize_input(user_input) @@ -37,16 +37,16 @@ from typing import List, Optional, Union def validate_prompt_text(text: str) -> bool: """ Validate prompt text input. - + Ensures the prompt text is a valid string with reasonable length limits. Empty or whitespace-only strings are rejected. - + Args: text: The prompt text to validate - + Returns: True if the text passes validation - + Raises: ValueError: If text is invalid with descriptive message including: - Not a string type @@ -55,28 +55,28 @@ def validate_prompt_text(text: str) -> bool: """ 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. - + Validates rating values on a 1-5 scale, with None allowed for unrated prompts. - + Args: rating: The rating to validate (1-5 scale or None for no rating) - + Returns: True if the rating is valid - + Raises: ValueError: If rating is invalid: - Not an integer (when not None) @@ -84,29 +84,29 @@ def validate_rating(rating: Optional[int]) -> bool: """ 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. - + Accepts tags as comma-separated string, list of strings, or None. Validates each tag for length and character restrictions. - + Args: tags: Tags as comma-separated string, list of strings, or None - + Returns: True if all tags are valid - + Raises: ValueError: If tags are invalid: - Wrong input type (not string, list, or None) @@ -117,48 +117,48 @@ def validate_tags(tags: Union[str, List[str], None]) -> bool: """ 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()] + 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") # Reject control characters and null bytes - if re.search(r'[\x00-\x1f]', tag.strip()): + if re.search(r"[\x00-\x1f]", tag.strip()): raise ValueError(f"Tag '{tag}' contains invalid control 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. - + Validates optional category strings with reasonable length limits and character restrictions. - + Args: category: The category string to validate (None allowed for no category) - + Returns: True if the category is valid or None - + Raises: ValueError: If category is invalid: - Not a string type (when not None) @@ -167,37 +167,37 @@ def validate_category(category: Optional[str]) -> bool: """ 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") # Reject control characters and null bytes - if re.search(r'[\x00-\x1f]', category): + if re.search(r"[\x00-\x1f]", category): raise ValueError("Category contains invalid control characters") - + return True def validate_workflow_name(workflow_name: Optional[str]) -> bool: """ Validate workflow name input. - + Validates optional workflow name strings with generous length limits to accommodate descriptive workflow names. - + Args: workflow_name: The workflow name to validate (None allowed for no workflow) - + Returns: True if the workflow name is valid or None - + Raises: ValueError: If workflow name is invalid: - Not a string type (when not None) @@ -205,31 +205,31 @@ def validate_workflow_name(workflow_name: Optional[str]) -> bool: """ 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. - + Cleans text input by removing control characters, normalizing whitespace, and limiting excessive empty lines. Preserves the semantic content while ensuring safe storage and display. - + Args: text: The text string to sanitize - + Returns: Sanitized text string with: - Null bytes and control characters removed @@ -240,18 +240,18 @@ def sanitize_input(text: str) -> str: """ if not isinstance(text, str): return "" - + # Remove null bytes and other control characters - sanitized = text.replace('\x00', '').replace('\r\n', '\n').replace('\r', '\n') - + sanitized = text.replace("\x00", "").replace("\r\n", "\n").replace("\r", "\n") + # Strip excessive whitespace but preserve single newlines - lines = sanitized.split('\n') + 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 @@ -260,32 +260,32 @@ def sanitize_input(text: str) -> str: else: empty_count = 0 result_lines.append(line) - - return '\n'.join(result_lines).strip() + + 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. - + Converts comma-separated tag strings into a clean, deduplicated list of tags. Each tag is sanitized and trimmed. - + Args: tags_string: Comma-separated tags string (e.g., "tag1, tag2, tag3") - + Returns: List of unique, cleaned tag strings. Empty input returns empty list. Limited to maximum 20 tags to prevent abuse. """ 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(','): + 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 + + return tags[:20] # Limit to 20 tags