From 7695240065a187bf2f2f361b7239665516274938 Mon Sep 17 00:00:00 2001 From: Vito Sansevero Date: Sat, 24 May 2025 12:26:07 -0700 Subject: [PATCH] feat: Add PromptManager with persistent storage --- prompt_manager.py | 320 ++++++++++++++++++++++++++++++++++++++++++++++ prompts.db | Bin 0 -> 53248 bytes 2 files changed, 320 insertions(+) create mode 100644 prompt_manager.py create mode 100644 prompts.db diff --git a/prompt_manager.py b/prompt_manager.py new file mode 100644 index 0000000..b6a5b83 --- /dev/null +++ b/prompt_manager.py @@ -0,0 +1,320 @@ +""" +PromptManager: Main custom node implementation that extends CLIPTextEncode +with persistent prompt storage and search capabilities. +""" + +import hashlib +import datetime +import json +import webbrowser +import os +from typing import Optional, Dict, Any, Tuple, List + +try: + from comfy.comfy_types import IO, ComfyNodeABC, InputTypeDict +except ImportError: + # Fallback for older ComfyUI versions + class ComfyNodeABC: + pass + + class IO: + STRING = "STRING" + CLIP = "CLIP" + CONDITIONING = "CONDITIONING" + + InputTypeDict = dict + +try: + from .database.operations import PromptDatabase + from .utils.prompt_tracker import PromptTracker, PromptExecutionContext + from .utils.image_monitor import ImageMonitor +except ImportError: + # For direct imports when not in a package + import sys + import os + sys.path.insert(0, os.path.dirname(os.path.abspath(__file__))) + from database.operations import PromptDatabase + from utils.prompt_tracker import PromptTracker, PromptExecutionContext + from utils.image_monitor import ImageMonitor + + +class PromptManager(ComfyNodeABC): + """ + A ComfyUI custom node that functions like CLIPTextEncode but adds: + - Persistent storage of all prompts in SQLite database + - Search and retrieval capabilities + - Metadata management (categories, tags, ratings, notes) + - Duplicate detection via SHA256 hashing + """ + + def __init__(self): + self.db = PromptDatabase() + self.prompt_tracker = PromptTracker(self.db) + self.image_monitor = ImageMonitor(self.db, self.prompt_tracker) + + # Start image monitoring automatically + self._start_gallery_system() + + @classmethod + def INPUT_TYPES(cls) -> InputTypeDict: + return { + "required": { + "text": (IO.STRING, { + "multiline": True, + "dynamicPrompts": True, + "tooltip": "The text prompt to be encoded and saved to database." + }), + "clip": (IO.CLIP, { + "tooltip": "The CLIP model used for encoding the text." + }) + }, + "optional": { + "category": (IO.STRING, { + "default": "", + "tooltip": "Optional category for organizing prompts (e.g., 'landscapes', 'portraits')" + }), + "tags": (IO.STRING, { + "default": "", + "tooltip": "Comma-separated tags for the prompt (e.g., 'anime, detailed, sunset')" + }), + "search_text": (IO.STRING, { + "default": "", + "tooltip": "Search for past prompts containing this text" + }) + } + } + + RETURN_TYPES = (IO.CONDITIONING,) + OUTPUT_TOOLTIPS = ( + "A conditioning containing the embedded text used to guide the diffusion model.", + ) + FUNCTION = "encode" + CATEGORY = "PromptManager/Text" + DESCRIPTION = ( + "Encodes a text prompt using a CLIP model into an embedding that can be used to guide " + "the diffusion model towards generating specific images. Additionally saves all prompts " + "to a local SQLite database with optional metadata for search and retrieval." + ) + + def encode( + self, + clip, + text: str, + category: str = "", + tags: str = "", + search_text: str = "" + ) -> Tuple[Any]: + """ + Encode the text prompt and save it to the database. + + Args: + clip: The CLIP model for encoding + text: The text prompt to encode + category: Optional category for organization + tags: Comma-separated tags + search_text: Text to search for in past prompts + + Returns: + Tuple containing the conditioning for the diffusion model + + Raises: + RuntimeError: If clip input is invalid + """ + # 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: + raise RuntimeError( + "ERROR: clip input is invalid: None\n\n" + "If the clip is from a checkpoint loader node your checkpoint does not " + "contain a valid clip or text encoder model." + ) + + # Save prompt to database and set execution context for gallery tracking + prompt_id = None + if text and text.strip(): + try: + prompt_id = self._save_prompt_to_database( + text=text.strip(), # Always strip whitespace + category=category.strip() if category else None, + tags=self._parse_tags(tags) + ) + + # Set current prompt for image tracking + if prompt_id: + execution_id = self.prompt_tracker.set_current_prompt( + prompt_text=text.strip(), + additional_data={ + 'category': category.strip() if category else None, + 'tags': self._parse_tags(tags), + 'prompt_id': prompt_id + } + ) + print(f"[PromptManager] Set execution context: {execution_id}") + + except Exception as e: + # Log error but don't fail the encoding + print(f"Warning: Failed to save prompt to database: {e}") + + # Perform standard CLIP text encoding + tokens = clip.tokenize(text) + conditioning = clip.encode_from_tokens_scheduled(tokens) + + return (conditioning,) + + 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) + + # 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]): + self.db.update_prompt_metadata( + prompt_id=existing['id'], + category=category, + tags=tags + ) + return existing['id'] + + # Save new prompt + prompt_id = self.db.save_prompt( + text=text, + category=category, + tags=tags, + prompt_hash=prompt_hash + ) + + return prompt_id + + except Exception as e: + print(f"Error saving prompt to database: {e}") + return None + + def _generate_hash(self, text: str) -> str: + """Generate SHA256 hash for the prompt text.""" + # 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.""" + 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.""" + 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: + print(f"Error searching prompts: {e}") + return [] + + def _open_web_interface(self): + """Open the web interface in the default browser.""" + 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}") + print("\n=== Web interface opened in browser ===") + else: + print("\n=== Web interface directory found but no index.html ===") + print(f"Please check {web_dir} for setup instructions") + else: + print("\n=== Web interface not yet implemented ===") + print("This feature will open a web-based prompt management interface") + print("when the web_interface directory is created.") + + except Exception as e: + print(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: + print(f"Error getting recent prompts: {e}") + return [] + + def _start_gallery_system(self): + """Initialize and start the gallery monitoring system.""" + try: + print("[PromptManager] Starting gallery system...") + + # Start image monitoring + self.image_monitor.start_monitoring() + + print("[PromptManager] Gallery system started successfully") + + except Exception as e: + print(f"[PromptManager] Failed to start gallery system: {e}") + print("[PromptManager] 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() + print("[PromptManager] Gallery system cleaned up") + except Exception as e: + print(f"[PromptManager] Error cleaning up gallery system: {e}") + + def __del__(self): + """Cleanup when object is destroyed.""" + self.cleanup_gallery_system() + + @classmethod + def IS_CHANGED(cls, **kwargs): + """Always process to ensure database saving.""" + return float("NaN") # Always execute \ No newline at end of file diff --git a/prompts.db b/prompts.db new file mode 100644 index 0000000000000000000000000000000000000000..a08e8863482038d9dbbc234cfb021b18ac473e68 GIT binary patch literal 53248 zcmeI*%WvDr9S3kylpk`WFzpuA2F4as1B+To6jOZY5p;{HiIN(dIF2JX8w7zcB8L`X zJsgsX<1TV=+HHHTMB&2&KZvRtm$ zW?kB0FUvEwg-Lz&R_$)RdTX6rt6i^d-mH_Wn;RQ^EgPehqpKo*4Yb=H?#G%x0}cdGB}-9%8HK@5Lczr-z6$9tjZ7%iT#%aS7Rz;o8qkEm~wc zT&$SRtt+$Hix=hl(NQ3H)aGYmjetGsv#!m?CX;c=gJZy0;Y8rlH7l>{bzQ$= zm{$z_vQg0sramdfv`R*?Fe-f5snzp`ry8TA2TNTV_IV;QJ)GGb zoTQM!XZbt{TV6o;i#=vtUZIe81n{)oZ&7x%B51ltHq+OAeFz8}a{j@7k2)(weGgRsp=uggX3 zw&%K}_n7(ZUT<4?g)N?BGuoYqVOY9Rn3_^@X*Z<7UZdIf+79pg^X(QtWwnLMgQY=S zK|u#=B1WSX5F=4cG16g<#~#!6(!=Bpaq=%NJ(9R!!Z)Ajya@DiQ?pH%>Q5PT$+|rDF(LJaJ_>@!>9AE$1QjiQCW8sSWJmv z;@#eBGJmJv-R6yDo%t>oj1)%m#a8ODfLcb;p&GBDsk_aJW>>T_pT(S#QE4^{Wv1y( zYH|m&q}fGRtGKS?TD0jf%e8fvK1(Iil3lH?y!Sx9v7cG*`Jqp}kn~*MH1>LZ-W+y& z9oil2QKFA4+M}rFPpq+4D4W{U6w7Y$ZZ4ZQ?nZQ$-}86kcn_9#c-PmU>QOvl(Be_V z`!JJe4!xc;a2S>}?Ca!R%$80Kvs|zZooTevEH*X0T(nKgG;P{ z$`F761Rwwb2tWV=5P$##AOL~?oq(Zi%ddFLO){t4lmDxGQIr+MmS5q1jYL`5O?{=i zr54mLq@?<5N&SQRT>T3l!2kgWKmY;|fB*y_009U<00Izz!0RFK4P_;D?D-_;l=tPs zPvl50DCgzD11XfW@}4|=(n5mQ;o!0UgX{mFNb29!->Sb-KY2YWh_!$K1Rwwb2tWV= z5P$##AOHafoI-(bDOGvmgB7FaaVDoGM^C3zrY1)ZY)rhTR8x}*A3PHk&;L7x?Fm(a z00bZa0SG_<0uX=z1R(G>3*h;G6aR_G;QIdyN&T1lLj6+x=iA&wP~`Tcp|CDsTpx6fUN#V;x7ykfB*y_009U<00Izz00bZa z0SLU3!2WmSvyxgV8A>AYPHsQ*HGL&yLAVmo&Q6;(ybj%Bmv_2bVrFA}=IG1%CMlG! z=*AVpP(PCr{AYdb#>Y!ynWZZv^!x1TV5%{Am_MJnRBw3!;V()=HWCFQZqFxMwB2U@ z9(l}s&-H8?dcE$_(~qBu>;JO)rNmztAOHafKmY;|fB*y_009U<00Iy=eF8~EP3eR4 z|L2nW{Pa~2HG%*HAOHafKmY;|fB*y_009WRT>|HoN^0^q5C-S}{QmzJZ