diff --git a/.env.example b/.env.example index 5da8bc8..eb1b3ab 100644 --- a/.env.example +++ b/.env.example @@ -116,7 +116,12 @@ SECRET_KEY=change-this-to-a-long-random-string-in-production CORS_ORIGINS=http://localhost:5173,http://localhost:3000,http://localhost:3080,http://localhost # Database path -DATABASE_URL=sqlite:///./data/photoedit.db +DATABASE_URL=sqlite:///./data/ai_photo_edit.db + +# Auto-download SAM model on startup (true/false) +# When true (default): Downloads SAM model (~375MB) on first startup for offline Smart Select +# When false: Skips download, Smart Select uses Replicate API (requires REPLICATE_API_KEY) +AUTO_DOWNLOAD_SAM=true # Allow users to select model per-edit ALLOW_MODEL_OVERRIDE=true diff --git a/backend/app/routers/tools.py b/backend/app/routers/tools.py index 062a8ae..86a81d5 100644 --- a/backend/app/routers/tools.py +++ b/backend/app/routers/tools.py @@ -6,6 +6,8 @@ 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 @@ -133,11 +135,12 @@ 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 for the selected object. + 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. @@ -163,13 +166,62 @@ async def smart_select( # 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 Response( - content=buffer.getvalue(), - media_type="image/png" - ) + 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) @@ -387,11 +439,12 @@ async def color_select( 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 of selected areas. + Returns a mask and polygon data for selected areas. """ project = db.query(Project).filter(Project.id == project_id).first() if not project: @@ -412,18 +465,33 @@ async def color_select( distance = np.sum(diff, axis=2) # Create mask where distance is within tolerance - mask = (distance <= tolerance * 3).astype(np.uint8) * 255 + mask = (distance <= tolerance * 3).astype(np.uint8) # Convert to PNG - mask_img = Image.fromarray(mask, mode='L') + 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 Response( - content=buffer.getvalue(), - media_type="image/png" - ) + 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") diff --git a/backend/entrypoint.sh b/backend/entrypoint.sh index 362ca44..b44ca4d 100644 --- a/backend/entrypoint.sh +++ b/backend/entrypoint.sh @@ -3,9 +3,10 @@ # AI Photo Edit - Container Startup Script # ============================================================================= # This script runs when the container starts. It: -# 1. Downloads sample eye images if the catalog is empty -# 2. Ensures all directories exist -# 3. Starts the FastAPI server +# 1. Initializes the database +# 2. Downloads SAM model automatically (can be disabled with AUTO_DOWNLOAD_SAM=false) +# 3. Downloads sample eye images if the catalog is empty +# 4. Starts the FastAPI server # ============================================================================= set -e @@ -18,18 +19,15 @@ echo "==========================================" mkdir -p /app/data/projects mkdir -p /app/data/patches mkdir -p /app/data/models +mkdir -p /app/data/patch_library -# Check if eye catalog needs to be populated -echo "Checking eye catalog..." -PATCHES_COUNT=$(find /app/data/patches -maxdepth 1 -type d | wc -l) - -if [ "$PATCHES_COUNT" -le 1 ]; then - echo "Eye catalog is empty. Downloading sample eyes..." - python /scripts/download_sample_eyes.py || echo "Warning: Could not download sample eyes (non-fatal)" -else - echo "Eye catalog has content, skipping download." -fi +# Initialize database FIRST (before eye import) +echo "" +echo "Initializing database..." +echo "------------------------------------------" +cd /app && python /scripts/init_database.py || echo "Warning: Database init failed (non-fatal)" +# Check and download SAM model automatically echo "" echo "Checking SAM model (Smart Select)..." echo "------------------------------------------" @@ -39,17 +37,42 @@ if [ -f "/app/data/models/sam_model.pth" ] || \ [ -f "/app/data/models/sam_vit_h_4b8939.pth" ]; then echo "✓ SAM model found - Smart Select will use local AI (free, offline)" else - echo "" - echo "⚠ SAM model not found" - echo "" - echo " Smart Select will use Replicate API (requires REPLICATE_API_KEY)" - echo "" - echo " To enable FREE offline Smart Select, run:" - echo " docker exec -it ai-photo-edit-backend python /scripts/download_sam_model.py" - echo "" - echo " Model sizes: vit_b (375MB), vit_l (1.2GB), vit_h (2.5GB)" - echo " The model persists across container rebuilds." - echo "" + # Auto-download SAM unless explicitly disabled + AUTO_DOWNLOAD_SAM="${AUTO_DOWNLOAD_SAM:-true}" + if [ "$AUTO_DOWNLOAD_SAM" = "true" ]; then + echo "SAM model not found. Downloading automatically..." + echo "(This is a one-time ~375MB download that persists across rebuilds)" + echo "" + python /scripts/download_sam_model.py vit_b || { + echo "" + echo "⚠ SAM download failed (non-fatal)" + echo " Smart Select will fall back to Replicate API (requires REPLICATE_API_KEY)" + echo " To retry later: docker exec -it ai-photo-edit-backend python /scripts/download_sam_model.py" + } + else + echo "" + echo "⚠ SAM model not found (AUTO_DOWNLOAD_SAM=false)" + echo "" + echo " Smart Select will use Replicate API (requires REPLICATE_API_KEY)" + echo "" + echo " To enable FREE offline Smart Select, run:" + echo " docker exec -it ai-photo-edit-backend python /scripts/download_sam_model.py" + echo "" + fi +fi + +# Check if eye catalog needs to be populated +echo "" +echo "Checking eye catalog..." +echo "------------------------------------------" +PATCHES_COUNT=$(find /app/data/patches -maxdepth 1 -type d 2>/dev/null | wc -l) +DB_PATCHES_COUNT=$(sqlite3 /app/data/ai_photo_edit.db "SELECT COUNT(*) FROM patches;" 2>/dev/null || echo "0") + +if [ "$DB_PATCHES_COUNT" = "0" ] || [ "$PATCHES_COUNT" -le 1 ]; then + echo "Eye catalog is empty. Downloading sample eyes..." + cd /app && python /scripts/download_sample_eyes.py || echo "Warning: Could not download sample eyes (non-fatal)" +else + echo "✓ Eye catalog has $DB_PATCHES_COUNT patches" fi echo "" diff --git a/frontend/src/components/AdvancedTools.jsx b/frontend/src/components/AdvancedTools.jsx index 76404b3..2f337c6 100644 --- a/frontend/src/components/AdvancedTools.jsx +++ b/frontend/src/components/AdvancedTools.jsx @@ -11,8 +11,9 @@ const AdvancedTools = ({ isProcessing, setIsProcessing, setError, + activeToolMode, + setActiveToolMode, }) => { - const [activeToolMode, setActiveToolMode] = useState(null); const [colorTolerance, setColorTolerance] = useState(30); const handleRemoveBackground = async () => { diff --git a/frontend/src/components/ImageCanvas.css b/frontend/src/components/ImageCanvas.css index 252fecc..f088433 100644 --- a/frontend/src/components/ImageCanvas.css +++ b/frontend/src/components/ImageCanvas.css @@ -46,3 +46,66 @@ .clear-selection-btn:hover { background-color: #dd4444; } + +/* Zoom controls */ +.zoom-controls { + position: absolute; + bottom: 10px; + right: 10px; + display: flex; + align-items: center; + gap: 4px; + background-color: rgba(0, 0, 0, 0.8); + padding: 6px 10px; + border-radius: 4px; + z-index: 10; +} + +.zoom-controls button { + width: 28px; + height: 28px; + padding: 0; + font-size: 18px; + font-weight: bold; + background-color: #444; + color: white; + border: 1px solid #666; + border-radius: 4px; + cursor: pointer; + display: flex; + align-items: center; + justify-content: center; +} + +.zoom-controls button:hover { + background-color: #555; +} + +.zoom-level { + color: white; + font-size: 12px; + min-width: 45px; + text-align: center; +} + +/* Tool mode indicator */ +.tool-mode-indicator { + position: absolute; + top: 10px; + left: 50%; + transform: translateX(-50%); + background-color: rgba(0, 120, 255, 0.9); + color: white; + padding: 10px 20px; + border-radius: 4px; + font-size: 14px; + font-weight: 500; + z-index: 10; + white-space: nowrap; + animation: pulse 2s infinite; +} + +@keyframes pulse { + 0%, 100% { opacity: 1; } + 50% { opacity: 0.7; } +} diff --git a/frontend/src/components/ImageCanvas.jsx b/frontend/src/components/ImageCanvas.jsx index 1dbda67..80969b9 100644 --- a/frontend/src/components/ImageCanvas.jsx +++ b/frontend/src/components/ImageCanvas.jsx @@ -61,11 +61,36 @@ const ImageCanvas = forwardRef(({ handleResize(); window.addEventListener('resize', handleResize); + // Mouse wheel zoom + const handleWheel = (opt) => { + const e = opt.e; + e.preventDefault(); + e.stopPropagation(); + + const delta = e.deltaY; + let newZoom = canvas.getZoom(); + newZoom *= 0.999 ** delta; + + // Clamp zoom between 0.1x and 10x + if (newZoom > 10) newZoom = 10; + if (newZoom < 0.1) newZoom = 0.1; + + // Zoom to point under cursor + const pointer = canvas.getPointer(e, true); + canvas.zoomToPoint({ x: pointer.x, y: pointer.y }, newZoom); + + setCurrentZoom(newZoom); + onZoomChange?.(newZoom); + }; + + canvas.on('mouse:wheel', handleWheel); + return () => { window.removeEventListener('resize', handleResize); + canvas.off('mouse:wheel', handleWheel); canvas.dispose(); }; - }, []); + }, [onZoomChange]); // Center and scale image const centerImage = (canvas, img, zoomFactor) => { @@ -601,6 +626,35 @@ const ImageCanvas = forwardRef(({ } }; + const handleZoomIn = () => { + if (!fabricCanvasRef.current) return; + const canvas = fabricCanvasRef.current; + let newZoom = canvas.getZoom() * 1.2; + if (newZoom > 10) newZoom = 10; + canvas.setZoom(newZoom); + setCurrentZoom(newZoom); + onZoomChange?.(newZoom); + }; + + const handleZoomOut = () => { + if (!fabricCanvasRef.current) return; + const canvas = fabricCanvasRef.current; + let newZoom = canvas.getZoom() / 1.2; + if (newZoom < 0.1) newZoom = 0.1; + canvas.setZoom(newZoom); + setCurrentZoom(newZoom); + onZoomChange?.(newZoom); + }; + + const handleZoomReset = () => { + if (!fabricCanvasRef.current) return; + const canvas = fabricCanvasRef.current; + canvas.setZoom(1); + canvas.setViewportTransform([1, 0, 0, 1, 0, 0]); + setCurrentZoom(1); + onZoomChange?.(1); + }; + return (