Automated formatting pass across 30 files to establish consistent code style enforced by CI. No logic changes.
317 lines
12 KiB
Python
317 lines
12 KiB
Python
"""ComfyUI integration utilities for PromptManager.
|
|
|
|
This module provides deep integration with ComfyUI's metadata system to ensure that
|
|
PromptManager-generated prompts appear correctly in standard ComfyUI image metadata.
|
|
It patches core ComfyUI components to bridge the gap between PromptManager's custom
|
|
nodes and ComfyUI's standard metadata extraction.
|
|
|
|
Key features:
|
|
- Automatic patching of ComfyUI's SaveImage node
|
|
- Thread-safe prompt registration and retrieval
|
|
- Cross-thread prompt context sharing
|
|
- Standard metadata format compatibility
|
|
- Automatic cleanup of old prompt registrations
|
|
|
|
The integration works by:
|
|
1. PromptManager nodes register their prompts during execution
|
|
2. SaveImage node is patched to include registered prompts in metadata
|
|
3. Standard tools can then extract prompts from the generated images
|
|
|
|
Typical usage:
|
|
from utils.comfyui_integration import get_comfyui_integration
|
|
|
|
integration = get_comfyui_integration()
|
|
integration.register_prompt(node_id, prompt_text, metadata)
|
|
# Generated images will now include this prompt in their metadata
|
|
|
|
This integration is essential for:
|
|
- Third-party tool compatibility
|
|
- Standard metadata parsers
|
|
- Workflow sharing and reproduction
|
|
- Integration with existing ComfyUI ecosystems
|
|
"""
|
|
|
|
import threading
|
|
import time
|
|
import json
|
|
from typing import Dict, Any, Optional
|
|
|
|
try:
|
|
from .logging_config import get_logger
|
|
except ImportError:
|
|
import sys
|
|
import os
|
|
|
|
sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
|
|
from utils.logging_config import get_logger
|
|
|
|
|
|
class ComfyUIMetadataIntegration:
|
|
"""Integrates PromptManager with ComfyUI's standard metadata system.
|
|
|
|
This singleton class manages the integration between PromptManager custom nodes
|
|
and ComfyUI's standard metadata system. It ensures that prompts generated by
|
|
PromptManager appear in the standard ComfyUI metadata format for compatibility
|
|
with third-party tools and parsers.
|
|
|
|
The integration uses thread-local storage combined with global tracking to
|
|
handle prompt context across different execution threads, which is necessary
|
|
because ComfyUI executions may span multiple threads.
|
|
|
|
Key responsibilities:
|
|
- Register prompts from PromptManager nodes during execution
|
|
- Patch SaveImage to include PromptManager prompts in metadata
|
|
- Manage prompt lifecycle and cleanup
|
|
"""
|
|
|
|
_instance = None
|
|
_lock = threading.Lock()
|
|
|
|
def __new__(cls):
|
|
"""Ensure singleton pattern with thread safety.
|
|
|
|
Returns:
|
|
The single instance of ComfyUIMetadataIntegration
|
|
"""
|
|
if cls._instance is None:
|
|
with cls._lock:
|
|
if cls._instance is None:
|
|
cls._instance = super().__new__(cls)
|
|
return cls._instance
|
|
|
|
def __init__(self):
|
|
"""Initialize the ComfyUI integration system.
|
|
|
|
Sets up prompt tracking data structures and attempts to patch the
|
|
SaveImage node. Uses _initialized flag to prevent duplicate initialization.
|
|
"""
|
|
if hasattr(self, "_initialized"):
|
|
return
|
|
|
|
self.logger = get_logger("prompt_manager.comfyui_integration")
|
|
self._current_prompts = {}
|
|
self._thread_local = threading.local()
|
|
self._saveimage_patched = False
|
|
self._initialized = True
|
|
|
|
# Try to patch SaveImage node on initialization
|
|
self._patch_saveimage_node()
|
|
|
|
def register_prompt(self, node_id: str, prompt_text: str, metadata: Dict[str, Any]):
|
|
"""
|
|
Register a prompt from PromptManager for inclusion in ComfyUI metadata.
|
|
|
|
This method is called by PromptManager nodes during execution to register
|
|
their prompt text for later inclusion in image metadata. The prompt is stored
|
|
in both thread-local and global storage to handle cross-thread access scenarios.
|
|
|
|
Args:
|
|
node_id: Unique identifier for this prompt (typically the node ID)
|
|
prompt_text: The actual prompt text that was encoded by PromptManager
|
|
metadata: Additional metadata from PromptManager (category, tags, etc.)
|
|
"""
|
|
thread_id = threading.current_thread().ident
|
|
|
|
# Store in thread-local storage
|
|
if not hasattr(self._thread_local, "prompts"):
|
|
self._thread_local.prompts = {}
|
|
|
|
self._thread_local.prompts[node_id] = {
|
|
"text": prompt_text,
|
|
"metadata": metadata,
|
|
"timestamp": time.time(),
|
|
"thread_id": thread_id,
|
|
}
|
|
|
|
# Also store globally for cross-thread access
|
|
with self._lock:
|
|
self._current_prompts[f"{thread_id}_{node_id}"] = {
|
|
"text": prompt_text,
|
|
"metadata": metadata,
|
|
"timestamp": time.time(),
|
|
"thread_id": thread_id,
|
|
}
|
|
|
|
self.logger.debug(
|
|
f"Registered prompt for node {node_id}: {prompt_text[:50]}..."
|
|
)
|
|
|
|
def get_current_prompt_text(self, node_id: str = None) -> Optional[str]:
|
|
"""
|
|
Get the current prompt text for metadata inclusion.
|
|
|
|
Retrieves the most appropriate prompt text for inclusion in image metadata.
|
|
Uses a fallback strategy: thread-local storage first, then global storage,
|
|
with recency checks to avoid stale prompts.
|
|
|
|
Args:
|
|
node_id: Optional specific node ID to retrieve prompt for
|
|
|
|
Returns:
|
|
The prompt text string if available, None if no suitable prompt found
|
|
"""
|
|
# First try thread-local storage
|
|
if hasattr(self._thread_local, "prompts"):
|
|
if node_id and node_id in self._thread_local.prompts:
|
|
return self._thread_local.prompts[node_id]["text"]
|
|
elif self._thread_local.prompts:
|
|
# Return the most recent prompt from this thread
|
|
latest = max(
|
|
self._thread_local.prompts.values(), key=lambda x: x["timestamp"]
|
|
)
|
|
return latest["text"]
|
|
|
|
# Fallback to global storage
|
|
thread_id = threading.current_thread().ident
|
|
with self._lock:
|
|
# Look for prompts from current thread
|
|
thread_prompts = {
|
|
k: v
|
|
for k, v in self._current_prompts.items()
|
|
if v["thread_id"] == thread_id
|
|
}
|
|
|
|
if thread_prompts:
|
|
latest = max(thread_prompts.values(), key=lambda x: x["timestamp"])
|
|
return latest["text"]
|
|
|
|
# Last resort: return the most recent prompt from any thread
|
|
if self._current_prompts:
|
|
latest = max(
|
|
self._current_prompts.values(), key=lambda x: x["timestamp"]
|
|
)
|
|
# Only return if it's recent (within last 5 minutes)
|
|
if time.time() - latest["timestamp"] < 300:
|
|
return latest["text"]
|
|
|
|
return None
|
|
|
|
def _patch_saveimage_node(self):
|
|
"""
|
|
Patch ComfyUI's SaveImage node to include PromptManager prompts in metadata.
|
|
|
|
This method modifies ComfyUI's SaveImage.save_images method to automatically
|
|
include PromptManager prompts in the image metadata. The patching:
|
|
|
|
1. Wraps the original save_images method
|
|
2. Retrieves current PromptManager prompt text
|
|
3. Updates the text input in PromptManager nodes to reflect actual prompt
|
|
4. Calls the original method with the updated data
|
|
|
|
NOTE: We intentionally do NOT change class_type to CLIPTextEncode anymore.
|
|
That approach was corrupting saved workflows - when users saved and reloaded
|
|
workflows, ComfyUI would instantiate CLIPTextEncode instead of PromptManager,
|
|
causing errors with prepend_text/append_text inputs.
|
|
"""
|
|
try:
|
|
import nodes
|
|
|
|
if not hasattr(nodes, "SaveImage"):
|
|
self.logger.warning("SaveImage node not found in ComfyUI nodes")
|
|
return
|
|
|
|
# Store original save_images method
|
|
original_save_images = nodes.SaveImage.save_images
|
|
integration = self # Capture self reference
|
|
|
|
def patched_save_images(
|
|
self_node,
|
|
images,
|
|
filename_prefix="ComfyUI",
|
|
prompt=None,
|
|
extra_pnginfo=None,
|
|
):
|
|
"""Patched save_images method that includes PromptManager prompts."""
|
|
|
|
# Get current prompt text from PromptManager
|
|
current_prompt_text = integration.get_current_prompt_text()
|
|
|
|
if current_prompt_text:
|
|
integration.logger.debug(
|
|
f"Including PromptManager prompt in SaveImage metadata: {current_prompt_text[:50]}..."
|
|
)
|
|
|
|
# If no prompt provided, create one with our text
|
|
if prompt is None:
|
|
prompt = {}
|
|
|
|
# Ensure prompt has the standard structure ComfyUI expects
|
|
if not isinstance(prompt, dict):
|
|
prompt = {}
|
|
|
|
# Find PromptManager nodes and ensure prompt text is captured
|
|
# NOTE: We do NOT change class_type anymore - that was corrupting saved workflows
|
|
# when users reload them. PromptManager stays as PromptManager.
|
|
for node_id, node_data in prompt.items():
|
|
if isinstance(node_data, dict):
|
|
class_type = node_data.get("class_type", "")
|
|
if "promptmanager" in class_type.lower():
|
|
# Ensure the text input reflects the actual prompt used
|
|
if "inputs" not in node_data:
|
|
node_data["inputs"] = {}
|
|
node_data["inputs"]["text"] = current_prompt_text
|
|
integration.logger.debug(
|
|
f"Updated PromptManager node {node_id} with prompt text"
|
|
)
|
|
|
|
# Call original method with potentially modified prompt
|
|
return original_save_images(
|
|
self_node, images, filename_prefix, prompt, extra_pnginfo
|
|
)
|
|
|
|
# Apply the patch
|
|
nodes.SaveImage.save_images = patched_save_images
|
|
self._saveimage_patched = True
|
|
self.logger.info(
|
|
"Successfully patched SaveImage node for PromptManager integration"
|
|
)
|
|
|
|
except Exception as e:
|
|
self.logger.error(f"Failed to patch SaveImage node: {e}")
|
|
self.logger.warning(
|
|
"PromptManager prompts may not appear in standard ComfyUI metadata"
|
|
)
|
|
|
|
def cleanup_old_prompts(self, max_age_seconds: int = 600):
|
|
"""
|
|
Clean up old prompt registrations.
|
|
|
|
Removes prompt registrations that are older than the specified age to
|
|
prevent memory leaks and ensure that only recent, relevant prompts are
|
|
used for metadata inclusion.
|
|
|
|
Args:
|
|
max_age_seconds: Maximum age in seconds before a prompt registration
|
|
is considered stale and removed (default: 600 seconds/10 minutes)
|
|
"""
|
|
current_time = time.time()
|
|
|
|
with self._lock:
|
|
old_keys = [
|
|
key
|
|
for key, prompt_data in self._current_prompts.items()
|
|
if current_time - prompt_data["timestamp"] > max_age_seconds
|
|
]
|
|
|
|
for key in old_keys:
|
|
del self._current_prompts[key]
|
|
|
|
if old_keys:
|
|
self.logger.debug(f"Cleaned up {len(old_keys)} old prompt registrations")
|
|
|
|
|
|
# Global instance
|
|
_integration_instance = None
|
|
|
|
|
|
def get_comfyui_integration() -> ComfyUIMetadataIntegration:
|
|
"""Get the global ComfyUI integration instance.
|
|
|
|
Returns:
|
|
The singleton ComfyUIMetadataIntegration instance, creating it if necessary
|
|
"""
|
|
global _integration_instance
|
|
if _integration_instance is None:
|
|
_integration_instance = ComfyUIMetadataIntegration()
|
|
return _integration_instance
|