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:
+1
-1
@@ -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
@@ -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
File diff suppressed because it is too large
Load Diff
+145
-138
@@ -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
File diff suppressed because it is too large
Load Diff
+69
-86
@@ -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
@@ -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
@@ -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
@@ -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}")
|
||||
|
||||
Reference in New Issue
Block a user