diff --git a/.coveragerc b/.coveragerc new file mode 100644 index 0000000..9eded8d --- /dev/null +++ b/.coveragerc @@ -0,0 +1,62 @@ +# Coverage configuration for pytest-cov +[run] +source = . +omit = + */tests/* + */venv/* + */env/* + */__pycache__/* + */migrations/* + */conftest.py + */setup.py + */manage.py + */.tox/* + */node_modules/* + */static/* + */media/* + +branch = true +parallel = true + +[report] +# Regexes for lines to exclude from consideration +exclude_lines = + # Have to re-enable the standard pragma + pragma: no cover + + # Don't complain about missing debug-only code: + def __repr__ + if self\.debug + + # Don't complain if tests don't hit defensive assertion code: + raise AssertionError + raise NotImplementedError + + # Don't complain if non-runnable code isn't run: + if 0: + if __name__ == .__main__.: + + # Don't complain about abstract methods, they aren't run: + @(abc\.)?abstractmethod + + # Don't complain about type checking imports + if TYPE_CHECKING: + + # Don't complain about logger calls + \.logger\.(debug|info|warning|error|critical) + +ignore_errors = true +show_missing = true +skip_covered = false +precision = 2 + +[html] +directory = htmlcov +title = ComfyUI PromptManager Test Coverage Report + +[xml] +output = coverage.xml + +[json] +output = coverage.json +pretty_print = true \ No newline at end of file diff --git a/.github/workflows/test.yml b/.github/workflows/test.yml new file mode 100644 index 0000000..65b722d --- /dev/null +++ b/.github/workflows/test.yml @@ -0,0 +1,30 @@ +name: Tests + +on: + pull_request: + branches: [main] + push: + branches: [main] + +jobs: + test: + runs-on: ubuntu-latest + strategy: + matrix: + python-version: ["3.10", "3.11", "3.12"] + + steps: + - uses: actions/checkout@v4 + + - name: Set up Python ${{ matrix.python-version }} + uses: actions/setup-python@v5 + with: + python-version: ${{ matrix.python-version }} + + - name: Install dependencies + run: | + python -m pip install --upgrade pip + pip install Pillow aiohttp watchdog + + - name: Run tests + run: python -m unittest discover tests/ -v diff --git a/.gitignore b/.gitignore index 0a19790..eea9d96 100644 --- a/.gitignore +++ b/.gitignore @@ -1,3 +1,19 @@ +# Tailwind standalone CLI (downloaded by `make css-setup`) +tailwindcss + +# AI tooling / local config +.claude/ +.playwright-mcp/ +.serena/ +AGENTS.md +CLAUDE.md + +# Local files +docs/ +logs/ +standalone_tagger.py +prompts.db + # Byte-compiled / optimized / DLL files __pycache__/ *.py[cod] @@ -16,6 +32,8 @@ eggs/ .eggs/ lib/ lib64/ +!web/lib/ +!web/lib/** parts/ sdist/ var/ diff --git a/Makefile b/Makefile new file mode 100644 index 0000000..657a310 --- /dev/null +++ b/Makefile @@ -0,0 +1,43 @@ +TAILWIND_VERSION := 3.4.17 +TAILWIND_CLI := ./tailwindcss +UNAME_S := $(shell uname -s) +UNAME_M := $(shell uname -m) + +ifeq ($(UNAME_S),Darwin) + ifeq ($(UNAME_M),arm64) + TAILWIND_PLATFORM := tailwindcss-macos-arm64 + else + TAILWIND_PLATFORM := tailwindcss-macos-x64 + endif +else + ifeq ($(UNAME_M),aarch64) + TAILWIND_PLATFORM := tailwindcss-linux-arm64 + else + TAILWIND_PLATFORM := tailwindcss-linux-x64 + endif +endif + +TAILWIND_URL := https://github.com/tailwindlabs/tailwindcss/releases/download/v$(TAILWIND_VERSION)/$(TAILWIND_PLATFORM) + +.PHONY: css css-watch css-setup clean-css + +## Download Tailwind standalone CLI (dev only, not committed) +css-setup: + @if [ ! -f $(TAILWIND_CLI) ]; then \ + echo "Downloading Tailwind CSS v$(TAILWIND_VERSION) CLI..."; \ + curl -sLo $(TAILWIND_CLI) $(TAILWIND_URL) && chmod +x $(TAILWIND_CLI); \ + else \ + echo "Tailwind CLI already present"; \ + fi + +## Rebuild compiled CSS (run after changing Tailwind classes in HTML/JS) +css: css-setup + $(TAILWIND_CLI) -i web/lib/tailwind/input.css -o web/lib/tailwind/styles.css --minify + +## Watch mode for development +css-watch: css-setup + $(TAILWIND_CLI) -i web/lib/tailwind/input.css -o web/lib/tailwind/styles.css --watch + +## Remove downloaded CLI binary +clean-css: + rm -f $(TAILWIND_CLI) diff --git a/__init__.py b/__init__.py index bc93503..159140a 100644 --- a/__init__.py +++ b/__init__.py @@ -79,7 +79,7 @@ except Exception as e: logger = get_logger("prompt_manager.init") logger.error(f"Failed to register API routes: {e}") - except: + except Exception: pass # Start image monitoring globally at module import time @@ -107,7 +107,7 @@ except Exception as e: 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: + except Exception: print(f"[ComfyUI-PromptManager] Warning: Failed to start image monitoring: {e}") __all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS", "WEB_DIRECTORY"] diff --git a/database/models.py b/database/models.py index e14541a..299380d 100644 --- a/database/models.py +++ b/database/models.py @@ -4,6 +4,7 @@ Database schema and models for KikoTextEncode prompt storage. import sqlite3 import os +import threading from typing import Optional # Import logging system @@ -29,6 +30,8 @@ class PromptModel: 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: @@ -43,6 +46,8 @@ class PromptModel: """ try: with sqlite3.connect(self.db_path) as conn: + conn.execute("PRAGMA journal_mode = WAL") + conn.execute("PRAGMA busy_timeout = 5000") conn.execute("PRAGMA foreign_keys = ON") self._create_tables(conn) self._create_indexes(conn) @@ -96,14 +101,34 @@ class PromptModel: ) """) + # Create normalized tag tables (junction table pattern) + conn.execute(""" + CREATE TABLE IF NOT EXISTS tags ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + name TEXT NOT NULL UNIQUE + ) + """) + conn.execute(""" + CREATE TABLE IF NOT EXISTS prompt_tags ( + prompt_id INTEGER NOT NULL, + tag_id INTEGER NOT NULL, + PRIMARY KEY (prompt_id, tag_id), + FOREIGN KEY (prompt_id) REFERENCES prompts(id) ON DELETE CASCADE, + FOREIGN KEY (tag_id) REFERENCES tags(id) ON DELETE CASCADE + ) + """) + # Add unique constraint to existing databases (migration) self._migrate_add_unique_constraint(conn) - + # Check if we need to migrate from old schema with workflow_name self._migrate_workflow_name_removal(conn) - + # Fix foreign key data type mismatch self._migrate_foreign_key_types(conn) + + # Migrate JSON tags to normalized junction tables + self._migrate_json_tags_to_junction(conn) def _create_indexes(self, conn: sqlite3.Connection) -> None: """ @@ -127,6 +152,8 @@ class PromptModel: "CREATE INDEX IF NOT EXISTS idx_prompt_images ON generated_images(prompt_id)", "CREATE INDEX IF NOT EXISTS idx_image_path ON generated_images(image_path)", "CREATE INDEX IF NOT EXISTS idx_generation_time ON generated_images(generation_time)", + "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: @@ -134,15 +161,26 @@ class PromptModel: def get_connection(self) -> sqlite3.Connection: """ - Get a database connection with proper configuration. - + Get the persistent database connection. + + Returns a thread-safe, reusable connection protected by a lock. + The connection is created on first call and reused thereafter. + Callers use it as a context manager; the lock is held for the + duration of the ``with`` block. + Returns: sqlite3.Connection: Configured database connection """ - conn = sqlite3.connect(self.db_path) - conn.row_factory = sqlite3.Row # Enable dict-like access to rows - conn.execute("PRAGMA foreign_keys = ON") - return conn + with self._conn_lock: + if self._conn is None: + 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: """ @@ -343,6 +381,52 @@ class PromptModel: self.logger.error(f"UNIQUE constraint migration error: {e}") # Continue with existing schema if migration fails + def _migrate_json_tags_to_junction(self, conn: sqlite3.Connection) -> None: + """ + Populate tags and prompt_tags tables from legacy JSON tags column. + + Runs once: skips if the tags table already has data. Uses json_each() + to extract tag names from the JSON arrays stored in prompts.tags. + """ + try: + cursor = conn.execute("SELECT COUNT(*) FROM tags") + if cursor.fetchone()[0] > 0: + return # Already migrated + + cursor = conn.execute( + "SELECT COUNT(*) FROM prompts " + "WHERE tags IS NOT NULL AND tags != '' AND tags != '[]'" + ) + if cursor.fetchone()[0] == 0: + return # No tags to migrate + + self.logger.info("Migrating JSON tags to normalized junction tables") + + # Insert all unique tag names + conn.execute( + "INSERT OR IGNORE INTO tags (name) " + "SELECT DISTINCT je.value FROM prompts, json_each(prompts.tags) AS je " + "WHERE prompts.tags IS NOT NULL AND prompts.tags != '' AND prompts.tags != '[]'" + ) + + # Populate junction table + conn.execute( + "INSERT OR IGNORE INTO prompt_tags (prompt_id, tag_id) " + "SELECT p.id, t.id " + "FROM prompts p, json_each(p.tags) AS je " + "JOIN tags t ON t.name = je.value " + "WHERE p.tags IS NOT NULL AND p.tags != '' AND p.tags != '[]'" + ) + + tag_count = conn.execute("SELECT COUNT(*) FROM tags").fetchone()[0] + link_count = conn.execute("SELECT COUNT(*) FROM prompt_tags").fetchone()[0] + self.logger.info( + f"Tag migration complete: {tag_count} unique tags, {link_count} prompt-tag links" + ) + + except Exception as e: + self.logger.error(f"Tag junction migration error: {e}") + def migrate_database(self) -> None: """ Apply any pending database migrations. diff --git a/database/operations.py b/database/operations.py index f98a3b4..440eaee 100644 --- a/database/operations.py +++ b/database/operations.py @@ -20,9 +20,17 @@ except ImportError: from utils.logging_config import get_logger +# Subquery to fetch tags from junction table, embedded in SELECT statements +TAG_SUBQUERY = ( + "(SELECT GROUP_CONCAT(t.name, '|||') " + "FROM prompt_tags pt JOIN tags t ON pt.tag_id = t.id " + "WHERE pt.prompt_id = prompts.id) AS _tag_list" +) + + class PromptDatabase: """Database operations class for managing prompts.""" - + def __init__(self, db_path: str = "prompts.db"): """ Initialize the database operations. @@ -68,21 +76,18 @@ class PromptDatabase: if rating is not None and (rating < 1 or rating > 5): raise ValueError("Rating must be between 1 and 5") - tags_json = json.dumps(tags) if tags else None - 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( """ INSERT INTO prompts ( text, category, tags, rating, notes, hash, created_at, updated_at - ) VALUES (?, ?, ?, ?, ?, ?, ?, ?) + ) VALUES (?, ?, NULL, ?, ?, ?, ?, ?) """, ( text.strip(), category, - tags_json, rating, notes, prompt_hash, @@ -90,8 +95,10 @@ class PromptDatabase: datetime.datetime.now(datetime.timezone.utc).isoformat() ) ) - conn.commit() prompt_id = cursor.lastrowid + if tags: + self._sync_prompt_tags(conn, prompt_id, tags) + conn.commit() self.logger.debug(f"Successfully saved prompt with ID: {prompt_id}") return prompt_id @@ -107,24 +114,24 @@ class PromptDatabase: """ with self.model.get_connection() as conn: cursor = conn.execute( - "SELECT * 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 - + def get_prompt_by_hash(self, prompt_hash: str) -> Optional[Dict[str, Any]]: """ Get a prompt by its hash. - + Args: prompt_hash: SHA256 hash of the prompt - + Returns: Dict containing prompt data or None if not found """ with self.model.get_connection() as conn: cursor = conn.execute( - "SELECT * 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 @@ -158,49 +165,58 @@ class PromptDatabase: Returns: List of dictionaries containing prompt data """ - query_parts = ["SELECT * FROM prompts WHERE 1=1"] + query_parts = [f"SELECT prompts.*, {TAG_SUBQUERY} FROM prompts WHERE 1=1"] params = [] - + if text: query_parts.append("AND text LIKE ?") params.append(f"%{text}%") - + if category: query_parts.append("AND category = ?") params.append(category) - + if tags: for tag in tags: - query_parts.append("AND tags LIKE ?") - params.append(f"%{tag}%") - + query_parts.append( + "AND prompts.id IN (" + " SELECT pt.prompt_id FROM prompt_tags pt" + " JOIN tags t ON pt.tag_id = t.id WHERE t.name = ?)" + ) + params.append(tag) + if rating_min is not None: query_parts.append("AND rating >= ?") params.append(rating_min) - + if rating_max is not None: query_parts.append("AND rating <= ?") params.append(rating_max) - + if date_from: query_parts.append("AND created_at >= ?") params.append(date_from) - + if date_to: query_parts.append("AND created_at <= ?") params.append(date_to) - - + query_parts.append("ORDER BY created_at DESC LIMIT ? OFFSET ?") params.extend([limit, offset]) - + query = " ".join(query_parts) with self.model.get_connection() as conn: cursor = conn.execute(query, params) rows = cursor.fetchall() - return [self._row_to_dict(row) for row in rows] - + prompts = [self._row_to_dict(row) for row in rows] + + if prompts: + prompt_ids = [p["id"] for p in prompts] + self._attach_preview_images(conn, prompts, prompt_ids) + + return prompts + def get_recent_prompts(self, limit: int = 10, offset: int = 0) -> Dict[str, Any]: """ Get the most recent prompts with pagination support. @@ -214,17 +230,22 @@ class PromptDatabase: """ with self.model.get_connection() as conn: # Get total count - cursor = conn.execute("SELECT COUNT(*) as total FROM prompts") - total_count = cursor.fetchone()["total"] + 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( - "SELECT * FROM prompts ORDER BY created_at DESC LIMIT ? OFFSET ?", + f"SELECT prompts.*, {TAG_SUBQUERY} FROM prompts ORDER BY created_at DESC LIMIT ? OFFSET ?", (limit, offset) ) rows = cursor.fetchall() prompts = [self._row_to_dict(row) for row in rows] - + + if prompts: + prompt_ids = [p["id"] for p in prompts] + self._attach_preview_images(conn, prompts, prompt_ids) + return { 'prompts': prompts, 'total': total_count, @@ -248,7 +269,7 @@ class PromptDatabase: """ with self.model.get_connection() as conn: cursor = conn.execute( - "SELECT * FROM prompts WHERE category = ? ORDER BY created_at DESC LIMIT ?", + f"SELECT prompts.*, {TAG_SUBQUERY} FROM prompts WHERE category = ? ORDER BY created_at DESC LIMIT ?", (category, limit) ) rows = cursor.fetchall() @@ -266,10 +287,10 @@ class PromptDatabase: """ with self.model.get_connection() as conn: cursor = conn.execute( - """ - SELECT * FROM prompts - WHERE rating IS NOT NULL - ORDER BY rating DESC, created_at DESC + f""" + SELECT prompts.*, {TAG_SUBQUERY} FROM prompts + WHERE rating IS NOT NULL + ORDER BY rating DESC, created_at DESC LIMIT ? """, (limit,) @@ -300,39 +321,39 @@ class PromptDatabase: """ updates = [] params = [] - + if category is not None: updates.append("category = ?") params.append(category) - - if tags is not None: - updates.append("tags = ?") - params.append(json.dumps(tags)) - + if rating is not None: if rating < 1 or rating > 5: raise ValueError("Rating must be between 1 and 5") updates.append("rating = ?") params.append(rating) - + if notes is not None: updates.append("notes = ?") params.append(notes) - - - if not updates: + + if not updates and tags is None: return False - - updates.append("updated_at = ?") - params.append(datetime.datetime.now(datetime.timezone.utc).isoformat()) - params.append(prompt_id) - - query = f"UPDATE prompts SET {', '.join(updates)} WHERE id = ?" - + with self.model.get_connection() as conn: - cursor = conn.execute(query, params) + if updates: + updates.append("updated_at = ?") + params.append(datetime.datetime.now(datetime.timezone.utc).isoformat()) + params.append(prompt_id) + query = f"UPDATE prompts SET {', '.join(updates)} WHERE id = ?" + cursor = conn.execute(query, params) + if tags is not None: + 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), + ) conn.commit() - return cursor.rowcount > 0 + return True def delete_prompt(self, prompt_id: int) -> bool: """ @@ -367,28 +388,19 @@ class PromptDatabase: def get_all_tags(self) -> List[str]: """ - Get all unique tags from the database. + Get all unique tags that are in use (linked to at least one prompt). Returns: - List of tag names + Sorted list of tag names """ - all_tags = set() - with self.model.get_connection() as conn: - cursor = conn.execute("SELECT tags FROM prompts WHERE tags IS NOT NULL") - for row in cursor.fetchall(): - try: - tags = json.loads(row['tags']) - if isinstance(tags, list): - all_tags.update(tags) - elif isinstance(tags, str): - # Handle corrupted data (comma-separated string stored as JSON) - parsed_tags = [t.strip() for t in tags.split(',') if t.strip()] - all_tags.update(parsed_tags) - except (json.JSONDecodeError, TypeError): - continue - - return sorted(list(all_tags)) + cursor = conn.execute( + "SELECT DISTINCT t.name AS tag " + "FROM prompt_tags pt " + "JOIN tags t ON pt.tag_id = t.id " + "ORDER BY t.name" + ) + return [row["tag"] for row in cursor.fetchall()] def get_tags_with_counts( self, @@ -398,7 +410,7 @@ class PromptDatabase: sort: str = "alpha_asc" ) -> Dict[str, Any]: """ - Get all unique tags with their usage counts, with search/sort/pagination. + Get all unique tags with their usage counts via junction table. Args: limit: Maximum number of tags to return @@ -409,51 +421,47 @@ class PromptDatabase: Returns: Dict with tags list, total count, and pagination info """ - tag_counts: Dict[str, int] = {} + search_clause = "" + params: list = [] + if search: + search_clause = "HAVING t.name LIKE ?" + params.append(f"%{search}%") + + sort_map = { + "alpha_desc": "tag COLLATE NOCASE DESC", + "count_desc": "count DESC, tag COLLATE NOCASE ASC", + "count_asc": "count ASC, tag COLLATE NOCASE ASC", + } + order = sort_map.get(sort, "tag COLLATE NOCASE ASC") with self.model.get_connection() as conn: - cursor = conn.execute( - "SELECT tags FROM prompts WHERE tags IS NOT NULL AND tags != '' AND tags != '[]'" + count_sql = ( + "SELECT COUNT(*) as total FROM (" + " SELECT t.name AS tag, COUNT(*) AS count" + " FROM prompt_tags pt JOIN tags t ON pt.tag_id = t.id" + f" GROUP BY t.name {search_clause}" + ")" ) - for row in cursor.fetchall(): - try: - tags = json.loads(row['tags']) - if isinstance(tags, list): - for tag in tags: - if isinstance(tag, str) and tag.strip(): - tag_counts[tag.strip()] = tag_counts.get(tag.strip(), 0) + 1 - elif isinstance(tags, str): - for t in tags.split(','): - t = t.strip() - if t: - tag_counts[t] = tag_counts.get(t, 0) + 1 - except (json.JSONDecodeError, TypeError): - continue + cursor = conn.execute(count_sql, params) + row = cursor.fetchone() + total = (row[0] if row else 0) or 0 - # Apply search filter - if search: - search_lower = search.lower() - tag_counts = {k: v for k, v in tag_counts.items() if search_lower in k.lower()} - - # Sort - if sort == "alpha_desc": - sorted_tags = sorted(tag_counts.items(), key=lambda x: x[0].lower(), reverse=True) - elif sort == "count_desc": - sorted_tags = sorted(tag_counts.items(), key=lambda x: (-x[1], x[0].lower())) - elif sort == "count_asc": - sorted_tags = sorted(tag_counts.items(), key=lambda x: (x[1], x[0].lower())) - else: # alpha_asc - sorted_tags = sorted(tag_counts.items(), key=lambda x: x[0].lower()) - - total = len(sorted_tags) - paginated = sorted_tags[offset:offset + limit] + data_sql = ( + "SELECT t.name AS tag, COUNT(*) AS count" + " FROM prompt_tags pt JOIN tags t ON pt.tag_id = t.id" + f" GROUP BY t.name {search_clause}" + f" ORDER BY {order}" + " LIMIT ? OFFSET ?" + ) + cursor = conn.execute(data_sql, params + [limit, offset]) + tags = [{"name": row["tag"], "count": row["count"]} for row in cursor.fetchall()] return { - 'tags': [{'name': name, 'count': count} for name, count in paginated], - 'total': total, - 'limit': limit, - 'offset': offset, - 'has_more': (offset + limit) < total + "tags": tags, + "total": total, + "limit": limit, + "offset": offset, + "has_more": (offset + limit) < total, } def get_prompts_by_tags( @@ -466,6 +474,9 @@ class PromptDatabase: """ Get prompts that match the given tags with AND/OR filtering. + Uses junction table for precise tag matching and a batched image + fetch to avoid N+1 queries. + Args: tags: List of tag names to filter by mode: 'and' (must have all tags) or 'or' (must have any tag) @@ -479,64 +490,61 @@ class PromptDatabase: return {'prompts': [], 'total': 0, 'limit': limit, 'offset': offset, 'has_more': False} with self.model.get_connection() as conn: - # Build tag filter conditions using quoted strings to prevent partial matches - tag_conditions = [] - params = [] - for tag in tags: - tag_conditions.append('tags LIKE ?') - params.append(f'%"{tag}"%') + placeholders = ",".join(["?"] * len(tags)) + if mode == "and": + where_clause = ( + "prompts.id IN (" + " SELECT pt.prompt_id FROM prompt_tags pt" + " JOIN tags t ON pt.tag_id = t.id" + f" WHERE t.name IN ({placeholders})" + " GROUP BY pt.prompt_id" + f" HAVING COUNT(DISTINCT t.name) = ?" + ")" + ) + tag_params = list(tags) + [len(tags)] + else: + where_clause = ( + "prompts.id IN (" + " SELECT pt.prompt_id FROM prompt_tags pt" + " JOIN tags t ON pt.tag_id = t.id" + f" WHERE t.name IN ({placeholders})" + ")" + ) + tag_params = list(tags) - joiner = ' AND ' if mode == 'and' else ' OR ' - where_clause = f"tags IS NOT NULL AND ({joiner.join(tag_conditions)})" + cursor = conn.execute( + f"SELECT COUNT(*) FROM prompts WHERE {where_clause}", + tag_params, + ) + row = cursor.fetchone() + total = (row[0] if row else 0) or 0 - # Get total count - count_query = f"SELECT COUNT(*) as total FROM prompts WHERE {where_clause}" - cursor = conn.execute(count_query, params) - total = cursor.fetchone()['total'] - - # Get paginated results - data_query = f"SELECT * FROM prompts WHERE {where_clause} ORDER BY created_at DESC LIMIT ? OFFSET ?" - cursor = conn.execute(data_query, params + [limit, offset]) + cursor = conn.execute( + f"SELECT prompts.*, {TAG_SUBQUERY} FROM prompts WHERE {where_clause} " + "ORDER BY created_at DESC LIMIT ? OFFSET ?", + tag_params + [limit, offset], + ) rows = cursor.fetchall() + prompts = [self._row_to_dict(row) for row in rows] - prompts = [] - for row in rows: - prompt = self._row_to_dict(row) - - # Get first 3 images (lightweight - skip heavy JSON fields) - img_cursor = conn.execute( - """SELECT id, prompt_id, image_path, filename, generation_time, width, height, format - FROM generated_images - WHERE prompt_id = ? - ORDER BY generation_time DESC - LIMIT 3""", - (prompt['id'],) - ) - images = [dict(img_row) for img_row in img_cursor.fetchall()] - - # Get total image count - count_cursor = conn.execute( - "SELECT COUNT(*) as cnt FROM generated_images WHERE prompt_id = ?", - (prompt['id'],) - ) - image_count = count_cursor.fetchone()['cnt'] - - prompt['images'] = images - prompt['image_count'] = image_count - prompts.append(prompt) + if prompts: + prompt_ids = [p["id"] for p in prompts] + 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 rename_tag_all_prompts(self, old_name: str, new_name: str) -> Dict[str, Any]: """ Rename a tag across all prompts that use it. + With junction tables this is a single UPDATE or a merge operation. + Args: old_name: Current tag name new_name: New tag name @@ -551,40 +559,49 @@ class PromptDatabase: old_name = old_name.strip() new_name = new_name.strip() - affected = 0 - skipped = 0 with self.model.get_connection() as conn: + 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} + + old_tag_id = old_tag['id'] + + # Count affected prompts before the operation cursor = conn.execute( - "SELECT id, tags FROM prompts WHERE tags LIKE ?", - (f'%"{old_name}"%',) + "SELECT COUNT(*) as c FROM prompt_tags WHERE tag_id = ?", (old_tag_id,) ) - for row in cursor.fetchall(): - try: - tags = json.loads(row['tags']) - if not isinstance(tags, list) or old_name not in tags: - continue - tags.remove(old_name) - if new_name not in tags: - tags.append(new_name) - conn.execute( - "UPDATE prompts SET tags = ?, updated_at = ? WHERE id = ?", - (json.dumps(tags), datetime.datetime.now(datetime.timezone.utc).isoformat(), row['id']) - ) - affected += 1 - except (json.JSONDecodeError, TypeError) as e: - skipped += 1 - self.logger.warning(f"Skipped prompt {row['id']} during tag rename: {e}") - continue + 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'] + conn.execute( + "UPDATE OR IGNORE prompt_tags SET tag_id = ? WHERE tag_id = ?", + (new_tag_id, old_tag_id), + ) + # Delete remaining links (conflicts = prompt already had the target tag) + conn.execute("DELETE FROM prompt_tags WHERE tag_id = ?", (old_tag_id,)) + 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.commit() - self.logger.info(f"Renamed tag '{old_name}' -> '{new_name}' in {affected} prompts (skipped {skipped})") - return {'success': True, 'affected_count': affected, 'skipped_count': skipped} + 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]: """ Remove a tag from all prompts that use it. + With junction tables, deleting from the tags table cascades to prompt_tags. + Args: tag_name: Tag name to remove @@ -595,38 +612,32 @@ class PromptDatabase: raise ValueError("Tag name cannot be empty") tag_name = tag_name.strip() - affected = 0 - skipped = 0 with self.model.get_connection() as conn: + 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} + + tag_id = tag_row['id'] cursor = conn.execute( - "SELECT id, tags FROM prompts WHERE tags LIKE ?", - (f'%"{tag_name}"%',) + "SELECT COUNT(*) as c FROM prompt_tags WHERE tag_id = ?", (tag_id,) ) - for row in cursor.fetchall(): - try: - tags = json.loads(row['tags']) - if not isinstance(tags, list) or tag_name not in tags: - continue - tags.remove(tag_name) - conn.execute( - "UPDATE prompts SET tags = ?, updated_at = ? WHERE id = ?", - (json.dumps(tags), datetime.datetime.now(datetime.timezone.utc).isoformat(), row['id']) - ) - affected += 1 - except (json.JSONDecodeError, TypeError) as e: - skipped += 1 - self.logger.warning(f"Skipped prompt {row['id']} during tag delete: {e}") - continue + 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 (skipped {skipped})") - return {'success': True, 'affected_count': affected, 'skipped_count': skipped} + self.logger.info(f"Deleted tag '{tag_name}' from {affected} prompts") + return {'success': True, 'affected_count': affected, 'skipped_count': 0} def merge_tags(self, source_tags: List[str], target_tag: str) -> Dict[str, Any]: """ Merge one or more source tags into a target tag across all prompts. + With junction tables, this is UPDATE + DELETE per source tag. + Args: source_tags: List of tag names to merge from target_tag: Tag name to merge into @@ -643,39 +654,43 @@ class PromptDatabase: source_tags = [t.strip() for t in source_tags if t.strip()] affected = 0 tags_merged = 0 - skipped = 0 with self.model.get_connection() as conn: + # 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'] + 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'] + cursor = conn.execute( - "SELECT id, tags FROM prompts WHERE tags LIKE ?", - (f'%"{src_tag}"%',) + "SELECT COUNT(*) as c FROM prompt_tags WHERE tag_id = ?", (src_id,) ) - src_affected = 0 - for row in cursor.fetchall(): - try: - tags = json.loads(row['tags']) - if not isinstance(tags, list) or src_tag not in tags: - continue - tags.remove(src_tag) - if target_tag not in tags: - tags.append(target_tag) - conn.execute( - "UPDATE prompts SET tags = ?, updated_at = ? WHERE id = ?", - (json.dumps(tags), datetime.datetime.now(datetime.timezone.utc).isoformat(), row['id']) - ) - src_affected += 1 - except (json.JSONDecodeError, TypeError) as e: - skipped += 1 - self.logger.warning(f"Skipped prompt {row['id']} during tag merge: {e}") - continue - if src_affected > 0: + src_count = cursor.fetchone()['c'] + + if src_count > 0: + # Move links from source to target (ignore conflicts) + conn.execute( + "UPDATE OR IGNORE prompt_tags SET tag_id = ? WHERE tag_id = ?", + (target_id, src_id), + ) + # Delete remaining (conflicts = prompt already had target tag) + conn.execute("DELETE FROM prompt_tags WHERE tag_id = ?", (src_id,)) tags_merged += 1 - affected += src_affected + affected += src_count + + # Delete the source tag + conn.execute("DELETE FROM tags WHERE id = ?", (src_id,)) + conn.commit() - self.logger.info(f"Merged {tags_merged} tags into '{target_tag}', affected {affected} prompts (skipped {skipped})") - return {'success': True, 'affected_count': affected, 'tags_merged': tags_merged, 'skipped_count': skipped} + 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: """ @@ -686,7 +701,8 @@ class PromptDatabase: """ with self.model.get_connection() as conn: cursor = conn.execute( - "SELECT COUNT(*) as total FROM prompts WHERE tags IS NULL OR tags = '' OR tags = '[]'" + "SELECT COUNT(*) as total FROM prompts " + "WHERE NOT EXISTS (SELECT 1 FROM prompt_tags WHERE prompt_id = prompts.id)" ) return cursor.fetchone()['total'] @@ -694,6 +710,8 @@ class PromptDatabase: """ Get prompts that have no tags, with pagination. + Uses batched image fetch to avoid N+1 queries. + Args: limit: Maximum number of prompts to return offset: Number of prompts to skip @@ -702,44 +720,105 @@ class PromptDatabase: Dict with prompts list, total count, and pagination """ with self.model.get_connection() as conn: - where = "tags IS NULL OR tags = '' OR tags = '[]'" + 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 * FROM prompts WHERE {where} ORDER BY created_at DESC LIMIT ? OFFSET ?", - (limit, offset) + f"SELECT prompts.*, {TAG_SUBQUERY} FROM prompts WHERE {where} " + "ORDER BY created_at DESC LIMIT ? OFFSET ?", + (limit, offset), ) rows = cursor.fetchall() + prompts = [self._row_to_dict(row) for row in rows] - prompts = [] - for row in rows: - prompt = self._row_to_dict(row) - img_cursor = conn.execute( - """SELECT id, prompt_id, image_path, filename, generation_time, width, height, format - FROM generated_images WHERE prompt_id = ? - ORDER BY generation_time DESC LIMIT 3""", - (prompt['id'],) - ) - images = [dict(img_row) for img_row in img_cursor.fetchall()] - count_cursor = conn.execute( - "SELECT COUNT(*) as cnt FROM generated_images WHERE prompt_id = ?", - (prompt['id'],) - ) - image_count = count_cursor.fetchone()['cnt'] - prompt['images'] = images - prompt['image_count'] = image_count - prompts.append(prompt) + if prompts: + prompt_ids = [p["id"] for p in prompts] + self._attach_preview_images(conn, prompts, prompt_ids) return { 'prompts': prompts, 'total': total, 'limit': limit, 'offset': offset, - 'has_more': (offset + limit) < total + 'has_more': (offset + limit) < total, } + def _attach_preview_images( + self, + conn: sqlite3.Connection, + prompts: List[Dict[str, Any]], + prompt_ids: List[int], + ) -> None: + """Batch-fetch up to 3 preview images + counts for a list of prompts.""" + id_placeholders = ",".join(["?"] * len(prompt_ids)) + + img_cursor = conn.execute( + f"""SELECT * FROM ( + SELECT id, prompt_id, image_path, filename, + generation_time, width, height, format, + ROW_NUMBER() OVER ( + PARTITION BY prompt_id ORDER BY generation_time DESC + ) AS rn + FROM generated_images + WHERE prompt_id IN ({id_placeholders}) + ) WHERE rn <= 3""", + prompt_ids, + ) + images_by_prompt: Dict[int, list] = {} + for img_row in img_cursor.fetchall(): + pid = img_row["prompt_id"] + images_by_prompt.setdefault(pid, []).append(dict(img_row)) + + cnt_cursor = conn.execute( + f"SELECT prompt_id, COUNT(*) as cnt FROM generated_images " + 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()} + + 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]: + """Ensure tag names exist in tags table, return name->id mapping.""" + if not tag_names: + return {} + for name in tag_names: + conn.execute("INSERT OR IGNORE INTO tags (name) VALUES (?)", (name,)) + placeholders = ",".join(["?"] * len(tag_names)) + cursor = conn.execute( + f"SELECT id, name FROM tags WHERE name IN ({placeholders})", + list(tag_names), + ) + 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: + """Replace all junction table entries for a prompt.""" + conn.execute("DELETE FROM prompt_tags WHERE prompt_id = ?", (prompt_id,)) + if tags: + tag_map = self._ensure_tags(conn, tags) + for tag_name in tags: + tag_id = tag_map.get(tag_name) + if tag_id: + conn.execute( + "INSERT OR IGNORE INTO prompt_tags (prompt_id, tag_id) VALUES (?, ?)", + (prompt_id, tag_id), + ) + + def _get_prompt_tags(self, conn: sqlite3.Connection, prompt_id: int) -> List[str]: + """Get tag names for a prompt from junction table.""" + cursor = conn.execute( + "SELECT t.name FROM prompt_tags pt " + "JOIN tags t ON pt.tag_id = t.id " + "WHERE pt.prompt_id = ?", + (prompt_id,), + ) + return [row['name'] for row in cursor.fetchall()] + def _row_to_dict(self, row: sqlite3.Row) -> Dict[str, Any]: """ Convert a database row to a dictionary with parsed JSON fields. @@ -752,14 +831,21 @@ class PromptDatabase: """ data = dict(row) - # Parse tags JSON - if data.get('tags'): + # 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: + # TAG_SUBQUERY present but NULL (no tags) + data['tags'] = [] + del data['_tag_list'] + elif data.get('tags'): + # Fallback: parse legacy JSON column try: parsed = json.loads(data['tags']) if isinstance(parsed, list): data['tags'] = parsed elif isinstance(parsed, str): - # Handle corrupted data (comma-separated string stored as JSON) data['tags'] = [t.strip() for t in parsed.split(',') if t.strip()] else: data['tags'] = [] @@ -848,24 +934,15 @@ class PromptDatabase: # Get full details for all prompts in this duplicate group prompts = [] for prompt_id in ids: - cursor = conn.execute(""" - SELECT id, text, category, rating, tags, created_at, updated_at - FROM prompts - WHERE id = ? - """, (int(prompt_id),)) - + cursor = conn.execute( + "SELECT id, text, category, rating, created_at, updated_at " + "FROM prompts WHERE id = ?", + (int(prompt_id),), + ) prompt_data = cursor.fetchone() if prompt_data: prompt_dict = dict(prompt_data) - # Parse tags if they exist - if prompt_dict['tags']: - try: - import json - prompt_dict['tags'] = json.loads(prompt_dict['tags']) - except: - prompt_dict['tags'] = [] - else: - prompt_dict['tags'] = [] + prompt_dict['tags'] = self._get_prompt_tags(conn, int(prompt_id)) prompts.append(prompt_dict) if prompts: @@ -1082,39 +1159,43 @@ class PromptDatabase: ) return [self._image_row_to_dict(row) for row in cursor.fetchall()] - def get_all_images(self) -> List[Dict[str, Any]]: + def get_all_images( + self, + limit: int = 0, + offset: int = 0, + ) -> List[Dict[str, Any]]: """ - Get all generated images with their linked prompts. + Get generated images with their linked prompts. + + Args: + limit: Maximum number of images to return (0 = all). + offset: Number of images to skip. Returns: - List of all image records with prompt text and tags + List of image records with prompt text and tags """ + sql = ( + "SELECT gi.*, p.text as prompt_text, " + "(SELECT GROUP_CONCAT(t.name, '|||') FROM prompt_tags pt " + "JOIN tags t ON pt.tag_id = t.id WHERE pt.prompt_id = p.id) AS _prompt_tags_list " + "FROM generated_images gi " + "INNER JOIN prompts p ON gi.prompt_id = p.id " + "WHERE gi.image_path IS NOT NULL AND gi.image_path != '' " + "ORDER BY gi.generation_time DESC" + ) + params: list = [] + if limit > 0: + sql += " LIMIT ? OFFSET ?" + params = [limit, offset] + with self.model.get_connection() as conn: - cursor = conn.execute( - """ - SELECT gi.*, p.text as prompt_text, p.tags as prompt_tags - FROM generated_images gi - INNER JOIN prompts p ON gi.prompt_id = p.id - WHERE gi.image_path IS NOT NULL AND gi.image_path != '' - ORDER BY gi.generation_time DESC - """ - ) + cursor = conn.execute(sql, params) rows = cursor.fetchall() result = [] for row in rows: data = self._image_row_to_dict(row) - # Parse prompt_tags JSON - if row['prompt_tags']: - try: - tags = json.loads(row['prompt_tags']) - if isinstance(tags, list): - data['prompt_tags'] = tags - elif isinstance(tags, str): - data['prompt_tags'] = [t.strip() for t in tags.split(',') if t.strip()] - else: - data['prompt_tags'] = [] - except (json.JSONDecodeError, TypeError): - data['prompt_tags'] = [] + if row['_prompt_tags_list']: + data['prompt_tags'] = row['_prompt_tags_list'].split('|||') else: data['prompt_tags'] = [] result.append(data) @@ -1233,54 +1314,47 @@ class PromptDatabase: Dict containing merged metadata """ try: - # Get primary prompt metadata - cursor = conn.execute("SELECT category, tags, rating, notes FROM prompts WHERE id = ?", (primary_id,)) + cursor = conn.execute( + "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': json.loads(primary_data['tags']) if primary_data['tags'] else [], + 'tags': self._get_prompt_tags(conn, primary_id), 'rating': primary_data['rating'], - 'notes': primary_data['notes'] or '' + 'notes': primary_data['notes'] or '', } - - # Merge data from duplicates + for dup_id in duplicate_ids: - cursor = conn.execute("SELECT category, tags, rating, notes FROM prompts WHERE id = ?", (dup_id,)) + cursor = conn.execute( + "SELECT category, rating, notes FROM prompts WHERE id = ?", (dup_id,) + ) dup_data = cursor.fetchone() if not dup_data: continue - - # Merge category (keep first non-empty) + if not merged['category'] and dup_data['category']: merged['category'] = dup_data['category'] - - # Merge tags (combine unique tags) - if dup_data['tags']: - try: - dup_tags = json.loads(dup_data['tags']) - if isinstance(dup_tags, list): - for tag in dup_tags: - if tag not in merged['tags']: - merged['tags'].append(tag) - except (json.JSONDecodeError, TypeError): - pass - - # Keep highest rating + + 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 dup_data['rating'] and (not merged['rating'] or dup_data['rating'] > merged['rating']): merged['rating'] = dup_data['rating'] - - # Combine 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'] - + return merged - + except Exception as e: self.logger.error(f"Error merging metadata: {e}", exc_info=True) return {} @@ -1328,24 +1402,295 @@ class PromptDatabase: try: conn.execute( """ - UPDATE prompts - SET category = ?, tags = ?, rating = ?, notes = ?, updated_at = ? + UPDATE prompts + SET category = ?, rating = ?, notes = ?, updated_at = ? WHERE id = ? """, ( merged_metadata.get('category'), - json.dumps(merged_metadata.get('tags', [])), merged_metadata.get('rating'), merged_metadata.get('notes'), datetime.datetime.now(datetime.timezone.utc).isoformat(), - primary_id - ) + 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") - + except Exception as e: self.logger.error(f"Error updating primary prompt metadata: {e}", exc_info=True) + # ------------------------------------------------------------------ + # Methods extracted from api.py raw SQL (Fix 3.1) + # ------------------------------------------------------------------ + + def get_statistics(self) -> Dict[str, Any]: + """Get comprehensive database statistics.""" + with self.model.get_connection() as conn: + cursor = conn.execute("SELECT COUNT(*) FROM prompts") + row = cursor.fetchone() + total_prompts = (row[0] if row else 0) or 0 + + cursor = conn.execute( + "SELECT COUNT(DISTINCT TRIM(category)) FROM prompts " + "WHERE category IS NOT NULL AND TRIM(category) != ''" + ) + row = cursor.fetchone() + total_categories = (row[0] if row else 0) or 0 + + cursor = conn.execute( + "SELECT AVG(rating) FROM prompts WHERE rating IS NOT NULL" + ) + row = cursor.fetchone() + avg_rating = row[0] if row else None + + cursor = conn.execute( + "SELECT COUNT(*) FROM tags t " + "WHERE EXISTS (SELECT 1 FROM prompt_tags pt WHERE pt.tag_id = t.id)" + ) + row = cursor.fetchone() + total_tags = (row[0] if row else 0) or 0 + + cursor = conn.execute("SELECT COUNT(*) FROM generated_images") + row = cursor.fetchone() + total_images = (row[0] if row else 0) or 0 + + cursor = conn.execute( + "SELECT COUNT(DISTINCT prompt_id) FROM generated_images" + ) + row = cursor.fetchone() + images_with_prompts = (row[0] if row else 0) or 0 + + avg = round(avg_rating, 2) if avg_rating else None + return { + "total_prompts": total_prompts, + "total_categories": total_categories, + "average_rating": avg, + "avg_rating": avg, + "total_tags": total_tags, + "total_images": total_images, + "images_with_prompts": images_with_prompts, + } + + def update_prompt_text(self, prompt_id: int, new_text: str) -> bool: + """Update only the text of a prompt. + + Returns: + True if a row was updated. + """ + 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), + ) + conn.commit() + return cursor.rowcount > 0 + + def update_prompt_rating(self, prompt_id: int, rating: int) -> bool: + """Update only the rating of a prompt. + + Returns: + True if a row was updated. + """ + 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), + ) + conn.commit() + return cursor.rowcount > 0 + + def set_prompt_tags(self, prompt_id: int, tags: List[str]) -> None: + """Overwrite the tags list for a prompt via junction table.""" + with self.model.get_connection() as conn: + 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), + ) + conn.commit() + + def bulk_delete_prompts(self, prompt_ids: List[int]) -> int: + """Delete multiple prompts and their associated images. + + Returns: + Number of prompts actually deleted. + """ + count = 0 + with self.model.get_connection() as conn: + for pid in prompt_ids: + conn.execute("DELETE FROM generated_images WHERE prompt_id = ?", (pid,)) + cursor = conn.execute("DELETE FROM prompts WHERE id = ?", (pid,)) + if cursor.rowcount > 0: + count += 1 + conn.commit() + return count + + def bulk_add_tags(self, prompt_ids: List[int], new_tags: List[str]) -> int: + """Add tags to multiple prompts via junction table (skipping duplicates). + + Returns: + Number of prompts that were actually modified. + """ + count = 0 + with self.model.get_connection() as conn: + tag_map = self._ensure_tags(conn, new_tags) + now = datetime.datetime.now(datetime.timezone.utc).isoformat() + for pid in prompt_ids: + cursor = conn.execute("SELECT id FROM prompts WHERE id = ?", (pid,)) + if not cursor.fetchone(): + continue + added = False + for tag_name in new_tags: + tag_id = tag_map.get(tag_name) + if tag_id: + cursor = conn.execute( + "INSERT OR IGNORE INTO prompt_tags (prompt_id, tag_id) VALUES (?, ?)", + (pid, tag_id), + ) + if cursor.rowcount > 0: + added = True + if added: + conn.execute( + "UPDATE prompts SET updated_at = ? WHERE id = ?", (now, pid) + ) + count += 1 + conn.commit() + return count + + def bulk_set_category(self, prompt_ids: List[int], category: str) -> int: + """Set category on multiple prompts. + + Returns: + Number of prompts actually updated. + """ + count = 0 + now = datetime.datetime.now(datetime.timezone.utc).isoformat() + with self.model.get_connection() as conn: + for pid in prompt_ids: + cursor = conn.execute( + "UPDATE prompts SET category = ?, updated_at = ? WHERE id = ?", + (category, now, pid), + ) + if cursor.rowcount > 0: + count += 1 + conn.commit() + return count + + def get_image_prompt_info(self, image_path: str) -> Optional[Dict[str, Any]]: + """Look up prompt + image metadata for a given image path. + + Returns a dict with prompt_id, text, category, tags, rating, notes, + workflow_data, prompt_metadata, generation_time — or None. + """ + with self.model.get_connection() as conn: + cursor = conn.execute( + """SELECT gi.prompt_id, p.text, p.category, p.rating, p.notes, + gi.workflow_data, gi.prompt_metadata, gi.generation_time + 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)}'), + ) + row = cursor.fetchone() + if not row: + return None + tags = self._get_prompt_tags(conn, row[0]) + try: + workflow = json.loads(row[5]) if row[5] else None + except (json.JSONDecodeError, TypeError): + workflow = None + try: + prompt_meta = json.loads(row[6]) if row[6] else None + except (json.JSONDecodeError, TypeError): + prompt_meta = None + return { + "prompt_id": row[0], + "text": row[1], + "category": row[2], + "tags": tags, + "rating": row[3], + "notes": row[4], + "workflow_data": workflow, + "prompt_metadata": prompt_meta, + "generation_time": row[7], + } + + def get_prompt_id_for_image(self, image_path: str) -> Optional[int]: + """Return the prompt_id linked to an image, or None.""" + with self.model.get_connection() as conn: + cursor = conn.execute( + "SELECT prompt_id FROM generated_images WHERE image_path = ?", + (image_path,), + ) + row = cursor.fetchone() + return row[0] if row else None + + def check_hash_duplicates(self) -> List[Dict[str, Any]]: + """Find prompts sharing the same hash. + + Returns list of dicts with 'hash' and 'count'. + """ + with self.model.get_connection() as conn: + cursor = conn.execute( + "SELECT hash, COUNT(*) as count FROM prompts " + "WHERE hash IS NOT NULL GROUP BY hash HAVING COUNT(*) > 1" + ) + return [dict(row) for row in cursor.fetchall()] + + def prune_orphaned_prompts(self) -> int: + """Delete prompts with no linked images and no __protected__ tag. + + Returns number of prompts removed. + """ + with self.model.get_connection() as conn: + cursor = conn.execute(""" + SELECT p.id FROM prompts p + LEFT JOIN generated_images gi ON p.id = gi.prompt_id + WHERE gi.prompt_id IS NULL + AND NOT EXISTS ( + SELECT 1 FROM prompt_tags pt + JOIN tags t ON pt.tag_id = t.id + WHERE pt.prompt_id = p.id AND t.name = '__protected__' + ) + """) + orphaned = [row["id"] for row in cursor.fetchall()] + if not orphaned: + return 0 + placeholders = ",".join(["?"] * len(orphaned)) + cursor = conn.execute( + f"DELETE FROM prompts WHERE id IN ({placeholders})", orphaned + ) + conn.commit() + return cursor.rowcount + + def check_consistency(self) -> List[str]: + """Run consistency checks and return a list of issue descriptions.""" + issues: List[str] = [] + with self.model.get_connection() as conn: + # Check for orphaned prompt_tags entries + cursor = conn.execute(""" + SELECT pt.prompt_id, pt.tag_id FROM prompt_tags pt + LEFT JOIN prompts p ON pt.prompt_id = p.id + WHERE p.id IS NULL + """) + for ref in cursor.fetchall(): + issues.append( + f"prompt_tags entry references non-existent prompt {ref['prompt_id']}" + ) + + # Check for orphaned image entries + cursor = conn.execute(""" + SELECT gi.id, gi.prompt_id FROM generated_images gi + LEFT JOIN prompts p ON gi.prompt_id = p.id + WHERE p.id IS NULL + """) + for ref in cursor.fetchall(): + issues.append( + f"Image {ref['id']} references non-existent prompt {ref['prompt_id']}" + ) + return issues + def _image_row_to_dict(self, row: sqlite3.Row) -> Dict[str, Any]: """ Convert an image database row to a dictionary with parsed JSON fields. diff --git a/prompt_manager.py b/prompt_manager.py index 8f0e9bb..c3ad1db 100644 --- a/prompt_manager.py +++ b/prompt_manager.py @@ -3,22 +3,8 @@ PromptManager: Main custom node implementation that extends CLIPTextEncode with persistent prompt storage and search capabilities. """ -import datetime -import hashlib -import json -import os import time -import webbrowser -from typing import Any, Dict, List, Optional, Tuple - -# Import logging system -try: - from .utils.logging_config import get_logger -except ImportError: - import sys - - sys.path.insert(0, os.path.dirname(os.path.abspath(__file__))) - from utils.logging_config import get_logger +from typing import Any, Tuple try: from comfy.comfy_types import IO, ComfyNodeABC, InputTypeDict @@ -35,23 +21,16 @@ except ImportError: InputTypeDict = dict try: - from .database.operations import PromptDatabase - from .utils.comfyui_integration import get_comfyui_integration - from .utils.image_monitor import get_image_monitor - from .utils.prompt_tracker import PromptExecutionContext, get_prompt_tracker + from .prompt_manager_base import PromptManagerBase except ImportError: - # For direct imports when not in a package import os import sys sys.path.insert(0, os.path.dirname(os.path.abspath(__file__))) - from database.operations import PromptDatabase - from utils.comfyui_integration import get_comfyui_integration - from utils.image_monitor import get_image_monitor - from utils.prompt_tracker import PromptExecutionContext, get_prompt_tracker + from prompt_manager_base import PromptManagerBase -class PromptManager(ComfyNodeABC): +class PromptManager(PromptManagerBase, ComfyNodeABC): """ A ComfyUI custom node that functions like CLIPTextEncode but adds: - Persistent storage of all prompts in SQLite database @@ -61,18 +40,7 @@ class PromptManager(ComfyNodeABC): """ def __init__(self): - self.logger = get_logger("prompt_manager.node") - self.logger.debug("Initializing PromptManager node") - - self.db = PromptDatabase() - # Use singleton getters to ensure only one tracker/monitor exists - self.prompt_tracker = get_prompt_tracker(self.db) - self.image_monitor = get_image_monitor(self.db, self.prompt_tracker) - self.comfyui_integration = get_comfyui_integration() - - # Start image monitoring automatically - self._start_gallery_system() - self.logger.debug("PromptManager node initialization completed") + super().__init__(logger_name="prompt_manager.node") @classmethod def INPUT_TYPES(cls) -> InputTypeDict: @@ -183,9 +151,6 @@ class PromptManager(ComfyNodeABC): # For database storage, save the original main text with metadata about prepend/append storage_text = text - # Search functionality is now handled by the JavaScript UI - # The search parameters are still available for backend processing if needed - # Validate CLIP model if clip is None: error_msg = ( @@ -210,7 +175,7 @@ class PromptManager(ComfyNodeABC): try: prompt_id = self._save_prompt_to_database( - text=storage_text.strip(), # Always strip whitespace + text=storage_text.strip(), category=category.strip() if category else None, tags=extended_tags if extended_tags else None, ) @@ -218,7 +183,7 @@ class PromptManager(ComfyNodeABC): # Set current prompt for image tracking if prompt_id: execution_id = self.prompt_tracker.set_current_prompt( - prompt_text=encoding_text.strip(), # Use final combined text for tracking + prompt_text=encoding_text.strip(), additional_data={ "category": category.strip() if category else None, "tags": extended_tags, @@ -227,7 +192,7 @@ class PromptManager(ComfyNodeABC): prepend_text.strip() if prepend_text else None ), "append_text": append_text.strip() if append_text else None, - "final_text": encoding_text.strip(), # Store final combined text + "final_text": encoding_text.strip(), }, ) self.logger.debug( @@ -237,7 +202,6 @@ class PromptManager(ComfyNodeABC): except Exception as e: # Log error but don't fail the encoding self.logger.warning(f"Failed to save prompt to database: {e}") - # Already logged above, no need for additional print # Perform standard CLIP text encoding using the combined text self.logger.debug( @@ -247,7 +211,7 @@ class PromptManager(ComfyNodeABC): conditioning = clip.encode_from_tokens_scheduled(tokens) # Register with ComfyUI integration for standard metadata compatibility - node_id = f"promptmanager_{int(time.time() * 1000)}" # Unique node ID + node_id = f"promptmanager_{int(time.time() * 1000)}" self.comfyui_integration.register_prompt( node_id, encoding_text.strip(), @@ -263,228 +227,6 @@ class PromptManager(ComfyNodeABC): self.logger.info(f"CLIP encoding completed, text: {repr(encoding_text)[:80]}") return (conditioning, encoding_text) - def _save_prompt_to_database( - self, text: str, category: Optional[str] = None, tags: Optional[list] = None - ) -> Optional[int]: - """ - Save the prompt to the SQLite database. - - Args: - text: The prompt text - category: Optional category - tags: List of tags - - Returns: - The prompt ID if saved successfully, None otherwise - """ - try: - # Generate hash for duplicate detection - prompt_hash = self._generate_hash(text) - self.logger.debug(f"Generated hash for prompt: {prompt_hash[:16]}...") - - # Check if prompt already exists - existing = self.db.get_prompt_by_hash(prompt_hash) - if existing: - self.logger.info( - f"Found existing prompt with ID {existing['id']}, updating metadata" - ) - # Update metadata if this is a duplicate with new info - if any([category, tags]): - self.db.update_prompt_metadata( - prompt_id=existing["id"], category=category, tags=tags - ) - self.logger.debug("Updated metadata for existing prompt") - return existing["id"] - - # Save new prompt - self.logger.debug( - f"Saving new prompt with category: {category}, tags: {tags}" - ) - prompt_id = self.db.save_prompt( - text=text, category=category, tags=tags, prompt_hash=prompt_hash - ) - - if prompt_id: - self.logger.debug(f"Successfully saved new prompt with ID: {prompt_id}") - else: - self.logger.warning("Failed to save prompt - no ID returned") - - return prompt_id - - except Exception as e: - self.logger.error(f"Error saving prompt to database: {e}") - # Already logged above, no need for additional print - return None - - def _generate_hash(self, text: str) -> str: - """ - Generate SHA256 hash for the prompt text. - - Args: - text: The prompt text to hash - - Returns: - Hexadecimal string representation of the SHA256 hash - """ - # Normalize text for consistent hashing (strip whitespace, normalize case) - normalized_text = text.strip().lower() - return hashlib.sha256(normalized_text.encode("utf-8")).hexdigest() - - def _parse_tags(self, tags_string: str) -> Optional[list]: - """ - Parse comma-separated tags string into a list. - - Args: - tags_string: Comma-separated string of tags - - Returns: - List of parsed tags, or None if no valid tags found - """ - if not tags_string or not tags_string.strip(): - return None - - tags = [tag.strip() for tag in tags_string.split(",") if tag.strip()] - return tags if tags else None - - def _search_prompts(self, search_text: str = "") -> List[Dict[str, Any]]: - """ - Search for past prompts by text content. - - Args: - search_text: Text to search for in prompt database - - Returns: - List of matching prompt dictionaries with metadata - """ - try: - if not search_text or not search_text.strip(): - return [] - - results = self.db.search_prompts( - text=search_text.strip(), - category=None, - tags=None, - rating_min=None, - limit=50, - ) - - return results - - except Exception as e: - self.logger.error(f"Error searching prompts: {e}") - return [] - - def _open_web_interface(self): - """ - Open the web interface in the default browser. - - Attempts to locate and open the web interface HTML file. - Logs warnings if the interface is not properly configured. - """ - try: - # Look for a web interface directory - current_dir = os.path.dirname(os.path.abspath(__file__)) - web_dir = os.path.join(current_dir, "web_interface") - - if os.path.exists(web_dir): - # If web interface exists, try to start it - index_path = os.path.join(web_dir, "index.html") - if os.path.exists(index_path): - webbrowser.open(f"file://{index_path}") - self.logger.info("Web interface opened in browser") - else: - self.logger.warning( - f"Web interface directory found but no index.html. Please check {web_dir} for setup instructions" - ) - else: - self.logger.info( - "Web interface not yet implemented. This feature will open a web-based prompt management interface when the web_interface directory is created." - ) - - except Exception as e: - self.logger.error(f"Error opening web interface: {e}") - - def search_prompts_api(self, search_text: str = "") -> List[Dict[str, Any]]: - """ - API method for JavaScript UI to search prompts. - - Args: - search_text: Text to search for in prompts - - Returns: - List of matching prompt dictionaries - """ - return self._search_prompts(search_text=search_text) - - def get_recent_prompts_api(self, limit: int = 20) -> List[Dict[str, Any]]: - """ - API method for JavaScript UI to get recent prompts. - - Args: - limit: Maximum number of recent prompts to retrieve - - Returns: - List of recent prompt dictionaries ordered by creation time - """ - try: - return self.db.get_recent_prompts(limit=limit) - except Exception as e: - self.logger.error(f"Error getting recent prompts: {e}") - return [] - - def _start_gallery_system(self): - """ - Initialize and start the gallery monitoring system. - - Starts the image monitor which watches for new generated images - and links them to their source prompts in the database. - """ - try: - self.logger.debug("Starting gallery system...") - - # Start image monitoring - self.image_monitor.start_monitoring() - - self.logger.debug("Gallery system started successfully") - - except Exception as e: - self.logger.error(f"Failed to start gallery system: {e}") - self.logger.warning("Gallery features will be disabled") - - def get_gallery_status(self) -> Dict[str, Any]: - """ - Get status of the gallery system. - - Returns: - Dictionary containing status information for image monitor and prompt tracker - """ - return { - "image_monitor": self.image_monitor.get_status(), - "prompt_tracker": self.prompt_tracker.get_status(), - } - - def cleanup_gallery_system(self): - """ - Clean up gallery system resources. - - Stops image monitoring and releases associated resources. - Called automatically during object destruction. - """ - try: - if hasattr(self, "image_monitor"): - self.image_monitor.stop_monitoring() - self.logger.debug("Gallery system cleaned up") - except Exception as e: - self.logger.error(f"Error cleaning up gallery system: {e}") - - def __del__(self): - """ - Cleanup when object is destroyed. - - Ensures proper resource cleanup by stopping the gallery system. - """ - self.cleanup_gallery_system() - @classmethod def IS_CHANGED(cls, clip, text="", category="", tags="", search_text="", prepend_text="", append_text="", **kwargs): @@ -496,8 +238,5 @@ class PromptManager(ComfyNodeABC): """ import hashlib - # Combine all text inputs that affect the conditioning output - # Note: search_text doesn't affect output, so it's excluded combined = f"{text}|{prepend_text}|{append_text}" - return hashlib.sha256(combined.encode()).hexdigest() diff --git a/prompt_manager_base.py b/prompt_manager_base.py new file mode 100644 index 0000000..db46a66 --- /dev/null +++ b/prompt_manager_base.py @@ -0,0 +1,231 @@ +""" +PromptManagerBase: Shared logic for PromptManager node variants. + +Provides database initialization, prompt saving, hashing, tag parsing, +search, gallery system management, and cleanup — extracted from the +duplicate code in prompt_manager.py and prompt_manager_text.py. +""" + +import hashlib +import os +import webbrowser +from typing import Any, Dict, List, Optional + +try: + from .utils.logging_config import get_logger +except ImportError: + import sys + + sys.path.insert(0, os.path.dirname(os.path.abspath(__file__))) + from utils.logging_config import get_logger + +try: + from .database.operations import PromptDatabase + from .utils.comfyui_integration import get_comfyui_integration + from .utils.image_monitor import get_image_monitor + from .utils.prompt_tracker import get_prompt_tracker +except ImportError: + import sys + + sys.path.insert(0, os.path.dirname(os.path.abspath(__file__))) + from database.operations import PromptDatabase + from utils.comfyui_integration import get_comfyui_integration + from utils.image_monitor import get_image_monitor + from utils.prompt_tracker import get_prompt_tracker + + +class PromptManagerBase: + """Mixin providing shared prompt management logic for ComfyUI nodes. + + Handles database connection, prompt saving with deduplication, + tag parsing, search, gallery system lifecycle, and cleanup. + Subclasses only need to define ComfyUI-specific class attributes + (INPUT_TYPES, RETURN_TYPES, FUNCTION) and their execution method. + """ + + def __init__(self, logger_name: str = "prompt_manager.node"): + self.logger = get_logger(logger_name) + self.logger.debug(f"Initializing {self.__class__.__name__} node") + + self.db = PromptDatabase() + self.prompt_tracker = get_prompt_tracker(self.db) + self.image_monitor = get_image_monitor(self.db, self.prompt_tracker) + self.comfyui_integration = get_comfyui_integration() + + self._start_gallery_system() + self.logger.debug(f"{self.__class__.__name__} node initialization completed") + + def _save_prompt_to_database( + self, text: str, category: Optional[str] = None, tags: Optional[list] = None + ) -> Optional[int]: + """Save the prompt to the SQLite database. + + Args: + text: The prompt text + category: Optional category + tags: List of tags + + Returns: + The prompt ID if saved successfully, None otherwise + """ + try: + prompt_hash = self._generate_hash(text) + self.logger.debug(f"Generated hash for prompt: {prompt_hash[:16]}...") + + existing = self.db.get_prompt_by_hash(prompt_hash) + if existing: + self.logger.info( + f"Found existing prompt with ID {existing['id']}, updating metadata" + ) + if any([category, tags]): + self.db.update_prompt_metadata( + prompt_id=existing["id"], category=category, tags=tags + ) + self.logger.debug("Updated metadata for existing prompt") + return existing["id"] + + self.logger.debug( + f"Saving new prompt with category: {category}, tags: {tags}" + ) + prompt_id = self.db.save_prompt( + text=text, category=category, tags=tags, prompt_hash=prompt_hash + ) + + if prompt_id: + self.logger.debug(f"Successfully saved new prompt with ID: {prompt_id}") + else: + self.logger.warning("Failed to save prompt - no ID returned") + + return prompt_id + + except Exception as e: + self.logger.error(f"Error saving prompt to database: {e}") + return None + + def _generate_hash(self, text: str) -> str: + """Generate SHA256 hash for the prompt text. + + Args: + text: The prompt text to hash + + Returns: + Hexadecimal string representation of the SHA256 hash + """ + normalized_text = text.strip().lower() + return hashlib.sha256(normalized_text.encode("utf-8")).hexdigest() + + def _parse_tags(self, tags_string: str) -> Optional[list]: + """Parse comma-separated tags string into a list. + + Args: + tags_string: Comma-separated string of tags + + Returns: + List of parsed tags, or None if no valid tags found + """ + if not tags_string or not tags_string.strip(): + return None + + tags = [tag.strip() for tag in tags_string.split(",") if tag.strip()] + return tags if tags else None + + def _search_prompts(self, search_text: str = "") -> List[Dict[str, Any]]: + """Search for past prompts by text content. + + Args: + search_text: Text to search for in prompt database + + Returns: + List of matching prompt dictionaries with metadata + """ + try: + if not search_text or not search_text.strip(): + return [] + + try: + from .py.config import PromptManagerConfig + max_results = PromptManagerConfig.MAX_SEARCH_RESULTS + except Exception: + max_results = 100 + + results = self.db.search_prompts( + text=search_text.strip(), + category=None, + tags=None, + rating_min=None, + limit=max_results, + ) + + return results + + except Exception as e: + self.logger.error(f"Error searching prompts: {e}") + return [] + + def _open_web_interface(self): + """Open the web interface in the default browser.""" + try: + current_dir = os.path.dirname(os.path.abspath(__file__)) + web_dir = os.path.join(current_dir, "web_interface") + + if os.path.exists(web_dir): + index_path = os.path.join(web_dir, "index.html") + if os.path.exists(index_path): + webbrowser.open(f"file://{index_path}") + self.logger.info("Web interface opened in browser") + else: + self.logger.warning( + f"Web interface directory found but no index.html. " + f"Please check {web_dir} for setup instructions" + ) + else: + self.logger.info( + "Web interface not yet implemented. This feature will open a " + "web-based prompt management interface when the web_interface " + "directory is created." + ) + + except Exception as e: + self.logger.error(f"Error opening web interface: {e}") + + def search_prompts_api(self, search_text: str = "") -> List[Dict[str, Any]]: + """API method for JavaScript UI to search prompts.""" + return self._search_prompts(search_text=search_text) + + def get_recent_prompts_api(self, limit: int = 20) -> List[Dict[str, Any]]: + """API method for JavaScript UI to get recent prompts.""" + try: + return self.db.get_recent_prompts(limit=limit) + except Exception as e: + self.logger.error(f"Error getting recent prompts: {e}") + return [] + + def _start_gallery_system(self): + """Initialize and start the gallery monitoring system.""" + try: + self.logger.debug("Starting gallery system...") + self.image_monitor.start_monitoring() + self.logger.debug("Gallery system started successfully") + except Exception as e: + self.logger.error(f"Failed to start gallery system: {e}") + self.logger.warning("Gallery features will be disabled") + + def get_gallery_status(self) -> Dict[str, Any]: + """Get status of the gallery system.""" + return { + "image_monitor": self.image_monitor.get_status(), + "prompt_tracker": self.prompt_tracker.get_status(), + } + + def cleanup_gallery_system(self): + """Clean up gallery system resources.""" + try: + if hasattr(self, "image_monitor"): + self.image_monitor.stop_monitoring() + self.logger.debug("Gallery system cleaned up") + except Exception as e: + self.logger.error(f"Error cleaning up gallery system: {e}") + + def __del__(self): + """Cleanup when object is destroyed.""" + self.cleanup_gallery_system() diff --git a/prompt_manager_text.py b/prompt_manager_text.py index f4bd6fd..7932bfe 100644 --- a/prompt_manager_text.py +++ b/prompt_manager_text.py @@ -3,22 +3,8 @@ PromptManagerText: A text-only version of PromptManager that outputs STRING without CLIP encoding, while maintaining all database and search features. """ -import datetime -import hashlib -import json -import os import time -import webbrowser -from typing import Any, Dict, List, Optional, Tuple - -# Import logging system -try: - from .utils.logging_config import get_logger -except ImportError: - import sys - - sys.path.insert(0, os.path.dirname(os.path.abspath(__file__))) - from utils.logging_config import get_logger +from typing import Tuple try: from comfy.comfy_types import IO, ComfyNodeABC, InputTypeDict @@ -35,23 +21,16 @@ except ImportError: InputTypeDict = dict try: - from .database.operations import PromptDatabase - from .utils.comfyui_integration import get_comfyui_integration - from .utils.image_monitor import get_image_monitor - from .utils.prompt_tracker import PromptExecutionContext, get_prompt_tracker + from .prompt_manager_base import PromptManagerBase except ImportError: - # For direct imports when not in a package import os import sys sys.path.insert(0, os.path.dirname(os.path.abspath(__file__))) - from database.operations import PromptDatabase - from utils.comfyui_integration import get_comfyui_integration - from utils.image_monitor import get_image_monitor - from utils.prompt_tracker import PromptExecutionContext, get_prompt_tracker + from prompt_manager_base import PromptManagerBase -class PromptManagerText(ComfyNodeABC): +class PromptManagerText(PromptManagerBase, ComfyNodeABC): """ A ComfyUI custom node that provides all PromptManager features but outputs only a STRING without CLIP encoding. Includes: @@ -63,18 +42,7 @@ class PromptManagerText(ComfyNodeABC): """ def __init__(self): - self.logger = get_logger("prompt_manager_text.node") - self.logger.debug("Initializing PromptManagerText node") - - self.db = PromptDatabase() - # Use singleton getters to ensure only one tracker/monitor exists - self.prompt_tracker = get_prompt_tracker(self.db) - self.image_monitor = get_image_monitor(self.db, self.prompt_tracker) - self.comfyui_integration = get_comfyui_integration() - - # Start image monitoring automatically - self._start_gallery_system() - self.logger.debug("PromptManagerText node initialization completed") + super().__init__(logger_name="prompt_manager_text.node") @classmethod def INPUT_TYPES(cls) -> InputTypeDict: @@ -174,9 +142,6 @@ class PromptManagerText(ComfyNodeABC): # For database storage, save the original main text with metadata about prepend/append storage_text = text - # Search functionality is now handled by the JavaScript UI - # The search parameters are still available for backend processing if needed - # Save prompt to database and set execution context for gallery tracking prompt_id = None if storage_text and storage_text.strip(): @@ -191,7 +156,7 @@ class PromptManagerText(ComfyNodeABC): try: prompt_id = self._save_prompt_to_database( - text=storage_text.strip(), # Always strip whitespace + text=storage_text.strip(), category=category.strip() if category else None, tags=extended_tags if extended_tags else None, ) @@ -199,7 +164,7 @@ class PromptManagerText(ComfyNodeABC): # Set current prompt for image tracking if prompt_id: execution_id = self.prompt_tracker.set_current_prompt( - prompt_text=final_text.strip(), # Use final combined text for tracking + prompt_text=final_text.strip(), additional_data={ "category": category.strip() if category else None, "tags": extended_tags, @@ -219,7 +184,7 @@ class PromptManagerText(ComfyNodeABC): self.logger.warning(f"Failed to save prompt to database: {e}") # Register with ComfyUI integration for standard metadata compatibility - node_id = f"promptmanagertext_{int(time.time() * 1000)}" # Unique node ID + node_id = f"promptmanagertext_{int(time.time() * 1000)}" self.comfyui_integration.register_prompt( node_id, final_text.strip(), @@ -235,227 +200,6 @@ class PromptManagerText(ComfyNodeABC): self.logger.debug(f"Text processing completed: {final_text[:100]}...") return (final_text,) - def _save_prompt_to_database( - self, text: str, category: Optional[str] = None, tags: Optional[list] = None - ) -> Optional[int]: - """ - Save the prompt to the SQLite database. - - Args: - text: The prompt text - category: Optional category - tags: List of tags - - Returns: - The prompt ID if saved successfully, None otherwise - """ - try: - # Generate hash for duplicate detection - prompt_hash = self._generate_hash(text) - self.logger.debug(f"Generated hash for prompt: {prompt_hash[:16]}...") - - # Check if prompt already exists - existing = self.db.get_prompt_by_hash(prompt_hash) - if existing: - self.logger.info( - f"Found existing prompt with ID {existing['id']}, updating metadata" - ) - # Update metadata if this is a duplicate with new info - if any([category, tags]): - self.db.update_prompt_metadata( - prompt_id=existing["id"], category=category, tags=tags - ) - self.logger.debug("Updated metadata for existing prompt") - return existing["id"] - - # Save new prompt - self.logger.debug( - f"Saving new prompt with category: {category}, tags: {tags}" - ) - prompt_id = self.db.save_prompt( - text=text, category=category, tags=tags, prompt_hash=prompt_hash - ) - - if prompt_id: - self.logger.debug(f"Successfully saved new prompt with ID: {prompt_id}") - else: - self.logger.warning("Failed to save prompt - no ID returned") - - return prompt_id - - except Exception as e: - self.logger.error(f"Error saving prompt to database: {e}") - return None - - def _generate_hash(self, text: str) -> str: - """ - Generate SHA256 hash for the prompt text. - - Args: - text: The prompt text to hash - - Returns: - Hexadecimal string representation of the SHA256 hash - """ - # Normalize text for consistent hashing (strip whitespace, normalize case) - normalized_text = text.strip().lower() - return hashlib.sha256(normalized_text.encode("utf-8")).hexdigest() - - def _parse_tags(self, tags_string: str) -> Optional[list]: - """ - Parse comma-separated tags string into a list. - - Args: - tags_string: Comma-separated string of tags - - Returns: - List of parsed tags, or None if no valid tags found - """ - if not tags_string or not tags_string.strip(): - return None - - tags = [tag.strip() for tag in tags_string.split(",") if tag.strip()] - return tags if tags else None - - def _search_prompts(self, search_text: str = "") -> List[Dict[str, Any]]: - """ - Search for past prompts by text content. - - Args: - search_text: Text to search for in prompt database - - Returns: - List of matching prompt dictionaries with metadata - """ - try: - if not search_text or not search_text.strip(): - return [] - - results = self.db.search_prompts( - text=search_text.strip(), - category=None, - tags=None, - rating_min=None, - limit=50, - ) - - return results - - except Exception as e: - self.logger.error(f"Error searching prompts: {e}") - return [] - - def _open_web_interface(self): - """ - Open the web interface in the default browser. - - Attempts to locate and open the web interface HTML file. - Logs warnings if the interface is not properly configured. - """ - try: - # Look for a web interface directory - current_dir = os.path.dirname(os.path.abspath(__file__)) - web_dir = os.path.join(current_dir, "web_interface") - - if os.path.exists(web_dir): - # If web interface exists, try to start it - index_path = os.path.join(web_dir, "index.html") - if os.path.exists(index_path): - webbrowser.open(f"file://{index_path}") - self.logger.info("Web interface opened in browser") - else: - self.logger.warning( - f"Web interface directory found but no index.html. Please check {web_dir} for setup instructions" - ) - else: - self.logger.info( - "Web interface not yet implemented. This feature will open a web-based prompt management interface when the web_interface directory is created." - ) - - except Exception as e: - self.logger.error(f"Error opening web interface: {e}") - - def search_prompts_api(self, search_text: str = "") -> List[Dict[str, Any]]: - """ - API method for JavaScript UI to search prompts. - - Args: - search_text: Text to search for in prompts - - Returns: - List of matching prompt dictionaries - """ - return self._search_prompts(search_text=search_text) - - def get_recent_prompts_api(self, limit: int = 20) -> List[Dict[str, Any]]: - """ - API method for JavaScript UI to get recent prompts. - - Args: - limit: Maximum number of recent prompts to retrieve - - Returns: - List of recent prompt dictionaries ordered by creation time - """ - try: - return self.db.get_recent_prompts(limit=limit) - except Exception as e: - self.logger.error(f"Error getting recent prompts: {e}") - return [] - - def _start_gallery_system(self): - """ - Initialize and start the gallery monitoring system. - - Starts the image monitor which watches for new generated images - and links them to their source prompts in the database. - """ - try: - self.logger.debug("Starting gallery system...") - - # Start image monitoring - self.image_monitor.start_monitoring() - - self.logger.debug("Gallery system started successfully") - - except Exception as e: - self.logger.error(f"Failed to start gallery system: {e}") - self.logger.warning("Gallery features will be disabled") - - def get_gallery_status(self) -> Dict[str, Any]: - """ - Get status of the gallery system. - - Returns: - Dictionary containing status information for image monitor and prompt tracker - """ - return { - "image_monitor": self.image_monitor.get_status(), - "prompt_tracker": self.prompt_tracker.get_status(), - } - - def cleanup_gallery_system(self): - """ - Clean up gallery system resources. - - Stops image monitoring and releases associated resources. - Called automatically during object destruction. - """ - try: - if hasattr(self, "image_monitor"): - self.image_monitor.stop_monitoring() - self.logger.debug("Gallery system cleaned up") - except Exception as e: - self.logger.error(f"Error cleaning up gallery system: {e}") - - def __del__(self): - """ - Cleanup when object is destroyed. - - Ensures proper resource cleanup by stopping the gallery system. - """ - self.cleanup_gallery_system() - @classmethod def IS_CHANGED(cls, text="", category="", tags="", search_text="", prepend_text="", append_text="", **kwargs): @@ -467,8 +211,5 @@ class PromptManagerText(ComfyNodeABC): """ import hashlib - # Combine all text inputs that affect the output - # Note: search_text doesn't affect output, so it's excluded combined = f"{text}|{prepend_text}|{append_text}" - return hashlib.sha256(combined.encode()).hexdigest() diff --git a/prompts.db b/prompts.db deleted file mode 100644 index a08e886..0000000 Binary files a/prompts.db and /dev/null differ diff --git a/py/api.py b/py/api.py deleted file mode 100644 index 87e09c2..0000000 --- a/py/api.py +++ /dev/null @@ -1,5293 +0,0 @@ -"""REST API module for ComfyUI PromptManager. - -This module provides a comprehensive REST API for managing prompts, images, and -gallery functionality within ComfyUI. The API handles CRUD operations for prompts, -image gallery management, system administration, logging, and real-time image -monitoring with metadata extraction. - -Key Features: -- Prompt management (create, read, update, delete, search) -- Image gallery with automatic ComfyUI output monitoring -- Bulk operations for efficiency -- Database maintenance and optimization -- System diagnostics and logging -- Thumbnail generation and management -- Metadata extraction from PNG files -- Real-time progress tracking - -The API integrates with ComfyUI's aiohttp server and provides endpoints for -both programmatic access and web UI functionality. - -Classes: - PromptManagerAPI: Main API class handling all REST endpoints - -Example: - api = PromptManagerAPI() - api.add_routes(server.routes) -""" - -# PromptManager/py/api.py - -import datetime -import json -import os -import traceback -from typing import Any, Dict, List, Optional - -import server -from aiohttp import web -from PIL import Image - -# Import database operations -try: - from ..database.operations import PromptDatabase - from ..utils.logging_config import get_logger -except ImportError: - # Fallback for when module isn't imported as package - import os - import sys - - sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__)))) - from database.operations import PromptDatabase - from utils.logging_config import get_logger - - -class PromptManagerAPI: - """REST API handler for PromptManager operations and web interface. - - This class provides comprehensive REST API endpoints for managing prompts, - images, and system operations. It handles database interactions, file - operations, image processing, and web UI serving. - - The API is designed to integrate seamlessly with ComfyUI's aiohttp server - and provides both JSON API endpoints and static file serving for the - web interface. - - Attributes: - logger: Configured logger instance for API operations - db (PromptDatabase): Database connection and operations handler - - Example: - api = PromptManagerAPI() - api.add_routes(server_routes) - # API endpoints are now available at /prompt_manager/* - """ - - def __init__(self): - """Initialize the PromptManager API with database connection and cleanup. - - Sets up logging, initializes the database connection, and performs - startup cleanup to remove any duplicate prompts that may exist. - - Raises: - Exception: If database initialization fails, logs error but continues. - If cleanup fails, logs error but continues operation. - """ - self.logger = get_logger('prompt_manager.api') - self.logger.info("Initializing PromptManager API") - - self.db = PromptDatabase() - - # Run cleanup on initialization to remove any existing duplicates - try: - removed = self.db.cleanup_duplicates() - if removed > 0: - self.logger.info(f"Startup cleanup: removed {removed} duplicate prompts") - except Exception as e: - self.logger.error(f"Startup cleanup failed: {e}") - - self.logger.info("PromptManager API initialization completed") - - def add_routes(self, routes): - """Add API routes to ComfyUI server using decorator pattern. - - Registers all API endpoints with the provided aiohttp routes object. - Uses decorator pattern to define routes inline with their handlers. - - Categories of routes registered: - - Core prompt operations (search, save, delete, categories, tags) - - Database maintenance (cleanup, duplicates, maintenance) - - Web UI serving (admin, gallery, metadata viewer) - - Image operations (gallery, thumbnails, metadata extraction) - - System operations (diagnostics, logging, statistics) - - Bulk operations (delete, tag management, export) - - Args: - routes: aiohttp RouteTableDef object from ComfyUI server instance. - Routes will be registered with this object for URL handling. - - Example: - from server import PromptServer - api = PromptManagerAPI() - api.add_routes(PromptServer.instance.routes) - """ - - # Test route to verify registration works - @routes.get("/prompt_manager/test") - async def test_route(request): - return web.json_response( - { - "success": True, - "message": "PromptManager API is working!", - "timestamp": str(datetime.datetime.now()), - } - ) - - @routes.get("/prompt_manager/search") - async def search_prompts_route(request): - return await self.search_prompts(request) - - @routes.get("/prompt_manager/recent") - async def get_recent_prompts_route(request): - return await self.get_recent_prompts(request) - - @routes.get("/prompt_manager/categories") - async def get_categories_route(request): - return await self.get_categories(request) - - # Tag management endpoints (must be registered BEFORE /prompt_manager/tags) - @routes.get("/prompt_manager/tags/stats") - async def get_tags_stats_route(request): - return await self.get_tags_stats(request) - - @routes.get("/prompt_manager/tags/filter") - async def get_tags_filter_route(request): - return await self.get_tags_filter(request) - - # Bulk tag operations (register BEFORE {tag_name} to avoid path param match) - @routes.post("/prompt_manager/tags/merge") - async def merge_tags_route(request): - return await self.merge_tags_endpoint(request) - - @routes.put("/prompt_manager/tags/{tag_name}") - async def rename_tag_route(request): - return await self.rename_tag_endpoint(request) - - @routes.delete("/prompt_manager/tags/{tag_name}") - async def delete_tag_route(request): - return await self.delete_tag_endpoint(request) - - @routes.get("/prompt_manager/tags/{tag_name}/prompts") - async def get_tag_prompts_route(request): - return await self.get_tag_prompts(request) - - @routes.get("/prompt_manager/tags") - async def get_tags_route(request): - return await self.get_tags(request) - - @routes.get("/prompt_manager/scan_duplicates") - async def scan_duplicates_route(request): - return await self.scan_duplicates_endpoint(request) - - @routes.post("/prompt_manager/delete_duplicate_images") - async def delete_duplicate_images_route(request): - return await self.delete_duplicate_images_endpoint(request) - - @routes.post("/prompt_manager/cleanup") - async def cleanup_duplicates_route(request): - return await self.cleanup_duplicates_endpoint(request) - - @routes.post("/prompt_manager/maintenance") - async def maintenance_route(request): - return await self.run_maintenance(request) - - @routes.post("/prompt_manager/save") - async def save_prompt_route(request): - return await self.save_prompt(request) - - @routes.delete("/prompt_manager/delete/{prompt_id}") - async def delete_prompt_route(request): - return await self.delete_prompt(request) - - # Serve the web UI HTML file - @routes.get("/prompt_manager/web") - async def serve_web_ui(request): - try: - import os - - # Get the path to the web directory - current_dir = os.path.dirname( - os.path.dirname(os.path.abspath(__file__)) - ) - html_path = os.path.join(current_dir, "web", "index.html") - - if os.path.exists(html_path): - with open(html_path, "r", encoding="utf-8") as f: - html_content = f.read() - - return web.Response( - text=html_content, content_type="text/html", charset="utf-8" - ) - else: - return web.Response( - text="

Web UI not found

HTML file not located at expected path.

", - content_type="text/html", - status=404, - ) - - except Exception as e: - return web.Response( - text=f"

Error

Failed to load web UI: {str(e)}

", - content_type="text/html", - status=500, - ) - - # Serve the gallery interface - @routes.get("/prompt_manager/gallery.html") - async def serve_gallery_ui(request): - try: - import os - - current_dir = os.path.dirname( - os.path.dirname(os.path.abspath(__file__)) - ) - html_path = os.path.join(current_dir, "web", "metadata.html") - - if os.path.exists(html_path): - with open(html_path, "r", encoding="utf-8") as f: - html_content = f.read() - - return web.Response( - text=html_content, content_type="text/html", charset="utf-8" - ) - else: - return web.Response( - text="

Gallery not found

gallery.html file not located at expected path.

", - content_type="text/html", - status=404, - ) - - except Exception as e: - return web.Response( - text=f"

Error

Failed to load gallery: {str(e)}

", - content_type="text/html", - status=500, - ) - - # Serve the admin interface - @routes.get("/prompt_manager/admin") - async def serve_admin_ui(request): - try: - import os - - current_dir = os.path.dirname( - os.path.dirname(os.path.abspath(__file__)) - ) - html_path = os.path.join(current_dir, "web", "admin.html") - - if os.path.exists(html_path): - with open(html_path, "r", encoding="utf-8") as f: - html_content = f.read() - - return web.Response( - text=html_content, content_type="text/html", charset="utf-8" - ) - else: - return web.Response( - text="

Admin UI not found

", - content_type="text/html", - status=404, - ) - - except Exception as e: - return web.Response( - text=f"

Error

Failed to load admin UI: {str(e)}

", - content_type="text/html", - status=500, - ) - - # Serve the gallery interface (new admin gallery) - @routes.get("/prompt_manager/gallery") - async def serve_gallery_admin_ui(request): - try: - import os - - current_dir = os.path.dirname( - os.path.dirname(os.path.abspath(__file__)) - ) - html_path = os.path.join(current_dir, "web", "gallery.html") - - if os.path.exists(html_path): - with open(html_path, "r", encoding="utf-8") as f: - html_content = f.read() - - return web.Response( - text=html_content, content_type="text/html", charset="utf-8" - ) - else: - return web.Response( - text="

Gallery not found

gallery.html file not located at expected path.

", - content_type="text/html", - status=404, - ) - - except Exception as e: - return web.Response( - text=f"

Error

Failed to load gallery: {str(e)}

", - content_type="text/html", - status=500, - ) - - # Serve static files from web/lib directory - @routes.get("/prompt_manager/lib/{filepath:.*}") - async def serve_lib_static(request): - """Serve static library files (JS, CSS) from web/lib directory.""" - import os - - # Explicit MIME type mapping - MIME_TYPES = { - ".js": "application/javascript", - ".css": "text/css", - ".json": "application/json", - ".map": "application/json", - } - - filepath = request.match_info.get("filepath", "") - - # Security: prevent directory traversal - if ".." in filepath or filepath.startswith("/"): - return web.Response(text="Forbidden", status=403) - - current_dir = os.path.dirname( - os.path.dirname(os.path.abspath(__file__)) - ) - file_path = os.path.join(current_dir, "web", "lib", filepath) - - if not os.path.exists(file_path) or not os.path.isfile(file_path): - return web.Response(text=f"Not Found: {filepath}", status=404) - - # Get extension and content type - ext = os.path.splitext(file_path)[1].lower() - content_type = MIME_TYPES.get(ext, "application/octet-stream") - - with open(file_path, "rb") as f: - content = f.read() - - return web.Response( - body=content, - content_type=content_type, - ) - - # Serve static JS files from web/js directory - @routes.get("/prompt_manager/js/{filepath:.*}") - async def serve_js_static(request): - """Serve static JavaScript files from web/js directory.""" - import os - - MIME_TYPES = { - ".js": "application/javascript", - ".css": "text/css", - ".json": "application/json", - ".map": "application/json", - } - - filepath = request.match_info.get("filepath", "") - - if ".." in filepath or filepath.startswith("/"): - return web.Response(text="Forbidden", status=403) - - current_dir = os.path.dirname( - os.path.dirname(os.path.abspath(__file__)) - ) - file_path = os.path.join(current_dir, "web", "js", filepath) - - if not os.path.exists(file_path) or not os.path.isfile(file_path): - return web.Response(text=f"Not Found: {filepath}", status=404) - - ext = os.path.splitext(file_path)[1].lower() - content_type = MIME_TYPES.get(ext, "application/octet-stream") - - with open(file_path, "rb") as f: - content = f.read() - - return web.Response( - body=content, - content_type=content_type, - ) - - # Statistics endpoint - @routes.get("/prompt_manager/stats") - async def get_stats_route(request): - return await self.get_statistics(request) - - # Settings endpoints - @routes.get("/prompt_manager/settings") - async def get_settings_route(request): - return await self.get_settings(request) - - @routes.post("/prompt_manager/settings") - async def save_settings_route(request): - return await self.save_settings(request) - - - # Individual prompt management - @routes.put("/prompt_manager/prompts/{prompt_id}") - async def update_prompt_route(request): - return await self.update_prompt(request) - - @routes.put("/prompt_manager/prompts/{prompt_id}/rating") - async def update_rating_route(request): - return await self.update_prompt_rating(request) - - @routes.post("/prompt_manager/prompts/{prompt_id}/tags") - async def add_tag_route(request): - return await self.add_prompt_tag(request) - - @routes.delete("/prompt_manager/prompts/{prompt_id}/tags") - async def remove_tag_route(request): - return await self.remove_prompt_tag(request) - - @routes.post("/prompt_manager/prompts/tags") - async def add_tags_to_prompt_route(request): - return await self.add_tags_to_prompt(request) - - # Bulk operations - @routes.post("/prompt_manager/bulk/delete") - async def bulk_delete_route(request): - return await self.bulk_delete_prompts(request) - - @routes.post("/prompt_manager/bulk/tags") - async def bulk_add_tags_route(request): - return await self.bulk_add_tags(request) - - @routes.post("/prompt_manager/bulk/category") - async def bulk_set_category_route(request): - return await self.bulk_set_category(request) - - # Export functionality - @routes.get("/prompt_manager/export") - async def export_prompts_route(request): - return await self.export_prompts(request) - - # Database backup and restore - @routes.get("/prompt_manager/backup") - async def backup_database_route(request): - return await self.backup_database(request) - - @routes.post("/prompt_manager/restore") - async def restore_database_route(request): - return await self.restore_database(request) - - # Gallery endpoints - @routes.get("/prompt_manager/prompts/{prompt_id}/images") - async def get_prompt_images_route(request): - return await self.get_prompt_images(request) - - @routes.get("/prompt_manager/images/recent") - async def get_recent_images_route(request): - return await self.get_recent_images(request) - - @routes.get("/prompt_manager/images/all") - async def get_all_images_route(request): - return await self.get_all_images(request) - - @routes.get("/prompt_manager/images/search") - async def search_images_route(request): - return await self.search_images(request) - - @routes.get("/prompt_manager/images/output") - async def get_output_images_route(request): - return await self.get_output_images(request) - - @routes.get("/prompt_manager/images/{image_id}/file") - async def serve_image_route(request): - return await self.serve_image(request) - - @routes.get("/prompt_manager/images/serve/{filepath:.*}") - async def serve_output_image_route(request): - return await self.serve_output_image(request) - - @routes.post("/prompt_manager/images/link") - async def link_image_route(request): - return await self.link_image_to_prompt(request) - - @routes.get("/prompt_manager/images/prompt/{image_path:.*}") - async def get_image_prompt_route(request): - return await self.get_image_prompt(request) - - @routes.delete("/prompt_manager/images/{image_id}") - async def delete_image_route(request): - return await self.delete_image(request) - - @routes.post("/prompt_manager/images/generate-thumbnails") - async def generate_thumbnails_route(request): - return await self.generate_thumbnails(request) - - @routes.get("/prompt_manager/images/generate-thumbnails/progress") - async def generate_thumbnails_progress_route(request): - return await self.generate_thumbnails_with_progress(request) - - @routes.post("/prompt_manager/images/clear-thumbnails") - async def clear_thumbnails_route(request): - return await self.clear_thumbnails(request) - - # Diagnostic endpoints - @routes.get("/prompt_manager/diagnostics") - async def run_diagnostics_route(request): - return await self.run_diagnostics(request) - - @routes.post("/prompt_manager/diagnostics/test-link") - async def test_image_link_route(request): - return await self.test_image_link(request) - - # Scan endpoint - @routes.post("/prompt_manager/scan") - async def scan_images_route(request): - return await self.scan_images(request) - - # Logging endpoints - @routes.get("/prompt_manager/logs") - async def get_logs_route(request): - return await self.get_logs(request) - - @routes.get("/prompt_manager/logs/files") - async def get_log_files_route(request): - return await self.get_log_files(request) - - @routes.get("/prompt_manager/logs/download/{filename}") - async def download_log_route(request): - return await self.download_log_file(request) - - @routes.post("/prompt_manager/logs/truncate") - async def truncate_logs_route(request): - return await self.truncate_logs(request) - - @routes.get("/prompt_manager/logs/config") - async def get_log_config_route(request): - return await self.get_log_config(request) - - @routes.post("/prompt_manager/logs/config") - async def update_log_config_route(request): - return await self.update_log_config(request) - - @routes.get("/prompt_manager/logs/stats") - async def get_log_stats_route(request): - return await self.get_log_stats(request) - - # AutoTag endpoints - @routes.get("/prompt_manager/autotag/models") - async def get_autotag_models_route(request): - return await self.get_autotag_models(request) - - @routes.get("/prompt_manager/autotag/download/{model_type}") - async def download_autotag_model_route(request): - return await self.download_autotag_model(request) - - @routes.get("/prompt_manager/autotag/start") - async def start_autotag_route(request): - return await self.start_autotag(request) - - @routes.post("/prompt_manager/autotag/single") - async def autotag_single_route(request): - return await self.autotag_single(request) - - @routes.post("/prompt_manager/autotag/apply") - async def apply_autotag_route(request): - return await self.apply_autotag(request) - - @routes.post("/prompt_manager/autotag/unload") - async def unload_autotag_model_route(request): - return await self.unload_autotag_model(request) - - @routes.get("/prompt_manager/scan_output_dir") - async def scan_output_dir_route(request): - return await self.scan_output_dir(request) - - self.logger.info("All routes registered with decorator pattern") - - async def search_prompts(self, request): - """Search for prompts using multiple filter criteria. - - Provides comprehensive search functionality across prompt text, categories, - tags, and ratings with configurable result limits. - - Query Parameters: - text (str, optional): Search text to match against prompt content - category (str, optional): Filter by specific category - tags (str, optional): Comma-separated list of tags to filter by - min_rating (int, optional): Minimum rating (1-5) to include - limit (int, optional): Maximum results to return (default: 50, max: 1000) - - Args: - request (aiohttp.web.Request): HTTP request object containing query parameters - - Returns: - aiohttp.web.Response: JSON response with structure: - { - "success": bool, - "results": List[Dict], # List of matching prompt objects - "count": int # Number of results returned - } - - Raises: - Returns 500 status with error details if search fails - - Example: - GET /prompt_manager/search?text=portrait&category=photography&min_rating=3 - """ - try: - # Get query parameters - text = request.query.get("text", "").strip() - category = request.query.get("category", "").strip() - tags_str = request.query.get("tags", "").strip() - min_rating = request.query.get("min_rating", 0) - limit = int(request.query.get("limit", 50)) - - # Parse tags - tags = None - if tags_str: - tags = [tag.strip() for tag in tags_str.split(",") if tag.strip()] - - # Parse min_rating - try: - min_rating = int(min_rating) if min_rating else None - except ValueError: - min_rating = None - - # Perform search - results = self.db.search_prompts( - text=text if text else None, - category=category if category else None, - tags=tags, - rating_min=min_rating, - limit=limit, - ) - - return web.json_response( - {"success": True, "results": results, "count": len(results)} - ) - - except Exception as e: - self.logger.error(f"Search error: {e}", exc_info=True) - return web.json_response( - {"success": False, "error": f"Search failed: {str(e)}", "results": []}, - status=500, - ) - - async def get_recent_prompts(self, request): - """Retrieve recently created prompts with pagination support. - - Returns prompts sorted by creation date (newest first) with configurable - pagination using either page-based or offset-based navigation. - - Query Parameters: - limit (int, optional): Number of prompts per page (default: 50, max: 1000) - page (int, optional): Page number for pagination (1-based, default: 1) - offset (int, optional): Offset for results (takes precedence over page) - - Args: - request (aiohttp.web.Request): HTTP request object containing query parameters - - Returns: - aiohttp.web.Response: JSON response with structure: - { - "success": bool, - "results": List[Dict], # List of prompt objects - "pagination": { - "total": int, # Total number of prompts - "limit": int, # Items per page - "offset": int, # Current offset - "page": int, # Current page number - "total_pages": int, # Total number of pages - "has_more": bool, # Whether more pages exist - "count": int # Number of items in current page - } - } - - Raises: - Returns 500 status with error details if retrieval fails - - Example: - GET /prompt_manager/recent?limit=20&page=2 - """ - try: - limit = int(request.query.get("limit", 50)) - page = int(request.query.get("page", 1)) - offset = int(request.query.get("offset", 0)) - - # If page is provided, calculate offset from page - if page > 1 and offset == 0: - offset = (page - 1) * limit - - # Ensure reasonable limits - if limit > 1000: - limit = 1000 - elif limit < 1: - limit = 1 - - results = self.db.get_recent_prompts(limit=limit, offset=offset) - - 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) - return web.json_response( - { - "success": False, - "error": f"Failed to get recent prompts: {str(e)}", - "results": [], - "pagination": {"total": 0, "page": 1, "total_pages": 0} - }, - status=500, - ) - - async def get_categories(self, request): - """Retrieve all available prompt categories. - - Returns a list of all unique categories found in the database. - - Args: - request (aiohttp.web.Request): HTTP request object - - Returns: - aiohttp.web.Response: JSON response with structure: - { - "success": bool, - "categories": List[str] # List of category names - } - - Raises: - Returns 500 status with error details if retrieval fails - - Example: - GET /prompt_manager/categories - """ - try: - categories = self.db.get_all_categories() - - return web.json_response({"success": True, "categories": categories}) - - 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": [], - }, - status=500, - ) - - async def get_tags(self, request): - """Retrieve all available prompt tags. - - Returns a list of all unique tags found across all prompts in the database. - - Args: - request (aiohttp.web.Request): HTTP request object - - Returns: - aiohttp.web.Response: JSON response with structure: - { - "success": bool, - "tags": List[str] # List of unique tag names - } - - Raises: - Returns 500 status with error details if retrieval fails - - Example: - GET /prompt_manager/tags - """ - try: - tags = self.db.get_all_tags() - - return web.json_response({"success": True, "tags": tags}) - - 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": [], - }, - status=500, - ) - - async def get_tags_stats(self, request): - """Get tags with usage counts, search, sort, and pagination.""" - try: - try: - limit = int(request.query.get("limit", 50)) - offset = int(request.query.get("offset", 0)) - except (ValueError, TypeError): - return web.json_response( - {"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 = self.db.get_tags_with_counts(limit, offset, search, sort) - untagged_count = 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'] - } - }) - 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, - ) - - def _enrich_prompt_images(self, prompts): - """Add url and thumbnail_url to each image in prompt results.""" - from pathlib import Path - from urllib.parse import quote as url_quote - - output_dir = self._find_comfyui_output_dir() - 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', '') - 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" - - # 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='/')}" - - # Check for thumbnail - 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='/')}" - except (ValueError, RuntimeError): - pass - - return prompts - - async def get_tag_prompts(self, request): - """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( - {"success": False, "error": "Tag name required"}, - status=400, - ) - - try: - limit = int(request.query.get("limit", 20)) - offset = int(request.query.get("offset", 0)) - except (ValueError, TypeError): - return web.json_response( - {"success": False, "error": "Invalid limit or offset parameter"}, status=400 - ) - - result = 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'] - } - }) - 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, - ) - - async def get_tags_filter(self, request): - """Get prompts matching multiple tags with AND/OR mode, or untagged prompts.""" - try: - untagged = request.query.get("untagged", "").lower() == "true" - - if untagged: - try: - limit = int(request.query.get("limit", 20)) - offset = int(request.query.get("offset", 0)) - except (ValueError, TypeError): - return web.json_response( - {"success": False, "error": "Invalid limit or offset parameter"}, status=400 - ) - result = 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: - return web.json_response( - {"success": False, "error": "Tags parameter required"}, - status=400, - ) - - tags_list = [t.strip() for t in tags_str.split(",") if t.strip()] - mode = request.query.get("mode", "and").lower() - if mode not in ("and", "or"): - mode = "and" - - try: - limit = int(request.query.get("limit", 20)) - offset = int(request.query.get("offset", 0)) - except (ValueError, TypeError): - return web.json_response( - {"success": False, "error": "Invalid limit or offset parameter"}, status=400 - ) - - result = 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'] - } - }) - 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, - ) - - async def rename_tag_endpoint(self, request): - """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( - {"success": False, "error": "Tag name required"}, status=400 - ) - - try: - body = await request.json() - except Exception: - return web.json_response( - {"success": False, "error": "Invalid JSON body"}, status=400 - ) - new_name = body.get("new_name", "").strip() - if not new_name: - return web.json_response( - {"success": False, "error": "New tag name required"}, status=400 - ) - - result = 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'] - } - 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) - return web.json_response( - {"success": False, "error": str(e)}, status=500 - ) - - async def delete_tag_endpoint(self, request): - """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 = self.db.delete_tag_all_prompts(tag_name) - resp = { - "success": True, - "tag_name": tag_name, - "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" - return web.json_response(resp) - except Exception as e: - self.logger.error(f"Delete tag error: {e}", exc_info=True) - return web.json_response( - {"success": False, "error": str(e)}, status=500 - ) - - async def merge_tags_endpoint(self, request): - """Merge source tags into a target tag.""" - try: - try: - body = await request.json() - except Exception: - return web.json_response( - {"success": False, "error": "Invalid JSON body"}, status=400 - ) - source_tags = body.get("source_tags", []) - target_tag = body.get("target_tag", "").strip() - - if not source_tags: - return web.json_response( - {"success": False, "error": "Source tags required"}, status=400 - ) - if not target_tag: - return web.json_response( - {"success": False, "error": "Target tag required"}, status=400 - ) - - result = 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'] - } - 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) - return web.json_response( - {"success": False, "error": str(e)}, status=500 - ) - - async def save_prompt(self, request): - """Save a new prompt with metadata and duplicate detection. - - Creates a new prompt record with automatic duplicate detection based on - SHA256 hash of the prompt text. If a duplicate is found, updates the - existing record's metadata instead of creating a new one. - - Request Body (JSON): - text (str, required): The prompt text content - category (str, optional): Category for organization - tags (List[str], optional): List of tags for classification - rating (int, optional): Rating from 1-5 - notes (str, optional): Additional notes or description - - Args: - request (aiohttp.web.Request): HTTP request with JSON body - - Returns: - aiohttp.web.Response: JSON response with structure: - { - "success": bool, - "prompt_id": int, # ID of created/updated prompt - "message": str, # Success/status message - "is_duplicate": bool # True if prompt already existed - } - - Raises: - Returns 400 status if required fields are missing - Returns 500 status with error details if save fails - - Example: - POST /prompt_manager/save - { - "text": "A beautiful sunset over mountains", - "category": "landscape", - "tags": ["nature", "scenic"], - "rating": 4, - "notes": "Great for wallpapers" - } - """ - try: - data = await request.json() - - text = data.get("text", "").strip() - if not text: - return web.json_response( - {"success": False, "error": "Text is required"}, status=400 - ) - - category = data.get("category", "").strip() or None - tags = data.get("tags", []) - rating = data.get("rating") or None - notes = data.get("notes", "").strip() or None - - # Generate hash for duplicate detection - import hashlib - - prompt_hash = hashlib.sha256(text.encode("utf-8")).hexdigest() - - # Check if prompt already exists - existing = self.db.get_prompt_by_hash(prompt_hash) - if existing: - # Update metadata if this is a duplicate with new info - if any([category, tags, rating, notes]): - self.db.update_prompt_metadata( - prompt_id=existing['id'], - category=category, - tags=tags, - rating=rating, - notes=notes - ) - return web.json_response( - { - "success": True, - "prompt_id": existing['id'], - "message": "Prompt already exists, metadata updated", - "is_duplicate": True - } - ) - - # Save new prompt - prompt_id = self.db.save_prompt( - text=text, - category=category, - tags=tags if tags else None, - rating=rating, - notes=notes, - prompt_hash=prompt_hash, - ) - - 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) - return web.json_response( - {"success": False, "error": f"Failed to save prompt: {str(e)}"}, - status=500, - ) - - async def delete_prompt(self, request): - """Delete a specific prompt by ID. - - Permanently removes a prompt and all associated metadata from the database. - Associated image links may also be removed depending on database configuration. - - URL Parameters: - prompt_id (int): The unique identifier of the prompt to delete - - Args: - request (aiohttp.web.Request): HTTP request with prompt_id in URL path - - Returns: - aiohttp.web.Response: JSON response with structure: - { - "success": bool, - "message": str # Success or error message - } - - Raises: - Returns 404 status if prompt not found - Returns 500 status with error details if deletion fails - - Example: - DELETE /prompt_manager/delete/123 - """ - try: - prompt_id = int(request.match_info["prompt_id"]) - - success = self.db.delete_prompt(prompt_id) - - if success: - return web.json_response( - {"success": True, "message": "Prompt deleted successfully"} - ) - else: - return web.json_response( - { - "success": False, - "error": "Prompt not found or could not be deleted", - }, - status=404, - ) - - except ValueError: - return web.json_response( - {"success": False, "error": "Invalid prompt ID"}, status=400 - ) - except Exception as e: - self.logger.error(f"Delete error: {e}") - return web.json_response( - {"success": False, "error": f"Failed to delete prompt: {str(e)}"}, - status=500, - ) - - async def scan_duplicates_endpoint(self, request): - """ - Scan for duplicate images without removing them. - GET /prompt_manager/scan_duplicates - """ - try: - duplicates = await self.find_duplicate_images() - - return web.json_response( - { - "success": True, - "duplicates": duplicates, - "message": f"Found {len(duplicates)} groups of duplicate images", - } - ) - - 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)}"}, - status=500, - ) - - async def cleanup_duplicates_endpoint(self, request): - """ - Cleanup duplicate prompts endpoint. - POST /prompt_manager/cleanup - """ - try: - removed_count = self.db.cleanup_duplicates() - - return web.json_response( - { - "success": True, - "message": f"Cleanup completed", - "duplicates_removed": removed_count, - } - ) - - except Exception as e: - self.logger.error(f"Cleanup error: {e}") - return web.json_response( - {"success": False, "error": f"Failed to cleanup duplicates: {str(e)}"}, - status=500, - ) - - async def find_duplicate_images(self): - """Find duplicate images in ComfyUI output directory using content hashing. - - Scans the ComfyUI output directory for image and video files, calculates - SHA256 hashes of file contents, and identifies groups of files with - identical content. Supports both images and videos with thumbnail detection. - - The method processes files efficiently, logging progress every 100 files, - and handles various media formats including PNG, JPG, JPEG, WebP, GIF, - and common video formats. - - Returns: - List[Dict[str, Any]]: List of duplicate groups, where each group contains: - - hash (str): SHA256 hash of the file content - - images (List[Dict]): List of file info dictionaries with: - - id (str): Unique identifier based on file path hash - - filename (str): Original filename - - path (str): Absolute file path - - relative_path (str): Path relative to output directory - - url (str): URL for serving the file - - thumbnail_url (str, optional): URL for thumbnail if available - - size (int): File size in bytes - - modified_time (float): Last modification timestamp - - extension (str): File extension - - media_type (str): 'image' or 'video' - - is_video (bool): True if file is a video - - hash (str): SHA256 content hash - - count (int): Number of duplicate files in this group - - Note: - Files within each duplicate group are sorted by modification time - (oldest first) to help users decide which files to keep. - - Example: - duplicates = await api.find_duplicate_images() - for group in duplicates: - print(f"Found {group['count']} duplicates with hash {group['hash']}") - for img in group['images']: - print(f" - {img['filename']} ({img['size']} bytes)") - """ - import hashlib - from pathlib import Path - - self.logger.info("Scanning for duplicate images") - - try: - # Find ComfyUI output directory - output_dir = self._find_comfyui_output_dir() - if not output_dir: - self.logger.warning("ComfyUI output directory not found") - return [] - - output_path = Path(output_dir) - - # Extensions to check - image_extensions = ['.png', '.jpg', '.jpeg', '.gif', '.webp', '.bmp', '.tiff'] - video_extensions = ['.mp4', '.avi', '.mov', '.mkv', '.webm', '.gif'] - all_extensions = image_extensions + video_extensions - - # Find all media files, excluding thumbnails directory - # Use a set to avoid duplicates on case-insensitive filesystems (Windows) - media_files = [] - seen_paths = set() - for ext in all_extensions: - # Search for both lowercase and uppercase extensions - for pattern in [f"*{ext.lower()}", f"*{ext.upper()}"]: - for media_path in output_path.rglob(pattern): - # Skip files in thumbnails directory - if 'thumbnails' not in media_path.parts: - # Normalize path for deduplication (case-insensitive on Windows) - normalized_path = str(media_path).lower() - if normalized_path not in seen_paths: - seen_paths.add(normalized_path) - media_files.append(media_path) - - self.logger.info(f"Found {len(media_files)} media files to analyze") - - # Calculate hash for each file - file_hashes = {} - processed = 0 - - for media_path in media_files: - try: - # Calculate file hash - file_hash = self._calculate_file_hash(media_path) - - if file_hash not in file_hashes: - file_hashes[file_hash] = [] - - # Create image info object similar to get_output_images - stat = media_path.stat() - 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' - - # Check if thumbnail exists - thumbnail_url = None - thumbnails_dir = output_path / "thumbnails" - if thumbnails_dir.exists(): - thumbnail_ext = '.jpg' if is_video else extension - # Preserve subdirectory structure in thumbnail path - 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 - } - - file_hashes[file_hash].append(image_info) - processed += 1 - - # Log progress every 100 files - if processed % 100 == 0: - 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}") - continue - - # Find duplicates (groups with more than one file) - duplicates = [] - for file_hash, images in file_hashes.items(): - if len(images) > 1: - # Sort by modification time (oldest first) to help users decide which to keep - 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 - - except Exception as e: - self.logger.error(f"Error finding duplicate images: {e}") - return [] - - def _calculate_file_hash(self, file_path): - """Calculate SHA-256 hash of a file's content. - - Reads the file in 4KB chunks to efficiently handle large files - without loading the entire content into memory. - - Args: - file_path (str): Path to the file to hash - - Returns: - str: Hexadecimal SHA-256 hash of the file content - - Raises: - IOError: If the file cannot be read - OSError: If the file path is invalid or inaccessible - """ - import hashlib - - hash_sha256 = hashlib.sha256() - with open(file_path, "rb") as f: - for chunk in iter(lambda: f.read(4096), b""): - hash_sha256.update(chunk) - return hash_sha256.hexdigest() - - async def delete_duplicate_images_endpoint(self, request): - """ - Delete duplicate image files from disk. - POST /prompt_manager/delete_duplicate_images - """ - try: - data = await request.json() - image_paths = data.get('image_paths', []) - - if not image_paths: - return web.json_response( - {"success": False, "error": "No image paths provided"}, - status=400, - ) - - deleted_count = 0 - failed_count = 0 - failed_files = [] - - for image_path in image_paths: - try: - from pathlib import Path - import os - - # 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_count += 1 - continue - - output_path = Path(output_dir) - file_path = Path(image_path) - - # Security check - ensure file is within output directory - try: - file_path.resolve().relative_to(output_path.resolve()) - except ValueError: - 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 - - if file_path.exists() and file_path.is_file(): - os.remove(file_path) - deleted_count += 1 - self.logger.info(f"Deleted duplicate image: {image_path}") - - # 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}" - if thumbnail_path.exists(): - os.remove(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}") - else: - failed_files.append(f"{image_path} (file not found)") - failed_count += 1 - - except Exception as e: - self.logger.error(f"Error deleting file {image_path}: {e}") - failed_files.append(f"{image_path} ({str(e)})") - failed_count += 1 - - response_data = { - "success": True, - "deleted_count": deleted_count, - "failed_count": failed_count, - "message": f"Deleted {deleted_count} files successfully" - } - - if failed_count > 0: - response_data["failed_files"] = failed_files - response_data["message"] += f", {failed_count} failed" - - return web.json_response(response_data) - - 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)}"}, - status=500, - ) - - async def get_statistics(self, request): - """Get database statistics.""" - try: - # Get basic stats from database - with self.db.model.get_connection() as conn: - cursor = conn.execute("SELECT COUNT(*) as total FROM prompts") - total_prompts = cursor.fetchone()["total"] - - cursor = conn.execute( - "SELECT COUNT(DISTINCT TRIM(category)) as total FROM prompts WHERE category IS NOT NULL AND TRIM(category) != ''" - ) - total_categories = cursor.fetchone()["total"] - - cursor = conn.execute( - "SELECT AVG(rating) as avg FROM prompts WHERE rating IS NOT NULL" - ) - avg_rating = cursor.fetchone()["avg"] - - # Count unique tags - cursor = conn.execute("SELECT tags FROM prompts WHERE tags IS NOT NULL") - all_tags = set() - for row in cursor.fetchall(): - try: - tags = json.loads(row["tags"]) - if isinstance(tags, list): - all_tags.update(tags) - except: - continue - - # Get all categories for debugging (including raw data) - cursor = conn.execute( - "SELECT DISTINCT category FROM prompts WHERE category IS NOT NULL ORDER BY category" - ) - raw_categories = [row['category'] for row in cursor.fetchall()] - - cursor = conn.execute( - "SELECT DISTINCT TRIM(category) as category FROM prompts WHERE category IS NOT NULL AND TRIM(category) != '' ORDER BY category" - ) - filtered_categories = [row['category'] for row in cursor.fetchall()] - - # Debug logging with detailed category info - self.logger.debug(f"Statistics calculated - Prompts: {total_prompts}, Categories: {total_categories}, Tags: {len(all_tags)}, Avg Rating: {avg_rating}") - self.logger.debug(f"Raw categories from DB: {[(cat, len(cat), repr(cat)) for cat in raw_categories]}") - self.logger.debug(f"Filtered categories: {[(cat, len(cat), repr(cat)) for cat in filtered_categories]}") - - return web.json_response( - { - "success": True, - "stats": { - "total_prompts": total_prompts, - "total_categories": total_categories, - "unique_categories": total_categories, # Keep both for compatibility - "total_tags": len(all_tags), - "average_rating": ( - round(avg_rating, 2) if avg_rating else None - ), - "avg_rating": ( - round(avg_rating, 2) if avg_rating else None - ), # Keep both for compatibility - "recent_prompts": total_prompts, # For now, use total as recent count - }, - } - ) - - except Exception as e: - self.logger.error(f"Stats error: {e}") - return web.json_response( - {"success": False, "error": f"Failed to get statistics: {str(e)}"}, - status=500, - ) - - async def get_settings(self, request): - """Get current settings.""" - try: - from .config import PromptManagerConfig, GalleryConfig - - # Get monitored directories from image monitor singleton if available - monitored_dirs = [] - try: - import sys - # Check for the module in sys.modules (handles different import paths) - 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'): - 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', []) - elif GalleryConfig.MONITORING_DIRECTORIES: - monitored_dirs = GalleryConfig.MONITORING_DIRECTORIES - except Exception: - # Fallback to config - if GalleryConfig.MONITORING_DIRECTORIES: - monitored_dirs = GalleryConfig.MONITORING_DIRECTORIES - - # Return actual configuration settings - 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)}"}, - status=500, - ) - - async def save_settings(self, request): - """Save settings.""" - try: - from .config import PromptManagerConfig, GalleryConfig - import json - import os - - data = await request.json() - 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'] - - # 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 new_path != old_path: - if new_path: - GalleryConfig.MONITORING_DIRECTORIES = [new_path] - else: - GalleryConfig.MONITORING_DIRECTORIES = [] # Reset to auto-detect - restart_required = True - - # Save to config file for persistence - config_dir = 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 - }, - 'gallery': { - 'monitoring': { - 'directories': GalleryConfig.MONITORING_DIRECTORIES - } - } - } - - try: - 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 - }) - except Exception as e: - return web.json_response( - {"success": False, "error": f"Failed to save settings: {str(e)}"}, - status=500, - ) - - async def update_prompt(self, request): - """Update prompt text.""" - try: - prompt_id = int(request.match_info["prompt_id"]) - data = await request.json() - new_text = data.get("text", "").strip() - - if not new_text: - return web.json_response( - {"success": False, "error": "Text cannot be empty"}, status=400 - ) - - # Update the prompt in database - with self.db.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, - ), - ) - conn.commit() - - if cursor.rowcount > 0: - return web.json_response( - {"success": True, "message": "Prompt updated successfully"} - ) - else: - return web.json_response( - {"success": False, "error": "Prompt not found"}, status=404 - ) - - except ValueError: - return web.json_response( - {"success": False, "error": "Invalid prompt ID"}, status=400 - ) - except Exception as e: - self.logger.error(f"Update prompt error: {e}") - return web.json_response( - {"success": False, "error": f"Failed to update prompt: {str(e)}"}, - status=500, - ) - - async def update_prompt_rating(self, request): - """Update prompt rating.""" - try: - prompt_id = int(request.match_info["prompt_id"]) - data = await request.json() - rating = data.get("rating") - - if rating is not None and (rating < 1 or rating > 5): - return web.json_response( - {"success": False, "error": "Rating must be between 1 and 5"}, - status=400, - ) - - with self.db.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, - ), - ) - conn.commit() - - if cursor.rowcount > 0: - return web.json_response( - {"success": True, "message": "Rating updated successfully"} - ) - else: - return web.json_response( - {"success": False, "error": "Prompt not found"}, status=404 - ) - - except ValueError: - return web.json_response( - {"success": False, "error": "Invalid prompt ID"}, status=400 - ) - except Exception as e: - self.logger.error(f"Update rating error: {e}") - return web.json_response( - {"success": False, "error": f"Failed to update rating: {str(e)}"}, - status=500, - ) - - async def add_prompt_tag(self, request): - """Add tag to prompt.""" - try: - prompt_id = int(request.match_info["prompt_id"]) - data = await request.json() - new_tag = data.get("tag", "").strip() - - if not new_tag: - return web.json_response( - {"success": False, "error": "Tag cannot be empty"}, status=400 - ) - - # Get current prompt - prompt = self.db.get_prompt_by_id(prompt_id) - if not prompt: - return web.json_response( - {"success": False, "error": "Prompt not found"}, status=404 - ) - - # Get current tags - current_tags = prompt.get("tags", []) - if not isinstance(current_tags, list): - current_tags = [] - - # Add new tag if not already present - if new_tag not in current_tags: - current_tags.append(new_tag) - - # Update database - with self.db.model.get_connection() as conn: - cursor = conn.execute( - "UPDATE prompts SET tags = ?, updated_at = ? WHERE id = ?", - ( - json.dumps(current_tags), - datetime.datetime.now(datetime.timezone.utc).isoformat(), - prompt_id, - ), - ) - conn.commit() - - return web.json_response( - {"success": True, "message": "Tag added successfully"} - ) - - except ValueError: - return web.json_response( - {"success": False, "error": "Invalid prompt ID"}, status=400 - ) - except Exception as e: - self.logger.error(f"Add tag error: {e}") - return web.json_response( - {"success": False, "error": f"Failed to add tag: {str(e)}"}, status=500 - ) - - async def add_tags_to_prompt(self, request): - """Add multiple tags to a single prompt.""" - try: - data = await request.json() - prompt_id = data.get("prompt_id") - new_tags = data.get("tags", []) - - if not prompt_id: - return web.json_response( - {"success": False, "error": "Prompt ID is required"}, status=400 - ) - - 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 - ) - - # Get current prompt - prompt = self.db.get_prompt_by_id(prompt_id) - if not prompt: - return web.json_response( - {"success": False, "error": "Prompt not found"}, status=404 - ) - - # Get current tags - current_tags = prompt.get("tags", []) - if not isinstance(current_tags, list): - current_tags = [] - - # Add new tags if not already present - tags_added = 0 - for new_tag in new_tags: - new_tag = new_tag.strip() - if new_tag and new_tag not in current_tags: - current_tags.append(new_tag) - tags_added += 1 - - # Update database if any tags were added - if tags_added > 0: - with self.db.model.get_connection() as conn: - cursor = conn.execute( - "UPDATE prompts SET tags = ?, updated_at = ? WHERE id = ?", - ( - json.dumps(current_tags), - datetime.datetime.now(datetime.timezone.utc).isoformat(), - prompt_id, - ), - ) - conn.commit() - - message = f"{tags_added} tag(s) added successfully" - if tags_added == 0: - message = "No new tags to add (all tags already exist)" - - return web.json_response( - {"success": True, "message": message, "tags_added": tags_added} - ) - - except ValueError: - return web.json_response( - {"success": False, "error": "Invalid prompt ID"}, status=400 - ) - except Exception as e: - self.logger.error(f"Add tags error: {e}") - return web.json_response( - {"success": False, "error": f"Failed to add tags: {str(e)}"}, status=500 - ) - - async def remove_prompt_tag(self, request): - """Remove tag from prompt.""" - try: - prompt_id = int(request.match_info["prompt_id"]) - data = await request.json() - tag_to_remove = data.get("tag", "").strip() - - # Get current prompt - prompt = self.db.get_prompt_by_id(prompt_id) - if not prompt: - return web.json_response( - {"success": False, "error": "Prompt not found"}, status=404 - ) - - # Get current tags - current_tags = prompt.get("tags", []) - if not isinstance(current_tags, list): - current_tags = [] - - # Remove tag if present - if tag_to_remove in current_tags: - current_tags.remove(tag_to_remove) - - # Update database - with self.db.model.get_connection() as conn: - cursor = conn.execute( - "UPDATE prompts SET tags = ?, updated_at = ? WHERE id = ?", - ( - json.dumps(current_tags), - datetime.datetime.now(datetime.timezone.utc).isoformat(), - prompt_id, - ), - ) - conn.commit() - - return web.json_response( - {"success": True, "message": "Tag removed successfully"} - ) - - except ValueError: - return web.json_response( - {"success": False, "error": "Invalid prompt ID"}, status=400 - ) - except Exception as e: - self.logger.error(f"Remove tag error: {e}") - return web.json_response( - {"success": False, "error": f"Failed to remove tag: {str(e)}"}, - status=500, - ) - - async def bulk_delete_prompts(self, request): - """Bulk delete prompts.""" - try: - data = await request.json() - prompt_ids = data.get("prompt_ids", []) - - if not prompt_ids: - return web.json_response( - {"success": False, "error": "No prompt IDs provided"}, status=400 - ) - - deleted_count = 0 - with self.db.model.get_connection() as conn: - for prompt_id in prompt_ids: - # First delete related images to avoid foreign key constraint - 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,) - ) - if cursor.rowcount > 0: - deleted_count += 1 - conn.commit() - - 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}") - return web.json_response( - {"success": False, "error": f"Failed to delete prompts: {str(e)}"}, - status=500, - ) - - async def bulk_add_tags(self, request): - """Bulk add tags to prompts.""" - try: - data = await request.json() - prompt_ids = data.get("prompt_ids", []) - new_tags = data.get("tags", []) - - if not prompt_ids or not new_tags: - return web.json_response( - {"success": False, "error": "No prompt IDs or tags provided"}, - status=400, - ) - - updated_count = 0 - with self.db.model.get_connection() as conn: - for prompt_id in prompt_ids: - # Get current tags - cursor = conn.execute( - "SELECT tags FROM prompts WHERE id = ?", (prompt_id,) - ) - row = cursor.fetchone() - if row: - current_tags = [] - if row["tags"]: - try: - current_tags = json.loads(row["tags"]) - if not isinstance(current_tags, list): - current_tags = [] - except: - current_tags = [] - - # Add new tags - for tag in new_tags: - if tag not in current_tags: - current_tags.append(tag) - - # Update database - cursor = conn.execute( - "UPDATE prompts SET tags = ?, updated_at = ? WHERE id = ?", - ( - json.dumps(current_tags), - datetime.datetime.now( - datetime.timezone.utc - ).isoformat(), - prompt_id, - ), - ) - if cursor.rowcount > 0: - updated_count += 1 - - conn.commit() - - 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}") - return web.json_response( - {"success": False, "error": f"Failed to add tags: {str(e)}"}, status=500 - ) - - async def bulk_set_category(self, request): - """Bulk set category for prompts.""" - try: - data = await request.json() - prompt_ids = data.get("prompt_ids", []) - category = data.get("category", "").strip() - - if not prompt_ids: - return web.json_response( - {"success": False, "error": "No prompt IDs provided"}, status=400 - ) - - updated_count = 0 - with self.db.model.get_connection() as conn: - for prompt_id in prompt_ids: - cursor = conn.execute( - "UPDATE prompts SET category = ?, updated_at = ? WHERE id = ?", - ( - category, - datetime.datetime.now(datetime.timezone.utc).isoformat(), - prompt_id, - ), - ) - if cursor.rowcount > 0: - updated_count += 1 - conn.commit() - - 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}") - return web.json_response( - {"success": False, "error": f"Failed to set category: {str(e)}"}, - status=500, - ) - - async def export_prompts(self, request): - """Export all prompts to JSON.""" - try: - # Get all prompts - prompts = self.db.search_prompts(limit=10000) - - # Create export data - export_data = { - "export_date": datetime.datetime.now(datetime.timezone.utc).isoformat(), - "total_prompts": len(prompts), - "prompts": prompts, - } - - # Return as JSON download - json_data = json.dumps(export_data, indent=2, ensure_ascii=False) - - return web.Response( - text=json_data, - content_type="application/json", - headers={ - "Content-Disposition": f'attachment; filename="prompt_manager_{datetime.datetime.now().strftime("%Y%m%d_%H%M%S")}.json"' - }, - ) - - except Exception as e: - self.logger.error(f"Export error: {e}") - return web.json_response( - {"success": False, "error": f"Failed to export prompts: {str(e)}"}, - status=500, - ) - - # Gallery-related endpoints - def _clean_nan_recursive(self, obj): - """Recursively clean NaN values from nested data structures.""" - if isinstance(obj, dict): - 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': - return None - else: - return obj - - async def get_prompt_images(self, request): - """Get all images for a specific prompt.""" - try: - prompt_id = request.match_info["prompt_id"] - images = self.db.get_prompt_images(prompt_id) - - # Clean up any NaN values that cause JSON parsing errors (recursive) - cleaned_images = [self._clean_nan_recursive(image) for image in images] - - # Additional fallback: convert to JSON string and clean NaN values manually - import json - import re - try: - 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) - - # Parse back to verify it's valid JSON - cleaned_data = json.loads(json_str) - - return web.json_response(cleaned_data) - 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 - }) - - except Exception as e: - self.logger.error(f"Get prompt images error: {e}") - 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)) - images = self.db.get_recent_images(limit) - - 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) - - async def get_all_images(self, request): - """Get all generated images with linked prompts.""" - try: - images = self.db.get_all_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) - - async def search_images(self, request): - """Search images by prompt text.""" - try: - query = request.query.get('q', '') - if not query: - return web.json_response({ - 'success': False, - 'error': 'Search query required' - }, status=400) - - images = self.db.search_images_by_prompt(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) - - async def get_output_images(self, request): - """Get all images from ComfyUI output folder.""" - try: - import os - from pathlib import Path - - # 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': [] - }) - - # Get pagination parameters - limit = int(request.query.get('limit', 100)) - offset = int(request.query.get('offset', 0)) - - # Find all media files (images and videos) - image_extensions = ['.png', '.jpg', '.jpeg', '.webp', '.gif'] - video_extensions = ['.mp4', '.webm', '.avi', '.mov', '.mkv', '.m4v', '.wmv'] - media_extensions = image_extensions + video_extensions - all_images = [] - - output_path = Path(output_dir) - thumbnails_dir = output_path / 'thumbnails' - - # Find all media files, excluding thumbnails directory - # Use a set to avoid duplicates on case-insensitive filesystems (Windows) - seen_paths = set() - for ext in media_extensions: - # Search for both lowercase and uppercase extensions - for pattern in [f"*{ext}", f"*{ext.upper()}"]: - for media_path in output_path.rglob(pattern): - # Skip files in thumbnails directory - if 'thumbnails' not in media_path.parts: - # Normalize path for deduplication (case-insensitive on Windows) - normalized_path = str(media_path).lower() - if normalized_path not in seen_paths: - seen_paths.add(normalized_path) - all_images.append(media_path) - - # Sort by modification time (newest first) - all_images.sort(key=lambda x: x.stat().st_mtime, reverse=True) - - # Apply pagination - paginated_images = all_images[offset:offset + limit] - - # Format media data (images and videos) - images = [] - for media_path in paginated_images: - try: - stat = media_path.stat() - # Create a relative path from output directory for the URL - rel_path = media_path.relative_to(output_path) - - # Determine media type - extension = media_path.suffix.lower() - is_video = extension in video_extensions - media_type = 'video' if is_video else 'image' - - # Check if thumbnail exists for this media - thumbnail_url = None - if thumbnails_dir.exists(): - # For videos, look for thumbnail with .jpg extension - # Use full relative path to support subdirectories - thumbnail_ext = '.jpg' if is_video else extension - # Preserve subdirectory structure in thumbnail path - 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="/")}' - - images.append({ - 'id': str(hash(str(media_path))), # Simple hash for ID - 'filename': media_path.name, - 'path': str(media_path), - 'relative_path': str(rel_path), - 'url': f'/prompt_manager/images/serve/{rel_path.as_posix()}', # Use forward slashes for URLs - 'thumbnail_url': thumbnail_url, # Thumbnail URL if exists - 'size': stat.st_size, - 'modified_time': stat.st_mtime, - 'extension': extension, - 'media_type': media_type, # 'image' or 'video' - 'is_video': is_video - }) - except Exception as e: - self.logger.error(f"Error processing media {media_path}: {e}") - continue - - 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) - - async def serve_image(self, request): - """Serve the actual image file.""" - try: - image_id = int(request.match_info["image_id"]) - image = self.db.get_image_by_id(image_id) - - if not image: - return web.json_response({'error': 'Image not found'}, status=404) - - import os - from pathlib import Path - - image_path = Path(image['image_path']) - if not image_path.exists(): - return web.json_response({'error': 'Image file not found'}, status=404) - - # Determine content type based on file extension - content_type = 'image/jpeg' - if image_path.suffix.lower() in ['.png']: - content_type = 'image/png' - elif image_path.suffix.lower() in ['.webp']: - content_type = 'image/webp' - elif image_path.suffix.lower() in ['.gif']: - content_type = 'image/gif' - - # Read and serve the file - with open(image_path, 'rb') as f: - file_data = f.read() - - return web.Response( - body=file_data, - content_type=content_type, - headers={ - 'Cache-Control': 'public, max-age=3600', - 'Content-Length': str(len(file_data)) - } - ) - - except ValueError: - return web.json_response({'error': 'Invalid image ID'}, status=400) - except Exception as e: - self.logger.error(f"Serve image error: {e}") - return web.json_response({'error': str(e)}, status=500) - - async def serve_output_image(self, request): - """Serve image file directly from ComfyUI output folder.""" - try: - import os - from pathlib import Path - - filepath = request.match_info["filepath"] - - # Find ComfyUI output directory - output_dir = self._find_comfyui_output_dir() - if not output_dir: - return web.json_response({'error': 'ComfyUI output directory not found'}, status=404) - - # Construct full image path - image_path = Path(output_dir) / filepath - - # Security check: make sure the path is within the output directory - try: - image_path = image_path.resolve() - output_path = Path(output_dir).resolve() - if not str(image_path).startswith(str(output_path)): - return web.json_response({'error': 'Access denied'}, status=403) - except Exception: - return web.json_response({'error': 'Invalid file path'}, status=400) - - if not image_path.exists(): - return web.json_response({'error': 'Image file not found'}, status=404) - - # Determine content type based on file extension - content_type = 'image/jpeg' - if image_path.suffix.lower() == '.png': - content_type = 'image/png' - elif image_path.suffix.lower() == '.webp': - content_type = 'image/webp' - elif image_path.suffix.lower() == '.gif': - content_type = 'image/gif' - - # Read and serve the file - with open(image_path, 'rb') as f: - file_data = f.read() - - return web.Response( - body=file_data, - content_type=content_type, - headers={ - 'Cache-Control': 'public, max-age=3600', - 'Content-Length': str(len(file_data)) - } - ) - - except Exception as e: - self.logger.error(f"Serve output image error: {e}") - return web.json_response({'error': str(e)}, status=500) - - async def generate_thumbnails(self, request): - """Generate thumbnails for all images and videos in the ComfyUI output directory. - - Creates optimized thumbnail versions of all media files in a separate - 'thumbnails' subdirectory. This process is safe and never modifies - original files. - - Safety Features: - - Read-only access to original files - - Writes only to separate thumbnails directory - - Uses PIL context managers for safe file handling - - All thumbnails clearly marked with '_thumb' suffix - - Proper error handling prevents corruption - - Supported Formats: - - Images: PNG, JPG, JPEG, WebP, GIF - - Videos: MP4, AVI, MOV, WMV (generates frame thumbnails) - - Query Parameters: - size (int, optional): Thumbnail size in pixels (default: 256) - quality (int, optional): JPEG quality 1-100 (default: 85) - overwrite (bool, optional): Regenerate existing thumbnails (default: false) - - Args: - request (aiohttp.web.Request): HTTP request with optional query parameters - - Returns: - aiohttp.web.Response: JSON response with structure: - { - \"success\": bool, - \"message\": str, - \"processed\": int, # Number of thumbnails created - \"skipped\": int, # Number of files skipped - \"errors\": int, # Number of processing errors - \"total_size\": str # Total size of created thumbnails - } - - Raises: - Returns 500 status with error details if thumbnail generation fails - - Example: - POST /prompt_manager/images/generate-thumbnails?size=512&quality=90 - """ - try: - import os - import time - from pathlib import Path - from PIL import Image - - # Get request parameters - data = await request.json() - quality = data.get('quality', 'medium') - report_progress = data.get('report_progress', False) - - # Map quality to size - 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) - - output_path = Path(output_dir) - thumbnails_dir = output_path / 'thumbnails' - thumbnails_dir.mkdir(exist_ok=True) - - # Find all media files (images and videos) - 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'] - media_extensions = image_extensions + video_extensions - media_files = [] - for root, dirs, files in os.walk(output_path): - # Skip thumbnails directory - if 'thumbnails' in Path(root).parts: - continue - for file in files: - if any(file.lower().endswith(ext) for ext in media_extensions): - media_files.append(Path(root) / file) - - total_images = len(media_files) - self.logger.info(f"Found {total_images} media files to process for thumbnails") - - if total_images == 0: - return web.json_response({ - 'success': True, - 'count': 0, - 'total_images': 0, - 'message': 'No media files found to process', - 'errors': [] - }) - - generated_count = 0 - skipped_count = 0 - errors = [] - start_time = time.time() - - for i, media_file in enumerate(media_files): - try: - # Create relative path for thumbnail - rel_path = media_file.relative_to(output_path) - - # Check if this is a video file - is_video = any(media_file.name.lower().endswith(ext) for ext in video_extensions) - - # For videos, always save thumbnail as .jpg - # Preserve subdirectory structure in thumbnail path - rel_path_no_ext = rel_path.with_suffix('') - if is_video: - 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}" - - # 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): - skipped_count += 1 - continue - - # Create thumbnail directory structure if needed - thumbnail_path.parent.mkdir(parents=True, exist_ok=True) - - # Generate thumbnail based on media type - if is_video: - # Generate video thumbnail - 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}") - else: - # Generate image thumbnail - with Image.open(media_file) as img: - # Convert to RGB if necessary (for PNG with transparency) - if img.mode in ('RGBA', 'LA', 'P'): - img = img.convert('RGB') - - # Create thumbnail maintaining aspect ratio - img.thumbnail(thumbnail_size, Image.Resampling.LANCZOS) - - # Save thumbnail - 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 - - # Log progress every 100 images or at 10% intervals - 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") - - except Exception as 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") - - return web.json_response({ - '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) - }) - - except ImportError: - 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) - - async def generate_thumbnails_with_progress(self, request): - """Generate thumbnails with Server-Sent Events progress updates.""" - try: - import os - import time - import json - import asyncio - from pathlib import Path - from PIL import Image - - # Parse query parameters - quality = request.query.get('quality', 'medium') - - # Map quality to size - 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': '*', - } - ) - await response.prepare(request) - - async def send_progress(event_type, data): - """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 asyncio.sleep(0.01) # Small delay to ensure delivery - except Exception as e: - self.logger.warning(f"Failed to send SSE message: {e}") - - try: - # 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' - }) - return response - - output_path = Path(output_dir) - 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...' - }) - - 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'] - media_extensions = image_extensions + video_extensions - media_files = [] - scanned_dirs = 0 - - for root, dirs, files in os.walk(output_path): - # Skip thumbnails directory - if 'thumbnails' in Path(root).parts: - continue - - scanned_dirs += 1 - # Send scanning progress for every directory - if scanned_dirs % 5 == 0: # Update every 5 directories - 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") - - 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)) - 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 - }) - - 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' - }) - return response - - generated_count = 0 - skipped_count = 0 - errors = [] - start_time = time.time() - - # Process images with progress updates - # SAFETY: We only READ from original images, NEVER modify them - for i, media_file in enumerate(media_files): - try: - # SAFETY: Verify we're only reading from original media files - if not media_file.exists() or not media_file.is_file(): - continue - - # Check if this is a video file - is_video = any(media_file.name.lower().endswith(ext) for ext in video_extensions) - - # Create relative path for thumbnail (in separate thumbnails directory) - rel_path = media_file.relative_to(output_path) - - # For videos, always save thumbnail as .jpg - # Preserve subdirectory structure in thumbnail path - rel_path_no_ext = rel_path.with_suffix('') - if is_video: - 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}" - - # SAFETY: 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}") - continue - except Exception as 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): - skipped_count += 1 - else: - # Create thumbnail directory structure if needed (only within thumbnails dir) - thumbnail_path.parent.mkdir(parents=True, exist_ok=True) - - # Generate thumbnail based on media type - if is_video: - # Generate video thumbnail - 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}") - else: - # Generate image thumbnail (READ-ONLY operation on original) - # SAFETY: Image.open() with context manager ensures read-only access - with Image.open(media_file) as img: - # SAFETY: Work with a copy of the image data, never modify original - # Convert to RGB if necessary (for PNG with transparency) - if img.mode in ('RGBA', 'LA', 'P'): - img = img.convert('RGB') - - # Create thumbnail maintaining aspect ratio - # SAFETY: thumbnail() modifies the in-memory copy, not the original file - img.thumbnail(thumbnail_size, Image.Resampling.LANCZOS) - - # Save thumbnail to separate location - # SAFETY: Only write to our thumbnails directory, never touch originals - 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 more frequently for better feedback - # Update every 5 images or every 1% for large sets - 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 - - # Include more detailed file info - 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' - } - - 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'] - }) - - # Log progress - if (i + 1) % 50 == 0: # Log every 50 files - 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)}" - errors.append(error_msg) - self.logger.warning(error_msg) - - # Send error notification for immediate feedback - if len(errors) <= 5: # Only send first 5 errors to avoid spam - await send_progress('status', { - 'phase': 'processing', - 'message': f'Error processing {media_file.name}: {str(e)}' - }) - - elapsed_time = time.time() - start_time - - # Send detailed completion event - completion_message = f'Successfully generated {generated_count} new thumbnails, skipped {skipped_count} existing' - if errors: - completion_message += f' ({len(errors)} errors occurred)' - - await send_progress('complete', { - 'count': generated_count, - 'skipped': skipped_count, - 'total_images': total_images, - 'errors': errors[:10], # Limit errors to first 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") - - except Exception as 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) - 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) - - def _generate_video_thumbnail(self, video_path, thumbnail_path, thumbnail_size): - """ - Generate thumbnail from video file. - Returns True if successful, False otherwise. - """ - try: - # Try using OpenCV first (most reliable) - try: - import cv2 - - # Open video - cap = cv2.VideoCapture(str(video_path)) - if not cap.isOpened(): - return False - - # Get frame from 10% into the video (avoid black intro frames) - frame_count = int(cap.get(cv2.CAP_PROP_FRAME_COUNT)) - target_frame = max(1, int(frame_count * 0.1)) - cap.set(cv2.CAP_PROP_POS_FRAMES, target_frame) - - # Read frame - ret, frame = cap.read() - cap.release() - - if not ret or frame is None: - return False - - # Convert BGR to RGB - frame_rgb = cv2.cvtColor(frame, cv2.COLOR_BGR2RGB) - - # Convert to PIL Image and create thumbnail - from PIL import Image - img = Image.fromarray(frame_rgb) - img.thumbnail(thumbnail_size, Image.Resampling.LANCZOS) - - # Save thumbnail - img.save(thumbnail_path, 'JPEG', quality=85, optimize=True) - self.logger.debug(f"Generated video thumbnail using OpenCV: {thumbnail_path}") - return True - - except ImportError: - # OpenCV not available, try ffmpeg - pass - - # Fallback to ffmpeg - try: - import subprocess - - # Use ffmpeg to extract frame at 10% duration - cmd = [ - 'ffmpeg', '-i', str(video_path), - '-ss', '00:00:01', # Skip first second to avoid black frames - '-vframes', '1', - '-s', f"{thumbnail_size[0]}x{thumbnail_size[1]}", - '-y', # Overwrite output - 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}") - return True - else: - self.logger.warning(f"ffmpeg failed for {video_path}: {result.stderr}") - - except (ImportError, subprocess.TimeoutExpired, FileNotFoundError): - # ffmpeg not available - pass - - # Last resort: create a placeholder thumbnail - try: - from PIL import Image, ImageDraw, ImageFont - - # Create a placeholder image - img = Image.new('RGB', thumbnail_size, color=(50, 50, 50)) - draw = ImageDraw.Draw(img) - - # Add play button icon - center_x, center_y = thumbnail_size[0] // 2, thumbnail_size[1] // 2 - triangle_size = min(thumbnail_size) // 4 - - # Draw play triangle - 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) - ] - draw.polygon(points, fill=(255, 255, 255)) - - # Add text - try: - # Try to get a font - font = ImageFont.load_default() - text = "VIDEO" - bbox = draw.textbbox((0, 0), text, font=font) - text_width = bbox[2] - bbox[0] - text_height = bbox[3] - bbox[1] - draw.text( - (center_x - text_width//2, center_y + triangle_size//2 + 10), - text, fill=(255, 255, 255), font=font - ) - except: - pass - - 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}") - return False - - except Exception as e: - self.logger.error(f"Video thumbnail generation failed for {video_path}: {e}") - return False - - async def clear_thumbnails(self, request): - """Safely clear only our generated thumbnails, never touch original images.""" - try: - import os - import shutil - from pathlib import Path - - # 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) - - output_path = Path(output_dir) - thumbnails_dir = output_path / 'thumbnails' - - # SAFETY: Only operate within our thumbnails directory - if not thumbnails_dir.exists(): - return web.json_response({ - 'success': True, - 'message': 'No thumbnails directory found - nothing to clear', - 'cleared_files': 0 - }) - - # SAFETY: Verify this is actually our thumbnails directory - try: - thumbnails_dir_resolved = thumbnails_dir.resolve() - output_path_resolved = output_path.resolve() - - # Ensure thumbnails dir is within output dir and named 'thumbnails' - 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) - - # Count files before deletion - cleared_count = 0 - cleared_size = 0 - - # SAFETY: Only delete files within thumbnails directory that match our naming pattern - for root, dirs, files in os.walk(thumbnails_dir): - for file in files: - file_path = Path(root) / file - - # SAFETY: Additional check - only delete files with '_thumb' in name - 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() # Delete the file - cleared_count += 1 - 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}") - - # SAFETY: Remove empty directories within thumbnails folder (but not the main thumbnails dir) - try: - for root, dirs, files in os.walk(thumbnails_dir, topdown=False): - if root != str(thumbnails_dir): # Don't remove the main thumbnails directory - try: - Path(root).rmdir() # Only removes if empty - except OSError: - pass # Directory not empty, that's fine - except Exception as e: - self.logger.debug(f"Directory cleanup info: {e}") - - # Convert size to human readable format - def format_size(bytes_size): - 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)})") - - 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) - - 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', {}) - - if not prompt_id or not image_path: - return web.json_response({ - 'success': False, - 'error': 'prompt_id and image_path are required' - }, status=400) - - # Check if image file exists - import os - if not os.path.exists(image_path): - return web.json_response({ - 'success': False, - 'error': 'Image file not found' - }, status=404) - - # Link image to prompt - image_id = 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' - }) - - except Exception as e: - self.logger.error(f"Link image error: {e}") - 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: - import urllib.parse - import os - import json - from pathlib import Path - - # Get the image path from URL - 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) - - # Convert relative path to absolute if needed - if not os.path.isabs(image_path): - # If it's a relative path from ComfyUI output, make it absolute - output_dir = self._find_comfyui_output_dir() - if output_dir: - image_path = str(Path(output_dir) / image_path) - - # Look up the image in generated_images table - try: - with self.db.model.get_connection() as conn: - cursor = conn.execute( - """SELECT gi.prompt_id, p.text, p.category, p.tags, p.rating, p.notes, - gi.workflow_data, gi.prompt_metadata, gi.generation_time - 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)}') - ) - result = cursor.fetchone() - - if result: - # Convert to dict - prompt_data = { - 'prompt_id': result[0], - 'text': result[1], - 'category': result[2], - 'tags': json.loads(result[3]) if result[3] else [], - 'rating': result[4], - 'notes': result[5], - 'workflow_data': json.loads(result[6]) if result[6] else None, - 'prompt_metadata': json.loads(result[7]) if result[7] else None, - 'generation_time': result[8], - 'image_path': image_path - } - - return web.json_response({ - 'success': True, - 'prompt': prompt_data - }) - else: - # No linked prompt found - this is normal for many images - 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) - - except Exception as e: - self.logger.error(f"Get image prompt error: {e}") - return web.json_response({ - 'success': False, - 'error': str(e) - }, status=500) - - async def delete_image(self, request): - """Delete an image record.""" - try: - image_id = int(request.match_info["image_id"]) - success = self.db.delete_image(image_id) - - if success: - return web.json_response({ - 'success': True, - 'message': 'Image deleted successfully' - }) - else: - return web.json_response({ - 'success': False, - 'error': 'Image not found' - }, status=404) - - except ValueError: - return web.json_response({'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) - - # Diagnostic endpoints - async def run_diagnostics(self, request): - """Run comprehensive system diagnostics and health checks. - - Performs various system health checks including database connectivity, - file system access, ComfyUI integration status, and configuration - validation. Useful for troubleshooting and system monitoring. - - Diagnostic Checks: - - Database connection and integrity - - ComfyUI output directory detection - - File system permissions - - Configuration validation - - Memory and performance metrics - - Extension loading status - - Args: - request (aiohttp.web.Request): HTTP request object - - Returns: - aiohttp.web.Response: JSON response with diagnostic results: - { - "success": bool, - "diagnostics": { - "database": Dict, # Database health info - "filesystem": Dict, # File system status - "comfyui": Dict, # ComfyUI integration status - "config": Dict, # Configuration validation - "performance": Dict # Performance metrics - }, - "summary": str, # Overall system status - "issues": List[str] # List of identified issues - } - - Example: - POST /prompt_manager/diagnostics - """ - try: - # Simple diagnostics without importing complex modules - import os - import sqlite3 - - results = {} - - # Check database - try: - db_path = "prompts.db" - if os.path.exists(db_path): - 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'] - - # Check if 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 - - if has_images_table: - 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 - } - else: - 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)}' - } - - # Check dependencies - dependencies = {} - try: - import watchdog - dependencies['watchdog'] = True - except ImportError: - dependencies['watchdog'] = False - - try: - from PIL import Image - dependencies['PIL'] = True - except ImportError: - dependencies['PIL'] = False - - try: - import sqlite3 - dependencies['sqlite3'] = True - except ImportError: - dependencies['sqlite3'] = False - - results['dependencies'] = { - 'status': 'ok' if all(dependencies.values()) else 'error', - 'dependencies': dependencies - } - - # Check output directories - output_dirs = [] - potential_dirs = ["output", "../output", "../../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) - - 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 - } - else: - 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)}' - } - - 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) - - 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') - - if not prompt_id: - return web.json_response({ - 'success': False, - 'error': 'prompt_id is required' - }, status=400) - - # Test linking directly using the database manager - test_metadata = { - 'file_info': { - 'size': 1024000, - 'dimensions': [512, 512], - 'format': 'PNG' - }, - 'workflow': {'test': True}, - 'prompt': {'test_prompt': 'This is a test image'} - } - - try: - image_id = self.db.link_image_to_prompt( - prompt_id=str(prompt_id), - image_path=test_image_path, - 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}' - } - }) - except Exception as 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) - - async def run_maintenance(self, request): - """Perform comprehensive database maintenance and optimization. - - Executes a series of database maintenance operations to optimize - performance, clean up orphaned data, and ensure database integrity. - This is a resource-intensive operation that should be run during - low-traffic periods. - - Maintenance Operations: - - Remove duplicate prompts based on content hash - - Clean up orphaned image references - - Remove missing file records from database - - Optimize database with VACUUM operation - - Update database statistics - - Validate data integrity - - Clean up temporary files - - Query Parameters: - full (bool, optional): Perform full maintenance including VACUUM - cleanup_images (bool, optional): Clean up missing image references - optimize (bool, optional): Run database optimization - - Args: - request (aiohttp.web.Request): HTTP request with optional parameters - - Returns: - aiohttp.web.Response: JSON response with maintenance results: - { - "success": bool, - "operations": { - "duplicates_removed": int, - "orphaned_images_cleaned": int, - "missing_files_removed": int, - "database_optimized": bool, - "integrity_check_passed": bool - }, - "before_stats": Dict, # Database stats before maintenance - "after_stats": Dict, # Database stats after maintenance - "duration": float, # Maintenance duration in seconds - "message": str - } - - Raises: - Returns 500 status with error details if maintenance fails - - Example: - POST /prompt_manager/maintenance?full=true&cleanup_images=true - """ - try: - data = await request.json() if request.content_type == 'application/json' else {} - operations = data.get('operations', ['cleanup_duplicates', 'vacuum', 'cleanup_orphaned_images']) - - results = {} - - # Clean up duplicate prompts - 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' - } - except Exception as e: - results['cleanup_duplicates'] = { - 'success': False, - 'error': str(e), - 'message': 'Failed to cleanup duplicates' - } - - # Vacuum database - if 'vacuum' in operations: - try: - self.db.model.vacuum_database() - 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' - } - - # Clean up orphaned image records - 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' - } - except Exception as e: - results['cleanup_orphaned_images'] = { - 'success': False, - 'error': str(e), - 'message': 'Failed to cleanup orphaned images' - } - - # Check for potential duplicate hashes - if 'check_hash_duplicates' in operations: - try: - with self.db.model.get_connection() as conn: - cursor = conn.execute(""" - SELECT hash, COUNT(*) as count - FROM prompts - WHERE hash IS NOT NULL - GROUP BY hash - HAVING COUNT(*) > 1 - """) - hash_duplicates = cursor.fetchall() - - 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' - } - - # Get database statistics - if 'statistics' in operations: - try: - db_info = self.db.model.get_database_info() - 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' - } - - # Prune orphaned prompts (prompts with no linked images) - # Note: Prompts with the __protected__ tag are excluded from orphan cleanup - # This allows users to add prompts manually without images that won't be deleted - if 'prune_orphaned_prompts' in operations: - try: - with self.db.model.get_connection() as conn: - # Find prompts that have no images linked to them - # Exclude prompts with the __protected__ tag - cursor = conn.execute(""" - SELECT p.id - FROM prompts p - LEFT JOIN generated_images gi ON p.id = gi.prompt_id - WHERE gi.prompt_id IS NULL - AND (p.tags IS NULL OR p.tags NOT LIKE '%"__protected__"%') - """) - orphaned_prompts = cursor.fetchall() - orphaned_count = len(orphaned_prompts) - - if orphaned_count > 0: - # Delete the orphaned prompts - orphaned_ids = [row['id'] for row in orphaned_prompts] - placeholders = ','.join(['?'] * len(orphaned_ids)) - cursor = conn.execute(f"DELETE FROM prompts WHERE id IN ({placeholders})", orphaned_ids) - conn.commit() - removed_count = cursor.rowcount - else: - removed_count = 0 - - 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' - } - - # Check for consistency issues - if 'check_consistency' in operations: - try: - consistency_issues = [] - - with self.db.model.get_connection() as conn: - # Check for prompts with invalid JSON in tags - cursor = conn.execute("SELECT id, tags FROM prompts WHERE tags IS NOT NULL") - for row in cursor.fetchall(): - try: - if row['tags']: - json.loads(row['tags']) - except json.JSONDecodeError: - consistency_issues.append(f"Prompt {row['id']} has invalid JSON in tags") - - # Check for orphaned foreign key references - cursor = conn.execute(""" - SELECT gi.id, gi.prompt_id - FROM generated_images gi - LEFT JOIN prompts p ON gi.prompt_id = p.id - WHERE p.id IS NULL - """) - orphaned_refs = cursor.fetchall() - for ref in orphaned_refs: - consistency_issues.append(f"Image {ref['id']} references non-existent prompt {ref['prompt_id']}") - - results['check_consistency'] = { - 'success': True, - 'issues_found': len(consistency_issues), - 'issues': consistency_issues[:10], # Limit to first 10 issues - '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' - } - - # Overall success status - 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' - }) - - 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) - - async def backup_database(self, request): - """ - Backup the entire prompts.db database file. - GET /prompt_manager/backup - """ - try: - import os - import shutil - import tempfile - from pathlib import Path - - # Get the database path - db_path = "prompts.db" - - if not os.path.exists(db_path): - return web.json_response({ - 'success': False, - 'error': 'Database file not found' - }, status=404) - - # Create a temporary copy of the database - with tempfile.NamedTemporaryFile(delete=False, suffix='.db') as temp_file: - temp_path = temp_file.name - - # Copy the database file - shutil.copy2(db_path, temp_path) - - # Read the file content - with open(temp_path, 'rb') as f: - file_data = f.read() - - # Clean up temporary file - os.unlink(temp_path) - - # Generate filename with timestamp - timestamp = datetime.datetime.now().strftime("%Y%m%d_%H%M%S") - filename = f"prompts_backup_{timestamp}.db" - - return web.Response( - body=file_data, - content_type='application/octet-stream', - headers={ - '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) - - async def restore_database(self, request): - """ - Restore the prompts.db database from uploaded file. - POST /prompt_manager/restore - """ - try: - import os - import shutil - import tempfile - import sqlite3 - from pathlib import Path - - # Get the uploaded file - 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) - - # Read the uploaded file content - file_data = await field.read() - - if not file_data: - return web.json_response({ - 'success': False, - 'error': 'Uploaded file is empty' - }, status=400) - - # Create a temporary file to validate the database - with tempfile.NamedTemporaryFile(delete=False, suffix='.db') as temp_file: - temp_path = temp_file.name - temp_file.write(file_data) - - try: - # Validate that it's a valid SQLite database with expected structure - with sqlite3.connect(temp_path) as conn: - conn.row_factory = sqlite3.Row - - # Check if it has the prompts table - 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") - - # Check basic structure of prompts table - cursor = conn.execute("PRAGMA table_info(prompts)") - 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}") - - # Get basic stats for validation - cursor = conn.execute("SELECT COUNT(*) as count FROM prompts") - prompt_count = cursor.fetchone()['count'] - - # If validation passes, backup current database and restore - db_path = "prompts.db" - backup_path = f"{db_path}.backup_{datetime.datetime.now().strftime('%Y%m%d_%H%M%S')}" - - # Create backup of current database if it exists - if os.path.exists(db_path): - shutil.copy2(db_path, backup_path) - self.logger.info(f"Current database backed up to: {backup_path}") - - # Replace current database with uploaded one - shutil.copy2(temp_path, db_path) - - # Reinitialize the database connection - 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 - }) - - except sqlite3.Error as e: - 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) - finally: - # Clean up temporary file - 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) - - async def scan_images(self, request): - """ - Scan ComfyUI output images for prompt metadata and add them to the database. - Streams progress updates to the client. - """ - import json - import asyncio - from pathlib import Path - from aiohttp import web - - async def stream_response(): - try: - self.logger.info("Starting image scan operation") - - # Clear any active prompt tracking timers to avoid conflicts - # Note: For now we'll skip this since we don't have easy access to the tracker instance - # This is mainly important if scans are run during active generation, which is rare - self.logger.info("Starting scan (timer clearing not implemented yet)") - - # Find ComfyUI output directory - output_dir = self._find_comfyui_output_dir() - if not output_dir: - self.logger.error("ComfyUI output directory not found") - yield f"data: {json.dumps({'type': 'error', 'message': 'ComfyUI output directory not found'})}\n\n" - return - - yield f"data: {json.dumps({'type': 'progress', 'progress': 0, 'status': 'Scanning for PNG files...', 'processed': 0, 'found': 0})}\n\n" - - # Find all PNG files - png_files = list(Path(output_dir).rglob("*.png")) - total_files = len(png_files) - - if total_files == 0: - yield f"data: {json.dumps({'type': 'complete', 'processed': 0, 'found': 0, 'added': 0})}\n\n" - return - - yield f"data: {json.dumps({'type': 'progress', 'progress': 5, 'status': f'Found {total_files} PNG files to process...', 'processed': 0, 'found': 0})}\n\n" - - processed_count = 0 - found_count = 0 - added_count = 0 - linked_count = 0 # Count of images linked to existing prompts - - for i, png_file in enumerate(png_files): - try: - # Extract metadata from PNG - metadata = self._extract_comfyui_metadata(str(png_file)) - processed_count += 1 - - if metadata: - self.logger.debug(f"Found metadata in {os.path.basename(png_file)}: {list(metadata.keys())}") - - # Parse ComfyUI prompt data - 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'))}") - - # Check if we found any meaningful prompt data - if parsed_data.get('prompt') or parsed_data.get('parameters'): - found_count += 1 - - # Extract readable prompt text - prompt_text = self._extract_readable_prompt(parsed_data) - - # Debug: print what we found and its type - if prompt_text: - 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())}") - - # Ensure prompt_text is a string - if prompt_text and not isinstance(prompt_text, str): - 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 to save to database (will skip duplicates) - try: - # Generate hash for duplicate detection - try: - from ..utils.hashing import generate_prompt_hash - except ImportError: - import sys - current_dir = os.path.dirname(os.path.dirname(os.path.abspath(__file__))) - sys.path.insert(0, current_dir) - from utils.hashing import generate_prompt_hash - - prompt_hash = generate_prompt_hash(prompt_text.strip()) - self.logger.debug(f"Generated hash for prompt: {prompt_hash[:16]}...") - - # Check if prompt already exists - existing = 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)}") - # Link image to existing prompt - try: - 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']}") - except Exception as e: - self.logger.error(f"Failed to link image {png_file} to existing prompt: {e}") - else: - # Save new prompt - self.logger.debug(f"Saving new prompt from {os.path.basename(png_file)}") - prompt_id = self.db.save_prompt( - text=prompt_text.strip(), - category='scanned', - tags=['auto-scanned'], - notes=f'Auto-scanned from {os.path.basename(png_file)}', - prompt_hash=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)}") - - # Link image to new prompt - try: - 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}") - except Exception as 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") - - except Exception as e: - self.logger.error(f"Failed to save prompt from {png_file}: {e}") - - # Update progress every 10 files or so - if i % 10 == 0 or i == total_files - 1: - progress = int((i + 1) / total_files * 100) - status = f"Processing file {i + 1}/{total_files}..." - - yield f"data: {json.dumps({'type': 'progress', 'progress': progress, 'status': status, 'processed': processed_count, 'found': found_count})}\n\n" - - # Small delay to allow UI updates - await asyncio.sleep(0.01) - - except Exception as e: - self.logger.error(f"Error processing {png_file}: {e}") - continue - - # Send completion message - 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: - self.logger.error(f"Scan error: {e}") - import traceback - self.logger.error(f"Scan error traceback: {traceback.format_exc()}") - yield f"data: {json.dumps({'type': 'error', 'message': str(e)})}\n\n" - - # Return streaming response - response = web.StreamResponse( - status=200, - reason='OK', - headers={ - '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_eof() - return response - - def _find_comfyui_output_dir(self): - """Locate the ComfyUI output directory using multiple detection strategies. - - Attempts to find the ComfyUI output directory by checking various - possible locations relative to the current file and common installation - patterns. Handles different ComfyUI installation types and structures. - - Detection Strategy: - 1. Look for 'output' directory in parent directories (up to 10 levels) - 2. Check common ComfyUI installation patterns - 3. Verify directory contains typical ComfyUI subdirectories - 4. Return first valid match found - - Returns: - str or None: Absolute path to ComfyUI output directory, or None - if no valid directory is found - - Example: - output_dir = api._find_comfyui_output_dir() - if output_dir: - print(f"Found ComfyUI output at: {output_dir}") - else: - print("ComfyUI output directory not found") - """ - import os - from pathlib import Path - - current_file = Path(__file__).resolve() - self.logger.debug(f"Starting ComfyUI output search from: {current_file}") - - # Method 1: Search upward from current file location - current_dir = current_file.parent - max_depth = 10 # Prevent infinite loops - - for i in range(max_depth): - # Check if current directory contains ComfyUI markers - 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}") - return str(output_dir) - - # Move up one directory - parent = current_dir.parent - if parent == current_dir: # Reached filesystem root - break - current_dir = parent - - # Method 2: Try common installation patterns relative to this file - base_dir = current_file.parent # /path/to/ComfyUI/custom_nodes/ComfyUI_PromptManager/py - possible_paths = [ - # Standard custom node installation: custom_nodes/ComfyUI_PromptManager/py -> output - base_dir.parent.parent.parent / "output", # ../../../output - # Nested custom node: custom_nodes/promptmanager/py -> output - base_dir.parent.parent / "output", # ../../output - # Direct in ComfyUI root - base_dir.parent / "output", # ../output - base_dir / "output", # ./output - ] - - # Method 3: Add common ComfyUI installation locations - common_locations = [ - Path.home() / "ComfyUI" / "output", - Path.cwd() / "output", - Path.cwd() / ".." / "output", - Path.cwd() / ".." / ".." / "output", - ] - - all_paths = possible_paths + common_locations - - for path in all_paths: - try: - abs_path = path.resolve() - if abs_path.exists() and abs_path.is_dir(): - self.logger.debug(f"Found ComfyUI output directory: {abs_path}") - return str(abs_path) - except (OSError, RuntimeError): - continue # Skip invalid paths - - self.logger.warning("ComfyUI output directory not found. Searched paths:") - for path in all_paths: - try: - self.logger.warning(f" - {path.resolve()} (exists: {path.exists()})") - except (OSError, RuntimeError): - self.logger.warning(f" - {path} (invalid path)") - - return None - - def _extract_comfyui_metadata(self, image_path): - """Extract ComfyUI workflow metadata from PNG image files. - - Reads embedded metadata from PNG files generated by ComfyUI, - extracting workflow information, parameters, and generation details - stored in PNG text chunks. - - ComfyUI stores metadata in standard PNG text chunks: - - 'workflow': Complete node graph workflow data - - 'prompt': Simplified prompt/parameter data - - Custom fields: Additional generation parameters - - Args: - image_path (str): Path to the PNG image file to analyze - - Returns: - Dict[str, Any]: Dictionary containing extracted metadata: - - 'workflow': Raw workflow JSON data (if present) - - 'prompt': Simplified prompt data (if present) - - Additional custom fields from PNG text chunks - - Empty dict if no metadata found or file is not PNG - - Raises: - Exception: If file cannot be opened or read (handled gracefully, - returns empty dict) - - Example: - metadata = api._extract_comfyui_metadata('output/image_001.png') - if 'workflow' in metadata: - print("Found ComfyUI workflow data") - if 'prompt' in metadata: - print(f"Prompt data: {metadata['prompt']}") - """ - try: - with Image.open(image_path) as img: - metadata = {} - if hasattr(img, 'text'): - for key, value in img.text.items(): - metadata[key] = value - return metadata - except Exception as e: - self.logger.error(f"Error reading {image_path}: {e}") - return {} - - 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 - } - - # Check for A1111 style parameters first (like parse-metadata.py) - if "parameters" in metadata: - params = metadata["parameters"] - lines = params.splitlines() - if lines: - result['positive_prompt'] = lines[0].strip() - for line in lines: - if line.lower().startswith("negative prompt:"): - result['negative_prompt'] = line.split(":", 1)[1].strip() - break - # Store raw parameters too - result['parameters']['parameters'] = params - - # If no A1111 format found, proceed with ComfyUI parsing - if result['positive_prompt'] is None: - # Check for direct prompt field - if 'prompt' in metadata: - try: - prompt_data = json.loads(metadata['prompt']) - result['prompt'] = prompt_data - except json.JSONDecodeError: - result['prompt'] = metadata['prompt'] - - # Check for workflow - if 'workflow' in metadata: - try: - workflow_data = json.loads(metadata['workflow']) - result['workflow'] = workflow_data - except json.JSONDecodeError: - result['workflow'] = metadata['workflow'] - - # Check for other common ComfyUI fields - 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]) - except json.JSONDecodeError: - result['parameters'][field] = metadata[field] - - return result - - def _extract_readable_prompt(self, parsed_data): - """Extract human-readable prompt text from ComfyUI/A1111 data using improved logic.""" - import json - - # Helper function to convert any value to string safely - def safe_to_string(value): - if isinstance(value, str): - return value - elif isinstance(value, list): - # Join list elements with spaces - 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'] - - # Check if prompt is already a string - 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']) - - 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) - if positive_prompt: - return positive_prompt - - # Check workflow data if available - workflow_data = parsed_data.get('workflow') - if isinstance(workflow_data, dict): - 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']) - - return None - - def _get_node_inputs(self, node): - """ - Safely get inputs from a node, handling both dict and list formats. - Returns a normalized dict format for consistent access. - - Old format: inputs is a dict with direct key-value pairs - inputs = {"text": "my prompt", "seed": 123} - - New format: inputs is a list of connection objects - inputs = [ - {"name": "text", "type": "STRING", "link": null, "widget": {"name": "text"}}, - {"name": "clip", "type": "CLIP", "link": 11} - ] - """ - if not isinstance(node, dict): - return {} - - inputs = node.get("inputs", {}) - - # If inputs is already a dict, return it - if isinstance(inputs, dict): - return inputs - - # If inputs is a list, convert to dict format - if isinstance(inputs, list): - inputs_dict = {} - for input_item in inputs: - if isinstance(input_item, dict) and "name" in input_item: - name = input_item["name"] - # For now, just mark that this input exists - # The actual value might be in widgets_values - inputs_dict[name] = input_item - return inputs_dict - - # If inputs is neither dict nor list, return empty dict - return {} - - def _find_text_in_node(self, node): - """ - Try to find text content in a node using various strategies. - Handles both old and new workflow formats. - """ - if not isinstance(node, dict): - return None - - # Strategy 1: Check normalized inputs for 'text' field - inputs = self._get_node_inputs(node) - if "text" in inputs and isinstance(inputs["text"], str): - return inputs["text"] - - # 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' - ] - - 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: - # First widget is usually the text for these nodes - if isinstance(widgets_values[0], str) and widgets_values[0].strip(): - return widgets_values[0] - - return None - - def _extract_positive_prompt_from_comfyui_data(self, data): - """Extract positive prompt from ComfyUI data, handling both old and new formats.""" - if not isinstance(data, dict): - return None - - # Build nodes dictionary - nodes_by_id = {} - if "nodes" in data: - # Handle nodes array format - for node in data["nodes"]: - if isinstance(node, dict): - nid = node.get("id") - if nid is not None: - nodes_by_id[nid] = node - else: - # Handle flat dictionary format (node_id -> node_data) - for nid_str, node in data.items(): - try: - nid = int(nid_str) - except: - nid = nid_str - if isinstance(node, dict): - if "id" in node: - nid = node["id"] - nodes_by_id[nid] = node - - if not nodes_by_id: - return None - - # First, try to find positive/negative connection pattern - pos_id = None - for node in nodes_by_id.values(): - if isinstance(node, dict): - inputs = self._get_node_inputs(node) # Use our safe function - if "positive" in inputs and "negative" in inputs: - try: - # Handle both old format (direct value) and new format (connection object) - pos_input = inputs["positive"] - if isinstance(pos_input, list) and len(pos_input) > 0: - pos_id = int(pos_input[0]) - break - except: - continue - - # Get text from the positive node - if pos_id is not None and pos_id in nodes_by_id: - text_val = self._find_text_in_node(nodes_by_id[pos_id]) - if text_val: - return text_val - - # Fallback: find any text encoder node with text content - # Collect all text encoder nodes - text_nodes = [] - for node in nodes_by_id.values(): - if isinstance(node, dict): - class_type = node.get('class_type', node.get('type', '')) - - # Check if this is a text encoder node - 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: - # Prioritize non-negative prompts - text_nodes.insert(0, text_val) - else: - text_nodes.append(text_val) - - # Return the first positive-looking prompt - if text_nodes: - return text_nodes[0] - - return None - - # Logging API endpoints - async def get_logs(self, request): - """ - Get recent log entries. - GET /prompt_manager/logs?limit=100&level=INFO - """ - try: - # Import logger here to avoid circular imports - try: - from ..utils.logging_config import get_logger_manager - 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_manager - - logger_manager = get_logger_manager() - - # Get query parameters - limit = int(request.query.get('limit', 100)) - level = request.query.get('level', None) - - # Validate limit - if limit > 1000: - limit = 1000 - elif limit < 1: - limit = 1 - - # Get logs from memory buffer - 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 - }) - - except Exception as e: - self.logger.error(f"Get logs error: {e}") - 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. - GET /prompt_manager/logs/files - """ - try: - try: - from ..utils.logging_config import get_logger_manager - 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_manager - - logger_manager = get_logger_manager() - log_files = logger_manager.get_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) - - async def download_log_file(self, request): - """ - Download a specific log file. - GET /prompt_manager/logs/download/{filename} - """ - try: - filename = request.match_info['filename'] - - try: - from ..utils.logging_config import get_logger_manager - 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_manager - - logger_manager = get_logger_manager() - - # Validate filename for security - if not filename or '..' in filename or '/' in filename or '\\' in filename: - return web.json_response({ - 'success': False, - 'error': 'Invalid filename' - }, status=400) - - # Get the file content - 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) - - # Read file content - with open(log_file_path, 'rb') as f: - file_content = f.read() - - return web.Response( - body=file_content, - content_type='text/plain', - headers={ - '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) - - async def truncate_logs(self, request): - """ - Truncate all log files. - POST /prompt_manager/logs/truncate - """ - try: - try: - from ..utils.logging_config import get_logger_manager - 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_manager - - logger_manager = get_logger_manager() - results = logger_manager.truncate_logs() - - 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) - - async def get_log_config(self, request): - """ - Get current logging configuration. - GET /prompt_manager/logs/config - """ - try: - try: - from ..utils.logging_config import get_logger_manager - 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_manager - - logger_manager = get_logger_manager() - config = logger_manager.get_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) - - async def update_log_config(self, request): - """ - Update logging configuration. - POST /prompt_manager/logs/config - Body: {"level": "DEBUG", "console_logging": true, ...} - """ - try: - data = await request.json() - - try: - from ..utils.logging_config import get_logger_manager - 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_manager - - logger_manager = get_logger_manager() - - # Validate level if provided - 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() - }) - - except Exception as e: - self.logger.error(f"Update log config error: {e}") - return web.json_response({ - 'success': False, - 'error': str(e) - }, status=500) - - async def get_log_stats(self, request): - """ - Get logging statistics. - GET /prompt_manager/logs/stats - """ - try: - try: - from ..utils.logging_config import get_logger_manager - 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_manager - - logger_manager = get_logger_manager() - stats = logger_manager.get_log_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) - - # ======================================================================== - # AutoTag Endpoints - # ======================================================================== - - async def get_autotag_models(self, request): - """ - Get status of available AutoTag models. - GET /prompt_manager/autotag/models - - Returns model availability, download status, loaded status, and configuration. - """ - try: - from .autotag import get_autotag_service - - 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() - }) - - except Exception as e: - self.logger.error(f"Get autotag models error: {e}") - 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. - POST /prompt_manager/autotag/download/{model_type} - - Streams SSE progress updates during download. - """ - import json - import asyncio - - model_type = request.match_info.get('model_type') - - async def stream_response(): - try: - from .autotag import get_autotag_service - - service = get_autotag_service() - - if model_type not in service.models_config: - yield f"data: {json.dumps({'type': 'error', 'message': f'Invalid model type: {model_type}'})}\n\n" - return - - yield f"data: {json.dumps({'type': 'progress', 'progress': 0, 'status': 'Starting download...'})}\n\n" - - # Define progress callback - progress_data = {'last_progress': 0} - - def progress_callback(status: str, progress: float): - progress_data['last_progress'] = progress - - # Perform download (blocking in thread) - loop = asyncio.get_event_loop() - success = await loop.run_in_executor( - None, - lambda: service.download_model(model_type, progress_callback) - ) - - if success: - yield f"data: {json.dumps({'type': 'complete', 'progress': 100, 'status': 'Download complete'})}\n\n" - else: - yield f"data: {json.dumps({'type': 'error', 'message': 'Download failed'})}\n\n" - - except Exception as e: - self.logger.error(f"Download model error: {e}") - yield f"data: {json.dumps({'type': 'error', 'message': str(e)})}\n\n" - - response = web.StreamResponse( - status=200, - reason='OK', - headers={ - '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_eof() - return response - - async def start_autotag(self, request): - """ - Start batch auto-tagging with streaming progress. - GET /prompt_manager/autotag/start - - Query params: - model_type: "gguf" or "hf" - prompt: custom prompt text - keep_in_memory: "true" or "false" (default true) - keep model loaded after processing - - Streams SSE progress updates during processing. - """ - import json - import asyncio - from pathlib import Path - - # Read from query params (EventSource only supports GET) - 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(): - try: - from .autotag import get_autotag_service - - service = get_autotag_service() - - # Check model is downloaded - status = service.get_models_status() - 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 - - yield f"data: {json.dumps({'type': 'progress', 'progress': 0, 'status': 'Loading model...'})}\n\n" - - # Load model in thread pool - loop = asyncio.get_event_loop() - try: - await loop.run_in_executor( - 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" - return - - # Set custom prompt if provided - if custom_prompt: - service.custom_prompt = custom_prompt - - yield f"data: {json.dumps({'type': 'progress', 'progress': 5, 'status': 'Model loaded. Fetching all images from database...'})}\n\n" - - # Get ALL images from database (only images with linked prompts) - images = self.db.get_all_images() - - total_files = len(images) - if total_files == 0: - yield f"data: {json.dumps({'type': 'complete', 'processed': 0, 'tagged': 0, 'skipped': 0, 'status': 'No images with linked prompts found'})}\n\n" - service.unload_model() - return - - yield f"data: {json.dumps({'type': 'progress', 'progress': 10, 'status': f'Found {total_files} images. Processing...'})}\n\n" - - processed = 0 - tagged = 0 - skipped = 0 - errors = 0 - tagged_prompt_ids = set() # Track prompts already tagged this run - - import time as _time - 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') - - if not image_path or not prompt_id: - skipped += 1 - continue - - # Skip if we already tagged this prompt during this run - if prompt_id in tagged_prompt_ids: - skipped += 1 - now = _time.monotonic() - if (now - last_update_time) >= 0.5 or i == total_files - 1: - progress = 10 + int((i + 1) / total_files * 85) - yield f"data: {json.dumps({'type': 'progress', 'progress': progress, 'status': f'Skipping {i+1}/{total_files} (prompt already processed)...', 'processed': processed, 'tagged': tagged, 'skipped': skipped})}\n\n" - await asyncio.sleep(0.01) - last_update_time = now - continue - - # Check if file exists - if not Path(image_path).exists(): - skipped += 1 - now = _time.monotonic() - if (now - last_update_time) >= 0.5 or i == total_files - 1: - progress = 10 + int((i + 1) / total_files * 85) - yield f"data: {json.dumps({'type': 'progress', 'progress': progress, 'status': f'Skipping {i+1}/{total_files} (file missing)...', 'processed': processed, 'tagged': tagged, 'skipped': skipped})}\n\n" - await asyncio.sleep(0.01) - last_update_time = now - continue - - # Check if image already has real tags (skip_tagged option) - if skip_tagged: - 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()] - # Filter out "auto-scanned" - it's not a real tag - real_tags = [t for t in prompt_tags if t != 'auto-scanned'] - if real_tags: - tagged_prompt_ids.add(prompt_id) # Don't re-check other images for this prompt - skipped += 1 - now = _time.monotonic() - if (now - last_update_time) >= 0.5 or i == total_files - 1: - progress = 10 + int((i + 1) / total_files * 85) - yield f"data: {json.dumps({'type': 'progress', 'progress': progress, 'status': f'Skipping {i+1}/{total_files} (already tagged)...', 'processed': processed, 'tagged': tagged, 'skipped': skipped})}\n\n" - await asyncio.sleep(0.01) - last_update_time = now - continue - - try: - # Generate tags - tags = await loop.run_in_executor( - None, - lambda p=str(image_path): service.generate_tags(p) - ) - - processed += 1 - - if tags: - # Get existing prompt (live read, not snapshot) - existing_prompt = self.db.get_prompt_by_id(prompt_id) - if existing_prompt: - existing_tags = existing_prompt.get('tags', []) - if isinstance(existing_tags, str): - 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: - all_tags = existing_tags + new_tags - self.db.update_prompt_metadata( - prompt_id, - tags=all_tags - ) - tagged += 1 - else: - skipped += 1 - else: - skipped += 1 - else: - skipped += 1 - - # Mark this prompt as done so other images for it are skipped - tagged_prompt_ids.add(prompt_id) - - # Always send progress after LLM inference (each call is slow) - progress = 10 + int((i + 1) / total_files * 85) - yield f"data: {json.dumps({'type': 'progress', 'progress': progress, 'status': f'Processing {i+1}/{total_files}...', 'processed': processed, 'tagged': tagged, 'skipped': skipped})}\n\n" - await asyncio.sleep(0.01) - last_update_time = _time.monotonic() - - except Exception as img_err: - self.logger.error(f"Error processing {image_path}: {img_err}") - errors += 1 - processed += 1 - - # Only unload model if keep_in_memory is False - if not keep_in_memory: - service.unload_model() - model_status = 'Model unloaded' - else: - 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', - headers={ - '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_eof() - return response - - async def autotag_single(self, request): - """ - Generate tags for a single image (for Review mode). - POST /prompt_manager/autotag/single - - Request body: - { - "image_path": "/path/to/image.png", - "model_type": "gguf", - "prompt": "optional custom prompt", - "use_gpu": true - } - """ - 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) - - if not image_path: - return web.json_response({ - 'success': False, - 'error': 'image_path is required' - }, status=400) - - from .autotag import get_autotag_service - import asyncio - - service = get_autotag_service() - - 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) - ) - - if custom_prompt: - service.custom_prompt = custom_prompt - - loop = asyncio.get_event_loop() - tags = await loop.run_in_executor( - None, - lambda: service.generate_tags(image_path) - ) - - prompt_id = None - try: - with self.db.model.get_connection() as conn: - cursor = conn.execute( - "SELECT prompt_id FROM generated_images WHERE image_path = ?", - (image_path,) - ) - row = cursor.fetchone() - if row: - prompt_id = row[0] - 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 - }) - - except Exception as e: - self.logger.error(f"AutoTag single error: {e}") - return web.json_response({ - 'success': False, - 'error': str(e) - }, status=500) - - async def apply_autotag(self, request): - """ - Apply selected tags to a prompt. - POST /prompt_manager/autotag/apply - - Request body: - { - "prompt_id": 123, - "tags": ["tag1", "tag2", ...] - } - """ - try: - data = await request.json() - 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) - - if not tags: - return web.json_response({ - 'success': True, - 'message': 'No tags to apply' - }) - - prompt = 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) - - existing_tags = prompt.get('tags', []) - if isinstance(existing_tags, str): - 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 - - 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) - }) - - except Exception as e: - self.logger.error(f"Apply autotag error: {e}") - 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. - POST /prompt_manager/autotag/unload - - Used to free up VRAM when the model is no longer needed. - """ - try: - from .autotag import get_autotag_service - - service = get_autotag_service() - - if not service.is_model_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 - }) - - except Exception as e: - self.logger.error(f"Unload autotag model error: {e}") - return web.json_response({ - 'success': False, - 'error': str(e) - }, status=500) - - async def scan_output_dir(self, request): - """ - Scan ComfyUI output directory for images. - GET /prompt_manager/scan_output_dir - - Returns a list of images in the output directory for autotag review mode. - """ - from pathlib import Path - - 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) - - output_path = Path(output_dir) - image_extensions = ['.png', '.jpg', '.jpeg', '.gif', '.webp', '.bmp', '.tiff'] - - images = [] - seen_paths = set() - - 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: - normalized_path = str(image_path).lower() - if normalized_path not in seen_paths: - seen_paths.add(normalized_path) - - rel_path = image_path.relative_to(output_path) - - # Check for thumbnail - thumbnail_url = None - thumbnails_dir = output_path / "thumbnails" - if thumbnails_dir.exists(): - # Preserve subdirectory structure in thumbnail path - 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}" - 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 - }) - - # Sort by filename - 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) - }) - - except Exception as e: - self.logger.error(f"Scan output dir error: {e}") - return web.json_response({ - 'success': False, - 'error': str(e) - }, status=500) diff --git a/py/api/__init__.py b/py/api/__init__.py new file mode 100644 index 0000000..c3c30af --- /dev/null +++ b/py/api/__init__.py @@ -0,0 +1,735 @@ +"""REST API module for ComfyUI PromptManager. + +Split into domain-specific mixins for maintainability: +- PromptRoutesMixin: prompt CRUD, tags, categories, bulk ops, export +- ImageRoutesMixin: gallery, thumbnails, image serving and linking +- AdminRoutesMixin: duplicates, stats, settings, diagnostics, maintenance, backup, scan +- LoggingRoutesMixin: log management endpoints +- AutotagRoutesMixin: auto-tagging model management and tagging +""" + +import asyncio +import datetime +import functools +import gzip as gzip_module +import json +import os +from pathlib import Path + +from aiohttp import web +from PIL import Image + +from .prompts import PromptRoutesMixin +from .images import ImageRoutesMixin +from .admin import AdminRoutesMixin +from .logging_routes import LoggingRoutesMixin +from .autotag_routes import AutotagRoutesMixin + +try: + from ...database.operations import PromptDatabase + from ...utils.logging_config import get_logger +except ImportError: + import sys + + sys.path.insert( + 0, os.path.dirname(os.path.dirname(os.path.dirname(os.path.abspath(__file__)))) + ) + from database.operations import PromptDatabase + from utils.logging_config import get_logger + + +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__))) + ) + + +# ── 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_registered = False + + +@web.middleware +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/'): + 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', ''): + return response + + 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 '' + if not any(ct in content_type for ct in _GZIP_TYPES): + return response + + compressed = gzip_module.compress(body, compresslevel=6) + if len(compressed) >= len(body): + return response + + response.body = compressed + response.headers['Content-Encoding'] = 'gzip' + response.headers['Vary'] = 'Accept-Encoding' + return response + + +class PromptManagerAPI( + PromptRoutesMixin, + ImageRoutesMixin, + AdminRoutesMixin, + LoggingRoutesMixin, + AutotagRoutesMixin, +): + """REST API handler for PromptManager operations and web interface. + + This class provides comprehensive REST API endpoints for managing prompts, + images, and system operations. It handles database interactions, file + operations, image processing, and web UI serving. + + The API is designed to integrate seamlessly with ComfyUI's aiohttp server + and provides both JSON API endpoints and static file serving for the + web interface. + + Attributes: + logger: Configured logger instance for API operations + db (PromptDatabase): Database connection and operations handler + """ + + def __init__(self): + """Initialize the PromptManager API with database connection and cleanup.""" + self.logger = get_logger('prompt_manager.api') + self.logger.info("Initializing PromptManager API") + + self.db = PromptDatabase() + self._cached_output_dir = None # Lazy-cached by _find_comfyui_output_dir() + self._html_cache = {} # Cached HTML file contents keyed by path + self._gallery_cache = None # Cached gallery file listing (Fix 2.4) + self._gallery_cache_time = 0 # Timestamp of last cache fill + self._gallery_cache_ttl = 30 # Cache TTL in seconds + + # Run cleanup on initialization to remove any existing duplicates + try: + removed = self.db.cleanup_duplicates() + if removed > 0: + self.logger.info(f"Startup cleanup: removed {removed} duplicate prompts") + except Exception as e: + self.logger.error(f"Startup cleanup failed: {e}") + + self.logger.info("PromptManager API initialization completed") + + async def _run_in_executor(self, func, *args, **kwargs): + """Run a blocking function in the default thread pool executor. + + Prevents synchronous database calls, file I/O, and PIL operations + from blocking the aiohttp event loop. + """ + loop = asyncio.get_event_loop() + if kwargs: + call = functools.partial(func, *args, **kwargs) + return await loop.run_in_executor(None, call) + return await loop.run_in_executor(None, func, *args) + + def invalidate_gallery_cache(self): + """Invalidate the cached gallery file listing. + + Called by the image monitor when new files are detected. + """ + self._gallery_cache = None + self._gallery_cache_time = 0 + + def add_routes(self, routes): + """Register all API routes with the ComfyUI server. + + Args: + routes: aiohttp RouteTableDef object from ComfyUI server instance. + """ + + # Test route to verify registration works + @routes.get("/prompt_manager/test") + async def test_route(request): + return web.json_response( + { + "success": True, + "message": "PromptManager API is working!", + "timestamp": str(datetime.datetime.now()), + } + ) + + # ── Web UI serving routes ───────────────────────────────────── + + @routes.get("/prompt_manager/web") + async def serve_web_ui(request): + try: + html_path = os.path.join(_get_project_root(), "web", "index.html") + + if os.path.exists(html_path): + with open(html_path, "r", encoding="utf-8") as f: + html_content = f.read() + + return web.Response( + text=html_content, content_type="text/html", charset="utf-8" + ) + else: + return web.Response( + text="

Web UI not found

HTML file not located at expected path.

", + content_type="text/html", + status=404, + ) + + except Exception as e: + return web.Response( + text=f"

Error

Failed to load web UI: {str(e)}

", + content_type="text/html", + status=500, + ) + + @routes.get("/prompt_manager/gallery.html") + async def serve_gallery_ui(request): + try: + html_path = os.path.join( + _get_project_root(), "web", "metadata.html", + ) + + if html_path not in self._html_cache: + if os.path.exists(html_path): + with open(html_path, "r", encoding="utf-8") as f: + self._html_cache[html_path] = f.read() + else: + return web.Response( + text="

Gallery not found

gallery.html file not located at expected path.

", + content_type="text/html", + status=404, + ) + + return web.Response( + text=self._html_cache[html_path], content_type="text/html", charset="utf-8" + ) + + except Exception as e: + return web.Response( + text=f"

Error

Failed to load gallery: {str(e)}

", + content_type="text/html", + status=500, + ) + + @routes.get("/prompt_manager/admin") + async def serve_admin_ui(request): + try: + html_path = os.path.join( + _get_project_root(), "web", "admin.html", + ) + + if html_path not in self._html_cache: + if os.path.exists(html_path): + with open(html_path, "r", encoding="utf-8") as f: + self._html_cache[html_path] = f.read() + else: + return web.Response( + text="

Admin UI not found

", + content_type="text/html", + status=404, + ) + + return web.Response( + text=self._html_cache[html_path], content_type="text/html", charset="utf-8" + ) + + except Exception as e: + return web.Response( + text=f"

Error

Failed to load admin UI: {str(e)}

", + content_type="text/html", + status=500, + ) + + @routes.get("/prompt_manager/gallery") + async def serve_gallery_admin_ui(request): + try: + html_path = os.path.join( + _get_project_root(), "web", "gallery.html", + ) + + if html_path not in self._html_cache: + if os.path.exists(html_path): + with open(html_path, "r", encoding="utf-8") as f: + self._html_cache[html_path] = f.read() + else: + return web.Response( + text="

Gallery not found

gallery.html file not located at expected path.

", + content_type="text/html", + status=404, + ) + + return web.Response( + text=self._html_cache[html_path], content_type="text/html", charset="utf-8" + ) + + except Exception as e: + return web.Response( + text=f"

Error

Failed to load gallery: {str(e)}

", + content_type="text/html", + status=500, + ) + + # ── Static file serving ─────────────────────────────────────── + + @routes.get("/prompt_manager/lib/{filepath:.*}") + async def serve_lib_static(request): + """Serve static library files (JS, CSS) from web/lib directory.""" + MIME_TYPES = { + ".js": "application/javascript", + ".css": "text/css", + ".json": "application/json", + ".map": "application/json", + } + + filepath = request.match_info.get("filepath", "") + + # Security: prevent directory traversal + if ".." in filepath or filepath.startswith("/"): + return web.Response(text="Forbidden", status=403) + + file_path = os.path.join(_get_project_root(), "web", "lib", filepath) + + if not os.path.exists(file_path) or not os.path.isfile(file_path): + return web.Response(text=f"Not Found: {filepath}", status=404) + + ext = os.path.splitext(file_path)[1].lower() + content_type = MIME_TYPES.get(ext, "application/octet-stream") + + with open(file_path, "rb") as f: + content = f.read() + + return web.Response(body=content, content_type=content_type) + + @routes.get("/prompt_manager/js/{filepath:.*}") + async def serve_js_static(request): + """Serve static JavaScript files from web/js directory.""" + MIME_TYPES = { + ".js": "application/javascript", + ".css": "text/css", + ".json": "application/json", + ".map": "application/json", + } + + filepath = request.match_info.get("filepath", "") + + if ".." in filepath or filepath.startswith("/"): + return web.Response(text="Forbidden", status=403) + + file_path = os.path.join(_get_project_root(), "web", "js", filepath) + + if not os.path.exists(file_path) or not os.path.isfile(file_path): + return web.Response(text=f"Not Found: {filepath}", status=404) + + ext = os.path.splitext(file_path)[1].lower() + content_type = MIME_TYPES.get(ext, "application/octet-stream") + + with open(file_path, "rb") as f: + content = f.read() + + return web.Response(body=content, content_type=content_type) + + # ── Register domain-specific routes from mixins ─────────────── + + self._register_prompt_routes(routes) + self._register_image_routes(routes) + self._register_admin_routes(routes) + self._register_logging_routes(routes) + self._register_autotag_routes(routes) + + # Register gzip compression middleware (once) + global _gzip_registered + 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") + except Exception as e: + self.logger.warning(f"Could not register gzip middleware: {e}") + + self.logger.info("All routes registered with decorator pattern") + + # ── Shared utilities used by multiple mixins ────────────────────── + + def _enrich_prompt_images(self, prompts): + """Add url and thumbnail_url to each image in prompt results.""" + from urllib.parse import quote as url_quote + + output_dir = self._find_comfyui_output_dir() + 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', '') + 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" + + # 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='/')}" + + # Check for thumbnail + 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='/')}" + except (ValueError, RuntimeError): + pass + + return prompts + + def _clean_nan_recursive(self, obj): + """Recursively clean NaN values from nested data structures.""" + if isinstance(obj, dict): + 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': + return None + else: + return obj + + def _find_comfyui_output_dir(self): + """Locate the ComfyUI output directory using multiple detection strategies. + + Results are cached after first successful lookup since the output + directory does not change during runtime. + + Returns: + str or None: Absolute path to ComfyUI output directory, or None + if no valid directory is found + """ + if self._cached_output_dir is not None: + return self._cached_output_dir + + current_file = Path(__file__).resolve() + self.logger.debug(f"Starting ComfyUI output search from: {current_file}") + + # Method 1: Search upward from current file location + current_dir = current_file.parent + max_depth = 10 # Prevent infinite loops + + for i in range(max_depth): + # Check if current directory contains ComfyUI markers + 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._cached_output_dir = str(output_dir) + return self._cached_output_dir + + # Move up one directory + parent = current_dir.parent + if parent == current_dir: # Reached filesystem root + break + current_dir = parent + + # Method 2: Try common installation patterns relative to this file + # File is at: custom_nodes/ComfyUI_PromptManager/py/api/__init__.py + base_dir = current_file.parent # .../py/api/ + possible_paths = [ + base_dir.parent.parent.parent.parent / "output", # ../../../../output + base_dir.parent.parent.parent / "output", # ../../../output + base_dir.parent.parent / "output", # ../../output + base_dir.parent / "output", # ../output + ] + + # Method 3: Add common ComfyUI installation locations + common_locations = [ + Path.home() / "ComfyUI" / "output", + Path.cwd() / "output", + Path.cwd() / ".." / "output", + Path.cwd() / ".." / ".." / "output", + ] + + all_paths = possible_paths + common_locations + + for path in all_paths: + try: + abs_path = path.resolve() + if abs_path.exists() and abs_path.is_dir(): + self.logger.debug(f"Found ComfyUI output directory: {abs_path}") + self._cached_output_dir = str(abs_path) + return self._cached_output_dir + except (OSError, RuntimeError): + continue # Skip invalid paths + + self.logger.warning("ComfyUI output directory not found. Searched paths:") + for path in all_paths: + try: + self.logger.warning(f" - {path.resolve()} (exists: {path.exists()})") + except (OSError, RuntimeError): + self.logger.warning(f" - {path} (invalid path)") + + return None + + def _extract_comfyui_metadata(self, image_path): + """Extract ComfyUI workflow metadata from PNG image files.""" + try: + with Image.open(image_path) as img: + metadata = {} + if hasattr(img, 'text'): + for key, value in img.text.items(): + metadata[key] = value + return metadata + except Exception as e: + self.logger.error(f"Error reading {image_path}: {e}") + return {} + + 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 + } + + # Check for A1111 style parameters first (like parse-metadata.py) + if "parameters" in metadata: + params = metadata["parameters"] + lines = params.splitlines() + if lines: + result['positive_prompt'] = lines[0].strip() + for line in lines: + if line.lower().startswith("negative prompt:"): + result['negative_prompt'] = line.split(":", 1)[1].strip() + break + # Store raw parameters too + result['parameters']['parameters'] = params + + # If no A1111 format found, proceed with ComfyUI parsing + if result['positive_prompt'] is None: + # Check for direct prompt field + if 'prompt' in metadata: + try: + prompt_data = json.loads(metadata['prompt']) + result['prompt'] = prompt_data + except json.JSONDecodeError: + result['prompt'] = metadata['prompt'] + + # Check for workflow + if 'workflow' in metadata: + try: + workflow_data = json.loads(metadata['workflow']) + result['workflow'] = workflow_data + except json.JSONDecodeError: + result['workflow'] = metadata['workflow'] + + # Check for other common ComfyUI fields + 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]) + except json.JSONDecodeError: + result['parameters'][field] = metadata[field] + + return result + + def _extract_readable_prompt(self, parsed_data): + """Extract human-readable prompt text from ComfyUI/A1111 data.""" + + def safe_to_string(value): + if isinstance(value, str): + return value + elif isinstance(value, list): + 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'] + + # Check if prompt is already a string + 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']) + + 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) + if positive_prompt: + return positive_prompt + + # Check workflow data if available + workflow_data = parsed_data.get('workflow') + if isinstance(workflow_data, dict): + 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']) + + return None + + def _get_node_inputs(self, node): + """Safely get inputs from a node, handling both dict and list formats. + + Old format: inputs is a dict with direct key-value pairs + inputs = {"text": "my prompt", "seed": 123} + + New format: inputs is a list of connection objects + inputs = [ + {"name": "text", "type": "STRING", "link": null, "widget": {"name": "text"}}, + {"name": "clip", "type": "CLIP", "link": 11} + ] + """ + if not isinstance(node, dict): + return {} + + inputs = node.get("inputs", {}) + + # If inputs is already a dict, return it + if isinstance(inputs, dict): + return inputs + + # If inputs is a list, convert to dict format + if isinstance(inputs, list): + inputs_dict = {} + for input_item in inputs: + if isinstance(input_item, dict) and "name" in input_item: + name = input_item["name"] + inputs_dict[name] = input_item + return inputs_dict + + return {} + + def _find_text_in_node(self, node): + """Try to find text content in a node using various strategies. + + Handles both old and new workflow formats. + """ + if not isinstance(node, dict): + return None + + # Strategy 1: Check normalized inputs for 'text' field + inputs = self._get_node_inputs(node) + if "text" in inputs and isinstance(inputs["text"], str): + return inputs["text"] + + # 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' + ] + + 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(): + return widgets_values[0] + + return None + + def _extract_positive_prompt_from_comfyui_data(self, data): + """Extract positive prompt from ComfyUI data, handling both old and new formats.""" + if not isinstance(data, dict): + return None + + # Build nodes dictionary + nodes_by_id = {} + if "nodes" in data: + # Handle nodes array format + for node in data["nodes"]: + if isinstance(node, dict): + nid = node.get("id") + if nid is not None: + nodes_by_id[nid] = node + else: + # Handle flat dictionary format (node_id -> node_data) + for nid_str, node in data.items(): + try: + nid = int(nid_str) + except (ValueError, TypeError): + nid = nid_str + if isinstance(node, dict): + if "id" in node: + nid = node["id"] + nodes_by_id[nid] = node + + if not nodes_by_id: + return None + + # First, try to find positive/negative connection pattern + pos_id = None + for node in nodes_by_id.values(): + if isinstance(node, dict): + inputs = self._get_node_inputs(node) + if "positive" in inputs and "negative" in inputs: + try: + pos_input = inputs["positive"] + if isinstance(pos_input, list) and len(pos_input) > 0: + pos_id = int(pos_input[0]) + break + except (ValueError, TypeError, IndexError): + continue + + # Get text from the positive node + if pos_id is not None and pos_id in nodes_by_id: + text_val = self._find_text_in_node(nodes_by_id[pos_id]) + if text_val: + return text_val + + # Fallback: find any text encoder node with text content + text_nodes = [] + for node in nodes_by_id.values(): + if isinstance(node, dict): + 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: + # Prioritize non-negative prompts + text_nodes.insert(0, text_val) + else: + text_nodes.append(text_val) + + # Return the first positive-looking prompt + if text_nodes: + return text_nodes[0] + + return None diff --git a/py/api/admin.py b/py/api/admin.py new file mode 100644 index 0000000..c64c1d7 --- /dev/null +++ b/py/api/admin.py @@ -0,0 +1,963 @@ +"""Admin and maintenance API routes for PromptManager.""" + +import asyncio +import datetime +import hashlib +import json +import os +import shutil +import sqlite3 +import tempfile +import traceback +from pathlib import Path + +from aiohttp import web + + +class AdminRoutesMixin: + """Mixin providing admin, diagnostics, and maintenance API endpoints.""" + + def _register_admin_routes(self, routes): + @routes.get("/prompt_manager/scan_duplicates") + async def scan_duplicates_route(request): + return await self.scan_duplicates_endpoint(request) + + @routes.post("/prompt_manager/delete_duplicate_images") + async def delete_duplicate_images_route(request): + return await self.delete_duplicate_images_endpoint(request) + + @routes.post("/prompt_manager/cleanup") + async def cleanup_duplicates_route(request): + return await self.cleanup_duplicates_endpoint(request) + + @routes.post("/prompt_manager/maintenance") + async def maintenance_route(request): + return await self.run_maintenance(request) + + @routes.get("/prompt_manager/stats") + async def get_stats_route(request): + return await self.get_statistics(request) + + @routes.get("/prompt_manager/settings") + async def get_settings_route(request): + return await self.get_settings(request) + + @routes.post("/prompt_manager/settings") + async def save_settings_route(request): + return await self.save_settings(request) + + @routes.get("/prompt_manager/backup") + async def backup_database_route(request): + return await self.backup_database(request) + + @routes.post("/prompt_manager/restore") + async def restore_database_route(request): + return await self.restore_database(request) + + @routes.get("/prompt_manager/diagnostics") + async def run_diagnostics_route(request): + return await self.run_diagnostics(request) + + @routes.post("/prompt_manager/diagnostics/test-link") + async def test_image_link_route(request): + return await self.test_image_link(request) + + @routes.post("/prompt_manager/scan") + async def scan_images_route(request): + return await self.scan_images(request) + + async def scan_duplicates_endpoint(self, request): + """Scan for duplicate images without removing them.""" + try: + duplicates = await self.find_duplicate_images() + + return web.json_response( + { + "success": True, + "duplicates": duplicates, + "message": f"Found {len(duplicates)} groups of duplicate images", + } + ) + + 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)}"}, + status=500, + ) + + async def cleanup_duplicates_endpoint(self, request): + """Cleanup duplicate prompts endpoint.""" + try: + removed_count = await self._run_in_executor(self.db.cleanup_duplicates) + + return web.json_response( + { + "success": True, + "message": "Cleanup completed", + "duplicates_removed": removed_count, + } + ) + + except Exception as e: + self.logger.error(f"Cleanup error: {e}") + return web.json_response( + {"success": False, "error": f"Failed to cleanup duplicates: {str(e)}"}, + status=500, + ) + + async def find_duplicate_images(self): + """Find duplicate images in ComfyUI output directory using content hashing.""" + self.logger.info("Scanning for duplicate images") + + try: + output_dir = self._find_comfyui_output_dir() + if not output_dir: + self.logger.warning("ComfyUI output directory not found") + return [] + + output_path = Path(output_dir) + + 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 = [] + seen_paths = set() + 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: + normalized_path = str(media_path).lower() + if normalized_path not in seen_paths: + seen_paths.add(normalized_path) + media_files.append(media_path) + + self.logger.info(f"Found {len(media_files)} media files to analyze") + + file_hashes = {} + processed = 0 + + for media_path in media_files: + try: + file_hash = self._calculate_file_hash(media_path) + + if file_hash not in file_hashes: + file_hashes[file_hash] = [] + + stat = media_path.stat() + 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' + + # 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_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 + } + + 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") + + except Exception as e: + self.logger.error(f"Error processing file {media_path}: {e}") + continue + + # Find duplicates (groups with more than one file) + 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) + }) + + self.logger.info(f"Found {len(duplicates)} groups of duplicate images") + return duplicates + + except Exception as e: + self.logger.error(f"Error finding duplicate images: {e}") + return [] + + def _calculate_file_hash(self, file_path): + """Calculate SHA-256 hash of a file's content.""" + hash_sha256 = hashlib.sha256() + with open(file_path, "rb") as f: + for chunk in iter(lambda: f.read(4096), b""): + hash_sha256.update(chunk) + return hash_sha256.hexdigest() + + async def delete_duplicate_images_endpoint(self, request): + """Delete duplicate image files from disk.""" + try: + data = await request.json() + image_paths = data.get('image_paths', []) + + if not image_paths: + return web.json_response( + {"success": False, "error": "No image paths provided"}, + status=400, + ) + + deleted_count = 0 + failed_count = 0 + failed_files = [] + + for image_path in image_paths: + try: + # 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_count += 1 + continue + + output_path = Path(output_dir) + file_path = Path(image_path) + + # Security check - ensure file is within output directory + try: + file_path.resolve().relative_to(output_path.resolve()) + except ValueError: + 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 + + if file_path.exists() and file_path.is_file(): + os.remove(file_path) + deleted_count += 1 + self.logger.info(f"Deleted duplicate image: {image_path}") + + # 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}" + if thumbnail_path.exists(): + os.remove(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}") + else: + failed_files.append(f"{image_path} (file not found)") + failed_count += 1 + + except Exception as e: + self.logger.error(f"Error deleting file {image_path}: {e}") + failed_files.append(f"{image_path} ({str(e)})") + failed_count += 1 + + response_data = { + "success": True, + "deleted_count": deleted_count, + "failed_count": failed_count, + "message": f"Deleted {deleted_count} files successfully" + } + + if failed_count > 0: + response_data["failed_files"] = failed_files + response_data["message"] += f", {failed_count} failed" + + return web.json_response(response_data) + + 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)}"}, + status=500, + ) + + async def get_statistics(self, request): + """Get database statistics.""" + try: + stats = await self._run_in_executor(self.db.get_statistics) + + return web.json_response({"success": True, "stats": stats}) + + except Exception as e: + self.logger.error(f"Stats error: {e}") + return web.json_response( + {"success": False, "error": f"Failed to get statistics: {str(e)}"}, + status=500, + ) + + async def get_settings(self, request): + """Get current settings.""" + try: + from ..config import PromptManagerConfig, GalleryConfig + + # Get monitored directories from image monitor singleton if available + 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'): + 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', []) + 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 + } + }) + except Exception as e: + return web.json_response( + {"success": False, "error": f"Failed to get settings: {str(e)}"}, + status=500, + ) + + async def save_settings(self, request): + """Save settings.""" + try: + from ..config import PromptManagerConfig, GalleryConfig + + data = await request.json() + 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'] + + # 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 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'] + 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) + 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_data = { + 'web_ui': { + 'result_timeout': PromptManagerConfig.RESULT_TIMEOUT, + 'webui_display_mode': PromptManagerConfig.WEBUI_DISPLAY_MODE + }, + 'gallery': { + 'monitoring': { + 'directories': GalleryConfig.MONITORING_DIRECTORIES + } + } + } + + try: + 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 + }) + except Exception as e: + return web.json_response( + {"success": False, "error": f"Failed to save settings: {str(e)}"}, + status=500, + ) + + async def run_diagnostics(self, request): + """Run comprehensive system diagnostics and health checks.""" + try: + results = {} + + # Check database + try: + db_path = "prompts.db" + if os.path.exists(db_path): + 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'] + + 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'] + else: + image_count = 0 + + 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}' + } + except Exception as e: + results['database'] = { + 'status': 'error', + 'message': f'Database error: {str(e)}' + } + + # Check dependencies + dependencies = {} + try: + import watchdog + dependencies['watchdog'] = True + except ImportError: + dependencies['watchdog'] = False + + try: + from PIL import Image + dependencies['PIL'] = True + except ImportError: + dependencies['PIL'] = False + + dependencies['sqlite3'] = True # Always available in Python + + results['dependencies'] = { + 'status': 'ok' if all(dependencies.values()) else 'error', + 'dependencies': dependencies + } + + # Check output directories + output_dirs = [] + potential_dirs = ["output", "../output", "../../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) + + 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 + } + else: + 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)}' + } + + 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) + + 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') + + if not prompt_id: + return web.json_response({ + 'success': False, + 'error': 'prompt_id is required' + }, status=400) + + test_metadata = { + 'file_info': { + 'size': 1024000, + 'dimensions': [512, 512], + 'format': 'PNG' + }, + 'workflow': {'test': True}, + 'prompt': {'test_prompt': 'This is a test image'} + } + + try: + image_id = await self._run_in_executor( + self.db.link_image_to_prompt, + prompt_id=str(prompt_id), + image_path=test_image_path, + 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}' + } + }) + except Exception as 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) + + 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']) + + results = {} + + def _run_maintenance(): + 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' + } + except Exception as e: + results['cleanup_duplicates'] = {'success': False, 'error': str(e), 'message': 'Failed to cleanup duplicates'} + + if 'vacuum' in operations: + try: + self.db.model.vacuum_database() + 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'} + + 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' + } + except Exception as e: + results['cleanup_orphaned_images'] = {'success': False, 'error': str(e), 'message': 'Failed to cleanup orphaned images'} + + 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' + } + except Exception as e: + results['check_hash_duplicates'] = {'success': False, 'error': str(e), 'message': 'Failed to check hash duplicates'} + + if 'statistics' in operations: + try: + db_info = self.db.model.get_database_info() + 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'} + + 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)' + } + except Exception as e: + results['prune_orphaned_prompts'] = {'success': False, 'error': str(e), 'message': 'Failed to prune orphaned prompts'} + + 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' + } + except Exception as e: + 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()) + + 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) + + async def backup_database(self, request): + """Backup the entire prompts.db database file.""" + try: + db_path = "prompts.db" + + if not os.path.exists(db_path): + return web.json_response({ + 'success': False, + 'error': 'Database file not found' + }, status=404) + + 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: + file_data = f.read() + + os.unlink(temp_path) + + timestamp = datetime.datetime.now().strftime("%Y%m%d_%H%M%S") + filename = f"prompts_backup_{timestamp}.db" + + return web.Response( + body=file_data, + content_type='application/octet-stream', + headers={ + '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) + + async def restore_database(self, request): + """Restore the prompts.db database from uploaded file.""" + try: + 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) + + 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) + + 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) + + with tempfile.NamedTemporaryFile(delete=False, suffix='.db') as temp_file: + temp_path = temp_file.name + temp_file.write(file_data) + + try: + 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'") + 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'] + + 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'] + + db_path = "prompts.db" + backup_path = f"{db_path}.backup_{datetime.datetime.now().strftime('%Y%m%d_%H%M%S')}" + + if os.path.exists(db_path): + shutil.copy2(db_path, backup_path) + self.logger.info(f"Current database backed up to: {backup_path}") + + shutil.copy2(temp_path, db_path) + + # Reinitialize the database connection + try: + 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__))))) + 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 + }) + + except sqlite3.Error as e: + 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) + 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) + + async def scan_images(self, request): + """Scan ComfyUI output images for prompt metadata and add them to the database.""" + + async def stream_response(): + try: + self.logger.info("Starting image scan operation") + self.logger.info("Starting scan (timer clearing not implemented yet)") + + output_dir = self._find_comfyui_output_dir() + if not output_dir: + self.logger.error("ComfyUI output directory not found") + yield f"data: {json.dumps({'type': 'error', 'message': 'ComfyUI output directory not found'})}\n\n" + return + + yield f"data: {json.dumps({'type': 'progress', 'progress': 0, 'status': 'Scanning for PNG files...', 'processed': 0, 'found': 0})}\n\n" + + png_files = await self._run_in_executor( + lambda: list(Path(output_dir).rglob("*.png")) + ) + total_files = len(png_files) + + if total_files == 0: + yield f"data: {json.dumps({'type': 'complete', 'processed': 0, 'found': 0, 'added': 0})}\n\n" + return + + yield f"data: {json.dumps({'type': 'progress', 'progress': 5, 'status': f'Found {total_files} PNG files to process...', 'processed': 0, 'found': 0})}\n\n" + + processed_count = 0 + found_count = 0 + added_count = 0 + linked_count = 0 + + try: + 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__)))) + sys.path.insert(0, current_dir) + from utils.hashing import generate_prompt_hash + + for i, png_file in enumerate(png_files): + try: + metadata = await self._run_in_executor( + self._extract_comfyui_metadata, str(png_file) + ) + processed_count += 1 + + if metadata: + 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'))}") + + 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]}...") + else: + 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") + 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]}...") + + 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)}") + try: + await self._run_in_executor( + 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']}") + except Exception as 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)}") + 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 + ) + + if prompt_id: + added_count += 1 + 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.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}") + else: + 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}") + + # Update progress every 10 files + if i % 10 == 0 or i == total_files - 1: + progress = int((i + 1) / total_files * 100) + status = f"Processing file {i + 1}/{total_files}..." + + yield f"data: {json.dumps({'type': 'progress', 'progress': progress, 'status': status, 'processed': processed_count, 'found': found_count})}\n\n" + + await asyncio.sleep(0.01) + + except Exception as e: + 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}") + 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: + self.logger.error(f"Scan error: {e}") + self.logger.error(f"Scan error traceback: {traceback.format_exc()}") + yield f"data: {json.dumps({'type': 'error', 'message': str(e)})}\n\n" + + response = web.StreamResponse( + status=200, + reason='OK', + headers={ + '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_eof() + return response diff --git a/py/api/autotag_routes.py b/py/api/autotag_routes.py new file mode 100644 index 0000000..bfac0c9 --- /dev/null +++ b/py/api/autotag_routes.py @@ -0,0 +1,489 @@ +"""AutoTag API routes for PromptManager.""" + +import asyncio +import json +from pathlib import Path + +from aiohttp import web + + +class AutotagRoutesMixin: + """Mixin providing auto-tagging API endpoints.""" + + def _register_autotag_routes(self, routes): + @routes.get("/prompt_manager/autotag/models") + async def get_autotag_models_route(request): + return await self.get_autotag_models(request) + + @routes.get("/prompt_manager/autotag/download/{model_type}") + async def download_autotag_model_route(request): + return await self.download_autotag_model(request) + + @routes.get("/prompt_manager/autotag/start") + async def start_autotag_route(request): + return await self.start_autotag(request) + + @routes.post("/prompt_manager/autotag/single") + async def autotag_single_route(request): + return await self.autotag_single(request) + + @routes.post("/prompt_manager/autotag/apply") + async def apply_autotag_route(request): + return await self.apply_autotag(request) + + @routes.post("/prompt_manager/autotag/unload") + async def unload_autotag_model_route(request): + return await self.unload_autotag_model(request) + + @routes.get("/prompt_manager/scan_output_dir") + async def scan_output_dir_route(request): + return await self.scan_output_dir(request) + + async def get_autotag_models(self, request): + """Get status of available AutoTag models.""" + try: + from ..autotag import get_autotag_service + + 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() + }) + + except Exception as e: + self.logger.error(f"Get autotag models error: {e}") + 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') + + async def stream_response(): + try: + from ..autotag import get_autotag_service + + service = get_autotag_service() + + if model_type not in service.models_config: + yield f"data: {json.dumps({'type': 'error', 'message': f'Invalid model type: {model_type}'})}\n\n" + return + + yield f"data: {json.dumps({'type': 'progress', 'progress': 0, 'status': 'Starting download...'})}\n\n" + + progress_data = {'last_progress': 0} + + def progress_callback(status: str, progress: float): + 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) + ) + + if success: + yield f"data: {json.dumps({'type': 'complete', 'progress': 100, 'status': 'Download complete'})}\n\n" + else: + yield f"data: {json.dumps({'type': 'error', 'message': 'Download failed'})}\n\n" + + except Exception as e: + self.logger.error(f"Download model error: {e}") + yield f"data: {json.dumps({'type': 'error', 'message': str(e)})}\n\n" + + response = web.StreamResponse( + status=200, + reason='OK', + headers={ + '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_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' + use_gpu = True + + async def stream_response(): + try: + from ..autotag import get_autotag_service + import time as _time + + service = get_autotag_service() + + status = service.get_models_status() + 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 + + yield f"data: {json.dumps({'type': 'progress', 'progress': 0, 'status': 'Loading model...'})}\n\n" + + loop = asyncio.get_event_loop() + try: + await loop.run_in_executor( + 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" + return + + if custom_prompt: + service.custom_prompt = custom_prompt + + yield f"data: {json.dumps({'type': 'progress', 'progress': 5, 'status': 'Model loaded. Fetching all images from database...'})}\n\n" + + images = await self._run_in_executor(self.db.get_all_images) + + total_files = len(images) + if total_files == 0: + yield f"data: {json.dumps({'type': 'complete', 'processed': 0, 'tagged': 0, 'skipped': 0, 'status': 'No images with linked prompts found'})}\n\n" + service.unload_model() + return + + yield f"data: {json.dumps({'type': 'progress', 'progress': 10, 'status': f'Found {total_files} images. Processing...'})}\n\n" + + processed = 0 + tagged = 0 + skipped = 0 + errors = 0 + tagged_prompt_ids = set() + + 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') + + if not image_path or not prompt_id: + skipped += 1 + continue + + if prompt_id in tagged_prompt_ids: + skipped += 1 + now = _time.monotonic() + if (now - last_update_time) >= 0.5 or i == total_files - 1: + progress = 10 + int((i + 1) / total_files * 85) + yield f"data: {json.dumps({'type': 'progress', 'progress': progress, 'status': f'Skipping {i+1}/{total_files} (prompt already processed)...', 'processed': processed, 'tagged': tagged, 'skipped': skipped})}\n\n" + await asyncio.sleep(0.01) + last_update_time = now + continue + + if not Path(image_path).exists(): + skipped += 1 + now = _time.monotonic() + if (now - last_update_time) >= 0.5 or i == total_files - 1: + progress = 10 + int((i + 1) / total_files * 85) + yield f"data: {json.dumps({'type': 'progress', 'progress': progress, 'status': f'Skipping {i+1}/{total_files} (file missing)...', 'processed': processed, 'tagged': tagged, 'skipped': skipped})}\n\n" + await asyncio.sleep(0.01) + last_update_time = now + continue + + if skip_tagged: + 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'] + if real_tags: + tagged_prompt_ids.add(prompt_id) + skipped += 1 + now = _time.monotonic() + if (now - last_update_time) >= 0.5 or i == total_files - 1: + progress = 10 + int((i + 1) / total_files * 85) + yield f"data: {json.dumps({'type': 'progress', 'progress': progress, 'status': f'Skipping {i+1}/{total_files} (already tagged)...', 'processed': processed, 'tagged': tagged, 'skipped': skipped})}\n\n" + await asyncio.sleep(0.01) + last_update_time = now + continue + + try: + tags = await loop.run_in_executor( + 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) + if existing_prompt: + existing_tags = existing_prompt.get('tags', []) + if isinstance(existing_tags, str): + 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: + all_tags = existing_tags + new_tags + await self._run_in_executor( + self.db.update_prompt_metadata, + prompt_id, + tags=all_tags + ) + tagged += 1 + else: + skipped += 1 + else: + skipped += 1 + else: + skipped += 1 + + tagged_prompt_ids.add(prompt_id) + + progress = 10 + int((i + 1) / total_files * 85) + yield f"data: {json.dumps({'type': 'progress', 'progress': progress, 'status': f'Processing {i+1}/{total_files}...', 'processed': processed, 'tagged': tagged, 'skipped': skipped})}\n\n" + await asyncio.sleep(0.01) + last_update_time = _time.monotonic() + + except Exception as img_err: + self.logger.error(f"Error processing {image_path}: {img_err}") + errors += 1 + processed += 1 + + if not keep_in_memory: + service.unload_model() + model_status = 'Model unloaded' + else: + 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', + headers={ + '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_eof() + return response + + async def autotag_single(self, request): + """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) + + if not image_path: + 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: + loop = asyncio.get_event_loop() + await loop.run_in_executor( + None, + lambda: service.load_model(model_type, use_gpu) + ) + + if custom_prompt: + service.custom_prompt = custom_prompt + + loop = asyncio.get_event_loop() + tags = await loop.run_in_executor( + 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) + 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 + }) + + except Exception as e: + self.logger.error(f"AutoTag single error: {e}") + 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', []) + + if not prompt_id: + 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' + }) + + 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) + + existing_tags = prompt.get('tags', []) + if isinstance(existing_tags, str): + 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 + ) + + 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) + + async def unload_autotag_model(self, request): + """Manually unload the AutoTag model from memory.""" + try: + from ..autotag import get_autotag_service + + service = get_autotag_service() + + if not service.is_model_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 + }) + + except Exception as e: + self.logger.error(f"Unload autotag model error: {e}") + 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) + + output_path = Path(output_dir) + image_extensions = ['.png', '.jpg', '.jpeg', '.gif', '.webp', '.bmp', '.tiff'] + + images = [] + seen_paths = set() + + 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: + normalized_path = str(image_path).lower() + if normalized_path not in seen_paths: + seen_paths.add(normalized_path) + + rel_path = image_path.relative_to(output_path) + + thumbnail_url = None + thumbnails_dir = output_path / "thumbnails" + if thumbnails_dir.exists(): + 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}" + 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']) + + self.logger.info(f"Found {len(images)} images in output directory") + + 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) diff --git a/py/api/images.py b/py/api/images.py new file mode 100644 index 0000000..c34a223 --- /dev/null +++ b/py/api/images.py @@ -0,0 +1,1045 @@ +"""Image and gallery API routes for PromptManager.""" + +import asyncio +import json +import os +import re +import time as _time +import urllib.parse +from pathlib import Path + +from aiohttp import web +from PIL import Image + + +class ImageRoutesMixin: + """Mixin providing image and gallery-related API endpoints.""" + + def _register_image_routes(self, routes): + @routes.get("/prompt_manager/prompts/{prompt_id}/images") + async def get_prompt_images_route(request): + return await self.get_prompt_images(request) + + @routes.get("/prompt_manager/images/recent") + async def get_recent_images_route(request): + return await self.get_recent_images(request) + + @routes.get("/prompt_manager/images/all") + async def get_all_images_route(request): + return await self.get_all_images(request) + + @routes.get("/prompt_manager/images/search") + async def search_images_route(request): + return await self.search_images(request) + + @routes.get("/prompt_manager/images/output") + async def get_output_images_route(request): + return await self.get_output_images(request) + + @routes.get("/prompt_manager/images/{image_id}/file") + async def serve_image_route(request): + return await self.serve_image(request) + + @routes.get("/prompt_manager/images/serve/{filepath:.*}") + async def serve_output_image_route(request): + return await self.serve_output_image(request) + + @routes.post("/prompt_manager/images/link") + async def link_image_route(request): + return await self.link_image_to_prompt(request) + + @routes.get("/prompt_manager/images/prompt/{image_path:.*}") + async def get_image_prompt_route(request): + return await self.get_image_prompt(request) + + @routes.delete("/prompt_manager/images/{image_id}") + async def delete_image_route(request): + return await self.delete_image(request) + + @routes.post("/prompt_manager/images/generate-thumbnails") + async def generate_thumbnails_route(request): + return await self.generate_thumbnails(request) + + @routes.get("/prompt_manager/images/generate-thumbnails/progress") + async def generate_thumbnails_progress_route(request): + return await self.generate_thumbnails_with_progress(request) + + @routes.post("/prompt_manager/images/clear-thumbnails") + async def clear_thumbnails_route(request): + return await self.clear_thumbnails(request) + + async def get_prompt_images(self, request): + """Get all images for a specific prompt.""" + try: + prompt_id = request.match_info["prompt_id"] + images = await self._run_in_executor(self.db.get_prompt_images, prompt_id) + + # Clean up any NaN values that cause JSON parsing errors (recursive) + cleaned_images = [self._clean_nan_recursive(image) for image in images] + + # Additional fallback: convert to JSON string and clean NaN values manually + try: + 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) + + # Parse back to verify it's valid JSON + cleaned_data = json.loads(json_str) + + return web.json_response(cleaned_data) + 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 + }) + + except Exception as e: + self.logger.error(f"Get prompt images error: {e}") + 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)) + images = await self._run_in_executor(self.db.get_recent_images, limit) + + 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) + + 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) + }) + except Exception as e: + self.logger.error(f"Get all images error: {e}") + 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', '') + if not query: + 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 + }) + except Exception as e: + self.logger.error(f"Search images error: {e}") + 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'] + media_extensions = image_extensions + video_extensions + all_images = [] + + seen_paths = set() + 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: + normalized_path = str(media_path).lower() + if normalized_path not in seen_paths: + seen_paths.add(normalized_path) + all_images.append(media_path) + + # Sort by modification time (newest first) + all_images.sort(key=lambda x: x.stat().st_mtime, reverse=True) + return all_images + + 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): + return self._gallery_cache + + all_images = await self._run_in_executor( + self._scan_gallery_files_sync, output_path + ) + self._gallery_cache = all_images + self._gallery_cache_time = now + return all_images + + async def get_output_images(self, request): + """Get all images from ComfyUI output folder.""" + try: + from urllib.parse import quote + + # 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': [] + }) + + # Get pagination parameters + 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'] + + # 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] + + # Format media data in executor (stat calls are blocking) + def _format_page(): + images = [] + for media_path in paginated_images: + try: + stat = media_path.stat() + 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' + + thumbnail_url = None + if thumbnails_dir.exists(): + 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 + }) + except Exception as e: + self.logger.error(f"Error processing media {media_path}: {e}") + continue + return images + + 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) + }) + + 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) + + async def serve_image(self, request): + """Serve the actual image file using streamed FileResponse.""" + try: + image_id = int(request.match_info["image_id"]) + 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) + + 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) + + if not image_path.exists(): + 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' + return response + + except ValueError: + 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) + + async def serve_output_image(self, request): + """Serve image file directly from ComfyUI output folder using streamed FileResponse.""" + try: + filepath = request.match_info["filepath"] + + # 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) + + # Construct full image path + image_path = Path(output_dir) / filepath + + # Security check: make sure the path is within the output directory + try: + 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) + except Exception: + 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) + + response = web.FileResponse(image_path) + 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) + + 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') + + # Map quality to size + 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) + + output_path = Path(output_dir) + 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 + ) + + # Invalidate gallery cache since thumbnails changed + self.invalidate_gallery_cache() + + 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) + except Exception as e: + self.logger.error(f"Generate thumbnails error: {e}") + 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).""" + import time + + 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'] + media_extensions = image_extensions + video_extensions + media_files = [] + for root, dirs, files in os.walk(output_path): + if 'thumbnails' in Path(root).parts: + continue + for file in files: + if any(file.lower().endswith(ext) for ext in media_extensions): + media_files.append(Path(root) / file) + + total_images = len(media_files) + self.logger.info(f"Found {total_images} media files to process for thumbnails") + + if total_images == 0: + return { + 'success': True, + 'count': 0, + 'total_images': 0, + 'message': 'No media files found to process', + 'errors': [] + } + + generated_count = 0 + skipped_count = 0 + errors = [] + start_time = time.time() + + 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) + + rel_path_no_ext = rel_path.with_suffix('') + if is_video: + 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}" + + 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): + generated_count += 1 + else: + 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} + img.save(thumbnail_path, **save_kwargs) + generated_count += 1 + + 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") + + except Exception as 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") + + 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) + } + + async def generate_thumbnails_with_progress(self, request): + """Generate thumbnails with Server-Sent Events progress updates.""" + try: + import time + + # Parse query parameters + quality = request.query.get('quality', 'medium') + + # Map quality to size + 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': '*', + } + ) + await response.prepare(request) + + async def send_progress(event_type, data): + """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 asyncio.sleep(0.01) + except Exception as e: + self.logger.warning(f"Failed to send SSE message: {e}") + + try: + # 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' + }) + return response + + output_path = Path(output_dir) + 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...' + }) + + 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'] + 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: + 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)' + }) + + 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") + + 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)) + 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 + }) + + 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' + }) + return response + + generated_count = 0 + skipped_count = 0 + errors = [] + start_time = time.time() + + for i, media_file in enumerate(media_files): + try: + 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) + + rel_path = media_file.relative_to(output_path) + + rel_path_no_ext = rel_path.with_suffix('') + if is_video: + 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}" + + # 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}") + continue + except Exception as 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): + 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): + generated_count += 1 + else: + 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} + 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: + 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' + } + + 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}") + + except Exception as e: + error_msg = f"Failed to generate thumbnail for {media_file.name}: {str(e)}" + errors.append(error_msg) + self.logger.warning(error_msg) + + if len(errors) <= 5: + 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' + if errors: + 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 + }) + + 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)}' + }) + + return response + + except ImportError: + 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) + + def _generate_video_thumbnail(self, video_path, thumbnail_path, thumbnail_size): + """Generate thumbnail from video file. Returns True if successful.""" + try: + # Try using OpenCV first (most reliable) + try: + import cv2 + + cap = cv2.VideoCapture(str(video_path)) + if not cap.isOpened(): + return False + + frame_count = int(cap.get(cv2.CAP_PROP_FRAME_COUNT)) + target_frame = max(1, int(frame_count * 0.1)) + cap.set(cv2.CAP_PROP_POS_FRAMES, target_frame) + + ret, frame = cap.read() + cap.release() + + if not ret or frame is None: + return False + + frame_rgb = cv2.cvtColor(frame, cv2.COLOR_BGR2RGB) + + 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}") + return True + + except ImportError: + pass + + # Fallback to ffmpeg + try: + 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) + ] + + 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}") + return True + else: + self.logger.warning(f"ffmpeg failed for {video_path}: {result.stderr}") + + except (ImportError, subprocess.TimeoutExpired, FileNotFoundError): + pass + + # Last resort: create a placeholder thumbnail + try: + from PIL import ImageDraw, ImageFont + + 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) + ] + draw.polygon(points, fill=(255, 255, 255)) + + try: + font = ImageFont.load_default() + text = "VIDEO" + 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 + ) + except (OSError, AttributeError): + pass + + 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}") + return False + + except Exception as e: + self.logger.error(f"Video thumbnail generation failed for {video_path}: {e}") + return False + + async def clear_thumbnails(self, request): + """Safely clear only our generated thumbnails, never touch original images.""" + try: + import shutil + + # 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) + + output_path = Path(output_dir) + 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 + }) + + # 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) + + except Exception as e: + self.logger.error(f"Path validation failed: {e}") + return web.json_response({ + 'success': False, + 'error': 'Path validation failed' + }, status=500) + + # Count files before deletion + cleared_count = 0 + cleared_size = 0 + + for root, dirs, files in os.walk(thumbnails_dir): + 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']): + try: + file_size = file_path.stat().st_size + file_path.unlink() + cleared_count += 1 + 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}") + + # Remove empty directories within thumbnails folder + try: + for root, dirs, files in os.walk(thumbnails_dir, topdown=False): + if root != str(thumbnails_dir): + try: + Path(root).rmdir() + except OSError: + pass + except Exception as e: + self.logger.debug(f"Directory cleanup info: {e}") + + def format_size(bytes_size): + 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)})") + + 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) + + 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', {}) + + if not prompt_id or not image_path: + 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) + + 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' + }) + + except Exception as e: + self.logger.error(f"Link image error: {e}") + 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', '') + 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) + + # Convert relative path to absolute if needed + if not os.path.isabs(image_path): + output_dir = self._find_comfyui_output_dir() + if output_dir: + image_path = str(Path(output_dir) / image_path) + + # Look up the image in generated_images table + try: + prompt_data = await self._run_in_executor(self.db.get_image_prompt_info, image_path) + if prompt_data: + prompt_data['image_path'] = image_path + if 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 + }) + + 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) + + except Exception as e: + self.logger.error(f"Get image prompt error: {e}") + return web.json_response({ + 'success': False, + 'error': str(e) + }, status=500) + + async def delete_image(self, request): + """Delete an image record.""" + try: + image_id = int(request.match_info["image_id"]) + success = await self._run_in_executor(self.db.delete_image, image_id) + + if success: + return web.json_response({ + 'success': True, + 'message': 'Image deleted successfully' + }) + else: + 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) + except Exception as e: + self.logger.error(f"Delete image error: {e}") + 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 new file mode 100644 index 0000000..829dd80 --- /dev/null +++ b/py/api/logging_routes.py @@ -0,0 +1,224 @@ +"""Logging API routes for PromptManager.""" + +import os + +from aiohttp import web + + +class LoggingRoutesMixin: + """Mixin providing logging-related API endpoints.""" + + def _register_logging_routes(self, routes): + @routes.get("/prompt_manager/logs") + async def get_logs_route(request): + return await self.get_logs(request) + + @routes.get("/prompt_manager/logs/files") + async def get_log_files_route(request): + return await self.get_log_files(request) + + @routes.get("/prompt_manager/logs/download/{filename}") + async def download_log_route(request): + return await self.download_log_file(request) + + @routes.post("/prompt_manager/logs/truncate") + async def truncate_logs_route(request): + return await self.truncate_logs(request) + + @routes.get("/prompt_manager/logs/config") + async def get_log_config_route(request): + return await self.get_log_config(request) + + @routes.post("/prompt_manager/logs/config") + async def update_log_config_route(request): + return await self.update_log_config(request) + + @routes.get("/prompt_manager/logs/stats") + async def get_log_stats_route(request): + return await self.get_log_stats(request) + + def _get_logger_manager(self): + """Get the logger manager instance.""" + try: + 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__)))) + sys.path.insert(0, current_dir) + from utils.logging_config import get_logger_manager + return get_logger_manager() + + async def get_logs(self, request): + """Get recent log entries.""" + try: + logger_manager = self._get_logger_manager() + + limit = int(request.query.get('limit', 100)) + level = request.query.get('level', None) + + if limit > 1000: + limit = 1000 + elif limit < 1: + limit = 1 + + 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 + }) + + except Exception as e: + self.logger.error(f"Get logs error: {e}") + 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.""" + try: + 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) + }) + + 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) + + async def download_log_file(self, request): + """Download a specific log file.""" + try: + 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) + + 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) + + with open(log_file_path, 'rb') as f: + file_content = f.read() + + return web.Response( + body=file_content, + content_type='text/plain', + headers={ + '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) + + async def truncate_logs(self, request): + """Truncate all log files.""" + try: + 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 + }) + + except Exception as e: + self.logger.error(f"Truncate logs error: {e}") + return web.json_response({ + 'success': False, + 'error': str(e) + }, status=500) + + async def get_log_config(self, request): + """Get current logging configuration.""" + try: + logger_manager = self._get_logger_manager() + config = logger_manager.get_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) + + async def update_log_config(self, request): + """Update logging configuration.""" + try: + 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() + + logger_manager.update_config(data) + + 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) + + async def get_log_stats(self, request): + """Get logging statistics.""" + try: + logger_manager = self._get_logger_manager() + stats = logger_manager.get_log_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) diff --git a/py/api/prompts.py b/py/api/prompts.py new file mode 100644 index 0000000..34cd31d --- /dev/null +++ b/py/api/prompts.py @@ -0,0 +1,860 @@ +"""Prompt API routes for PromptManager.""" + +import datetime +import json + +from aiohttp import web + + +class PromptRoutesMixin: + """Mixin providing prompt-related API endpoints.""" + + def _register_prompt_routes(self, routes): + @routes.get("/prompt_manager/search") + async def search_prompts_route(request): + return await self.search_prompts(request) + + @routes.get("/prompt_manager/recent") + async def get_recent_prompts_route(request): + return await self.get_recent_prompts(request) + + @routes.get("/prompt_manager/categories") + async def get_categories_route(request): + return await self.get_categories(request) + + # Tag management endpoints (must be registered BEFORE /prompt_manager/tags) + @routes.get("/prompt_manager/tags/stats") + async def get_tags_stats_route(request): + return await self.get_tags_stats(request) + + @routes.get("/prompt_manager/tags/filter") + async def get_tags_filter_route(request): + return await self.get_tags_filter(request) + + # Bulk tag operations (register BEFORE {tag_name} to avoid path param match) + @routes.post("/prompt_manager/tags/merge") + async def merge_tags_route(request): + return await self.merge_tags_endpoint(request) + + @routes.put("/prompt_manager/tags/{tag_name}") + async def rename_tag_route(request): + return await self.rename_tag_endpoint(request) + + @routes.delete("/prompt_manager/tags/{tag_name}") + async def delete_tag_route(request): + return await self.delete_tag_endpoint(request) + + @routes.get("/prompt_manager/tags/{tag_name}/prompts") + async def get_tag_prompts_route(request): + return await self.get_tag_prompts(request) + + @routes.get("/prompt_manager/tags") + async def get_tags_route(request): + return await self.get_tags(request) + + @routes.post("/prompt_manager/save") + async def save_prompt_route(request): + return await self.save_prompt(request) + + @routes.delete("/prompt_manager/delete/{prompt_id}") + async def delete_prompt_route(request): + return await self.delete_prompt(request) + + # Individual prompt management + @routes.put("/prompt_manager/prompts/{prompt_id}") + async def update_prompt_route(request): + return await self.update_prompt(request) + + @routes.put("/prompt_manager/prompts/{prompt_id}/rating") + async def update_rating_route(request): + return await self.update_prompt_rating(request) + + @routes.post("/prompt_manager/prompts/{prompt_id}/tags") + async def add_tag_route(request): + return await self.add_prompt_tag(request) + + @routes.delete("/prompt_manager/prompts/{prompt_id}/tags") + async def remove_tag_route(request): + return await self.remove_prompt_tag(request) + + @routes.post("/prompt_manager/prompts/tags") + async def add_tags_to_prompt_route(request): + return await self.add_tags_to_prompt(request) + + # Bulk operations + @routes.post("/prompt_manager/bulk/delete") + async def bulk_delete_route(request): + return await self.bulk_delete_prompts(request) + + @routes.post("/prompt_manager/bulk/tags") + async def bulk_add_tags_route(request): + return await self.bulk_add_tags(request) + + @routes.post("/prompt_manager/bulk/category") + async def bulk_set_category_route(request): + return await self.bulk_set_category(request) + + # Export functionality + @routes.get("/prompt_manager/export") + async def export_prompts_route(request): + return await self.export_prompts(request) + + async def search_prompts(self, request): + """Search for prompts using multiple filter criteria.""" + try: + text = request.query.get("text", "").strip() + category = request.query.get("category", "").strip() + tags_str = request.query.get("tags", "").strip() + min_rating = request.query.get("min_rating", 0) + limit = int(request.query.get("limit", 50)) + + tags = None + if tags_str: + tags = [tag.strip() for tag in tags_str.split(",") if tag.strip()] + + try: + min_rating = int(min_rating) if min_rating else None + except ValueError: + min_rating = None + + results = await self._run_in_executor( + self.db.search_prompts, + text=text if text else None, + category=category if category else None, + tags=tags, + rating_min=min_rating, + limit=limit, + ) + self._enrich_prompt_images(results) + + return web.json_response( + {"success": True, "results": results, "count": len(results)} + ) + + except Exception as e: + self.logger.error(f"Search error: {e}", exc_info=True) + return web.json_response( + {"success": False, "error": f"Search failed: {str(e)}", "results": []}, + status=500, + ) + + async def get_recent_prompts(self, request): + """Retrieve recently created prompts with pagination support.""" + try: + limit = int(request.query.get("limit", 50)) + page = int(request.query.get("page", 1)) + offset = int(request.query.get("offset", 0)) + + if page > 1 and offset == 0: + offset = (page - 1) * limit + + if limit > 1000: + limit = 1000 + 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']) + + 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) + return web.json_response( + { + "success": False, + "error": f"Failed to get recent prompts: {str(e)}", + "results": [], + "pagination": {"total": 0, "page": 1, "total_pages": 0} + }, + status=500, + ) + + async def get_categories(self, request): + """Retrieve all available prompt categories.""" + try: + categories = await self._run_in_executor(self.db.get_all_categories) + return web.json_response({"success": True, "categories": categories}) + 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": []}, + status=500, + ) + + async def get_tags(self, request): + """Retrieve all available prompt tags.""" + try: + tags = await self._run_in_executor(self.db.get_all_tags) + return web.json_response({"success": True, "tags": tags}) + 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": []}, + status=500, + ) + + async def get_tags_stats(self, request): + """Get tags with usage counts, search, sort, and pagination.""" + try: + try: + limit = int(request.query.get("limit", 50)) + offset = int(request.query.get("offset", 0)) + except (ValueError, TypeError): + return web.json_response( + {"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) + + 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) + + async def get_tag_prompts(self, request): + """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( + {"success": False, "error": "Tag name required"}, status=400 + ) + + try: + limit = int(request.query.get("limit", 20)) + offset = int(request.query.get("offset", 0)) + except (ValueError, TypeError): + return web.json_response( + {"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']) + + 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) + + async def get_tags_filter(self, request): + """Get prompts matching multiple tags with AND/OR mode, or untagged prompts.""" + try: + untagged = request.query.get("untagged", "").lower() == "true" + + if untagged: + try: + limit = int(request.query.get("limit", 20)) + offset = int(request.query.get("offset", 0)) + except (ValueError, TypeError): + return web.json_response( + {"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'] + } + }) + + tags_str = request.query.get("tags", "").strip() + if not tags_str: + return web.json_response( + {"success": False, "error": "Tags parameter required"}, status=400 + ) + + tags_list = [t.strip() for t in tags_str.split(",") if t.strip()] + mode = request.query.get("mode", "and").lower() + if mode not in ("and", "or"): + mode = "and" + + try: + limit = int(request.query.get("limit", 20)) + offset = int(request.query.get("offset", 0)) + except (ValueError, TypeError): + return web.json_response( + {"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']) + + 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) + + async def rename_tag_endpoint(self, request): + """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( + {"success": False, "error": "Tag name required"}, status=400 + ) + + try: + body = await request.json() + except Exception: + return web.json_response( + {"success": False, "error": "Invalid JSON body"}, status=400 + ) + new_name = body.get("new_name", "").strip() + if not new_name: + return web.json_response( + {"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) + resp = { + "success": True, + "old_name": tag_name, + "new_name": new_name, + "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" + return web.json_response(resp) + except Exception as e: + self.logger.error(f"Rename tag error: {e}", exc_info=True) + return web.json_response({"success": False, "error": str(e)}, status=500) + + async def delete_tag_endpoint(self, request): + """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) + resp = { + "success": True, + "tag_name": tag_name, + "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" + return web.json_response(resp) + except Exception as e: + self.logger.error(f"Delete tag error: {e}", exc_info=True) + return web.json_response({"success": False, "error": str(e)}, status=500) + + async def merge_tags_endpoint(self, request): + """Merge source tags into a target tag.""" + try: + try: + body = await request.json() + except Exception: + return web.json_response( + {"success": False, "error": "Invalid JSON body"}, status=400 + ) + source_tags = body.get("source_tags", []) + target_tag = body.get("target_tag", "").strip() + + if not source_tags: + return web.json_response( + {"success": False, "error": "Source tags required"}, status=400 + ) + if not target_tag: + return web.json_response( + {"success": False, "error": "Target tag required"}, status=400 + ) + + 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'] + } + 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) + return web.json_response({"success": False, "error": str(e)}, status=500) + + async def save_prompt(self, request): + """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, + ) + data = await request.json() + + text = data.get("text", "").strip() + if not text: + return web.json_response( + {"success": False, "error": "Text is required"}, status=400 + ) + + category = data.get("category", "").strip() or None + tags = data.get("tags", []) + rating = data.get("rating") or None + notes = data.get("notes", "").strip() or None + + try: + validate_prompt_text(text) + validate_category(category) + validate_tags(tags) + validate_rating(rating) + except ValueError as ve: + return web.json_response( + {"success": False, "error": str(ve)}, status=400 + ) + + 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) + if existing: + if any([category, tags, rating, notes]): + await self._run_in_executor( + self.db.update_prompt_metadata, + prompt_id=existing['id'], + category=category, + tags=tags, + rating=rating, + notes=notes + ) + 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, + text=text, + category=category, + tags=tags if tags else None, + rating=rating, + notes=notes, + prompt_hash=prompt_hash, + ) + + 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) + return web.json_response( + {"success": False, "error": f"Failed to save prompt: {str(e)}"}, + status=500, + ) + + async def delete_prompt(self, request): + """Delete a specific prompt by ID.""" + try: + prompt_id = int(request.match_info["prompt_id"]) + success = await self._run_in_executor(self.db.delete_prompt, prompt_id) + + if success: + return web.json_response( + {"success": True, "message": "Prompt deleted successfully"} + ) + else: + return web.json_response( + {"success": False, "error": "Prompt not found or could not be deleted"}, + status=404, + ) + + except ValueError: + return web.json_response( + {"success": False, "error": "Invalid prompt ID"}, status=400 + ) + except Exception as e: + self.logger.error(f"Delete error: {e}") + return web.json_response( + {"success": False, "error": f"Failed to delete prompt: {str(e)}"}, + status=500, + ) + + async def update_prompt(self, request): + """Update prompt text.""" + try: + from utils.validators import validate_prompt_text, sanitize_input + + prompt_id = int(request.match_info["prompt_id"]) + data = await request.json() + new_text = data.get("text", "").strip() + + if not new_text: + return web.json_response( + {"success": False, "error": "Text cannot be empty"}, status=400 + ) + + try: + validate_prompt_text(new_text) + except ValueError as ve: + return web.json_response( + {"success": False, "error": str(ve)}, status=400 + ) + + new_text = sanitize_input(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"} + ) + else: + return web.json_response( + {"success": False, "error": "Prompt not found"}, status=404 + ) + + except ValueError: + return web.json_response( + {"success": False, "error": "Invalid prompt ID"}, status=400 + ) + except Exception as e: + self.logger.error(f"Update prompt error: {e}") + return web.json_response( + {"success": False, "error": f"Failed to update prompt: {str(e)}"}, + status=500, + ) + + async def update_prompt_rating(self, request): + """Update prompt rating.""" + try: + from utils.validators import validate_rating + + prompt_id = int(request.match_info["prompt_id"]) + data = await request.json() + rating = data.get("rating") + + try: + validate_rating(rating) + except ValueError as ve: + return web.json_response( + {"success": False, "error": str(ve)}, status=400 + ) + + 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"} + ) + else: + return web.json_response( + {"success": False, "error": "Prompt not found"}, status=404 + ) + + except ValueError: + return web.json_response( + {"success": False, "error": "Invalid prompt ID"}, status=400 + ) + except Exception as e: + self.logger.error(f"Update rating error: {e}") + return web.json_response( + {"success": False, "error": f"Failed to update rating: {str(e)}"}, + status=500, + ) + + async def add_prompt_tag(self, request): + """Add tag to prompt.""" + try: + prompt_id = int(request.match_info["prompt_id"]) + data = await request.json() + new_tag = data.get("tag", "").strip() + + if not new_tag: + return web.json_response( + {"success": False, "error": "Tag cannot be empty"}, status=400 + ) + + prompt = await self._run_in_executor(self.db.get_prompt_by_id, prompt_id) + if not prompt: + return web.json_response( + {"success": False, "error": "Prompt not found"}, status=404 + ) + + current_tags = prompt.get("tags", []) + if not isinstance(current_tags, list): + current_tags = [] + + 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) + + return web.json_response( + {"success": True, "message": "Tag added successfully"} + ) + + except ValueError: + return web.json_response( + {"success": False, "error": "Invalid prompt ID"}, status=400 + ) + except Exception as e: + self.logger.error(f"Add tag error: {e}") + return web.json_response( + {"success": False, "error": f"Failed to add tag: {str(e)}"}, status=500 + ) + + async def add_tags_to_prompt(self, request): + """Add multiple tags to a single prompt.""" + try: + data = await request.json() + prompt_id = data.get("prompt_id") + new_tags = data.get("tags", []) + + if not prompt_id: + return web.json_response( + {"success": False, "error": "Prompt ID is required"}, status=400 + ) + + 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 + ) + + prompt = await self._run_in_executor(self.db.get_prompt_by_id, prompt_id) + if not prompt: + return web.json_response( + {"success": False, "error": "Prompt not found"}, status=404 + ) + + current_tags = prompt.get("tags", []) + if not isinstance(current_tags, list): + current_tags = [] + + tags_added = 0 + for new_tag in new_tags: + new_tag = new_tag.strip() + if new_tag and new_tag not in current_tags: + current_tags.append(new_tag) + tags_added += 1 + + if tags_added > 0: + 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: + message = "No new tags to add (all tags already exist)" + + return web.json_response( + {"success": True, "message": message, "tags_added": tags_added} + ) + + except ValueError: + return web.json_response( + {"success": False, "error": "Invalid prompt ID"}, status=400 + ) + except Exception as e: + self.logger.error(f"Add tags error: {e}") + return web.json_response( + {"success": False, "error": f"Failed to add tags: {str(e)}"}, status=500 + ) + + async def remove_prompt_tag(self, request): + """Remove tag from prompt.""" + try: + prompt_id = int(request.match_info["prompt_id"]) + data = await request.json() + tag_to_remove = data.get("tag", "").strip() + + prompt = await self._run_in_executor(self.db.get_prompt_by_id, prompt_id) + if not prompt: + return web.json_response( + {"success": False, "error": "Prompt not found"}, status=404 + ) + + current_tags = prompt.get("tags", []) + if not isinstance(current_tags, list): + current_tags = [] + + 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) + + return web.json_response( + {"success": True, "message": "Tag removed successfully"} + ) + + except ValueError: + return web.json_response( + {"success": False, "error": "Invalid prompt ID"}, status=400 + ) + except Exception as e: + self.logger.error(f"Remove tag error: {e}") + return web.json_response( + {"success": False, "error": f"Failed to remove tag: {str(e)}"}, + status=500, + ) + + async def bulk_delete_prompts(self, request): + """Bulk delete prompts.""" + try: + data = await request.json() + prompt_ids = data.get("prompt_ids", []) + + if not prompt_ids: + return web.json_response( + {"success": False, "error": "No prompt IDs provided"}, status=400 + ) + + 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, + }) + + except Exception as e: + self.logger.error(f"Bulk delete error: {e}") + return web.json_response( + {"success": False, "error": f"Failed to delete prompts: {str(e)}"}, + status=500, + ) + + async def bulk_add_tags(self, request): + """Bulk add tags to prompts.""" + try: + data = await request.json() + prompt_ids = data.get("prompt_ids", []) + new_tags = data.get("tags", []) + + if not prompt_ids or not new_tags: + return web.json_response( + {"success": False, "error": "No prompt IDs or tags provided"}, + status=400, + ) + + 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, + }) + + except Exception as e: + self.logger.error(f"Bulk add tags error: {e}") + return web.json_response( + {"success": False, "error": f"Failed to add tags: {str(e)}"}, status=500 + ) + + async def bulk_set_category(self, request): + """Bulk set category for prompts.""" + try: + data = await request.json() + prompt_ids = data.get("prompt_ids", []) + category = data.get("category", "").strip() + + if not prompt_ids: + return web.json_response( + {"success": False, "error": "No prompt IDs provided"}, status=400 + ) + + 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, + }) + + except Exception as e: + self.logger.error(f"Bulk set category error: {e}") + return web.json_response( + {"success": False, "error": f"Failed to set category: {str(e)}"}, + status=500, + ) + + async def export_prompts(self, request): + """Export all prompts to JSON.""" + try: + prompts = await self._run_in_executor(self.db.search_prompts, limit=10000) + + export_data = { + "export_date": datetime.datetime.now(datetime.timezone.utc).isoformat(), + "total_prompts": len(prompts), + "prompts": prompts, + } + + json_data = json.dumps(export_data, indent=2, ensure_ascii=False) + + return web.Response( + text=json_data, + content_type="application/json", + headers={ + "Content-Disposition": f'attachment; filename="prompt_manager_{datetime.datetime.now().strftime("%Y%m%d_%H%M%S")}.json"' + }, + ) + + except Exception as e: + self.logger.error(f"Export error: {e}") + return web.json_response( + {"success": False, "error": f"Failed to export prompts: {str(e)}"}, + status=500, + ) diff --git a/py/config.py b/py/config.py index 5de3937..10ece56 100644 --- a/py/config.py +++ b/py/config.py @@ -79,7 +79,7 @@ class GalleryConfig: PROCESSING_DELAY = 2.0 # Seconds to wait before processing new files # Prompt tracking settings - PROMPT_TIMEOUT = 120 # Seconds to keep prompt context active + 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 diff --git a/tailwind.config.js b/tailwind.config.js new file mode 100644 index 0000000..bb8b02d --- /dev/null +++ b/tailwind.config.js @@ -0,0 +1,11 @@ +module.exports = { + content: ["./web/**/*.html", "./web/js/**/*.js"], + darkMode: 'class', + theme: { + extend: { + colors: { + gray: { 850: '#1f2937', 875: '#1a202c', 925: '#0f172a' } + } + } + }, +} diff --git a/tests/test_api.py b/tests/test_api.py new file mode 100644 index 0000000..07664cb --- /dev/null +++ b/tests/test_api.py @@ -0,0 +1,320 @@ +""" +API endpoint integration tests for PromptManager. + +Uses aiohttp test client to test request/response contracts +without requiring a full ComfyUI server. +""" + +import json +import os +import sys +import tempfile +import unittest + +sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__)))) + +from aiohttp import web +from aiohttp.test_utils import AioHTTPTestCase + +from database.operations import PromptDatabase +from py.api import PromptManagerAPI +from utils.hashing import generate_prompt_hash + + +class APITestCase(AioHTTPTestCase): + """Base class that stands up a test aiohttp app with PromptManager routes.""" + + async def get_application(self): + self._temp_db = tempfile.NamedTemporaryFile(delete=False, suffix=".db") + self._temp_db.close() + + app = web.Application() + routes = web.RouteTableDef() + + self.api = PromptManagerAPI() + self.api.db = PromptDatabase(self._temp_db.name) + self.api.add_routes(routes) + app.router.add_routes(routes) + return app + + async def tearDownAsync(self): + 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) + + # ── helpers ──────────────────────────────────────────────────────── + + def _save_prompt(self, text="Test prompt", category=None, tags=None, rating=None): + return self.api.db.save_prompt( + text=text, + category=category, + tags=tags or [], + rating=rating, + prompt_hash=generate_prompt_hash(text), + ) + + +class TestHealthEndpoint(APITestCase): + + async def test_test_route(self): + resp = await self.client.request("GET", "/prompt_manager/test") + self.assertEqual(resp.status, 200) + data = await resp.json() + self.assertTrue(data["success"]) + self.assertIn("timestamp", data) + + +class TestRecentPrompts(APITestCase): + + async def test_recent_empty(self): + resp = await self.client.request("GET", "/prompt_manager/recent") + self.assertEqual(resp.status, 200) + data = await resp.json() + self.assertTrue(data["success"]) + self.assertEqual(data["pagination"]["total"], 0) + self.assertEqual(len(data["results"]), 0) + + async def test_recent_with_data(self): + self._save_prompt("First prompt") + self._save_prompt("Second prompt") + resp = await self.client.request("GET", "/prompt_manager/recent") + self.assertEqual(resp.status, 200) + data = await resp.json() + self.assertEqual(data["pagination"]["total"], 2) + self.assertEqual(len(data["results"]), 2) + + 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") + data = await resp.json() + self.assertEqual(len(data["results"]), 5) + self.assertEqual(data["pagination"]["total"], 15) + self.assertTrue(data["pagination"]["has_more"]) + + 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") + data = await resp.json() + self.assertEqual(len(data["results"]), 2) + self.assertFalse(data["pagination"]["has_more"]) + + +class TestSearch(APITestCase): + + async def test_search_by_text(self): + self._save_prompt("Beautiful mountain landscape") + self._save_prompt("City skyline at night") + resp = await self.client.request("GET", "/prompt_manager/search?text=mountain") + self.assertEqual(resp.status, 200) + data = await resp.json() + self.assertTrue(data["success"]) + self.assertEqual(len(data["results"]), 1) + self.assertIn("mountain", data["results"][0]["text"]) + + 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") + data = await resp.json() + self.assertEqual(len(data["results"]), 1) + + async def test_search_by_tag(self): + self._save_prompt("Tagged prompt", tags=["landscape", "sunset"]) + self._save_prompt("Other prompt", tags=["portrait"]) + resp = await self.client.request("GET", "/prompt_manager/search?tags=landscape") + data = await resp.json() + 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") + data = await resp.json() + self.assertTrue(data["success"]) + self.assertEqual(len(data["results"]), 0) + + +class TestSaveAndDelete(APITestCase): + + async def test_save_prompt(self): + resp = await self.client.request( + "POST", + "/prompt_manager/save", + json={"text": "New prompt via API", "category": "test", "tags": ["api", "test"]}, + ) + self.assertEqual(resp.status, 200) + data = await resp.json() + self.assertTrue(data["success"]) + self.assertIn("prompt_id", data) + + async def test_save_empty_text(self): + resp = await self.client.request("POST", "/prompt_manager/save", json={"text": ""}) + self.assertEqual(resp.status, 400) + data = await resp.json() + self.assertFalse(data["success"]) + + async def test_delete_prompt(self): + pid = self._save_prompt("To delete") + resp = await self.client.request("DELETE", f"/prompt_manager/delete/{pid}") + self.assertEqual(resp.status, 200) + data = await resp.json() + self.assertTrue(data["success"]) + + async def test_delete_nonexistent(self): + resp = await self.client.request("DELETE", "/prompt_manager/delete/99999") + self.assertEqual(resp.status, 404) + data = await resp.json() + self.assertFalse(data["success"]) + + +class TestUpdatePrompt(APITestCase): + + async def test_update_text(self): + pid = self._save_prompt("Original text") + resp = await self.client.request( + "PUT", + f"/prompt_manager/prompts/{pid}", + json={"text": "Updated text"}, + ) + self.assertEqual(resp.status, 200) + data = await resp.json() + self.assertTrue(data["success"]) + + async def test_update_rating(self): + pid = self._save_prompt("Ratable") + resp = await self.client.request( + "PUT", + f"/prompt_manager/prompts/{pid}/rating", + json={"rating": 4}, + ) + self.assertEqual(resp.status, 200) + data = await resp.json() + self.assertTrue(data["success"]) + + async def test_add_tag(self): + pid = self._save_prompt("Taggable", tags=["old"]) + resp = await self.client.request( + "POST", + f"/prompt_manager/prompts/{pid}/tags", + json={"tag": "new_tag"}, + ) + self.assertEqual(resp.status, 200) + data = await resp.json() + self.assertTrue(data["success"]) + + +class TestTags(APITestCase): + + async def test_get_all_tags(self): + self._save_prompt("P1", tags=["alpha", "beta"]) + self._save_prompt("P2", tags=["beta", "gamma"]) + resp = await self.client.request("GET", "/prompt_manager/tags") + self.assertEqual(resp.status, 200) + data = await resp.json() + self.assertTrue(data["success"]) + self.assertIn("tags", data) + + async def test_get_tag_stats(self): + self._save_prompt("P1", tags=["common", "rare"]) + self._save_prompt("P2", tags=["common"]) + resp = await self.client.request("GET", "/prompt_manager/tags/stats") + self.assertEqual(resp.status, 200) + data = await resp.json() + self.assertTrue(data["success"]) + + async def test_rename_tag(self): + self._save_prompt("P1", tags=["old_name"]) + resp = await self.client.request( + "PUT", + "/prompt_manager/tags/old_name", + json={"new_name": "new_name"}, + ) + self.assertEqual(resp.status, 200) + data = await resp.json() + self.assertTrue(data["success"]) + + async def test_delete_tag(self): + self._save_prompt("P1", tags=["removable"]) + resp = await self.client.request("DELETE", "/prompt_manager/tags/removable") + self.assertEqual(resp.status, 200) + data = await resp.json() + self.assertTrue(data["success"]) + + async def test_merge_tags(self): + self._save_prompt("P1", tags=["target"]) + self._save_prompt("P2", tags=["source"]) + resp = await self.client.request( + "POST", + "/prompt_manager/tags/merge", + json={"source_tags": ["source"], "target_tag": "target"}, + ) + self.assertEqual(resp.status, 200) + data = await resp.json() + self.assertTrue(data["success"]) + + +class TestCategories(APITestCase): + + async def test_get_categories(self): + self._save_prompt("P1", category="nature") + self._save_prompt("P2", category="urban") + resp = await self.client.request("GET", "/prompt_manager/categories") + self.assertEqual(resp.status, 200) + data = await resp.json() + self.assertTrue(data["success"]) + self.assertIn("categories", data) + + +class TestStats(APITestCase): + + async def test_get_stats(self): + self._save_prompt("P1", category="test", tags=["t1"], rating=4) + resp = await self.client.request("GET", "/prompt_manager/stats") + self.assertEqual(resp.status, 200) + data = await resp.json() + self.assertTrue(data["success"]) + self.assertIn("stats", data) + + +class TestExport(APITestCase): + + async def test_export_json(self): + self._save_prompt("Exportable", category="test", tags=["export"]) + resp = await self.client.request("GET", "/prompt_manager/export?format=json") + self.assertEqual(resp.status, 200) + # Export returns raw JSON file download, not a success envelope + self.assertIn("Content-Disposition", resp.headers) + text = await resp.text() + data = json.loads(text) + self.assertIn("prompts", data) + self.assertEqual(len(data["prompts"]), 1) + + +class TestResponseEnvelope(APITestCase): + """Verify all responses follow the {success: bool, ...} envelope.""" + + async def test_success_responses_have_success_true(self): + self._save_prompt("Envelope test") + endpoints = [ + ("GET", "/prompt_manager/test"), + ("GET", "/prompt_manager/recent"), + ("GET", "/prompt_manager/search"), + ("GET", "/prompt_manager/categories"), + ("GET", "/prompt_manager/tags"), + ("GET", "/prompt_manager/stats"), + ] + for method, path in endpoints: + 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}") + + async def test_error_responses_have_success_false(self): + resp = await self.client.request("DELETE", "/prompt_manager/delete/99999") + data = await resp.json() + self.assertIn("success", data) + self.assertFalse(data["success"]) + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/test_basic.py b/tests/test_basic.py index 22ab420..89c3159 100644 --- a/tests/test_basic.py +++ b/tests/test_basic.py @@ -218,36 +218,38 @@ class TestNodeIntegration(unittest.TestCase): if os.path.exists(self.temp_db.name): os.unlink(self.temp_db.name) - @patch('prompt_manager.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 mock_db = Mock() mock_db_class.return_value = mock_db - + mock_db.save_prompt.return_value = 1 + mock_db.get_prompt_by_hash.return_value = None + # Mock CLIP model mock_clip = Mock() mock_clip.tokenize.return_value = "mock_tokens" mock_clip.encode_from_tokens_scheduled.return_value = "mock_conditioning" - + # Import and test the node from prompt_manager import PromptManager - + node = PromptManager() - + # Test encoding - result = node.encode( + 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") mock_clip.encode_from_tokens_scheduled.assert_called_once_with("mock_tokens") - - # Verify result - PromptManager returns both conditioning and prompt text - self.assertEqual(result, ("mock_conditioning", "Test prompt")) + + # 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.""" diff --git a/tests/test_database.py b/tests/test_database.py new file mode 100644 index 0000000..666cc15 --- /dev/null +++ b/tests/test_database.py @@ -0,0 +1,467 @@ +""" +Comprehensive database layer tests for PromptManager. + +Tests all CRUD operations, tag junction tables, pagination, +search, statistics, image linking, and edge cases using +an in-memory SQLite database. +""" + +import os +import sys +import tempfile +import unittest + +sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__)))) + +from database.operations import PromptDatabase +from utils.hashing import generate_prompt_hash + + +class DatabaseTestCase(unittest.TestCase): + """Base class with temp database setup/teardown.""" + + def setUp(self): + self.temp_db = tempfile.NamedTemporaryFile(delete=False, suffix=".db") + self.temp_db.close() + self.db = PromptDatabase(self.temp_db.name) + + def tearDown(self): + if os.path.exists(self.temp_db.name): + os.unlink(self.temp_db.name) + wal = self.temp_db.name + "-wal" + shm = self.temp_db.name + "-shm" + for f in (wal, shm): + if os.path.exists(f): + os.unlink(f) + + 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, + category=category, + tags=tags or [], + rating=rating, + notes=notes, + prompt_hash=generate_prompt_hash(text), + ) + + +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) + prompt = self.db.get_prompt_by_id(pid) + self.assertEqual(prompt["text"], "A beautiful sunset") + self.assertEqual(prompt["category"], "nature") + self.assertIn("sunset", prompt["tags"]) + self.assertIn("sky", prompt["tags"]) + self.assertEqual(prompt["rating"], 5) + + def test_save_minimal(self): + pid = self._save("Minimal prompt") + prompt = self.db.get_prompt_by_id(pid) + self.assertEqual(prompt["text"], "Minimal prompt") + self.assertIsNone(prompt["category"]) + self.assertEqual(prompt["tags"], []) + self.assertIsNone(prompt["rating"]) + + def test_get_nonexistent_prompt(self): + result = self.db.get_prompt_by_id(99999) + self.assertIsNone(result) + + def test_get_by_hash(self): + text = "Hash test prompt" + h = generate_prompt_hash(text) + self._save(text) + found = self.db.get_prompt_by_hash(h) + self.assertIsNotNone(found) + self.assertEqual(found["text"], text) + + def test_get_by_hash_nonexistent(self): + result = self.db.get_prompt_by_hash("nonexistent_hash_value") + self.assertIsNone(result) + + 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") + prompt = self.db.get_prompt_by_id(pid) + self.assertEqual(prompt["category"], "new") + self.assertIn("new_tag", prompt["tags"]) + self.assertNotIn("old_tag", prompt["tags"]) + self.assertEqual(prompt["rating"], 5) + self.assertEqual(prompt["notes"], "updated") + + def test_update_partial_metadata(self): + pid = self._save("Partial update", category="keep", tags=["keep_tag"], rating=3) + self.db.update_prompt_metadata(pid, rating=1) + prompt = self.db.get_prompt_by_id(pid) + self.assertEqual(prompt["rating"], 1) + # Category and tags should be unchanged + self.assertEqual(prompt["category"], "keep") + self.assertIn("keep_tag", prompt["tags"]) + + def test_delete_prompt(self): + pid = self._save("To be deleted") + result = self.db.delete_prompt(pid) + self.assertTrue(result) + self.assertIsNone(self.db.get_prompt_by_id(pid)) + + def test_delete_nonexistent(self): + result = self.db.delete_prompt(99999) + self.assertFalse(result) + + +class TestTagJunctionTables(DatabaseTestCase): + """Test normalized tag storage via junction tables.""" + + def test_tags_stored_in_junction_table(self): + pid = self._save("Tagged prompt", tags=["alpha", "beta"]) + prompt = self.db.get_prompt_by_id(pid) + self.assertEqual(sorted(prompt["tags"]), ["alpha", "beta"]) + + def test_get_all_tags(self): + self._save("P1", tags=["a", "b"]) + self._save("P2", tags=["b", "c"]) + all_tags = self.db.get_all_tags() + self.assertEqual(sorted(all_tags), ["a", "b", "c"]) + + def test_get_tags_with_counts(self): + self._save("P1", tags=["common", "rare"]) + self._save("P2", tags=["common"]) + self._save("P3", tags=["common", "other"]) + result = self.db.get_tags_with_counts() + counts_dict = {t["name"]: t["count"] for t in result["tags"]} + self.assertEqual(counts_dict["common"], 3) + self.assertEqual(counts_dict["rare"], 1) + self.assertEqual(counts_dict["other"], 1) + + def test_set_prompt_tags(self): + pid = self._save("Retaggable", tags=["old"]) + self.db.set_prompt_tags(pid, ["new1", "new2"]) + prompt = self.db.get_prompt_by_id(pid) + self.assertEqual(sorted(prompt["tags"]), ["new1", "new2"]) + + def test_set_empty_tags(self): + pid = self._save("Clear tags", tags=["remove_me"]) + self.db.set_prompt_tags(pid, []) + prompt = self.db.get_prompt_by_id(pid) + self.assertEqual(prompt["tags"], []) + + def test_rename_tag(self): + self._save("P1", tags=["old_name"]) + self._save("P2", tags=["old_name", "other"]) + self.db.rename_tag_all_prompts("old_name", "new_name") + all_tags = self.db.get_all_tags() + self.assertIn("new_name", all_tags) + self.assertNotIn("old_name", all_tags) + + def test_delete_tag(self): + self._save("P1", tags=["keep", "remove"]) + self._save("P2", tags=["remove"]) + self.db.delete_tag_all_prompts("remove") + all_tags = self.db.get_all_tags() + self.assertIn("keep", all_tags) + self.assertNotIn("remove", all_tags) + + def test_merge_tags(self): + self._save("P1", tags=["target"]) + self._save("P2", tags=["source1"]) + self._save("P3", tags=["source2", "target"]) + self.db.merge_tags(["source1", "source2"], "target") + all_tags = self.db.get_all_tags() + self.assertIn("target", all_tags) + self.assertNotIn("source1", all_tags) + self.assertNotIn("source2", all_tags) + # All prompts should have the target tag + result = self.db.get_tags_with_counts() + counts = {t["name"]: t["count"] for t in result["tags"]} + self.assertEqual(counts["target"], 3) + + def test_bulk_add_tags(self): + p1 = self._save("P1", tags=["existing"]) + p2 = self._save("P2") + self.db.bulk_add_tags([p1, p2], ["bulk1", "bulk2"]) + prompt1 = self.db.get_prompt_by_id(p1) + prompt2 = self.db.get_prompt_by_id(p2) + self.assertIn("bulk1", prompt1["tags"]) + self.assertIn("bulk2", prompt1["tags"]) + self.assertIn("existing", prompt1["tags"]) + self.assertIn("bulk1", prompt2["tags"]) + + def test_untagged_prompts(self): + self._save("Tagged", tags=["has_tag"]) + self._save("Untagged1") + self._save("Untagged2") + count = self.db.get_untagged_prompts_count() + self.assertEqual(count, 2) + result = self.db.get_untagged_prompts() + self.assertEqual(len(result["prompts"]), 2) + + +class TestSearch(DatabaseTestCase): + """Test search and filter operations.""" + + 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) + + def test_search_by_text(self): + results = self.db.search_prompts(text="mountain") + self.assertEqual(len(results), 1) + self.assertIn("mountain", results[0]["text"]) + + def test_search_by_text_case_insensitive(self): + results = self.db.search_prompts(text="MOUNTAIN") + self.assertEqual(len(results), 1) + + def test_search_by_category(self): + results = self.db.search_prompts(category="urban") + self.assertEqual(len(results), 1) + self.assertEqual(results[0]["category"], "urban") + + def test_search_by_tag(self): + results = self.db.search_prompts(tags=["art"]) + self.assertEqual(len(results), 2) + + def test_search_by_multiple_tags(self): + results = self.db.search_prompts(tags=["art", "geometric"]) + self.assertEqual(len(results), 1) + self.assertIn("geometric", results[0]["text"].lower()) + + def test_search_by_rating_min(self): + results = self.db.search_prompts(rating_min=4) + self.assertEqual(len(results), 2) + + def test_search_by_rating_range(self): + results = self.db.search_prompts(rating_min=3, rating_max=4) + self.assertEqual(len(results), 2) + + def test_search_combined_filters(self): + results = self.db.search_prompts(text="landscape", rating_min=4) + self.assertEqual(len(results), 1) + + def test_search_no_results(self): + results = self.db.search_prompts(text="nonexistent_query_xyz") + self.assertEqual(len(results), 0) + + def test_search_with_limit(self): + results = self.db.search_prompts(limit=2) + self.assertEqual(len(results), 2) + + +class TestPagination(DatabaseTestCase): + """Test pagination in get_recent_prompts.""" + + def setUp(self): + super().setUp() + for i in range(15): + self._save(f"Prompt number {i:02d}") + + def test_first_page(self): + result = self.db.get_recent_prompts(limit=5, offset=0) + self.assertEqual(len(result["prompts"]), 5) + self.assertEqual(result["total"], 15) + self.assertTrue(result["has_more"]) + self.assertEqual(result["page"], 1) + self.assertEqual(result["total_pages"], 3) + + def test_middle_page(self): + result = self.db.get_recent_prompts(limit=5, offset=5) + self.assertEqual(len(result["prompts"]), 5) + self.assertTrue(result["has_more"]) + self.assertEqual(result["page"], 2) + + def test_last_page(self): + result = self.db.get_recent_prompts(limit=5, offset=10) + self.assertEqual(len(result["prompts"]), 5) + self.assertFalse(result["has_more"]) + self.assertEqual(result["page"], 3) + + def test_beyond_last_page(self): + result = self.db.get_recent_prompts(limit=5, offset=20) + self.assertEqual(len(result["prompts"]), 0) + self.assertFalse(result["has_more"]) + + def test_empty_database(self): + # Use a fresh empty DB + empty_db_file = tempfile.NamedTemporaryFile(delete=False, suffix=".db") + empty_db_file.close() + try: + empty_db = PromptDatabase(empty_db_file.name) + result = empty_db.get_recent_prompts(limit=10, offset=0) + self.assertEqual(result["total"], 0) + self.assertEqual(len(result["prompts"]), 0) + self.assertFalse(result["has_more"]) + finally: + os.unlink(empty_db_file.name) + + def test_total_count_is_integer(self): + result = self.db.get_recent_prompts(limit=5) + self.assertIsInstance(result["total"], int) + self.assertIsInstance(result["has_more"], bool) + + +class TestStatistics(DatabaseTestCase): + """Test get_statistics method.""" + + def test_empty_database_statistics(self): + stats = self.db.get_statistics() + self.assertEqual(stats["total_prompts"], 0) + self.assertEqual(stats["total_categories"], 0) + self.assertIsNone(stats.get("average_rating") or stats.get("avg_rating")) + self.assertEqual(stats["total_tags"], 0) + + def test_populated_statistics(self): + self._save("P1", category="cat_a", tags=["t1", "t2"], rating=4) + self._save("P2", category="cat_b", tags=["t2", "t3"], rating=2) + self._save("P3", category="cat_a", tags=["t1"]) + stats = self.db.get_statistics() + self.assertEqual(stats["total_prompts"], 3) + self.assertEqual(stats["total_categories"], 2) + self.assertEqual(stats["total_tags"], 3) + + +class TestCategories(DatabaseTestCase): + """Test category operations.""" + + def test_get_prompts_by_category(self): + self._save("P1", category="nature") + self._save("P2", category="nature") + self._save("P3", category="urban") + results = self.db.get_prompts_by_category("nature") + self.assertEqual(len(results), 2) + + def test_get_all_categories(self): + self._save("P1", category="nature") + self._save("P2", category="urban") + self._save("P3", category="nature") + categories = self.db.get_all_categories() + self.assertEqual(sorted(categories), ["nature", "urban"]) + + +class TestTopRated(DatabaseTestCase): + """Test top-rated prompt retrieval.""" + + def test_get_top_rated(self): + self._save("Low", rating=1) + self._save("High", rating=5) + self._save("Mid", rating=3) + self._save("Unrated") + results = self.db.get_top_rated_prompts(limit=2) + self.assertEqual(len(results), 2) + self.assertEqual(results[0]["rating"], 5) + self.assertEqual(results[1]["rating"], 3) + + +class TestDuplicateDetection(DatabaseTestCase): + """Test duplicate handling.""" + + def test_same_hash_detected(self): + text = "Duplicate content" + h = generate_prompt_hash(text) + pid1 = self.db.save_prompt(text=text, prompt_hash=h) + existing = self.db.get_prompt_by_hash(h) + self.assertIsNotNone(existing) + self.assertEqual(existing["id"], pid1) + + def test_cleanup_duplicates(self): + # Save multiple prompts first + self._save("Unique 1") + self._save("Unique 2") + removed = self.db.cleanup_duplicates() + self.assertEqual(removed, 0) + + +class TestImageOperations(DatabaseTestCase): + """Test image linking and retrieval.""" + + def _link_image(self, prompt_id, path="/fake/path/image.png"): + return self.db.link_image_to_prompt( + prompt_id=str(prompt_id), + image_path=path, + ) + + def test_save_and_get_image(self): + pid = self._save("Prompt with image") + img_id = self._link_image(pid, "/fake/path/image.png") + self.assertIsNotNone(img_id) + images = self.db.get_prompt_images(str(pid)) + self.assertEqual(len(images), 1) + self.assertEqual(images[0]["filename"], "image.png") + + def test_image_count(self): + pid = self._save("Multi-image prompt") + for i in range(3): + self._link_image(pid, f"/fake/path/img{i}.png") + images = self.db.get_prompt_images(str(pid)) + self.assertEqual(len(images), 3) + + def test_delete_prompt_cascades_images(self): + pid = self._save("Cascade test") + self._link_image(pid, "/fake/path/img.png") + self.db.delete_prompt(pid) + images = self.db.get_prompt_images(str(pid)) + self.assertEqual(len(images), 0) + + +class TestEdgeCases(DatabaseTestCase): + """Test edge cases and boundary conditions.""" + + def test_special_characters_in_text(self): + pid = self._save("Prompt with 'quotes' and \"double quotes\" and ") + prompt = self.db.get_prompt_by_id(pid) + self.assertIn("quotes", prompt["text"]) + + def test_unicode_text(self): + pid = self._save("日本語テスト prompt with émojis 🎨") + prompt = self.db.get_prompt_by_id(pid) + self.assertIn("日本語", prompt["text"]) + + def test_very_long_text(self): + long_text = "word " * 1000 + pid = self._save(long_text.strip()) + prompt = self.db.get_prompt_by_id(pid) + 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"]) + prompt = self.db.get_prompt_by_id(pid) + self.assertEqual(len(prompt["tags"]), 3) + + def test_empty_category_string(self): + pid = self._save("Empty cat", category="") + prompt = self.db.get_prompt_by_id(pid) + # Empty string category is stored as-is + self.assertIn(prompt["category"], ["", None]) + + def test_rating_boundary_values(self): + p1 = self._save("Rating 1", rating=1) + p5 = self._save("Rating 5", rating=5) + self.assertEqual(self.db.get_prompt_by_id(p1)["rating"], 1) + self.assertEqual(self.db.get_prompt_by_id(p5)["rating"], 5) + + +class TestPreviewImages(DatabaseTestCase): + """Test _attach_preview_images functionality.""" + + def test_preview_images_attached(self): + pid = self._save("Preview test") + for i in range(5): + self.db.link_image_to_prompt( + prompt_id=str(pid), + image_path=f"/fake/img{i}.png", + ) + result = self.db.get_recent_prompts(limit=10) + prompt = result["prompts"][0] + # Should have preview images (max 3) and total count + self.assertIn("images", prompt) + self.assertLessEqual(len(prompt["images"]), 3) + self.assertEqual(prompt["image_count"], 5) + + +if __name__ == "__main__": + unittest.main() diff --git a/utils/diagnostics.py b/utils/diagnostics.py index 88e67cf..d6bccb4 100644 --- a/utils/diagnostics.py +++ b/utils/diagnostics.py @@ -246,7 +246,7 @@ class GalleryDiagnostics: f.write("test") os.remove(test_file) can_write = True - except: + except (IOError, OSError): can_write = False self.logger.info(f" [EDIT] Can write to directory: {can_write}") diff --git a/utils/image_monitor.py b/utils/image_monitor.py index 271330f..32015d8 100644 --- a/utils/image_monitor.py +++ b/utils/image_monitor.py @@ -61,8 +61,16 @@ class ImageGenerationHandler(FileSystemEventHandler): self.db_manager = db_manager self.prompt_tracker = prompt_tracker self.metadata_extractor = ComfyUIMetadataExtractor() - self.processing_delay = 2.0 # Wait 2 seconds before processing 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') def on_created(self, event): """Handle filesystem creation events. @@ -96,7 +104,7 @@ class ImageGenerationHandler(FileSystemEventHandler): # Skip files in thumbnails directory - those are derivatives, not generated images if '/thumbnails/' in filepath or '\\thumbnails\\' in filepath: return False - return filepath.lower().endswith(('.png', '.jpg', '.jpeg', '.webp', '.gif')) + return filepath.lower().endswith(self.supported_extensions) def process_new_image(self, image_path: str): """Process a newly created image file for gallery integration. @@ -288,6 +296,15 @@ class ImageMonitor: self.logger.warning("Image monitoring already running") return + # 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 + except Exception: + pass + # Check config first, then auto-detect if not configured if not output_directories: try: diff --git a/utils/logging_config.py b/utils/logging_config.py index ef29122..2d4535d 100644 --- a/utils/logging_config.py +++ b/utils/logging_config.py @@ -9,6 +9,7 @@ This module provides centralized logging configuration with support for: - Log viewer API integration """ +import collections import logging import logging.handlers import os @@ -70,7 +71,7 @@ class PromptManagerLogger: } # Memory buffer for recent logs (for web viewer) - self._log_buffer = [] + self._log_buffer = collections.deque(maxlen=self.config['buffer_size']) self._buffer_lock = threading.Lock() # Initialize loggers @@ -176,10 +177,6 @@ class PromptManagerLogger: } self._log_buffer.append(log_entry) - - # Maintain buffer size limit - if len(self._log_buffer) > self.config['buffer_size']: - self._log_buffer = self._log_buffer[-self.config['buffer_size']:] def get_recent_logs(self, limit: int = 100, level: Optional[str] = None) -> List[Dict[str, Any]]: """Get recent log entries from memory buffer. @@ -192,14 +189,14 @@ class PromptManagerLogger: List of log entry dictionaries, most recent first """ with self._buffer_lock: - logs = self._log_buffer[:] - + logs = list(self._log_buffer) + # Filter by level if specified 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] - + # Return most recent first return list(reversed(logs[-limit:])) diff --git a/utils/prompt_tracker.py b/utils/prompt_tracker.py index 597933c..3ab1523 100644 --- a/utils/prompt_tracker.py +++ b/utils/prompt_tracker.py @@ -75,8 +75,15 @@ class PromptTracker: self._local = threading.local() self.active_prompts = {} # Global tracking for multiple threads self.lock = threading.Lock() - self.cleanup_interval = 300 # 5 minutes - self.prompt_timeout = 600 # 10 minutes (increased for longer generations) + + # 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 # Start cleanup thread self.cleanup_thread = threading.Thread(target=self._cleanup_expired_prompts, daemon=True) diff --git a/utils/validators.py b/utils/validators.py index 0d369ea..6796d2a 100644 --- a/utils/validators.py +++ b/utils/validators.py @@ -135,10 +135,10 @@ def validate_tags(tags: Union[str, List[str], None]) -> bool: if len(tag.strip()) > 50: raise ValueError("Individual tags cannot exceed 50 characters") - - # Check for invalid characters (optional - you can adjust this) - if not re.match(r'^[a-zA-Z0-9\s\-_]+$', tag.strip()): - raise ValueError(f"Tag '{tag}' contains invalid characters") + + # Reject control characters and null bytes + 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") @@ -177,10 +177,10 @@ def validate_category(category: Optional[str]) -> bool: if len(category) > 100: raise ValueError("Category cannot exceed 100 characters") - - # Check for invalid characters (adjust as needed) - if not re.match(r'^[a-zA-Z0-9\s\-_]+$', category): - raise ValueError("Category contains invalid characters") + + # Reject control characters and null bytes + if re.search(r'[\x00-\x1f]', category): + raise ValueError("Category contains invalid control characters") return True diff --git a/web/admin.html b/web/admin.html index ab5614f..f05083f 100644 --- a/web/admin.html +++ b/web/admin.html @@ -4,25 +4,9 @@ PromptManager - Admin Dashboard - + - - + \ No newline at end of file diff --git a/web/gallery.html b/web/gallery.html index c0b9ba6..04c2e62 100644 --- a/web/gallery.html +++ b/web/gallery.html @@ -4,25 +4,9 @@ PromptManager - Gallery - + - - + \ No newline at end of file diff --git a/web/js/admin.js b/web/js/admin.js new file mode 100644 index 0000000..11f0b5e --- /dev/null +++ b/web/js/admin.js @@ -0,0 +1,4199 @@ + class PromptAdmin { + constructor() { + this.prompts = []; + this.selectedPrompts = new Set(); + this.settings = { + resultTimeout: 5, + webuiDisplayMode: 'popup', + }; + this.categories = []; + this.tags = []; + this.imageViewMode = 'fit'; // 'fit' or 'full' + this.naturalImageSize = { width: 0, height: 0 }; + + // Pagination state + this.pagination = { + currentPage: 1, + limit: 50, + total: 0, + totalPages: 1 + }; + + this.tagsPage = null; + + this.init(); + } + + init() { + this.bindEvents(); + this.initRouter(); + this.loadInitialData(); + } + + initRouter() { + window.addEventListener('hashchange', () => this.handleRoute()); + this.handleRoute(); + } + + handleRoute() { + const hash = window.location.hash; + if (hash === '#/tags' || hash.startsWith('#/tags/')) { + this.showTagsPage(); + } else { + this.showDashboard(); + } + } + + showDashboard() { + const dashboard = document.getElementById('dashboardView'); + const tagsPage = document.getElementById('tagsPageView'); + if (dashboard) dashboard.classList.remove('hidden'); + if (tagsPage) tagsPage.classList.add('hidden'); + if (this.tagsPage) { + this.tagsPage.destroy(); + this.tagsPage = null; + } + } + + showTagsPage() { + const dashboard = document.getElementById('dashboardView'); + const tagsPage = document.getElementById('tagsPageView'); + if (dashboard) dashboard.classList.add('hidden'); + if (tagsPage) tagsPage.classList.remove('hidden'); + + // Lazy-initialize TagsPageManager + if (!this.tagsPage && typeof TagsPageManager !== 'undefined') { + this.tagsPage = new TagsPageManager(this); + } else if (this.tagsPage) { + this.tagsPage.loadTagsList(); + } + } + + bindEvents() { + // Search + document.getElementById("searchBtn").addEventListener("click", () => this.search()); + document.getElementById("searchText").addEventListener("keyup", (e) => { + if (e.key === "Enter") this.search(); + }); + + // Bulk actions + document.getElementById("selectAll").addEventListener("change", (e) => + this.toggleSelectAll(e.target.checked) + ); + document.getElementById("bulkDeleteBtn").addEventListener("click", () => this.bulkDelete()); + document.getElementById("bulkTagBtn").addEventListener("click", () => this.showBulkTagModal()); + document.getElementById("bulkCategoryBtn").addEventListener("click", () => this.showBulkCategoryModal()); + document.getElementById("exportBtn").addEventListener("click", () => this.exportPrompts()); + document.getElementById("settingsBtn").addEventListener("click", () => this.showSettingsModal()); + document.getElementById("statsBtn").addEventListener("click", () => this.showStats()); + document.getElementById("diagnosticsBtn").addEventListener("click", () => this.showDiagnosticsModal()); + document.getElementById("maintenanceBtn").addEventListener("click", () => this.showMaintenanceModal()); + document.getElementById("backupBtn").addEventListener("click", () => this.backupDatabase()); + document.getElementById("restoreBtn").addEventListener("click", () => this.showRestoreModal()); + document.getElementById("scanBtn").addEventListener("click", () => this.showScanModal()); + document.getElementById("autoTagBtn").addEventListener("click", () => this.showAutoTagModal()); + document.getElementById("addPromptBtn").addEventListener("click", () => this.showAddPromptModal()); + document.getElementById("logsBtn").addEventListener("click", () => this.showLogsModal()); + document.getElementById("metadataBtn").addEventListener("click", () => this.openMetadataViewer()); + document.getElementById("galleryBtn").addEventListener("click", () => this.openGallery()); + + // Pagination controls + document.getElementById("limitSelector").addEventListener("change", (e) => this.changeLimit(parseInt(e.target.value))); + document.getElementById("firstPageBtn").addEventListener("click", () => this.goToPage(1)); + document.getElementById("prevPageBtn").addEventListener("click", () => this.goToPage(this.pagination.currentPage - 1)); + document.getElementById("nextPageBtn").addEventListener("click", () => this.goToPage(this.pagination.currentPage + 1)); + document.getElementById("lastPageBtn").addEventListener("click", () => this.goToPage(this.pagination.totalPages)); + + // Modals + this.bindModalEvents(); + + // Auto-search on filter changes + ["searchCategory"].forEach((id) => { + document.getElementById(id).addEventListener("change", () => this.search()); + }); + + // Sort dropdown change + document.getElementById("sortBy").addEventListener("change", () => this.renderPrompts()); + } + + bindModalEvents() { + // Settings modal + document.getElementById("saveSettings").addEventListener("click", () => this.saveSettings()); + document.getElementById("cancelSettings").addEventListener("click", () => this.hideModal("settingsModal")); + document.getElementById("refreshMonitoringStatus").addEventListener("click", () => this.updateMonitoringStatus()); + + // Bulk tag modal + document.getElementById("confirmBulkTag").addEventListener("click", () => this.confirmBulkTag()); + document.getElementById("cancelBulkTag").addEventListener("click", () => this.hideModal("bulkTagModal")); + + // Individual tag modal + document.getElementById("confirmIndividualTag").addEventListener("click", () => this.confirmIndividualTag()); + document.getElementById("cancelIndividualTag").addEventListener("click", () => this.hideModal("individualTagModal")); + + // Bulk category modal + document.getElementById("confirmBulkCategory").addEventListener("click", () => this.confirmBulkCategory()); + document.getElementById("cancelBulkCategory").addEventListener("click", () => this.hideModal("bulkCategoryModal")); + + // Restore modal + document.getElementById("confirmRestore").addEventListener("click", () => this.confirmRestore()); + document.getElementById("cancelRestore").addEventListener("click", () => this.hideModal("restoreModal")); + + // Diagnostics modal + document.getElementById("runDiagnosticsBtn").addEventListener("click", () => this.runDiagnostics()); + document.getElementById("testImageLinkBtn").addEventListener("click", () => this.testImageLink()); + + // Maintenance modal + document.getElementById("runMaintenanceBtn").addEventListener("click", () => this.runMaintenance()); + document.getElementById("selectAllMaintenanceBtn").addEventListener("click", () => this.selectAllMaintenance()); + document.getElementById("clearAllMaintenanceBtn").addEventListener("click", () => this.clearAllMaintenance()); + + // Scan modal + document.getElementById("startScan").addEventListener("click", () => this.startScan()); + document.getElementById("cancelScan").addEventListener("click", () => this.hideModal("scanModal")); + document.getElementById("quickBackupBtn").addEventListener("click", () => this.quickBackup()); + + // Logs modal + document.getElementById("refreshLogsBtn").addEventListener("click", () => this.refreshLogs()); + document.getElementById("downloadLogsBtn").addEventListener("click", () => this.downloadLogs()); + document.getElementById("clearLogsBtn").addEventListener("click", () => this.clearLogs()); + document.getElementById("updateLogConfigBtn").addEventListener("click", () => this.updateLogConfig()); + + // Auto Tag modals + document.getElementById("startAutoTagBtn").addEventListener("click", () => this.startAutoTag()); + document.getElementById("startReviewBtn").addEventListener("click", () => this.startReview()); + document.getElementById("skipReviewBtn").addEventListener("click", () => this.skipReviewImage()); + document.getElementById("applyReviewBtn").addEventListener("click", () => this.applyReviewTags()); + document.getElementById("cancelAutoTagBtn").addEventListener("click", () => this.cancelAutoTag()); + document.getElementById("cancelDownloadBtn").addEventListener("click", () => this.cancelDownload()); + + // Re-tag confirmation modal buttons + document.getElementById("retagSkipBtn").addEventListener("click", () => this.handleRetagChoice('skip')); + document.getElementById("retagSkipAllBtn").addEventListener("click", () => this.handleRetagChoice('skipAll')); + document.getElementById("retagConfirmBtn").addEventListener("click", () => this.handleRetagChoice('retag')); + + // Tags accordion toggle + document.getElementById("reviewTagsToggle").addEventListener("click", () => this.toggleReviewTags()); + + // Add Prompt modal + document.getElementById("saveNewPromptBtn").addEventListener("click", () => this.saveNewPrompt()); + this.setupAddPromptModal(); + + // Close modals on backdrop click + document.querySelectorAll("[id$='Modal']").forEach((modal) => { + modal.addEventListener("click", (e) => { + if (e.target === modal) { + modal.classList.add("hidden"); + modal.classList.remove("flex"); + document.body.style.overflow = ""; + } + }); + }); + } + + async loadInitialData() { + try { + await Promise.all([ + this.loadStatistics(), + this.loadCategories(), + this.loadTags(), + this.loadRecentPrompts(), + this.loadSettings(), + ]); + } catch (error) { + this.showNotification("Failed to load initial data", "error"); + console.error("Load error:", error); + } + } + + async loadSettings() { + try { + const response = await fetch("/prompt_manager/settings"); + if (response.ok) { + const data = await response.json(); + if (data.success && data.settings) { + this.settings.resultTimeout = data.settings.result_timeout || 5; + this.settings.webuiDisplayMode = data.settings.webui_display_mode || 'popup'; + this.settings.galleryRootPath = data.settings.gallery_root_path || ''; + this.settings.monitoredDirectories = data.settings.monitored_directories || []; + } + } + } catch (error) { + console.error("Settings load error:", error); + } + } + + async loadStatistics() { + try { + const response = await fetch("/prompt_manager/stats"); + if (response.ok) { + const data = await response.json(); + console.log("Stats response:", data); // Debug log + + if (data.success) { + // Try different possible response structures + const stats = data.stats || data; + document.getElementById("totalPrompts").textContent = stats.total_prompts || data.total_prompts || 0; + document.getElementById("totalCategories").textContent = stats.total_categories || data.total_categories || 0; + document.getElementById("totalTags").textContent = stats.total_tags || data.total_tags || 0; + document.getElementById("avgRating").textContent = + (stats.avg_rating || data.avg_rating) ? (stats.avg_rating || data.avg_rating).toFixed(1) : "N/A"; + } else { + console.error("Stats API returned success=false:", data); + } + } else { + console.error("Stats API HTTP error:", response.status, response.statusText); + } + } catch (error) { + console.error("Stats error:", error); + // Set some default values if API fails + document.getElementById("totalPrompts").textContent = this.prompts.length || 0; + document.getElementById("totalCategories").textContent = this.categories.length || 0; + document.getElementById("totalTags").textContent = this.tags.length || 0; + document.getElementById("avgRating").textContent = "N/A"; + } + } + + async loadCategories() { + try { + const response = await fetch("/prompt_manager/categories"); + if (response.ok) { + const data = await response.json(); + if (data.success) { + this.categories = data.categories; + this.populateCategoryDropdown(); + } + } + } catch (error) { + console.error("Categories error:", error); + } + } + + async loadTags() { + try { + const response = await fetch("/prompt_manager/tags"); + if (response.ok) { + const data = await response.json(); + if (data.success) { + this.tags = data.tags; + } + } + } catch (error) { + console.error("Tags error:", error); + } + } + + populateCategoryDropdown() { + const select = document.getElementById("searchCategory"); + select.innerHTML = ''; + this.categories.forEach((category) => { + const option = document.createElement("option"); + option.value = category; + option.textContent = category; + select.appendChild(option); + }); + } + + async loadRecentPrompts(page = 1) { + try { + this.pagination.currentPage = page; + const offset = (page - 1) * this.pagination.limit; + const response = await fetch(`/prompt_manager/recent?limit=${this.pagination.limit}&offset=${offset}&page=${page}`); + if (response.ok) { + const data = await response.json(); + if (data.success) { + this.prompts = data.results; + this.pagination = { + ...this.pagination, + total: data.pagination.total, + totalPages: data.pagination.total_pages, + currentPage: data.pagination.page + }; + this.renderPrompts(); + this.updatePaginationControls(); + // Don't update stats from local data as it's only a subset of prompts + } + } + } catch (error) { + console.error("Recent prompts error:", error); + } finally { + document.getElementById("loadingState").classList.add("hidden"); + document.getElementById("resultsList").classList.remove("hidden"); + document.getElementById("paginationControls").classList.remove("hidden"); + } + } + + updateLocalStats() { + // Only update local stats if we haven't loaded proper statistics yet + // This prevents overwriting the correct total counts from the API + const currentTotal = document.getElementById("totalPrompts").textContent; + + // Only update if the current total is still the default "-" + if (currentTotal === "-") { + // Show local counts as fallback only when API hasn't loaded yet + document.getElementById("totalPrompts").textContent = this.prompts.length; + + const uniqueCategories = new Set(this.prompts.map(p => p.category).filter(Boolean)); + document.getElementById("totalCategories").textContent = uniqueCategories.size; + + const allTags = this.prompts.flatMap(p => Array.isArray(p.tags) ? p.tags : []); + const uniqueTags = new Set(allTags.filter(Boolean)); + document.getElementById("totalTags").textContent = uniqueTags.size; + + const ratings = this.prompts.map(p => p.rating).filter(r => r && r > 0); + if (ratings.length > 0) { + const avgRating = ratings.reduce((a, b) => a + b, 0) / ratings.length; + document.getElementById("avgRating").textContent = avgRating.toFixed(1); + } + } + } + + async search() { + const searchText = document.getElementById("searchText").value; + const category = document.getElementById("searchCategory").value; + const tags = document.getElementById("searchTags").value; + + document.getElementById("loadingState").classList.remove("hidden"); + document.getElementById("resultsList").classList.add("hidden"); + + try { + const params = new URLSearchParams(); + if (searchText) params.append("text", searchText); + if (category) params.append("category", category); + if (tags) params.append("tags", tags); + params.append("limit", "100"); + + const response = await fetch(`/prompt_manager/search?${params}`); + if (response.ok) { + const data = await response.json(); + if (data.success) { + this.prompts = data.results; + this.renderPrompts(); + document.getElementById("resultsTitle").textContent = "Search Results"; + } + } + } catch (error) { + this.showNotification("Search failed", "error"); + console.error("Search error:", error); + } finally { + document.getElementById("loadingState").classList.add("hidden"); + document.getElementById("resultsList").classList.remove("hidden"); + } + } + + sortPrompts(prompts) { + const sortBy = document.getElementById("sortBy")?.value || "created_desc"; + const [field, direction] = sortBy.split("_"); + + return [...prompts].sort((a, b) => { + let aVal, bVal; + + switch (field) { + case "rating": + aVal = a.rating || 0; + bVal = b.rating || 0; + break; + case "created": + aVal = new Date(a.created_at); + bVal = new Date(b.created_at); + break; + case "text": + aVal = a.text.toLowerCase(); + bVal = b.text.toLowerCase(); + break; + default: + return 0; + } + + if (direction === "asc") { + return aVal > bVal ? 1 : aVal < bVal ? -1 : 0; + } else { + return aVal < bVal ? 1 : aVal > bVal ? -1 : 0; + } + }); + } + + renderPrompts() { + const container = document.getElementById("resultsList"); + const count = document.getElementById("resultsCount"); + + count.textContent = `${this.prompts.length} prompts`; + + if (this.prompts.length === 0) { + container.innerHTML = ` +
+
+ 📭 +
+

No prompts found

+

Try adjusting your search criteria or create some prompts in ComfyUI

+
+ `; + return; + } + + // Sort prompts before rendering + const sortedPrompts = this.sortPrompts(this.prompts); + container.innerHTML = sortedPrompts.map((prompt) => this.renderPromptItem(prompt)).join(""); + this.selectedPrompts.clear(); + this.updateBulkActionButtons(); + + // Add hover behavior to all star ratings + sortedPrompts.forEach(prompt => { + this.addStarHoverBehavior(prompt.id); + }); + + // Delegated click handler for remove-tag buttons (avoids inline onclick XSS) + container.querySelectorAll('.remove-tag-btn').forEach(btn => { + btn.addEventListener('click', (e) => { + const promptId = parseInt(e.target.dataset.promptId); + const tag = e.target.dataset.tag; + window.admin.removeTag(promptId, tag); + }); + }); + + // Load film strips for each prompt (async, non-blocking) + this.loadAllFilmStrips(sortedPrompts); + } + + async loadAllFilmStrips(prompts) { + for (const prompt of prompts) { + const container = document.getElementById(`filmStrip-${prompt.id}`); + if (!container) continue; + + // Use pre-loaded images from API response if available + if (prompt.images && prompt.images.length > 0) { + container.innerHTML = this.createFilmStrip(prompt.images, prompt.id); // trusted internal HTML + } else if (prompt.image_count > 0) { + // Images exist but weren't pre-loaded — fetch individually + this.loadFilmStripForPrompt(prompt.id); + } + } + } + + async loadFilmStripForPrompt(promptId) { + const container = document.getElementById(`filmStrip-${promptId}`); + if (!container) return; + + try { + const images = await this.loadFilmStripImages(promptId); + container.innerHTML = this.createFilmStrip(images, promptId); // trusted internal HTML + } catch (error) { + console.error(`Error loading filmstrip for prompt ${promptId}:`, error); + container.innerHTML = `
Failed to load images
`; // static HTML + } + } + + renderPromptItem(prompt) { + const tags = Array.isArray(prompt.tags) ? prompt.tags : []; + const category = prompt.category || "No category"; + const rating = prompt.rating || 0; + const created = new Date(prompt.created_at).toLocaleDateString(); + + return ` +
+
+
+ + +
+
+
${this.escapeHtml(prompt.text)}
+ +
+ +
+
+ 📁 + ${this.escapeHtml(category)} +
+
+ 📅 + ${created} +
+
+
+ ${this.renderStars(rating, prompt.id)} +
+
+
+ +
+
+ ${tags.slice(0, 10).map(tag => ` + + ${this.escapeHtml(tag)} + + + `).join("")} + +
+ ${tags.length > 10 ? ` + + + ` : ''} +
+ + +
+
Loading images...
+
+
+ +
+ + + +
+
+
+
+ `; + } + + renderStars(rating, promptId) { + return Array.from({ length: 5 }, (_, i) => { + const starNum = i + 1; + const isActive = starNum <= rating; + const starIcon = isActive ? '⭐' : '☆'; // filled star vs outline star + const color = isActive ? '#fbbf24' : '#9ca3af'; // yellow-400 : gray-400 + return ``; + }).join(""); + } + + addStarHoverBehavior(promptId) { + const starButtons = document.querySelectorAll(`[data-prompt-id="${promptId}"].star-btn`); + const ratingElement = document.querySelector(`[data-id="${promptId}"][data-rating]`); + const currentRating = parseInt(ratingElement?.dataset.rating || 0); + + starButtons.forEach((starBtn, index) => { + const starNum = index + 1; + + starBtn.addEventListener('mouseenter', () => { + // Light up all stars up to this one on hover + starButtons.forEach((btn, i) => { + if (i < starNum) { + btn.textContent = '⭐'; + btn.style.color = '#fcd34d'; // bright yellow on hover + } else { + btn.textContent = '☆'; + btn.style.color = '#9ca3af'; // gray + } + }); + }); + + starBtn.addEventListener('mouseleave', () => { + // Reset to current rating + starButtons.forEach((btn, i) => { + if (i < currentRating) { + btn.textContent = '⭐'; + btn.style.color = '#fbbf24'; // yellow + } else { + btn.textContent = '☆'; + btn.style.color = '#9ca3af'; // gray + } + }); + }); + }); + } + + escapeHtml(text) { + const div = document.createElement("div"); + div.textContent = text; + return div.innerHTML; + } + + showNotification(message, type = "info") { + const notification = document.createElement("div"); + + // Check if ViewerJS is open and use higher z-index + const viewerContainer = document.querySelector('.viewer-container'); + const isViewerOpen = viewerContainer && viewerContainer.style.display !== 'none'; + const zIndex = isViewerOpen ? 'z-[50000]' : 'z-50'; + + notification.className = `fixed top-4 right-4 px-6 py-4 rounded-lg shadow-lg ${zIndex} transition-all duration-300 transform translate-x-full`; + + const colors = { + success: "bg-green-600 text-white", + error: "bg-red-600 text-white", + warning: "bg-yellow-600 text-white", + info: "bg-blue-600 text-white" + }; + + notification.className += ` ${colors[type] || colors.info}`; + notification.textContent = message; + + document.body.appendChild(notification); + + setTimeout(() => notification.classList.remove("translate-x-full"), 100); + setTimeout(() => { + notification.classList.add("translate-x-full"); + setTimeout(() => { + if (notification.parentNode) { + document.body.removeChild(notification); + } + }, 300); + }, 3000); + } + + showModal(modalId) { + document.getElementById(modalId).classList.remove("hidden"); + document.getElementById(modalId).classList.add("flex"); + document.body.style.overflow = "hidden"; + } + + hideModal(modalId) { + document.getElementById(modalId).classList.add("hidden"); + document.getElementById(modalId).classList.remove("flex"); + document.body.style.overflow = ""; + } + + showSettingsModal() { + document.getElementById("resultTimeout").value = this.settings.resultTimeout; + document.getElementById("webuiDisplayMode").value = this.settings.webuiDisplayMode; + document.getElementById("galleryRootPath").value = this.settings.galleryRootPath || ''; + this.updateMonitoringStatus(); + this.showModal("settingsModal"); + } + + async updateMonitoringStatus() { + const statusEl = document.getElementById("monitoringStatus"); + try { + const response = await fetch("/prompt_manager/settings"); + if (response.ok) { + const data = await response.json(); + const dirs = data.settings?.monitored_directories || []; + if (dirs.length > 0) { + statusEl.innerHTML = dirs.map(d => `
✓ ${d}
`).join(''); + } else { + statusEl.textContent = 'No directories being monitored (auto-detect on restart)'; + } + } + } catch (error) { + statusEl.textContent = 'Unable to fetch status'; + } + } + + async saveSettings() { + const timeout = parseInt(document.getElementById("resultTimeout").value); + const displayMode = document.getElementById("webuiDisplayMode").value; + const galleryPath = document.getElementById("galleryRootPath").value.trim(); + + this.settings.resultTimeout = timeout; + this.settings.webuiDisplayMode = displayMode; + this.settings.galleryRootPath = galleryPath; + + try { + const response = await fetch("/prompt_manager/settings", { + method: "POST", + headers: { "Content-Type": "application/json" }, + body: JSON.stringify({ + result_timeout: timeout, + webui_display_mode: displayMode, + gallery_root_path: galleryPath + }), + }); + + if (response.ok) { + const data = await response.json(); + if (data.restart_required) { + this.showNotification("Settings saved. Restart ComfyUI for gallery path changes to take effect.", "warning"); + } else { + this.showNotification("Settings saved successfully", "success"); + } + this.hideModal("settingsModal"); + } else { + throw new Error("Failed to save settings"); + } + } catch (error) { + this.showNotification("Failed to save settings", "error"); + } + } + + toggleSelectAll(checked) { + document.querySelectorAll(".prompt-checkbox").forEach((checkbox) => { + checkbox.checked = checked; + if (checked) { + this.selectedPrompts.add(parseInt(checkbox.dataset.id)); + } else { + this.selectedPrompts.clear(); + } + }); + this.updateBulkActionButtons(); + } + + updateBulkActionButtons() { + const hasSelected = this.selectedPrompts.size > 0; + document.getElementById("bulkDeleteBtn").disabled = !hasSelected; + document.getElementById("bulkTagBtn").disabled = !hasSelected; + document.getElementById("bulkCategoryBtn").disabled = !hasSelected; + } + + async setRating(promptId, rating) { + console.log(`Setting rating for prompt ${promptId} to ${rating}`); + + // Update UI immediately for better UX + const ratingElement = document.querySelector(`[data-id="${promptId}"][data-rating]`); + if (ratingElement) { + ratingElement.dataset.rating = rating; + ratingElement.innerHTML = this.renderStars(rating, promptId); + // Add hover behavior to new stars + this.addStarHoverBehavior(promptId); + } + + try { + const response = await fetch(`/prompt_manager/prompts/${promptId}/rating`, { + method: "PUT", + headers: { "Content-Type": "application/json" }, + body: JSON.stringify({ rating }), + }); + + if (response.ok) { + this.showNotification("Rating updated", "success"); + + // Update the prompt data in memory + const promptIndex = this.prompts.findIndex(p => p.id === promptId); + if (promptIndex !== -1) { + this.prompts[promptIndex].rating = rating; + } + + // Refresh stats + setTimeout(() => { + this.loadStatistics(); + }, 500); + } else { + // Revert UI change on API failure + const originalPrompt = this.prompts.find(p => p.id === promptId); + const originalRating = originalPrompt ? originalPrompt.rating : 0; + if (ratingElement) { + ratingElement.dataset.rating = originalRating; + ratingElement.innerHTML = this.renderStars(originalRating, promptId); + } + this.showNotification("Failed to update rating", "error"); + } + } catch (error) { + console.error("Rating update error:", error); + // Revert UI change on error + const originalPrompt = this.prompts.find(p => p.id === promptId); + const originalRating = originalPrompt ? originalPrompt.rating : 0; + if (ratingElement) { + ratingElement.dataset.rating = originalRating; + ratingElement.innerHTML = this.renderStars(originalRating, promptId); + } + this.showNotification("Failed to update rating", "error"); + } + } + + addTag(promptId) { + // Store the prompt ID for later use + this.currentTagPromptId = promptId; + // Clear the input and show the modal + document.getElementById("individualTagInput").value = ""; + this.showModal("individualTagModal"); + } + + async confirmIndividualTag() { + const tags = document.getElementById("individualTagInput").value.split(",") + .map((tag) => tag.trim()).filter((tag) => tag); + + if (tags.length === 0) return; + + try { + const response = await fetch("/prompt_manager/prompts/tags", { + method: "POST", + headers: { "Content-Type": "application/json" }, + body: JSON.stringify({ + prompt_id: this.currentTagPromptId, + tags: tags, + }), + }); + + if (response.ok) { + this.showNotification(`Tags added to prompt`, "success"); + this.hideModal("individualTagModal"); + // Refresh both stats and results + await this.loadStatistics(); + this.search(); + } + } catch (error) { + this.showNotification("Failed to add tags", "error"); + } + } + + async removeTag(promptId, tag) { + try { + const response = await fetch(`/prompt_manager/prompts/${promptId}/tags`, { + method: "DELETE", + headers: { "Content-Type": "application/json" }, + body: JSON.stringify({ tag }), + }); + + if (response.ok) { + this.showNotification("Tag removed", "success"); + // Refresh both stats and results + await this.loadStatistics(); + this.search(); + } + } catch (error) { + this.showNotification("Failed to remove tag", "error"); + } + } + + async editPrompt(promptId) { + const promptElement = document.querySelector(`[data-id="${promptId}"].prompt-text`); + if (!promptElement) return; + + const originalText = promptElement.textContent; + promptElement.contentEditable = true; + promptElement.focus(); + promptElement.classList.add("bg-gray-700", "border", "border-blue-500", "rounded-lg", "p-3"); + + const saveEdit = async () => { + const newText = promptElement.textContent.trim(); + if (newText !== originalText && newText) { + try { + const response = await fetch(`/prompt_manager/prompts/${promptId}`, { + method: "PUT", + headers: { "Content-Type": "application/json" }, + body: JSON.stringify({ text: newText }), + }); + + if (response.ok) { + this.showNotification("Prompt updated", "success"); + // Refresh the page to show updated content + setTimeout(() => { + this.loadStatistics(); + setTimeout(() => window.location.reload(), 500); + }, 1000); + } else { + throw new Error("Update failed"); + } + } catch (error) { + this.showNotification("Failed to update prompt", "error"); + promptElement.textContent = originalText; + } + } + promptElement.contentEditable = false; + promptElement.classList.remove("bg-gray-700", "border", "border-blue-500", "rounded-lg", "p-3"); + }; + + promptElement.addEventListener("blur", saveEdit, { once: true }); + promptElement.addEventListener("keydown", (e) => { + if (e.key === "Enter" && !e.shiftKey) { + e.preventDefault(); + saveEdit(); + } + if (e.key === "Escape") { + promptElement.textContent = originalText; + promptElement.contentEditable = false; + promptElement.classList.remove("bg-gray-700", "border", "border-blue-500", "rounded-lg", "p-3"); + } + }); + } + + async deletePrompt(promptId) { + if (!confirm("Are you sure you want to delete this prompt?")) return; + + try { + const response = await fetch(`/prompt_manager/delete/${promptId}`, { + method: "DELETE", + }); + + if (response.ok) { + this.showNotification("Prompt deleted", "success"); + // Refresh stats and reload page + setTimeout(() => { + this.loadStatistics(); + setTimeout(() => window.location.reload(), 500); + }, 1000); + } + } catch (error) { + this.showNotification("Failed to delete prompt", "error"); + } + } + + showBulkTagModal() { + this.showModal("bulkTagModal"); + } + + showBulkCategoryModal() { + this.showModal("bulkCategoryModal"); + } + + async confirmBulkTag() { + const tags = document.getElementById("bulkTagInput").value.split(",") + .map((tag) => tag.trim()).filter((tag) => tag); + + if (tags.length === 0) return; + + try { + const response = await fetch("/prompt_manager/bulk/tags", { + method: "POST", + headers: { "Content-Type": "application/json" }, + body: JSON.stringify({ + prompt_ids: Array.from(this.selectedPrompts), + tags, + }), + }); + + if (response.ok) { + this.showNotification(`Tags added to ${this.selectedPrompts.size} prompts`, "success"); + this.hideModal("bulkTagModal"); + // Refresh both stats and results + await this.loadStatistics(); + this.search(); + } + } catch (error) { + this.showNotification("Failed to add tags", "error"); + } + } + + async confirmBulkCategory() { + const category = document.getElementById("bulkCategoryInput").value.trim(); + if (!category) return; + + try { + const response = await fetch("/prompt_manager/bulk/category", { + method: "POST", + headers: { "Content-Type": "application/json" }, + body: JSON.stringify({ + prompt_ids: Array.from(this.selectedPrompts), + category, + }), + }); + + if (response.ok) { + this.showNotification(`Category set for ${this.selectedPrompts.size} prompts`, "success"); + this.hideModal("bulkCategoryModal"); + // Refresh both stats and results + await this.loadStatistics(); + this.search(); + } + } catch (error) { + this.showNotification("Failed to set category", "error"); + } + } + + async bulkDelete() { + if (!confirm(`Are you sure you want to delete ${this.selectedPrompts.size} prompts?`)) return; + + try { + const response = await fetch("/prompt_manager/bulk/delete", { + method: "POST", + headers: { "Content-Type": "application/json" }, + body: JSON.stringify({ + prompt_ids: Array.from(this.selectedPrompts), + }), + }); + + if (response.ok) { + this.showNotification(`${this.selectedPrompts.size} prompts deleted`, "success"); + // Refresh stats and reload page + setTimeout(() => { + this.loadStatistics(); + setTimeout(() => window.location.reload(), 500); + }, 1000); + } + } catch (error) { + this.showNotification("Failed to delete prompts", "error"); + } + } + + async exportPrompts() { + try { + const response = await fetch("/prompt_manager/export"); + if (response.ok) { + const blob = await response.blob(); + const url = window.URL.createObjectURL(blob); + const a = document.createElement("a"); + a.href = url; + a.download = `prompts-${new Date().toISOString().split("T")[0]}.json`; + document.body.appendChild(a); + a.click(); + window.URL.revokeObjectURL(url); + document.body.removeChild(a); + this.showNotification("Prompts exported", "success"); + } + } catch (error) { + this.showNotification("Failed to export prompts", "error"); + } + } + + async showStats() { + await this.loadStatistics(); + this.showNotification("📊 Stats refreshed successfully!", "success"); + } + + // Gallery functionality + async viewGallery(promptId) { + this.showModal("galleryModal"); + + // Show loading state + document.getElementById("galleryLoading").classList.remove("hidden"); + document.getElementById("galleryContent").classList.add("hidden"); + document.getElementById("galleryEmpty").classList.add("hidden"); + + try { + const response = await fetch(`/prompt_manager/prompts/${promptId}/images`); + if (response.ok) { + const data = await response.json(); + if (data.success) { + this.renderGallery(data.images, promptId); + } else { + throw new Error(data.error || 'Failed to load images'); + } + } else { + throw new Error(`HTTP ${response.status}: ${response.statusText}`); + } + } catch (error) { + console.error('Gallery load error:', error); + this.showNotification('Failed to load gallery', 'error'); + this.closeGallery(); + } finally { + document.getElementById("galleryLoading").classList.add("hidden"); + } + } + + renderGallery(images, promptId) { + const content = document.getElementById("galleryContent"); + const empty = document.getElementById("galleryEmpty"); + + // Update title + const prompt = this.prompts.find(p => p.id === promptId); + const promptText = prompt ? prompt.text.substring(0, 50) + (prompt.text.length > 50 ? '...' : '') : 'Unknown Prompt'; + document.getElementById("galleryTitle").textContent = `Gallery: ${promptText}`; + + if (!images || images.length === 0) { + content.classList.add("hidden"); + empty.classList.remove("hidden"); + return; + } + + content.classList.remove("hidden"); + empty.classList.add("hidden"); + + content.innerHTML = images.map((image, index) => ` +
+
+ Generated image ${index + 1} +
+ + + +
+
+
+
+ ${new Date(image.generation_time).toLocaleDateString()} ${new Date(image.generation_time).toLocaleTimeString()} +
+
+ ${image.width && image.height ? `${image.width}×${image.height}` : 'Unknown size'} + ${image.file_size ? ` • ${this.formatFileSize(image.file_size)}` : ''} +
+
+
+ `).join(''); + + // Store images for navigation + this.currentGalleryImages = images; + + // Initialize ViewerJS for the gallery after DOM is updated + setTimeout(() => { + if (this.galleryViewer) { + this.galleryViewer.destroy(); + } + + this.galleryViewer = new Viewer(content, { + toolbar: { + zoomIn: 1, + zoomOut: 1, + oneToOne: 1, + reset: 1, + prev: 1, + play: { + show: 1, + size: 'large', + }, + next: 1, + rotateLeft: 1, + rotateRight: 1, + flipHorizontal: 1, + flipVertical: 1, + }, + navbar: true, + title: true, + transition: true, + keyboard: true, + backdrop: 'static', + loading: true, + loop: true, + tooltip: true, + zoomRatio: 0.1, + minZoomRatio: 0.1, + maxZoomRatio: 4, + zoomOnTouch: true, + zoomOnWheel: true, + slideOnTouch: true, + toggleOnDblclick: false, + className: 'viewer-dark-theme', + shown: () => { + this.addMetadataSidebar(); + }, + viewed: (event) => { + this.loadMetadataForImage(event.detail.image); + } + }); + }, 100); + } + + + + + + closeGallery() { + this.hideModal("galleryModal"); + } + + + + + formatFileSize(bytes) { + if (!bytes) return 'Unknown'; + const units = ['B', 'KB', 'MB', 'GB']; + let size = bytes; + let unitIndex = 0; + + while (size >= 1024 && unitIndex < units.length - 1) { + size /= 1024; + unitIndex++; + } + + return `${size.toFixed(1)} ${units[unitIndex]}`; + } + + // Diagnostics functionality + showDiagnosticsModal() { + this.showModal("diagnosticsModal"); + } + + closeDiagnostics() { + this.hideModal("diagnosticsModal"); + } + + // Maintenance functionality + showMaintenanceModal() { + this.showModal("maintenanceModal"); + } + + closeMaintenance() { + this.hideModal("maintenanceModal"); + } + + selectAllMaintenance() { + document.querySelectorAll('.maintenance-option').forEach(checkbox => { + checkbox.checked = true; + }); + } + + clearAllMaintenance() { + document.querySelectorAll('.maintenance-option').forEach(checkbox => { + checkbox.checked = false; + }); + } + + async runMaintenance() { + const resultsContainer = document.getElementById("maintenanceResults"); + const runButton = document.getElementById("runMaintenanceBtn"); + + // Get selected operations + const selectedOps = Array.from(document.querySelectorAll('.maintenance-option:checked')) + .map(checkbox => checkbox.id); + + if (selectedOps.length === 0) { + this.showNotification('Please select at least one maintenance operation', 'warning'); + return; + } + + // Show loading state + runButton.disabled = true; + runButton.innerHTML = '🔄 Running...'; + resultsContainer.innerHTML = ` +
+
+

Running maintenance operations...

+

Operations: ${selectedOps.map(op => op.replace(/_/g, ' ')).join(', ')}

+
+ `; + + try { + const response = await fetch('/prompt_manager/maintenance', { + method: 'POST', + headers: { 'Content-Type': 'application/json' }, + body: JSON.stringify({ operations: selectedOps }) + }); + + if (response.ok) { + const data = await response.json(); + if (data.success) { + this.renderMaintenanceResults(data); + + // Show overall success notification + if (data.all_successful) { + this.showNotification('✅ All maintenance operations completed successfully!', 'success'); + } else { + this.showNotification('⚠️ Some maintenance operations had issues', 'warning'); + } + + // Refresh stats after maintenance + setTimeout(() => this.loadStatistics(), 1000); + } else { + throw new Error(data.error || 'Maintenance failed'); + } + } else { + throw new Error(`HTTP ${response.status}: ${response.statusText}`); + } + } catch (error) { + console.error('Maintenance error:', error); + resultsContainer.innerHTML = ` +
+

❌ Maintenance Failed

+

${error.message}

+
+ `; + this.showNotification('❌ Maintenance failed', 'error'); + } finally { + runButton.disabled = false; + runButton.innerHTML = '🔧 Run Maintenance'; + } + } + + renderMaintenanceResults(data) { + const resultsContainer = document.getElementById("maintenanceResults"); + + let html = '
'; + + // Summary + html += '
'; + html += '

📋 Maintenance Summary

'; + html += '
'; + html += ` +
+ Operations Completed + ${data.operations_completed} +
+ `; + html += ` +
+ Overall Success + + ${data.all_successful ? '✅ ALL PASSED' : '⚠️ SOME ISSUES'} + +
+ `; + html += '
'; + + // Detailed results + for (const [operation, result] of Object.entries(data.results)) { + const bgColor = result.success ? 'bg-green-900/20 border-green-600' : 'bg-red-900/20 border-red-600'; + const statusIcon = result.success ? '✅' : '❌'; + const statusText = result.success ? 'SUCCESS' : 'FAILED'; + + html += `
`; + html += `

`; + html += `${statusIcon}`; + html += `${operation.replace(/_/g, ' ')}`; + html += `${statusText}`; + html += `

`; + + if (result.message) { + html += `

${result.message}

`; + } + + // Show specific details + if (result.removed_count !== undefined) { + html += `
Items removed: ${result.removed_count}
`; + } + + if (result.duplicate_hashes !== undefined) { + html += `
Duplicate hash groups found: ${result.duplicate_hashes}
`; + } + + if (result.issues_found !== undefined) { + html += `
Issues found: ${result.issues_found}
`; + if (result.issues && result.issues.length > 0) { + html += ``; + } + } + + if (result.info) { + html += `
`; + html += `

Total prompts: ${result.info.total_prompts || 'N/A'}

`; + html += `

Database size: ${result.info.database_size_bytes ? this.formatFileSize(result.info.database_size_bytes) : 'N/A'}

`; + html += `
`; + } + + if (result.error) { + html += `
`; + html += `Error: ${result.error}`; + html += `
`; + } + + html += '
'; + } + + html += '
'; + resultsContainer.innerHTML = html; + } + + async runDiagnostics() { + const content = document.getElementById("diagnosticsContent"); + content.innerHTML = ` +
+
+

Running diagnostics...

+
+ `; + + try { + const response = await fetch('/prompt_manager/diagnostics'); + if (response.ok) { + const data = await response.json(); + if (data.success) { + this.renderDiagnostics(data.diagnostics); + } else { + throw new Error(data.error || 'Diagnostics failed'); + } + } else { + throw new Error(`HTTP ${response.status}: ${response.statusText}`); + } + } catch (error) { + console.error('Diagnostics error:', error); + content.innerHTML = ` +
+

❌ Diagnostics Failed

+

${error.message}

+
+ `; + } + } + + renderDiagnostics(diagnostics) { + const content = document.getElementById("diagnosticsContent"); + + let html = '
'; + + // Summary + html += '
'; + html += '

📋 Diagnostic Summary

'; + html += '
'; + + for (const [category, result] of Object.entries(diagnostics)) { + const status = result.status === 'ok' ? '✅ PASS' : (result.status === 'warning' ? '⚠️ WARNING' : '❌ FAIL'); + const statusColor = result.status === 'ok' ? 'text-green-400' : (result.status === 'warning' ? 'text-yellow-400' : 'text-red-400'); + + html += ` +
+ ${this.escapeHtml(category)} + ${status} +
+ `; + } + + html += '
'; + + // Detailed results + for (const [category, result] of Object.entries(diagnostics)) { + const bgColor = result.status === 'ok' ? 'bg-green-900/20 border-green-600' : + (result.status === 'warning' ? 'bg-yellow-900/20 border-yellow-600' : 'bg-red-900/20 border-red-600'); + + html += `
`; + html += `

${this.escapeHtml(category)}

`; + + if (result.message) { + html += `

${result.message}

`; + } + + // Show specific details based on category + if (category === 'database' && result.status === 'ok') { + html += `
`; + html += `

Prompts: ${result.prompt_count || 0}

`; + html += `

Images table: ${result.has_images_table ? 'Yes' : 'No'}

`; + html += `
`; + } + + if (category === 'images_table' && result.status === 'ok') { + html += `
`; + html += `

Images: ${result.image_count || 0}

`; + if (result.recent_images && result.recent_images.length > 0) { + html += `

Recent images:

`; + html += `
    `; + result.recent_images.slice(0, 3).forEach(img => { + html += `
  • ${img.filename} → Prompt ${img.prompt_id}
  • `; + }); + html += `
`; + } + html += `
`; + } + + if (category === 'comfyui_output' && result.output_dirs) { + html += `
`; + html += `

Output directories found:

`; + html += `
    `; + result.output_dirs.forEach(dir => { + html += `
  • ${dir}
  • `; + }); + html += `
`; + html += `
`; + } + + if (category === 'dependencies' && result.dependencies) { + html += `
`; + for (const [dep, available] of Object.entries(result.dependencies)) { + const status = available ? '✅' : '❌'; + html += `

${status} ${dep}

`; + } + html += `
`; + } + + html += '
'; + } + + html += '
'; + content.innerHTML = html; + } + + async testImageLink() { + if (this.prompts.length === 0) { + this.showNotification('No prompts available for testing', 'warning'); + return; + } + + // Use the first prompt for testing + const testPrompt = this.prompts[0]; + + try { + const response = await fetch('/prompt_manager/diagnostics/test-link', { + method: 'POST', + headers: { 'Content-Type': 'application/json' }, + body: JSON.stringify({ + prompt_id: testPrompt.id, + image_path: '/test/fake/image.png' + }) + }); + + if (response.ok) { + const data = await response.json(); + if (data.success) { + this.showNotification('✅ Test image link created successfully!', 'success'); + // Refresh the gallery to show the test image + setTimeout(() => this.search(), 1000); + } else { + throw new Error(data.result?.message || 'Test failed'); + } + } else { + throw new Error(`HTTP ${response.status}: ${response.statusText}`); + } + } catch (error) { + console.error('Test link error:', error); + this.showNotification(`❌ Test failed: ${error.message}`, 'error'); + } + } + + // Database backup and restore functionality + async backupDatabase() { + try { + const response = await fetch('/prompt_manager/backup'); + + if (response.ok) { + // Get the filename from the Content-Disposition header + const contentDisposition = response.headers.get('Content-Disposition'); + let filename = 'prompts_backup.db'; + + if (contentDisposition) { + const filenameMatch = contentDisposition.match(/filename="(.+)"/); + if (filenameMatch) { + filename = filenameMatch[1]; + } + } + + // Create blob and download + const blob = await response.blob(); + const url = window.URL.createObjectURL(blob); + const a = document.createElement('a'); + a.href = url; + a.download = filename; + document.body.appendChild(a); + a.click(); + window.URL.revokeObjectURL(url); + document.body.removeChild(a); + + this.showNotification('💾 Database backup downloaded successfully!', 'success'); + } else { + const errorData = await response.json(); + throw new Error(errorData.error || 'Backup failed'); + } + } catch (error) { + console.error('Backup error:', error); + this.showNotification(`❌ Backup failed: ${error.message}`, 'error'); + } + } + + showRestoreModal() { + // Clear previous file selection + document.getElementById('restoreFileInput').value = ''; + this.showModal('restoreModal'); + } + + async confirmRestore() { + const fileInput = document.getElementById('restoreFileInput'); + const file = fileInput.files[0]; + + if (!file) { + this.showNotification('Please select a database file to restore', 'warning'); + return; + } + + if (!file.name.endsWith('.db')) { + this.showNotification('Please select a valid .db file', 'warning'); + return; + } + + // Show confirmation + const confirmed = confirm( + `Are you sure you want to restore from "${file.name}"?\n\n` + + 'This will replace your current database. A backup will be created automatically.\n\n' + + 'This action cannot be undone.' + ); + + if (!confirmed) { + return; + } + + try { + // Disable the restore button and show loading + const confirmBtn = document.getElementById('confirmRestore'); + confirmBtn.disabled = true; + confirmBtn.innerHTML = '🔄 Restoring...'; + + // Create FormData for file upload + const formData = new FormData(); + formData.append('database_file', file); + + const response = await fetch('/prompt_manager/restore', { + method: 'POST', + body: formData + }); + + const data = await response.json(); + + if (response.ok && data.success) { + this.showNotification( + `✅ Database restored successfully! Found ${data.prompt_count} prompts.`, + 'success' + ); + this.hideModal('restoreModal'); + + // Refresh the interface to show new data + setTimeout(() => { + window.location.reload(); + }, 2000); + } else { + throw new Error(data.error || 'Restore failed'); + } + } catch (error) { + console.error('Restore error:', error); + this.showNotification(`❌ Restore failed: ${error.message}`, 'error'); + } finally { + // Re-enable the restore button + const confirmBtn = document.getElementById('confirmRestore'); + confirmBtn.disabled = false; + confirmBtn.innerHTML = 'Restore Database'; + } + } + + // Copy prompt to clipboard functionality + async copyPromptToClipboard(promptId) { + // Find the prompt text from the prompts array + const prompt = this.prompts.find(p => p.id === promptId); + if (!prompt) { + this.showNotification('❌ Prompt not found', 'error'); + return; + } + + const promptText = prompt.text; + + try { + // Use the modern Clipboard API if available + if (navigator.clipboard && window.isSecureContext) { + await navigator.clipboard.writeText(promptText); + this.showNotification('📋 Prompt copied to clipboard!', 'success'); + } else { + // Fallback for older browsers or non-secure contexts + const textArea = document.createElement('textarea'); + textArea.value = promptText; + textArea.style.position = 'fixed'; + textArea.style.left = '-999999px'; + textArea.style.top = '-999999px'; + document.body.appendChild(textArea); + textArea.focus(); + textArea.select(); + + if (document.execCommand('copy')) { + this.showNotification('📋 Prompt copied to clipboard!', 'success'); + } else { + throw new Error('Copy command failed'); + } + + document.body.removeChild(textArea); + } + } catch (error) { + console.error('Copy to clipboard failed:', error); + this.showNotification('❌ Failed to copy prompt to clipboard', 'error'); + + // Show a fallback modal with the text for manual copying + this.showCopyFallbackModal(promptText); + } + } + + showCopyFallbackModal(text) { + // Create a temporary modal for manual copy + const modal = document.createElement('div'); + modal.className = 'fixed inset-0 bg-black bg-opacity-50 flex items-center justify-center z-50'; + modal.innerHTML = ` +
+

📋 Copy Prompt Text

+

Please manually copy the text below:

+ +
+ +
+
+ `; + document.body.appendChild(modal); + document.body.style.overflow = "hidden"; + + // Auto-select the text in the textarea + const textarea = modal.querySelector('textarea'); + textarea.focus(); + textarea.select(); + + // Close modal when clicking outside + modal.addEventListener('click', (e) => { + if (e.target === modal) { + modal.remove(); + document.body.style.overflow = ""; + } + }); + + // Also handle the close button + modal.querySelector('button').addEventListener('click', () => { + document.body.style.overflow = ""; + }); + } + + // Scan functionality + showScanModal() { + // Reset progress display + document.getElementById("scanProgress").classList.add("hidden"); + document.getElementById("startScan").disabled = false; + document.getElementById("startScan").textContent = "Start Scan"; + this.showModal("scanModal"); + } + + async quickBackup() { + try { + const response = await fetch("/prompt_manager/backup"); + if (response.ok) { + const blob = await response.blob(); + const url = window.URL.createObjectURL(blob); + const a = document.createElement("a"); + a.href = url; + a.download = `prompts-backup-${new Date().toISOString().split("T")[0]}.db`; + document.body.appendChild(a); + a.click(); + window.URL.revokeObjectURL(url); + document.body.removeChild(a); + this.showNotification("Database backup created successfully!", "success"); + } else { + throw new Error("Backup failed"); + } + } catch (error) { + this.showNotification("Failed to create backup", "error"); + } + } + + async startScan() { + // Show progress section + document.getElementById("scanProgress").classList.remove("hidden"); + document.getElementById("startScan").disabled = true; + document.getElementById("startScan").textContent = "Scanning..."; + + // Reset progress + this.updateScanProgress(0, "Initializing scan...", 0, 0); + + try { + const response = await fetch("/prompt_manager/scan", { + method: "POST", + headers: { "Content-Type": "application/json" }, + body: JSON.stringify({}) + }); + + if (!response.ok) { + throw new Error(`HTTP ${response.status}: ${response.statusText}`); + } + + // Handle streaming response + const reader = response.body.getReader(); + const decoder = new TextDecoder(); + + while (true) { + const { done, value } = await reader.read(); + if (done) break; + + const chunk = decoder.decode(value); + const lines = chunk.split('\n'); + + for (const line of lines) { + if (line.trim().startsWith('data: ')) { + try { + const data = JSON.parse(line.substring(6)); + if (data.type === 'progress') { + this.updateScanProgress( + data.progress, + data.status, + data.processed, + data.found + ); + } else if (data.type === 'complete') { + this.completeScan(data.processed, data.found, data.added, data.linked); + return; + } else if (data.type === 'error') { + throw new Error(data.message); + } + } catch (e) { + console.log('Non-JSON line:', line); + } + } + } + } + } catch (error) { + this.showNotification(`Scan failed: ${error.message}`, "error"); + document.getElementById("startScan").disabled = false; + document.getElementById("startScan").textContent = "Start Scan"; + } + } + + updateScanProgress(progress, status, processed, found) { + document.getElementById("scanProgressBar").style.width = `${progress}%`; + document.getElementById("scanStatusText").textContent = status; + document.getElementById("scanCount").textContent = `${processed} files processed`; + document.getElementById("scanFound").textContent = `${found} prompts found`; + } + + completeScan(processed, found, added, linked = 0) { + this.updateScanProgress(100, "Scan completed!", processed, found); + document.getElementById("startScan").disabled = false; + document.getElementById("startScan").textContent = "Start New Scan"; + + // Show detailed notification with all counts + const linkedText = linked > 0 ? `, linked ${linked} images to existing prompts` : ''; + this.showNotification( + `Scan completed! Processed ${processed} files, found ${found} prompts, added ${added} new prompts to database${linkedText}.`, + "success" + ); + + // Auto-close modal after a short delay to let user see the completion message + setTimeout(() => { + this.hideModal("scanModal"); + + // Refresh the statistics immediately + this.loadStatistics(); + + // Also reload the page after a short delay to ensure everything is fully updated + setTimeout(() => { + window.location.reload(); + }, 1000); + }, 2000); + } + + // Logs functionality + showLogsModal() { + this.showModal("logsModal"); + this.loadLogStats(); + this.loadLogFiles(); + this.loadLogConfig(); + } + + closeLogs() { + this.hideModal("logsModal"); + } + + async loadLogStats() { + try { + const response = await fetch('/prompt_manager/logs/stats'); + const data = await response.json(); + + if (data.success) { + const stats = data.stats; + document.getElementById("logStatsTotal").textContent = stats.buffer_count || 0; + document.getElementById("logStatsErrors").textContent = stats.level_counts?.ERROR || 0; + document.getElementById("logStatsWarnings").textContent = stats.level_counts?.WARNING || 0; + document.getElementById("logStatsLevel").textContent = stats.current_level || "INFO"; + document.getElementById("logStatsSize").textContent = this.formatBytes(stats.total_log_size || 0); + } + } catch (error) { + console.error("Failed to load log stats:", error); + } + } + + async loadLogConfig() { + try { + const response = await fetch('/prompt_manager/logs/config'); + const data = await response.json(); + + if (data.success) { + const config = data.config; + document.getElementById("setLogLevel").value = config.level || "INFO"; + document.getElementById("consoleLogging").checked = config.console_logging !== false; + } + } catch (error) { + console.error("Failed to load log config:", error); + } + } + + async loadLogFiles() { + try { + const response = await fetch('/prompt_manager/logs/files'); + const data = await response.json(); + + if (data.success) { + const container = document.getElementById("logFilesList"); + + if (data.files.length === 0) { + container.innerHTML = '
No log files found
'; + return; + } + + container.innerHTML = data.files.map(file => ` +
+
+ ${file.filename} + ${file.is_main ? 'Active' : ''} +
+
+ Size: ${this.formatBytes(file.size)} | Modified: ${new Date(file.modified).toLocaleString()} +
+ +
+ `).join(''); + } + } catch (error) { + console.error("Failed to load log files:", error); + } + } + + async refreshLogs() { + const level = document.getElementById("logLevel").value; + const limit = parseInt(document.getElementById("logLimit").value); + + try { + const url = new URL('/prompt_manager/logs', window.location.origin); + if (level) url.searchParams.set('level', level); + url.searchParams.set('limit', limit.toString()); + + const response = await fetch(url); + const data = await response.json(); + + if (data.success) { + this.displayLogs(data.logs); + this.showNotification(`Loaded ${data.logs.length} log entries`, "success"); + } else { + this.showNotification(`Failed to load logs: ${data.error}`, "error"); + } + } catch (error) { + console.error("Failed to refresh logs:", error); + this.showNotification("Failed to refresh logs", "error"); + } + } + + displayLogs(logs) { + const container = document.getElementById("logsContainer"); + + if (logs.length === 0) { + container.innerHTML = '
No logs found
'; + return; + } + + container.innerHTML = logs.map(log => { + const levelColors = { + DEBUG: 'text-gray-400', + INFO: 'text-blue-400', + WARNING: 'text-yellow-400', + ERROR: 'text-red-400', + CRITICAL: 'text-red-600' + }; + + const levelColor = levelColors[log.level] || 'text-gray-400'; + const timestamp = new Date(log.timestamp).toLocaleString(); + + return ` +
+
+ ${timestamp.split(' ')[1]} + ${log.level} + ${log.logger} + ${log.message} +
+
+ ${log.filename}:${log.lineno} +
+
+ `; + }).join(''); + + // Auto-scroll to bottom + container.scrollTop = container.scrollHeight; + } + + async downloadLogs() { + const level = document.getElementById("logLevel").value; + const limit = parseInt(document.getElementById("logLimit").value); + + try { + const url = new URL('/prompt_manager/logs', window.location.origin); + if (level) url.searchParams.set('level', level); + url.searchParams.set('limit', limit.toString()); + + const response = await fetch(url); + const data = await response.json(); + + if (data.success) { + const logText = data.logs.map(log => + `${log.timestamp} [${log.level}] ${log.logger}: ${log.message} (${log.filename}:${log.lineno})` + ).join('\n'); + + const blob = new Blob([logText], { type: 'text/plain' }); + const url2 = URL.createObjectURL(blob); + const a = document.createElement('a'); + a.href = url2; + a.download = `prompt_manager_logs_${new Date().toISOString().slice(0, 19).replace(/:/g, '-')}.txt`; + document.body.appendChild(a); + a.click(); + document.body.removeChild(a); + URL.revokeObjectURL(url2); + + this.showNotification("Logs downloaded successfully", "success"); + } else { + this.showNotification(`Failed to download logs: ${data.error}`, "error"); + } + } catch (error) { + console.error("Failed to download logs:", error); + this.showNotification("Failed to download logs", "error"); + } + } + + async downloadLogFile(filename) { + try { + const response = await fetch(`/prompt_manager/logs/download/${encodeURIComponent(filename)}`); + + if (response.ok) { + const blob = await response.blob(); + const url = URL.createObjectURL(blob); + const a = document.createElement('a'); + a.href = url; + a.download = filename; + document.body.appendChild(a); + a.click(); + document.body.removeChild(a); + URL.revokeObjectURL(url); + + this.showNotification(`Downloaded ${filename}`, "success"); + } else { + this.showNotification(`Failed to download ${filename}`, "error"); + } + } catch (error) { + console.error("Failed to download log file:", error); + this.showNotification("Failed to download log file", "error"); + } + } + + async clearLogs() { + if (!confirm("Are you sure you want to clear all log files? This cannot be undone.")) { + return; + } + + try { + const response = await fetch('/prompt_manager/logs/truncate', { + method: 'POST' + }); + const data = await response.json(); + + if (data.success) { + this.showNotification(data.message, "success"); + this.refreshLogs(); + this.loadLogStats(); + this.loadLogFiles(); + } else { + this.showNotification(`Failed to clear logs: ${data.error}`, "error"); + } + } catch (error) { + console.error("Failed to clear logs:", error); + this.showNotification("Failed to clear logs", "error"); + } + } + + async updateLogConfig() { + const level = document.getElementById("setLogLevel").value; + const consoleLogging = document.getElementById("consoleLogging").checked; + + try { + const response = await fetch('/prompt_manager/logs/config', { + method: 'POST', + headers: { + 'Content-Type': 'application/json' + }, + body: JSON.stringify({ + level: level, + console_logging: consoleLogging + }) + }); + + const data = await response.json(); + + if (data.success) { + this.showNotification("Log configuration updated", "success"); + this.loadLogStats(); + } else { + this.showNotification(`Failed to update config: ${data.error}`, "error"); + } + } catch (error) { + console.error("Failed to update log config:", error); + this.showNotification("Failed to update log configuration", "error"); + } + } + + formatBytes(bytes) { + if (bytes === 0) return '0 B'; + const k = 1024; + const sizes = ['B', 'KB', 'MB', 'GB']; + const i = Math.floor(Math.log(bytes) / Math.log(k)); + return parseFloat((bytes / Math.pow(k, i)).toFixed(1)) + ' ' + sizes[i]; + } + + // Pagination methods + updatePaginationControls() { + const { currentPage, totalPages, total, limit } = this.pagination; + + // Update pagination info + document.getElementById("currentPage").textContent = currentPage; + document.getElementById("totalPages").textContent = totalPages; + document.getElementById("totalResults").textContent = total; + + // Calculate showing range + const start = ((currentPage - 1) * limit) + 1; + const end = Math.min(currentPage * limit, total); + document.getElementById("showingStart").textContent = start; + document.getElementById("showingEnd").textContent = end; + + // Update results count + document.getElementById("resultsCount").textContent = `${this.prompts.length} of ${total}`; + + // Update button states + document.getElementById("firstPageBtn").disabled = currentPage === 1; + document.getElementById("prevPageBtn").disabled = currentPage === 1; + document.getElementById("nextPageBtn").disabled = currentPage === totalPages; + document.getElementById("lastPageBtn").disabled = currentPage === totalPages; + } + + async goToPage(page) { + if (page < 1 || page > this.pagination.totalPages || page === this.pagination.currentPage) { + return; + } + await this.loadRecentPrompts(page); + // Scroll to top of results after page loads + window.scrollTo({ top: 0, behavior: 'smooth' }); + } + + async changeLimit(newLimit) { + this.pagination.limit = newLimit; + this.pagination.currentPage = 1; // Reset to first page + await this.loadRecentPrompts(1); + } + + // Metadata functionality + openMetadataViewer() { + window.open('/prompt_manager/gallery.html', '_blank'); + } + + // Gallery functionality + openGallery() { + window.open('/prompt_manager/gallery', '_blank'); + } + + + + async parsePNGMetadata(arrayBuffer) { + const dataView = new DataView(arrayBuffer); + let offset = 8; // Skip PNG signature + const metadata = {}; + let chunkCount = 0; + + console.log('Starting PNG metadata parsing...'); + + while (offset < arrayBuffer.byteLength - 8) { + const length = dataView.getUint32(offset); + const type = new TextDecoder().decode(arrayBuffer.slice(offset + 4, offset + 8)); + + chunkCount++; + console.log(`Chunk ${chunkCount}: type=${type}, length=${length}`); + + if (type === 'tEXt' || type === 'iTXt' || type === 'zTXt') { + const chunkData = arrayBuffer.slice(offset + 8, offset + 8 + length); + let text; + + if (type === 'tEXt') { + text = new TextDecoder().decode(chunkData); + } else if (type === 'iTXt') { + // iTXt format: keyword\0compression\0language\0translated_keyword\0text + const textData = new TextDecoder().decode(chunkData); + const parts = textData.split('\0'); + console.log(`iTXt parts count: ${parts.length}, first part: ${parts[0]}`); + if (parts.length >= 5) { + metadata[parts[0]] = parts[4]; + } + text = textData; + } else if (type === 'zTXt') { + // zTXt is compressed - basic parsing (might need proper decompression) + text = new TextDecoder().decode(chunkData); + } + + // Parse the text chunk for key-value pairs + const nullIndex = text.indexOf('\0'); + if (nullIndex !== -1) { + const key = text.substring(0, nullIndex); + const value = text.substring(nullIndex + 1); + console.log(`Found metadata: ${key} = ${value.substring(0, 100)}...`); + metadata[key] = value; + } + } + + offset += 8 + length + 4; // Move to next chunk (8 = length + type, 4 = CRC) + } + + console.log(`Parsed ${chunkCount} chunks, found ${Object.keys(metadata).length} metadata items`); + return metadata; + } + + extractComfyUIData(metadata) { + // Look for ComfyUI workflow data in various possible fields + let workflowData = null; + let promptData = null; + + // Common ComfyUI metadata field names + const workflowFields = ['workflow', 'Workflow', 'comfy', 'ComfyUI']; + const promptFields = ['prompt', 'Prompt', 'parameters', 'Parameters']; + + for (const field of workflowFields) { + if (metadata[field]) { + try { + // Clean NaN values from JSON string before parsing + let cleanedJson = metadata[field]; + cleanedJson = cleanedJson.replace(/:\s*NaN\b/g, ': null'); + cleanedJson = cleanedJson.replace(/\bNaN\b/g, 'null'); + + workflowData = JSON.parse(cleanedJson); + console.log(`Successfully parsed workflow field: ${field}`); + break; + } catch (e) { + console.log('Failed to parse workflow field:', field, e.message); + } + } + } + + for (const field of promptFields) { + if (metadata[field]) { + try { + // Clean NaN values from JSON string before parsing + let cleanedJson = metadata[field]; + cleanedJson = cleanedJson.replace(/:\s*NaN\b/g, ': null'); + cleanedJson = cleanedJson.replace(/\bNaN\b/g, 'null'); + + promptData = JSON.parse(cleanedJson); + console.log(`Successfully parsed prompt field: ${field}`); + break; + } catch (e) { + console.log('Failed to parse prompt field:', field, e.message); + console.log('Raw data:', metadata[field].substring(0, 200) + '...'); + } + } + } + + return { workflow: workflowData, prompt: promptData }; + } + + updateMetadataPanel(comfyData, imageSrc) { + const metadataContent = document.getElementById('metadataContent'); + if (!metadataContent) return; + + // Get the actual file path from current image data + let filePath = imageSrc; // fallback to URL + if (this.currentGalleryImages && this.currentImageIndex !== null && this.currentGalleryImages[this.currentImageIndex]) { + const currentImage = this.currentGalleryImages[this.currentImageIndex]; + filePath = currentImage.image_path || currentImage.filename || imageSrc; + } + + // Extract information from ComfyUI workflow + let checkpoint = 'Unknown'; + let positivePrompt = 'No prompt found'; + let negativePrompt = 'No negative prompt found'; + let steps = 'Unknown'; + let cfgScale = 'Unknown'; + let sampler = 'Unknown'; + let seed = 'Unknown'; + + // Parse prompt data first (more reliable for actual generation parameters) + if (comfyData.prompt) { + console.log('Parsing prompt data...', Object.keys(comfyData.prompt).length, 'nodes'); + const promptNodes = comfyData.prompt; + + for (const nodeId in promptNodes) { + const node = promptNodes[nodeId]; + console.log(`Node ${nodeId}:`, node.class_type, node); + + // Checkpoint + if (node.class_type === 'CheckpointLoaderSimple' && node.inputs) { + checkpoint = node.inputs.ckpt_name || checkpoint; + console.log('Found checkpoint:', checkpoint); + } + + // Prompts - need to identify which is positive vs negative + if (node.class_type === 'PromptManager' && node.inputs && node.inputs.text) { + // PromptManager typically contains the positive prompt + const promptValue = typeof node.inputs.text === 'string' ? node.inputs.text : (Array.isArray(node.inputs.text) ? node.inputs.text[0] : String(node.inputs.text)); + positivePrompt = promptValue; + console.log('Found PromptManager with text:', promptValue.substring(0, 100)); + } + + if (node.class_type === 'CLIPTextEncode' && node.inputs && node.inputs.text) { + // Check if this looks like a negative prompt + const textValue = typeof node.inputs.text === 'string' ? node.inputs.text : (Array.isArray(node.inputs.text) ? node.inputs.text[0] : String(node.inputs.text)); + console.log('Found CLIPTextEncode:', textValue.substring(0, 50)); + const text = textValue.toLowerCase(); + if (text.includes('bad anatomy') || text.includes('unfinished') || + text.includes('censored') || text.includes('weird anatomy') || + text.includes('negative') || text.includes('embedding:')) { + negativePrompt = textValue; + console.log('Found negative prompt:', textValue.substring(0, 100)); + } else if (positivePrompt === 'No prompt found') { + // If we haven't found a positive prompt yet, this might be it + positivePrompt = textValue; + console.log('Found potential positive prompt:', textValue.substring(0, 100)); + } + } + + // Sampling parameters + if (node.class_type === 'KSampler' && node.inputs) { + seed = node.inputs.seed || seed; + steps = node.inputs.steps || steps; + cfgScale = node.inputs.cfg || cfgScale; + sampler = node.inputs.sampler_name || sampler; + console.log('Found sampler params:', {seed, steps, cfgScale, sampler}); + } + } + } else { + console.log('No prompt data found, will try workflow fallback'); + } + + // Store the current metadata for copying + this.currentMetadata = { + positivePrompt, + negativePrompt, + checkpoint, + steps, + cfgScale, + sampler, + seed, + workflow: comfyData.workflow, + prompt: comfyData.prompt + }; + + // Update the HTML + metadataContent.innerHTML = ` + +
+

File Path

+
+ ${filePath} +
+
+ + +
+

Resources used

+
+
+
${checkpoint}
+
ComfyUI Generated
+
+ CHECKPOINT +
+
+ + +
+
+

Prompt

+ COMFYUI + +
+
+ ${positivePrompt.substring(0, 200)}${positivePrompt.length > 200 ? '...' : ''} +
+ ${positivePrompt.length > 200 ? '' : ''} +
+ + +
+
+

Negative prompt

+ +
+
+ ${negativePrompt.substring(0, 200)}${negativePrompt.length > 200 ? '...' : ''} +
+ ${negativePrompt.length > 200 ? '' : ''} +
+ + +
+

Other metadata

+
+ CFG SCALE: ${cfgScale} + STEPS: ${steps} + SAMPLER: ${sampler} +
+
+ SEED: ${seed} +
+
+ + +
+

ComfyUI Workflow

+
+ + +
+
+ `; + } + + showMetadataError() { + const metadataContent = document.getElementById('metadataContent'); + if (!metadataContent) return; + + metadataContent.innerHTML = ` +
+ + + +

Error Loading Metadata

+

Could not extract ComfyUI metadata from this image

+
+ `; + } + + async copyPrompt(type) { + if (!this.currentMetadata) { + this.showNotification('❌ No metadata available for copying', 'error'); + return; + } + + const text = type === 'positive' ? this.currentMetadata.positivePrompt : this.currentMetadata.negativePrompt; + + if (!text || text === 'No prompt found' || text === 'No negative prompt found') { + this.showNotification(`❌ No ${type} prompt available`, 'error'); + return; + } + + await this.copyToClipboard(text); + } + + async tryFallbackMetadata(imageSrc) { + console.log('Trying fallback metadata extraction for:', imageSrc); + + try { + // Try to get prompt from current gallery image data + let fallbackPrompt = 'No prompt found'; + + if (this.currentGalleryImages && this.currentImageIndex !== null && this.currentGalleryImages[this.currentImageIndex]) { + const currentImage = this.currentGalleryImages[this.currentImageIndex]; + console.log('Current image data:', currentImage); + + // If we have prompt_id, try to get the prompt from our local prompts array + if (currentImage.prompt_id && this.prompts) { + const prompt = this.prompts.find(p => p.id === currentImage.prompt_id); + if (prompt) { + fallbackPrompt = prompt.text; + console.log('Found prompt from database:', fallbackPrompt.substring(0, 100)); + } + } + } + + // Set fallback metadata + this.currentMetadata = { + positivePrompt: fallbackPrompt, + negativePrompt: 'No negative prompt found', + checkpoint: 'Unknown', + steps: 'Unknown', + cfgScale: 'Unknown', + sampler: 'Unknown', + seed: 'Unknown', + workflow: null, + prompt: null + }; + + // Update metadata panel with fallback data + this.updateMetadataPanel({}, imageSrc); + + } catch (error) { + console.error('Fallback metadata extraction failed:', error); + this.showMetadataError(); + + // Set empty metadata as last resort + this.currentMetadata = { + positivePrompt: 'No prompt found', + negativePrompt: 'No negative prompt found', + checkpoint: 'Unknown', + steps: 'Unknown', + cfgScale: 'Unknown', + sampler: 'Unknown', + seed: 'Unknown', + workflow: null, + prompt: null + }; + } + } + + async copyAllMetadata() { + if (!this.currentMetadata) return; + + const allData = `Checkpoint: ${this.currentMetadata.checkpoint || 'Unknown'} +Positive Prompt: ${this.currentMetadata.positivePrompt} +Negative Prompt: ${this.currentMetadata.negativePrompt} +Steps: ${this.currentMetadata.steps || 'Unknown'} +CFG Scale: ${this.currentMetadata.cfgScale || 'Unknown'} +Sampler: ${this.currentMetadata.sampler || 'Unknown'} +Seed: ${this.currentMetadata.seed || 'Unknown'}`; + + await this.copyToClipboard(allData); + } + + showFullPrompt(type) { + if (!this.currentMetadata) return; + + const prompt = type === 'positive' ? this.currentMetadata.positivePrompt : this.currentMetadata.negativePrompt; + const newWindow = window.open('', '_blank'); + newWindow.document.write(` + + ${type.charAt(0).toUpperCase() + type.slice(1)} Prompt + +

${type.charAt(0).toUpperCase() + type.slice(1)} Prompt

+
${prompt}
+ + + + `); + } + + showWorkflowData() { + if (!this.currentMetadata || !this.currentMetadata.workflow) return; + + const newWindow = window.open('', '_blank'); + newWindow.document.write(` + + ComfyUI Workflow Data + +

ComfyUI Workflow JSON

+
${JSON.stringify(this.currentMetadata.workflow, null, 2)}
+ + + `); + } + + downloadWorkflowJSON() { + if (!this.currentMetadata || !this.currentMetadata.workflow) return; + + const dataStr = JSON.stringify(this.currentMetadata.workflow, null, 2); + const dataBlob = new Blob([dataStr], {type: 'application/json'}); + const url = URL.createObjectURL(dataBlob); + const link = document.createElement('a'); + link.href = url; + link.download = 'comfyui_workflow.json'; + document.body.appendChild(link); + link.click(); + document.body.removeChild(link); + URL.revokeObjectURL(url); + } + + async copyToClipboard(text) { + try { + await navigator.clipboard.writeText(text); + this.showNotification('Copied to clipboard!', 'success'); + } catch (err) { + console.error('Copy failed:', err); + this.showNotification('Copy failed', 'error'); + } + } + + async parsePNGMetadata(arrayBuffer) { + const dataView = new DataView(arrayBuffer); + let offset = 8; // Skip PNG signature + const metadata = {}; + + while (offset < arrayBuffer.byteLength - 8) { + const length = dataView.getUint32(offset); + const type = new TextDecoder().decode(arrayBuffer.slice(offset + 4, offset + 8)); + + if (type === 'tEXt' || type === 'iTXt' || type === 'zTXt') { + const chunkData = arrayBuffer.slice(offset + 8, offset + 8 + length); + let text; + + if (type === 'tEXt') { + text = new TextDecoder().decode(chunkData); + } else if (type === 'iTXt') { + // iTXt format: keyword\0compression\0language\0translated_keyword\0text + const textData = new TextDecoder().decode(chunkData); + const parts = textData.split('\0'); + if (parts.length >= 5) { + metadata[parts[0]] = parts[4]; + } + text = textData; + } else if (type === 'zTXt') { + // zTXt is compressed - basic parsing (might need proper decompression) + text = new TextDecoder().decode(chunkData); + } + + // Parse the text chunk for key-value pairs + const nullIndex = text.indexOf('\0'); + if (nullIndex !== -1) { + const key = text.substring(0, nullIndex); + const value = text.substring(nullIndex + 1); + metadata[key] = value; + } + } + + offset += 8 + length + 4; // Move to next chunk (8 = length + type, 4 = CRC) + } + + return metadata; + } + + extractComfyUIData(metadata) { + // Look for ComfyUI workflow data in various possible fields + let workflowData = null; + let promptData = null; + + // Common ComfyUI metadata field names + const workflowFields = ['workflow', 'Workflow', 'comfy', 'ComfyUI']; + const promptFields = ['prompt', 'Prompt', 'parameters', 'Parameters']; + + for (const field of workflowFields) { + if (metadata[field]) { + try { + // Clean NaN values from JSON string before parsing + let cleanedJson = metadata[field]; + cleanedJson = cleanedJson.replace(/:\s*NaN\b/g, ': null'); + cleanedJson = cleanedJson.replace(/\bNaN\b/g, 'null'); + + workflowData = JSON.parse(cleanedJson); + break; + } catch (e) { + console.log('Failed to parse workflow field:', field); + } + } + } + + for (const field of promptFields) { + if (metadata[field]) { + try { + // Clean NaN values from JSON string before parsing + let cleanedJson = metadata[field]; + cleanedJson = cleanedJson.replace(/:\s*NaN\b/g, ': null'); + cleanedJson = cleanedJson.replace(/\bNaN\b/g, 'null'); + + promptData = JSON.parse(cleanedJson); + break; + } catch (e) { + console.log('Failed to parse prompt field:', field); + } + } + } + + return { workflow: workflowData, prompt: promptData }; + } + + async extractImageMetadata(imageUrl) { + try { + // Fetch the image as an array buffer + const response = await fetch(imageUrl); + const arrayBuffer = await response.arrayBuffer(); + + // Parse PNG metadata + const metadata = await this.parsePNGMetadata(arrayBuffer); + const comfyData = this.extractComfyUIData(metadata); + + return this.parseWorkflowData(comfyData, imageUrl); + } catch (error) { + console.error('Error extracting metadata:', error); + return null; + } + } + + parseWorkflowData(comfyData, imagePath) { + // Extract information from ComfyUI workflow + let checkpoint = 'Unknown'; + let positivePrompt = 'No prompt found'; + let negativePrompt = 'No negative prompt found'; + let steps = 'Unknown'; + let cfgScale = 'Unknown'; + let sampler = 'Unknown'; + let seed = 'Unknown'; + + // Parse prompt data first (more reliable for actual generation parameters) + if (comfyData.prompt) { + const promptNodes = comfyData.prompt; + + for (const nodeId in promptNodes) { + const node = promptNodes[nodeId]; + + // Checkpoint + if (node.class_type === 'CheckpointLoaderSimple' && node.inputs) { + checkpoint = node.inputs.ckpt_name || checkpoint; + } + + // Prompts - need to identify which is positive vs negative + if (node.class_type === 'PromptManager' && node.inputs && node.inputs.text) { + // PromptManager typically contains the positive prompt + const promptValue = typeof node.inputs.text === 'string' ? node.inputs.text : (Array.isArray(node.inputs.text) ? node.inputs.text[0] : String(node.inputs.text)); + positivePrompt = promptValue; + } + + if (node.class_type === 'CLIPTextEncode' && node.inputs && node.inputs.text) { + // Check if this looks like a negative prompt + const textValue = typeof node.inputs.text === 'string' ? node.inputs.text : (Array.isArray(node.inputs.text) ? node.inputs.text[0] : String(node.inputs.text)); + const text = textValue.toLowerCase(); + if (text.includes('bad anatomy') || text.includes('unfinished') || + text.includes('censored') || text.includes('weird anatomy') || + text.includes('negative') || text.includes('embedding:')) { + negativePrompt = textValue; + } else if (positivePrompt === 'No prompt found') { + // If we haven't found a positive prompt yet, this might be it + positivePrompt = textValue; + } + } + + // Sampling parameters + if (node.class_type === 'KSampler' && node.inputs) { + seed = node.inputs.seed || seed; + steps = node.inputs.steps || steps; + cfgScale = node.inputs.cfg || cfgScale; + sampler = node.inputs.sampler_name || sampler; + } + } + } + + // Parse workflow data (this is what we actually have in your case) + if (comfyData.workflow && comfyData.workflow.nodes) { + const nodes = comfyData.workflow.nodes; + console.log('DEBUG: Total nodes found:', nodes.length); + + // Debug: Log all nodes to understand the structure + nodes.forEach((node, index) => { + console.log(`Node ${index} (ID: ${node.id}):`, { + type: node.type, + title: node.title, + widgets_values: node.widgets_values + }); + }); + + // Look for checkpoint loader + const checkpointNode = nodes.find(node => + node.type === 'CheckpointLoaderSimple' || + node.type === 'CheckpointLoader' + ); + if (checkpointNode && checkpointNode.widgets_values && checkpointNode.widgets_values[0]) { + checkpoint = checkpointNode.widgets_values[0]; + } + + // Look for prompts in workflow nodes + const textEncodeNodes = nodes.filter(node => node.type === 'CLIPTextEncode'); + const promptManagerNodes = nodes.filter(node => node.type === 'PromptManager'); + + // Check PromptManager first for positive prompt + if (promptManagerNodes.length > 0 && promptManagerNodes[0].widgets_values && promptManagerNodes[0].widgets_values[0]) { + positivePrompt = promptManagerNodes[0].widgets_values[0]; + } + + // For negative prompt, look for CLIPTextEncode that contains negative keywords + for (const node of textEncodeNodes) { + if (node.widgets_values && node.widgets_values[0]) { + const text = node.widgets_values[0].toLowerCase(); + if (text.includes('bad anatomy') || text.includes('unfinished') || + text.includes('censored') || text.includes('weird anatomy') || + text.includes('negative') || text.includes('embedding:')) { + negativePrompt = node.widgets_values[0]; + break; + } + } + } + + // Extract generation parameters from their source nodes + // Look for specific node types with better fallback logic + + // CFG Scale - prioritize Float nodes with title 'CFG', then other sources + console.log('DEBUG: Looking for CFG scale...'); + + // First priority: Float node with title 'CFG' + const cfgFloatNode = nodes.find(node => node.type === 'Float' && node.title === 'CFG'); + if (cfgFloatNode && cfgFloatNode.widgets_values && cfgFloatNode.widgets_values[0]) { + cfgScale = cfgFloatNode.widgets_values[0]; + console.log(`DEBUG: Found CFG in Float node: ${cfgScale}`); + } + + // If not found, look for specific value of 7 (your target CFG) + if (cfgScale === 'Unknown') { + for (const node of nodes) { + if (node.widgets_values && Array.isArray(node.widgets_values)) { + for (let i = 0; i < node.widgets_values.length; i++) { + const value = node.widgets_values[i]; + if (value === 7) { // Target the specific value you mentioned + cfgScale = value; + console.log(`DEBUG: Found target CFG value 7 in node ${node.id} (${node.type}) at index ${i}`); + break; + } + } + if (cfgScale !== 'Unknown') break; + } + } + } + + // Fallback: any reasonable CFG value + if (cfgScale === 'Unknown') { + for (const node of nodes) { + if (node.widgets_values && Array.isArray(node.widgets_values)) { + for (let i = 0; i < node.widgets_values.length; i++) { + const value = node.widgets_values[i]; + if (typeof value === 'number' && value > 1 && value <= 30) { + console.log(`DEBUG: Found potential CFG ${value} in node ${node.id} (${node.type}) at index ${i}`); + cfgScale = value; + console.log(`DEBUG: Set CFG scale to ${value}`); + break; + } + } + if (cfgScale !== 'Unknown') break; + } + } + } + + // Steps - prioritize Int nodes with title 'Steps', then target value 30 + console.log('DEBUG: Looking for steps...'); + + // First priority: Int node with title 'Steps' + const stepsIntNode = nodes.find(node => node.type === 'Int' && node.title === 'Steps'); + if (stepsIntNode && stepsIntNode.widgets_values && stepsIntNode.widgets_values[0]) { + steps = stepsIntNode.widgets_values[0]; + console.log(`DEBUG: Found steps in Int node: ${steps}`); + } + + // If not found, look for specific value of 30 (your target steps) + if (steps === 'Unknown') { + for (const node of nodes) { + if (node.widgets_values && Array.isArray(node.widgets_values)) { + for (let i = 0; i < node.widgets_values.length; i++) { + const value = node.widgets_values[i]; + if (value === 30) { // Target the specific value you mentioned + steps = value; + console.log(`DEBUG: Found target steps value 30 in node ${node.id} (${node.type}) at index ${i}`); + break; + } + } + if (steps !== 'Unknown') break; + } + } + } + + // Fallback: reasonable step values, prefer 10-150 range + if (steps === 'Unknown') { + for (const node of nodes) { + if (node.widgets_values && Array.isArray(node.widgets_values)) { + for (let i = 0; i < node.widgets_values.length; i++) { + const value = node.widgets_values[i]; + if (typeof value === 'number' && value >= 10 && value <= 150) { + console.log(`DEBUG: Found good steps value ${value} in node ${node.id} (${node.type}) at index ${i}`); + steps = value; + console.log(`DEBUG: Set steps to ${value}`); + break; + } + } + if (steps !== 'Unknown') break; + } + } + } + + // Sampler - look for valid ComfyUI samplers + console.log('DEBUG: Looking for sampler...'); + const validSamplers = [ + 'euler', 'euler_ancestral', 'heun', 'dpm_2', 'dpm_2_ancestral', 'lms', + 'dpm_fast', 'dpm_adaptive', 'dpmpp_2s_ancestral', 'dpmpp_sde', 'dpmpp_sde_gpu', + 'dpmpp_2m', 'dpmpp_2m_sde', 'dpmpp_2m_sde_gpu', 'dpmpp_3m_sde', 'dpmpp_3m_sde_gpu', + 'ddim', 'uni_pc', 'uni_pc_bh2' + ]; + for (const node of nodes) { + if (node.widgets_values && Array.isArray(node.widgets_values)) { + for (let i = 0; i < node.widgets_values.length; i++) { + const value = node.widgets_values[i]; + if (typeof value === 'string') { + const samplerValue = value.toLowerCase(); + if (validSamplers.some(validSampler => samplerValue.includes(validSampler))) { + console.log(`DEBUG: Found sampler ${value} in node ${node.id} (${node.type}) at index ${i}`); + if (sampler === 'Unknown') { + sampler = value; + console.log(`DEBUG: Set sampler to ${value}`); + } + } + } + } + } + } + + // Seed - look for large numbers that could be seeds + console.log('DEBUG: Looking for seed...'); + for (const node of nodes) { + if (node.widgets_values && Array.isArray(node.widgets_values)) { + for (let i = 0; i < node.widgets_values.length; i++) { + const value = node.widgets_values[i]; + if (typeof value === 'number' && value > 100000) { + console.log(`DEBUG: Found potential seed ${value} in node ${node.id} (${node.type}) at index ${i}`); + if (seed === 'Unknown') { + seed = value; + console.log(`DEBUG: Set seed to ${value}`); + } + } + } + } + } + + // If we still don't have sampling params, look at KSampler widgets_values directly + const ksamplerNode = nodes.find(node => + node.type === 'KSampler' || node.type === 'KSamplerAdvanced' + ); + if (ksamplerNode && ksamplerNode.widgets_values) { + const widgets = ksamplerNode.widgets_values; + // KSampler widgets_values order: [seed, control_mode, steps, cfg, sampler_name, scheduler, denoise, ...] + if (seed === 'Unknown' && widgets[0]) seed = widgets[0]; + if (steps === 'Unknown' && widgets[2]) steps = widgets[2]; + if (cfgScale === 'Unknown' && widgets[3]) cfgScale = widgets[3]; + if (sampler === 'Unknown' && widgets[4]) sampler = widgets[4]; + } + } + + return { + positivePrompt, + negativePrompt, + checkpoint, + steps, + cfgScale, + sampler, + seed, + workflow: comfyData.workflow, + prompt: comfyData.prompt, + imagePath + }; + } + + addMetadataSidebar() { + const viewerContainer = document.querySelector('.viewer-container'); + if (!viewerContainer || document.getElementById('metadata-sidebar')) return; + + const sidebar = document.createElement('div'); + sidebar.id = 'metadata-sidebar'; + sidebar.className = 'metadata-sidebar'; + sidebar.innerHTML = ` + +
+
+ + + +

Generation data

+
+ +
+ + +
+ +
+ + + +

Loading Metadata...

+

Extracting ComfyUI workflow data

+
+
+ `; + + viewerContainer.appendChild(sidebar); + this.attachMetadataEventListeners(); + } + + async loadMetadataForImage(imgElement) { + const imageUrl = imgElement.getAttribute('data-original') || imgElement.src; + if (!imageUrl) return; + + const metadataContent = document.getElementById('metadata-content'); + if (!metadataContent) return; + + // Show loading state + metadataContent.innerHTML = ` +
+
+ + + +

Loading Metadata...

+

Extracting ComfyUI workflow data

+
+
+ `; + + try { + const metadata = await this.extractImageMetadata(imageUrl); + this.currentMetadata = metadata; + this.updateMetadataSidebar(metadata); + } catch (error) { + console.error('Failed to load metadata:', error); + metadataContent.innerHTML = ` +
+

Failed to Load Metadata

+

Could not extract ComfyUI workflow data

+
+ `; + } + } + + updateMetadataSidebar(metadata) { + const metadataContent = document.getElementById('metadata-content'); + if (!metadataContent || !metadata) return; + + metadataContent.innerHTML = ` + +
+

File Path

+
+ ${metadata.imagePath} +
+
+ + +
+

Resources used

+
+
+
${metadata.checkpoint}
+
ComfyUI Generated
+
+ CHECKPOINT +
+
+ + +
+
+

Prompt

+ COMFYUI + +
+
+ ${metadata.positivePrompt.substring(0, 200)}${metadata.positivePrompt.length > 200 ? '...' : ''} +
+ ${metadata.positivePrompt.length > 200 ? '' : ''} +
+ + +
+
+

Negative prompt

+ +
+
+ ${metadata.negativePrompt.substring(0, 200)}${metadata.negativePrompt.length > 200 ? '...' : ''} +
+ ${metadata.negativePrompt.length > 200 ? '' : ''} +
+ + +
+

Other metadata

+
+ CFG SCALE: ${metadata.cfgScale} + STEPS: ${metadata.steps} + SAMPLER: ${metadata.sampler} +
+
+ SEED: ${metadata.seed} +
+
+ + +
+

ComfyUI Workflow

+
+ + +
+
+ `; + + // Re-attach event listeners after updating content + this.attachMetadataEventListeners(); + } + + attachMetadataEventListeners() { + const sidebar = document.getElementById('metadata-sidebar'); + if (!sidebar) return; + + // Copy buttons + const copyBtns = sidebar.querySelectorAll('[data-copy-type]'); + copyBtns.forEach(btn => { + btn.addEventListener('click', (e) => { + const type = btn.getAttribute('data-copy-type'); + if (this.currentMetadata) { + const text = type === 'positive' ? this.currentMetadata.positivePrompt : this.currentMetadata.negativePrompt; + this.copyToClipboard(text); + } + }); + }); + + // File path copy + const pathEl = sidebar.querySelector('[data-copy-path]'); + if (pathEl) { + pathEl.addEventListener('click', () => { + if (this.currentMetadata) { + this.copyToClipboard(this.currentMetadata.imagePath); + } + }); + } + + // Show more buttons + const showMoreBtns = sidebar.querySelectorAll('[data-show-type]'); + showMoreBtns.forEach(btn => { + btn.addEventListener('click', (e) => { + const type = btn.getAttribute('data-show-type'); + this.showFullPrompt(type); + }); + }); + + // Workflow action buttons + const actionBtns = sidebar.querySelectorAll('[data-action]'); + actionBtns.forEach(btn => { + btn.addEventListener('click', (e) => { + const action = btn.getAttribute('data-action'); + if (action === 'show-workflow') { + this.showWorkflowData(); + } else if (action === 'download-workflow') { + this.downloadWorkflowJSON(); + } + }); + }); + + // Copy all button + const copyAllBtn = sidebar.querySelector('.metadata-copy-all'); + if (copyAllBtn) { + copyAllBtn.addEventListener('click', () => { + this.copyAllMetadata(); + }); + } + } + + showFullPrompt(type) { + if (this.currentMetadata) { + const prompt = type === 'positive' ? this.currentMetadata.positivePrompt : this.currentMetadata.negativePrompt; + const newWindow = window.open('', '_blank'); + newWindow.document.write(` + + ${type.charAt(0).toUpperCase() + type.slice(1)} Prompt + +

${type.charAt(0).toUpperCase() + type.slice(1)} Prompt

+
${prompt}
+ + + + `); + } + } + + showWorkflowData() { + if (this.currentMetadata && this.currentMetadata.workflow) { + const newWindow = window.open('', '_blank'); + newWindow.document.write(` + + ComfyUI Workflow Data + +

ComfyUI Workflow JSON

+
${JSON.stringify(this.currentMetadata.workflow, null, 2)}
+ + + `); + } + } + + downloadWorkflowJSON() { + if (this.currentMetadata && this.currentMetadata.workflow) { + const dataStr = JSON.stringify(this.currentMetadata.workflow, null, 2); + const dataBlob = new Blob([dataStr], {type: 'application/json'}); + const url = URL.createObjectURL(dataBlob); + const link = document.createElement('a'); + link.href = url; + link.download = 'comfyui_workflow.json'; + document.body.appendChild(link); + link.click(); + document.body.removeChild(link); + URL.revokeObjectURL(url); + } + } + + copyAllMetadata() { + if (!this.currentMetadata) return; + + const allData = `Checkpoint: ${this.currentMetadata.checkpoint || 'Unknown'} +Positive Prompt: ${this.currentMetadata.positivePrompt} +Negative Prompt: ${this.currentMetadata.negativePrompt} +Steps: ${this.currentMetadata.steps || 'Unknown'} +CFG Scale: ${this.currentMetadata.cfgScale || 'Unknown'} +Sampler: ${this.currentMetadata.sampler || 'Unknown'} +Seed: ${this.currentMetadata.seed || 'Unknown'}`; + + this.copyToClipboard(allData); + } + + // ========================================== + // Auto Tag Feature Methods + // ========================================== + + // State for AutoTag feature + autoTagState = { + eventSource: null, + downloadEventSource: null, + reviewImages: [], + reviewIndex: 0, + currentTags: [], + cancelled: false, + skipAllTagged: false, + retagResolve: null, // Promise resolver for retag confirmation + tagsExpanded: false // Track accordion state + }; + + // ==================== Add Prompt Modal ==================== + + showAddPromptModal() { + // Clear form + document.getElementById('addPromptText').value = ''; + document.getElementById('addPromptCategory').value = ''; + document.getElementById('addPromptRatingValue').value = ''; + document.getElementById('addPromptNotes').value = ''; + document.getElementById('addPromptProtected').checked = true; + + // Clear tags + this.addPromptTags = []; + this.renderAddPromptTags(); + + // Reset rating stars + this.updateAddPromptRatingStars(0); + + // Populate categories datalist + this.populateAddPromptCategories(); + + this.showModal('addPromptModal'); + } + + setupAddPromptModal() { + this.addPromptTags = []; + + // Rating stars + const ratingContainer = document.getElementById('addPromptRating'); + const ratingStars = ratingContainer.querySelectorAll('.rating-star'); + ratingStars.forEach(star => { + star.addEventListener('click', () => { + const rating = parseInt(star.dataset.rating); + document.getElementById('addPromptRatingValue').value = rating; + this.updateAddPromptRatingStars(rating); + }); + }); + + // Clear rating button + document.getElementById('clearAddPromptRatingBtn').addEventListener('click', () => { + document.getElementById('addPromptRatingValue').value = ''; + this.updateAddPromptRatingStars(0); + }); + + // Tag input + const tagInput = document.getElementById('addPromptTagInput'); + tagInput.addEventListener('keydown', (e) => { + if (e.key === 'Enter' || e.key === ',') { + e.preventDefault(); + const tag = tagInput.value.trim().replace(/,/g, ''); + if (tag && !this.addPromptTags.includes(tag)) { + this.addPromptTags.push(tag); + this.renderAddPromptTags(); + } + tagInput.value = ''; + } + }); + + // Tag suggestions + tagInput.addEventListener('input', () => { + this.showAddPromptTagSuggestions(tagInput.value); + }); + } + + updateAddPromptRatingStars(rating) { + const stars = document.querySelectorAll('#addPromptRating .rating-star'); + stars.forEach((star, index) => { + if (index < rating) { + star.textContent = '★'; + star.classList.add('text-yellow-400'); + star.classList.remove('text-gray-500'); + } else { + star.textContent = '☆'; + star.classList.remove('text-yellow-400'); + star.classList.add('text-gray-500'); + } + }); + } + + renderAddPromptTags() { + const container = document.getElementById('addPromptTagsContainer'); + const input = document.getElementById('addPromptTagInput'); + + // Remove existing tag chips + container.querySelectorAll('.tag-chip').forEach(chip => chip.remove()); + + // Add tag chips + this.addPromptTags.forEach(tag => { + const chip = document.createElement('span'); + chip.className = 'tag-chip inline-flex items-center px-2 py-1 bg-blue-600 text-white text-xs rounded cursor-pointer hover:bg-blue-700'; + chip.innerHTML = `${tag} ×`; + chip.addEventListener('click', () => { + this.addPromptTags = this.addPromptTags.filter(t => t !== tag); + this.renderAddPromptTags(); + }); + container.insertBefore(chip, input); + }); + } + + showAddPromptTagSuggestions(query) { + const suggestionsContainer = document.getElementById('addPromptTagSuggestions'); + if (!query) { + suggestionsContainer.classList.add('hidden'); + return; + } + + const matchingTags = this.tags.filter(tag => + tag.toLowerCase().includes(query.toLowerCase()) && + !this.addPromptTags.includes(tag) + ).slice(0, 10); + + if (matchingTags.length === 0) { + suggestionsContainer.classList.add('hidden'); + return; + } + + suggestionsContainer.innerHTML = matchingTags.map(tag => ` +
+ ${tag} +
+ `).join(''); + + suggestionsContainer.querySelectorAll('[data-tag]').forEach(el => { + el.addEventListener('click', () => { + const tag = el.dataset.tag; + if (!this.addPromptTags.includes(tag)) { + this.addPromptTags.push(tag); + this.renderAddPromptTags(); + } + document.getElementById('addPromptTagInput').value = ''; + suggestionsContainer.classList.add('hidden'); + }); + }); + + suggestionsContainer.classList.remove('hidden'); + } + + populateAddPromptCategories() { + const datalist = document.getElementById('addPromptCategoryList'); + datalist.innerHTML = this.categories.map(cat => `