refactor: split api.py into domain modules and extract frontend JS
Phase 4 structural refactoring: - Split monolithic py/api.py (5.3k lines) into domain mixins: prompts.py, images.py, admin.py, logging_routes.py, autotag_routes.py - Extract inline JS from admin.html into web/js/admin.js - Extract inline JS from gallery.html into web/js/gallery.js - Add gzip compression middleware to API responses (Phase 5.4) - Standardize API error envelope with success: false (Phase 5.3) - Add filmstrip image enrichment to prompt list responses (Phase 5.2)
This commit is contained in:
@@ -0,0 +1,735 @@
|
||||
"""REST API module for ComfyUI PromptManager.
|
||||
|
||||
Split into domain-specific mixins for maintainability:
|
||||
- PromptRoutesMixin: prompt CRUD, tags, categories, bulk ops, export
|
||||
- ImageRoutesMixin: gallery, thumbnails, image serving and linking
|
||||
- AdminRoutesMixin: duplicates, stats, settings, diagnostics, maintenance, backup, scan
|
||||
- LoggingRoutesMixin: log management endpoints
|
||||
- AutotagRoutesMixin: auto-tagging model management and tagging
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import datetime
|
||||
import functools
|
||||
import gzip as gzip_module
|
||||
import json
|
||||
import os
|
||||
from pathlib import Path
|
||||
|
||||
from aiohttp import web
|
||||
from PIL import Image
|
||||
|
||||
from .prompts import PromptRoutesMixin
|
||||
from .images import ImageRoutesMixin
|
||||
from .admin import AdminRoutesMixin
|
||||
from .logging_routes import LoggingRoutesMixin
|
||||
from .autotag_routes import AutotagRoutesMixin
|
||||
|
||||
try:
|
||||
from ...database.operations import PromptDatabase
|
||||
from ...utils.logging_config import get_logger
|
||||
except ImportError:
|
||||
import sys
|
||||
|
||||
sys.path.insert(
|
||||
0, os.path.dirname(os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
|
||||
)
|
||||
from database.operations import PromptDatabase
|
||||
from utils.logging_config import get_logger
|
||||
|
||||
|
||||
def _get_project_root():
|
||||
"""Get the project root directory (3 levels up from py/api/__init__.py)."""
|
||||
return os.path.dirname(
|
||||
os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
|
||||
)
|
||||
|
||||
|
||||
# ── Gzip compression middleware ────────────────────────────────────────
|
||||
_GZIP_MIN_SIZE = 1024 # Only compress bodies larger than 1 KB
|
||||
_GZIP_TYPES = frozenset((
|
||||
'application/json',
|
||||
'text/html',
|
||||
'text/css',
|
||||
'application/javascript',
|
||||
'text/plain',
|
||||
))
|
||||
_gzip_registered = False
|
||||
|
||||
|
||||
@web.middleware
|
||||
async def _gzip_middleware(request, handler):
|
||||
"""Compress PromptManager responses when the client accepts gzip."""
|
||||
response = await handler(request)
|
||||
|
||||
if not request.path.startswith('/prompt_manager/'):
|
||||
return response
|
||||
|
||||
# Only handle regular Response objects (not StreamResponse/WebSocket)
|
||||
if not isinstance(response, web.Response):
|
||||
return response
|
||||
|
||||
if 'gzip' not in request.headers.get('Accept-Encoding', ''):
|
||||
return response
|
||||
|
||||
if 'Content-Encoding' in response.headers:
|
||||
return response
|
||||
|
||||
body = response.body
|
||||
if body is None or len(body) < _GZIP_MIN_SIZE:
|
||||
return response
|
||||
|
||||
content_type = response.content_type or ''
|
||||
if not any(ct in content_type for ct in _GZIP_TYPES):
|
||||
return response
|
||||
|
||||
compressed = gzip_module.compress(body, compresslevel=6)
|
||||
if len(compressed) >= len(body):
|
||||
return response
|
||||
|
||||
response.body = compressed
|
||||
response.headers['Content-Encoding'] = 'gzip'
|
||||
response.headers['Vary'] = 'Accept-Encoding'
|
||||
return response
|
||||
|
||||
|
||||
class PromptManagerAPI(
|
||||
PromptRoutesMixin,
|
||||
ImageRoutesMixin,
|
||||
AdminRoutesMixin,
|
||||
LoggingRoutesMixin,
|
||||
AutotagRoutesMixin,
|
||||
):
|
||||
"""REST API handler for PromptManager operations and web interface.
|
||||
|
||||
This class provides comprehensive REST API endpoints for managing prompts,
|
||||
images, and system operations. It handles database interactions, file
|
||||
operations, image processing, and web UI serving.
|
||||
|
||||
The API is designed to integrate seamlessly with ComfyUI's aiohttp server
|
||||
and provides both JSON API endpoints and static file serving for the
|
||||
web interface.
|
||||
|
||||
Attributes:
|
||||
logger: Configured logger instance for API operations
|
||||
db (PromptDatabase): Database connection and operations handler
|
||||
"""
|
||||
|
||||
def __init__(self):
|
||||
"""Initialize the PromptManager API with database connection and cleanup."""
|
||||
self.logger = get_logger('prompt_manager.api')
|
||||
self.logger.info("Initializing PromptManager API")
|
||||
|
||||
self.db = PromptDatabase()
|
||||
self._cached_output_dir = None # Lazy-cached by _find_comfyui_output_dir()
|
||||
self._html_cache = {} # Cached HTML file contents keyed by path
|
||||
self._gallery_cache = None # Cached gallery file listing (Fix 2.4)
|
||||
self._gallery_cache_time = 0 # Timestamp of last cache fill
|
||||
self._gallery_cache_ttl = 30 # Cache TTL in seconds
|
||||
|
||||
# Run cleanup on initialization to remove any existing duplicates
|
||||
try:
|
||||
removed = self.db.cleanup_duplicates()
|
||||
if removed > 0:
|
||||
self.logger.info(f"Startup cleanup: removed {removed} duplicate prompts")
|
||||
except Exception as e:
|
||||
self.logger.error(f"Startup cleanup failed: {e}")
|
||||
|
||||
self.logger.info("PromptManager API initialization completed")
|
||||
|
||||
async def _run_in_executor(self, func, *args, **kwargs):
|
||||
"""Run a blocking function in the default thread pool executor.
|
||||
|
||||
Prevents synchronous database calls, file I/O, and PIL operations
|
||||
from blocking the aiohttp event loop.
|
||||
"""
|
||||
loop = asyncio.get_event_loop()
|
||||
if kwargs:
|
||||
call = functools.partial(func, *args, **kwargs)
|
||||
return await loop.run_in_executor(None, call)
|
||||
return await loop.run_in_executor(None, func, *args)
|
||||
|
||||
def invalidate_gallery_cache(self):
|
||||
"""Invalidate the cached gallery file listing.
|
||||
|
||||
Called by the image monitor when new files are detected.
|
||||
"""
|
||||
self._gallery_cache = None
|
||||
self._gallery_cache_time = 0
|
||||
|
||||
def add_routes(self, routes):
|
||||
"""Register all API routes with the ComfyUI server.
|
||||
|
||||
Args:
|
||||
routes: aiohttp RouteTableDef object from ComfyUI server instance.
|
||||
"""
|
||||
|
||||
# Test route to verify registration works
|
||||
@routes.get("/prompt_manager/test")
|
||||
async def test_route(request):
|
||||
return web.json_response(
|
||||
{
|
||||
"success": True,
|
||||
"message": "PromptManager API is working!",
|
||||
"timestamp": str(datetime.datetime.now()),
|
||||
}
|
||||
)
|
||||
|
||||
# ── Web UI serving routes ─────────────────────────────────────
|
||||
|
||||
@routes.get("/prompt_manager/web")
|
||||
async def serve_web_ui(request):
|
||||
try:
|
||||
html_path = os.path.join(_get_project_root(), "web", "index.html")
|
||||
|
||||
if os.path.exists(html_path):
|
||||
with open(html_path, "r", encoding="utf-8") as f:
|
||||
html_content = f.read()
|
||||
|
||||
return web.Response(
|
||||
text=html_content, content_type="text/html", charset="utf-8"
|
||||
)
|
||||
else:
|
||||
return web.Response(
|
||||
text="<h1>Web UI not found</h1><p>HTML file not located at expected path.</p>",
|
||||
content_type="text/html",
|
||||
status=404,
|
||||
)
|
||||
|
||||
except Exception as e:
|
||||
return web.Response(
|
||||
text=f"<h1>Error</h1><p>Failed to load web UI: {str(e)}</p>",
|
||||
content_type="text/html",
|
||||
status=500,
|
||||
)
|
||||
|
||||
@routes.get("/prompt_manager/gallery.html")
|
||||
async def serve_gallery_ui(request):
|
||||
try:
|
||||
html_path = os.path.join(
|
||||
_get_project_root(), "web", "metadata.html",
|
||||
)
|
||||
|
||||
if html_path not in self._html_cache:
|
||||
if os.path.exists(html_path):
|
||||
with open(html_path, "r", encoding="utf-8") as f:
|
||||
self._html_cache[html_path] = f.read()
|
||||
else:
|
||||
return web.Response(
|
||||
text="<h1>Gallery not found</h1><p>gallery.html file not located at expected path.</p>",
|
||||
content_type="text/html",
|
||||
status=404,
|
||||
)
|
||||
|
||||
return web.Response(
|
||||
text=self._html_cache[html_path], content_type="text/html", charset="utf-8"
|
||||
)
|
||||
|
||||
except Exception as e:
|
||||
return web.Response(
|
||||
text=f"<h1>Error</h1><p>Failed to load gallery: {str(e)}</p>",
|
||||
content_type="text/html",
|
||||
status=500,
|
||||
)
|
||||
|
||||
@routes.get("/prompt_manager/admin")
|
||||
async def serve_admin_ui(request):
|
||||
try:
|
||||
html_path = os.path.join(
|
||||
_get_project_root(), "web", "admin.html",
|
||||
)
|
||||
|
||||
if html_path not in self._html_cache:
|
||||
if os.path.exists(html_path):
|
||||
with open(html_path, "r", encoding="utf-8") as f:
|
||||
self._html_cache[html_path] = f.read()
|
||||
else:
|
||||
return web.Response(
|
||||
text="<h1>Admin UI not found</h1>",
|
||||
content_type="text/html",
|
||||
status=404,
|
||||
)
|
||||
|
||||
return web.Response(
|
||||
text=self._html_cache[html_path], content_type="text/html", charset="utf-8"
|
||||
)
|
||||
|
||||
except Exception as e:
|
||||
return web.Response(
|
||||
text=f"<h1>Error</h1><p>Failed to load admin UI: {str(e)}</p>",
|
||||
content_type="text/html",
|
||||
status=500,
|
||||
)
|
||||
|
||||
@routes.get("/prompt_manager/gallery")
|
||||
async def serve_gallery_admin_ui(request):
|
||||
try:
|
||||
html_path = os.path.join(
|
||||
_get_project_root(), "web", "gallery.html",
|
||||
)
|
||||
|
||||
if html_path not in self._html_cache:
|
||||
if os.path.exists(html_path):
|
||||
with open(html_path, "r", encoding="utf-8") as f:
|
||||
self._html_cache[html_path] = f.read()
|
||||
else:
|
||||
return web.Response(
|
||||
text="<h1>Gallery not found</h1><p>gallery.html file not located at expected path.</p>",
|
||||
content_type="text/html",
|
||||
status=404,
|
||||
)
|
||||
|
||||
return web.Response(
|
||||
text=self._html_cache[html_path], content_type="text/html", charset="utf-8"
|
||||
)
|
||||
|
||||
except Exception as e:
|
||||
return web.Response(
|
||||
text=f"<h1>Error</h1><p>Failed to load gallery: {str(e)}</p>",
|
||||
content_type="text/html",
|
||||
status=500,
|
||||
)
|
||||
|
||||
# ── Static file serving ───────────────────────────────────────
|
||||
|
||||
@routes.get("/prompt_manager/lib/{filepath:.*}")
|
||||
async def serve_lib_static(request):
|
||||
"""Serve static library files (JS, CSS) from web/lib directory."""
|
||||
MIME_TYPES = {
|
||||
".js": "application/javascript",
|
||||
".css": "text/css",
|
||||
".json": "application/json",
|
||||
".map": "application/json",
|
||||
}
|
||||
|
||||
filepath = request.match_info.get("filepath", "")
|
||||
|
||||
# Security: prevent directory traversal
|
||||
if ".." in filepath or filepath.startswith("/"):
|
||||
return web.Response(text="Forbidden", status=403)
|
||||
|
||||
file_path = os.path.join(_get_project_root(), "web", "lib", filepath)
|
||||
|
||||
if not os.path.exists(file_path) or not os.path.isfile(file_path):
|
||||
return web.Response(text=f"Not Found: {filepath}", status=404)
|
||||
|
||||
ext = os.path.splitext(file_path)[1].lower()
|
||||
content_type = MIME_TYPES.get(ext, "application/octet-stream")
|
||||
|
||||
with open(file_path, "rb") as f:
|
||||
content = f.read()
|
||||
|
||||
return web.Response(body=content, content_type=content_type)
|
||||
|
||||
@routes.get("/prompt_manager/js/{filepath:.*}")
|
||||
async def serve_js_static(request):
|
||||
"""Serve static JavaScript files from web/js directory."""
|
||||
MIME_TYPES = {
|
||||
".js": "application/javascript",
|
||||
".css": "text/css",
|
||||
".json": "application/json",
|
||||
".map": "application/json",
|
||||
}
|
||||
|
||||
filepath = request.match_info.get("filepath", "")
|
||||
|
||||
if ".." in filepath or filepath.startswith("/"):
|
||||
return web.Response(text="Forbidden", status=403)
|
||||
|
||||
file_path = os.path.join(_get_project_root(), "web", "js", filepath)
|
||||
|
||||
if not os.path.exists(file_path) or not os.path.isfile(file_path):
|
||||
return web.Response(text=f"Not Found: {filepath}", status=404)
|
||||
|
||||
ext = os.path.splitext(file_path)[1].lower()
|
||||
content_type = MIME_TYPES.get(ext, "application/octet-stream")
|
||||
|
||||
with open(file_path, "rb") as f:
|
||||
content = f.read()
|
||||
|
||||
return web.Response(body=content, content_type=content_type)
|
||||
|
||||
# ── Register domain-specific routes from mixins ───────────────
|
||||
|
||||
self._register_prompt_routes(routes)
|
||||
self._register_image_routes(routes)
|
||||
self._register_admin_routes(routes)
|
||||
self._register_logging_routes(routes)
|
||||
self._register_autotag_routes(routes)
|
||||
|
||||
# Register gzip compression middleware (once)
|
||||
global _gzip_registered
|
||||
if not _gzip_registered:
|
||||
try:
|
||||
from ..config import server_instance
|
||||
server_instance.app.middlewares.append(_gzip_middleware)
|
||||
_gzip_registered = True
|
||||
self.logger.info("Gzip compression middleware registered")
|
||||
except Exception as e:
|
||||
self.logger.warning(f"Could not register gzip middleware: {e}")
|
||||
|
||||
self.logger.info("All routes registered with decorator pattern")
|
||||
|
||||
# ── Shared utilities used by multiple mixins ──────────────────────
|
||||
|
||||
def _enrich_prompt_images(self, prompts):
|
||||
"""Add url and thumbnail_url to each image in prompt results."""
|
||||
from urllib.parse import quote as url_quote
|
||||
|
||||
output_dir = self._find_comfyui_output_dir()
|
||||
output_path = Path(output_dir) if output_dir else None
|
||||
|
||||
for prompt in prompts:
|
||||
for image in prompt.get('images', []):
|
||||
image_path_str = image.get('image_path', '')
|
||||
if not image_path_str:
|
||||
continue
|
||||
|
||||
img_path = Path(image_path_str)
|
||||
|
||||
# Set fallback url via image ID
|
||||
if image.get('id'):
|
||||
image['url'] = f"/prompt_manager/images/{image['id']}/file"
|
||||
|
||||
# Try to compute relative path and thumbnail URL
|
||||
if output_path:
|
||||
try:
|
||||
rel_path = img_path.resolve().relative_to(output_path.resolve())
|
||||
image['relative_path'] = str(rel_path)
|
||||
image['url'] = f"/prompt_manager/images/serve/{url_quote(rel_path.as_posix(), safe='/')}"
|
||||
|
||||
# Check for thumbnail
|
||||
rel_no_ext = rel_path.with_suffix('')
|
||||
thumb_rel = f"thumbnails/{rel_no_ext.as_posix()}_thumb{rel_path.suffix}"
|
||||
thumb_abs = output_path / thumb_rel
|
||||
if thumb_abs.exists():
|
||||
image['thumbnail_url'] = f"/prompt_manager/images/serve/{url_quote(thumb_rel, safe='/')}"
|
||||
except (ValueError, RuntimeError):
|
||||
pass
|
||||
|
||||
return prompts
|
||||
|
||||
def _clean_nan_recursive(self, obj):
|
||||
"""Recursively clean NaN values from nested data structures."""
|
||||
if isinstance(obj, dict):
|
||||
return {key: self._clean_nan_recursive(value) for key, value in obj.items()}
|
||||
elif isinstance(obj, list):
|
||||
return [self._clean_nan_recursive(item) for item in obj]
|
||||
elif isinstance(obj, float) and str(obj) == 'nan':
|
||||
return None
|
||||
else:
|
||||
return obj
|
||||
|
||||
def _find_comfyui_output_dir(self):
|
||||
"""Locate the ComfyUI output directory using multiple detection strategies.
|
||||
|
||||
Results are cached after first successful lookup since the output
|
||||
directory does not change during runtime.
|
||||
|
||||
Returns:
|
||||
str or None: Absolute path to ComfyUI output directory, or None
|
||||
if no valid directory is found
|
||||
"""
|
||||
if self._cached_output_dir is not None:
|
||||
return self._cached_output_dir
|
||||
|
||||
current_file = Path(__file__).resolve()
|
||||
self.logger.debug(f"Starting ComfyUI output search from: {current_file}")
|
||||
|
||||
# Method 1: Search upward from current file location
|
||||
current_dir = current_file.parent
|
||||
max_depth = 10 # Prevent infinite loops
|
||||
|
||||
for i in range(max_depth):
|
||||
# Check if current directory contains ComfyUI markers
|
||||
comfyui_markers = ['main.py', 'nodes.py', 'server.py']
|
||||
if any((current_dir / marker).exists() for marker in comfyui_markers):
|
||||
output_dir = current_dir / "output"
|
||||
if output_dir.exists() and output_dir.is_dir():
|
||||
self.logger.debug(f"Found ComfyUI output directory via upward search: {output_dir}")
|
||||
self._cached_output_dir = str(output_dir)
|
||||
return self._cached_output_dir
|
||||
|
||||
# Move up one directory
|
||||
parent = current_dir.parent
|
||||
if parent == current_dir: # Reached filesystem root
|
||||
break
|
||||
current_dir = parent
|
||||
|
||||
# Method 2: Try common installation patterns relative to this file
|
||||
# File is at: custom_nodes/ComfyUI_PromptManager/py/api/__init__.py
|
||||
base_dir = current_file.parent # .../py/api/
|
||||
possible_paths = [
|
||||
base_dir.parent.parent.parent.parent / "output", # ../../../../output
|
||||
base_dir.parent.parent.parent / "output", # ../../../output
|
||||
base_dir.parent.parent / "output", # ../../output
|
||||
base_dir.parent / "output", # ../output
|
||||
]
|
||||
|
||||
# Method 3: Add common ComfyUI installation locations
|
||||
common_locations = [
|
||||
Path.home() / "ComfyUI" / "output",
|
||||
Path.cwd() / "output",
|
||||
Path.cwd() / ".." / "output",
|
||||
Path.cwd() / ".." / ".." / "output",
|
||||
]
|
||||
|
||||
all_paths = possible_paths + common_locations
|
||||
|
||||
for path in all_paths:
|
||||
try:
|
||||
abs_path = path.resolve()
|
||||
if abs_path.exists() and abs_path.is_dir():
|
||||
self.logger.debug(f"Found ComfyUI output directory: {abs_path}")
|
||||
self._cached_output_dir = str(abs_path)
|
||||
return self._cached_output_dir
|
||||
except (OSError, RuntimeError):
|
||||
continue # Skip invalid paths
|
||||
|
||||
self.logger.warning("ComfyUI output directory not found. Searched paths:")
|
||||
for path in all_paths:
|
||||
try:
|
||||
self.logger.warning(f" - {path.resolve()} (exists: {path.exists()})")
|
||||
except (OSError, RuntimeError):
|
||||
self.logger.warning(f" - {path} (invalid path)")
|
||||
|
||||
return None
|
||||
|
||||
def _extract_comfyui_metadata(self, image_path):
|
||||
"""Extract ComfyUI workflow metadata from PNG image files."""
|
||||
try:
|
||||
with Image.open(image_path) as img:
|
||||
metadata = {}
|
||||
if hasattr(img, 'text'):
|
||||
for key, value in img.text.items():
|
||||
metadata[key] = value
|
||||
return metadata
|
||||
except Exception as e:
|
||||
self.logger.error(f"Error reading {image_path}: {e}")
|
||||
return {}
|
||||
|
||||
def _parse_comfyui_prompt(self, metadata):
|
||||
"""Parse ComfyUI and A1111 prompt data from metadata."""
|
||||
result = {
|
||||
'prompt': None,
|
||||
'workflow': None,
|
||||
'parameters': {},
|
||||
'positive_prompt': None,
|
||||
'negative_prompt': None
|
||||
}
|
||||
|
||||
# Check for A1111 style parameters first (like parse-metadata.py)
|
||||
if "parameters" in metadata:
|
||||
params = metadata["parameters"]
|
||||
lines = params.splitlines()
|
||||
if lines:
|
||||
result['positive_prompt'] = lines[0].strip()
|
||||
for line in lines:
|
||||
if line.lower().startswith("negative prompt:"):
|
||||
result['negative_prompt'] = line.split(":", 1)[1].strip()
|
||||
break
|
||||
# Store raw parameters too
|
||||
result['parameters']['parameters'] = params
|
||||
|
||||
# If no A1111 format found, proceed with ComfyUI parsing
|
||||
if result['positive_prompt'] is None:
|
||||
# Check for direct prompt field
|
||||
if 'prompt' in metadata:
|
||||
try:
|
||||
prompt_data = json.loads(metadata['prompt'])
|
||||
result['prompt'] = prompt_data
|
||||
except json.JSONDecodeError:
|
||||
result['prompt'] = metadata['prompt']
|
||||
|
||||
# Check for workflow
|
||||
if 'workflow' in metadata:
|
||||
try:
|
||||
workflow_data = json.loads(metadata['workflow'])
|
||||
result['workflow'] = workflow_data
|
||||
except json.JSONDecodeError:
|
||||
result['workflow'] = metadata['workflow']
|
||||
|
||||
# Check for other common ComfyUI fields
|
||||
common_fields = ['positive', 'negative', 'steps', 'cfg', 'sampler', 'scheduler', 'seed']
|
||||
for field in common_fields:
|
||||
if field in metadata:
|
||||
try:
|
||||
result['parameters'][field] = json.loads(metadata[field])
|
||||
except json.JSONDecodeError:
|
||||
result['parameters'][field] = metadata[field]
|
||||
|
||||
return result
|
||||
|
||||
def _extract_readable_prompt(self, parsed_data):
|
||||
"""Extract human-readable prompt text from ComfyUI/A1111 data."""
|
||||
|
||||
def safe_to_string(value):
|
||||
if isinstance(value, str):
|
||||
return value
|
||||
elif isinstance(value, list):
|
||||
return ' '.join(str(item) for item in value if item)
|
||||
elif value is not None:
|
||||
return str(value)
|
||||
return None
|
||||
|
||||
# First check if we already extracted a positive prompt (A1111 format)
|
||||
if parsed_data.get('positive_prompt'):
|
||||
return parsed_data['positive_prompt']
|
||||
|
||||
# Check if prompt is already a string
|
||||
if isinstance(parsed_data.get('prompt'), str):
|
||||
return parsed_data['prompt']
|
||||
|
||||
# Check if prompt is a simple value that can be converted
|
||||
if parsed_data.get('prompt') and not isinstance(parsed_data.get('prompt'), dict):
|
||||
return safe_to_string(parsed_data['prompt'])
|
||||
|
||||
prompt_data = parsed_data.get('prompt')
|
||||
if isinstance(prompt_data, dict):
|
||||
# Use the enhanced logic from parse-metadata.py
|
||||
positive_prompt = self._extract_positive_prompt_from_comfyui_data(prompt_data)
|
||||
if positive_prompt:
|
||||
return positive_prompt
|
||||
|
||||
# Check workflow data if available
|
||||
workflow_data = parsed_data.get('workflow')
|
||||
if isinstance(workflow_data, dict):
|
||||
positive_prompt = self._extract_positive_prompt_from_comfyui_data(workflow_data)
|
||||
if positive_prompt:
|
||||
return positive_prompt
|
||||
|
||||
# Check parameters for positive prompt
|
||||
if parsed_data.get('parameters', {}).get('positive'):
|
||||
return safe_to_string(parsed_data['parameters']['positive'])
|
||||
|
||||
return None
|
||||
|
||||
def _get_node_inputs(self, node):
|
||||
"""Safely get inputs from a node, handling both dict and list formats.
|
||||
|
||||
Old format: inputs is a dict with direct key-value pairs
|
||||
inputs = {"text": "my prompt", "seed": 123}
|
||||
|
||||
New format: inputs is a list of connection objects
|
||||
inputs = [
|
||||
{"name": "text", "type": "STRING", "link": null, "widget": {"name": "text"}},
|
||||
{"name": "clip", "type": "CLIP", "link": 11}
|
||||
]
|
||||
"""
|
||||
if not isinstance(node, dict):
|
||||
return {}
|
||||
|
||||
inputs = node.get("inputs", {})
|
||||
|
||||
# If inputs is already a dict, return it
|
||||
if isinstance(inputs, dict):
|
||||
return inputs
|
||||
|
||||
# If inputs is a list, convert to dict format
|
||||
if isinstance(inputs, list):
|
||||
inputs_dict = {}
|
||||
for input_item in inputs:
|
||||
if isinstance(input_item, dict) and "name" in input_item:
|
||||
name = input_item["name"]
|
||||
inputs_dict[name] = input_item
|
||||
return inputs_dict
|
||||
|
||||
return {}
|
||||
|
||||
def _find_text_in_node(self, node):
|
||||
"""Try to find text content in a node using various strategies.
|
||||
|
||||
Handles both old and new workflow formats.
|
||||
"""
|
||||
if not isinstance(node, dict):
|
||||
return None
|
||||
|
||||
# Strategy 1: Check normalized inputs for 'text' field
|
||||
inputs = self._get_node_inputs(node)
|
||||
if "text" in inputs and isinstance(inputs["text"], str):
|
||||
return inputs["text"]
|
||||
|
||||
# Strategy 2: For text encoder nodes, check widgets_values
|
||||
class_type = node.get("class_type", node.get("type", ""))
|
||||
text_encoder_types = [
|
||||
'CLIPTextEncode', 'CLIPTextEncodeSDXL', 'CLIPTextEncodeSDXLRefiner',
|
||||
'CLIPTextEncodeFlux', 'PromptManager', 'PromptManagerText',
|
||||
'BNK_CLIPTextEncoder', 'Text Encoder', 'CLIP Text Encode'
|
||||
]
|
||||
|
||||
if any(encoder_type.lower() in class_type.lower() for encoder_type in text_encoder_types):
|
||||
widgets_values = node.get("widgets_values", [])
|
||||
if widgets_values and len(widgets_values) > 0:
|
||||
if isinstance(widgets_values[0], str) and widgets_values[0].strip():
|
||||
return widgets_values[0]
|
||||
|
||||
return None
|
||||
|
||||
def _extract_positive_prompt_from_comfyui_data(self, data):
|
||||
"""Extract positive prompt from ComfyUI data, handling both old and new formats."""
|
||||
if not isinstance(data, dict):
|
||||
return None
|
||||
|
||||
# Build nodes dictionary
|
||||
nodes_by_id = {}
|
||||
if "nodes" in data:
|
||||
# Handle nodes array format
|
||||
for node in data["nodes"]:
|
||||
if isinstance(node, dict):
|
||||
nid = node.get("id")
|
||||
if nid is not None:
|
||||
nodes_by_id[nid] = node
|
||||
else:
|
||||
# Handle flat dictionary format (node_id -> node_data)
|
||||
for nid_str, node in data.items():
|
||||
try:
|
||||
nid = int(nid_str)
|
||||
except (ValueError, TypeError):
|
||||
nid = nid_str
|
||||
if isinstance(node, dict):
|
||||
if "id" in node:
|
||||
nid = node["id"]
|
||||
nodes_by_id[nid] = node
|
||||
|
||||
if not nodes_by_id:
|
||||
return None
|
||||
|
||||
# First, try to find positive/negative connection pattern
|
||||
pos_id = None
|
||||
for node in nodes_by_id.values():
|
||||
if isinstance(node, dict):
|
||||
inputs = self._get_node_inputs(node)
|
||||
if "positive" in inputs and "negative" in inputs:
|
||||
try:
|
||||
pos_input = inputs["positive"]
|
||||
if isinstance(pos_input, list) and len(pos_input) > 0:
|
||||
pos_id = int(pos_input[0])
|
||||
break
|
||||
except (ValueError, TypeError, IndexError):
|
||||
continue
|
||||
|
||||
# Get text from the positive node
|
||||
if pos_id is not None and pos_id in nodes_by_id:
|
||||
text_val = self._find_text_in_node(nodes_by_id[pos_id])
|
||||
if text_val:
|
||||
return text_val
|
||||
|
||||
# Fallback: find any text encoder node with text content
|
||||
text_nodes = []
|
||||
for node in nodes_by_id.values():
|
||||
if isinstance(node, dict):
|
||||
text_val = self._find_text_in_node(node)
|
||||
if text_val:
|
||||
# Try to determine if this is positive or negative
|
||||
node_title = node.get('title', '').lower()
|
||||
if 'neg' not in node_title and 'negative' not in node_title:
|
||||
# Prioritize non-negative prompts
|
||||
text_nodes.insert(0, text_val)
|
||||
else:
|
||||
text_nodes.append(text_val)
|
||||
|
||||
# Return the first positive-looking prompt
|
||||
if text_nodes:
|
||||
return text_nodes[0]
|
||||
|
||||
return None
|
||||
+963
@@ -0,0 +1,963 @@
|
||||
"""Admin and maintenance API routes for PromptManager."""
|
||||
|
||||
import asyncio
|
||||
import datetime
|
||||
import hashlib
|
||||
import json
|
||||
import os
|
||||
import shutil
|
||||
import sqlite3
|
||||
import tempfile
|
||||
import traceback
|
||||
from pathlib import Path
|
||||
|
||||
from aiohttp import web
|
||||
|
||||
|
||||
class AdminRoutesMixin:
|
||||
"""Mixin providing admin, diagnostics, and maintenance API endpoints."""
|
||||
|
||||
def _register_admin_routes(self, routes):
|
||||
@routes.get("/prompt_manager/scan_duplicates")
|
||||
async def scan_duplicates_route(request):
|
||||
return await self.scan_duplicates_endpoint(request)
|
||||
|
||||
@routes.post("/prompt_manager/delete_duplicate_images")
|
||||
async def delete_duplicate_images_route(request):
|
||||
return await self.delete_duplicate_images_endpoint(request)
|
||||
|
||||
@routes.post("/prompt_manager/cleanup")
|
||||
async def cleanup_duplicates_route(request):
|
||||
return await self.cleanup_duplicates_endpoint(request)
|
||||
|
||||
@routes.post("/prompt_manager/maintenance")
|
||||
async def maintenance_route(request):
|
||||
return await self.run_maintenance(request)
|
||||
|
||||
@routes.get("/prompt_manager/stats")
|
||||
async def get_stats_route(request):
|
||||
return await self.get_statistics(request)
|
||||
|
||||
@routes.get("/prompt_manager/settings")
|
||||
async def get_settings_route(request):
|
||||
return await self.get_settings(request)
|
||||
|
||||
@routes.post("/prompt_manager/settings")
|
||||
async def save_settings_route(request):
|
||||
return await self.save_settings(request)
|
||||
|
||||
@routes.get("/prompt_manager/backup")
|
||||
async def backup_database_route(request):
|
||||
return await self.backup_database(request)
|
||||
|
||||
@routes.post("/prompt_manager/restore")
|
||||
async def restore_database_route(request):
|
||||
return await self.restore_database(request)
|
||||
|
||||
@routes.get("/prompt_manager/diagnostics")
|
||||
async def run_diagnostics_route(request):
|
||||
return await self.run_diagnostics(request)
|
||||
|
||||
@routes.post("/prompt_manager/diagnostics/test-link")
|
||||
async def test_image_link_route(request):
|
||||
return await self.test_image_link(request)
|
||||
|
||||
@routes.post("/prompt_manager/scan")
|
||||
async def scan_images_route(request):
|
||||
return await self.scan_images(request)
|
||||
|
||||
async def scan_duplicates_endpoint(self, request):
|
||||
"""Scan for duplicate images without removing them."""
|
||||
try:
|
||||
duplicates = await self.find_duplicate_images()
|
||||
|
||||
return web.json_response(
|
||||
{
|
||||
"success": True,
|
||||
"duplicates": duplicates,
|
||||
"message": f"Found {len(duplicates)} groups of duplicate images",
|
||||
}
|
||||
)
|
||||
|
||||
except Exception as e:
|
||||
self.logger.error(f"Scan duplicates error: {e}")
|
||||
return web.json_response(
|
||||
{"success": False, "error": f"Failed to scan duplicate images: {str(e)}"},
|
||||
status=500,
|
||||
)
|
||||
|
||||
async def cleanup_duplicates_endpoint(self, request):
|
||||
"""Cleanup duplicate prompts endpoint."""
|
||||
try:
|
||||
removed_count = await self._run_in_executor(self.db.cleanup_duplicates)
|
||||
|
||||
return web.json_response(
|
||||
{
|
||||
"success": True,
|
||||
"message": "Cleanup completed",
|
||||
"duplicates_removed": removed_count,
|
||||
}
|
||||
)
|
||||
|
||||
except Exception as e:
|
||||
self.logger.error(f"Cleanup error: {e}")
|
||||
return web.json_response(
|
||||
{"success": False, "error": f"Failed to cleanup duplicates: {str(e)}"},
|
||||
status=500,
|
||||
)
|
||||
|
||||
async def find_duplicate_images(self):
|
||||
"""Find duplicate images in ComfyUI output directory using content hashing."""
|
||||
self.logger.info("Scanning for duplicate images")
|
||||
|
||||
try:
|
||||
output_dir = self._find_comfyui_output_dir()
|
||||
if not output_dir:
|
||||
self.logger.warning("ComfyUI output directory not found")
|
||||
return []
|
||||
|
||||
output_path = Path(output_dir)
|
||||
|
||||
image_extensions = ['.png', '.jpg', '.jpeg', '.gif', '.webp', '.bmp', '.tiff']
|
||||
video_extensions = ['.mp4', '.avi', '.mov', '.mkv', '.webm', '.gif']
|
||||
all_extensions = image_extensions + video_extensions
|
||||
|
||||
media_files = []
|
||||
seen_paths = set()
|
||||
for ext in all_extensions:
|
||||
for pattern in [f"*{ext.lower()}", f"*{ext.upper()}"]:
|
||||
for media_path in output_path.rglob(pattern):
|
||||
if 'thumbnails' not in media_path.parts:
|
||||
normalized_path = str(media_path).lower()
|
||||
if normalized_path not in seen_paths:
|
||||
seen_paths.add(normalized_path)
|
||||
media_files.append(media_path)
|
||||
|
||||
self.logger.info(f"Found {len(media_files)} media files to analyze")
|
||||
|
||||
file_hashes = {}
|
||||
processed = 0
|
||||
|
||||
for media_path in media_files:
|
||||
try:
|
||||
file_hash = self._calculate_file_hash(media_path)
|
||||
|
||||
if file_hash not in file_hashes:
|
||||
file_hashes[file_hash] = []
|
||||
|
||||
stat = media_path.stat()
|
||||
rel_path = media_path.relative_to(output_path)
|
||||
extension = media_path.suffix.lower()
|
||||
is_video = extension in [ext.lower() for ext in video_extensions]
|
||||
media_type = 'video' if is_video else 'image'
|
||||
|
||||
# Check if thumbnail exists
|
||||
thumbnail_url = None
|
||||
thumbnails_dir = output_path / "thumbnails"
|
||||
if thumbnails_dir.exists():
|
||||
thumbnail_ext = '.jpg' if is_video else extension
|
||||
rel_path_no_ext = rel_path.with_suffix('')
|
||||
thumbnail_rel_path = f"thumbnails/{rel_path_no_ext.as_posix()}_thumb{thumbnail_ext}"
|
||||
thumbnail_abs_path = output_path / thumbnail_rel_path
|
||||
|
||||
if thumbnail_abs_path.exists():
|
||||
from urllib.parse import quote
|
||||
thumbnail_url = f'/prompt_manager/images/serve/{quote(thumbnail_rel_path, safe="/")}'
|
||||
|
||||
image_info = {
|
||||
'id': str(hash(str(media_path))),
|
||||
'filename': media_path.name,
|
||||
'path': str(media_path),
|
||||
'relative_path': str(rel_path),
|
||||
'url': f'/prompt_manager/images/serve/{rel_path.as_posix()}',
|
||||
'thumbnail_url': thumbnail_url,
|
||||
'size': stat.st_size,
|
||||
'modified_time': stat.st_mtime,
|
||||
'extension': extension,
|
||||
'media_type': media_type,
|
||||
'is_video': is_video,
|
||||
'hash': file_hash
|
||||
}
|
||||
|
||||
file_hashes[file_hash].append(image_info)
|
||||
processed += 1
|
||||
|
||||
if processed % 100 == 0:
|
||||
self.logger.info(f"Processed {processed}/{len(media_files)} files for duplicate detection")
|
||||
|
||||
except Exception as e:
|
||||
self.logger.error(f"Error processing file {media_path}: {e}")
|
||||
continue
|
||||
|
||||
# Find duplicates (groups with more than one file)
|
||||
duplicates = []
|
||||
for file_hash, images in file_hashes.items():
|
||||
if len(images) > 1:
|
||||
images.sort(key=lambda x: x['modified_time'])
|
||||
duplicates.append({
|
||||
'hash': file_hash,
|
||||
'images': images,
|
||||
'count': len(images)
|
||||
})
|
||||
|
||||
self.logger.info(f"Found {len(duplicates)} groups of duplicate images")
|
||||
return duplicates
|
||||
|
||||
except Exception as e:
|
||||
self.logger.error(f"Error finding duplicate images: {e}")
|
||||
return []
|
||||
|
||||
def _calculate_file_hash(self, file_path):
|
||||
"""Calculate SHA-256 hash of a file's content."""
|
||||
hash_sha256 = hashlib.sha256()
|
||||
with open(file_path, "rb") as f:
|
||||
for chunk in iter(lambda: f.read(4096), b""):
|
||||
hash_sha256.update(chunk)
|
||||
return hash_sha256.hexdigest()
|
||||
|
||||
async def delete_duplicate_images_endpoint(self, request):
|
||||
"""Delete duplicate image files from disk."""
|
||||
try:
|
||||
data = await request.json()
|
||||
image_paths = data.get('image_paths', [])
|
||||
|
||||
if not image_paths:
|
||||
return web.json_response(
|
||||
{"success": False, "error": "No image paths provided"},
|
||||
status=400,
|
||||
)
|
||||
|
||||
deleted_count = 0
|
||||
failed_count = 0
|
||||
failed_files = []
|
||||
|
||||
for image_path in image_paths:
|
||||
try:
|
||||
# Ensure the path is within the output directory for security
|
||||
output_dir = self._find_comfyui_output_dir()
|
||||
if not output_dir:
|
||||
failed_files.append(f"{image_path} (output directory not found)")
|
||||
failed_count += 1
|
||||
continue
|
||||
|
||||
output_path = Path(output_dir)
|
||||
file_path = Path(image_path)
|
||||
|
||||
# Security check - ensure file is within output directory
|
||||
try:
|
||||
file_path.resolve().relative_to(output_path.resolve())
|
||||
except ValueError:
|
||||
self.logger.warning(f"Attempted to delete file outside output directory: {image_path}")
|
||||
failed_files.append(f"{image_path} (outside output directory)")
|
||||
failed_count += 1
|
||||
continue
|
||||
|
||||
if file_path.exists() and file_path.is_file():
|
||||
os.remove(file_path)
|
||||
deleted_count += 1
|
||||
self.logger.info(f"Deleted duplicate image: {image_path}")
|
||||
|
||||
# Also try to delete associated thumbnail if it exists
|
||||
try:
|
||||
rel_path = file_path.relative_to(output_path)
|
||||
rel_path_no_ext = rel_path.with_suffix('')
|
||||
thumbnail_path = output_path / "thumbnails" / f"{rel_path_no_ext.as_posix()}_thumb{file_path.suffix}"
|
||||
if thumbnail_path.exists():
|
||||
os.remove(thumbnail_path)
|
||||
self.logger.debug(f"Deleted associated thumbnail: {thumbnail_path}")
|
||||
except Exception as e:
|
||||
self.logger.warning(f"Could not delete thumbnail for {image_path}: {e}")
|
||||
else:
|
||||
failed_files.append(f"{image_path} (file not found)")
|
||||
failed_count += 1
|
||||
|
||||
except Exception as e:
|
||||
self.logger.error(f"Error deleting file {image_path}: {e}")
|
||||
failed_files.append(f"{image_path} ({str(e)})")
|
||||
failed_count += 1
|
||||
|
||||
response_data = {
|
||||
"success": True,
|
||||
"deleted_count": deleted_count,
|
||||
"failed_count": failed_count,
|
||||
"message": f"Deleted {deleted_count} files successfully"
|
||||
}
|
||||
|
||||
if failed_count > 0:
|
||||
response_data["failed_files"] = failed_files
|
||||
response_data["message"] += f", {failed_count} failed"
|
||||
|
||||
return web.json_response(response_data)
|
||||
|
||||
except Exception as e:
|
||||
self.logger.error(f"Delete duplicate images error: {e}")
|
||||
return web.json_response(
|
||||
{"success": False, "error": f"Failed to delete duplicate images: {str(e)}"},
|
||||
status=500,
|
||||
)
|
||||
|
||||
async def get_statistics(self, request):
|
||||
"""Get database statistics."""
|
||||
try:
|
||||
stats = await self._run_in_executor(self.db.get_statistics)
|
||||
|
||||
return web.json_response({"success": True, "stats": stats})
|
||||
|
||||
except Exception as e:
|
||||
self.logger.error(f"Stats error: {e}")
|
||||
return web.json_response(
|
||||
{"success": False, "error": f"Failed to get statistics: {str(e)}"},
|
||||
status=500,
|
||||
)
|
||||
|
||||
async def get_settings(self, request):
|
||||
"""Get current settings."""
|
||||
try:
|
||||
from ..config import PromptManagerConfig, GalleryConfig
|
||||
|
||||
# Get monitored directories from image monitor singleton if available
|
||||
monitored_dirs = []
|
||||
try:
|
||||
import sys
|
||||
monitor_module = None
|
||||
for mod_name in list(sys.modules.keys()):
|
||||
if 'image_monitor' in mod_name and hasattr(sys.modules[mod_name], '_monitor_instance'):
|
||||
monitor_module = sys.modules[mod_name]
|
||||
break
|
||||
|
||||
if monitor_module and monitor_module._monitor_instance is not None:
|
||||
monitored_dirs = getattr(monitor_module._monitor_instance, 'monitored_directories', [])
|
||||
elif GalleryConfig.MONITORING_DIRECTORIES:
|
||||
monitored_dirs = GalleryConfig.MONITORING_DIRECTORIES
|
||||
except Exception:
|
||||
if GalleryConfig.MONITORING_DIRECTORIES:
|
||||
monitored_dirs = GalleryConfig.MONITORING_DIRECTORIES
|
||||
|
||||
return web.json_response({
|
||||
"success": True,
|
||||
"settings": {
|
||||
"result_timeout": PromptManagerConfig.RESULT_TIMEOUT,
|
||||
"webui_display_mode": PromptManagerConfig.WEBUI_DISPLAY_MODE,
|
||||
"gallery_root_path": GalleryConfig.MONITORING_DIRECTORIES[0] if GalleryConfig.MONITORING_DIRECTORIES else "",
|
||||
"monitored_directories": monitored_dirs
|
||||
}
|
||||
})
|
||||
except Exception as e:
|
||||
return web.json_response(
|
||||
{"success": False, "error": f"Failed to get settings: {str(e)}"},
|
||||
status=500,
|
||||
)
|
||||
|
||||
async def save_settings(self, request):
|
||||
"""Save settings."""
|
||||
try:
|
||||
from ..config import PromptManagerConfig, GalleryConfig
|
||||
|
||||
data = await request.json()
|
||||
restart_required = False
|
||||
|
||||
# Update in-memory config
|
||||
if 'result_timeout' in data:
|
||||
PromptManagerConfig.RESULT_TIMEOUT = data['result_timeout']
|
||||
if 'webui_display_mode' in data:
|
||||
PromptManagerConfig.WEBUI_DISPLAY_MODE = data['webui_display_mode']
|
||||
|
||||
# Handle gallery root path
|
||||
if 'gallery_root_path' in data:
|
||||
new_path = data['gallery_root_path'].strip()
|
||||
old_path = GalleryConfig.MONITORING_DIRECTORIES[0] if GalleryConfig.MONITORING_DIRECTORIES else ""
|
||||
|
||||
if new_path != old_path:
|
||||
if new_path:
|
||||
from pathlib import Path as _Path
|
||||
resolved = _Path(new_path).resolve()
|
||||
if not resolved.is_dir():
|
||||
return web.json_response({
|
||||
'success': False,
|
||||
'error': f'Gallery path does not exist or is not a directory: {new_path}'
|
||||
}, status=400)
|
||||
blocked = ['/etc', '/usr', '/bin', '/sbin', '/boot', '/proc', '/sys', '/dev',
|
||||
'/var/log', '/root', 'C:\\Windows', 'C:\\Program Files']
|
||||
for b in blocked:
|
||||
if str(resolved).startswith(b):
|
||||
return web.json_response({
|
||||
'success': False,
|
||||
'error': 'Gallery path cannot point to a system directory'
|
||||
}, status=400)
|
||||
GalleryConfig.MONITORING_DIRECTORIES = [new_path]
|
||||
else:
|
||||
GalleryConfig.MONITORING_DIRECTORIES = []
|
||||
restart_required = True
|
||||
|
||||
# Save to config file for persistence
|
||||
config_dir = os.path.dirname(os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
|
||||
config_file = os.path.join(config_dir, 'config.json')
|
||||
|
||||
config_data = {
|
||||
'web_ui': {
|
||||
'result_timeout': PromptManagerConfig.RESULT_TIMEOUT,
|
||||
'webui_display_mode': PromptManagerConfig.WEBUI_DISPLAY_MODE
|
||||
},
|
||||
'gallery': {
|
||||
'monitoring': {
|
||||
'directories': GalleryConfig.MONITORING_DIRECTORIES
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
try:
|
||||
with open(config_file, 'w') as f:
|
||||
json.dump(config_data, f, indent=2)
|
||||
self.logger.info(f"Settings saved to {config_file}")
|
||||
except Exception as save_err:
|
||||
self.logger.warning(f"Could not save config file: {save_err}")
|
||||
|
||||
return web.json_response({
|
||||
"success": True,
|
||||
"message": "Settings saved successfully",
|
||||
"restart_required": restart_required
|
||||
})
|
||||
except Exception as e:
|
||||
return web.json_response(
|
||||
{"success": False, "error": f"Failed to save settings: {str(e)}"},
|
||||
status=500,
|
||||
)
|
||||
|
||||
async def run_diagnostics(self, request):
|
||||
"""Run comprehensive system diagnostics and health checks."""
|
||||
try:
|
||||
results = {}
|
||||
|
||||
# Check database
|
||||
try:
|
||||
db_path = "prompts.db"
|
||||
if os.path.exists(db_path):
|
||||
with sqlite3.connect(db_path) as conn:
|
||||
conn.row_factory = sqlite3.Row
|
||||
cursor = conn.execute("SELECT COUNT(*) as count FROM prompts")
|
||||
prompt_count = cursor.fetchone()['count']
|
||||
|
||||
cursor = conn.execute("SELECT name FROM sqlite_master WHERE type='table' AND name='generated_images'")
|
||||
has_images_table = cursor.fetchone() is not None
|
||||
|
||||
if has_images_table:
|
||||
cursor = conn.execute("SELECT COUNT(*) as count FROM generated_images")
|
||||
image_count = cursor.fetchone()['count']
|
||||
else:
|
||||
image_count = 0
|
||||
|
||||
results['database'] = {
|
||||
'status': 'ok',
|
||||
'prompt_count': prompt_count,
|
||||
'has_images_table': has_images_table,
|
||||
'image_count': image_count
|
||||
}
|
||||
else:
|
||||
results['database'] = {
|
||||
'status': 'error',
|
||||
'message': f'Database file not found: {db_path}'
|
||||
}
|
||||
except Exception as e:
|
||||
results['database'] = {
|
||||
'status': 'error',
|
||||
'message': f'Database error: {str(e)}'
|
||||
}
|
||||
|
||||
# Check dependencies
|
||||
dependencies = {}
|
||||
try:
|
||||
import watchdog
|
||||
dependencies['watchdog'] = True
|
||||
except ImportError:
|
||||
dependencies['watchdog'] = False
|
||||
|
||||
try:
|
||||
from PIL import Image
|
||||
dependencies['PIL'] = True
|
||||
except ImportError:
|
||||
dependencies['PIL'] = False
|
||||
|
||||
dependencies['sqlite3'] = True # Always available in Python
|
||||
|
||||
results['dependencies'] = {
|
||||
'status': 'ok' if all(dependencies.values()) else 'error',
|
||||
'dependencies': dependencies
|
||||
}
|
||||
|
||||
# Check output directories
|
||||
output_dirs = []
|
||||
potential_dirs = ["output", "../output", "../../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)
|
||||
|
||||
results['comfyui_output'] = {
|
||||
'status': 'ok' if output_dirs else 'warning',
|
||||
'output_dirs': output_dirs
|
||||
}
|
||||
|
||||
# Check image monitor status
|
||||
try:
|
||||
from ...utils.image_monitor import _monitor_instance
|
||||
if _monitor_instance is not None:
|
||||
monitor_status = _monitor_instance.get_status()
|
||||
results['image_monitor'] = {
|
||||
'status': 'ok' if monitor_status.get('observer_alive') else 'error',
|
||||
**monitor_status
|
||||
}
|
||||
else:
|
||||
results['image_monitor'] = {
|
||||
'status': 'error',
|
||||
'message': 'Image monitor not initialized'
|
||||
}
|
||||
except Exception as e:
|
||||
results['image_monitor'] = {
|
||||
'status': 'error',
|
||||
'message': f'Failed to get monitor status: {str(e)}'
|
||||
}
|
||||
|
||||
return web.json_response({
|
||||
'success': True,
|
||||
'diagnostics': results
|
||||
})
|
||||
|
||||
except Exception as e:
|
||||
self.logger.error(f"Diagnostics error: {e}", exc_info=True)
|
||||
return web.json_response({
|
||||
'success': False,
|
||||
'error': str(e)
|
||||
}, status=500)
|
||||
|
||||
async def test_image_link(self, request):
|
||||
"""Test creating an image link."""
|
||||
try:
|
||||
data = await request.json()
|
||||
prompt_id = data.get('prompt_id')
|
||||
test_image_path = data.get('image_path', '/test/fake/image.png')
|
||||
|
||||
if not prompt_id:
|
||||
return web.json_response({
|
||||
'success': False,
|
||||
'error': 'prompt_id is required'
|
||||
}, status=400)
|
||||
|
||||
test_metadata = {
|
||||
'file_info': {
|
||||
'size': 1024000,
|
||||
'dimensions': [512, 512],
|
||||
'format': 'PNG'
|
||||
},
|
||||
'workflow': {'test': True},
|
||||
'prompt': {'test_prompt': 'This is a test image'}
|
||||
}
|
||||
|
||||
try:
|
||||
image_id = await self._run_in_executor(
|
||||
self.db.link_image_to_prompt,
|
||||
prompt_id=str(prompt_id),
|
||||
image_path=test_image_path,
|
||||
metadata=test_metadata
|
||||
)
|
||||
|
||||
return web.json_response({
|
||||
'success': True,
|
||||
'result': {
|
||||
'status': 'ok',
|
||||
'image_id': image_id,
|
||||
'message': f'Test image linked successfully with ID {image_id}'
|
||||
}
|
||||
})
|
||||
except Exception as e:
|
||||
return web.json_response({
|
||||
'success': False,
|
||||
'result': {
|
||||
'status': 'error',
|
||||
'message': f'Failed to create test link: {str(e)}'
|
||||
}
|
||||
})
|
||||
|
||||
except Exception as e:
|
||||
self.logger.error(f"Test link error: {e}", exc_info=True)
|
||||
return web.json_response({
|
||||
'success': False,
|
||||
'error': str(e)
|
||||
}, status=500)
|
||||
|
||||
async def run_maintenance(self, request):
|
||||
"""Perform comprehensive database maintenance and optimization."""
|
||||
try:
|
||||
data = await request.json() if request.content_type == 'application/json' else {}
|
||||
operations = data.get('operations', ['cleanup_duplicates', 'vacuum', 'cleanup_orphaned_images'])
|
||||
|
||||
results = {}
|
||||
|
||||
def _run_maintenance():
|
||||
if 'cleanup_duplicates' in operations:
|
||||
try:
|
||||
duplicates_removed = self.db.cleanup_duplicates()
|
||||
results['cleanup_duplicates'] = {
|
||||
'success': True, 'removed_count': duplicates_removed,
|
||||
'message': f'Removed {duplicates_removed} duplicate prompts'
|
||||
}
|
||||
except Exception as e:
|
||||
results['cleanup_duplicates'] = {'success': False, 'error': str(e), 'message': 'Failed to cleanup duplicates'}
|
||||
|
||||
if 'vacuum' in operations:
|
||||
try:
|
||||
self.db.model.vacuum_database()
|
||||
results['vacuum'] = {'success': True, 'message': 'Database vacuum completed successfully'}
|
||||
except Exception as e:
|
||||
results['vacuum'] = {'success': False, 'error': str(e), 'message': 'Failed to vacuum database'}
|
||||
|
||||
if 'cleanup_orphaned_images' in operations:
|
||||
try:
|
||||
orphaned_removed = self.db.cleanup_missing_images()
|
||||
results['cleanup_orphaned_images'] = {
|
||||
'success': True, 'removed_count': orphaned_removed,
|
||||
'message': f'Removed {orphaned_removed} orphaned image records'
|
||||
}
|
||||
except Exception as e:
|
||||
results['cleanup_orphaned_images'] = {'success': False, 'error': str(e), 'message': 'Failed to cleanup orphaned images'}
|
||||
|
||||
if 'check_hash_duplicates' in operations:
|
||||
try:
|
||||
hash_duplicates = self.db.check_hash_duplicates()
|
||||
results['check_hash_duplicates'] = {
|
||||
'success': True, 'duplicate_hashes': len(hash_duplicates),
|
||||
'message': f'Found {len(hash_duplicates)} duplicate hash groups'
|
||||
}
|
||||
except Exception as e:
|
||||
results['check_hash_duplicates'] = {'success': False, 'error': str(e), 'message': 'Failed to check hash duplicates'}
|
||||
|
||||
if 'statistics' in operations:
|
||||
try:
|
||||
db_info = self.db.model.get_database_info()
|
||||
results['statistics'] = {'success': True, 'info': db_info, 'message': 'Database statistics retrieved'}
|
||||
except Exception as e:
|
||||
results['statistics'] = {'success': False, 'error': str(e), 'message': 'Failed to get database statistics'}
|
||||
|
||||
if 'prune_orphaned_prompts' in operations:
|
||||
try:
|
||||
removed_count = self.db.prune_orphaned_prompts()
|
||||
results['prune_orphaned_prompts'] = {
|
||||
'success': True, 'removed_count': removed_count,
|
||||
'message': f'Removed {removed_count} orphaned prompts (prompts with no linked images, excluding protected prompts)'
|
||||
}
|
||||
except Exception as e:
|
||||
results['prune_orphaned_prompts'] = {'success': False, 'error': str(e), 'message': 'Failed to prune orphaned prompts'}
|
||||
|
||||
if 'check_consistency' in operations:
|
||||
try:
|
||||
consistency_issues = self.db.check_consistency()
|
||||
results['check_consistency'] = {
|
||||
'success': True, 'issues_found': len(consistency_issues),
|
||||
'issues': consistency_issues[:10],
|
||||
'message': f'Found {len(consistency_issues)} consistency issues'
|
||||
}
|
||||
except Exception as e:
|
||||
results['check_consistency'] = {'success': False, 'error': str(e), 'message': 'Failed to check database consistency'}
|
||||
|
||||
await self._run_in_executor(_run_maintenance)
|
||||
|
||||
all_successful = all(result.get('success', False) for result in results.values())
|
||||
|
||||
return web.json_response({
|
||||
'success': True,
|
||||
'operations_completed': len(results),
|
||||
'all_successful': all_successful,
|
||||
'results': results,
|
||||
'message': f'Maintenance completed: {len(results)} operations processed'
|
||||
})
|
||||
|
||||
except Exception as e:
|
||||
self.logger.error(f"Maintenance error: {e}", exc_info=True)
|
||||
return web.json_response({
|
||||
'success': False,
|
||||
'error': f'Maintenance failed: {str(e)}'
|
||||
}, status=500)
|
||||
|
||||
async def backup_database(self, request):
|
||||
"""Backup the entire prompts.db database file."""
|
||||
try:
|
||||
db_path = "prompts.db"
|
||||
|
||||
if not os.path.exists(db_path):
|
||||
return web.json_response({
|
||||
'success': False,
|
||||
'error': 'Database file not found'
|
||||
}, status=404)
|
||||
|
||||
with tempfile.NamedTemporaryFile(delete=False, suffix='.db') as temp_file:
|
||||
temp_path = temp_file.name
|
||||
|
||||
shutil.copy2(db_path, temp_path)
|
||||
|
||||
with open(temp_path, 'rb') as f:
|
||||
file_data = f.read()
|
||||
|
||||
os.unlink(temp_path)
|
||||
|
||||
timestamp = datetime.datetime.now().strftime("%Y%m%d_%H%M%S")
|
||||
filename = f"prompts_backup_{timestamp}.db"
|
||||
|
||||
return web.Response(
|
||||
body=file_data,
|
||||
content_type='application/octet-stream',
|
||||
headers={
|
||||
'Content-Disposition': f'attachment; filename="{filename}"',
|
||||
'Content-Length': str(len(file_data))
|
||||
}
|
||||
)
|
||||
|
||||
except Exception as e:
|
||||
self.logger.error(f"Backup error: {e}", exc_info=True)
|
||||
return web.json_response({
|
||||
'success': False,
|
||||
'error': f'Failed to backup database: {str(e)}'
|
||||
}, status=500)
|
||||
|
||||
async def restore_database(self, request):
|
||||
"""Restore the prompts.db database from uploaded file."""
|
||||
try:
|
||||
reader = await request.multipart()
|
||||
field = await reader.next()
|
||||
|
||||
if not field or field.name != 'database_file':
|
||||
return web.json_response({
|
||||
'success': False,
|
||||
'error': 'No database file uploaded. Expected field name: database_file'
|
||||
}, status=400)
|
||||
|
||||
MAX_RESTORE_SIZE = 100 * 1024 * 1024 # 100MB
|
||||
file_data = await field.read()
|
||||
|
||||
if not file_data:
|
||||
return web.json_response({
|
||||
'success': False,
|
||||
'error': 'Uploaded file is empty'
|
||||
}, status=400)
|
||||
|
||||
if len(file_data) > MAX_RESTORE_SIZE:
|
||||
return web.json_response({
|
||||
'success': False,
|
||||
'error': f'File too large. Maximum size is {MAX_RESTORE_SIZE // (1024*1024)}MB'
|
||||
}, status=400)
|
||||
|
||||
with tempfile.NamedTemporaryFile(delete=False, suffix='.db') as temp_file:
|
||||
temp_path = temp_file.name
|
||||
temp_file.write(file_data)
|
||||
|
||||
try:
|
||||
with sqlite3.connect(temp_path) as conn:
|
||||
conn.row_factory = sqlite3.Row
|
||||
|
||||
cursor = conn.execute("SELECT name FROM sqlite_master WHERE type='table' AND name='prompts'")
|
||||
if not cursor.fetchone():
|
||||
raise ValueError("Database does not contain a 'prompts' table")
|
||||
|
||||
cursor = conn.execute("PRAGMA table_info(prompts)")
|
||||
columns = [row['name'] for row in cursor.fetchall()]
|
||||
required_columns = ['id', 'text', 'created_at']
|
||||
|
||||
for col in required_columns:
|
||||
if col not in columns:
|
||||
raise ValueError(f"Database missing required column: {col}")
|
||||
|
||||
cursor = conn.execute("SELECT COUNT(*) as count FROM prompts")
|
||||
prompt_count = cursor.fetchone()['count']
|
||||
|
||||
db_path = "prompts.db"
|
||||
backup_path = f"{db_path}.backup_{datetime.datetime.now().strftime('%Y%m%d_%H%M%S')}"
|
||||
|
||||
if os.path.exists(db_path):
|
||||
shutil.copy2(db_path, backup_path)
|
||||
self.logger.info(f"Current database backed up to: {backup_path}")
|
||||
|
||||
shutil.copy2(temp_path, db_path)
|
||||
|
||||
# Reinitialize the database connection
|
||||
try:
|
||||
from ...database.operations import PromptDatabase
|
||||
except ImportError:
|
||||
import sys
|
||||
sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.dirname(os.path.abspath(__file__)))))
|
||||
from database.operations import PromptDatabase
|
||||
self.db = PromptDatabase()
|
||||
|
||||
return web.json_response({
|
||||
'success': True,
|
||||
'message': f'Database restored successfully. Found {prompt_count} prompts.',
|
||||
'prompt_count': prompt_count,
|
||||
'backup_created': backup_path if os.path.exists(db_path) else None
|
||||
})
|
||||
|
||||
except sqlite3.Error as e:
|
||||
return web.json_response({
|
||||
'success': False,
|
||||
'error': f'Invalid SQLite database: {str(e)}'
|
||||
}, status=400)
|
||||
except ValueError as e:
|
||||
return web.json_response({
|
||||
'success': False,
|
||||
'error': str(e)
|
||||
}, status=400)
|
||||
finally:
|
||||
if os.path.exists(temp_path):
|
||||
os.unlink(temp_path)
|
||||
|
||||
except Exception as e:
|
||||
self.logger.error(f"Restore error: {e}", exc_info=True)
|
||||
return web.json_response({
|
||||
'success': False,
|
||||
'error': f'Failed to restore database: {str(e)}'
|
||||
}, status=500)
|
||||
|
||||
async def scan_images(self, request):
|
||||
"""Scan ComfyUI output images for prompt metadata and add them to the database."""
|
||||
|
||||
async def stream_response():
|
||||
try:
|
||||
self.logger.info("Starting image scan operation")
|
||||
self.logger.info("Starting scan (timer clearing not implemented yet)")
|
||||
|
||||
output_dir = self._find_comfyui_output_dir()
|
||||
if not output_dir:
|
||||
self.logger.error("ComfyUI output directory not found")
|
||||
yield f"data: {json.dumps({'type': 'error', 'message': 'ComfyUI output directory not found'})}\n\n"
|
||||
return
|
||||
|
||||
yield f"data: {json.dumps({'type': 'progress', 'progress': 0, 'status': 'Scanning for PNG files...', 'processed': 0, 'found': 0})}\n\n"
|
||||
|
||||
png_files = await self._run_in_executor(
|
||||
lambda: list(Path(output_dir).rglob("*.png"))
|
||||
)
|
||||
total_files = len(png_files)
|
||||
|
||||
if total_files == 0:
|
||||
yield f"data: {json.dumps({'type': 'complete', 'processed': 0, 'found': 0, 'added': 0})}\n\n"
|
||||
return
|
||||
|
||||
yield f"data: {json.dumps({'type': 'progress', 'progress': 5, 'status': f'Found {total_files} PNG files to process...', 'processed': 0, 'found': 0})}\n\n"
|
||||
|
||||
processed_count = 0
|
||||
found_count = 0
|
||||
added_count = 0
|
||||
linked_count = 0
|
||||
|
||||
try:
|
||||
from ...utils.hashing import generate_prompt_hash
|
||||
except ImportError:
|
||||
import sys
|
||||
current_dir = os.path.dirname(os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
|
||||
sys.path.insert(0, current_dir)
|
||||
from utils.hashing import generate_prompt_hash
|
||||
|
||||
for i, png_file in enumerate(png_files):
|
||||
try:
|
||||
metadata = await self._run_in_executor(
|
||||
self._extract_comfyui_metadata, str(png_file)
|
||||
)
|
||||
processed_count += 1
|
||||
|
||||
if metadata:
|
||||
self.logger.debug(f"Found metadata in {os.path.basename(png_file)}: {list(metadata.keys())}")
|
||||
|
||||
parsed_data = self._parse_comfyui_prompt(metadata)
|
||||
self.logger.debug(f"Parsed data keys: {list(parsed_data.keys())}, has prompt: {bool(parsed_data.get('prompt'))}, has parameters: {bool(parsed_data.get('parameters'))}")
|
||||
|
||||
if parsed_data.get('prompt') or parsed_data.get('parameters'):
|
||||
found_count += 1
|
||||
prompt_text = self._extract_readable_prompt(parsed_data)
|
||||
|
||||
if prompt_text:
|
||||
self.logger.debug(f"Found prompt in {os.path.basename(png_file)} (type: {type(prompt_text)}): {str(prompt_text)[:100]}...")
|
||||
else:
|
||||
self.logger.debug(f"No readable prompt found in {os.path.basename(png_file)}, parsed_data keys: {list(parsed_data.keys())}")
|
||||
|
||||
if prompt_text and not isinstance(prompt_text, str):
|
||||
self.logger.debug(f"Converting prompt_text from {type(prompt_text)} to string")
|
||||
prompt_text = str(prompt_text)
|
||||
|
||||
if prompt_text and prompt_text.strip():
|
||||
try:
|
||||
prompt_hash = generate_prompt_hash(prompt_text.strip())
|
||||
self.logger.debug(f"Generated hash for prompt: {prompt_hash[:16]}...")
|
||||
|
||||
existing = await self._run_in_executor(
|
||||
self.db.get_prompt_by_hash, prompt_hash
|
||||
)
|
||||
if existing:
|
||||
self.logger.debug(f"Found existing prompt ID {existing['id']} for image {os.path.basename(png_file)}")
|
||||
try:
|
||||
await self._run_in_executor(
|
||||
self.db.link_image_to_prompt, existing['id'], str(png_file)
|
||||
)
|
||||
linked_count += 1
|
||||
self.logger.debug(f"Linked image {os.path.basename(png_file)} to existing prompt {existing['id']}")
|
||||
except Exception as e:
|
||||
self.logger.error(f"Failed to link image {png_file} to existing prompt: {e}")
|
||||
else:
|
||||
self.logger.debug(f"Saving new prompt from {os.path.basename(png_file)}")
|
||||
prompt_id = await self._run_in_executor(
|
||||
self.db.save_prompt,
|
||||
prompt_text.strip(), 'scanned', ['auto-scanned'],
|
||||
f'Auto-scanned from {os.path.basename(png_file)}',
|
||||
prompt_hash
|
||||
)
|
||||
|
||||
if prompt_id:
|
||||
added_count += 1
|
||||
self.logger.info(f"Successfully saved new prompt with ID {prompt_id} from {os.path.basename(png_file)}")
|
||||
try:
|
||||
await self._run_in_executor(
|
||||
self.db.link_image_to_prompt, prompt_id, str(png_file)
|
||||
)
|
||||
self.logger.debug(f"Linked image {os.path.basename(png_file)} to new prompt {prompt_id}")
|
||||
except Exception as e:
|
||||
self.logger.error(f"Failed to link image {png_file} to new prompt: {e}")
|
||||
else:
|
||||
self.logger.error(f"Failed to save prompt from {os.path.basename(png_file)} - no ID returned")
|
||||
|
||||
except Exception as e:
|
||||
self.logger.error(f"Failed to save prompt from {png_file}: {e}")
|
||||
|
||||
# Update progress every 10 files
|
||||
if i % 10 == 0 or i == total_files - 1:
|
||||
progress = int((i + 1) / total_files * 100)
|
||||
status = f"Processing file {i + 1}/{total_files}..."
|
||||
|
||||
yield f"data: {json.dumps({'type': 'progress', 'progress': progress, 'status': status, 'processed': processed_count, 'found': found_count})}\n\n"
|
||||
|
||||
await asyncio.sleep(0.01)
|
||||
|
||||
except Exception as e:
|
||||
self.logger.error(f"Error processing {png_file}: {e}")
|
||||
continue
|
||||
|
||||
self.logger.info(f"Scan completed: processed={processed_count}, found={found_count}, new_prompts_added={added_count}, images_linked_to_existing={linked_count}")
|
||||
yield f"data: {json.dumps({'type': 'complete', 'processed': processed_count, 'found': found_count, 'added': added_count, 'linked': linked_count})}\n\n"
|
||||
|
||||
except Exception as e:
|
||||
self.logger.error(f"Scan error: {e}")
|
||||
self.logger.error(f"Scan error traceback: {traceback.format_exc()}")
|
||||
yield f"data: {json.dumps({'type': 'error', 'message': str(e)})}\n\n"
|
||||
|
||||
response = web.StreamResponse(
|
||||
status=200,
|
||||
reason='OK',
|
||||
headers={
|
||||
'Content-Type': 'text/event-stream',
|
||||
'Cache-Control': 'no-cache',
|
||||
'Connection': 'keep-alive'
|
||||
}
|
||||
)
|
||||
|
||||
await response.prepare(request)
|
||||
|
||||
async for chunk in stream_response():
|
||||
await response.write(chunk.encode('utf-8'))
|
||||
|
||||
await response.write_eof()
|
||||
return response
|
||||
@@ -0,0 +1,489 @@
|
||||
"""AutoTag API routes for PromptManager."""
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
from pathlib import Path
|
||||
|
||||
from aiohttp import web
|
||||
|
||||
|
||||
class AutotagRoutesMixin:
|
||||
"""Mixin providing auto-tagging API endpoints."""
|
||||
|
||||
def _register_autotag_routes(self, routes):
|
||||
@routes.get("/prompt_manager/autotag/models")
|
||||
async def get_autotag_models_route(request):
|
||||
return await self.get_autotag_models(request)
|
||||
|
||||
@routes.get("/prompt_manager/autotag/download/{model_type}")
|
||||
async def download_autotag_model_route(request):
|
||||
return await self.download_autotag_model(request)
|
||||
|
||||
@routes.get("/prompt_manager/autotag/start")
|
||||
async def start_autotag_route(request):
|
||||
return await self.start_autotag(request)
|
||||
|
||||
@routes.post("/prompt_manager/autotag/single")
|
||||
async def autotag_single_route(request):
|
||||
return await self.autotag_single(request)
|
||||
|
||||
@routes.post("/prompt_manager/autotag/apply")
|
||||
async def apply_autotag_route(request):
|
||||
return await self.apply_autotag(request)
|
||||
|
||||
@routes.post("/prompt_manager/autotag/unload")
|
||||
async def unload_autotag_model_route(request):
|
||||
return await self.unload_autotag_model(request)
|
||||
|
||||
@routes.get("/prompt_manager/scan_output_dir")
|
||||
async def scan_output_dir_route(request):
|
||||
return await self.scan_output_dir(request)
|
||||
|
||||
async def get_autotag_models(self, request):
|
||||
"""Get status of available AutoTag models."""
|
||||
try:
|
||||
from ..autotag import get_autotag_service
|
||||
|
||||
service = get_autotag_service()
|
||||
models_status = service.get_models_status()
|
||||
|
||||
return web.json_response({
|
||||
'success': True,
|
||||
'models': models_status,
|
||||
'default_prompt': service.default_prompt,
|
||||
'model_loaded': service.is_model_loaded(),
|
||||
'loaded_model_type': service.get_loaded_model_type()
|
||||
})
|
||||
|
||||
except Exception as e:
|
||||
self.logger.error(f"Get autotag models error: {e}")
|
||||
return web.json_response({
|
||||
'success': False,
|
||||
'error': str(e)
|
||||
}, status=500)
|
||||
|
||||
async def download_autotag_model(self, request):
|
||||
"""Download an AutoTag model with streaming progress."""
|
||||
model_type = request.match_info.get('model_type')
|
||||
|
||||
async def stream_response():
|
||||
try:
|
||||
from ..autotag import get_autotag_service
|
||||
|
||||
service = get_autotag_service()
|
||||
|
||||
if model_type not in service.models_config:
|
||||
yield f"data: {json.dumps({'type': 'error', 'message': f'Invalid model type: {model_type}'})}\n\n"
|
||||
return
|
||||
|
||||
yield f"data: {json.dumps({'type': 'progress', 'progress': 0, 'status': 'Starting download...'})}\n\n"
|
||||
|
||||
progress_data = {'last_progress': 0}
|
||||
|
||||
def progress_callback(status: str, progress: float):
|
||||
progress_data['last_progress'] = progress
|
||||
|
||||
loop = asyncio.get_event_loop()
|
||||
success = await loop.run_in_executor(
|
||||
None,
|
||||
lambda: service.download_model(model_type, progress_callback)
|
||||
)
|
||||
|
||||
if success:
|
||||
yield f"data: {json.dumps({'type': 'complete', 'progress': 100, 'status': 'Download complete'})}\n\n"
|
||||
else:
|
||||
yield f"data: {json.dumps({'type': 'error', 'message': 'Download failed'})}\n\n"
|
||||
|
||||
except Exception as e:
|
||||
self.logger.error(f"Download model error: {e}")
|
||||
yield f"data: {json.dumps({'type': 'error', 'message': str(e)})}\n\n"
|
||||
|
||||
response = web.StreamResponse(
|
||||
status=200,
|
||||
reason='OK',
|
||||
headers={
|
||||
'Content-Type': 'text/event-stream',
|
||||
'Cache-Control': 'no-cache',
|
||||
'Connection': 'keep-alive'
|
||||
}
|
||||
)
|
||||
|
||||
await response.prepare(request)
|
||||
|
||||
async for chunk in stream_response():
|
||||
await response.write(chunk.encode('utf-8'))
|
||||
|
||||
await response.write_eof()
|
||||
return response
|
||||
|
||||
async def start_autotag(self, request):
|
||||
"""Start batch auto-tagging with streaming progress."""
|
||||
model_type = request.query.get('model_type', 'gguf')
|
||||
custom_prompt = request.query.get('prompt', '')
|
||||
skip_tagged = request.query.get('skip_tagged', 'true').lower() == 'true'
|
||||
keep_in_memory = request.query.get('keep_in_memory', 'true').lower() == 'true'
|
||||
use_gpu = True
|
||||
|
||||
async def stream_response():
|
||||
try:
|
||||
from ..autotag import get_autotag_service
|
||||
import time as _time
|
||||
|
||||
service = get_autotag_service()
|
||||
|
||||
status = service.get_models_status()
|
||||
if not status.get(model_type, {}).get('downloaded'):
|
||||
yield f"data: {json.dumps({'type': 'error', 'message': f'Model {model_type} not downloaded'})}\n\n"
|
||||
return
|
||||
|
||||
yield f"data: {json.dumps({'type': 'progress', 'progress': 0, 'status': 'Loading model...'})}\n\n"
|
||||
|
||||
loop = asyncio.get_event_loop()
|
||||
try:
|
||||
await loop.run_in_executor(
|
||||
None,
|
||||
lambda: service.load_model(model_type, use_gpu)
|
||||
)
|
||||
except Exception as e:
|
||||
yield f"data: {json.dumps({'type': 'error', 'message': f'Failed to load model: {str(e)}'})}\n\n"
|
||||
return
|
||||
|
||||
if custom_prompt:
|
||||
service.custom_prompt = custom_prompt
|
||||
|
||||
yield f"data: {json.dumps({'type': 'progress', 'progress': 5, 'status': 'Model loaded. Fetching all images from database...'})}\n\n"
|
||||
|
||||
images = await self._run_in_executor(self.db.get_all_images)
|
||||
|
||||
total_files = len(images)
|
||||
if total_files == 0:
|
||||
yield f"data: {json.dumps({'type': 'complete', 'processed': 0, 'tagged': 0, 'skipped': 0, 'status': 'No images with linked prompts found'})}\n\n"
|
||||
service.unload_model()
|
||||
return
|
||||
|
||||
yield f"data: {json.dumps({'type': 'progress', 'progress': 10, 'status': f'Found {total_files} images. Processing...'})}\n\n"
|
||||
|
||||
processed = 0
|
||||
tagged = 0
|
||||
skipped = 0
|
||||
errors = 0
|
||||
tagged_prompt_ids = set()
|
||||
|
||||
last_update_time = _time.monotonic()
|
||||
|
||||
for i, image_data in enumerate(images):
|
||||
image_path = image_data.get('image_path')
|
||||
prompt_id = image_data.get('prompt_id')
|
||||
|
||||
if not image_path or not prompt_id:
|
||||
skipped += 1
|
||||
continue
|
||||
|
||||
if prompt_id in tagged_prompt_ids:
|
||||
skipped += 1
|
||||
now = _time.monotonic()
|
||||
if (now - last_update_time) >= 0.5 or i == total_files - 1:
|
||||
progress = 10 + int((i + 1) / total_files * 85)
|
||||
yield f"data: {json.dumps({'type': 'progress', 'progress': progress, 'status': f'Skipping {i+1}/{total_files} (prompt already processed)...', 'processed': processed, 'tagged': tagged, 'skipped': skipped})}\n\n"
|
||||
await asyncio.sleep(0.01)
|
||||
last_update_time = now
|
||||
continue
|
||||
|
||||
if not Path(image_path).exists():
|
||||
skipped += 1
|
||||
now = _time.monotonic()
|
||||
if (now - last_update_time) >= 0.5 or i == total_files - 1:
|
||||
progress = 10 + int((i + 1) / total_files * 85)
|
||||
yield f"data: {json.dumps({'type': 'progress', 'progress': progress, 'status': f'Skipping {i+1}/{total_files} (file missing)...', 'processed': processed, 'tagged': tagged, 'skipped': skipped})}\n\n"
|
||||
await asyncio.sleep(0.01)
|
||||
last_update_time = now
|
||||
continue
|
||||
|
||||
if skip_tagged:
|
||||
prompt_tags = image_data.get('prompt_tags', [])
|
||||
if isinstance(prompt_tags, str):
|
||||
prompt_tags = [t.strip() for t in prompt_tags.split(',') if t.strip()]
|
||||
real_tags = [t for t in prompt_tags if t != 'auto-scanned']
|
||||
if real_tags:
|
||||
tagged_prompt_ids.add(prompt_id)
|
||||
skipped += 1
|
||||
now = _time.monotonic()
|
||||
if (now - last_update_time) >= 0.5 or i == total_files - 1:
|
||||
progress = 10 + int((i + 1) / total_files * 85)
|
||||
yield f"data: {json.dumps({'type': 'progress', 'progress': progress, 'status': f'Skipping {i+1}/{total_files} (already tagged)...', 'processed': processed, 'tagged': tagged, 'skipped': skipped})}\n\n"
|
||||
await asyncio.sleep(0.01)
|
||||
last_update_time = now
|
||||
continue
|
||||
|
||||
try:
|
||||
tags = await loop.run_in_executor(
|
||||
None,
|
||||
lambda p=str(image_path): service.generate_tags(p)
|
||||
)
|
||||
|
||||
processed += 1
|
||||
|
||||
if tags:
|
||||
existing_prompt = await self._run_in_executor(self.db.get_prompt_by_id, prompt_id)
|
||||
if existing_prompt:
|
||||
existing_tags = existing_prompt.get('tags', [])
|
||||
if isinstance(existing_tags, str):
|
||||
existing_tags = [t.strip() for t in existing_tags.split(',') if t.strip()]
|
||||
|
||||
new_tags = [t for t in tags if t not in existing_tags]
|
||||
if new_tags:
|
||||
all_tags = existing_tags + new_tags
|
||||
await self._run_in_executor(
|
||||
self.db.update_prompt_metadata,
|
||||
prompt_id,
|
||||
tags=all_tags
|
||||
)
|
||||
tagged += 1
|
||||
else:
|
||||
skipped += 1
|
||||
else:
|
||||
skipped += 1
|
||||
else:
|
||||
skipped += 1
|
||||
|
||||
tagged_prompt_ids.add(prompt_id)
|
||||
|
||||
progress = 10 + int((i + 1) / total_files * 85)
|
||||
yield f"data: {json.dumps({'type': 'progress', 'progress': progress, 'status': f'Processing {i+1}/{total_files}...', 'processed': processed, 'tagged': tagged, 'skipped': skipped})}\n\n"
|
||||
await asyncio.sleep(0.01)
|
||||
last_update_time = _time.monotonic()
|
||||
|
||||
except Exception as img_err:
|
||||
self.logger.error(f"Error processing {image_path}: {img_err}")
|
||||
errors += 1
|
||||
processed += 1
|
||||
|
||||
if not keep_in_memory:
|
||||
service.unload_model()
|
||||
model_status = 'Model unloaded'
|
||||
else:
|
||||
model_status = 'Model kept in memory'
|
||||
|
||||
yield f"data: {json.dumps({'type': 'complete', 'progress': 100, 'processed': processed, 'tagged': tagged, 'skipped': skipped, 'errors': errors, 'status': 'Complete', 'model_status': model_status, 'model_loaded': keep_in_memory})}\n\n"
|
||||
|
||||
except Exception as e:
|
||||
self.logger.error(f"AutoTag error: {e}")
|
||||
import traceback
|
||||
self.logger.error(f"AutoTag traceback: {traceback.format_exc()}")
|
||||
yield f"data: {json.dumps({'type': 'error', 'message': str(e)})}\n\n"
|
||||
|
||||
response = web.StreamResponse(
|
||||
status=200,
|
||||
reason='OK',
|
||||
headers={
|
||||
'Content-Type': 'text/event-stream',
|
||||
'Cache-Control': 'no-cache',
|
||||
'Connection': 'keep-alive'
|
||||
}
|
||||
)
|
||||
|
||||
await response.prepare(request)
|
||||
|
||||
async for chunk in stream_response():
|
||||
await response.write(chunk.encode('utf-8'))
|
||||
|
||||
await response.write_eof()
|
||||
return response
|
||||
|
||||
async def autotag_single(self, request):
|
||||
"""Generate tags for a single image."""
|
||||
try:
|
||||
data = await request.json()
|
||||
image_path = data.get('image_path')
|
||||
model_type = data.get('model_type', 'gguf')
|
||||
custom_prompt = data.get('prompt')
|
||||
use_gpu = data.get('use_gpu', True)
|
||||
|
||||
if not image_path:
|
||||
return web.json_response({
|
||||
'success': False,
|
||||
'error': 'image_path is required'
|
||||
}, status=400)
|
||||
|
||||
from ..autotag import get_autotag_service
|
||||
|
||||
service = get_autotag_service()
|
||||
|
||||
if not service.is_model_loaded() or service.get_loaded_model_type() != model_type:
|
||||
loop = asyncio.get_event_loop()
|
||||
await loop.run_in_executor(
|
||||
None,
|
||||
lambda: service.load_model(model_type, use_gpu)
|
||||
)
|
||||
|
||||
if custom_prompt:
|
||||
service.custom_prompt = custom_prompt
|
||||
|
||||
loop = asyncio.get_event_loop()
|
||||
tags = await loop.run_in_executor(
|
||||
None,
|
||||
lambda: service.generate_tags(image_path)
|
||||
)
|
||||
|
||||
prompt_id = None
|
||||
try:
|
||||
prompt_id = await self._run_in_executor(self.db.get_prompt_id_for_image, image_path)
|
||||
except Exception as e:
|
||||
self.logger.warning(f"Could not find linked prompt: {e}")
|
||||
|
||||
return web.json_response({
|
||||
'success': True,
|
||||
'tags': tags,
|
||||
'prompt_id': prompt_id,
|
||||
'image_path': image_path
|
||||
})
|
||||
|
||||
except Exception as e:
|
||||
self.logger.error(f"AutoTag single error: {e}")
|
||||
return web.json_response({
|
||||
'success': False,
|
||||
'error': str(e)
|
||||
}, status=500)
|
||||
|
||||
async def apply_autotag(self, request):
|
||||
"""Apply selected tags to a prompt."""
|
||||
try:
|
||||
data = await request.json()
|
||||
prompt_id = data.get('prompt_id')
|
||||
tags = data.get('tags', [])
|
||||
|
||||
if not prompt_id:
|
||||
return web.json_response({
|
||||
'success': False,
|
||||
'error': 'prompt_id is required'
|
||||
}, status=400)
|
||||
|
||||
if not tags:
|
||||
return web.json_response({
|
||||
'success': True,
|
||||
'message': 'No tags to apply'
|
||||
})
|
||||
|
||||
prompt = await self._run_in_executor(self.db.get_prompt_by_id, prompt_id)
|
||||
if not prompt:
|
||||
return web.json_response({
|
||||
'success': False,
|
||||
'error': f'Prompt {prompt_id} not found'
|
||||
}, status=404)
|
||||
|
||||
existing_tags = prompt.get('tags', [])
|
||||
if isinstance(existing_tags, str):
|
||||
existing_tags = [t.strip() for t in existing_tags.split(',') if t.strip()]
|
||||
|
||||
new_tags = [t for t in tags if t not in existing_tags]
|
||||
all_tags = existing_tags + new_tags
|
||||
|
||||
await self._run_in_executor(
|
||||
self.db.update_prompt_metadata,
|
||||
prompt_id,
|
||||
tags=all_tags
|
||||
)
|
||||
|
||||
return web.json_response({
|
||||
'success': True,
|
||||
'added_tags': new_tags,
|
||||
'total_tags': len(all_tags)
|
||||
})
|
||||
|
||||
except Exception as e:
|
||||
self.logger.error(f"Apply autotag error: {e}")
|
||||
return web.json_response({
|
||||
'success': False,
|
||||
'error': str(e)
|
||||
}, status=500)
|
||||
|
||||
async def unload_autotag_model(self, request):
|
||||
"""Manually unload the AutoTag model from memory."""
|
||||
try:
|
||||
from ..autotag import get_autotag_service
|
||||
|
||||
service = get_autotag_service()
|
||||
|
||||
if not service.is_model_loaded():
|
||||
return web.json_response({
|
||||
'success': True,
|
||||
'message': 'No model was loaded'
|
||||
})
|
||||
|
||||
model_type = service.get_loaded_model_type()
|
||||
service.unload_model()
|
||||
|
||||
return web.json_response({
|
||||
'success': True,
|
||||
'message': f'{model_type.upper()} model unloaded successfully',
|
||||
'model_loaded': False
|
||||
})
|
||||
|
||||
except Exception as e:
|
||||
self.logger.error(f"Unload autotag model error: {e}")
|
||||
return web.json_response({
|
||||
'success': False,
|
||||
'error': str(e)
|
||||
}, status=500)
|
||||
|
||||
async def scan_output_dir(self, request):
|
||||
"""Scan ComfyUI output directory for images."""
|
||||
try:
|
||||
output_dir = self._find_comfyui_output_dir()
|
||||
if not output_dir:
|
||||
return web.json_response({
|
||||
'success': False,
|
||||
'error': 'ComfyUI output directory not found'
|
||||
}, status=404)
|
||||
|
||||
output_path = Path(output_dir)
|
||||
image_extensions = ['.png', '.jpg', '.jpeg', '.gif', '.webp', '.bmp', '.tiff']
|
||||
|
||||
images = []
|
||||
seen_paths = set()
|
||||
|
||||
for ext in image_extensions:
|
||||
for pattern in [f"*{ext.lower()}", f"*{ext.upper()}"]:
|
||||
for image_path in output_path.rglob(pattern):
|
||||
if 'thumbnails' not in image_path.parts:
|
||||
normalized_path = str(image_path).lower()
|
||||
if normalized_path not in seen_paths:
|
||||
seen_paths.add(normalized_path)
|
||||
|
||||
rel_path = image_path.relative_to(output_path)
|
||||
|
||||
thumbnail_url = None
|
||||
thumbnails_dir = output_path / "thumbnails"
|
||||
if thumbnails_dir.exists():
|
||||
rel_path_no_ext = rel_path.with_suffix('')
|
||||
thumbnail_rel_path = f"thumbnails/{rel_path_no_ext.as_posix()}_thumb{image_path.suffix}"
|
||||
thumbnail_abs_path = thumbnails_dir / f"{rel_path_no_ext.as_posix()}_thumb{image_path.suffix}"
|
||||
if thumbnail_abs_path.exists():
|
||||
from urllib.parse import quote
|
||||
thumbnail_url = f'/prompt_manager/images/serve/{quote(thumbnail_rel_path, safe="/")}'
|
||||
|
||||
from urllib.parse import quote as url_quote
|
||||
images.append({
|
||||
'filename': image_path.name,
|
||||
'path': str(image_path),
|
||||
'relative_path': str(rel_path),
|
||||
'url': f'/prompt_manager/images/serve/{url_quote(rel_path.as_posix(), safe="/")}',
|
||||
'thumbnail_url': thumbnail_url
|
||||
})
|
||||
|
||||
images.sort(key=lambda x: x['filename'])
|
||||
|
||||
self.logger.info(f"Found {len(images)} images in output directory")
|
||||
|
||||
return web.json_response({
|
||||
'success': True,
|
||||
'images': images,
|
||||
'count': len(images)
|
||||
})
|
||||
|
||||
except Exception as e:
|
||||
self.logger.error(f"Scan output dir error: {e}")
|
||||
return web.json_response({
|
||||
'success': False,
|
||||
'error': str(e)
|
||||
}, status=500)
|
||||
+1045
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,224 @@
|
||||
"""Logging API routes for PromptManager."""
|
||||
|
||||
import os
|
||||
|
||||
from aiohttp import web
|
||||
|
||||
|
||||
class LoggingRoutesMixin:
|
||||
"""Mixin providing logging-related API endpoints."""
|
||||
|
||||
def _register_logging_routes(self, routes):
|
||||
@routes.get("/prompt_manager/logs")
|
||||
async def get_logs_route(request):
|
||||
return await self.get_logs(request)
|
||||
|
||||
@routes.get("/prompt_manager/logs/files")
|
||||
async def get_log_files_route(request):
|
||||
return await self.get_log_files(request)
|
||||
|
||||
@routes.get("/prompt_manager/logs/download/{filename}")
|
||||
async def download_log_route(request):
|
||||
return await self.download_log_file(request)
|
||||
|
||||
@routes.post("/prompt_manager/logs/truncate")
|
||||
async def truncate_logs_route(request):
|
||||
return await self.truncate_logs(request)
|
||||
|
||||
@routes.get("/prompt_manager/logs/config")
|
||||
async def get_log_config_route(request):
|
||||
return await self.get_log_config(request)
|
||||
|
||||
@routes.post("/prompt_manager/logs/config")
|
||||
async def update_log_config_route(request):
|
||||
return await self.update_log_config(request)
|
||||
|
||||
@routes.get("/prompt_manager/logs/stats")
|
||||
async def get_log_stats_route(request):
|
||||
return await self.get_log_stats(request)
|
||||
|
||||
def _get_logger_manager(self):
|
||||
"""Get the logger manager instance."""
|
||||
try:
|
||||
from ...utils.logging_config import get_logger_manager
|
||||
except ImportError:
|
||||
import sys
|
||||
current_dir = os.path.dirname(os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
|
||||
sys.path.insert(0, current_dir)
|
||||
from utils.logging_config import get_logger_manager
|
||||
return get_logger_manager()
|
||||
|
||||
async def get_logs(self, request):
|
||||
"""Get recent log entries."""
|
||||
try:
|
||||
logger_manager = self._get_logger_manager()
|
||||
|
||||
limit = int(request.query.get('limit', 100))
|
||||
level = request.query.get('level', None)
|
||||
|
||||
if limit > 1000:
|
||||
limit = 1000
|
||||
elif limit < 1:
|
||||
limit = 1
|
||||
|
||||
logs = logger_manager.get_recent_logs(limit=limit, level=level)
|
||||
|
||||
return web.json_response({
|
||||
'success': True,
|
||||
'logs': logs,
|
||||
'count': len(logs),
|
||||
'level_filter': level,
|
||||
'limit': limit
|
||||
})
|
||||
|
||||
except Exception as e:
|
||||
self.logger.error(f"Get logs error: {e}")
|
||||
return web.json_response({
|
||||
'success': False,
|
||||
'error': str(e),
|
||||
'logs': []
|
||||
}, status=500)
|
||||
|
||||
async def get_log_files(self, request):
|
||||
"""Get information about available log files."""
|
||||
try:
|
||||
logger_manager = self._get_logger_manager()
|
||||
log_files = logger_manager.get_log_files()
|
||||
|
||||
return web.json_response({
|
||||
'success': True,
|
||||
'files': log_files,
|
||||
'count': len(log_files)
|
||||
})
|
||||
|
||||
except Exception as e:
|
||||
self.logger.error(f"Get log files error: {e}")
|
||||
return web.json_response({
|
||||
'success': False,
|
||||
'error': str(e),
|
||||
'files': []
|
||||
}, status=500)
|
||||
|
||||
async def download_log_file(self, request):
|
||||
"""Download a specific log file."""
|
||||
try:
|
||||
filename = request.match_info['filename']
|
||||
logger_manager = self._get_logger_manager()
|
||||
|
||||
if not filename or '..' in filename or '/' in filename or '\\' in filename:
|
||||
return web.json_response({
|
||||
'success': False,
|
||||
'error': 'Invalid filename'
|
||||
}, status=400)
|
||||
|
||||
log_file_path = logger_manager.log_dir / filename
|
||||
|
||||
if not log_file_path.exists():
|
||||
return web.json_response({
|
||||
'success': False,
|
||||
'error': 'Log file not found'
|
||||
}, status=404)
|
||||
|
||||
with open(log_file_path, 'rb') as f:
|
||||
file_content = f.read()
|
||||
|
||||
return web.Response(
|
||||
body=file_content,
|
||||
content_type='text/plain',
|
||||
headers={
|
||||
'Content-Disposition': f'attachment; filename="{filename}"',
|
||||
'Content-Length': str(len(file_content))
|
||||
}
|
||||
)
|
||||
|
||||
except Exception as e:
|
||||
self.logger.error(f"Download log file error: {e}")
|
||||
return web.json_response({
|
||||
'success': False,
|
||||
'error': str(e)
|
||||
}, status=500)
|
||||
|
||||
async def truncate_logs(self, request):
|
||||
"""Truncate all log files."""
|
||||
try:
|
||||
logger_manager = self._get_logger_manager()
|
||||
results = logger_manager.truncate_logs()
|
||||
|
||||
return web.json_response({
|
||||
'success': True,
|
||||
'message': f"Truncated {len(results['truncated'])} log files",
|
||||
'results': results
|
||||
})
|
||||
|
||||
except Exception as e:
|
||||
self.logger.error(f"Truncate logs error: {e}")
|
||||
return web.json_response({
|
||||
'success': False,
|
||||
'error': str(e)
|
||||
}, status=500)
|
||||
|
||||
async def get_log_config(self, request):
|
||||
"""Get current logging configuration."""
|
||||
try:
|
||||
logger_manager = self._get_logger_manager()
|
||||
config = logger_manager.get_config()
|
||||
|
||||
return web.json_response({
|
||||
'success': True,
|
||||
'config': config
|
||||
})
|
||||
|
||||
except Exception as e:
|
||||
self.logger.error(f"Get log config error: {e}")
|
||||
return web.json_response({
|
||||
'success': False,
|
||||
'error': str(e)
|
||||
}, status=500)
|
||||
|
||||
async def update_log_config(self, request):
|
||||
"""Update logging configuration."""
|
||||
try:
|
||||
data = await request.json()
|
||||
logger_manager = self._get_logger_manager()
|
||||
|
||||
if 'level' in data:
|
||||
valid_levels = ['DEBUG', 'INFO', 'WARNING', 'ERROR', 'CRITICAL']
|
||||
if data['level'].upper() not in valid_levels:
|
||||
return web.json_response({
|
||||
'success': False,
|
||||
'error': f'Invalid log level. Must be one of: {valid_levels}'
|
||||
}, status=400)
|
||||
data['level'] = data['level'].upper()
|
||||
|
||||
logger_manager.update_config(data)
|
||||
|
||||
return web.json_response({
|
||||
'success': True,
|
||||
'message': 'Logging configuration updated',
|
||||
'config': logger_manager.get_config()
|
||||
})
|
||||
|
||||
except Exception as e:
|
||||
self.logger.error(f"Update log config error: {e}")
|
||||
return web.json_response({
|
||||
'success': False,
|
||||
'error': str(e)
|
||||
}, status=500)
|
||||
|
||||
async def get_log_stats(self, request):
|
||||
"""Get logging statistics."""
|
||||
try:
|
||||
logger_manager = self._get_logger_manager()
|
||||
stats = logger_manager.get_log_stats()
|
||||
|
||||
return web.json_response({
|
||||
'success': True,
|
||||
'stats': stats
|
||||
})
|
||||
|
||||
except Exception as e:
|
||||
self.logger.error(f"Get log stats error: {e}")
|
||||
return web.json_response({
|
||||
'success': False,
|
||||
'error': str(e)
|
||||
}, status=500)
|
||||
@@ -0,0 +1,860 @@
|
||||
"""Prompt API routes for PromptManager."""
|
||||
|
||||
import datetime
|
||||
import json
|
||||
|
||||
from aiohttp import web
|
||||
|
||||
|
||||
class PromptRoutesMixin:
|
||||
"""Mixin providing prompt-related API endpoints."""
|
||||
|
||||
def _register_prompt_routes(self, routes):
|
||||
@routes.get("/prompt_manager/search")
|
||||
async def search_prompts_route(request):
|
||||
return await self.search_prompts(request)
|
||||
|
||||
@routes.get("/prompt_manager/recent")
|
||||
async def get_recent_prompts_route(request):
|
||||
return await self.get_recent_prompts(request)
|
||||
|
||||
@routes.get("/prompt_manager/categories")
|
||||
async def get_categories_route(request):
|
||||
return await self.get_categories(request)
|
||||
|
||||
# Tag management endpoints (must be registered BEFORE /prompt_manager/tags)
|
||||
@routes.get("/prompt_manager/tags/stats")
|
||||
async def get_tags_stats_route(request):
|
||||
return await self.get_tags_stats(request)
|
||||
|
||||
@routes.get("/prompt_manager/tags/filter")
|
||||
async def get_tags_filter_route(request):
|
||||
return await self.get_tags_filter(request)
|
||||
|
||||
# Bulk tag operations (register BEFORE {tag_name} to avoid path param match)
|
||||
@routes.post("/prompt_manager/tags/merge")
|
||||
async def merge_tags_route(request):
|
||||
return await self.merge_tags_endpoint(request)
|
||||
|
||||
@routes.put("/prompt_manager/tags/{tag_name}")
|
||||
async def rename_tag_route(request):
|
||||
return await self.rename_tag_endpoint(request)
|
||||
|
||||
@routes.delete("/prompt_manager/tags/{tag_name}")
|
||||
async def delete_tag_route(request):
|
||||
return await self.delete_tag_endpoint(request)
|
||||
|
||||
@routes.get("/prompt_manager/tags/{tag_name}/prompts")
|
||||
async def get_tag_prompts_route(request):
|
||||
return await self.get_tag_prompts(request)
|
||||
|
||||
@routes.get("/prompt_manager/tags")
|
||||
async def get_tags_route(request):
|
||||
return await self.get_tags(request)
|
||||
|
||||
@routes.post("/prompt_manager/save")
|
||||
async def save_prompt_route(request):
|
||||
return await self.save_prompt(request)
|
||||
|
||||
@routes.delete("/prompt_manager/delete/{prompt_id}")
|
||||
async def delete_prompt_route(request):
|
||||
return await self.delete_prompt(request)
|
||||
|
||||
# Individual prompt management
|
||||
@routes.put("/prompt_manager/prompts/{prompt_id}")
|
||||
async def update_prompt_route(request):
|
||||
return await self.update_prompt(request)
|
||||
|
||||
@routes.put("/prompt_manager/prompts/{prompt_id}/rating")
|
||||
async def update_rating_route(request):
|
||||
return await self.update_prompt_rating(request)
|
||||
|
||||
@routes.post("/prompt_manager/prompts/{prompt_id}/tags")
|
||||
async def add_tag_route(request):
|
||||
return await self.add_prompt_tag(request)
|
||||
|
||||
@routes.delete("/prompt_manager/prompts/{prompt_id}/tags")
|
||||
async def remove_tag_route(request):
|
||||
return await self.remove_prompt_tag(request)
|
||||
|
||||
@routes.post("/prompt_manager/prompts/tags")
|
||||
async def add_tags_to_prompt_route(request):
|
||||
return await self.add_tags_to_prompt(request)
|
||||
|
||||
# Bulk operations
|
||||
@routes.post("/prompt_manager/bulk/delete")
|
||||
async def bulk_delete_route(request):
|
||||
return await self.bulk_delete_prompts(request)
|
||||
|
||||
@routes.post("/prompt_manager/bulk/tags")
|
||||
async def bulk_add_tags_route(request):
|
||||
return await self.bulk_add_tags(request)
|
||||
|
||||
@routes.post("/prompt_manager/bulk/category")
|
||||
async def bulk_set_category_route(request):
|
||||
return await self.bulk_set_category(request)
|
||||
|
||||
# Export functionality
|
||||
@routes.get("/prompt_manager/export")
|
||||
async def export_prompts_route(request):
|
||||
return await self.export_prompts(request)
|
||||
|
||||
async def search_prompts(self, request):
|
||||
"""Search for prompts using multiple filter criteria."""
|
||||
try:
|
||||
text = request.query.get("text", "").strip()
|
||||
category = request.query.get("category", "").strip()
|
||||
tags_str = request.query.get("tags", "").strip()
|
||||
min_rating = request.query.get("min_rating", 0)
|
||||
limit = int(request.query.get("limit", 50))
|
||||
|
||||
tags = None
|
||||
if tags_str:
|
||||
tags = [tag.strip() for tag in tags_str.split(",") if tag.strip()]
|
||||
|
||||
try:
|
||||
min_rating = int(min_rating) if min_rating else None
|
||||
except ValueError:
|
||||
min_rating = None
|
||||
|
||||
results = await self._run_in_executor(
|
||||
self.db.search_prompts,
|
||||
text=text if text else None,
|
||||
category=category if category else None,
|
||||
tags=tags,
|
||||
rating_min=min_rating,
|
||||
limit=limit,
|
||||
)
|
||||
self._enrich_prompt_images(results)
|
||||
|
||||
return web.json_response(
|
||||
{"success": True, "results": results, "count": len(results)}
|
||||
)
|
||||
|
||||
except Exception as e:
|
||||
self.logger.error(f"Search error: {e}", exc_info=True)
|
||||
return web.json_response(
|
||||
{"success": False, "error": f"Search failed: {str(e)}", "results": []},
|
||||
status=500,
|
||||
)
|
||||
|
||||
async def get_recent_prompts(self, request):
|
||||
"""Retrieve recently created prompts with pagination support."""
|
||||
try:
|
||||
limit = int(request.query.get("limit", 50))
|
||||
page = int(request.query.get("page", 1))
|
||||
offset = int(request.query.get("offset", 0))
|
||||
|
||||
if page > 1 and offset == 0:
|
||||
offset = (page - 1) * limit
|
||||
|
||||
if limit > 1000:
|
||||
limit = 1000
|
||||
elif limit < 1:
|
||||
limit = 1
|
||||
|
||||
results = await self._run_in_executor(self.db.get_recent_prompts, limit=limit, offset=offset)
|
||||
self._enrich_prompt_images(results['prompts'])
|
||||
|
||||
return web.json_response({
|
||||
"success": True,
|
||||
"results": results['prompts'],
|
||||
"pagination": {
|
||||
"total": results['total'],
|
||||
"limit": results['limit'],
|
||||
"offset": results['offset'],
|
||||
"page": results['page'],
|
||||
"total_pages": results['total_pages'],
|
||||
"has_more": results['has_more'],
|
||||
"count": len(results['prompts'])
|
||||
}
|
||||
})
|
||||
|
||||
except Exception as e:
|
||||
self.logger.error(f"Recent prompts error: {e}", exc_info=True)
|
||||
return web.json_response(
|
||||
{
|
||||
"success": False,
|
||||
"error": f"Failed to get recent prompts: {str(e)}",
|
||||
"results": [],
|
||||
"pagination": {"total": 0, "page": 1, "total_pages": 0}
|
||||
},
|
||||
status=500,
|
||||
)
|
||||
|
||||
async def get_categories(self, request):
|
||||
"""Retrieve all available prompt categories."""
|
||||
try:
|
||||
categories = await self._run_in_executor(self.db.get_all_categories)
|
||||
return web.json_response({"success": True, "categories": categories})
|
||||
except Exception as e:
|
||||
self.logger.error(f"Categories error: {e}")
|
||||
return web.json_response(
|
||||
{"success": False, "error": f"Failed to get categories: {str(e)}", "categories": []},
|
||||
status=500,
|
||||
)
|
||||
|
||||
async def get_tags(self, request):
|
||||
"""Retrieve all available prompt tags."""
|
||||
try:
|
||||
tags = await self._run_in_executor(self.db.get_all_tags)
|
||||
return web.json_response({"success": True, "tags": tags})
|
||||
except Exception as e:
|
||||
self.logger.error(f"Tags error: {e}")
|
||||
return web.json_response(
|
||||
{"success": False, "error": f"Failed to get tags: {str(e)}", "tags": []},
|
||||
status=500,
|
||||
)
|
||||
|
||||
async def get_tags_stats(self, request):
|
||||
"""Get tags with usage counts, search, sort, and pagination."""
|
||||
try:
|
||||
try:
|
||||
limit = int(request.query.get("limit", 50))
|
||||
offset = int(request.query.get("offset", 0))
|
||||
except (ValueError, TypeError):
|
||||
return web.json_response(
|
||||
{"success": False, "error": "Invalid limit or offset parameter"}, status=400
|
||||
)
|
||||
search = request.query.get("search", "").strip() or None
|
||||
sort = request.query.get("sort", "alpha_asc")
|
||||
|
||||
result = await self._run_in_executor(self.db.get_tags_with_counts, limit, offset, search, sort)
|
||||
untagged_count = await self._run_in_executor(self.db.get_untagged_prompts_count)
|
||||
|
||||
return web.json_response({
|
||||
"success": True,
|
||||
"tags": result['tags'],
|
||||
"untagged_count": untagged_count,
|
||||
"pagination": {
|
||||
"total": result['total'],
|
||||
"limit": result['limit'],
|
||||
"offset": result['offset'],
|
||||
"has_more": result['has_more']
|
||||
}
|
||||
})
|
||||
except Exception as e:
|
||||
self.logger.error(f"Tags stats error: {e}", exc_info=True)
|
||||
return web.json_response({"success": False, "error": str(e)}, status=500)
|
||||
|
||||
async def get_tag_prompts(self, request):
|
||||
"""Get prompts for a single tag."""
|
||||
try:
|
||||
from urllib.parse import unquote
|
||||
tag_name = unquote(request.match_info.get("tag_name", ""))
|
||||
if not tag_name:
|
||||
return web.json_response(
|
||||
{"success": False, "error": "Tag name required"}, status=400
|
||||
)
|
||||
|
||||
try:
|
||||
limit = int(request.query.get("limit", 20))
|
||||
offset = int(request.query.get("offset", 0))
|
||||
except (ValueError, TypeError):
|
||||
return web.json_response(
|
||||
{"success": False, "error": "Invalid limit or offset parameter"}, status=400
|
||||
)
|
||||
|
||||
result = await self._run_in_executor(self.db.get_prompts_by_tags, [tag_name], 'and', limit, offset)
|
||||
self._enrich_prompt_images(result['prompts'])
|
||||
|
||||
return web.json_response({
|
||||
"success": True,
|
||||
"tag": tag_name,
|
||||
"prompts": result['prompts'],
|
||||
"pagination": {
|
||||
"total": result['total'],
|
||||
"limit": result['limit'],
|
||||
"offset": result['offset'],
|
||||
"has_more": result['has_more']
|
||||
}
|
||||
})
|
||||
except Exception as e:
|
||||
self.logger.error(f"Tag prompts error: {e}", exc_info=True)
|
||||
return web.json_response({"success": False, "error": str(e)}, status=500)
|
||||
|
||||
async def get_tags_filter(self, request):
|
||||
"""Get prompts matching multiple tags with AND/OR mode, or untagged prompts."""
|
||||
try:
|
||||
untagged = request.query.get("untagged", "").lower() == "true"
|
||||
|
||||
if untagged:
|
||||
try:
|
||||
limit = int(request.query.get("limit", 20))
|
||||
offset = int(request.query.get("offset", 0))
|
||||
except (ValueError, TypeError):
|
||||
return web.json_response(
|
||||
{"success": False, "error": "Invalid limit or offset parameter"}, status=400
|
||||
)
|
||||
result = await self._run_in_executor(self.db.get_untagged_prompts, limit, offset)
|
||||
self._enrich_prompt_images(result['prompts'])
|
||||
return web.json_response({
|
||||
"success": True,
|
||||
"tags": [],
|
||||
"mode": "untagged",
|
||||
"prompts": result['prompts'],
|
||||
"pagination": {
|
||||
"total": result['total'],
|
||||
"limit": result['limit'],
|
||||
"offset": result['offset'],
|
||||
"has_more": result['has_more']
|
||||
}
|
||||
})
|
||||
|
||||
tags_str = request.query.get("tags", "").strip()
|
||||
if not tags_str:
|
||||
return web.json_response(
|
||||
{"success": False, "error": "Tags parameter required"}, status=400
|
||||
)
|
||||
|
||||
tags_list = [t.strip() for t in tags_str.split(",") if t.strip()]
|
||||
mode = request.query.get("mode", "and").lower()
|
||||
if mode not in ("and", "or"):
|
||||
mode = "and"
|
||||
|
||||
try:
|
||||
limit = int(request.query.get("limit", 20))
|
||||
offset = int(request.query.get("offset", 0))
|
||||
except (ValueError, TypeError):
|
||||
return web.json_response(
|
||||
{"success": False, "error": "Invalid limit or offset parameter"}, status=400
|
||||
)
|
||||
|
||||
result = await self._run_in_executor(self.db.get_prompts_by_tags, tags_list, mode, limit, offset)
|
||||
self._enrich_prompt_images(result['prompts'])
|
||||
|
||||
return web.json_response({
|
||||
"success": True,
|
||||
"tags": tags_list,
|
||||
"mode": mode,
|
||||
"prompts": result['prompts'],
|
||||
"pagination": {
|
||||
"total": result['total'],
|
||||
"limit": result['limit'],
|
||||
"offset": result['offset'],
|
||||
"has_more": result['has_more']
|
||||
}
|
||||
})
|
||||
except Exception as e:
|
||||
self.logger.error(f"Tags filter error: {e}", exc_info=True)
|
||||
return web.json_response({"success": False, "error": str(e)}, status=500)
|
||||
|
||||
async def rename_tag_endpoint(self, request):
|
||||
"""Rename a tag across all prompts."""
|
||||
try:
|
||||
from urllib.parse import unquote
|
||||
tag_name = unquote(request.match_info.get("tag_name", ""))
|
||||
if not tag_name:
|
||||
return web.json_response(
|
||||
{"success": False, "error": "Tag name required"}, status=400
|
||||
)
|
||||
|
||||
try:
|
||||
body = await request.json()
|
||||
except Exception:
|
||||
return web.json_response(
|
||||
{"success": False, "error": "Invalid JSON body"}, status=400
|
||||
)
|
||||
new_name = body.get("new_name", "").strip()
|
||||
if not new_name:
|
||||
return web.json_response(
|
||||
{"success": False, "error": "New tag name required"}, status=400
|
||||
)
|
||||
|
||||
result = await self._run_in_executor(self.db.rename_tag_all_prompts, tag_name, new_name)
|
||||
resp = {
|
||||
"success": True,
|
||||
"old_name": tag_name,
|
||||
"new_name": new_name,
|
||||
"affected_count": result['affected_count']
|
||||
}
|
||||
if result.get('skipped_count', 0) > 0:
|
||||
resp['skipped_count'] = result['skipped_count']
|
||||
resp['warning'] = f"{result['skipped_count']} prompt(s) had corrupted tag data and were skipped"
|
||||
return web.json_response(resp)
|
||||
except Exception as e:
|
||||
self.logger.error(f"Rename tag error: {e}", exc_info=True)
|
||||
return web.json_response({"success": False, "error": str(e)}, status=500)
|
||||
|
||||
async def delete_tag_endpoint(self, request):
|
||||
"""Delete a tag from all prompts."""
|
||||
try:
|
||||
from urllib.parse import unquote
|
||||
tag_name = unquote(request.match_info.get("tag_name", ""))
|
||||
if not tag_name:
|
||||
return web.json_response(
|
||||
{"success": False, "error": "Tag name required"}, status=400
|
||||
)
|
||||
|
||||
result = await self._run_in_executor(self.db.delete_tag_all_prompts, tag_name)
|
||||
resp = {
|
||||
"success": True,
|
||||
"tag_name": tag_name,
|
||||
"affected_count": result['affected_count']
|
||||
}
|
||||
if result.get('skipped_count', 0) > 0:
|
||||
resp['skipped_count'] = result['skipped_count']
|
||||
resp['warning'] = f"{result['skipped_count']} prompt(s) had corrupted tag data and were skipped"
|
||||
return web.json_response(resp)
|
||||
except Exception as e:
|
||||
self.logger.error(f"Delete tag error: {e}", exc_info=True)
|
||||
return web.json_response({"success": False, "error": str(e)}, status=500)
|
||||
|
||||
async def merge_tags_endpoint(self, request):
|
||||
"""Merge source tags into a target tag."""
|
||||
try:
|
||||
try:
|
||||
body = await request.json()
|
||||
except Exception:
|
||||
return web.json_response(
|
||||
{"success": False, "error": "Invalid JSON body"}, status=400
|
||||
)
|
||||
source_tags = body.get("source_tags", [])
|
||||
target_tag = body.get("target_tag", "").strip()
|
||||
|
||||
if not source_tags:
|
||||
return web.json_response(
|
||||
{"success": False, "error": "Source tags required"}, status=400
|
||||
)
|
||||
if not target_tag:
|
||||
return web.json_response(
|
||||
{"success": False, "error": "Target tag required"}, status=400
|
||||
)
|
||||
|
||||
result = await self._run_in_executor(self.db.merge_tags, source_tags, target_tag)
|
||||
resp = {
|
||||
"success": True,
|
||||
"target_tag": target_tag,
|
||||
"affected_count": result['affected_count'],
|
||||
"tags_merged": result['tags_merged']
|
||||
}
|
||||
if result.get('skipped_count', 0) > 0:
|
||||
resp['skipped_count'] = result['skipped_count']
|
||||
resp['warning'] = f"{result['skipped_count']} prompt(s) had corrupted tag data and were skipped"
|
||||
return web.json_response(resp)
|
||||
except Exception as e:
|
||||
self.logger.error(f"Merge tags error: {e}", exc_info=True)
|
||||
return web.json_response({"success": False, "error": str(e)}, status=500)
|
||||
|
||||
async def save_prompt(self, request):
|
||||
"""Save a new prompt with metadata and duplicate detection."""
|
||||
try:
|
||||
from utils.validators import (
|
||||
validate_prompt_text, validate_rating, validate_tags,
|
||||
validate_category, sanitize_input,
|
||||
)
|
||||
data = await request.json()
|
||||
|
||||
text = data.get("text", "").strip()
|
||||
if not text:
|
||||
return web.json_response(
|
||||
{"success": False, "error": "Text is required"}, status=400
|
||||
)
|
||||
|
||||
category = data.get("category", "").strip() or None
|
||||
tags = data.get("tags", [])
|
||||
rating = data.get("rating") or None
|
||||
notes = data.get("notes", "").strip() or None
|
||||
|
||||
try:
|
||||
validate_prompt_text(text)
|
||||
validate_category(category)
|
||||
validate_tags(tags)
|
||||
validate_rating(rating)
|
||||
except ValueError as ve:
|
||||
return web.json_response(
|
||||
{"success": False, "error": str(ve)}, status=400
|
||||
)
|
||||
|
||||
text = sanitize_input(text)
|
||||
|
||||
from utils.hashing import generate_prompt_hash
|
||||
prompt_hash = generate_prompt_hash(text)
|
||||
|
||||
existing = await self._run_in_executor(self.db.get_prompt_by_hash, prompt_hash)
|
||||
if existing:
|
||||
if any([category, tags, rating, notes]):
|
||||
await self._run_in_executor(
|
||||
self.db.update_prompt_metadata,
|
||||
prompt_id=existing['id'],
|
||||
category=category,
|
||||
tags=tags,
|
||||
rating=rating,
|
||||
notes=notes
|
||||
)
|
||||
return web.json_response({
|
||||
"success": True,
|
||||
"prompt_id": existing['id'],
|
||||
"message": "Prompt already exists, metadata updated",
|
||||
"is_duplicate": True
|
||||
})
|
||||
|
||||
prompt_id = await self._run_in_executor(
|
||||
self.db.save_prompt,
|
||||
text=text,
|
||||
category=category,
|
||||
tags=tags if tags else None,
|
||||
rating=rating,
|
||||
notes=notes,
|
||||
prompt_hash=prompt_hash,
|
||||
)
|
||||
|
||||
return web.json_response({
|
||||
"success": True,
|
||||
"prompt_id": prompt_id,
|
||||
"message": "Prompt saved successfully",
|
||||
})
|
||||
|
||||
except Exception as e:
|
||||
self.logger.error(f"Save error: {e}", exc_info=True)
|
||||
return web.json_response(
|
||||
{"success": False, "error": f"Failed to save prompt: {str(e)}"},
|
||||
status=500,
|
||||
)
|
||||
|
||||
async def delete_prompt(self, request):
|
||||
"""Delete a specific prompt by ID."""
|
||||
try:
|
||||
prompt_id = int(request.match_info["prompt_id"])
|
||||
success = await self._run_in_executor(self.db.delete_prompt, prompt_id)
|
||||
|
||||
if success:
|
||||
return web.json_response(
|
||||
{"success": True, "message": "Prompt deleted successfully"}
|
||||
)
|
||||
else:
|
||||
return web.json_response(
|
||||
{"success": False, "error": "Prompt not found or could not be deleted"},
|
||||
status=404,
|
||||
)
|
||||
|
||||
except ValueError:
|
||||
return web.json_response(
|
||||
{"success": False, "error": "Invalid prompt ID"}, status=400
|
||||
)
|
||||
except Exception as e:
|
||||
self.logger.error(f"Delete error: {e}")
|
||||
return web.json_response(
|
||||
{"success": False, "error": f"Failed to delete prompt: {str(e)}"},
|
||||
status=500,
|
||||
)
|
||||
|
||||
async def update_prompt(self, request):
|
||||
"""Update prompt text."""
|
||||
try:
|
||||
from utils.validators import validate_prompt_text, sanitize_input
|
||||
|
||||
prompt_id = int(request.match_info["prompt_id"])
|
||||
data = await request.json()
|
||||
new_text = data.get("text", "").strip()
|
||||
|
||||
if not new_text:
|
||||
return web.json_response(
|
||||
{"success": False, "error": "Text cannot be empty"}, status=400
|
||||
)
|
||||
|
||||
try:
|
||||
validate_prompt_text(new_text)
|
||||
except ValueError as ve:
|
||||
return web.json_response(
|
||||
{"success": False, "error": str(ve)}, status=400
|
||||
)
|
||||
|
||||
new_text = sanitize_input(new_text)
|
||||
|
||||
updated = await self._run_in_executor(self.db.update_prompt_text, prompt_id, new_text)
|
||||
if updated:
|
||||
return web.json_response(
|
||||
{"success": True, "message": "Prompt updated successfully"}
|
||||
)
|
||||
else:
|
||||
return web.json_response(
|
||||
{"success": False, "error": "Prompt not found"}, status=404
|
||||
)
|
||||
|
||||
except ValueError:
|
||||
return web.json_response(
|
||||
{"success": False, "error": "Invalid prompt ID"}, status=400
|
||||
)
|
||||
except Exception as e:
|
||||
self.logger.error(f"Update prompt error: {e}")
|
||||
return web.json_response(
|
||||
{"success": False, "error": f"Failed to update prompt: {str(e)}"},
|
||||
status=500,
|
||||
)
|
||||
|
||||
async def update_prompt_rating(self, request):
|
||||
"""Update prompt rating."""
|
||||
try:
|
||||
from utils.validators import validate_rating
|
||||
|
||||
prompt_id = int(request.match_info["prompt_id"])
|
||||
data = await request.json()
|
||||
rating = data.get("rating")
|
||||
|
||||
try:
|
||||
validate_rating(rating)
|
||||
except ValueError as ve:
|
||||
return web.json_response(
|
||||
{"success": False, "error": str(ve)}, status=400
|
||||
)
|
||||
|
||||
updated = await self._run_in_executor(self.db.update_prompt_rating, prompt_id, rating)
|
||||
if updated:
|
||||
return web.json_response(
|
||||
{"success": True, "message": "Rating updated successfully"}
|
||||
)
|
||||
else:
|
||||
return web.json_response(
|
||||
{"success": False, "error": "Prompt not found"}, status=404
|
||||
)
|
||||
|
||||
except ValueError:
|
||||
return web.json_response(
|
||||
{"success": False, "error": "Invalid prompt ID"}, status=400
|
||||
)
|
||||
except Exception as e:
|
||||
self.logger.error(f"Update rating error: {e}")
|
||||
return web.json_response(
|
||||
{"success": False, "error": f"Failed to update rating: {str(e)}"},
|
||||
status=500,
|
||||
)
|
||||
|
||||
async def add_prompt_tag(self, request):
|
||||
"""Add tag to prompt."""
|
||||
try:
|
||||
prompt_id = int(request.match_info["prompt_id"])
|
||||
data = await request.json()
|
||||
new_tag = data.get("tag", "").strip()
|
||||
|
||||
if not new_tag:
|
||||
return web.json_response(
|
||||
{"success": False, "error": "Tag cannot be empty"}, status=400
|
||||
)
|
||||
|
||||
prompt = await self._run_in_executor(self.db.get_prompt_by_id, prompt_id)
|
||||
if not prompt:
|
||||
return web.json_response(
|
||||
{"success": False, "error": "Prompt not found"}, status=404
|
||||
)
|
||||
|
||||
current_tags = prompt.get("tags", [])
|
||||
if not isinstance(current_tags, list):
|
||||
current_tags = []
|
||||
|
||||
if new_tag not in current_tags:
|
||||
current_tags.append(new_tag)
|
||||
await self._run_in_executor(self.db.set_prompt_tags, prompt_id, current_tags)
|
||||
|
||||
return web.json_response(
|
||||
{"success": True, "message": "Tag added successfully"}
|
||||
)
|
||||
|
||||
except ValueError:
|
||||
return web.json_response(
|
||||
{"success": False, "error": "Invalid prompt ID"}, status=400
|
||||
)
|
||||
except Exception as e:
|
||||
self.logger.error(f"Add tag error: {e}")
|
||||
return web.json_response(
|
||||
{"success": False, "error": f"Failed to add tag: {str(e)}"}, status=500
|
||||
)
|
||||
|
||||
async def add_tags_to_prompt(self, request):
|
||||
"""Add multiple tags to a single prompt."""
|
||||
try:
|
||||
data = await request.json()
|
||||
prompt_id = data.get("prompt_id")
|
||||
new_tags = data.get("tags", [])
|
||||
|
||||
if not prompt_id:
|
||||
return web.json_response(
|
||||
{"success": False, "error": "Prompt ID is required"}, status=400
|
||||
)
|
||||
|
||||
if not new_tags or not isinstance(new_tags, list):
|
||||
return web.json_response(
|
||||
{"success": False, "error": "Tags must be a non-empty list"}, status=400
|
||||
)
|
||||
|
||||
prompt = await self._run_in_executor(self.db.get_prompt_by_id, prompt_id)
|
||||
if not prompt:
|
||||
return web.json_response(
|
||||
{"success": False, "error": "Prompt not found"}, status=404
|
||||
)
|
||||
|
||||
current_tags = prompt.get("tags", [])
|
||||
if not isinstance(current_tags, list):
|
||||
current_tags = []
|
||||
|
||||
tags_added = 0
|
||||
for new_tag in new_tags:
|
||||
new_tag = new_tag.strip()
|
||||
if new_tag and new_tag not in current_tags:
|
||||
current_tags.append(new_tag)
|
||||
tags_added += 1
|
||||
|
||||
if tags_added > 0:
|
||||
await self._run_in_executor(self.db.set_prompt_tags, prompt_id, current_tags)
|
||||
|
||||
message = f"{tags_added} tag(s) added successfully"
|
||||
if tags_added == 0:
|
||||
message = "No new tags to add (all tags already exist)"
|
||||
|
||||
return web.json_response(
|
||||
{"success": True, "message": message, "tags_added": tags_added}
|
||||
)
|
||||
|
||||
except ValueError:
|
||||
return web.json_response(
|
||||
{"success": False, "error": "Invalid prompt ID"}, status=400
|
||||
)
|
||||
except Exception as e:
|
||||
self.logger.error(f"Add tags error: {e}")
|
||||
return web.json_response(
|
||||
{"success": False, "error": f"Failed to add tags: {str(e)}"}, status=500
|
||||
)
|
||||
|
||||
async def remove_prompt_tag(self, request):
|
||||
"""Remove tag from prompt."""
|
||||
try:
|
||||
prompt_id = int(request.match_info["prompt_id"])
|
||||
data = await request.json()
|
||||
tag_to_remove = data.get("tag", "").strip()
|
||||
|
||||
prompt = await self._run_in_executor(self.db.get_prompt_by_id, prompt_id)
|
||||
if not prompt:
|
||||
return web.json_response(
|
||||
{"success": False, "error": "Prompt not found"}, status=404
|
||||
)
|
||||
|
||||
current_tags = prompt.get("tags", [])
|
||||
if not isinstance(current_tags, list):
|
||||
current_tags = []
|
||||
|
||||
if tag_to_remove in current_tags:
|
||||
current_tags.remove(tag_to_remove)
|
||||
await self._run_in_executor(self.db.set_prompt_tags, prompt_id, current_tags)
|
||||
|
||||
return web.json_response(
|
||||
{"success": True, "message": "Tag removed successfully"}
|
||||
)
|
||||
|
||||
except ValueError:
|
||||
return web.json_response(
|
||||
{"success": False, "error": "Invalid prompt ID"}, status=400
|
||||
)
|
||||
except Exception as e:
|
||||
self.logger.error(f"Remove tag error: {e}")
|
||||
return web.json_response(
|
||||
{"success": False, "error": f"Failed to remove tag: {str(e)}"},
|
||||
status=500,
|
||||
)
|
||||
|
||||
async def bulk_delete_prompts(self, request):
|
||||
"""Bulk delete prompts."""
|
||||
try:
|
||||
data = await request.json()
|
||||
prompt_ids = data.get("prompt_ids", [])
|
||||
|
||||
if not prompt_ids:
|
||||
return web.json_response(
|
||||
{"success": False, "error": "No prompt IDs provided"}, status=400
|
||||
)
|
||||
|
||||
deleted_count = await self._run_in_executor(self.db.bulk_delete_prompts, prompt_ids)
|
||||
|
||||
return web.json_response({
|
||||
"success": True,
|
||||
"message": f"Deleted {deleted_count} prompts",
|
||||
"deleted_count": deleted_count,
|
||||
})
|
||||
|
||||
except Exception as e:
|
||||
self.logger.error(f"Bulk delete error: {e}")
|
||||
return web.json_response(
|
||||
{"success": False, "error": f"Failed to delete prompts: {str(e)}"},
|
||||
status=500,
|
||||
)
|
||||
|
||||
async def bulk_add_tags(self, request):
|
||||
"""Bulk add tags to prompts."""
|
||||
try:
|
||||
data = await request.json()
|
||||
prompt_ids = data.get("prompt_ids", [])
|
||||
new_tags = data.get("tags", [])
|
||||
|
||||
if not prompt_ids or not new_tags:
|
||||
return web.json_response(
|
||||
{"success": False, "error": "No prompt IDs or tags provided"},
|
||||
status=400,
|
||||
)
|
||||
|
||||
updated_count = await self._run_in_executor(self.db.bulk_add_tags, prompt_ids, new_tags)
|
||||
|
||||
return web.json_response({
|
||||
"success": True,
|
||||
"message": f"Added tags to {updated_count} prompts",
|
||||
"updated_count": updated_count,
|
||||
})
|
||||
|
||||
except Exception as e:
|
||||
self.logger.error(f"Bulk add tags error: {e}")
|
||||
return web.json_response(
|
||||
{"success": False, "error": f"Failed to add tags: {str(e)}"}, status=500
|
||||
)
|
||||
|
||||
async def bulk_set_category(self, request):
|
||||
"""Bulk set category for prompts."""
|
||||
try:
|
||||
data = await request.json()
|
||||
prompt_ids = data.get("prompt_ids", [])
|
||||
category = data.get("category", "").strip()
|
||||
|
||||
if not prompt_ids:
|
||||
return web.json_response(
|
||||
{"success": False, "error": "No prompt IDs provided"}, status=400
|
||||
)
|
||||
|
||||
updated_count = await self._run_in_executor(self.db.bulk_set_category, prompt_ids, category)
|
||||
|
||||
return web.json_response({
|
||||
"success": True,
|
||||
"message": f"Set category for {updated_count} prompts",
|
||||
"updated_count": updated_count,
|
||||
})
|
||||
|
||||
except Exception as e:
|
||||
self.logger.error(f"Bulk set category error: {e}")
|
||||
return web.json_response(
|
||||
{"success": False, "error": f"Failed to set category: {str(e)}"},
|
||||
status=500,
|
||||
)
|
||||
|
||||
async def export_prompts(self, request):
|
||||
"""Export all prompts to JSON."""
|
||||
try:
|
||||
prompts = await self._run_in_executor(self.db.search_prompts, limit=10000)
|
||||
|
||||
export_data = {
|
||||
"export_date": datetime.datetime.now(datetime.timezone.utc).isoformat(),
|
||||
"total_prompts": len(prompts),
|
||||
"prompts": prompts,
|
||||
}
|
||||
|
||||
json_data = json.dumps(export_data, indent=2, ensure_ascii=False)
|
||||
|
||||
return web.Response(
|
||||
text=json_data,
|
||||
content_type="application/json",
|
||||
headers={
|
||||
"Content-Disposition": f'attachment; filename="prompt_manager_{datetime.datetime.now().strftime("%Y%m%d_%H%M%S")}.json"'
|
||||
},
|
||||
)
|
||||
|
||||
except Exception as e:
|
||||
self.logger.error(f"Export error: {e}")
|
||||
return web.json_response(
|
||||
{"success": False, "error": f"Failed to export prompts: {str(e)}"},
|
||||
status=500,
|
||||
)
|
||||
Reference in New Issue
Block a user