diff --git a/backend/app/main.py b/backend/app/main.py index 70d9b5a..082ec1e 100644 --- a/backend/app/main.py +++ b/backend/app/main.py @@ -21,6 +21,9 @@ async def lifespan(app: FastAPI): caps = probe_upscale_capabilities() if not caps["realesrgan_pytorch"] and not caps["realesrgan_ncnn"]: asyncio.create_task(ensure_ncnn_installed()) + # Pre-download SAM model in background so first click is fast + from app.services.sam_service import ensure_sam_installed + asyncio.create_task(ensure_sam_installed()) yield diff --git a/backend/app/routers/ai_tools.py b/backend/app/routers/ai_tools.py index 243c5ba..b99be59 100644 --- a/backend/app/routers/ai_tools.py +++ b/backend/app/routers/ai_tools.py @@ -348,3 +348,63 @@ async def get_config(): }, } } + + +# ─── SAM (Segment Anything) ────────────────────────────────────────────────── + +class SegmentPointRequest(BaseModel): + image: str # base64 PNG/JPEG + points: list[list[int]] # [[x, y], ...] original image coords + labels: list[int] # 1=include, 0=exclude — same length as points + + +@router.post("/segment/point") +async def segment_point(req: SegmentPointRequest): + """ + Run SAM point-prompt segmentation. + Returns a binary mask PNG (white = selected area). + Auto-downloads the SAM ViT-B model (~375 MB) on first call. + """ + if not req.points: + raise HTTPException(status_code=400, detail="At least one point required.") + if len(req.points) != len(req.labels): + raise HTTPException(status_code=400, detail="points and labels must have the same length.") + + try: + image_bytes = base64.b64decode(req.image) + except Exception as e: + raise HTTPException(status_code=400, detail=f"Could not decode image: {e}") + + from app.services.sam_service import predict_points, get_install_status + try: + mask_bytes = await predict_points( + image_bytes, + [tuple(p) for p in req.points], + req.labels, + ) + return { + "mask": base64.b64encode(mask_bytes).decode(), + "sam_install": get_install_status(), + } + except RuntimeError as e: + raise HTTPException(status_code=503, detail=str(e)) + except Exception as e: + import traceback; traceback.print_exc() + raise HTTPException(status_code=500, detail=str(e)) + + +@router.get("/segment/install-status") +def segment_install_status(): + """Poll SAM model download progress.""" + from app.services.sam_service import get_install_status, sam_model_available + status = get_install_status() + status["model_ready"] = sam_model_available() + return status + + +@router.post("/segment/install") +async def segment_install(): + """Trigger SAM model download explicitly (also auto-triggered on first /segment/point call).""" + from app.services.sam_service import ensure_sam_installed, get_install_status + asyncio.create_task(ensure_sam_installed()) + return get_install_status() diff --git a/backend/app/services/sam_service.py b/backend/app/services/sam_service.py new file mode 100644 index 0000000..b10e7ea --- /dev/null +++ b/backend/app/services/sam_service.py @@ -0,0 +1,184 @@ +""" +SAM (Segment Anything Model) service. + +Auto-downloads the ViT-B checkpoint (~375 MB) on first use. +Caches the loaded model in memory; re-uses predictor across calls. + +Prediction API: + predict_points(image_bytes, points, labels) -> mask_bytes (PNG, white=selected) + points: list of (x, y) in original image pixels + labels: list of 1 (include) or 0 (exclude), same length as points +""" + +import asyncio +import io +import os +import urllib.request +from dataclasses import dataclass +from enum import Enum +from pathlib import Path +from typing import Optional + +import numpy as np +from PIL import Image + +# ── Model download ──────────────────────────────────────────────────────────── + +SAM_DIR = Path("/app/data/models/sam") +SAM_FILENAME = "sam_vit_b_01ec64.pth" +SAM_URL = f"https://dl.fbaipublicfiles.com/segment_anything/{SAM_FILENAME}" +SAM_PATH = SAM_DIR / SAM_FILENAME + + +class SamInstallState(str, Enum): + idle = "idle" + downloading = "downloading" + done = "done" + failed = "failed" + + +@dataclass +class SamInstallStatus: + state: SamInstallState = SamInstallState.idle + progress: int = 0 + message: str = "" + error: str = "" + + +_install_status = SamInstallStatus() +_install_lock = asyncio.Lock() + + +def get_install_status() -> dict: + s = _install_status + return {"state": s.state.value, "progress": s.progress, + "message": s.message, "error": s.error} + + +def sam_model_available() -> bool: + return SAM_PATH.exists() and SAM_PATH.stat().st_size > 100_000_000 + + +async def ensure_sam_installed() -> bool: + """Download SAM ViT-B checkpoint if not present. Returns True on success.""" + global _install_status + + if sam_model_available(): + _install_status = SamInstallStatus(state=SamInstallState.done, progress=100, + message="SAM model ready.") + return True + + async with _install_lock: + if sam_model_available(): + _install_status = SamInstallStatus(state=SamInstallState.done, progress=100, + message="SAM model ready.") + return True + + if _install_status.state == SamInstallState.downloading: + return False + + try: + SAM_DIR.mkdir(parents=True, exist_ok=True) + _install_status = SamInstallStatus( + state=SamInstallState.downloading, progress=0, + message="Downloading SAM ViT-B model (~375 MB)…", + ) + + def _download(): + def _progress(count, block, total): + if total > 0: + _install_status.progress = min(99, int(count * block * 99 / total)) + tmp = SAM_PATH.with_suffix(".tmp") + urllib.request.urlretrieve(SAM_URL, tmp, _progress) + tmp.rename(SAM_PATH) + + loop = asyncio.get_event_loop() + await loop.run_in_executor(None, _download) + + _install_status = SamInstallStatus(state=SamInstallState.done, progress=100, + message="SAM model ready.") + return True + + except Exception as exc: + _install_status = SamInstallStatus( + state=SamInstallState.failed, error=str(exc), + message="SAM download failed.", + ) + print(f"[sam] Download failed: {exc}") + return False + + +# ── Model cache ─────────────────────────────────────────────────────────────── + +_predictor = None +_predictor_lock = asyncio.Lock() + + +def _load_predictor(): + """Load SAM model and return a SamPredictor. Called in thread pool.""" + global _predictor + if _predictor is not None: + return _predictor + + import torch + from segment_anything import sam_model_registry, SamPredictor + + if torch.cuda.is_available(): + device = "cuda" + elif hasattr(torch.backends, "mps") and torch.backends.mps.is_available(): + device = "mps" + else: + device = "cpu" + + print(f"[sam] Loading SAM ViT-B on {device}…") + sam = sam_model_registry["vit_b"](checkpoint=str(SAM_PATH)) + sam.to(device) + _predictor = SamPredictor(sam) + print("[sam] Model loaded.") + return _predictor + + +# ── Prediction ──────────────────────────────────────────────────────────────── + +def _predict_sync(image_bytes: bytes, + points: list[tuple[int, int]], + labels: list[int]) -> bytes: + """ + Run SAM prediction synchronously (call via run_in_executor). + Returns PNG bytes: white = selected, black = background. + """ + predictor = _load_predictor() + + image = Image.open(io.BytesIO(image_bytes)).convert("RGB") + img_array = np.array(image) + + predictor.set_image(img_array) + + pt_array = np.array(points, dtype=np.float32) # [[x, y], ...] + lbl_array = np.array(labels, dtype=np.int32) # [1=fg, 0=bg, ...] + + masks, scores, _ = predictor.predict( + point_coords=pt_array, + point_labels=lbl_array, + multimask_output=True, + ) + + # Pick the highest-confidence mask + best = masks[int(np.argmax(scores))] # bool array H×W + + mask_img = Image.fromarray((best * 255).astype(np.uint8), mode="L") + buf = io.BytesIO() + mask_img.save(buf, format="PNG") + return buf.getvalue() + + +async def predict_points(image_bytes: bytes, + points: list[tuple[int, int]], + labels: list[int]) -> bytes: + """Async wrapper for SAM point prediction.""" + if not sam_model_available(): + ok = await ensure_sam_installed() + if not ok: + raise RuntimeError("SAM model not available.") + loop = asyncio.get_event_loop() + return await loop.run_in_executor(None, _predict_sync, image_bytes, points, labels) diff --git a/frontend/src/js/tools/ai_edit.js b/frontend/src/js/tools/ai_edit.js index e0947ff..14105da 100644 --- a/frontend/src/js/tools/ai_edit.js +++ b/frontend/src/js/tools/ai_edit.js @@ -1,16 +1,16 @@ /** - * AI Edit — unified inpainting tool. + * AI Edit — unified smart selection + inpainting tool. * * Workflow: - * 1. Brush over the area you want to change (red overlay) - * 2. Floating action bar appears: Erase | Replace | Upscale | Expand | Clear - * 3. Erase → LaMa-removes masked content - * Replace → inline prompt → AI replaces masked area - * Upscale → opens upscale dialog (whole image) - * Expand → opens outpaint/expand dialog (whole image) - * Clear → wipe the mask and start over + * 1. CLICK mode (default): click any object → SAM auto-selects it (red overlay) + * • Alt+click → subtract from selection (deselect over-selected area) + * • Multiple clicks accumulate on the mask + * 2. BRUSH + / BRUSH − tabs: paint to add or erase from the mask by hand + * (refine what SAM missed or got wrong) + * 3. Action bar: Erase | Replace… | Upscale | Expand | Clear * - * Tool target: tools/ai_edit (auto-registered by webpack require.context) + * SAM model (~375 MB) auto-downloads on first click; progress shown inline. + * Falls back gracefully to brush-only if SAM is unavailable. */ import app from './../app.js'; @@ -20,29 +20,32 @@ import alertify from './../../../node_modules/alertifyjs/build/alertify.min.js'; var instance = null; -// ── constants ────────────────────────────────────────────────────────────── -const BRUSH_COLOR = 'rgba(255, 60, 60, 0.55)'; const BRUSH_DEFAULT = 30; +const OVERLAY_COLOR = 'rgba(255, 55, 55, 0.50)'; +const ERASE_COLOR = 'rgba(0, 0, 0, 0.70)'; // brush-erase preview class Tools_ai_edit_class { constructor() { if (instance) return instance; instance = this; - this.Base_layers = new Base_layers_class(); - this.name = 'ai_edit'; - this.title = 'AI Edit'; - // brush state - this._painting = false; - this._maskCanvas = null; // same size as layer original - this._maskCtx = null; - this._overlayEl = null; // red overlay on top of main canvas - this._panel = null; // floating action bar DOM element - this._hasMask = false; - this._isRunning = false; + this.Base_layers = new Base_layers_class(); + this.name = 'ai_edit'; + this.title = 'AI Edit'; + // interaction state + this._mode = 'sam'; // 'sam' | 'brush_add' | 'brush_sub' + this._painting = false; + this._samWorking = false; + this._isRunning = false; + this._hasMask = false; + // DOM elements + this._maskCanvas = null; + this._maskCtx = null; + this._overlayEl = null; + this._panel = null; } - // ── Tool lifecycle ─────────────────────────────────────────────────────── + // ── Tool lifecycle ──────────────────────────────────────────────────────── on_activate() { if (!config.layer || config.layer.type !== 'image') { @@ -60,72 +63,222 @@ class Tools_ai_edit_class { this._painting = false; } - // ── Mouse / touch ──────────────────────────────────────────────────────── + // ── Input routing ───────────────────────────────────────────────────────── mousedown(e) { if (!config.layer || config.layer.type !== 'image') return; - this._painting = true; - this._paint(e); + if (this._mode === 'sam') { + this._handleSamClick(e); + } else { + this._painting = true; + this._brushPaint(e); + } } mousemove(e) { - if (!this._painting) return; - this._paint(e); + if (this._mode !== 'sam' && this._painting) this._brushPaint(e); } mouseup() { - this._painting = false; - if (this._hasMask) this._showPanel(); + if (this._painting) { + this._painting = false; + if (this._hasMask) this._showActions(); + } } - // ── Mask painting ──────────────────────────────────────────────────────── + // ── Coordinate mapping ──────────────────────────────────────────────────── - _initMask() { - const w = config.layer.width_original; - const h = config.layer.height_original; - this._maskCanvas = document.createElement('canvas'); - this._maskCanvas.width = w; - this._maskCanvas.height = h; - this._maskCtx = this._maskCanvas.getContext('2d'); - this._hasMask = false; - } - - _paint(e) { - if (!this._maskCtx || !this._overlayEl) return; - - // Map screen coords → original image coords + _screenToImage(e) { const canvasEl = document.getElementById('canvas_minipaint') || document.querySelector('canvas'); - if (!canvasEl) return; + if (!canvasEl) return null; const rect = canvasEl.getBoundingClientRect(); const scaleX = config.layer.width_original / (config.WIDTH * config.ZOOM); const scaleY = config.layer.height_original / (config.HEIGHT * config.ZOOM); - const x = ((e.clientX - rect.left) - config.layer.x * config.ZOOM) * scaleX; - const y = ((e.clientY - rect.top) - config.layer.y * config.ZOOM) * scaleY; + const ix = ((e.clientX - rect.left) - config.layer.x * config.ZOOM) * scaleX; + const iy = ((e.clientY - rect.top) - config.layer.y * config.ZOOM) * scaleY; + return { ix, iy, scaleX, scaleY }; + } + // ── SAM click selection ─────────────────────────────────────────────────── + + async _handleSamClick(e) { + if (this._samWorking) return; + const coords = this._screenToImage(e); + if (!coords) return; + + const label = e.altKey ? 0 : 1; // alt = exclude, normal = include + const x = Math.round(coords.ix); + const y = Math.round(coords.iy); + + // Clamp to image bounds + const w = config.layer.width_original; + const h = config.layer.height_original; + if (x < 0 || y < 0 || x >= w || y >= h) return; + + this._samWorking = true; + this._setSamCursor('wait'); + + // Collect any existing points for multi-click accumulation + if (!this._samPoints) this._samPoints = []; + if (!this._samLabels) this._samLabels = []; + this._samPoints.push([x, y]); + this._samLabels.push(label); + + try { + const imageB64 = this._getLayerB64(); + const base = window.API_BASE_URL || ''; + const r = await fetch(`${base}/api/segment/point`, { + method: 'POST', + headers: { 'Content-Type': 'application/json' }, + body: JSON.stringify({ + image: imageB64, + points: this._samPoints, + labels: this._samLabels, + }), + }); + + if (r.status === 503) { + // SAM model downloading — poll and retry + const data = await r.json().catch(() => ({})); + await this._waitForSamModel(data.detail || ''); + // Remove the point we just added so user can retry cleanly + this._samPoints.pop(); + this._samLabels.pop(); + this._samWorking = false; + this._setSamCursor('crosshair'); + return; + } + + if (!r.ok) { + const err = await r.json().catch(() => ({})); + throw new Error(err.detail || 'SAM failed'); + } + + const data = await r.json(); + await this._applySamMask(data.mask, label === 0); + this._hasMask = true; + this._showActions(); + + } catch (err) { + alertify.error('SAM failed: ' + (err.message || err)); + // Pop failed point + this._samPoints.pop(); + this._samLabels.pop(); + } + + this._samWorking = false; + this._setSamCursor('crosshair'); + } + + async _applySamMask(maskB64, isSubtract) { + return new Promise((resolve) => { + const img = new Image(); + img.onload = () => { + // Draw SAM mask onto our persistent mask canvas + const tmp = document.createElement('canvas'); + tmp.width = this._maskCanvas.width; + tmp.height = this._maskCanvas.height; + const tctx = tmp.getContext('2d'); + tctx.drawImage(img, 0, 0, tmp.width, tmp.height); + + if (isSubtract) { + // Erase mask where SAM says to subtract + this._maskCtx.globalCompositeOperation = 'destination-out'; + this._maskCtx.drawImage(tmp, 0, 0); + this._maskCtx.globalCompositeOperation = 'source-over'; + } else { + this._maskCtx.drawImage(tmp, 0, 0); + } + this._redrawOverlay(); + resolve(); + }; + img.src = 'data:image/png;base64,' + maskB64; + }); + } + + async _waitForSamModel(detail) { + // SAM model is downloading — show progress bar and poll + return new Promise((resolve) => { + alertify.message( + `
Downloading SAM model (~375 MB)…
+ + 0%
+ This happens once — click the object again when done. +
`, 0 + ); + const poll = setInterval(async () => { + try { + const base = window.API_BASE_URL || ''; + const r = await fetch(`${base}/api/segment/install-status`); + if (!r.ok) return; + const s = await r.json(); + const bar = document.getElementById('sam-dl-progress'); + const pct = document.getElementById('sam-dl-pct'); + if (bar) bar.value = s.progress || 0; + if (pct) pct.textContent = `${s.progress || 0}%`; + if (s.state === 'done' || s.model_ready) { + clearInterval(poll); + alertify.dismissAll(); + alertify.success('SAM model ready — click the object now.'); + resolve(); + } else if (s.state === 'failed') { + clearInterval(poll); + alertify.dismissAll(); + alertify.error('SAM model download failed. Use Brush mode instead.'); + resolve(); + } + } catch { /* keep polling */ } + }, 1500); + }); + } + + _setSamCursor(cursor) { + const canvasEl = document.getElementById('canvas_minipaint') || document.querySelector('canvas'); + if (canvasEl) canvasEl.style.cursor = cursor; + } + + // ── Brush painting ──────────────────────────────────────────────────────── + + _brushPaint(e) { + const coords = this._screenToImage(e); + if (!coords) return; + const { ix, iy, scaleX, scaleY } = coords; const r = (config.tools[this.name]?.size ?? BRUSH_DEFAULT) / 2; - // Draw on mask (white = area to process) + // Paint on mask canvas + this._maskCtx.globalCompositeOperation = + this._mode === 'brush_sub' ? 'destination-out' : 'source-over'; this._maskCtx.fillStyle = '#ffffff'; this._maskCtx.beginPath(); - this._maskCtx.arc(x, y, r, 0, Math.PI * 2); + this._maskCtx.arc(ix, iy, r, 0, Math.PI * 2); this._maskCtx.fill(); + this._maskCtx.globalCompositeOperation = 'source-over'; - // Mirror onto overlay canvas (red tint for user feedback) + // Mirror on overlay const oc = this._overlayEl; + if (!oc) return; const oct = oc.getContext('2d'); - oct.fillStyle = BRUSH_COLOR; - // Map back: overlay is sized to the visible canvas area - const ox = (x / scaleX) + config.layer.x * config.ZOOM; - const oy = (y / scaleY) + config.layer.y * config.ZOOM; + const ox = (ix / scaleX) + config.layer.x * config.ZOOM; + const oy = (iy / scaleY) + config.layer.y * config.ZOOM; const or_ = r / scaleX; + + if (this._mode === 'brush_sub') { + oct.globalCompositeOperation = 'destination-out'; + oct.fillStyle = '#000'; + } else { + oct.globalCompositeOperation = 'source-over'; + oct.fillStyle = OVERLAY_COLOR; + } oct.beginPath(); oct.arc(ox, oy, or_, 0, Math.PI * 2); oct.fill(); + oct.globalCompositeOperation = 'source-over'; this._hasMask = true; } - // ── Overlay canvas (red mask feedback) ─────────────────────────────────── + // ── Overlay ─────────────────────────────────────────────────────────────── _mountOverlay() { this._removeOverlay(); @@ -136,46 +289,52 @@ class Tools_ai_edit_class { oc.width = base.offsetWidth; oc.height = base.offsetHeight; Object.assign(oc.style, { - position: 'absolute', - top: base.offsetTop + 'px', - left: base.offsetLeft + 'px', - pointerEvents: 'none', - zIndex: '50', + position: 'absolute', top: base.offsetTop + 'px', left: base.offsetLeft + 'px', + pointerEvents: 'none', zIndex: '50', }); base.parentElement.appendChild(oc); this._overlayEl = oc; } + _redrawOverlay() { + if (!this._overlayEl || !this._maskCanvas) return; + const oc = this._overlayEl; + const oct = oc.getContext('2d'); + oct.clearRect(0, 0, oc.width, oc.height); + + // Scale mask to overlay size and tint red + const tmp = document.createElement('canvas'); + tmp.width = oc.width; + tmp.height = oc.height; + const tctx = tmp.getContext('2d'); + tctx.drawImage(this._maskCanvas, 0, 0, oc.width, oc.height); + + // Multiply white mask pixels → red tint using composite + oct.globalCompositeOperation = 'source-over'; + oct.fillStyle = OVERLAY_COLOR; + oct.fillRect(0, 0, oc.width, oc.height); + oct.globalCompositeOperation = 'destination-in'; + oct.drawImage(tmp, 0, 0); + oct.globalCompositeOperation = 'source-over'; + } + _removeOverlay() { if (this._overlayEl) { this._overlayEl.remove(); this._overlayEl = null; } } - // ── Floating action panel ───────────────────────────────────────────────── + // ── Panel ───────────────────────────────────────────────────────────────── _mountPanel() { this._removePanel(); const panel = document.createElement('div'); panel.id = 'ai_edit_panel'; Object.assign(panel.style, { - position: 'fixed', - bottom: '80px', - left: '50%', - transform: 'translateX(-50%)', - background: '#1e1e1e', - border: '1px solid #444', - borderRadius: '10px', - padding: '10px 14px', - display: 'flex', - alignItems: 'center', - gap: '8px', - zIndex: '9999', - boxShadow: '0 4px 20px rgba(0,0,0,0.5)', - fontFamily: 'sans-serif', - fontSize: '13px', - color: '#eee', - userSelect: 'none', - flexWrap: 'wrap', - maxWidth: '600px', + position: 'fixed', bottom: '72px', left: '50%', transform: 'translateX(-50%)', + background: '#1a1a1a', border: '1px solid #3a3a3a', borderRadius: '12px', + padding: '10px 14px', display: 'flex', flexDirection: 'column', + gap: '8px', zIndex: '9999', boxShadow: '0 6px 24px rgba(0,0,0,0.6)', + fontFamily: 'sans-serif', fontSize: '13px', color: '#eee', + userSelect: 'none', minWidth: '460px', }); panel.innerHTML = this._panelHTML(); document.body.appendChild(panel); @@ -185,55 +344,89 @@ class Tools_ai_edit_class { _panelHTML() { return ` - Paint mask, then: - -
- - - -
- - - `; + .aie-btn:hover { background:#333; } + .aie-btn.active { background:#1e3a5f;border-color:#3b82f6;color:#93c5fd; } + .aie-btn--go { background:#2563eb;border-color:#3b82f6;color:#fff; } + .aie-btn--go:hover { background:#1d4ed8; } + .aie-btn--danger { border-color:#5a1a1a;color:#f87171; } + .aie-btn--danger:hover { background:#2a1010; } + .aie-divider { width:1px;background:#3a3a3a;align-self:stretch; } + + + +
+ Select: + + + +
+ Click an object to select it. Alt+click to deselect. +
+ + +
+ Then: + +
+ + + +
+ + + +
`; } - _showPanel() { - if (this._panel) this._panel.style.opacity = '1'; + _showActions() { + // No-op — actions are always visible; just a hook for future animation } _wirePanel() { if (!this._panel) return; const _this = this; + const hints = { + sam: 'Click an object to select it. Alt+click to deselect an area.', + brush_add: 'Paint over areas to add them to the selection.', + brush_sub: 'Paint over areas to remove them from the selection.', + }; - // Erase / Upscale / Expand / Clear buttons - this._panel.querySelectorAll('[data-ai-action]').forEach(btn => { + // Mode buttons + this._panel.querySelectorAll('[data-mode]').forEach(btn => { btn.addEventListener('click', () => { - const action = btn.dataset.aiAction; - if (action === 'erase') _this._doErase(); - if (action === 'upscale') _this._doUpscale(); - if (action === 'expand') _this._doExpand(); - if (action === 'clear') _this._doClear(); - if (action === 'replace') _this._toggleReplaceInput(); + _this._mode = btn.dataset.mode; + _this._panel.querySelectorAll('[data-mode]').forEach(b => + b.classList.toggle('active', b === btn)); + const hint = _this._panel.querySelector('#aie-hint'); + if (hint) hint.textContent = hints[_this._mode] || ''; + _this._setSamCursor(_this._mode === 'sam' ? 'crosshair' : 'cell'); }); }); - // Replace → Go - const goBtn = this._panel.querySelector('#ai_replace_go'); - const promptEl = this._panel.querySelector('#ai_replace_prompt'); + // Action buttons + this._panel.querySelectorAll('[data-action]').forEach(btn => { + btn.addEventListener('click', () => { + const a = btn.dataset.action; + if (a === 'erase') _this._doErase(); + if (a === 'replace') _this._toggleReplace(); + if (a === 'upscale') _this._doUpscale(); + if (a === 'expand') _this._doExpand(); + if (a === 'clear') _this._doClear(); + }); + }); + + // Replace prompt + const goBtn = this._panel.querySelector('#aie-go'); + const promptEl = this._panel.querySelector('#aie-prompt'); if (goBtn && promptEl) { goBtn.addEventListener('click', () => _this._doReplace(promptEl.value.trim())); promptEl.addEventListener('keydown', e => { @@ -242,61 +435,64 @@ class Tools_ai_edit_class { } } - _toggleReplaceInput() { - const promptEl = this._panel && this._panel.querySelector('#ai_replace_prompt'); - const goBtn = this._panel && this._panel.querySelector('#ai_replace_go'); - if (!promptEl || !goBtn) return; - const shown = promptEl.style.display !== 'none'; + _toggleReplace() { + const p = this._panel; + if (!p) return; + const promptEl = p.querySelector('#aie-prompt'); + const goBtn = p.querySelector('#aie-go'); + const shown = promptEl.style.display !== 'none'; promptEl.style.display = shown ? 'none' : 'inline-block'; goBtn.style.display = shown ? 'none' : 'inline-block'; - if (!shown) setTimeout(() => promptEl.focus(), 50); + if (!shown) setTimeout(() => promptEl.focus(), 40); } _removePanel() { if (this._panel) { this._panel.remove(); this._panel = null; } } - // ── Helpers ─────────────────────────────────────────────────────────────── + // ── Mask + image helpers ────────────────────────────────────────────────── + + _initMask() { + const w = config.layer.width_original; + const h = config.layer.height_original; + this._maskCanvas = document.createElement('canvas'); + this._maskCanvas.width = w; + this._maskCanvas.height = h; + this._maskCtx = this._maskCanvas.getContext('2d'); + this._hasMask = false; + this._samPoints = []; + this._samLabels = []; + } + + _getLayerB64() { + const layer = config.layer; + const c = document.createElement('canvas'); + c.width = layer.width_original; c.height = layer.height_original; + c.getContext('2d').drawImage(layer.link, 0, 0); + return c.toDataURL('image/png').split(',')[1]; + } _requireMask() { if (!this._hasMask) { - alertify.error('Paint over the area you want to change first.'); + alertify.error('Select an area first — click an object or use Brush.'); return false; } return true; } - /** Returns { imageB64, maskB64 } from current layer + painted mask. */ - _getImageAndMask() { - const layer = config.layer; - const w = layer.width_original; - const h = layer.height_original; - - // Image - const imgCanvas = document.createElement('canvas'); - imgCanvas.width = w; imgCanvas.height = h; - imgCanvas.getContext('2d').drawImage(layer.link, 0, 0); - const imageB64 = imgCanvas.toDataURL('image/png').split(',')[1]; - - // Mask (white = selected, black = keep) - const maskB64 = this._maskCanvas.toDataURL('image/png').split(',')[1]; - - return { imageB64, maskB64 }; - } - - _applyResult(resultB64, actionLabel) { + _applyResult(resultB64, label) { const img = new Image(); img.onload = () => { const rc = document.createElement('canvas'); rc.width = img.naturalWidth; rc.height = img.naturalHeight; rc.getContext('2d').drawImage(img, 0, 0); app.State.do_action( - new app.Actions.Bundle_action('ai_edit', actionLabel, [ + new app.Actions.Bundle_action('ai_edit', label, [ new app.Actions.Update_layer_image_action(rc) ]) ); alertify.dismissAll(); - alertify.success(`${actionLabel} applied.`); + alertify.success(label + ' applied.'); this._isRunning = false; this._doClear(); }; @@ -315,16 +511,15 @@ class Tools_ai_edit_class { this._isRunning = true; alertify.message('Erasing…', 0); try { - const { imageB64, maskB64 } = this._getImageAndMask(); - const base = window.API_BASE_URL || ''; + const maskB64 = this._maskCanvas.toDataURL('image/png').split(',')[1]; + const base = window.API_BASE_URL || ''; const r = await fetch(`${base}/api/erase`, { method: 'POST', headers: { 'Content-Type': 'application/json' }, - body: JSON.stringify({ image: imageB64, mask: maskB64 }), + body: JSON.stringify({ image: this._getLayerB64(), mask: maskB64 }), }); - if (!r.ok) throw new Error((await r.json().catch(() => ({}))).detail || 'Erase failed'); - const data = await r.json(); - this._applyResult(data.result, 'Erase'); + if (!r.ok) throw new Error((await r.json().catch(() => ({}))).detail || 'Failed'); + this._applyResult((await r.json()).result, 'Erase'); } catch (err) { alertify.dismissAll(); alertify.error('Erase failed: ' + (err.message || err)); @@ -334,20 +529,19 @@ class Tools_ai_edit_class { async _doReplace(prompt) { if (!this._requireMask() || this._isRunning) return; - if (!prompt) { alertify.error('Type what you want to put there.'); return; } + if (!prompt) { alertify.error('Describe what you want to put there.'); return; } this._isRunning = true; alertify.message(`Replacing: "${prompt}"…`, 0); try { - const { imageB64, maskB64 } = this._getImageAndMask(); - const base = window.API_BASE_URL || ''; + const maskB64 = this._maskCanvas.toDataURL('image/png').split(',')[1]; + const base = window.API_BASE_URL || ''; const r = await fetch(`${base}/api/inpaint/remote`, { method: 'POST', headers: { 'Content-Type': 'application/json' }, - body: JSON.stringify({ image: imageB64, mask: maskB64, prompt }), + body: JSON.stringify({ image: this._getLayerB64(), mask: maskB64, prompt }), }); - if (!r.ok) throw new Error((await r.json().catch(() => ({}))).detail || 'Replace failed'); - const data = await r.json(); - this._applyResult(data.result, `Replace: ${prompt}`); + if (!r.ok) throw new Error((await r.json().catch(() => ({}))).detail || 'Failed'); + this._applyResult((await r.json()).result, `Replace: ${prompt}`); } catch (err) { alertify.dismissAll(); alertify.error('Replace failed: ' + (err.message || err)); @@ -356,35 +550,28 @@ class Tools_ai_edit_class { } _doUpscale() { - // Delegate to the existing Upscale module - import('./../modules/image/upscale.js').then(m => { - const cls = m.default; - new cls().upscale(); - }); + import('./../modules/image/upscale.js').then(m => new m.default().upscale()); } _doExpand() { - import('./../modules/generate/outpaint.js').then(m => { - const cls = m.default; - new cls().outpaint(); - }); + import('./../modules/generate/outpaint.js').then(m => new m.default().outpaint()); } _doClear() { - if (this._maskCtx) { + if (this._maskCtx) this._maskCtx.clearRect(0, 0, this._maskCanvas.width, this._maskCanvas.height); + if (this._overlayEl) + this._overlayEl.getContext('2d').clearRect(0, 0, this._overlayEl.width, this._overlayEl.height); + const p = this._panel; + if (p) { + const promptEl = p.querySelector('#aie-prompt'); + const goBtn = p.querySelector('#aie-go'); + if (promptEl) { promptEl.style.display = 'none'; promptEl.value = ''; } + if (goBtn) goBtn.style.display = 'none'; } - if (this._overlayEl) { - this._overlayEl.getContext('2d').clearRect( - 0, 0, this._overlayEl.width, this._overlayEl.height - ); - } - // Hide the replace input - const promptEl = this._panel && this._panel.querySelector('#ai_replace_prompt'); - const goBtn = this._panel && this._panel.querySelector('#ai_replace_go'); - if (promptEl) { promptEl.style.display = 'none'; promptEl.value = ''; } - if (goBtn) goBtn.style.display = 'none'; - this._hasMask = false; + this._hasMask = false; + this._samPoints = []; + this._samLabels = []; } }