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.