style: apply Black formatter (line-length=88) to all Python files

Automated formatting pass across 30 files to establish consistent code
style enforced by CI. No logic changes.
This commit is contained in:
Vito Sansevero
2026-02-07 08:30:09 -08:00
parent 224fcc7803
commit 2737a4b1ef
30 changed files with 3278 additions and 2531 deletions
+6 -1
View File
@@ -8,4 +8,9 @@ image monitoring, metadata extraction, logging, and system diagnostics.
from .hashing import generate_prompt_hash
from .validators import validate_prompt_text, validate_rating, validate_tags
__all__ = ["generate_prompt_hash", "validate_prompt_text", "validate_rating", "validate_tags"]
__all__ = [
"generate_prompt_hash",
"validate_prompt_text",
"validate_rating",
"validate_tags",
]
+112 -84
View File
@@ -19,7 +19,7 @@ The integration works by:
Typical usage:
from utils.comfyui_integration import get_comfyui_integration
integration = get_comfyui_integration()
integration.register_prompt(node_id, prompt_text, metadata)
# Generated images will now include this prompt in their metadata
@@ -41,34 +41,35 @@ try:
except ImportError:
import sys
import os
sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
from utils.logging_config import get_logger
class ComfyUIMetadataIntegration:
"""Integrates PromptManager with ComfyUI's standard metadata system.
This singleton class manages the integration between PromptManager custom nodes
and ComfyUI's standard metadata system. It ensures that prompts generated by
PromptManager appear in the standard ComfyUI metadata format for compatibility
with third-party tools and parsers.
The integration uses thread-local storage combined with global tracking to
handle prompt context across different execution threads, which is necessary
because ComfyUI executions may span multiple threads.
Key responsibilities:
- Register prompts from PromptManager nodes during execution
- Patch SaveImage to include PromptManager prompts in metadata
- Manage prompt lifecycle and cleanup
"""
_instance = None
_lock = threading.Lock()
def __new__(cls):
"""Ensure singleton pattern with thread safety.
Returns:
The single instance of ComfyUIMetadataIntegration
"""
@@ -77,109 +78,118 @@ class ComfyUIMetadataIntegration:
if cls._instance is None:
cls._instance = super().__new__(cls)
return cls._instance
def __init__(self):
"""Initialize the ComfyUI integration system.
Sets up prompt tracking data structures and attempts to patch the
SaveImage node. Uses _initialized flag to prevent duplicate initialization.
"""
if hasattr(self, '_initialized'):
if hasattr(self, "_initialized"):
return
self.logger = get_logger('prompt_manager.comfyui_integration')
self.logger = get_logger("prompt_manager.comfyui_integration")
self._current_prompts = {}
self._thread_local = threading.local()
self._saveimage_patched = False
self._initialized = True
# Try to patch SaveImage node on initialization
self._patch_saveimage_node()
def register_prompt(self, node_id: str, prompt_text: str, metadata: Dict[str, Any]):
"""
Register a prompt from PromptManager for inclusion in ComfyUI metadata.
This method is called by PromptManager nodes during execution to register
their prompt text for later inclusion in image metadata. The prompt is stored
in both thread-local and global storage to handle cross-thread access scenarios.
Args:
node_id: Unique identifier for this prompt (typically the node ID)
prompt_text: The actual prompt text that was encoded by PromptManager
metadata: Additional metadata from PromptManager (category, tags, etc.)
"""
thread_id = threading.current_thread().ident
# Store in thread-local storage
if not hasattr(self._thread_local, 'prompts'):
if not hasattr(self._thread_local, "prompts"):
self._thread_local.prompts = {}
self._thread_local.prompts[node_id] = {
'text': prompt_text,
'metadata': metadata,
'timestamp': time.time(),
'thread_id': thread_id
"text": prompt_text,
"metadata": metadata,
"timestamp": time.time(),
"thread_id": thread_id,
}
# Also store globally for cross-thread access
with self._lock:
self._current_prompts[f"{thread_id}_{node_id}"] = {
'text': prompt_text,
'metadata': metadata,
'timestamp': time.time(),
'thread_id': thread_id
"text": prompt_text,
"metadata": metadata,
"timestamp": time.time(),
"thread_id": thread_id,
}
self.logger.debug(f"Registered prompt for node {node_id}: {prompt_text[:50]}...")
self.logger.debug(
f"Registered prompt for node {node_id}: {prompt_text[:50]}..."
)
def get_current_prompt_text(self, node_id: str = None) -> Optional[str]:
"""
Get the current prompt text for metadata inclusion.
Retrieves the most appropriate prompt text for inclusion in image metadata.
Uses a fallback strategy: thread-local storage first, then global storage,
with recency checks to avoid stale prompts.
Args:
node_id: Optional specific node ID to retrieve prompt for
Returns:
The prompt text string if available, None if no suitable prompt found
"""
# First try thread-local storage
if hasattr(self._thread_local, 'prompts'):
if hasattr(self._thread_local, "prompts"):
if node_id and node_id in self._thread_local.prompts:
return self._thread_local.prompts[node_id]['text']
return self._thread_local.prompts[node_id]["text"]
elif self._thread_local.prompts:
# Return the most recent prompt from this thread
latest = max(self._thread_local.prompts.values(), key=lambda x: x['timestamp'])
return latest['text']
latest = max(
self._thread_local.prompts.values(), key=lambda x: x["timestamp"]
)
return latest["text"]
# Fallback to global storage
thread_id = threading.current_thread().ident
with self._lock:
# Look for prompts from current thread
thread_prompts = {k: v for k, v in self._current_prompts.items()
if v['thread_id'] == thread_id}
thread_prompts = {
k: v
for k, v in self._current_prompts.items()
if v["thread_id"] == thread_id
}
if thread_prompts:
latest = max(thread_prompts.values(), key=lambda x: x['timestamp'])
return latest['text']
latest = max(thread_prompts.values(), key=lambda x: x["timestamp"])
return latest["text"]
# Last resort: return the most recent prompt from any thread
if self._current_prompts:
latest = max(self._current_prompts.values(), key=lambda x: x['timestamp'])
latest = max(
self._current_prompts.values(), key=lambda x: x["timestamp"]
)
# Only return if it's recent (within last 5 minutes)
if time.time() - latest['timestamp'] < 300:
return latest['text']
if time.time() - latest["timestamp"] < 300:
return latest["text"]
return None
def _patch_saveimage_node(self):
"""
Patch ComfyUI's SaveImage node to include PromptManager prompts in metadata.
This method modifies ComfyUI's SaveImage.save_images method to automatically
include PromptManager prompts in the image metadata. The patching:
@@ -195,80 +205,97 @@ class ComfyUIMetadataIntegration:
"""
try:
import nodes
if not hasattr(nodes, 'SaveImage'):
if not hasattr(nodes, "SaveImage"):
self.logger.warning("SaveImage node not found in ComfyUI nodes")
return
# Store original save_images method
original_save_images = nodes.SaveImage.save_images
integration = self # Capture self reference
def patched_save_images(self_node, images, filename_prefix="ComfyUI", prompt=None, extra_pnginfo=None):
def patched_save_images(
self_node,
images,
filename_prefix="ComfyUI",
prompt=None,
extra_pnginfo=None,
):
"""Patched save_images method that includes PromptManager prompts."""
# Get current prompt text from PromptManager
current_prompt_text = integration.get_current_prompt_text()
if current_prompt_text:
integration.logger.debug(f"Including PromptManager prompt in SaveImage metadata: {current_prompt_text[:50]}...")
integration.logger.debug(
f"Including PromptManager prompt in SaveImage metadata: {current_prompt_text[:50]}..."
)
# If no prompt provided, create one with our text
if prompt is None:
prompt = {}
# Ensure prompt has the standard structure ComfyUI expects
if not isinstance(prompt, dict):
prompt = {}
# Find PromptManager nodes and ensure prompt text is captured
# NOTE: We do NOT change class_type anymore - that was corrupting saved workflows
# when users reload them. PromptManager stays as PromptManager.
for node_id, node_data in prompt.items():
if isinstance(node_data, dict):
class_type = node_data.get('class_type', '')
if 'promptmanager' in class_type.lower():
class_type = node_data.get("class_type", "")
if "promptmanager" in class_type.lower():
# Ensure the text input reflects the actual prompt used
if 'inputs' not in node_data:
node_data['inputs'] = {}
node_data['inputs']['text'] = current_prompt_text
integration.logger.debug(f"Updated PromptManager node {node_id} with prompt text")
if "inputs" not in node_data:
node_data["inputs"] = {}
node_data["inputs"]["text"] = current_prompt_text
integration.logger.debug(
f"Updated PromptManager node {node_id} with prompt text"
)
# Call original method with potentially modified prompt
return original_save_images(self_node, images, filename_prefix, prompt, extra_pnginfo)
return original_save_images(
self_node, images, filename_prefix, prompt, extra_pnginfo
)
# Apply the patch
nodes.SaveImage.save_images = patched_save_images
self._saveimage_patched = True
self.logger.info("Successfully patched SaveImage node for PromptManager integration")
self.logger.info(
"Successfully patched SaveImage node for PromptManager integration"
)
except Exception as e:
self.logger.error(f"Failed to patch SaveImage node: {e}")
self.logger.warning("PromptManager prompts may not appear in standard ComfyUI metadata")
self.logger.warning(
"PromptManager prompts may not appear in standard ComfyUI metadata"
)
def cleanup_old_prompts(self, max_age_seconds: int = 600):
"""
Clean up old prompt registrations.
Removes prompt registrations that are older than the specified age to
prevent memory leaks and ensure that only recent, relevant prompts are
used for metadata inclusion.
Args:
max_age_seconds: Maximum age in seconds before a prompt registration
is considered stale and removed (default: 600 seconds/10 minutes)
"""
current_time = time.time()
with self._lock:
old_keys = [
key for key, prompt_data in self._current_prompts.items()
if current_time - prompt_data['timestamp'] > max_age_seconds
key
for key, prompt_data in self._current_prompts.items()
if current_time - prompt_data["timestamp"] > max_age_seconds
]
for key in old_keys:
del self._current_prompts[key]
if old_keys:
self.logger.debug(f"Cleaned up {len(old_keys)} old prompt registrations")
@@ -276,13 +303,14 @@ class ComfyUIMetadataIntegration:
# Global instance
_integration_instance = None
def get_comfyui_integration() -> ComfyUIMetadataIntegration:
"""Get the global ComfyUI integration instance.
Returns:
The singleton ComfyUIMetadataIntegration instance, creating it if necessary
"""
global _integration_instance
if _integration_instance is None:
_integration_instance = ComfyUIMetadataIntegration()
return _integration_instance
return _integration_instance
+146 -152
View File
@@ -15,10 +15,10 @@ The diagnostic system checks:
Typical usage:
from utils.diagnostics import GalleryDiagnostics, run_diagnostics
# Run full diagnostic suite
results = run_diagnostics()
# Or create custom diagnostic instance
diagnostics = GalleryDiagnostics("custom_db.db")
database_status = diagnostics.check_database()
@@ -40,39 +40,39 @@ from .logging_config import get_logger
class GalleryDiagnostics:
"""Diagnostics for the gallery system.
This class provides comprehensive diagnostic capabilities for the PromptManager
gallery system. It systematically checks all components and dependencies to
identify potential issues and provide actionable feedback.
The diagnostics cover:
- Database structure and connectivity
- Image tracking table integrity
- File system access and permissions
- ComfyUI integration points
- Python dependency availability
Each diagnostic method returns a standardized result dictionary with:
- status: 'ok', 'warning', or 'error'
- message: Descriptive message (for warnings/errors)
- Additional data specific to the diagnostic
"""
def __init__(self, db_path: str = "prompts.db"):
"""Initialize the diagnostics system.
Args:
db_path: Path to the SQLite database file to diagnose
"""
self.db_path = db_path
self.logger = get_logger('prompt_manager.diagnostics')
self.logger = get_logger("prompt_manager.diagnostics")
def run_full_diagnostic(self) -> Dict[str, Any]:
"""Run a complete diagnostic check.
Executes all diagnostic checks in sequence and provides a comprehensive
report of system status. Logs detailed information during the process.
Returns:
Dictionary mapping diagnostic categories to their results:
- database: Database connectivity and structure check
@@ -81,37 +81,37 @@ class GalleryDiagnostics:
- comfyui_output: ComfyUI output directory detection
- dependencies: Python dependency availability check
"""
self.logger.info("\n" + "="*60)
self.logger.info("\n" + "=" * 60)
self.logger.info("[DIAG] PROMPTMANAGER GALLERY DIAGNOSTICS")
self.logger.info("="*60)
self.logger.info("=" * 60)
results = {
'database': self.check_database(),
'images_table': self.check_images_table(),
'file_system': self.check_file_system(),
'comfyui_output': self.check_comfyui_output(),
'dependencies': self.check_dependencies()
"database": self.check_database(),
"images_table": self.check_images_table(),
"file_system": self.check_file_system(),
"comfyui_output": self.check_comfyui_output(),
"dependencies": self.check_dependencies(),
}
self.logger.info("\n" + "="*60)
self.logger.info("\n" + "=" * 60)
self.logger.info("[SUMMARY] DIAGNOSTIC SUMMARY")
self.logger.info("="*60)
self.logger.info("=" * 60)
for category, result in results.items():
status = "[PASS] PASS" if result['status'] == 'ok' else "[FAIL] FAIL"
status = "[PASS] PASS" if result["status"] == "ok" else "[FAIL] FAIL"
self.logger.info(f"{category.upper():<20} {status}")
if result['status'] != 'ok':
if result["status"] != "ok":
self.logger.warning(f" Issue: {result['message']}")
self.logger.info("\n" + "="*60)
self.logger.info("\n" + "=" * 60)
return results
def check_database(self) -> Dict[str, Any]:
"""Check database connection and structure.
Verifies that the database file exists, is accessible, and contains
the expected prompt table structure.
Returns:
Dictionary containing:
- status: 'ok' or 'error'
@@ -120,49 +120,46 @@ class GalleryDiagnostics:
- has_images_table: Whether generated_images table exists (if successful)
"""
self.logger.info("\n[DB] Checking Database...")
try:
if not os.path.exists(self.db_path):
return {
'status': 'error',
'message': f'Database file not found: {self.db_path}'
"status": "error",
"message": f"Database file not found: {self.db_path}",
}
with sqlite3.connect(self.db_path) as conn:
conn.row_factory = sqlite3.Row
# Check prompts table
cursor = conn.execute("SELECT COUNT(*) as count FROM prompts")
prompt_count = cursor.fetchone()['count']
prompt_count = cursor.fetchone()["count"]
self.logger.info(f" [NOTE] Prompts in database: {prompt_count}")
# Check if generated_images table exists
cursor = conn.execute("""
SELECT name FROM sqlite_master
WHERE type='table' AND name='generated_images'
""")
has_images_table = cursor.fetchone() is not None
self.logger.info(f" [IMG] Images table exists: {has_images_table}")
return {
'status': 'ok',
'prompt_count': prompt_count,
'has_images_table': has_images_table
"status": "ok",
"prompt_count": prompt_count,
"has_images_table": has_images_table,
}
except Exception as e:
return {
'status': 'error',
'message': f'Database error: {str(e)}'
}
return {"status": "error", "message": f"Database error: {str(e)}"}
def check_images_table(self) -> Dict[str, Any]:
"""Check the generated_images table specifically.
Examines the generated_images table structure and content to verify
the gallery system can function properly.
Returns:
Dictionary containing:
- status: 'ok' or 'error'
@@ -171,28 +168,28 @@ class GalleryDiagnostics:
- recent_images: List of recent image records (if successful)
"""
self.logger.info("\n[IMG] Checking Images Table...")
try:
with sqlite3.connect(self.db_path) as conn:
conn.row_factory = sqlite3.Row
# Check if table exists
cursor = conn.execute("""
SELECT name FROM sqlite_master
WHERE type='table' AND name='generated_images'
""")
if not cursor.fetchone():
return {
'status': 'error',
'message': 'generated_images table does not exist - run the updated code to create it'
"status": "error",
"message": "generated_images table does not exist - run the updated code to create it",
}
# Check image records
cursor = conn.execute("SELECT COUNT(*) as count FROM generated_images")
image_count = cursor.fetchone()['count']
image_count = cursor.fetchone()["count"]
self.logger.info(f" [STATS] Images in database: {image_count}")
# Get recent images
cursor = conn.execute("""
SELECT gi.*, p.text
@@ -202,29 +199,28 @@ class GalleryDiagnostics:
LIMIT 5
""")
recent_images = [dict(row) for row in cursor.fetchall()]
self.logger.info(f" [TIME] Recent images: {len(recent_images)}")
for img in recent_images:
self.logger.info(f" - {img['filename']} -> Prompt {img['prompt_id']}")
self.logger.info(
f" - {img['filename']} -> Prompt {img['prompt_id']}"
)
return {
'status': 'ok',
'image_count': image_count,
'recent_images': recent_images
"status": "ok",
"image_count": image_count,
"recent_images": recent_images,
}
except Exception as e:
return {
'status': 'error',
'message': f'Images table error: {str(e)}'
}
return {"status": "error", "message": f"Images table error: {str(e)}"}
def check_file_system(self) -> Dict[str, Any]:
"""Check file system and permissions.
Verifies that the application has appropriate file system access
for reading images and writing database files.
Returns:
Dictionary containing:
- status: 'ok' or 'error'
@@ -233,42 +229,35 @@ class GalleryDiagnostics:
- can_write: Whether write access is available
"""
self.logger.info("\n[DIR] Checking File System...")
try:
# Check current directory
current_dir = os.getcwd()
self.logger.info(f" [FOLDER] Current directory: {current_dir}")
# Check if we can write to current directory
test_file = "test_write.tmp"
try:
with open(test_file, 'w') as f:
with open(test_file, "w") as f:
f.write("test")
os.remove(test_file)
can_write = True
except (IOError, OSError):
can_write = False
self.logger.info(f" [EDIT] Can write to directory: {can_write}")
return {
'status': 'ok',
'current_dir': current_dir,
'can_write': can_write
}
return {"status": "ok", "current_dir": current_dir, "can_write": can_write}
except Exception as e:
return {
'status': 'error',
'message': f'File system error: {str(e)}'
}
return {"status": "error", "message": f"File system error: {str(e)}"}
def check_comfyui_output(self) -> Dict[str, Any]:
"""Check ComfyUI output directories.
Attempts to locate ComfyUI output directories where generated images
would be stored. Checks both common locations and ComfyUI's configured paths.
Returns:
Dictionary containing:
- status: 'ok' if directories found, 'warning' if none found
@@ -276,56 +265,57 @@ class GalleryDiagnostics:
- output_dirs: List of detected output directory paths
"""
self.logger.info("\n[STYLE] Checking ComfyUI Output...")
output_dirs = []
# Try to detect ComfyUI output directories
potential_dirs = [
"output",
"../output",
"../../output",
"../../output",
"ComfyUI/output",
"../ComfyUI/output",
"../../ComfyUI/output"
"../../ComfyUI/output",
]
for dir_path in potential_dirs:
abs_path = os.path.abspath(dir_path)
if os.path.exists(abs_path):
output_dirs.append(abs_path)
self.logger.info(f" [DIR] Found output dir: {abs_path}")
# Count images in this directory
try:
image_files = []
for ext in ['.png', '.jpg', '.jpeg', '.webp']:
image_files.extend(Path(abs_path).rglob(f'*{ext}'))
for ext in [".png", ".jpg", ".jpeg", ".webp"]:
image_files.extend(Path(abs_path).rglob(f"*{ext}"))
self.logger.info(f" [IMG] Images found: {len(image_files)}")
except Exception as e:
self.logger.error(f" [FAIL] Error scanning: {e}")
# Try ComfyUI's folder_paths
try:
import folder_paths
comfyui_output = folder_paths.get_output_directory()
if comfyui_output and comfyui_output not in output_dirs:
output_dirs.append(comfyui_output)
self.logger.info(f" [DIR] ComfyUI output dir: {comfyui_output}")
except ImportError:
self.logger.warning(" [WARN] ComfyUI folder_paths not available")
return {
'status': 'ok' if output_dirs else 'warning',
'message': 'No output directories found' if not output_dirs else None,
'output_dirs': output_dirs
"status": "ok" if output_dirs else "warning",
"message": "No output directories found" if not output_dirs else None,
"output_dirs": output_dirs,
}
def check_dependencies(self) -> Dict[str, Any]:
"""Check required dependencies.
Verifies that all required Python packages are available for the
PromptManager system to function properly.
Returns:
Dictionary containing:
- status: 'ok' if all dependencies available, 'error' if any missing
@@ -333,55 +323,56 @@ class GalleryDiagnostics:
- dependencies: Dictionary mapping package names to availability status
"""
self.logger.info("\n[PKG] Checking Dependencies...")
dependencies = {
'watchdog': False,
'PIL': False,
'sqlite3': False
}
dependencies = {"watchdog": False, "PIL": False, "sqlite3": False}
# Check watchdog
try:
import watchdog
dependencies['watchdog'] = True
dependencies["watchdog"] = True
self.logger.info(f" [PASS] watchdog: {watchdog.__version__}")
except ImportError:
self.logger.error(" [FAIL] watchdog: NOT INSTALLED")
# Check PIL
try:
from PIL import Image
dependencies['PIL'] = True
dependencies["PIL"] = True
self.logger.info(f" [PASS] PIL (Pillow): Available")
except ImportError:
self.logger.error(" [FAIL] PIL (Pillow): NOT AVAILABLE")
# Check sqlite3
try:
import sqlite3
dependencies['sqlite3'] = True
dependencies["sqlite3"] = True
self.logger.info(f" [PASS] sqlite3: {sqlite3.sqlite_version}")
except ImportError:
self.logger.error(" [FAIL] sqlite3: NOT AVAILABLE")
all_deps_ok = all(dependencies.values())
return {
'status': 'ok' if all_deps_ok else 'error',
'message': 'Missing dependencies' if not all_deps_ok else None,
'dependencies': dependencies
"status": "ok" if all_deps_ok else "error",
"message": "Missing dependencies" if not all_deps_ok else None,
"dependencies": dependencies,
}
def create_test_image_link(self, prompt_id: int, test_image_path: str = None) -> Dict[str, Any]:
def create_test_image_link(
self, prompt_id: int, test_image_path: str = None
) -> Dict[str, Any]:
"""Create a test image link to verify the system works.
Creates a test entry in the generated_images table to verify that
the image linking functionality is working correctly.
Args:
prompt_id: ID of an existing prompt to link the test image to
test_image_path: Optional path for the test image (uses fake path if None)
Returns:
Dictionary containing:
- status: 'ok' if test link created successfully, 'error' otherwise
@@ -389,59 +380,62 @@ class GalleryDiagnostics:
- image_id: ID of created test image record (if successful)
"""
self.logger.info(f"\n[TEST] Creating test image link for prompt {prompt_id}...")
try:
# Import database operations
import sys
import os
sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
sys.path.insert(
0, os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
)
from database.operations import PromptDatabase
db = PromptDatabase()
# Create a fake image path if none provided
if not test_image_path:
test_image_path = "/fake/test/image.png"
# Create test metadata
test_metadata = {
'file_info': {
'size': 1024000,
'dimensions': [512, 512],
'format': 'PNG'
"file_info": {
"size": 1024000,
"dimensions": [512, 512],
"format": "PNG",
},
'workflow': {'test': True},
'prompt': {'test_prompt': 'This is a test'}
"workflow": {"test": True},
"prompt": {"test_prompt": "This is a test"},
}
# Link the test image
image_id = db.link_image_to_prompt(
prompt_id=str(prompt_id),
image_path=test_image_path,
metadata=test_metadata
metadata=test_metadata,
)
self.logger.info(f" [PASS] Test image linked with ID: {image_id}")
return {
'status': 'ok',
'image_id': image_id,
'message': f'Test image linked successfully with ID {image_id}'
"status": "ok",
"image_id": image_id,
"message": f"Test image linked successfully with ID {image_id}",
}
except Exception as e:
return {
'status': 'error',
'message': f'Failed to create test link: {str(e)}'
"status": "error",
"message": f"Failed to create test link: {str(e)}",
}
def run_diagnostics() -> Dict[str, Any]:
"""Run diagnostics from command line.
Convenience function to create a GalleryDiagnostics instance and run
the full diagnostic suite with default settings.
Returns:
Complete diagnostic results dictionary
"""
@@ -450,4 +444,4 @@ def run_diagnostics() -> Dict[str, Any]:
if __name__ == "__main__":
run_diagnostics()
run_diagnostics()
+38 -28
View File
@@ -13,11 +13,11 @@ Key features:
Typical usage:
from utils.hashing import generate_prompt_hash, is_duplicate_prompt
hash1 = generate_prompt_hash("Beautiful landscape")
hash2 = generate_prompt_hash(" beautiful landscape ")
# hash1 == hash2 (normalization makes them identical)
if is_duplicate_prompt(text1, text2):
print("Duplicate prompts detected")
@@ -34,20 +34,20 @@ import hashlib
def generate_prompt_hash(text: str) -> str:
"""
Generate a SHA256 hash for prompt text to enable deduplication.
Normalizes the input text (strips whitespace, converts to lowercase)
before hashing to ensure consistent results for functionally identical
prompts with minor formatting differences.
Args:
text: The prompt text to hash
Returns:
SHA256 hexadecimal digest of the normalized text
Raises:
TypeError: If the input is not a string
Example:
>>> generate_prompt_hash("Beautiful Landscape")
'c3d5f8a9b2e1...'
@@ -57,69 +57,79 @@ def generate_prompt_hash(text: str) -> str:
"""
if not isinstance(text, str):
raise TypeError("Text must be a string")
# Normalize the text by stripping whitespace and converting to lowercase
# for consistent hashing regardless of minor formatting differences
normalized_text = text.strip().lower()
return hashlib.sha256(normalized_text.encode('utf-8')).hexdigest()
return hashlib.sha256(normalized_text.encode("utf-8")).hexdigest()
def generate_content_hash(content: dict) -> str:
"""
Generate a hash for prompt content including metadata.
Creates a comprehensive hash that includes not just the prompt text
but also associated metadata like category, tags, and workflow name.
This enables detection of prompts that are identical in all aspects.
Args:
content: Dictionary containing prompt data with optional keys:
- text: The prompt text
- category: Prompt category
- tags: List of tags
- workflow_name: Associated workflow name
Returns:
SHA256 hexadecimal digest of the normalized content structure
Note:
The hash is generated from a normalized JSON representation with
sorted keys and normalized text fields to ensure consistency.
"""
import json
# Create a normalized representation of the content
normalized = {
'text': content.get('text', '').strip().lower(),
'category': content.get('category', '').strip().lower() if content.get('category') else '',
'tags': sorted([tag.strip().lower() for tag in content.get('tags', []) if tag.strip()]),
'workflow_name': content.get('workflow_name', '').strip().lower() if content.get('workflow_name') else ''
"text": content.get("text", "").strip().lower(),
"category": (
content.get("category", "").strip().lower()
if content.get("category")
else ""
),
"tags": sorted(
[tag.strip().lower() for tag in content.get("tags", []) if tag.strip()]
),
"workflow_name": (
content.get("workflow_name", "").strip().lower()
if content.get("workflow_name")
else ""
),
}
# Convert to JSON string for consistent hashing
content_str = json.dumps(normalized, sort_keys=True)
return hashlib.sha256(content_str.encode('utf-8')).hexdigest()
return hashlib.sha256(content_str.encode("utf-8")).hexdigest()
def is_duplicate_prompt(text1: str, text2: str, threshold: float = 0.95) -> bool:
"""
Check if two prompts are likely duplicates using hash comparison.
Compares the normalized hashes of two prompt texts to determine if
they are functionally identical. This is an exact match after normalization,
not a similarity measure.
Args:
text1: First prompt text to compare
text2: Second prompt text to compare
threshold: Similarity threshold (not used - kept for API compatibility)
Returns:
True if the prompts have identical normalized hashes (i.e., are duplicates),
False otherwise
Note:
The threshold parameter is not used in the current implementation as
this performs exact hash matching rather than similarity scoring.
@@ -127,5 +137,5 @@ def is_duplicate_prompt(text1: str, text2: str, threshold: float = 0.95) -> bool
"""
hash1 = generate_prompt_hash(text1)
hash2 = generate_prompt_hash(text2)
return hash1 == hash2
return hash1 == hash2
+105 -88
View File
@@ -11,7 +11,7 @@ The main components are:
Typical usage:
from utils.image_monitor import ImageMonitor
monitor = ImageMonitor(db_manager, prompt_tracker)
monitor.start_monitoring(['/path/to/comfyui/output'])
@@ -38,22 +38,22 @@ from .logging_config import get_logger
class ImageGenerationHandler(FileSystemEventHandler):
"""Filesystem event handler for detecting new image generation.
This handler extends watchdog's FileSystemEventHandler to specifically handle
new image file creation events in ComfyUI output directories. When a new image
is detected, it attempts to:
1. Extract ComfyUI metadata from the image
2. Associate the image with the currently active prompt
3. Store the relationship in the database
The handler implements a small delay before processing to ensure files are
completely written before attempting to read them.
"""
def __init__(self, db_manager, prompt_tracker):
"""
Initialize the image generation handler.
Args:
db_manager: Database manager instance for storing image-prompt relationships
prompt_tracker: Prompt tracking instance for getting current active prompts
@@ -61,24 +61,25 @@ class ImageGenerationHandler(FileSystemEventHandler):
self.db_manager = db_manager
self.prompt_tracker = prompt_tracker
self.metadata_extractor = ComfyUIMetadataExtractor()
self.logger = get_logger('prompt_manager.image_monitor')
self.logger = get_logger("prompt_manager.image_monitor")
# Read from GalleryConfig if available, otherwise use defaults
try:
from ..py.config import GalleryConfig
self.processing_delay = GalleryConfig.PROCESSING_DELAY
self.supported_extensions = tuple(GalleryConfig.SUPPORTED_EXTENSIONS)
except Exception:
self.processing_delay = 2.0
self.supported_extensions = ('.png', '.jpg', '.jpeg', '.webp', '.gif')
self.supported_extensions = (".png", ".jpg", ".jpeg", ".webp", ".gif")
def on_created(self, event):
"""Handle filesystem creation events.
This method is called by watchdog when a new file is created in a monitored
directory. It filters for image files and schedules them for processing after
a small delay to ensure the file is fully written.
Args:
event: FileSystemEvent object containing event details
"""
@@ -86,11 +87,9 @@ class ImageGenerationHandler(FileSystemEventHandler):
self.logger.info(f"New image detected: {event.src_path}")
# Small delay to ensure file is fully written
threading.Timer(
self.processing_delay,
self.process_new_image,
args=[event.src_path]
self.processing_delay, self.process_new_image, args=[event.src_path]
).start()
def is_image_file(self, filepath: str) -> bool:
"""Check if file is a supported image format and not a thumbnail.
@@ -102,20 +101,20 @@ class ImageGenerationHandler(FileSystemEventHandler):
inside a thumbnails directory, False otherwise
"""
# Skip files in thumbnails directory - those are derivatives, not generated images
if '/thumbnails/' in filepath or '\\thumbnails\\' in filepath:
if "/thumbnails/" in filepath or "\\thumbnails\\" in filepath:
return False
return filepath.lower().endswith(self.supported_extensions)
def process_new_image(self, image_path: str):
"""Process a newly created image file for gallery integration.
This method handles the complete processing pipeline for a new image:
1. Verifies the file still exists
2. Gets the current prompt context from the tracker
3. Extracts ComfyUI metadata from the image
4. Links the image to the appropriate prompt in the database
5. Handles fallback scenarios when no active prompt is available
Args:
image_path: Full path to the newly created image file
"""
@@ -128,8 +127,10 @@ class ImageGenerationHandler(FileSystemEventHandler):
# Get current prompt context first
current_prompt = self.prompt_tracker.get_current_prompt()
self.logger.info(f"Current prompt context: {current_prompt['id'] if current_prompt else 'None'}")
self.logger.info(
f"Current prompt context: {current_prompt['id'] if current_prompt else 'None'}"
)
if not current_prompt:
self.logger.debug(f"No active prompt context for image: {image_path}")
# Fallback: try to link to the most recent prompt in database
@@ -141,9 +142,11 @@ class ImageGenerationHandler(FileSystemEventHandler):
return
else:
# Extend the timeout for this prompt since we're still getting images
if 'execution_id' in current_prompt:
self.prompt_tracker.extend_prompt_timeout(current_prompt['execution_id'], 300) # Add 5 more minutes
if "execution_id" in current_prompt:
self.prompt_tracker.extend_prompt_timeout(
current_prompt["execution_id"], 300
) # Add 5 more minutes
# Extract ComfyUI metadata
try:
metadata = self.metadata_extractor.extract_metadata(image_path)
@@ -151,30 +154,37 @@ class ImageGenerationHandler(FileSystemEventHandler):
except Exception as meta_error:
self.logger.warning(f"Metadata extraction failed: {meta_error}")
metadata = None
if metadata:
self.logger.debug(f"Linking image with full metadata to prompt {current_prompt['id']}")
self.logger.debug(
f"Linking image with full metadata to prompt {current_prompt['id']}"
)
self.link_image_to_prompt(image_path, current_prompt, metadata)
else:
self.logger.debug(f"Linking image with basic info to prompt {current_prompt['id']}")
self.logger.debug(
f"Linking image with basic info to prompt {current_prompt['id']}"
)
# Link with basic file info even without metadata
basic_metadata = self.get_basic_file_info(image_path)
self.link_image_to_prompt(image_path, current_prompt, {'file_info': basic_metadata})
self.link_image_to_prompt(
image_path, current_prompt, {"file_info": basic_metadata}
)
except Exception as e:
self.logger.error(f"Error processing image {image_path}: {e}")
import traceback
self.logger.error(traceback.format_exc())
def get_basic_file_info(self, image_path: str) -> Dict[str, Any]:
"""Get basic file information when metadata extraction fails.
Provides fallback file information when ComfyUI metadata cannot be extracted
from the image. Includes file size, format, and dimensions when possible.
Args:
image_path: Path to the image file
Returns:
Dictionary containing basic file information:
- size: File size in bytes
@@ -183,34 +193,30 @@ class ImageGenerationHandler(FileSystemEventHandler):
"""
try:
from PIL import Image
stat = os.stat(image_path)
file_info = {
'size': stat.st_size,
'format': None,
'dimensions': None
}
file_info = {"size": stat.st_size, "format": None, "dimensions": None}
# Try to get image dimensions
try:
with Image.open(image_path) as img:
file_info['dimensions'] = list(img.size)
file_info['format'] = img.format
file_info["dimensions"] = list(img.size)
file_info["format"] = img.format
except Exception:
pass
return file_info
except Exception as e:
self.logger.error(f"Error getting file info: {e}")
return {}
def _get_fallback_prompt(self) -> Optional[Dict[str, Any]]:
"""Get the most recent prompt from database as fallback.
When no active prompt is available from the tracker, this method attempts
to find the most recently created prompt in the database to use as a fallback
for image linking.
Returns:
Dictionary containing prompt information with 'fallback' flag set to True,
or None if no recent prompt is available
@@ -220,21 +226,23 @@ class ImageGenerationHandler(FileSystemEventHandler):
if recent_prompts:
prompt = recent_prompts[0]
return {
'id': prompt['id'],
'text': prompt['text'],
'timestamp': prompt.get('created_at'),
'fallback': True
"id": prompt["id"],
"text": prompt["text"],
"timestamp": prompt.get("created_at"),
"fallback": True,
}
except Exception as e:
self.logger.error(f"Error getting fallback prompt: {e}")
return None
def link_image_to_prompt(self, image_path: str, prompt_context: Dict, metadata: Dict):
def link_image_to_prompt(
self, image_path: str, prompt_context: Dict, metadata: Dict
):
"""Link an image to a prompt in the database.
Creates a database record associating the generated image with its source prompt,
including any extracted metadata from the image file.
Args:
image_path: Full path to the image file
prompt_context: Dictionary containing prompt information including ID and text
@@ -242,33 +250,33 @@ class ImageGenerationHandler(FileSystemEventHandler):
"""
try:
image_id = self.db_manager.link_image_to_prompt(
prompt_id=prompt_context['id'],
image_path=image_path,
metadata=metadata
prompt_id=prompt_context["id"], image_path=image_path, metadata=metadata
)
fallback_note = " (fallback)" if prompt_context.get("fallback") else ""
self.logger.debug(
f"Successfully linked image {image_id} to prompt {prompt_context['id']}{fallback_note}"
)
fallback_note = " (fallback)" if prompt_context.get('fallback') else ""
self.logger.debug(f"Successfully linked image {image_id} to prompt {prompt_context['id']}{fallback_note}")
except Exception as e:
self.logger.error(f"Failed to link image to prompt: {e}")
class ImageMonitor:
"""Main image monitoring system for ComfyUI gallery integration.
This class manages the overall image monitoring system, including:
- Setting up filesystem watchers for output directories
- Auto-detecting ComfyUI output locations
- Managing the lifecycle of monitoring operations
- Providing status information
The monitor uses watchdog to efficiently watch filesystem changes and can
monitor multiple directories simultaneously with recursive subdirectory support.
"""
def __init__(self, db_manager, prompt_tracker):
"""
Initialize the image monitor.
Args:
db_manager: Database manager instance for storing image relationships
prompt_tracker: Prompt tracking instance for getting active prompt context
@@ -278,8 +286,8 @@ class ImageMonitor:
self.observer = None
self.handler = None
self.monitored_directories = []
self.logger = get_logger('prompt_manager.image_monitor')
self.logger = get_logger("prompt_manager.image_monitor")
def start_monitoring(self, output_directories: Optional[list] = None):
"""
Start monitoring ComfyUI output directories for new images.
@@ -299,6 +307,7 @@ class ImageMonitor:
# Check if monitoring is enabled in config
try:
from ..py.config import GalleryConfig
if not GalleryConfig.MONITORING_ENABLED:
self.logger.info("Image monitoring disabled in config")
return
@@ -309,26 +318,29 @@ class ImageMonitor:
if not output_directories:
try:
from py.config import GalleryConfig
if GalleryConfig.MONITORING_DIRECTORIES:
output_directories = GalleryConfig.MONITORING_DIRECTORIES
self.logger.info(f"Using configured monitoring directories: {output_directories}")
self.logger.info(
f"Using configured monitoring directories: {output_directories}"
)
except ImportError:
pass
# Auto-detect ComfyUI output directory if still none
if not output_directories:
output_directories = self.detect_comfyui_output_dirs()
if not output_directories:
self.logger.warning("No output directories found to monitor")
return
# Create event handler
self.handler = ImageGenerationHandler(self.db_manager, self.prompt_tracker)
# Start observer
self.observer = Observer()
for output_dir in output_directories:
if os.path.exists(output_dir):
self.observer.schedule(self.handler, output_dir, recursive=True)
@@ -339,13 +351,15 @@ class ImageMonitor:
if self.monitored_directories:
self.observer.start()
self.logger.info(f"Image monitoring started for {len(self.monitored_directories)} directories")
self.logger.info(
f"Image monitoring started for {len(self.monitored_directories)} directories"
)
else:
self.logger.warning("No valid directories to monitor")
def stop_monitoring(self):
"""Stop the image monitoring system.
Cleanly shuts down the filesystem watcher and clears all monitoring state.
This method should be called before program exit to ensure proper cleanup.
"""
@@ -356,47 +370,50 @@ class ImageMonitor:
self.handler = None
self.monitored_directories = []
self.logger.debug("Image monitoring stopped")
def detect_comfyui_output_dirs(self) -> list:
"""Auto-detect ComfyUI output directories.
Attempts to locate ComfyUI output directories using multiple strategies:
1. Import ComfyUI's folder_paths module to get the configured output directory
2. Search common relative paths where ComfyUI output directories are typically located
3. Verify that detected directories actually exist
Returns:
List of absolute paths to detected output directories
"""
potential_dirs = []
try:
# Try to import ComfyUI's folder_paths
import folder_paths
output_dir = folder_paths.get_output_directory()
if output_dir and os.path.exists(output_dir):
potential_dirs.append(output_dir)
self.logger.debug(f"Detected ComfyUI output directory: {output_dir}")
except ImportError:
self.logger.debug("ComfyUI folder_paths not available, using fallback detection")
self.logger.debug(
"ComfyUI folder_paths not available, using fallback detection"
)
# Fallback: Look for common ComfyUI directory structures
fallback_paths = [
"output",
"../output",
"../output",
"../../output",
"ComfyUI/output",
"../ComfyUI/output"
"../ComfyUI/output",
]
for path in fallback_paths:
abs_path = os.path.abspath(path)
if os.path.exists(abs_path) and abs_path not in potential_dirs:
potential_dirs.append(abs_path)
self.logger.debug(f"Found output directory: {abs_path}")
return potential_dirs
def get_status(self) -> Dict[str, Any]:
"""Get monitoring status information.
@@ -414,10 +431,10 @@ class ImageMonitor:
except Exception:
pass
return {
'running': self.observer is not None,
'monitored_directories': self.monitored_directories,
'handler_active': self.handler is not None,
'observer_alive': observer_alive
"running": self.observer is not None,
"monitored_directories": self.monitored_directories,
"handler_active": self.handler is not None,
"observer_alive": observer_alive,
}
@@ -444,4 +461,4 @@ def get_image_monitor(db_manager, prompt_tracker) -> ImageMonitor:
with _monitor_lock:
if _monitor_instance is None:
_monitor_instance = ImageMonitor(db_manager, prompt_tracker)
return _monitor_instance
return _monitor_instance
+155 -139
View File
@@ -23,7 +23,7 @@ from pathlib import Path
class PromptManagerLogger:
"""
Centralized logging system for PromptManager.
Features:
- Configurable log levels
- File rotation with size limits
@@ -31,13 +31,13 @@ class PromptManagerLogger:
- Thread-safe operations
- Memory buffer for recent logs
"""
_instance = None
_lock = threading.Lock()
def __new__(cls):
"""Ensure singleton pattern with thread safety.
Returns:
The single instance of PromptManagerLogger
"""
@@ -46,145 +46,149 @@ class PromptManagerLogger:
if cls._instance is None:
cls._instance = super().__new__(cls)
return cls._instance
def __init__(self):
"""Initialize the logging system if not already initialized.
Sets up log directory, configuration, memory buffer, and all handlers.
Uses _initialized flag to prevent duplicate initialization.
"""
if hasattr(self, '_initialized'):
if hasattr(self, "_initialized"):
return
self._initialized = True
self.log_dir = Path(__file__).parent.parent / "logs"
self.log_dir.mkdir(exist_ok=True)
# Configuration
self.config = {
'level': 'INFO',
'max_file_size': 10 * 1024 * 1024, # 10MB
'backup_count': 5,
'console_logging': True,
'file_logging': True,
'buffer_size': 1000 # Keep last 1000 log entries in memory
"level": "INFO",
"max_file_size": 10 * 1024 * 1024, # 10MB
"backup_count": 5,
"console_logging": True,
"file_logging": True,
"buffer_size": 1000, # Keep last 1000 log entries in memory
}
# Memory buffer for recent logs (for web viewer)
self._log_buffer = collections.deque(maxlen=self.config['buffer_size'])
self._log_buffer = collections.deque(maxlen=self.config["buffer_size"])
self._buffer_lock = threading.Lock()
# Initialize loggers
self._setup_loggers()
def _setup_loggers(self):
"""Set up all loggers with appropriate handlers.
Configures the main logger with console, file, and memory handlers.
Sets up file rotation and safe encoding for cross-platform compatibility.
"""
# Main logger
self.logger = logging.getLogger('prompt_manager')
self.logger.setLevel(getattr(logging, self.config['level']))
self.logger = logging.getLogger("prompt_manager")
self.logger.setLevel(getattr(logging, self.config["level"]))
# Clear existing handlers
for handler in self.logger.handlers[:]:
self.logger.removeHandler(handler)
# Custom formatter
formatter = logging.Formatter(
'%(asctime)s - %(name)s - %(levelname)s - %(filename)s:%(lineno)d - %(message)s',
datefmt='%Y-%m-%d %H:%M:%S'
"%(asctime)s - %(name)s - %(levelname)s - %(filename)s:%(lineno)d - %(message)s",
datefmt="%Y-%m-%d %H:%M:%S",
)
# Console handler
if self.config['console_logging']:
if self.config["console_logging"]:
console_handler = logging.StreamHandler()
console_handler.setFormatter(formatter)
self.logger.addHandler(console_handler)
# File handler with rotation and safe encoding for Windows
if self.config['file_logging']:
if self.config["file_logging"]:
log_file = self.log_dir / "prompt_manager.log"
file_handler = logging.handlers.RotatingFileHandler(
log_file,
maxBytes=self.config['max_file_size'],
backupCount=self.config['backup_count'],
encoding='utf-8',
errors='replace' # Replace problematic characters instead of crashing
maxBytes=self.config["max_file_size"],
backupCount=self.config["backup_count"],
encoding="utf-8",
errors="replace", # Replace problematic characters instead of crashing
)
file_handler.setFormatter(formatter)
self.logger.addHandler(file_handler)
# Memory handler for web viewer
memory_handler = MemoryBufferHandler(self)
memory_handler.setFormatter(formatter)
self.logger.addHandler(memory_handler)
# Component-specific loggers
self._setup_component_loggers()
def _setup_component_loggers(self):
"""Set up loggers for specific components.
Creates child loggers for different PromptManager components.
Child loggers inherit handlers from the main logger.
"""
components = [
'prompt_manager.database',
'prompt_manager.api',
'prompt_manager.image_monitor',
'prompt_manager.prompt_tracker',
'prompt_manager.web_ui'
"prompt_manager.database",
"prompt_manager.api",
"prompt_manager.image_monitor",
"prompt_manager.prompt_tracker",
"prompt_manager.web_ui",
]
for component in components:
comp_logger = logging.getLogger(component)
comp_logger.setLevel(getattr(logging, self.config['level']))
comp_logger.setLevel(getattr(logging, self.config["level"]))
# Child loggers inherit handlers from parent
def get_logger(self, name: str = 'prompt_manager') -> logging.Logger:
def get_logger(self, name: str = "prompt_manager") -> logging.Logger:
"""Get a logger instance for a specific component.
Args:
name: Logger name, typically in format 'prompt_manager.component'
Returns:
Configured logger instance
"""
return logging.getLogger(name)
def add_to_buffer(self, record: logging.LogRecord, formatted_message: str):
"""Add a log entry to the memory buffer for web viewer.
Args:
record: The LogRecord object from the logging system
formatted_message: The formatted log message string
"""
with self._buffer_lock:
log_entry = {
'timestamp': datetime.fromtimestamp(record.created).isoformat(),
'level': record.levelname,
'logger': record.name,
'message': record.getMessage(),
'formatted': formatted_message,
'module': record.module,
'filename': record.filename,
'lineno': record.lineno,
'thread': record.thread,
'thread_name': record.threadName if hasattr(record, 'threadName') else '',
'process': record.process
"timestamp": datetime.fromtimestamp(record.created).isoformat(),
"level": record.levelname,
"logger": record.name,
"message": record.getMessage(),
"formatted": formatted_message,
"module": record.module,
"filename": record.filename,
"lineno": record.lineno,
"thread": record.thread,
"thread_name": (
record.threadName if hasattr(record, "threadName") else ""
),
"process": record.process,
}
self._log_buffer.append(log_entry)
def get_recent_logs(self, limit: int = 100, level: Optional[str] = None) -> List[Dict[str, Any]]:
def get_recent_logs(
self, limit: int = 100, level: Optional[str] = None
) -> List[Dict[str, Any]]:
"""Get recent log entries from memory buffer.
Args:
limit: Maximum number of log entries to return
level: Optional log level filter (DEBUG, INFO, WARNING, ERROR, CRITICAL)
Returns:
List of log entry dictionaries, most recent first
"""
@@ -195,14 +199,16 @@ class PromptManagerLogger:
if level:
level_num = getattr(logging, level.upper(), None)
if level_num:
logs = [log for log in logs if getattr(logging, log['level']) >= level_num]
logs = [
log for log in logs if getattr(logging, log["level"]) >= level_num
]
# Return most recent first
return list(reversed(logs[-limit:]))
def get_log_files(self) -> List[Dict[str, Any]]:
"""Get information about available log files.
Returns:
List of dictionaries containing file information:
- filename: Name of the log file
@@ -212,133 +218,140 @@ class PromptManagerLogger:
- is_main: True if this is the main log file
"""
log_files = []
for log_file in self.log_dir.glob("*.log*"):
try:
stat = log_file.stat()
log_files.append({
'filename': log_file.name,
'path': str(log_file),
'size': stat.st_size,
'modified': datetime.fromtimestamp(stat.st_mtime).isoformat(),
'is_main': log_file.name == 'prompt_manager.log'
})
log_files.append(
{
"filename": log_file.name,
"path": str(log_file),
"size": stat.st_size,
"modified": datetime.fromtimestamp(stat.st_mtime).isoformat(),
"is_main": log_file.name == "prompt_manager.log",
}
)
except Exception as e:
self.logger.error(f"Error getting info for log file {log_file}: {e}")
# Sort by modification time, newest first
log_files.sort(key=lambda x: x['modified'], reverse=True)
log_files.sort(key=lambda x: x["modified"], reverse=True)
return log_files
def read_log_file(self, filename: str, lines: int = 100) -> List[str]:
"""Read lines from a specific log file.
Args:
filename: Name of the log file to read (must be in log directory)
lines: Number of lines to read from the end of the file (0 for all)
Returns:
List of line strings from the log file
Raises:
FileNotFoundError: If the log file doesn't exist
ValueError: If the file path is outside the log directory
"""
log_file = self.log_dir / filename
if not log_file.exists() or not log_file.is_file():
raise FileNotFoundError(f"Log file not found: {filename}")
# Security check - ensure file is in log directory
if not str(log_file.resolve()).startswith(str(self.log_dir.resolve())):
raise ValueError("Invalid log file path")
try:
# Use UTF-8 encoding with error handling for Windows compatibility
with open(log_file, 'r', encoding='utf-8', errors='replace') as f:
with open(log_file, "r", encoding="utf-8", errors="replace") as f:
all_lines = f.readlines()
return all_lines[-lines:] if lines > 0 else all_lines
except Exception as e:
self.logger.error(f"Error reading log file {filename}: {e}")
raise
def truncate_logs(self) -> Dict[str, Any]:
"""Truncate all log files.
Clears the main log file and deletes rotated log files.
Also clears the memory buffer.
Returns:
Dictionary with:
- truncated: List of successfully truncated files
- errors: List of error messages for failed operations
"""
results = {
'truncated': [],
'errors': []
}
results = {"truncated": [], "errors": []}
for log_file in self.log_dir.glob("prompt_manager.log*"):
try:
if log_file.name == 'prompt_manager.log':
if log_file.name == "prompt_manager.log":
# For main log file, just clear it with safe encoding
with open(log_file, 'w', encoding='utf-8', errors='replace') as f:
with open(log_file, "w", encoding="utf-8", errors="replace") as f:
f.write("")
else:
# For rotated files, delete them
log_file.unlink()
results['truncated'].append(log_file.name)
results["truncated"].append(log_file.name)
self.logger.info(f"Truncated log file: {log_file.name}")
except Exception as e:
error_msg = f"Failed to truncate {log_file.name}: {str(e)}"
results['errors'].append(error_msg)
results["errors"].append(error_msg)
self.logger.error(error_msg)
# Clear memory buffer
with self._buffer_lock:
self._log_buffer.clear()
return results
def update_config(self, new_config: Dict[str, Any]):
"""Update logging configuration.
Args:
new_config: Dictionary of configuration updates
(level, console_logging, file_logging, etc.)
"""
self.config.update(new_config)
# Reconfigure loggers if level changed
if 'level' in new_config:
level = getattr(logging, new_config['level'].upper(), logging.INFO)
if "level" in new_config:
level = getattr(logging, new_config["level"].upper(), logging.INFO)
self.logger.setLevel(level)
# Update all component loggers
for logger_name in logging.Logger.manager.loggerDict:
if logger_name.startswith('prompt_manager'):
if logger_name.startswith("prompt_manager"):
logger = logging.getLogger(logger_name)
logger.setLevel(level)
# Re-setup loggers if handlers changed
if any(key in new_config for key in ['console_logging', 'file_logging', 'max_file_size', 'backup_count']):
if any(
key in new_config
for key in [
"console_logging",
"file_logging",
"max_file_size",
"backup_count",
]
):
self._setup_loggers()
self.logger.info(f"Updated logging configuration: {new_config}")
def get_config(self) -> Dict[str, Any]:
"""Get current logging configuration.
Returns:
Copy of the current configuration dictionary
"""
return self.config.copy()
def get_log_stats(self) -> Dict[str, Any]:
"""Get logging statistics.
Returns:
Dictionary containing:
- buffer_count: Number of entries in memory buffer
@@ -352,45 +365,45 @@ class PromptManagerLogger:
buffer_count = len(self._log_buffer)
level_counts = {}
for entry in self._log_buffer:
level = entry['level']
level = entry["level"]
level_counts[level] = level_counts.get(level, 0) + 1
log_files = self.get_log_files()
total_log_size = sum(f['size'] for f in log_files)
total_log_size = sum(f["size"] for f in log_files)
return {
'buffer_count': buffer_count,
'level_counts': level_counts,
'log_files_count': len(log_files),
'total_log_size': total_log_size,
'log_directory': str(self.log_dir),
'current_level': self.config['level']
"buffer_count": buffer_count,
"level_counts": level_counts,
"log_files_count": len(log_files),
"total_log_size": total_log_size,
"log_directory": str(self.log_dir),
"current_level": self.config["level"],
}
class MemoryBufferHandler(logging.Handler):
"""Custom logging handler that stores logs in memory for web viewer.
This handler extends the standard logging.Handler to capture log records
and store them in the PromptManagerLogger's memory buffer. This enables
the web interface to display recent log entries without reading from files.
The handler is designed to be failure-safe - if any error occurs during
log processing, it silently continues without disrupting the application.
"""
def __init__(self, logger_manager: PromptManagerLogger):
"""Initialize the memory buffer handler.
Args:
logger_manager: The PromptManagerLogger instance to store logs in
"""
super().__init__()
self.logger_manager = logger_manager
def emit(self, record: logging.LogRecord):
"""Handle a log record by adding it to the memory buffer.
Args:
record: The LogRecord to process and store
"""
@@ -405,9 +418,10 @@ class MemoryBufferHandler(logging.Handler):
# Global logger instance
_logger_manager = None
def get_logger_manager() -> PromptManagerLogger:
"""Get the global logger manager instance.
Returns:
The singleton PromptManagerLogger instance, creating it if necessary
"""
@@ -416,16 +430,18 @@ def get_logger_manager() -> PromptManagerLogger:
_logger_manager = PromptManagerLogger()
return _logger_manager
def get_logger(name: str = 'prompt_manager') -> logging.Logger:
def get_logger(name: str = "prompt_manager") -> logging.Logger:
"""Convenience function to get a logger.
Args:
name: Logger name, defaults to 'prompt_manager'
Returns:
Configured logger instance for the specified name
"""
return get_logger_manager().get_logger(name)
# Initialize logging on import
get_logger_manager()
get_logger_manager()
+128 -108
View File
@@ -13,10 +13,10 @@ The extractor supports:
Typical usage:
from utils.metadata_extractor import ComfyUIMetadataExtractor
extractor = ComfyUIMetadataExtractor()
metadata = extractor.extract_metadata('/path/to/image.png')
if metadata:
workflow = metadata.get('workflow', {})
text_nodes = metadata.get('text_encoder_nodes', [])
@@ -41,33 +41,33 @@ from .logging_config import get_logger
class ComfyUIMetadataExtractor:
"""Extracts ComfyUI metadata from generated images.
This class handles extraction of ComfyUI workflow and generation metadata
from PNG images. It parses the text chunks embedded by ComfyUI and extracts
structured information about workflows, prompts, and generation parameters.
The extractor is designed to be robust and handle various ComfyUI workflow
formats, including custom nodes and different metadata structures.
"""
def __init__(self):
"""Initialize the metadata extractor.
Sets up logging and prepares the extractor for metadata parsing operations.
"""
self.logger = get_logger('prompt_manager.metadata_extractor')
self.logger = get_logger("prompt_manager.metadata_extractor")
def extract_metadata(self, image_path: str) -> Optional[Dict[str, Any]]:
"""
Extract ComfyUI workflow and prompt metadata from an image.
This method opens a PNG image and extracts all available ComfyUI metadata
from the embedded text chunks. It handles JSON parsing, workflow analysis,
and parameter extraction.
Args:
image_path: Path to the PNG image file to analyze
Returns:
Dictionary containing extracted metadata with keys:
- file_info: Basic file information (always present)
@@ -76,7 +76,7 @@ class ComfyUIMetadataExtractor:
- prompt: ComfyUI prompt data (if available)
- Additional fields: parameters, model, sampler, steps, etc. (if present)
Returns None if no ComfyUI metadata is found or extraction fails
Raises:
Exception: Re-raises any exception that occurs during extraction,
with appropriate error logging
@@ -84,40 +84,49 @@ class ComfyUIMetadataExtractor:
try:
with Image.open(image_path) as image:
metadata = {}
# Add basic file information
metadata['file_info'] = self.get_file_info(image_path, image)
metadata["file_info"] = self.get_file_info(image_path, image)
# Extract ComfyUI-specific metadata from PNG text chunks
if hasattr(image, 'text') and image.text:
if hasattr(image, "text") and image.text:
# Look for ComfyUI workflow data
if 'workflow' in image.text:
if "workflow" in image.text:
try:
workflow_data = json.loads(image.text['workflow'])
metadata['workflow'] = workflow_data
workflow_data = json.loads(image.text["workflow"])
metadata["workflow"] = workflow_data
# Extract text encoder nodes from workflow
text_encoder_nodes = self.find_text_encoder_nodes(workflow_data)
text_encoder_nodes = self.find_text_encoder_nodes(
workflow_data
)
if text_encoder_nodes:
metadata['text_encoder_nodes'] = text_encoder_nodes
metadata["text_encoder_nodes"] = text_encoder_nodes
except json.JSONDecodeError as e:
self.logger.warning(f"Failed to parse workflow JSON: {e}")
# Look for prompt data
if 'prompt' in image.text:
if "prompt" in image.text:
try:
prompt_data = json.loads(image.text['prompt'])
metadata['prompt'] = prompt_data
prompt_data = json.loads(image.text["prompt"])
metadata["prompt"] = prompt_data
except json.JSONDecodeError as e:
self.logger.warning(f"Failed to parse prompt JSON: {e}")
# Extract other common metadata fields
metadata_fields = [
'parameters', 'model', 'sampler', 'steps', 'cfg_scale',
'seed', 'scheduler', 'positive', 'negative'
"parameters",
"model",
"sampler",
"steps",
"cfg_scale",
"seed",
"scheduler",
"positive",
"negative",
]
for field in metadata_fields:
if field in image.text:
try:
@@ -126,24 +135,28 @@ class ComfyUIMetadataExtractor:
except (json.JSONDecodeError, TypeError):
# Store as string if not valid JSON
metadata[field] = image.text[field]
return metadata if any(key != 'file_info' for key in metadata.keys()) else None
return (
metadata
if any(key != "file_info" for key in metadata.keys())
else None
)
except Exception as e:
self.logger.error(f"Error extracting metadata from {image_path}: {e}")
return None
def get_file_info(self, image_path: str, image: Image.Image) -> Dict[str, Any]:
"""
Get basic file information.
Extracts fundamental file properties including size, dimensions, format,
and filesystem timestamps.
Args:
image_path: Path to the image file
image: Opened PIL Image object
Returns:
Dictionary containing:
- size: File size in bytes
@@ -156,53 +169,53 @@ class ComfyUIMetadataExtractor:
try:
stat = os.stat(image_path)
return {
'size': stat.st_size,
'dimensions': list(image.size),
'format': image.format,
'mode': image.mode,
'created_time': stat.st_ctime,
'modified_time': stat.st_mtime
"size": stat.st_size,
"dimensions": list(image.size),
"format": image.format,
"mode": image.mode,
"created_time": stat.st_ctime,
"modified_time": stat.st_mtime,
}
except Exception as e:
self.logger.error(f"Error getting file info: {e}")
return {}
def find_text_encoder_nodes(self, workflow_data: Dict) -> list:
"""
Find text encoder nodes in the workflow data.
Searches through the workflow structure to identify nodes that perform
text encoding operations. This includes standard ComfyUI nodes like
CLIPTextEncode as well as custom nodes and extensions.
Args:
workflow_data: ComfyUI workflow data structure (can be in various formats)
Returns:
List of dictionaries representing text encoder nodes, each containing
the node's configuration, inputs, and metadata. For dictionary-based
workflows, adds 'node_id' field to each node.
"""
text_encoder_nodes = []
if not isinstance(workflow_data, dict):
return text_encoder_nodes
# Check different possible workflow structures
nodes_data = None
# Try different keys where nodes might be stored
if 'nodes' in workflow_data:
nodes_data = workflow_data['nodes']
elif 'workflow' in workflow_data and 'nodes' in workflow_data['workflow']:
nodes_data = workflow_data['workflow']['nodes']
if "nodes" in workflow_data:
nodes_data = workflow_data["nodes"]
elif "workflow" in workflow_data and "nodes" in workflow_data["workflow"]:
nodes_data = workflow_data["workflow"]["nodes"]
elif isinstance(workflow_data, dict):
# Sometimes the workflow data is just a flat dict of node IDs
nodes_data = workflow_data
if not nodes_data:
return text_encoder_nodes
# Handle different node data structures
if isinstance(nodes_data, list):
# Nodes as a list
@@ -213,98 +226,98 @@ class ComfyUIMetadataExtractor:
# Nodes as a dictionary (node_id -> node_data)
for node_id, node_data in nodes_data.items():
if self.is_text_encoder_node(node_data):
node_data['node_id'] = node_id
node_data["node_id"] = node_id
text_encoder_nodes.append(node_data)
return text_encoder_nodes
def is_text_encoder_node(self, node_data: Any) -> bool:
"""
Check if a node is a text encoder node.
Determines whether a given node performs text encoding by examining
its type, class, and title for known text encoding patterns. Supports
various ComfyUI node types and custom extensions.
Args:
node_data: Node configuration dictionary to analyze
Returns:
True if the node is identified as a text encoder, False otherwise
"""
if not isinstance(node_data, dict):
return False
# Check for common text encoder node types
text_encoder_types = [
'CLIPTextEncode',
'CLIPTextEncodeSDXL',
'CLIPTextEncodeSDXLRefiner',
'PromptManager', # Our custom node
'BNK_CLIPTextEncoder',
'Text Encoder',
'CLIP Text Encode'
"CLIPTextEncode",
"CLIPTextEncodeSDXL",
"CLIPTextEncodeSDXLRefiner",
"PromptManager", # Our custom node
"BNK_CLIPTextEncoder",
"Text Encoder",
"CLIP Text Encode",
]
# Check node type/class_type
node_type = node_data.get('type') or node_data.get('class_type') or ''
node_type = node_data.get("type") or node_data.get("class_type") or ""
for encoder_type in text_encoder_types:
if encoder_type.lower() in node_type.lower():
return True
# Check node title/name for text encoding keywords
node_title = (node_data.get('title') or node_data.get('name') or '').lower()
text_keywords = ['text', 'prompt', 'encode', 'clip']
node_title = (node_data.get("title") or node_data.get("name") or "").lower()
text_keywords = ["text", "prompt", "encode", "clip"]
if any(keyword in node_title for keyword in text_keywords):
return True
return False
def extract_prompt_text_from_workflow(self, workflow_data: Dict) -> Optional[str]:
"""
Extract the actual prompt text from workflow data.
Searches through text encoder nodes to find and extract the actual
prompt text that was used for generation. Handles various input field
names and data structures.
Args:
workflow_data: ComfyUI workflow data structure
Returns:
The extracted prompt text string, or None if no prompt text is found
"""
text_encoder_nodes = self.find_text_encoder_nodes(workflow_data)
for node in text_encoder_nodes:
# Try different ways to get the prompt text
inputs = node.get('inputs', {})
inputs = node.get("inputs", {})
# Common input field names for prompt text
text_fields = ['text', 'prompt', 'positive', 'conditioning']
text_fields = ["text", "prompt", "positive", "conditioning"]
for field in text_fields:
if field in inputs and inputs[field]:
if isinstance(inputs[field], str):
return inputs[field]
elif isinstance(inputs[field], list) and inputs[field]:
return str(inputs[field][0])
return None
def get_generation_parameters(self, metadata: Dict[str, Any]) -> Dict[str, Any]:
"""
Extract generation parameters from metadata.
Parses the metadata to extract key generation parameters used for
image creation, including sampling settings, model information,
and other configuration values.
Args:
metadata: Full metadata dictionary from extract_metadata()
Returns:
Dictionary containing generation parameters such as:
- steps: Number of sampling steps
@@ -317,42 +330,49 @@ class ComfyUIMetadataExtractor:
- batch_size: Number of images generated
"""
parameters = {}
# Common generation parameters to extract
param_fields = [
'steps', 'cfg_scale', 'sampler', 'scheduler', 'seed',
'model', 'width', 'height', 'batch_size'
"steps",
"cfg_scale",
"sampler",
"scheduler",
"seed",
"model",
"width",
"height",
"batch_size",
]
for field in param_fields:
if field in metadata:
parameters[field] = metadata[field]
# Extract from workflow if available
if 'workflow' in metadata:
workflow_params = self.extract_params_from_workflow(metadata['workflow'])
if "workflow" in metadata:
workflow_params = self.extract_params_from_workflow(metadata["workflow"])
parameters.update(workflow_params)
return parameters
def extract_params_from_workflow(self, workflow_data: Dict) -> Dict[str, Any]:
"""
Extract generation parameters from workflow data.
Analyzes the workflow structure to identify and extract generation
parameters from various node types. This method can be extended
to support specific workflow patterns and custom nodes.
Args:
workflow_data: ComfyUI workflow data structure
Returns:
Dictionary of extracted parameters. Currently returns empty dict
but can be extended based on specific workflow analysis needs.
"""
parameters = {}
# This would need to be customized based on your specific workflow structure
# For now, return empty dict - can be expanded based on specific needs
return parameters
return parameters
+144 -111
View File
@@ -14,11 +14,11 @@ Key features:
Typical usage:
from utils.prompt_tracker import PromptTracker
tracker = PromptTracker(db_manager)
execution_id = tracker.set_current_prompt("beautiful landscape")
# Images generated after this point will be linked to this prompt
Or using the context manager:
with PromptExecutionContext(tracker, "beautiful landscape"):
# Generate images here
@@ -38,6 +38,7 @@ try:
except ImportError:
import sys
import os
current_dir = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
sys.path.insert(0, current_dir)
from utils.logging_config import get_logger
@@ -45,32 +46,32 @@ except ImportError:
class PromptTracker:
"""Thread-safe tracking of current prompt executions.
This class manages active prompt contexts across multiple threads, enabling
the image monitoring system to correctly associate generated images with their
source prompts. It uses both thread-local storage and global tracking to handle
various execution scenarios.
The tracker automatically cleans up expired prompts and provides fallback
mechanisms when prompts are accessed from different threads (e.g., image
monitoring running in a separate thread from prompt execution).
Attributes:
active_prompts: Global dictionary of active prompt contexts
cleanup_interval: Seconds between cleanup operations (default: 300)
prompt_timeout: Seconds before a prompt expires (default: 600)
"""
def __init__(self, db_manager):
"""
Initialize the prompt tracker.
Args:
db_manager: Database manager instance for prompt operations
"""
self.logger = get_logger('prompt_manager.prompt_tracker')
self.logger = get_logger("prompt_manager.prompt_tracker")
self.logger.debug("Initializing PromptTracker")
self.db_manager = db_manager
self._local = threading.local()
self.active_prompts = {} # Global tracking for multiple threads
@@ -79,51 +80,59 @@ class PromptTracker:
# Read from GalleryConfig if available, otherwise use defaults
try:
from ..py.config import GalleryConfig
self.cleanup_interval = GalleryConfig.CLEANUP_INTERVAL
self.prompt_timeout = GalleryConfig.PROMPT_TIMEOUT
except Exception:
self.cleanup_interval = 300 # 5 minutes
self.prompt_timeout = 600 # 10 minutes
self.cleanup_interval = 300 # 5 minutes
self.prompt_timeout = 600 # 10 minutes
# Start cleanup thread
self.cleanup_thread = threading.Thread(target=self._cleanup_expired_prompts, daemon=True)
self.cleanup_thread = threading.Thread(
target=self._cleanup_expired_prompts, daemon=True
)
self.cleanup_thread.start()
self.logger.debug("PromptTracker initialization completed")
def set_current_prompt(self, prompt_text: str, additional_data: Optional[Dict[str, Any]] = None) -> str:
def set_current_prompt(
self, prompt_text: str, additional_data: Optional[Dict[str, Any]] = None
) -> str:
"""
Set the current prompt for this thread for image tracking.
This method establishes the prompt context that will be used to link any
subsequently generated images. The prompt is stored in both thread-local
storage and global tracking to handle cross-thread access scenarios.
Note: Prompt saving is handled by PromptManager to avoid duplicates.
Args:
prompt_text: The prompt text being executed
additional_data: Additional prompt metadata, should include prompt_id from PromptManager
Returns:
Unique execution ID for this prompt execution
"""
execution_id = self.generate_execution_id()
# Get prompt_id from additional_data (passed from PromptManager)
# PromptTracker should NOT save prompts to avoid duplicates
prompt_id = additional_data.get('prompt_id') if additional_data else None
prompt_id = additional_data.get("prompt_id") if additional_data else None
if not prompt_id:
# Fallback: try to find existing prompt using consistent hash calculation
try:
# Use consistent hash calculation (same as PromptManager)
import hashlib
normalized_text = prompt_text.strip().lower()
prompt_hash = hashlib.sha256(normalized_text.encode('utf-8')).hexdigest()
prompt_hash = hashlib.sha256(
normalized_text.encode("utf-8")
).hexdigest()
existing_prompt = self.db_manager.get_prompt_by_hash(prompt_hash)
if existing_prompt:
prompt_id = existing_prompt['id']
prompt_id = existing_prompt["id"]
self.logger.debug(f"Found existing prompt ID: {prompt_id}")
else:
# Generate a temporary ID for tracking (prompt should be saved by PromptManager)
@@ -132,36 +141,38 @@ class PromptTracker:
except Exception as e:
self.logger.error(f"Error finding prompt: {e}")
prompt_id = f"temp_{int(time.time())}"
# Create execution context
execution_context = {
'id': prompt_id,
'execution_id': execution_id,
'text': prompt_text,
'timestamp': time.time(),
'thread_id': threading.current_thread().ident,
'additional_data': additional_data or {}
"id": prompt_id,
"execution_id": execution_id,
"text": prompt_text,
"timestamp": time.time(),
"thread_id": threading.current_thread().ident,
"additional_data": additional_data or {},
}
# Store in thread-local storage
self._local.current_prompt = execution_context
# Store in global tracking with thread safety
with self.lock:
self.active_prompts[execution_id] = execution_context
self.logger.debug(f"Set current prompt: {execution_id} -> {prompt_text[:50]}... (thread: {threading.current_thread().ident})")
self.logger.debug(
f"Set current prompt: {execution_id} -> {prompt_text[:50]}... (thread: {threading.current_thread().ident})"
)
self.logger.debug(f"Active prompts count: {len(self.active_prompts)}")
return execution_id
def get_current_prompt(self) -> Optional[Dict[str, Any]]:
"""
Get the current prompt context for this thread.
Attempts to retrieve the current prompt context, first from thread-local
storage, then from global tracking as a fallback. This enables image
monitoring (which may run in a different thread) to access prompt context.
Returns:
Dictionary containing prompt context with keys:
- id: Prompt ID in database
@@ -172,94 +183,112 @@ class PromptTracker:
- additional_data: Any additional metadata
Returns None if no valid prompt context is found
"""
current = getattr(self._local, 'current_prompt', None)
current = getattr(self._local, "current_prompt", None)
if current:
# Check if prompt hasn't expired
if time.time() - current['timestamp'] < self.prompt_timeout:
if time.time() - current["timestamp"] < self.prompt_timeout:
return current
else:
self.logger.debug(f"Current prompt expired: {current['execution_id']}")
self.clear_current_prompt()
# Fallback: try to find recent prompt from global tracking
# This is crucial for image monitoring which runs in different threads
recent_prompt = self._find_recent_prompt()
if recent_prompt:
self.logger.debug(f"Using recent prompt from global tracking: {recent_prompt['execution_id']}")
self.logger.debug(
f"Using recent prompt from global tracking: {recent_prompt['execution_id']}"
)
return recent_prompt
self.logger.debug(f"No prompt context found (thread: {threading.current_thread().ident})")
self.logger.debug(
f"No prompt context found (thread: {threading.current_thread().ident})"
)
return None
def _find_recent_prompt(self) -> Optional[Dict[str, Any]]:
"""Find the most recent prompt that's still valid.
Searches the global prompt tracking for the most recently set prompt
that hasn't expired. This is used as a fallback when thread-local
storage doesn't contain a current prompt.
Returns:
Most recent valid prompt context, or None if no valid prompts found
"""
with self.lock:
current_time = time.time()
self.logger.debug(f"Searching for recent prompt among {len(self.active_prompts)} active prompts")
self.logger.debug(
f"Searching for recent prompt among {len(self.active_prompts)} active prompts"
)
recent_prompts = []
for exec_id, prompt in self.active_prompts.items():
age_seconds = current_time - prompt['timestamp']
self.logger.debug(f"Prompt {exec_id}: age={age_seconds:.1f}s, timeout={self.prompt_timeout}s")
age_seconds = current_time - prompt["timestamp"]
self.logger.debug(
f"Prompt {exec_id}: age={age_seconds:.1f}s, timeout={self.prompt_timeout}s"
)
if age_seconds < self.prompt_timeout:
recent_prompts.append(prompt)
else:
self.logger.debug(f"Prompt {exec_id} is expired")
if recent_prompts:
# Return the most recent one
most_recent = max(recent_prompts, key=lambda p: p['timestamp'])
most_recent = max(recent_prompts, key=lambda p: p["timestamp"])
self.logger.debug(f"Found recent prompt: {most_recent['execution_id']}")
return most_recent
else:
self.logger.debug(f"No recent prompts found")
return None
def clear_current_prompt(self):
"""Clear the current prompt context for this thread.
Removes the prompt context from both thread-local storage and global
tracking. This should be called when prompt execution is complete,
though the system also handles automatic cleanup via timeouts.
"""
current = getattr(self._local, 'current_prompt', None)
current = getattr(self._local, "current_prompt", None)
if current:
execution_id = current['execution_id']
self.logger.debug(f"Clearing current prompt: {execution_id} (thread: {threading.current_thread().ident})")
execution_id = current["execution_id"]
self.logger.debug(
f"Clearing current prompt: {execution_id} (thread: {threading.current_thread().ident})"
)
# Clear from thread-local storage
self._local.current_prompt = None
# Remove from global tracking
with self.lock:
removed = self.active_prompts.pop(execution_id, None)
if removed:
self.logger.debug(f"Removed prompt from global tracking: {execution_id}")
self.logger.debug(f"Remaining active prompts: {len(self.active_prompts)}")
self.logger.debug(
f"Removed prompt from global tracking: {execution_id}"
)
self.logger.debug(
f"Remaining active prompts: {len(self.active_prompts)}"
)
else:
self.logger.debug(f"Prompt {execution_id} was not in global tracking")
self.logger.debug(
f"Prompt {execution_id} was not in global tracking"
)
else:
self.logger.debug(f"No current prompt to clear (thread: {threading.current_thread().ident})")
self.logger.debug(
f"No current prompt to clear (thread: {threading.current_thread().ident})"
)
def extend_prompt_timeout(self, execution_id: str, additional_seconds: int = 60):
"""
Extend the timeout for a specific prompt execution.
This is useful for long-running image generation processes where images
may continue to be generated well after the initial prompt execution.
Updates the timestamp to prevent the prompt from being cleaned up.
Args:
execution_id: The execution ID to extend
additional_seconds: Additional seconds to add to timeout (currently unused,
@@ -267,22 +296,22 @@ class PromptTracker:
"""
with self.lock:
if execution_id in self.active_prompts:
self.active_prompts[execution_id]['timestamp'] = time.time()
self.active_prompts[execution_id]["timestamp"] = time.time()
self.logger.debug(f"Extended timeout for prompt: {execution_id}")
def generate_execution_id(self) -> str:
"""Generate a unique execution ID.
Creates a unique identifier for this prompt execution using UUID and timestamp.
Returns:
Unique execution ID in format 'exec_{uuid8}_{timestamp}'
"""
return f"exec_{uuid.uuid4().hex[:8]}_{int(time.time())}"
def _cleanup_expired_prompts(self):
"""Background thread to clean up expired prompts.
Runs continuously as a daemon thread, periodically removing expired
prompt contexts from the global tracking dictionary. This prevents
memory leaks from accumulating old prompt data.
@@ -291,59 +320,62 @@ class PromptTracker:
try:
current_time = time.time()
expired_ids = []
with self.lock:
for exec_id, prompt_data in self.active_prompts.items():
if current_time - prompt_data['timestamp'] > self.prompt_timeout:
if (
current_time - prompt_data["timestamp"]
> self.prompt_timeout
):
expired_ids.append(exec_id)
for exec_id in expired_ids:
self.active_prompts.pop(exec_id, None)
if expired_ids:
self.logger.debug(f"Cleaned up {len(expired_ids)} expired prompts")
time.sleep(self.cleanup_interval)
except Exception as e:
self.logger.error(f"Error in cleanup thread: {e}")
time.sleep(60) # Wait a minute before retrying
def get_active_prompts(self) -> Dict[str, Dict[str, Any]]:
"""
Get all currently active prompts.
Returns:
Dictionary mapping execution IDs to prompt context dictionaries.
This is a copy of the internal tracking dictionary.
"""
with self.lock:
return self.active_prompts.copy()
def clear_all_active_prompts(self) -> int:
"""
Clear all active prompts.
Removes all prompt contexts from both global and thread-local storage.
This is useful before batch operations or when resetting the tracking state.
Returns:
Number of prompts that were cleared from global tracking
"""
with self.lock:
cleared_count = len(self.active_prompts)
self.active_prompts.clear()
# Clear thread-local storage as well
self._local.current_prompt = None
self.logger.debug(f"Cleared {cleared_count} active prompts")
return cleared_count
def get_status(self) -> Dict[str, Any]:
"""
Get tracker status information.
Returns:
Dictionary containing:
- active_prompts_count: Number of currently active prompts
@@ -355,28 +387,30 @@ class PromptTracker:
"""
with self.lock:
active_count = len(self.active_prompts)
current_prompt = self.get_current_prompt()
return {
'active_prompts_count': active_count,
'current_prompt_id': current_prompt['id'] if current_prompt else None,
'current_execution_id': current_prompt['execution_id'] if current_prompt else None,
'thread_id': threading.current_thread().ident,
'prompt_timeout': self.prompt_timeout,
'cleanup_interval': self.cleanup_interval
"active_prompts_count": active_count,
"current_prompt_id": current_prompt["id"] if current_prompt else None,
"current_execution_id": (
current_prompt["execution_id"] if current_prompt else None
),
"thread_id": threading.current_thread().ident,
"prompt_timeout": self.prompt_timeout,
"cleanup_interval": self.cleanup_interval,
}
class PromptExecutionContext:
"""Context manager for prompt executions.
This context manager provides automatic prompt lifecycle management,
ensuring prompt context is properly set and cleaned up. While the current
implementation doesn't automatically clear prompts on exit (to allow for
images generated after prompt execution), it provides a clean interface
for prompt management.
Example:
with PromptExecutionContext(tracker, "beautiful landscape") as exec_id:
# Generate images here
@@ -384,11 +418,11 @@ class PromptExecutionContext:
pass
# Prompt context remains active for timeout period
"""
def __init__(self, prompt_tracker: PromptTracker, prompt_text: str, **kwargs):
"""
Initialize execution context.
Args:
prompt_tracker: PromptTracker instance to use for tracking
prompt_text: The prompt text to track
@@ -398,21 +432,20 @@ class PromptExecutionContext:
self.prompt_text = prompt_text
self.additional_data = kwargs
self.execution_id = None
def __enter__(self):
"""Enter the execution context.
Sets the current prompt in the tracker and returns the execution ID.
Returns:
Unique execution ID for this prompt
"""
self.execution_id = self.prompt_tracker.set_current_prompt(
self.prompt_text,
self.additional_data
self.prompt_text, self.additional_data
)
return self.execution_id
def __exit__(self, exc_type, exc_val, exc_tb):
"""Exit the execution context.
@@ -452,4 +485,4 @@ def get_prompt_tracker(db_manager) -> PromptTracker:
with _tracker_lock:
if _tracker_instance is None:
_tracker_instance = PromptTracker(db_manager)
return _tracker_instance
return _tracker_instance
+63 -63
View File
@@ -20,7 +20,7 @@ All validation functions follow a consistent pattern:
Typical usage:
from utils.validators import validate_prompt_text, sanitize_input
try:
validate_prompt_text(user_input)
clean_text = sanitize_input(user_input)
@@ -37,16 +37,16 @@ from typing import List, Optional, Union
def validate_prompt_text(text: str) -> bool:
"""
Validate prompt text input.
Ensures the prompt text is a valid string with reasonable length limits.
Empty or whitespace-only strings are rejected.
Args:
text: The prompt text to validate
Returns:
True if the text passes validation
Raises:
ValueError: If text is invalid with descriptive message including:
- Not a string type
@@ -55,28 +55,28 @@ def validate_prompt_text(text: str) -> bool:
"""
if not isinstance(text, str):
raise ValueError("Prompt text must be a string")
if not text or not text.strip():
raise ValueError("Prompt text cannot be empty")
if len(text.strip()) > 10000: # Reasonable limit for prompt length
raise ValueError("Prompt text is too long (maximum 10,000 characters)")
return True
def validate_rating(rating: Optional[int]) -> bool:
"""
Validate rating input.
Validates rating values on a 1-5 scale, with None allowed for unrated prompts.
Args:
rating: The rating to validate (1-5 scale or None for no rating)
Returns:
True if the rating is valid
Raises:
ValueError: If rating is invalid:
- Not an integer (when not None)
@@ -84,29 +84,29 @@ def validate_rating(rating: Optional[int]) -> bool:
"""
if rating is None:
return True
if not isinstance(rating, int):
raise ValueError("Rating must be an integer")
if rating < 1 or rating > 5:
raise ValueError("Rating must be between 1 and 5")
return True
def validate_tags(tags: Union[str, List[str], None]) -> bool:
"""
Validate tags input.
Accepts tags as comma-separated string, list of strings, or None.
Validates each tag for length and character restrictions.
Args:
tags: Tags as comma-separated string, list of strings, or None
Returns:
True if all tags are valid
Raises:
ValueError: If tags are invalid:
- Wrong input type (not string, list, or None)
@@ -117,48 +117,48 @@ def validate_tags(tags: Union[str, List[str], None]) -> bool:
"""
if tags is None:
return True
if isinstance(tags, str):
# Parse comma-separated tags
tag_list = [tag.strip() for tag in tags.split(',') if tag.strip()]
tag_list = [tag.strip() for tag in tags.split(",") if tag.strip()]
tags = tag_list
if not isinstance(tags, list):
raise ValueError("Tags must be a string, list, or None")
for tag in tags:
if not isinstance(tag, str):
raise ValueError("All tags must be strings")
if not tag.strip():
raise ValueError("Tags cannot be empty")
if len(tag.strip()) > 50:
raise ValueError("Individual tags cannot exceed 50 characters")
# Reject control characters and null bytes
if re.search(r'[\x00-\x1f]', tag.strip()):
if re.search(r"[\x00-\x1f]", tag.strip()):
raise ValueError(f"Tag '{tag}' contains invalid control characters")
if len(tags) > 20: # Reasonable limit
raise ValueError("Maximum 20 tags allowed")
return True
def validate_category(category: Optional[str]) -> bool:
"""
Validate category input.
Validates optional category strings with reasonable length limits
and character restrictions.
Args:
category: The category string to validate (None allowed for no category)
Returns:
True if the category is valid or None
Raises:
ValueError: If category is invalid:
- Not a string type (when not None)
@@ -167,37 +167,37 @@ def validate_category(category: Optional[str]) -> bool:
"""
if category is None:
return True
if not isinstance(category, str):
raise ValueError("Category must be a string")
category = category.strip()
if not category:
return True # Empty category is valid (same as None)
if len(category) > 100:
raise ValueError("Category cannot exceed 100 characters")
# Reject control characters and null bytes
if re.search(r'[\x00-\x1f]', category):
if re.search(r"[\x00-\x1f]", category):
raise ValueError("Category contains invalid control characters")
return True
def validate_workflow_name(workflow_name: Optional[str]) -> bool:
"""
Validate workflow name input.
Validates optional workflow name strings with generous length limits
to accommodate descriptive workflow names.
Args:
workflow_name: The workflow name to validate (None allowed for no workflow)
Returns:
True if the workflow name is valid or None
Raises:
ValueError: If workflow name is invalid:
- Not a string type (when not None)
@@ -205,31 +205,31 @@ def validate_workflow_name(workflow_name: Optional[str]) -> bool:
"""
if workflow_name is None:
return True
if not isinstance(workflow_name, str):
raise ValueError("Workflow name must be a string")
workflow_name = workflow_name.strip()
if not workflow_name:
return True # Empty workflow name is valid
if len(workflow_name) > 200:
raise ValueError("Workflow name cannot exceed 200 characters")
return True
def sanitize_input(text: str) -> str:
"""
Sanitize text input by removing potentially harmful content.
Cleans text input by removing control characters, normalizing whitespace,
and limiting excessive empty lines. Preserves the semantic content while
ensuring safe storage and display.
Args:
text: The text string to sanitize
Returns:
Sanitized text string with:
- Null bytes and control characters removed
@@ -240,18 +240,18 @@ def sanitize_input(text: str) -> str:
"""
if not isinstance(text, str):
return ""
# Remove null bytes and other control characters
sanitized = text.replace('\x00', '').replace('\r\n', '\n').replace('\r', '\n')
sanitized = text.replace("\x00", "").replace("\r\n", "\n").replace("\r", "\n")
# Strip excessive whitespace but preserve single newlines
lines = sanitized.split('\n')
lines = sanitized.split("\n")
sanitized_lines = [line.strip() for line in lines]
# Remove excessive empty lines (keep max 2 consecutive)
result_lines = []
empty_count = 0
for line in sanitized_lines:
if not line:
empty_count += 1
@@ -260,32 +260,32 @@ def sanitize_input(text: str) -> str:
else:
empty_count = 0
result_lines.append(line)
return '\n'.join(result_lines).strip()
return "\n".join(result_lines).strip()
def parse_tags_string(tags_string: str) -> List[str]:
"""
Parse a comma-separated tags string into a clean list.
Converts comma-separated tag strings into a clean, deduplicated list
of tags. Each tag is sanitized and trimmed.
Args:
tags_string: Comma-separated tags string (e.g., "tag1, tag2, tag3")
Returns:
List of unique, cleaned tag strings. Empty input returns empty list.
Limited to maximum 20 tags to prevent abuse.
"""
if not tags_string or not isinstance(tags_string, str):
return []
# Split by comma and clean each tag
tags = []
for tag in tags_string.split(','):
for tag in tags_string.split(","):
clean_tag = sanitize_input(tag).strip()
if clean_tag and clean_tag not in tags: # Avoid duplicates
tags.append(clean_tag)
return tags[:20] # Limit to 20 tags
return tags[:20] # Limit to 20 tags