From 56eed388a8fa8f0658303f3928b5f1fec2dd2309 Mon Sep 17 00:00:00 2001 From: Vito Sansevero Date: Sun, 7 Dec 2025 17:25:06 -0800 Subject: [PATCH] feat(api): add AutoTag endpoints and image routes --- py/api.py | 515 +++++++++++++++++++++++++++++++++++++++++++++++++++++- 1 file changed, 511 insertions(+), 4 deletions(-) diff --git a/py/api.py b/py/api.py index 100b1d1..f177be1 100644 --- a/py/api.py +++ b/py/api.py @@ -374,6 +374,10 @@ class PromptManagerAPI: async def get_recent_images_route(request): return await self.get_recent_images(request) + @routes.get("/prompt_manager/images/all") + async def get_all_images_route(request): + return await self.get_all_images(request) + @routes.get("/prompt_manager/images/search") async def search_images_route(request): return await self.search_images(request) @@ -457,6 +461,31 @@ class PromptManagerAPI: async def get_log_stats_route(request): return await self.get_log_stats(request) + # AutoTag endpoints + @routes.get("/prompt_manager/autotag/models") + async def get_autotag_models_route(request): + return await self.get_autotag_models(request) + + @routes.get("/prompt_manager/autotag/download/{model_type}") + async def download_autotag_model_route(request): + return await self.download_autotag_model(request) + + @routes.get("/prompt_manager/autotag/start") + async def start_autotag_route(request): + return await self.start_autotag(request) + + @routes.post("/prompt_manager/autotag/single") + async def autotag_single_route(request): + return await self.autotag_single(request) + + @routes.post("/prompt_manager/autotag/apply") + async def apply_autotag_route(request): + return await self.apply_autotag(request) + + @routes.get("/prompt_manager/scan_output_dir") + async def scan_output_dir_route(request): + return await self.scan_output_dir(request) + self.logger.info("All routes registered with decorator pattern") async def search_prompts(self, request): @@ -1754,7 +1783,7 @@ class PromptManagerAPI: try: limit = int(request.query.get('limit', 50)) images = self.db.get_recent_images(limit) - + return web.json_response({ 'success': True, 'images': images @@ -1766,6 +1795,23 @@ class PromptManagerAPI: 'error': str(e) }, status=500) + async def get_all_images(self, request): + """Get all generated images with linked prompts.""" + try: + images = self.db.get_all_images() + + return web.json_response({ + 'success': True, + 'images': images, + 'count': len(images) + }) + except Exception as e: + self.logger.error(f"Get all images error: {e}") + return web.json_response({ + 'success': False, + 'error': str(e) + }, status=500) + async def search_images(self, request): """Search images by prompt text.""" try: @@ -4195,18 +4241,479 @@ class PromptManagerAPI: 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_manager - + logger_manager = get_logger_manager() stats = logger_manager.get_log_stats() - + return web.json_response({ 'success': True, 'stats': stats }) - + except Exception as e: self.logger.error(f"Get log stats error: {e}") return web.json_response({ 'success': False, 'error': str(e) }, status=500) + + # ======================================================================== + # AutoTag Endpoints + # ======================================================================== + + async def get_autotag_models(self, request): + """ + Get status of available AutoTag models. + GET /prompt_manager/autotag/models + + Returns model availability, download status, and configuration. + """ + try: + from .autotag import get_autotag_service + + service = get_autotag_service() + models_status = service.get_models_status() + + return web.json_response({ + 'success': True, + 'models': models_status, + 'default_prompt': service.default_prompt + }) + + except Exception as e: + self.logger.error(f"Get autotag models error: {e}") + return web.json_response({ + 'success': False, + 'error': str(e) + }, status=500) + + async def download_autotag_model(self, request): + """ + Download an AutoTag model with streaming progress. + POST /prompt_manager/autotag/download/{model_type} + + Streams SSE progress updates during download. + """ + import json + import asyncio + + model_type = request.match_info.get('model_type') + + async def stream_response(): + try: + from .autotag import get_autotag_service + + service = get_autotag_service() + + if model_type not in service.models_config: + yield f"data: {json.dumps({'type': 'error', 'message': f'Invalid model type: {model_type}'})}\n\n" + return + + yield f"data: {json.dumps({'type': 'progress', 'progress': 0, 'status': 'Starting download...'})}\n\n" + + # Define progress callback + progress_data = {'last_progress': 0} + + def progress_callback(status: str, progress: float): + progress_data['last_progress'] = progress + + # Perform download (blocking in thread) + loop = asyncio.get_event_loop() + success = await loop.run_in_executor( + None, + lambda: service.download_model(model_type, progress_callback) + ) + + if success: + yield f"data: {json.dumps({'type': 'complete', 'progress': 100, 'status': 'Download complete'})}\n\n" + else: + yield f"data: {json.dumps({'type': 'error', 'message': 'Download failed'})}\n\n" + + except Exception as e: + self.logger.error(f"Download model error: {e}") + yield f"data: {json.dumps({'type': 'error', 'message': str(e)})}\n\n" + + response = web.StreamResponse( + status=200, + reason='OK', + headers={ + 'Content-Type': 'text/event-stream', + 'Cache-Control': 'no-cache', + 'Connection': 'keep-alive' + } + ) + + await response.prepare(request) + + async for chunk in stream_response(): + await response.write(chunk.encode('utf-8')) + + await response.write_eof() + return response + + async def start_autotag(self, request): + """ + Start batch auto-tagging with streaming progress. + GET /prompt_manager/autotag/start + + Query params: + model_type: "gguf" or "hf" + prompt: custom prompt text + + Streams SSE progress updates during processing. + """ + import json + import asyncio + from pathlib import Path + + # Read from query params (EventSource only supports GET) + model_type = request.query.get('model_type', 'gguf') + custom_prompt = request.query.get('prompt', '') + skip_tagged = request.query.get('skip_tagged', 'true').lower() == 'true' + use_gpu = True + + async def stream_response(): + try: + from .autotag import get_autotag_service + + service = get_autotag_service() + + # Check model is downloaded + status = service.get_models_status() + if not status.get(model_type, {}).get('downloaded'): + yield f"data: {json.dumps({'type': 'error', 'message': f'Model {model_type} not downloaded'})}\n\n" + return + + yield f"data: {json.dumps({'type': 'progress', 'progress': 0, 'status': 'Loading model...'})}\n\n" + + # Load model in thread pool + loop = asyncio.get_event_loop() + try: + await loop.run_in_executor( + None, + lambda: service.load_model(model_type, use_gpu) + ) + except Exception as e: + yield f"data: {json.dumps({'type': 'error', 'message': f'Failed to load model: {str(e)}'})}\n\n" + return + + # Set custom prompt if provided + if custom_prompt: + service.custom_prompt = custom_prompt + + yield f"data: {json.dumps({'type': 'progress', 'progress': 5, 'status': 'Model loaded. Fetching all images from database...'})}\n\n" + + # Get ALL images from database (only images with linked prompts) + images = self.db.get_all_images() + + total_files = len(images) + if total_files == 0: + yield f"data: {json.dumps({'type': 'complete', 'processed': 0, 'tagged': 0, 'skipped': 0, 'status': 'No images with linked prompts found'})}\n\n" + service.unload_model() + return + + yield f"data: {json.dumps({'type': 'progress', 'progress': 10, 'status': f'Found {total_files} images. Processing...'})}\n\n" + + processed = 0 + tagged = 0 + skipped = 0 + errors = 0 + + for i, image_data in enumerate(images): + image_path = image_data.get('image_path') + prompt_id = image_data.get('prompt_id') + + if not image_path or not prompt_id: + skipped += 1 + continue + + # Check if file exists + if not Path(image_path).exists(): + skipped += 1 + continue + + # Check if image already has real tags (skip_tagged option) + if skip_tagged: + prompt_tags = image_data.get('prompt_tags', []) + if isinstance(prompt_tags, str): + prompt_tags = [t.strip() for t in prompt_tags.split(',') if t.strip()] + # Filter out "auto-scanned" - it's not a real tag + real_tags = [t for t in prompt_tags if t != 'auto-scanned'] + if real_tags: + skipped += 1 + if i % 5 == 0 or i == total_files - 1: + progress = 10 + int((i + 1) / total_files * 85) + yield f"data: {json.dumps({'type': 'progress', 'progress': progress, 'status': f'Skipping {i+1}/{total_files} (already tagged)...', 'processed': processed, 'tagged': tagged, 'skipped': skipped})}\n\n" + await asyncio.sleep(0.01) + continue + + try: + # Generate tags + tags = await loop.run_in_executor( + None, + lambda p=str(image_path): service.generate_tags(p) + ) + + processed += 1 + + if tags: + # Get existing prompt + existing_prompt = self.db.get_prompt_by_id(prompt_id) + if existing_prompt: + existing_tags = existing_prompt.get('tags', []) + if isinstance(existing_tags, str): + existing_tags = [t.strip() for t in existing_tags.split(',') if t.strip()] + + new_tags = [t for t in tags if t not in existing_tags] + if new_tags: + all_tags = existing_tags + new_tags + self.db.update_prompt_metadata( + prompt_id, + tags=all_tags + ) + tagged += 1 + else: + skipped += 1 + else: + skipped += 1 + else: + skipped += 1 + + if i % 5 == 0 or i == total_files - 1: + progress = 10 + int((i + 1) / total_files * 85) + yield f"data: {json.dumps({'type': 'progress', 'progress': progress, 'status': f'Processing {i+1}/{total_files}...', 'processed': processed, 'tagged': tagged, 'skipped': skipped})}\n\n" + await asyncio.sleep(0.01) + + except Exception as img_err: + self.logger.error(f"Error processing {image_path}: {img_err}") + errors += 1 + processed += 1 + + service.unload_model() + + yield f"data: {json.dumps({'type': 'complete', 'progress': 100, 'processed': processed, 'tagged': tagged, 'skipped': skipped, 'errors': errors, 'status': 'Complete'})}\n\n" + + except Exception as e: + self.logger.error(f"AutoTag error: {e}") + import traceback + self.logger.error(f"AutoTag traceback: {traceback.format_exc()}") + yield f"data: {json.dumps({'type': 'error', 'message': str(e)})}\n\n" + + response = web.StreamResponse( + status=200, + reason='OK', + headers={ + 'Content-Type': 'text/event-stream', + 'Cache-Control': 'no-cache', + 'Connection': 'keep-alive' + } + ) + + await response.prepare(request) + + async for chunk in stream_response(): + await response.write(chunk.encode('utf-8')) + + await response.write_eof() + return response + + async def autotag_single(self, request): + """ + Generate tags for a single image (for Review mode). + POST /prompt_manager/autotag/single + + Request body: + { + "image_path": "/path/to/image.png", + "model_type": "gguf", + "prompt": "optional custom prompt", + "use_gpu": true + } + """ + try: + data = await request.json() + image_path = data.get('image_path') + model_type = data.get('model_type', 'gguf') + custom_prompt = data.get('prompt') + use_gpu = data.get('use_gpu', True) + + if not image_path: + return web.json_response({ + 'success': False, + 'error': 'image_path is required' + }, status=400) + + from .autotag import get_autotag_service + import asyncio + + service = get_autotag_service() + + if not service.is_model_loaded() or service.get_loaded_model_type() != model_type: + loop = asyncio.get_event_loop() + await loop.run_in_executor( + None, + lambda: service.load_model(model_type, use_gpu) + ) + + if custom_prompt: + service.custom_prompt = custom_prompt + + loop = asyncio.get_event_loop() + tags = await loop.run_in_executor( + None, + lambda: service.generate_tags(image_path) + ) + + prompt_id = None + try: + with self.db.model.get_connection() as conn: + cursor = conn.execute( + "SELECT prompt_id FROM generated_images WHERE image_path = ?", + (image_path,) + ) + row = cursor.fetchone() + if row: + prompt_id = row[0] + except Exception as e: + self.logger.warning(f"Could not find linked prompt: {e}") + + return web.json_response({ + 'success': True, + 'tags': tags, + 'prompt_id': prompt_id, + 'image_path': image_path + }) + + except Exception as e: + self.logger.error(f"AutoTag single error: {e}") + return web.json_response({ + 'success': False, + 'error': str(e) + }, status=500) + + async def apply_autotag(self, request): + """ + Apply selected tags to a prompt. + POST /prompt_manager/autotag/apply + + Request body: + { + "prompt_id": 123, + "tags": ["tag1", "tag2", ...] + } + """ + try: + data = await request.json() + prompt_id = data.get('prompt_id') + tags = data.get('tags', []) + + if not prompt_id: + return web.json_response({ + 'success': False, + 'error': 'prompt_id is required' + }, status=400) + + if not tags: + return web.json_response({ + 'success': True, + 'message': 'No tags to apply' + }) + + prompt = self.db.get_prompt_by_id(prompt_id) + if not prompt: + return web.json_response({ + 'success': False, + 'error': f'Prompt {prompt_id} not found' + }, status=404) + + existing_tags = prompt.get('tags', []) + if isinstance(existing_tags, str): + existing_tags = [t.strip() for t in existing_tags.split(',') if t.strip()] + + new_tags = [t for t in tags if t not in existing_tags] + all_tags = existing_tags + new_tags + + self.db.update_prompt_metadata( + prompt_id, + tags=all_tags + ) + + return web.json_response({ + 'success': True, + 'added_tags': new_tags, + 'total_tags': len(all_tags) + }) + + except Exception as e: + self.logger.error(f"Apply autotag error: {e}") + return web.json_response({ + 'success': False, + 'error': str(e) + }, status=500) + + async def scan_output_dir(self, request): + """ + Scan ComfyUI output directory for images. + GET /prompt_manager/scan_output_dir + + Returns a list of images in the output directory for autotag review mode. + """ + from pathlib import Path + + try: + output_dir = self._find_comfyui_output_dir() + if not output_dir: + return web.json_response({ + 'success': False, + 'error': 'ComfyUI output directory not found' + }, status=404) + + output_path = Path(output_dir) + image_extensions = ['.png', '.jpg', '.jpeg', '.gif', '.webp', '.bmp', '.tiff'] + + images = [] + seen_paths = set() + + for ext in image_extensions: + for pattern in [f"*{ext.lower()}", f"*{ext.upper()}"]: + for image_path in output_path.rglob(pattern): + if 'thumbnails' not in image_path.parts: + normalized_path = str(image_path).lower() + if normalized_path not in seen_paths: + seen_paths.add(normalized_path) + + rel_path = image_path.relative_to(output_path) + + # Check for thumbnail + thumbnail_url = None + thumbnails_dir = output_path / "thumbnails" + if thumbnails_dir.exists(): + thumbnail_path = thumbnails_dir / f"{image_path.stem}_thumb{image_path.suffix}" + if thumbnail_path.exists(): + thumbnail_url = f'/prompt_manager/images/serve/thumbnails/{image_path.stem}_thumb{image_path.suffix}' + + images.append({ + 'filename': image_path.name, + 'path': str(image_path), + 'relative_path': str(rel_path), + 'url': f'/prompt_manager/images/serve/{rel_path.as_posix()}', + 'thumbnail_url': thumbnail_url + }) + + # Sort by filename + images.sort(key=lambda x: x['filename']) + + self.logger.info(f"Found {len(images)} images in output directory") + + return web.json_response({ + 'success': True, + 'images': images, + 'count': len(images) + }) + + except Exception as e: + self.logger.error(f"Scan output dir error: {e}") + return web.json_response({ + 'success': False, + 'error': str(e) + }, status=500)