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:
+6
-1
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user