"""Retouch service: upload, process via Flux.1-Pro, and retrieve results.""" from __future__ import annotations import base64 import logging import os import uuid from typing import Any import httpx from sqlalchemy import func, select from sqlalchemy.ext.asyncio import AsyncSession from app.config import settings from app.models.retouch import RetouchResult, RetouchStatus logger = logging.getLogger(__name__) ALLOWED_MIME_TYPES = {"image/png", "image/jpeg", "image/jpg", "image/webp"} RETOUCH_SYSTEM_PROMPT = ( "You are an expert automotive photo retoucher. " "Enhance the provided vehicle image with the following improvements: " "1. Remove or replace the background with a clean, professional studio backdrop. " "2. Correct color balance and enhance saturation for a vibrant, realistic look. " "3. Remove reflections and glare from windows and bodywork. " "4. Clean up minor blemishes, dust, and scratches on the vehicle surface. " "5. Ensure the vehicle is well-lit and centered. " "Return the enhanced image." ) def validate_mime_type(mime_type: str) -> bool: """Check if the MIME type is an allowed image type.""" return mime_type.lower() in ALLOWED_MIME_TYPES def validate_file_size(file_size: int) -> bool: """Check if file size is within the configured limit.""" max_bytes = settings.MAX_FILE_SIZE_MB * 1024 * 1024 return file_size <= max_bytes def _get_file_extension(mime_type: str) -> str: """Map MIME type to file extension.""" mapping = { "image/png": ".png", "image/jpeg": ".jpg", "image/jpg": ".jpg", "image/webp": ".webp", } return mapping.get(mime_type.lower(), ".png") def generate_retouch_prompt(vehicle_info: dict[str, Any] | None = None) -> str: """Generate a retouch prompt based on optional vehicle info. Args: vehicle_info: Optional dict with make, model, year, color etc. Returns: A detailed prompt string for the image retouch model. """ base_prompt = ( "Professional automotive photograph retouching: " "Remove the original background and replace with a clean white studio backdrop. " "Enhance color correction for realistic, vibrant tones. " "Remove reflections and glare from windows and paint. " "Clean up dust, scratches, and minor blemishes on the bodywork. " "Ensure even, professional lighting across the entire vehicle. " "Sharpen details on wheels, grille, and badges." ) if vehicle_info: parts = [base_prompt] make = vehicle_info.get("make") model = vehicle_info.get("model") color = vehicle_info.get("color") if make and model: parts.append(f"The vehicle is a {make} {model}.") if color: parts.append( f"The vehicle color is {color}; ensure it looks accurate and rich." ) return " ".join(parts) return base_prompt async def upload_retouch_file( db: AsyncSession, file_bytes: bytes, file_name: str, mime_type: str, vehicle_id: uuid.UUID | None = None, ) -> RetouchResult: """Save uploaded image and create a RetouchResult record. Raises ValueError for invalid MIME type or file size. """ if not validate_mime_type(mime_type): raise ValueError( f"Invalid MIME type: {mime_type}. Only image/* types are allowed." ) if not validate_file_size(len(file_bytes)): raise ValueError(f"File size exceeds limit of {settings.MAX_FILE_SIZE_MB} MB") upload_dir = settings.UPLOAD_DIR os.makedirs(upload_dir, exist_ok=True) ext = _get_file_extension(mime_type) stored_filename = f"retouch_{uuid.uuid4().hex}{ext}" file_path = os.path.join(upload_dir, stored_filename) with open(file_path, "wb") as f: f.write(file_bytes) result = RetouchResult( original_file_path=file_path, original_file_name=file_name, mime_type=mime_type, status=RetouchStatus.pending.value, vehicle_id=vehicle_id, ) db.add(result) await db.flush() await db.refresh(result) return result async def get_result(db: AsyncSession, result_id: uuid.UUID) -> RetouchResult | None: """Get a retouch result by ID.""" stmt = select(RetouchResult).where(RetouchResult.id == result_id) res = await db.execute(stmt) return res.scalar_one_or_none() async def list_results( db: AsyncSession, vehicle_id: uuid.UUID | None = None, page: int = 1, page_size: int = 20, ) -> tuple[list[RetouchResult], int]: """List retouch results with optional vehicle filter and pagination.""" count_stmt = select(func.count(RetouchResult.id)) data_stmt = select(RetouchResult) if vehicle_id is not None: count_stmt = count_stmt.where(RetouchResult.vehicle_id == vehicle_id) data_stmt = data_stmt.where(RetouchResult.vehicle_id == vehicle_id) total_result = await db.execute(count_stmt) total = total_result.scalar_one() offset = (page - 1) * page_size data_stmt = ( data_stmt.order_by(RetouchResult.created_at.desc()) .offset(offset) .limit(page_size) ) result = await db.execute(data_stmt) items = list(result.scalars().all()) return items, total async def process_retouch( db: AsyncSession, result_id: uuid.UUID, vehicle_info: dict[str, Any] | None = None, ) -> RetouchResult: """Process a retouch result by sending the image to Flux.1-Pro via OpenRouter. Updates the result status to processing, calls the API, saves the retouched image, and sets status to completed. On failure, sets status to failed with an error message. """ result = await get_result(db, result_id) if result is None: raise ValueError(f"Retouch result {result_id} not found") result.status = RetouchStatus.processing.value await db.flush() try: with open(result.original_file_path, "rb") as f: image_bytes = f.read() retouched_bytes = await _call_flux_pro( image_bytes=image_bytes, mime_type=result.mime_type, prompt=generate_retouch_prompt(vehicle_info), ) upload_dir = settings.UPLOAD_DIR ext = _get_file_extension(result.mime_type) retouched_filename = f"retouched_{result_id.hex}{ext}" retouched_path = os.path.join(upload_dir, retouched_filename) with open(retouched_path, "wb") as f: f.write(retouched_bytes) result.retouched_file_path = retouched_path result.status = RetouchStatus.completed.value result.error_message = None await db.flush() await db.refresh(result) logger.info("Retouch completed for result %s", result_id) return result except Exception as exc: logger.error("Retouch failed for result %s: %s", result_id, exc) result.status = RetouchStatus.failed.value result.error_message = str(exc) await db.flush() await db.refresh(result) return result async def _call_flux_pro( image_bytes: bytes, mime_type: str, prompt: str, api_key: str | None = None, model: str | None = None, ) -> bytes: """Send image to Flux.1-Pro via OpenRouter for retouching. Returns the retouched image bytes. Raises httpx.HTTPStatusError or ValueError on failure. """ key = api_key or settings.OPENROUTER_API_KEY if not key: raise ValueError("OPENROUTER_API_KEY is not configured") model_name = model or "black-forest-labs/flux-1-pro" b64_image = base64.b64encode(image_bytes).decode("utf-8") data_uri = f"data:{mime_type};base64,{b64_image}" headers = { "Authorization": f"Bearer {key}", "Content-Type": "application/json", } payload: dict[str, Any] = { "model": model_name, "messages": [ { "role": "system", "content": RETOUCH_SYSTEM_PROMPT, }, { "role": "user", "content": [ { "type": "text", "text": prompt, }, { "type": "image_url", "image_url": {"url": data_uri}, }, ], }, ], "temperature": 0.7, "max_tokens": 4096, } base_url = settings.OPENROUTER_BASE_URL.rstrip("/") url = f"{base_url}/chat/completions" async with httpx.AsyncClient(timeout=httpx.Timeout(120.0)) as client: response = await client.post(url, headers=headers, json=payload) response.raise_for_status() body = response.json() content = body.get("choices", [{}])[0].get("message", {}).get("content", "") # The model may return a base64-encoded image or a URL if isinstance(content, list): for part in content: if isinstance(part, dict) and part.get("type") == "image_url": url_or_data = part.get("image_url", {}).get("url", "") if url_or_data.startswith("data:"): b64_data = url_or_data.split(",", 1)[1] return base64.b64decode(b64_data) elif url_or_data.startswith("http"): async with httpx.AsyncClient(timeout=60.0) as dl_client: dl_resp = await dl_client.get(url_or_data) dl_resp.raise_for_status() return dl_resp.content if isinstance(content, str) and content.startswith("data:"): b64_data = content.split(",", 1)[1] return base64.b64decode(b64_data) # Fallback: if no image returned, return original bytes logger.warning("Flux.1-Pro did not return an image, returning original bytes") return image_bytes