from fastapi import APIRouter, Depends, HTTPException, UploadFile, File, Form from fastapi.responses import Response from sqlalchemy.orm import Session from typing import Optional from PIL import Image from io import BytesIO import numpy as np import json import base64 import cv2 from app.database import get_db from app.models.project import Project from app.schemas import StatusResponse router = APIRouter(prefix="/tools", tags=["tools"]) @router.post("/remove-background") async def remove_background( project_id: Optional[int] = Form(None), file: Optional[UploadFile] = File(None), db: Session = Depends(get_db) ): """ Remove background from an image using rembg. Either provide project_id to use current project image, or upload a file directly. Returns PNG with transparent background. """ try: from rembg import remove except ImportError: raise HTTPException( status_code=500, detail="rembg not installed. Run: pip install rembg" ) # Get image bytes if file: image_bytes = await file.read() elif project_id: project = db.query(Project).filter(Project.id == project_id).first() if not project: raise HTTPException(status_code=404, detail="Project not found") from app.services.edit_service import EditService edit_service = EditService() image_path = edit_service.get_current_image_path(project_id) with open(image_path, 'rb') as f: image_bytes = f.read() else: raise HTTPException( status_code=400, detail="Provide either project_id or file" ) # Remove background result_bytes = remove(image_bytes) return Response( content=result_bytes, media_type="image/png", headers={"Content-Disposition": "inline; filename=no-background.png"} ) @router.post("/remove-background-to-layer") async def remove_background_to_layer( project_id: int = Form(...), db: Session = Depends(get_db) ): """ Remove background and save as a new layer in the project. Returns layer info that can be added to frontend layer system. """ try: from rembg import remove except ImportError: raise HTTPException( status_code=500, detail="rembg not installed. Run: pip install rembg" ) project = db.query(Project).filter(Project.id == project_id).first() if not project: raise HTTPException(status_code=404, detail="Project not found") from app.services.edit_service import EditService from pathlib import Path edit_service = EditService() image_path = edit_service.get_current_image_path(project_id) with open(image_path, 'rb') as f: image_bytes = f.read() # Remove background result_bytes = remove(image_bytes) # Save as layer file project_dir = edit_service.get_project_dir(project_id) layers_dir = project_dir / 'layers' layers_dir.mkdir(exist_ok=True) # Find next layer number existing_layers = list(layers_dir.glob('layer_*.png')) layer_num = len(existing_layers) + 1 layer_path = layers_dir / f'layer_{layer_num}.png' with open(layer_path, 'wb') as f: f.write(result_bytes) # Get dimensions img = Image.open(BytesIO(result_bytes)) return { "status": "success", "layer": { "id": layer_num, "name": f"No Background {layer_num}", "path": str(layer_path), "width": img.width, "height": img.height, "type": "background_removed" } } @router.post("/smart-select") async def smart_select( project_id: int = Form(...), point_x: int = Form(...), point_y: int = Form(...), return_format: str = Form("json"), # "json" (default) or "image" db: Session = Depends(get_db) ): """ Use SAM (Segment Anything) to select object at given point. Returns mask and polygon data for the selected object. Note: Requires SAM model to be downloaded. Falls back to simple flood-fill selection if SAM unavailable. """ project = db.query(Project).filter(Project.id == project_id).first() if not project: raise HTTPException(status_code=404, detail="Project not found") from app.services.edit_service import EditService edit_service = EditService() image_path = edit_service.get_current_image_path(project_id) img = Image.open(image_path).convert('RGB') img_array = np.array(img) # Try SAM first, fall back to flood fill try: mask = await _sam_select(img_array, point_x, point_y) except Exception as e: print(f"SAM not available, using flood fill: {e}") mask = _flood_fill_select(img_array, point_x, point_y) # Convert mask to PNG mask_img = Image.fromarray((mask * 255).astype(np.uint8), mode='L') if return_format == "image": buffer = BytesIO() mask_img.save(buffer, format='PNG') return Response( content=buffer.getvalue(), media_type="image/png" ) # Return JSON with polygon and bbox polygon, bbox = _mask_to_polygon(mask) # Also return mask as base64 for potential use buffer = BytesIO() mask_img.save(buffer, format='PNG') mask_b64 = base64.b64encode(buffer.getvalue()).decode('utf-8') return { "polygon": polygon, "bbox": bbox, "mask_base64": mask_b64, } def _mask_to_polygon(mask: np.ndarray) -> tuple: """ Convert a binary mask to a simplified polygon and bounding box. Returns: (polygon, bbox) where: - polygon: list of [x, y] points (simplified contour) - bbox: dict with x, y, width, height """ # Ensure mask is binary uint8 mask_uint8 = (mask * 255).astype(np.uint8) if mask.max() <= 1 else mask.astype(np.uint8) # Find contours contours, _ = cv2.findContours(mask_uint8, cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_SIMPLE) if not contours: return [], {"x": 0, "y": 0, "width": 0, "height": 0} # Get largest contour largest = max(contours, key=cv2.contourArea) # Get bounding box x, y, w, h = cv2.boundingRect(largest) bbox = {"x": int(x), "y": int(y), "width": int(w), "height": int(h)} # Simplify contour to reduce points (epsilon = 1% of arc length) epsilon = 0.01 * cv2.arcLength(largest, True) simplified = cv2.approxPolyDP(largest, epsilon, True) # Convert to list of [x, y] points polygon = [[int(pt[0][0]), int(pt[0][1])] for pt in simplified] return polygon, bbox # Global SAM model cache (loaded once, reused) _sam_model = None _sam_predictor = None def _get_sam_model(): """Load SAM model from local file (cached after first load)""" global _sam_model, _sam_predictor if _sam_predictor is not None: return _sam_predictor from pathlib import Path # Check for SAM model in models directory models_dir = Path('/app/data/models') model_path = models_dir / 'sam_model.pth' # Also check for specific model files if not model_path.exists(): for filename in ['sam_vit_b_01ec64.pth', 'sam_vit_l_0b3195.pth', 'sam_vit_h_4b8939.pth']: alt_path = models_dir / filename if alt_path.exists(): model_path = alt_path break if not model_path.exists(): raise FileNotFoundError( f"SAM model not found. Download it with:\n" f" docker exec -it ai-photo-edit-backend python /scripts/download_sam_model.py" ) # Determine model type from filename model_type = 'vit_b' # default if 'vit_l' in model_path.name: model_type = 'vit_l' elif 'vit_h' in model_path.name: model_type = 'vit_h' print(f"Loading SAM model: {model_path} (type: {model_type})") import torch from segment_anything import sam_model_registry, SamPredictor # Use CPU by default (works everywhere), GPU if available device = 'cuda' if torch.cuda.is_available() else 'cpu' _sam_model = sam_model_registry[model_type](checkpoint=str(model_path)) _sam_model.to(device) _sam_predictor = SamPredictor(_sam_model) print(f"SAM model loaded on {device}") return _sam_predictor def _sam_select_local(img_array: np.ndarray, x: int, y: int) -> np.ndarray: """Use local SAM model for selection (no API calls, runs offline)""" predictor = _get_sam_model() # Set image predictor.set_image(img_array) # Point coordinates (x, y) and label (1 = foreground) input_point = np.array([[x, y]]) input_label = np.array([1]) # Get mask prediction masks, scores, _ = predictor.predict( point_coords=input_point, point_labels=input_label, multimask_output=True, # Get multiple mask options ) # Use the mask with highest score best_mask_idx = np.argmax(scores) mask = masks[best_mask_idx] return mask.astype(np.uint8) async def _sam_select(img_array: np.ndarray, x: int, y: int) -> np.ndarray: """ Smart object selection using SAM (Segment Anything Model). Priority: 1. Local SAM model (free, fast, offline) 2. Replicate API (if local not available and API key set) 3. Raises exception if neither available """ # Try local SAM first (free, no API calls) try: return _sam_select_local(img_array, x, y) except FileNotFoundError as e: print(f"Local SAM not available: {e}") except ImportError as e: print(f"SAM dependencies not installed: {e}") except Exception as e: print(f"Local SAM failed: {e}") # Fall back to Replicate API from app.config import settings if not settings.replicate_api_key: raise ValueError( "SAM model not available. Either:\n" " 1. Download local model: docker exec -it ai-photo-edit-backend python /scripts/download_sam_model.py\n" " 2. Or set REPLICATE_API_KEY in .env for cloud SAM" ) return await _sam_select_replicate(img_array, x, y) async def _sam_select_replicate(img_array: np.ndarray, x: int, y: int) -> np.ndarray: """Fallback: Use SAM via Replicate API (requires API key, costs ~$0.002/call)""" import httpx import base64 import asyncio from app.config import settings # Convert image to base64 img = Image.fromarray(img_array) buffer = BytesIO() img.save(buffer, format='PNG') img_b64 = base64.b64encode(buffer.getvalue()).decode('utf-8') async with httpx.AsyncClient(timeout=120.0) as client: prediction_data = { "version": "meta/sam-2-image:fe97b453d6525baeeb530595c74a3c4f567c1f655ee2a0fee11f76bd1d31e495", "input": { "image": f"data:image/png;base64,{img_b64}", "point_coords": f"{x},{y}", "point_labels": "1", } } headers = { 'Authorization': f'Bearer {settings.replicate_api_key}', 'Content-Type': 'application/json' } response = await client.post( "https://api.replicate.com/v1/predictions", json=prediction_data, headers=headers ) if response.status_code != 201: raise Exception(f"Replicate API error: {response.text}") prediction = response.json() prediction_url = prediction['urls']['get'] # Poll for completion for _ in range(60): await asyncio.sleep(2) status_response = await client.get(prediction_url, headers=headers) status_data = status_response.json() if status_data['status'] == 'succeeded': mask_url = status_data['output'] if isinstance(mask_url, list): mask_url = mask_url[0] mask_response = await client.get(mask_url) mask_img = Image.open(BytesIO(mask_response.content)).convert('L') if mask_img.size != (img_array.shape[1], img_array.shape[0]): mask_img = mask_img.resize( (img_array.shape[1], img_array.shape[0]), Image.Resampling.LANCZOS ) return np.array(mask_img) // 255 elif status_data['status'] == 'failed': raise Exception(f"SAM prediction failed: {status_data.get('error')}") raise Exception("SAM prediction timed out") def _flood_fill_select(img_array: np.ndarray, x: int, y: int, tolerance: int = 32) -> np.ndarray: """Simple flood-fill based selection with color tolerance""" import cv2 h, w = img_array.shape[:2] # Ensure point is within bounds x = max(0, min(x, w - 1)) y = max(0, min(y, h - 1)) # Create mask for flood fill (needs to be 2 pixels larger) mask = np.zeros((h + 2, w + 2), np.uint8) # Flood fill cv2.floodFill( img_array.copy(), mask, (x, y), (255, 255, 255), (tolerance, tolerance, tolerance), (tolerance, tolerance, tolerance), cv2.FLOODFILL_MASK_ONLY ) # Extract the actual mask (remove padding) return mask[1:-1, 1:-1] @router.post("/color-select") async def color_select( project_id: int = Form(...), color_r: int = Form(...), color_g: int = Form(...), color_b: int = Form(...), tolerance: int = Form(30), return_format: str = Form("json"), # "json" (default) or "image" db: Session = Depends(get_db) ): """ Select all pixels similar to the given color. Returns a mask and polygon data for selected areas. """ project = db.query(Project).filter(Project.id == project_id).first() if not project: raise HTTPException(status_code=404, detail="Project not found") from app.services.edit_service import EditService edit_service = EditService() image_path = edit_service.get_current_image_path(project_id) img = Image.open(image_path).convert('RGB') img_array = np.array(img) # Target color target = np.array([color_r, color_g, color_b]) # Calculate color distance diff = np.abs(img_array.astype(np.int16) - target.astype(np.int16)) distance = np.sum(diff, axis=2) # Create mask where distance is within tolerance mask = (distance <= tolerance * 3).astype(np.uint8) # Convert to PNG mask_img = Image.fromarray(mask * 255, mode='L') if return_format == "image": buffer = BytesIO() mask_img.save(buffer, format='PNG') return Response( content=buffer.getvalue(), media_type="image/png" ) # Return JSON with polygon and bbox polygon, bbox = _mask_to_polygon(mask) buffer = BytesIO() mask_img.save(buffer, format='PNG') mask_b64 = base64.b64encode(buffer.getvalue()).decode('utf-8') return { "polygon": polygon, "bbox": bbox, "mask_base64": mask_b64, "color": {"r": color_r, "g": color_g, "b": color_b}, "tolerance": tolerance, } @router.post("/extract-object") async def extract_object( project_id: int = Form(...), mask: UploadFile = File(...), db: Session = Depends(get_db) ): """ Extract object using provided mask. Returns PNG with transparent background containing only the masked area. """ project = db.query(Project).filter(Project.id == project_id).first() if not project: raise HTTPException(status_code=404, detail="Project not found") from app.services.edit_service import EditService edit_service = EditService() image_path = edit_service.get_current_image_path(project_id) # Load image and mask img = Image.open(image_path).convert('RGBA') mask_bytes = await mask.read() mask_img = Image.open(BytesIO(mask_bytes)).convert('L') # Resize mask if needed if mask_img.size != img.size: mask_img = mask_img.resize(img.size, Image.Resampling.LANCZOS) # Apply mask as alpha channel img_array = np.array(img) mask_array = np.array(mask_img) # Set alpha channel based on mask img_array[:, :, 3] = mask_array result = Image.fromarray(img_array, mode='RGBA') buffer = BytesIO() result.save(buffer, format='PNG') return Response( content=buffer.getvalue(), media_type="image/png" ) @router.get("/layers/{project_id}") async def list_layers( project_id: int, db: Session = Depends(get_db) ): """List all layers for a project""" project = db.query(Project).filter(Project.id == project_id).first() if not project: raise HTTPException(status_code=404, detail="Project not found") from app.services.edit_service import EditService from pathlib import Path edit_service = EditService() project_dir = edit_service.get_project_dir(project_id) layers_dir = project_dir / 'layers' if not layers_dir.exists(): return {"layers": []} layers = [] for layer_file in sorted(layers_dir.glob('layer_*.png')): img = Image.open(layer_file) layer_num = int(layer_file.stem.split('_')[1]) layers.append({ "id": layer_num, "name": f"Layer {layer_num}", "path": str(layer_file), "width": img.width, "height": img.height }) return {"layers": layers} @router.post("/flatten-layers") async def flatten_layers( project_id: int = Form(...), layer_order: str = Form(...), # JSON array of layer IDs in order db: Session = Depends(get_db) ): """ Flatten all layers into a single image and save as current. layer_order is a JSON array like [1, 2, 3] from bottom to top. """ project = db.query(Project).filter(Project.id == project_id).first() if not project: raise HTTPException(status_code=404, detail="Project not found") from app.services.edit_service import EditService from pathlib import Path edit_service = EditService() project_dir = edit_service.get_project_dir(project_id) layers_dir = project_dir / 'layers' order = json.loads(layer_order) # Start with original image as base base_path = edit_service.get_current_image_path(project_id) result = Image.open(base_path).convert('RGBA') # Composite layers in order for layer_id in order: layer_path = layers_dir / f'layer_{layer_id}.png' if layer_path.exists(): layer = Image.open(layer_path).convert('RGBA') # Resize if needed if layer.size != result.size: layer = layer.resize(result.size, Image.Resampling.LANCZOS) result = Image.alpha_composite(result, layer) # Save as current result.save(base_path, 'PNG') return StatusResponse( status="success", message="Layers flattened successfully" )