diff --git a/backend/Dockerfile b/backend/Dockerfile index cd7fe09..6f85f98 100644 --- a/backend/Dockerfile +++ b/backend/Dockerfile @@ -2,10 +2,15 @@ FROM python:3.11-slim WORKDIR /app -# Install system dependencies +# Install system dependencies for OpenCV, rembg, and image processing RUN apt-get update && apt-get install -y \ libgl1 \ libglib2.0-0 \ + libsm6 \ + libxext6 \ + libxrender-dev \ + libgomp1 \ + wget \ && rm -rf /var/lib/apt/lists/* # Copy requirements @@ -14,11 +19,14 @@ COPY requirements.txt . # Install Python dependencies RUN pip install --no-cache-dir -r requirements.txt +# Pre-download rembg model (u2net) to avoid first-run delay +RUN python -c "from rembg import remove; print('rembg model downloaded')" || true + # Copy application COPY . . -# Create data directory -RUN mkdir -p /app/data/projects +# Create data directories +RUN mkdir -p /app/data/projects /app/data/patches /app/data/models # Expose port EXPOSE 8000 diff --git a/backend/app/main.py b/backend/app/main.py index f408e8f..f6a95fd 100644 --- a/backend/app/main.py +++ b/backend/app/main.py @@ -5,7 +5,7 @@ from contextlib import asynccontextmanager from app.config import settings from app.database import init_db -from app.routers import projects, edits, images, patches, generate +from app.routers import projects, edits, images, patches, generate, tools @asynccontextmanager @@ -37,6 +37,7 @@ app.include_router(edits.router) app.include_router(images.router) app.include_router(patches.router) app.include_router(generate.router) +app.include_router(tools.router) @app.get("/") diff --git a/backend/app/routers/tools.py b/backend/app/routers/tools.py new file mode 100644 index 0000000..6145390 --- /dev/null +++ b/backend/app/routers/tools.py @@ -0,0 +1,454 @@ +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 + +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(...), + db: Session = Depends(get_db) +): + """ + Use SAM (Segment Anything) to select object at given point. + Returns mask 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') + + buffer = BytesIO() + mask_img.save(buffer, format='PNG') + + return Response( + content=buffer.getvalue(), + media_type="image/png" + ) + + +async def _sam_select(img_array: np.ndarray, x: int, y: int) -> np.ndarray: + """Use Segment Anything Model for selection via Replicate API""" + import httpx + import base64 + import asyncio + from app.config import settings + + if not settings.replicate_api_key: + raise ValueError("REPLICATE_API_KEY not configured") + + # 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') + + # Call SAM via Replicate + async with httpx.AsyncClient(timeout=120.0) as client: + # Use SAM model on Replicate + 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", # 1 = foreground point + } + } + + headers = { + 'Authorization': f'Bearer {settings.replicate_api_key}', + 'Content-Type': 'application/json' + } + + # Start prediction + 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() + + # Poll for completion + prediction_url = prediction['urls']['get'] + max_attempts = 60 + attempt = 0 + + while attempt < max_attempts: + await asyncio.sleep(2) + + status_response = await client.get(prediction_url, headers=headers) + status_data = status_response.json() + + if status_data['status'] == 'succeeded': + # Download mask image + 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') + + # Resize if needed + 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 # Normalize to 0-1 + + elif status_data['status'] == 'failed': + raise Exception(f"SAM prediction failed: {status_data.get('error')}") + + attempt += 1 + + 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), + db: Session = Depends(get_db) +): + """ + Select all pixels similar to the given color. + Returns a mask of 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) * 255 + + # Convert to PNG + mask_img = Image.fromarray(mask, mode='L') + + buffer = BytesIO() + mask_img.save(buffer, format='PNG') + + return Response( + content=buffer.getvalue(), + media_type="image/png" + ) + + +@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" + ) diff --git a/backend/app/services/edit_service.py b/backend/app/services/edit_service.py index 93965fb..4007b9d 100644 --- a/backend/app/services/edit_service.py +++ b/backend/app/services/edit_service.py @@ -188,9 +188,11 @@ class EditService: if not result_path.exists(): raise FileNotFoundError(f"Edit {edit_id} result not found") - # Copy result to current + # Copy result to current (preserve alpha channel) current_path = self.get_current_image_path(project_id) - Image.open(result_path).save(current_path) + img = Image.open(result_path) + # Preserve original mode to maintain transparency + img.save(current_path, format='PNG') return str(current_path) @@ -210,7 +212,9 @@ class EditService: if not original_path.exists(): raise FileNotFoundError(f"Original image for project {project_id} not found") - # Copy original to current - Image.open(original_path).save(current_path) + # Copy original to current (preserve alpha channel) + img = Image.open(original_path) + # Preserve original mode to maintain transparency + img.save(current_path, format='PNG') return str(current_path) diff --git a/backend/app/utils/image_processing.py b/backend/app/utils/image_processing.py index 4f6ae25..978d2e6 100644 --- a/backend/app/utils/image_processing.py +++ b/backend/app/utils/image_processing.py @@ -55,19 +55,22 @@ def blend_patch( original_patch: Image.Image, regenerated_patch: Image.Image, mask: Image.Image, - feather_px: int = 0 + feather_px: int = 0, + preserve_alpha: bool = True ) -> Image.Image: """ - Blend regenerated patch with original using mask + Blend regenerated patch with original using mask. + Preserves original alpha channel for semi-transparent areas (veils, glass, etc). Args: original_patch: Original cropped patch regenerated_patch: AI-regenerated patch mask: Binary mask (same size as patches) feather_px: Feather radius for smooth blending + preserve_alpha: If True, preserves original alpha channel Returns: - Blended patch + Blended patch with preserved transparency """ # Ensure all images are the same size if regenerated_patch.size != original_patch.size: @@ -83,12 +86,20 @@ def blend_patch( # Apply feathering to mask feathered_mask = create_feathered_mask(mask, feather_px) - # Convert images to RGBA - original_patch = original_patch.convert('RGBA') - regenerated_patch = regenerated_patch.convert('RGBA') + # Convert images to RGBA, storing original alpha + original_rgba = original_patch.convert('RGBA') + original_alpha = original_rgba.split()[3] # Store original alpha channel + + regenerated_rgba = regenerated_patch.convert('RGBA') # Blend using the feathered mask - blended = Image.composite(regenerated_patch, original_patch, feathered_mask) + blended = Image.composite(regenerated_rgba, original_rgba, feathered_mask) + + # Restore original alpha channel to preserve transparency + # This keeps semi-transparent areas (veils, glass, smoke) intact + if preserve_alpha: + r, g, b, _ = blended.split() + blended = Image.merge('RGBA', (r, g, b, original_alpha)) return blended diff --git a/backend/requirements.txt b/backend/requirements.txt index 1114661..d1be434 100644 --- a/backend/requirements.txt +++ b/backend/requirements.txt @@ -14,3 +14,5 @@ pydantic-settings==2.1.0 email-validator==2.1.0 opencv-python-headless==4.9.0.80 scikit-image==0.22.0 +rembg==2.0.50 +onnxruntime==1.16.3 diff --git a/frontend/src/App.css b/frontend/src/App.css index bf7c1fc..86634e7 100644 --- a/frontend/src/App.css +++ b/frontend/src/App.css @@ -118,7 +118,10 @@ .right-panel { display: flex; flex-direction: column; - gap: 20px; + gap: 12px; + overflow-y: auto; + max-height: calc(100vh - 180px); + padding-right: 8px; } .history-wrapper { diff --git a/frontend/src/App.jsx b/frontend/src/App.jsx index 19b23ee..a61287e 100644 --- a/frontend/src/App.jsx +++ b/frontend/src/App.jsx @@ -3,7 +3,9 @@ import ImageCanvas from './components/ImageCanvas'; import Controls from './components/Controls'; import History from './components/History'; import EyeCatalog from './components/EyeCatalog'; -import { projectsApi, editsApi } from './utils/api'; +import AdvancedTools from './components/AdvancedTools'; +import Layers from './components/Layers'; +import { projectsApi, editsApi, toolsApi } from './utils/api'; import './App.css'; function App() { @@ -21,6 +23,9 @@ function App() { const [projectName, setProjectName] = useState(''); const [showProjectInput, setShowProjectInput] = useState(true); const [currentEditIndex, setCurrentEditIndex] = useState(-1); + const [layers, setLayers] = useState([]); + const [activeLayer, setActiveLayer] = useState('background'); + const [generatedMask, setGeneratedMask] = useState(null); const editsRef = useRef([]); // Create project and upload image @@ -273,6 +278,32 @@ function App() { return () => window.removeEventListener('beforeunload', handleBeforeUnload); }, [project, edits]); + // Handle layer creation from advanced tools + const handleLayerCreated = (layer) => { + setLayers((prev) => [...prev, { ...layer, visible: true }]); + }; + + // Handle mask generation from smart select / color select + const handleMaskGenerated = async (maskBlob, source) => { + setGeneratedMask({ blob: maskBlob, source }); + // The mask can be used for various operations + }; + + // Handle flatten layers + const handleFlattenLayers = async (layerOrder) => { + if (!project) return; + try { + setIsProcessing(true); + await toolsApi.flattenLayers(project.id, layerOrder); + setCurrentImageUrl(projectsApi.getCurrentImageUrl(project.id)); + setLayers([]); + } catch (err) { + setError(`Failed to flatten layers: ${err.message}`); + } finally { + setIsProcessing(false); + } + }; + return (
+ {activeToolMode === 'smart-select' + ? 'Click on an object to select it' + : 'Click on a color to select all similar pixels'} +
+ )} +