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

Automated formatting pass across 30 files to establish consistent code
style enforced by CI. No logic changes.
This commit is contained in:
Vito Sansevero
2026-02-07 08:30:09 -08:00
parent 224fcc7803
commit 2737a4b1ef
30 changed files with 3278 additions and 2531 deletions
+1 -1
View File
@@ -3,4 +3,4 @@ ComfyUI_PromptManager Python API modules.
This package contains the core API components for the web interface and configuration
management, including REST endpoints, configuration handling, and server integration.
"""
"""
+122 -76
View File
@@ -40,20 +40,20 @@ except ImportError:
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__)))
)
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_TYPES = frozenset(
(
"application/json",
"text/html",
"text/css",
"application/javascript",
"text/plain",
)
)
_gzip_registered = False
@@ -62,24 +62,24 @@ 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/'):
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', ''):
if "gzip" not in request.headers.get("Accept-Encoding", ""):
return response
if 'Content-Encoding' in response.headers:
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 ''
content_type = response.content_type or ""
if not any(ct in content_type for ct in _GZIP_TYPES):
return response
@@ -88,8 +88,8 @@ async def _gzip_middleware(request, handler):
return response
response.body = compressed
response.headers['Content-Encoding'] = 'gzip'
response.headers['Vary'] = 'Accept-Encoding'
response.headers["Content-Encoding"] = "gzip"
response.headers["Vary"] = "Accept-Encoding"
return response
@@ -117,7 +117,7 @@ class PromptManagerAPI(
def __init__(self):
"""Initialize the PromptManager API with database connection and cleanup."""
self.logger = get_logger('prompt_manager.api')
self.logger = get_logger("prompt_manager.api")
self.logger.info("Initializing PromptManager API")
self.db = PromptDatabase()
@@ -131,7 +131,9 @@ class PromptManagerAPI(
try:
removed = self.db.cleanup_duplicates()
if removed > 0:
self.logger.info(f"Startup cleanup: removed {removed} duplicate prompts")
self.logger.info(
f"Startup cleanup: removed {removed} duplicate prompts"
)
except Exception as e:
self.logger.error(f"Startup cleanup failed: {e}")
@@ -207,7 +209,9 @@ class PromptManagerAPI(
async def serve_gallery_ui(request):
try:
html_path = os.path.join(
_get_project_root(), "web", "metadata.html",
_get_project_root(),
"web",
"metadata.html",
)
if html_path not in self._html_cache:
@@ -222,7 +226,9 @@ class PromptManagerAPI(
)
return web.Response(
text=self._html_cache[html_path], content_type="text/html", charset="utf-8"
text=self._html_cache[html_path],
content_type="text/html",
charset="utf-8",
)
except Exception as e:
@@ -236,7 +242,9 @@ class PromptManagerAPI(
async def serve_admin_ui(request):
try:
html_path = os.path.join(
_get_project_root(), "web", "admin.html",
_get_project_root(),
"web",
"admin.html",
)
if html_path not in self._html_cache:
@@ -251,7 +259,9 @@ class PromptManagerAPI(
)
return web.Response(
text=self._html_cache[html_path], content_type="text/html", charset="utf-8"
text=self._html_cache[html_path],
content_type="text/html",
charset="utf-8",
)
except Exception as e:
@@ -265,7 +275,9 @@ class PromptManagerAPI(
async def serve_gallery_admin_ui(request):
try:
html_path = os.path.join(
_get_project_root(), "web", "gallery.html",
_get_project_root(),
"web",
"gallery.html",
)
if html_path not in self._html_cache:
@@ -280,7 +292,9 @@ class PromptManagerAPI(
)
return web.Response(
text=self._html_cache[html_path], content_type="text/html", charset="utf-8"
text=self._html_cache[html_path],
content_type="text/html",
charset="utf-8",
)
except Exception as e:
@@ -362,6 +376,7 @@ class PromptManagerAPI(
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")
@@ -380,30 +395,36 @@ class PromptManagerAPI(
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', '')
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"
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='/')}"
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}"
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='/')}"
image["thumbnail_url"] = (
f"/prompt_manager/images/serve/{url_quote(thumb_rel, safe='/')}"
)
except (ValueError, RuntimeError):
pass
@@ -415,7 +436,7 @@ class PromptManagerAPI(
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':
elif isinstance(obj, float) and str(obj) == "nan":
return None
else:
return obj
@@ -442,11 +463,13 @@ class PromptManagerAPI(
for i in range(max_depth):
# Check if current directory contains ComfyUI markers
comfyui_markers = ['main.py', 'nodes.py', 'server.py']
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.logger.debug(
f"Found ComfyUI output directory via upward search: {output_dir}"
)
self._cached_output_dir = str(output_dir)
return self._cached_output_dir
@@ -500,7 +523,7 @@ class PromptManagerAPI(
try:
with Image.open(image_path) as img:
metadata = {}
if hasattr(img, 'text'):
if hasattr(img, "text"):
for key, value in img.text.items():
metadata[key] = value
return metadata
@@ -511,11 +534,11 @@ class PromptManagerAPI(
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
"prompt": None,
"workflow": None,
"parameters": {},
"positive_prompt": None,
"negative_prompt": None,
}
# Check for A1111 style parameters first (like parse-metadata.py)
@@ -523,40 +546,48 @@ class PromptManagerAPI(
params = metadata["parameters"]
lines = params.splitlines()
if lines:
result['positive_prompt'] = lines[0].strip()
result["positive_prompt"] = lines[0].strip()
for line in lines:
if line.lower().startswith("negative prompt:"):
result['negative_prompt'] = line.split(":", 1)[1].strip()
result["negative_prompt"] = line.split(":", 1)[1].strip()
break
# Store raw parameters too
result['parameters']['parameters'] = params
result["parameters"]["parameters"] = params
# If no A1111 format found, proceed with ComfyUI parsing
if result['positive_prompt'] is None:
if result["positive_prompt"] is None:
# Check for direct prompt field
if 'prompt' in metadata:
if "prompt" in metadata:
try:
prompt_data = json.loads(metadata['prompt'])
result['prompt'] = prompt_data
prompt_data = json.loads(metadata["prompt"])
result["prompt"] = prompt_data
except json.JSONDecodeError:
result['prompt'] = metadata['prompt']
result["prompt"] = metadata["prompt"]
# Check for workflow
if 'workflow' in metadata:
if "workflow" in metadata:
try:
workflow_data = json.loads(metadata['workflow'])
result['workflow'] = workflow_data
workflow_data = json.loads(metadata["workflow"])
result["workflow"] = workflow_data
except json.JSONDecodeError:
result['workflow'] = metadata['workflow']
result["workflow"] = metadata["workflow"]
# Check for other common ComfyUI fields
common_fields = ['positive', 'negative', 'steps', 'cfg', 'sampler', 'scheduler', 'seed']
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])
result["parameters"][field] = json.loads(metadata[field])
except json.JSONDecodeError:
result['parameters'][field] = metadata[field]
result["parameters"][field] = metadata[field]
return result
@@ -567,40 +598,46 @@ class PromptManagerAPI(
if isinstance(value, str):
return value
elif isinstance(value, list):
return ' '.join(str(item) for item in value if item)
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']
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']
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'])
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')
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)
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')
workflow_data = parsed_data.get("workflow")
if isinstance(workflow_data, dict):
positive_prompt = self._extract_positive_prompt_from_comfyui_data(workflow_data)
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'])
if parsed_data.get("parameters", {}).get("positive"):
return safe_to_string(parsed_data["parameters"]["positive"])
return None
@@ -652,12 +689,21 @@ class PromptManagerAPI(
# 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'
"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):
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():
@@ -721,8 +767,8 @@ class PromptManagerAPI(
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:
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:
+454 -278
View File
File diff suppressed because it is too large Load Diff
+145 -138
View File
@@ -47,24 +47,23 @@ class AutotagRoutesMixin:
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()
})
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)
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')
model_type = request.match_info.get("model_type")
async def stream_response():
try:
@@ -78,15 +77,14 @@ class AutotagRoutesMixin:
yield f"data: {json.dumps({'type': 'progress', 'progress': 0, 'status': 'Starting download...'})}\n\n"
progress_data = {'last_progress': 0}
progress_data = {"last_progress": 0}
def progress_callback(status: str, progress: float):
progress_data['last_progress'] = progress
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)
None, lambda: service.download_model(model_type, progress_callback)
)
if success:
@@ -100,28 +98,28 @@ class AutotagRoutesMixin:
response = web.StreamResponse(
status=200,
reason='OK',
reason="OK",
headers={
'Content-Type': 'text/event-stream',
'Cache-Control': 'no-cache',
'Connection': 'keep-alive'
}
"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(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'
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():
@@ -132,7 +130,7 @@ class AutotagRoutesMixin:
service = get_autotag_service()
status = service.get_models_status()
if not status.get(model_type, {}).get('downloaded'):
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
@@ -141,8 +139,7 @@ class AutotagRoutesMixin:
loop = asyncio.get_event_loop()
try:
await loop.run_in_executor(
None,
lambda: service.load_model(model_type, use_gpu)
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"
@@ -172,8 +169,8 @@ class AutotagRoutesMixin:
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')
image_path = image_data.get("image_path")
prompt_id = image_data.get("prompt_id")
if not image_path or not prompt_id:
skipped += 1
@@ -200,10 +197,12 @@ class AutotagRoutesMixin:
continue
if skip_tagged:
prompt_tags = image_data.get('prompt_tags', [])
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']
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
@@ -217,18 +216,23 @@ class AutotagRoutesMixin:
try:
tags = await loop.run_in_executor(
None,
lambda p=str(image_path): service.generate_tags(p)
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)
existing_prompt = await self._run_in_executor(
self.db.get_prompt_by_id, prompt_id
)
if existing_prompt:
existing_tags = existing_prompt.get('tags', [])
existing_tags = existing_prompt.get("tags", [])
if isinstance(existing_tags, str):
existing_tags = [t.strip() for t in existing_tags.split(',') if t.strip()]
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:
@@ -236,7 +240,7 @@ class AutotagRoutesMixin:
await self._run_in_executor(
self.db.update_prompt_metadata,
prompt_id,
tags=all_tags
tags=all_tags,
)
tagged += 1
else:
@@ -260,32 +264,33 @@ class AutotagRoutesMixin:
if not keep_in_memory:
service.unload_model()
model_status = 'Model unloaded'
model_status = "Model unloaded"
else:
model_status = 'Model kept in memory'
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',
reason="OK",
headers={
'Content-Type': 'text/event-stream',
'Cache-Control': 'no-cache',
'Connection': 'keep-alive'
}
"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(chunk.encode("utf-8"))
await response.write_eof()
return response
@@ -294,26 +299,27 @@ class AutotagRoutesMixin:
"""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)
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)
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:
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)
None, lambda: service.load_model(model_type, use_gpu)
)
if custom_prompt:
@@ -321,81 +327,74 @@ class AutotagRoutesMixin:
loop = asyncio.get_event_loop()
tags = await loop.run_in_executor(
None,
lambda: service.generate_tags(image_path)
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)
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
})
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)
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', [])
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)
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'
})
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)
return web.json_response(
{"success": False, "error": f"Prompt {prompt_id} not found"},
status=404,
)
existing_tags = prompt.get('tags', [])
existing_tags = prompt.get("tags", [])
if isinstance(existing_tags, str):
existing_tags = [t.strip() for t in existing_tags.split(',') if t.strip()]
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
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)
})
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)
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."""
@@ -405,39 +404,45 @@ class AutotagRoutesMixin:
service = get_autotag_service()
if not service.is_model_loaded():
return web.json_response({
'success': True,
'message': 'No model was 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
})
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)
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)
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']
image_extensions = [
".png",
".jpg",
".jpeg",
".gif",
".webp",
".bmp",
".tiff",
]
images = []
seen_paths = set()
@@ -445,7 +450,7 @@ class AutotagRoutesMixin:
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:
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)
@@ -455,35 +460,37 @@ class AutotagRoutesMixin:
thumbnail_url = None
thumbnails_dir = output_path / "thumbnails"
if thumbnails_dir.exists():
rel_path_no_ext = rel_path.with_suffix('')
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}"
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'])
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)
})
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)
return web.json_response({"success": False, "error": str(e)}, status=500)
+515 -383
View File
File diff suppressed because it is too large Load Diff
+69 -86
View File
@@ -43,7 +43,10 @@ class LoggingRoutesMixin:
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__))))
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()
@@ -53,8 +56,8 @@ class LoggingRoutesMixin:
try:
logger_manager = self._get_logger_manager()
limit = int(request.query.get('limit', 100))
level = request.query.get('level', None)
limit = int(request.query.get("limit", 100))
level = request.query.get("level", None)
if limit > 1000:
limit = 1000
@@ -63,21 +66,21 @@ class LoggingRoutesMixin:
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
})
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)
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."""
@@ -85,58 +88,49 @@ class LoggingRoutesMixin:
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)
})
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)
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']
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)
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)
return web.json_response(
{"success": False, "error": "Log file not found"}, status=404
)
with open(log_file_path, 'rb') as f:
with open(log_file_path, "rb") as f:
file_content = f.read()
return web.Response(
body=file_content,
content_type='text/plain',
content_type="text/plain",
headers={
'Content-Disposition': f'attachment; filename="{filename}"',
'Content-Length': str(len(file_content))
}
"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)
return web.json_response({"success": False, "error": str(e)}, status=500)
async def truncate_logs(self, request):
"""Truncate all log files."""
@@ -144,18 +138,17 @@ class LoggingRoutesMixin:
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
})
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)
return web.json_response({"success": False, "error": str(e)}, status=500)
async def get_log_config(self, request):
"""Get current logging configuration."""
@@ -163,17 +156,11 @@ class LoggingRoutesMixin:
logger_manager = self._get_logger_manager()
config = logger_manager.get_config()
return web.json_response({
'success': True,
'config': 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)
return web.json_response({"success": False, "error": str(e)}, status=500)
async def update_log_config(self, request):
"""Update logging configuration."""
@@ -181,29 +168,31 @@ class LoggingRoutesMixin:
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()
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()
})
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)
return web.json_response({"success": False, "error": str(e)}, status=500)
async def get_log_stats(self, request):
"""Get logging statistics."""
@@ -211,14 +200,8 @@ class LoggingRoutesMixin:
logger_manager = self._get_logger_manager()
stats = logger_manager.get_log_stats()
return web.json_response({
'success': True,
'stats': 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)
return web.json_response({"success": False, "error": str(e)}, status=500)
+217 -128
View File
@@ -153,22 +153,26 @@ class PromptRoutesMixin:
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'])
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'])
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)
@@ -177,7 +181,7 @@ class PromptRoutesMixin:
"success": False,
"error": f"Failed to get recent prompts: {str(e)}",
"results": [],
"pagination": {"total": 0, "page": 1, "total_pages": 0}
"pagination": {"total": 0, "page": 1, "total_pages": 0},
},
status=500,
)
@@ -190,7 +194,11 @@ class PromptRoutesMixin:
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": []},
{
"success": False,
"error": f"Failed to get categories: {str(e)}",
"categories": [],
},
status=500,
)
@@ -202,7 +210,11 @@ class PromptRoutesMixin:
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": []},
{
"success": False,
"error": f"Failed to get tags: {str(e)}",
"tags": [],
},
status=500,
)
@@ -214,25 +226,32 @@ class PromptRoutesMixin:
offset = int(request.query.get("offset", 0))
except (ValueError, TypeError):
return web.json_response(
{"success": False, "error": "Invalid limit or offset parameter"}, status=400
{"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)
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']
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)
@@ -241,6 +260,7 @@ class PromptRoutesMixin:
"""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(
@@ -252,23 +272,28 @@ class PromptRoutesMixin:
offset = int(request.query.get("offset", 0))
except (ValueError, TypeError):
return web.json_response(
{"success": False, "error": "Invalid limit or offset parameter"}, status=400
{"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'])
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']
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)
@@ -284,22 +309,30 @@ class PromptRoutesMixin:
offset = int(request.query.get("offset", 0))
except (ValueError, TypeError):
return web.json_response(
{"success": False, "error": "Invalid limit or offset parameter"}, status=400
{
"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']
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:
@@ -317,24 +350,29 @@ class PromptRoutesMixin:
offset = int(request.query.get("offset", 0))
except (ValueError, TypeError):
return web.json_response(
{"success": False, "error": "Invalid limit or offset parameter"}, status=400
{"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'])
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']
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)
@@ -343,6 +381,7 @@ class PromptRoutesMixin:
"""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(
@@ -361,16 +400,20 @@ class PromptRoutesMixin:
{"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)
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']
"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"
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)
@@ -380,21 +423,26 @@ class PromptRoutesMixin:
"""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)
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']
"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"
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)
@@ -421,16 +469,20 @@ class PromptRoutesMixin:
{"success": False, "error": "Target tag required"}, status=400
)
result = await self._run_in_executor(self.db.merge_tags, source_tags, target_tag)
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']
"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"
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)
@@ -440,9 +492,13 @@ class PromptRoutesMixin:
"""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,
validate_prompt_text,
validate_rating,
validate_tags,
validate_category,
sanitize_input,
)
data = await request.json()
text = data.get("text", "").strip()
@@ -469,25 +525,30 @@ class PromptRoutesMixin:
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)
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'],
prompt_id=existing["id"],
category=category,
tags=tags,
rating=rating,
notes=notes
notes=notes,
)
return web.json_response({
"success": True,
"prompt_id": existing['id'],
"message": "Prompt already exists, metadata updated",
"is_duplicate": True
})
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,
@@ -499,11 +560,13 @@ class PromptRoutesMixin:
prompt_hash=prompt_hash,
)
return web.json_response({
"success": True,
"prompt_id": prompt_id,
"message": "Prompt saved successfully",
})
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)
@@ -524,7 +587,10 @@ class PromptRoutesMixin:
)
else:
return web.json_response(
{"success": False, "error": "Prompt not found or could not be deleted"},
{
"success": False,
"error": "Prompt not found or could not be deleted",
},
status=404,
)
@@ -562,7 +628,9 @@ class PromptRoutesMixin:
new_text = sanitize_input(new_text)
updated = await self._run_in_executor(self.db.update_prompt_text, prompt_id, 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"}
@@ -599,7 +667,9 @@ class PromptRoutesMixin:
{"success": False, "error": str(ve)}, status=400
)
updated = await self._run_in_executor(self.db.update_prompt_rating, prompt_id, rating)
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"}
@@ -644,7 +714,9 @@ class PromptRoutesMixin:
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)
await self._run_in_executor(
self.db.set_prompt_tags, prompt_id, current_tags
)
return web.json_response(
{"success": True, "message": "Tag added successfully"}
@@ -674,7 +746,8 @@ class PromptRoutesMixin:
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
{"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)
@@ -695,7 +768,9 @@ class PromptRoutesMixin:
tags_added += 1
if tags_added > 0:
await self._run_in_executor(self.db.set_prompt_tags, prompt_id, current_tags)
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:
@@ -734,7 +809,9 @@ class PromptRoutesMixin:
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)
await self._run_in_executor(
self.db.set_prompt_tags, prompt_id, current_tags
)
return web.json_response(
{"success": True, "message": "Tag removed successfully"}
@@ -762,13 +839,17 @@ class PromptRoutesMixin:
{"success": False, "error": "No prompt IDs provided"}, status=400
)
deleted_count = await self._run_in_executor(self.db.bulk_delete_prompts, prompt_ids)
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,
})
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}")
@@ -790,13 +871,17 @@ class PromptRoutesMixin:
status=400,
)
updated_count = await self._run_in_executor(self.db.bulk_add_tags, prompt_ids, new_tags)
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,
})
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}")
@@ -816,13 +901,17 @@ class PromptRoutesMixin:
{"success": False, "error": "No prompt IDs provided"}, status=400
)
updated_count = await self._run_in_executor(self.db.bulk_set_category, prompt_ids, category)
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,
})
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}")
+106 -97
View File
@@ -16,13 +16,16 @@ try:
from ..utils.logging_config import get_logger
except ImportError:
import logging
def get_logger(name: str) -> logging.Logger:
logger = logging.getLogger(name)
if not logger.handlers:
handler = logging.StreamHandler()
handler.setFormatter(logging.Formatter(
'%(asctime)s - %(name)s - %(levelname)s - %(message)s'
))
handler.setFormatter(
logging.Formatter(
"%(asctime)s - %(name)s - %(levelname)s - %(message)s"
)
)
logger.addHandler(handler)
logger.setLevel(logging.DEBUG)
return logger
@@ -46,7 +49,7 @@ MODELS = {
"size": "~16GB",
"repo": "fancyfeast/llama-joycaption-beta-one-hf-llava",
"subdir": "llama-joycaption-beta-one-hf-llava",
}
},
}
# Default prompts
@@ -84,7 +87,7 @@ class AutoTagService:
models_dir: Directory for storing models. If None, uses ComfyUI's
folder_paths.models_dir / "LLM" path.
"""
self.logger = get_logger('autotag.service')
self.logger = get_logger("autotag.service")
# Determine models directory
if models_dir:
@@ -93,11 +96,14 @@ class AutoTagService:
# Use ComfyUI's folder_paths system
try:
import folder_paths
self.models_dir = Path(folder_paths.models_dir) / "LLM"
except ImportError:
# Fallback for standalone usage (not running in ComfyUI)
self.logger.warning("folder_paths not available, using fallback path")
self.models_dir = Path(__file__).parent.parent.parent.parent / "models" / "LLM"
self.models_dir = (
Path(__file__).parent.parent.parent.parent / "models" / "LLM"
)
self.logger.info(f"AutoTag service initialized. Models dir: {self.models_dir}")
@@ -148,29 +154,27 @@ class AutoTagService:
for model_type, config in MODELS.items():
model_status = {
'name': config['name'],
'description': config['description'],
'size': config['size'],
'downloaded': False,
'model_path': None
"name": config["name"],
"description": config["description"],
"size": config["size"],
"downloaded": False,
"model_path": None,
}
if model_type == 'gguf':
if model_type == "gguf":
model_exists, mmproj_exists = self._check_gguf_models()
model_status['model_exists'] = model_exists
model_status['mmproj_exists'] = mmproj_exists
model_status['downloaded'] = model_exists and mmproj_exists
if model_status['downloaded']:
model_status['model_path'] = str(
self.models_dir / config['subdir'] / config['filename']
model_status["model_exists"] = model_exists
model_status["mmproj_exists"] = mmproj_exists
model_status["downloaded"] = model_exists and mmproj_exists
if model_status["downloaded"]:
model_status["model_path"] = str(
self.models_dir / config["subdir"] / config["filename"]
)
else: # hf
model_status['downloaded'] = self._check_hf_model()
if model_status['downloaded']:
model_status["downloaded"] = self._check_hf_model()
if model_status["downloaded"]:
# Get the actual path (local or cache)
model_status['model_path'] = str(
self._get_hf_model_path()
)
model_status["model_path"] = str(self._get_hf_model_path())
status[model_type] = model_status
@@ -182,10 +186,10 @@ class AutoTagService:
Returns:
Tuple of (model_exists, mmproj_exists)
"""
config = MODELS['gguf']
gguf_dir = self.models_dir / config['subdir']
model_path = gguf_dir / config['filename']
mmproj_path = gguf_dir / config['mmproj_filename']
config = MODELS["gguf"]
gguf_dir = self.models_dir / config["subdir"]
model_path = gguf_dir / config["filename"]
mmproj_path = gguf_dir / config["mmproj_filename"]
return model_path.exists(), mmproj_path.exists()
def _check_hf_model(self) -> bool:
@@ -196,17 +200,17 @@ class AutoTagService:
Returns:
True if model directory contains config.json
"""
config = MODELS['hf']
config = MODELS["hf"]
# Check local directory first
model_dir = self.models_dir / config['subdir']
model_dir = self.models_dir / config["subdir"]
self.logger.debug(f"Checking local HF model path: {model_dir}")
if (model_dir / "config.json").exists():
self.logger.debug("Found HF model in local directory")
return True
# Check HuggingFace cache as fallback
hf_cache_path = self._get_hf_cache_path(config['repo'])
hf_cache_path = self._get_hf_cache_path(config["repo"])
if hf_cache_path:
self.logger.debug(f"Found HF model in cache: {hf_cache_path}")
return True
@@ -251,15 +255,15 @@ class AutoTagService:
Returns:
Path to the model directory, or None if not found
"""
config = MODELS['hf']
config = MODELS["hf"]
# Check local directory first
model_dir = self.models_dir / config['subdir']
model_dir = self.models_dir / config["subdir"]
if (model_dir / "config.json").exists():
return model_dir
# Check HuggingFace cache as fallback
cache_path = self._get_hf_cache_path(config['repo'])
cache_path = self._get_hf_cache_path(config["repo"])
if cache_path:
return cache_path
@@ -268,7 +272,7 @@ class AutoTagService:
def download_model(
self,
model_type: str,
progress_callback: Optional[Callable[[str, float], None]] = None
progress_callback: Optional[Callable[[str, float], None]] = None,
) -> bool:
"""Download a model with optional progress updates.
@@ -283,7 +287,9 @@ class AutoTagService:
ValueError: If model_type is not valid
"""
if model_type not in MODELS:
raise ValueError(f"Invalid model type: {model_type}. Must be 'gguf' or 'hf'")
raise ValueError(
f"Invalid model type: {model_type}. Must be 'gguf' or 'hf'"
)
try:
from huggingface_hub import hf_hub_download, snapshot_download
@@ -294,7 +300,7 @@ class AutoTagService:
return False
try:
if model_type == 'gguf':
if model_type == "gguf":
return self._download_gguf_models(progress_callback)
else:
return self._download_hf_model(progress_callback)
@@ -305,18 +311,17 @@ class AutoTagService:
return False
def _download_gguf_models(
self,
progress_callback: Optional[Callable[[str, float], None]] = None
self, progress_callback: Optional[Callable[[str, float], None]] = None
) -> bool:
"""Download GGUF model and mmproj files."""
from huggingface_hub import hf_hub_download
config = MODELS['gguf']
gguf_dir = self.models_dir / config['subdir']
config = MODELS["gguf"]
gguf_dir = self.models_dir / config["subdir"]
gguf_dir.mkdir(parents=True, exist_ok=True)
model_path = gguf_dir / config['filename']
mmproj_path = gguf_dir / config['mmproj_filename']
model_path = gguf_dir / config["filename"]
mmproj_path = gguf_dir / config["mmproj_filename"]
# Download main model
if not model_path.exists():
@@ -325,10 +330,10 @@ class AutoTagService:
self.logger.info(f"Downloading GGUF model: {config['filename']}")
hf_hub_download(
repo_id=config['repo'],
filename=config['filename'],
repo_id=config["repo"],
filename=config["filename"],
local_dir=str(gguf_dir),
local_dir_use_symlinks=False
local_dir_use_symlinks=False,
)
self.logger.info("GGUF model downloaded")
@@ -342,10 +347,10 @@ class AutoTagService:
self.logger.info(f"Downloading mmproj: {config['mmproj_filename']}")
hf_hub_download(
repo_id=config['mmproj_repo'],
filename=config['mmproj_filename'],
repo_id=config["mmproj_repo"],
filename=config["mmproj_filename"],
local_dir=str(gguf_dir),
local_dir_use_symlinks=False
local_dir_use_symlinks=False,
)
self.logger.info("mmproj downloaded")
@@ -355,14 +360,13 @@ class AutoTagService:
return True
def _download_hf_model(
self,
progress_callback: Optional[Callable[[str, float], None]] = None
self, progress_callback: Optional[Callable[[str, float], None]] = None
) -> bool:
"""Download HuggingFace model."""
from huggingface_hub import snapshot_download
config = MODELS['hf']
model_dir = self.models_dir / config['subdir']
config = MODELS["hf"]
model_dir = self.models_dir / config["subdir"]
if not self._check_hf_model():
if progress_callback:
@@ -370,9 +374,9 @@ class AutoTagService:
self.logger.info(f"Downloading HF model: {config['repo']}")
snapshot_download(
repo_id=config['repo'],
repo_id=config["repo"],
local_dir=str(model_dir),
local_dir_use_symlinks=False
local_dir_use_symlinks=False,
)
self.logger.info("HF model downloaded")
@@ -403,11 +407,11 @@ class AutoTagService:
self.unload_model()
status = self.get_models_status()
if not status[model_type]['downloaded']:
if not status[model_type]["downloaded"]:
raise RuntimeError(f"Model {model_type} not downloaded")
try:
if model_type == 'gguf':
if model_type == "gguf":
self._tagger = self._load_gguf_tagger(use_gpu)
else:
self._tagger = self._load_hf_tagger()
@@ -427,10 +431,10 @@ class AutoTagService:
from llama_cpp import Llama
from llama_cpp.llama_chat_format import Llava15ChatHandler
config = MODELS['gguf']
gguf_dir = self.models_dir / config['subdir']
model_path = gguf_dir / config['filename']
mmproj_path = gguf_dir / config['mmproj_filename']
config = MODELS["gguf"]
gguf_dir = self.models_dir / config["subdir"]
model_path = gguf_dir / config["filename"]
mmproj_path = gguf_dir / config["mmproj_filename"]
self.logger.info("Loading GGUF model...")
n_gpu_layers = -1 if use_gpu else 0
@@ -447,17 +451,23 @@ class AutoTagService:
)
self.logger.info("GGUF model loaded")
return ('gguf', tagger)
return ("gguf", tagger)
def _load_hf_tagger(self, quantization: str = "8bit"):
"""Load HuggingFace-based tagger."""
import torch
from transformers import AutoProcessor, LlavaForConditionalGeneration, BitsAndBytesConfig
from transformers import (
AutoProcessor,
LlavaForConditionalGeneration,
BitsAndBytesConfig,
)
# Get the actual model path (local or cache)
model_path = self._get_hf_model_path()
if model_path is None:
raise RuntimeError("HuggingFace model not found in local directory or cache")
raise RuntimeError(
"HuggingFace model not found in local directory or cache"
)
self.logger.info(f"Loading HuggingFace model from {model_path}...")
device = "cuda" if torch.cuda.is_available() else "cpu"
@@ -476,20 +486,18 @@ class AutoTagService:
str(model_path),
torch_dtype=torch.float16,
quantization_config=qnt_config,
**model_kwargs
**model_kwargs,
)
else:
model = LlavaForConditionalGeneration.from_pretrained(
str(model_path),
torch_dtype=torch.bfloat16,
**model_kwargs
str(model_path), torch_dtype=torch.bfloat16, **model_kwargs
)
model.eval()
self.logger.info("HuggingFace model loaded")
# Track the compute dtype for pixel_values conversion
compute_dtype = torch.float16 if quantization == "8bit" else torch.bfloat16
return ('hf', (model, processor, device, compute_dtype))
return ("hf", (model, processor, device, compute_dtype))
def unload_model(self):
"""Unload the current model and free memory."""
@@ -503,6 +511,7 @@ class AutoTagService:
# Try to clear CUDA cache if available
try:
import torch
if torch.cuda.is_available():
torch.cuda.empty_cache()
except ImportError:
@@ -518,11 +527,7 @@ class AutoTagService:
"""Get the type of currently loaded model."""
return self._current_model_type
def generate_tags(
self,
image_path: str,
prompt: Optional[str] = None
) -> List[str]:
def generate_tags(self, image_path: str, prompt: Optional[str] = None) -> List[str]:
"""Generate tags for an image.
Args:
@@ -546,17 +551,19 @@ class AutoTagService:
# Load image
image = Image.open(image_path)
if image.mode != 'RGB':
image = image.convert('RGB')
if image.mode != "RGB":
image = image.convert("RGB")
# Generate based on model type
model_type, tagger_obj = self._tagger
if model_type == 'gguf':
if model_type == "gguf":
raw_tags = self._generate_gguf(tagger_obj, image, use_prompt)
else:
model, processor, device, compute_dtype = tagger_obj
raw_tags = self._generate_hf(model, processor, device, compute_dtype, image, use_prompt)
raw_tags = self._generate_hf(
model, processor, device, compute_dtype, image, use_prompt
)
# Parse tags from response
tags = self._parse_tags(raw_tags)
@@ -572,9 +579,9 @@ class AutoTagService:
# Encode to base64
buffer = io.BytesIO()
image.save(buffer, format='PNG')
image.save(buffer, format="PNG")
buffer.seek(0)
img_base64 = base64.b64encode(buffer.read()).decode('utf-8')
img_base64 = base64.b64encode(buffer.read()).decode("utf-8")
data_uri = f"data:image/png;base64,{img_base64}"
# Create message
@@ -584,9 +591,9 @@ class AutoTagService:
"role": "user",
"content": [
{"type": "text", "text": prompt},
{"type": "image_url", "image_url": {"url": data_uri}}
]
}
{"type": "image_url", "image_url": {"url": data_uri}},
],
},
]
# Generate
@@ -608,7 +615,7 @@ class AutoTagService:
device: str,
compute_dtype,
image: Image.Image,
prompt: str
prompt: str,
) -> str:
"""Generate tags using HuggingFace model."""
import torch
@@ -624,14 +631,14 @@ class AutoTagService:
convo, tokenize=False, add_generation_prompt=True
)
inputs = processor(
text=[convo_string], images=[image], return_tensors="pt"
).to(device)
inputs = processor(text=[convo_string], images=[image], return_tensors="pt").to(
device
)
# Convert pixel_values to the model's compute dtype
# (float16 for 8-bit quantized, bfloat16 for non-quantized)
if 'pixel_values' in inputs and inputs['pixel_values'] is not None:
inputs['pixel_values'] = inputs['pixel_values'].to(compute_dtype)
if "pixel_values" in inputs and inputs["pixel_values"] is not None:
inputs["pixel_values"] = inputs["pixel_values"].to(compute_dtype)
with torch.inference_mode(), torch.cuda.amp.autocast(enabled=True):
generate_ids = model.generate(
@@ -643,8 +650,10 @@ class AutoTagService:
use_cache=True,
)[0]
generate_ids = generate_ids[inputs['input_ids'].shape[1]:]
return processor.tokenizer.decode(generate_ids, skip_special_tokens=True).strip()
generate_ids = generate_ids[inputs["input_ids"].shape[1] :]
return processor.tokenizer.decode(
generate_ids, skip_special_tokens=True
).strip()
def _parse_tags(self, raw_output: str) -> List[str]:
"""Parse raw model output into a list of clean tags.
@@ -657,22 +666,22 @@ class AutoTagService:
"""
# Patterns to exclude
exclude_prefixes = (
'copyright:',
'meta:',
'photo_',
'photo:',
"copyright:",
"meta:",
"photo_",
"photo:",
)
# Split by common delimiters
tags = []
# Handle comma-separated tags
for part in raw_output.split(','):
for part in raw_output.split(","):
tag = part.strip().lower()
# Remove any quotes or extra characters
tag = tag.strip('"\'')
tag = tag.strip("\"'")
# Replace spaces with underscores (Danbooru style)
tag = tag.replace(' ', '_')
tag = tag.replace(" ", "_")
# Remove empty tags
if tag and len(tag) > 1:
# Filter out unwanted tag patterns
+150 -146
View File
@@ -23,6 +23,7 @@ extension_name = "PromptManager"
# Get server instance and routes (same pattern as ComfyUI_Assets)
from server import PromptServer
server_instance = PromptServer.instance
routes = server_instance.routes
@@ -37,24 +38,25 @@ try:
from ..utils.logging_config import get_logger
except ImportError:
import sys
current_dir = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
sys.path.insert(0, current_dir)
from utils.logging_config import get_logger
# Initialize logger for config operations
config_logger = get_logger('prompt_manager.config')
config_logger = get_logger("prompt_manager.config")
class GalleryConfig:
"""Configuration class for the gallery monitoring and image processing system.
This class manages all settings related to automatic image monitoring,
prompt tracking, database cleanup, web interface display, and performance
optimization for the gallery functionality.
All configuration values are class attributes that can be modified at runtime
or loaded from external configuration files.
Attributes:
MONITORING_ENABLED (bool): Enable/disable automatic image monitoring
MONITORING_DIRECTORIES (List[str]): Directories to monitor for new images
@@ -71,86 +73,88 @@ class GalleryConfig:
MAX_CONCURRENT_PROCESSING (int): Maximum concurrent image processing tasks
METADATA_EXTRACTION_TIMEOUT (int): Timeout for metadata extraction operations
"""
# Image monitoring settings
MONITORING_ENABLED = True
MONITORING_DIRECTORIES = [] # Auto-detect if empty
SUPPORTED_EXTENSIONS = ['.png', '.jpg', '.jpeg', '.webp', '.gif']
SUPPORTED_EXTENSIONS = [".png", ".jpg", ".jpeg", ".webp", ".gif"]
PROCESSING_DELAY = 2.0 # Seconds to wait before processing new files
# Prompt tracking settings
PROMPT_TIMEOUT = 600 # Seconds to keep prompt context active (10 min for long generations)
PROMPT_TIMEOUT = (
600 # Seconds to keep prompt context active (10 min for long generations)
)
CLEANUP_INTERVAL = 300 # Seconds between cleanup of expired prompts
# Database settings
AUTO_CLEANUP_MISSING_FILES = True
MAX_IMAGE_AGE_DAYS = 365 # Clean up images older than this
# Web interface settings
IMAGES_PER_PAGE = 20
THUMBNAIL_SIZE = 256
ENABLE_SEARCH = True
ENABLE_METADATA_VIEW = True
# Performance settings
MAX_CONCURRENT_PROCESSING = 3
METADATA_EXTRACTION_TIMEOUT = 10 # Seconds
@classmethod
def get_config(cls) -> Dict[str, Any]:
"""Get the complete gallery configuration as a structured dictionary.
Returns:
Dict[str, Any]: A nested dictionary containing all gallery configuration
sections: monitoring, tracking, database, web_interface, and performance.
Each section contains the relevant configuration parameters as key-value pairs.
Example:
config = GalleryConfig.get_config()
monitoring_enabled = config['monitoring']['enabled']
images_per_page = config['web_interface']['images_per_page']
"""
return {
'monitoring': {
'enabled': cls.MONITORING_ENABLED,
'directories': cls.MONITORING_DIRECTORIES,
'extensions': cls.SUPPORTED_EXTENSIONS,
'processing_delay': cls.PROCESSING_DELAY
"monitoring": {
"enabled": cls.MONITORING_ENABLED,
"directories": cls.MONITORING_DIRECTORIES,
"extensions": cls.SUPPORTED_EXTENSIONS,
"processing_delay": cls.PROCESSING_DELAY,
},
'tracking': {
'prompt_timeout': cls.PROMPT_TIMEOUT,
'cleanup_interval': cls.CLEANUP_INTERVAL
"tracking": {
"prompt_timeout": cls.PROMPT_TIMEOUT,
"cleanup_interval": cls.CLEANUP_INTERVAL,
},
'database': {
'auto_cleanup': cls.AUTO_CLEANUP_MISSING_FILES,
'max_image_age_days': cls.MAX_IMAGE_AGE_DAYS
"database": {
"auto_cleanup": cls.AUTO_CLEANUP_MISSING_FILES,
"max_image_age_days": cls.MAX_IMAGE_AGE_DAYS,
},
'web_interface': {
'images_per_page': cls.IMAGES_PER_PAGE,
'thumbnail_size': cls.THUMBNAIL_SIZE,
'enable_search': cls.ENABLE_SEARCH,
'enable_metadata_view': cls.ENABLE_METADATA_VIEW
"web_interface": {
"images_per_page": cls.IMAGES_PER_PAGE,
"thumbnail_size": cls.THUMBNAIL_SIZE,
"enable_search": cls.ENABLE_SEARCH,
"enable_metadata_view": cls.ENABLE_METADATA_VIEW,
},
"performance": {
"max_concurrent_processing": cls.MAX_CONCURRENT_PROCESSING,
"metadata_extraction_timeout": cls.METADATA_EXTRACTION_TIMEOUT,
},
'performance': {
'max_concurrent_processing': cls.MAX_CONCURRENT_PROCESSING,
'metadata_extraction_timeout': cls.METADATA_EXTRACTION_TIMEOUT
}
}
@classmethod
def update_config(cls, new_config: Dict[str, Any]):
"""Update gallery configuration attributes from a dictionary.
Takes a nested dictionary with gallery configuration sections and updates
the corresponding class attributes. Only updates attributes that are
present in the input dictionary, leaving others unchanged.
Args:
new_config (Dict[str, Any]): Nested dictionary containing gallery
configuration updates. Should follow the same structure as returned
by get_config(). Valid top-level keys are: 'monitoring', 'tracking',
'database', 'web_interface', 'performance'.
Example:
gallery_settings = {
'monitoring': {'enabled': False},
@@ -158,58 +162,58 @@ class GalleryConfig:
}
GalleryConfig.update_config(gallery_settings)
"""
monitoring = new_config.get('monitoring', {})
if 'enabled' in monitoring:
cls.MONITORING_ENABLED = monitoring['enabled']
if 'directories' in monitoring:
cls.MONITORING_DIRECTORIES = monitoring['directories']
if 'extensions' in monitoring:
cls.SUPPORTED_EXTENSIONS = monitoring['extensions']
if 'processing_delay' in monitoring:
cls.PROCESSING_DELAY = monitoring['processing_delay']
tracking = new_config.get('tracking', {})
if 'prompt_timeout' in tracking:
cls.PROMPT_TIMEOUT = tracking['prompt_timeout']
if 'cleanup_interval' in tracking:
cls.CLEANUP_INTERVAL = tracking['cleanup_interval']
database = new_config.get('database', {})
if 'auto_cleanup' in database:
cls.AUTO_CLEANUP_MISSING_FILES = database['auto_cleanup']
if 'max_image_age_days' in database:
cls.MAX_IMAGE_AGE_DAYS = database['max_image_age_days']
web_interface = new_config.get('web_interface', {})
if 'images_per_page' in web_interface:
cls.IMAGES_PER_PAGE = web_interface['images_per_page']
if 'thumbnail_size' in web_interface:
cls.THUMBNAIL_SIZE = web_interface['thumbnail_size']
if 'enable_search' in web_interface:
cls.ENABLE_SEARCH = web_interface['enable_search']
if 'enable_metadata_view' in web_interface:
cls.ENABLE_METADATA_VIEW = web_interface['enable_metadata_view']
performance = new_config.get('performance', {})
if 'max_concurrent_processing' in performance:
cls.MAX_CONCURRENT_PROCESSING = performance['max_concurrent_processing']
if 'metadata_extraction_timeout' in performance:
cls.METADATA_EXTRACTION_TIMEOUT = performance['metadata_extraction_timeout']
monitoring = new_config.get("monitoring", {})
if "enabled" in monitoring:
cls.MONITORING_ENABLED = monitoring["enabled"]
if "directories" in monitoring:
cls.MONITORING_DIRECTORIES = monitoring["directories"]
if "extensions" in monitoring:
cls.SUPPORTED_EXTENSIONS = monitoring["extensions"]
if "processing_delay" in monitoring:
cls.PROCESSING_DELAY = monitoring["processing_delay"]
tracking = new_config.get("tracking", {})
if "prompt_timeout" in tracking:
cls.PROMPT_TIMEOUT = tracking["prompt_timeout"]
if "cleanup_interval" in tracking:
cls.CLEANUP_INTERVAL = tracking["cleanup_interval"]
database = new_config.get("database", {})
if "auto_cleanup" in database:
cls.AUTO_CLEANUP_MISSING_FILES = database["auto_cleanup"]
if "max_image_age_days" in database:
cls.MAX_IMAGE_AGE_DAYS = database["max_image_age_days"]
web_interface = new_config.get("web_interface", {})
if "images_per_page" in web_interface:
cls.IMAGES_PER_PAGE = web_interface["images_per_page"]
if "thumbnail_size" in web_interface:
cls.THUMBNAIL_SIZE = web_interface["thumbnail_size"]
if "enable_search" in web_interface:
cls.ENABLE_SEARCH = web_interface["enable_search"]
if "enable_metadata_view" in web_interface:
cls.ENABLE_METADATA_VIEW = web_interface["enable_metadata_view"]
performance = new_config.get("performance", {})
if "max_concurrent_processing" in performance:
cls.MAX_CONCURRENT_PROCESSING = performance["max_concurrent_processing"]
if "metadata_extraction_timeout" in performance:
cls.METADATA_EXTRACTION_TIMEOUT = performance["metadata_extraction_timeout"]
class PromptManagerConfig:
"""Main configuration class for PromptManager core functionality.
This class manages configuration for database operations, web UI behavior,
performance settings, and integrates gallery configuration. It provides
methods for loading and saving configuration from/to JSON files.
The configuration is organized into logical sections:
- Database: Settings for SQLite operations and data management
- Web UI: User interface behavior and display options
- Performance: Optimization and resource management settings
- Gallery: Embedded gallery configuration (via GalleryConfig)
Attributes:
DEFAULT_DB_PATH (str): Default path for the SQLite database file
ENABLE_DUPLICATE_DETECTION (bool): Enable automatic duplicate detection
@@ -221,82 +225,82 @@ class PromptManagerConfig:
ENABLE_FUZZY_SEARCH (bool): Enable fuzzy search capabilities
AUTO_BACKUP_INTERVAL (int): Hours between automatic database backups
"""
# Database settings
DEFAULT_DB_PATH = "prompts.db"
ENABLE_DUPLICATE_DETECTION = True
ENABLE_AUTO_SAVE = True
# Web UI settings
RESULT_TIMEOUT = 5 # Seconds to auto-hide results in ComfyUI node
SHOW_TEST_BUTTON = False # Show API test button in node UI
WEBUI_DISPLAY_MODE = 'newtab' # 'popup' or 'newtab'
WEBUI_DISPLAY_MODE = "newtab" # 'popup' or 'newtab'
# Performance settings
MAX_SEARCH_RESULTS = 100
ENABLE_FUZZY_SEARCH = False # Requires fuzzywuzzy
AUTO_BACKUP_INTERVAL = 24 # Hours
@classmethod
def get_config(cls) -> Dict[str, Any]:
"""Get the complete PromptManager configuration as a structured dictionary.
Returns:
Dict[str, Any]: A nested dictionary containing all configuration sections:
- database: Database-related settings
- web_ui: Web interface configuration
- performance: Performance and optimization settings
- gallery: Complete gallery configuration (from GalleryConfig)
Example:
config = PromptManagerConfig.get_config()
db_path = config['database']['default_path']
max_results = config['performance']['max_search_results']
"""
return {
'database': {
'default_path': cls.DEFAULT_DB_PATH,
'enable_duplicate_detection': cls.ENABLE_DUPLICATE_DETECTION,
'enable_auto_save': cls.ENABLE_AUTO_SAVE
"database": {
"default_path": cls.DEFAULT_DB_PATH,
"enable_duplicate_detection": cls.ENABLE_DUPLICATE_DETECTION,
"enable_auto_save": cls.ENABLE_AUTO_SAVE,
},
'web_ui': {
'result_timeout': cls.RESULT_TIMEOUT,
'show_test_button': cls.SHOW_TEST_BUTTON,
'webui_display_mode': cls.WEBUI_DISPLAY_MODE
"web_ui": {
"result_timeout": cls.RESULT_TIMEOUT,
"show_test_button": cls.SHOW_TEST_BUTTON,
"webui_display_mode": cls.WEBUI_DISPLAY_MODE,
},
'performance': {
'max_search_results': cls.MAX_SEARCH_RESULTS,
'enable_fuzzy_search': cls.ENABLE_FUZZY_SEARCH,
'auto_backup_interval': cls.AUTO_BACKUP_INTERVAL
"performance": {
"max_search_results": cls.MAX_SEARCH_RESULTS,
"enable_fuzzy_search": cls.ENABLE_FUZZY_SEARCH,
"auto_backup_interval": cls.AUTO_BACKUP_INTERVAL,
},
'gallery': GalleryConfig.get_config()
"gallery": GalleryConfig.get_config(),
}
@classmethod
def load_from_file(cls, config_path: str):
"""Load configuration settings from a JSON file.
Reads configuration from the specified JSON file and updates the current
configuration attributes. If the file doesn't exist or contains invalid
JSON, logs an appropriate message and continues with default values.
Args:
config_path (str): Path to the JSON configuration file to load.
Can be relative or absolute path.
Raises:
The method handles all exceptions internally and logs errors rather
than propagating them, ensuring the system continues with defaults.
Example:
PromptManagerConfig.load_from_file('custom_config.json')
PromptManagerConfig.load_from_file('/path/to/config.json')
"""
import json
if os.path.exists(config_path):
try:
with open(config_path, 'r') as f:
with open(config_path, "r") as f:
config = json.load(f)
cls.update_config(config)
config_logger.info(f"Loaded configuration from {config_path}")
@@ -304,52 +308,52 @@ class PromptManagerConfig:
config_logger.error(f"Error loading config from {config_path}: {e}")
else:
config_logger.info(f"Config file not found: {config_path}, using defaults")
@classmethod
def save_to_file(cls, config_path: str):
"""Save the current configuration to a JSON file.
Serializes the complete configuration (including gallery settings) to
a JSON file. Creates the directory structure if it doesn't exist.
Args:
config_path (str): Path where the JSON configuration file should be saved.
Parent directories will be created if they don't exist.
Raises:
The method handles all exceptions internally and logs errors rather
than propagating them.
Example:
PromptManagerConfig.save_to_file('backup_config.json')
PromptManagerConfig.save_to_file('/etc/comfyui/prompt_manager.json')
"""
import json
try:
config = cls.get_config()
os.makedirs(os.path.dirname(config_path), exist_ok=True)
with open(config_path, 'w') as f:
with open(config_path, "w") as f:
json.dump(config, f, indent=2)
config_logger.info(f"Saved configuration to {config_path}")
except Exception as e:
config_logger.error(f"Error saving config to {config_path}: {e}")
@classmethod
def update_config(cls, new_config: Dict[str, Any]):
"""Update configuration attributes from a dictionary.
Takes a nested dictionary with configuration sections and updates
the corresponding class attributes. Only updates attributes that
are present in the input dictionary, leaving others unchanged.
Args:
new_config (Dict[str, Any]): Nested dictionary containing configuration
updates. Should follow the same structure as returned by get_config().
Valid top-level keys are: 'database', 'web_ui', 'performance', 'gallery'.
Example:
new_settings = {
'database': {'default_path': 'custom.db'},
@@ -357,39 +361,39 @@ class PromptManagerConfig:
}
PromptManagerConfig.update_config(new_settings)
"""
database = new_config.get('database', {})
if 'default_path' in database:
cls.DEFAULT_DB_PATH = database['default_path']
if 'enable_duplicate_detection' in database:
cls.ENABLE_DUPLICATE_DETECTION = database['enable_duplicate_detection']
if 'enable_auto_save' in database:
cls.ENABLE_AUTO_SAVE = database['enable_auto_save']
web_ui = new_config.get('web_ui', {})
if 'result_timeout' in web_ui:
cls.RESULT_TIMEOUT = web_ui['result_timeout']
if 'show_test_button' in web_ui:
cls.SHOW_TEST_BUTTON = web_ui['show_test_button']
if 'webui_display_mode' in web_ui:
cls.WEBUI_DISPLAY_MODE = web_ui['webui_display_mode']
performance = new_config.get('performance', {})
if 'max_search_results' in performance:
cls.MAX_SEARCH_RESULTS = performance['max_search_results']
if 'enable_fuzzy_search' in performance:
cls.ENABLE_FUZZY_SEARCH = performance['enable_fuzzy_search']
if 'auto_backup_interval' in performance:
cls.AUTO_BACKUP_INTERVAL = performance['auto_backup_interval']
database = new_config.get("database", {})
if "default_path" in database:
cls.DEFAULT_DB_PATH = database["default_path"]
if "enable_duplicate_detection" in database:
cls.ENABLE_DUPLICATE_DETECTION = database["enable_duplicate_detection"]
if "enable_auto_save" in database:
cls.ENABLE_AUTO_SAVE = database["enable_auto_save"]
web_ui = new_config.get("web_ui", {})
if "result_timeout" in web_ui:
cls.RESULT_TIMEOUT = web_ui["result_timeout"]
if "show_test_button" in web_ui:
cls.SHOW_TEST_BUTTON = web_ui["show_test_button"]
if "webui_display_mode" in web_ui:
cls.WEBUI_DISPLAY_MODE = web_ui["webui_display_mode"]
performance = new_config.get("performance", {})
if "max_search_results" in performance:
cls.MAX_SEARCH_RESULTS = performance["max_search_results"]
if "enable_fuzzy_search" in performance:
cls.ENABLE_FUZZY_SEARCH = performance["enable_fuzzy_search"]
if "auto_backup_interval" in performance:
cls.AUTO_BACKUP_INTERVAL = performance["auto_backup_interval"]
# Update gallery config
if 'gallery' in new_config:
GalleryConfig.update_config(new_config['gallery'])
if "gallery" in new_config:
GalleryConfig.update_config(new_config["gallery"])
# Load configuration on import
try:
config_dir = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
config_file = os.path.join(config_dir, 'config.json')
config_file = os.path.join(config_dir, "config.json")
PromptManagerConfig.load_from_file(config_file)
except Exception as e:
config_logger.error(f"Error during config initialization: {e}")
config_logger.error(f"Error during config initialization: {e}")