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/config.js b/frontend/src/js/config.js index 466112c..6678895 100644 --- a/frontend/src/js/config.js +++ b/frontend/src/js/config.js @@ -105,25 +105,8 @@ config.TOOLS = [ attributes: {}, }, { - name: 'ai_inpaint', - title: 'AI Inpaint', - on_activate: 'on_activate', - attributes: {}, - }, - { - name: 'ai_lama_erase', - title: 'AI Magic Erase (LaMa) - Paint over to erase', - attributes: { - size: { - value: 30, - min: 5, - max: 200, - }, - }, - }, - { - name: 'ai_smart_inpaint', - title: 'AI Smart Inpaint - Paint + describe replacement', + name: 'ai_edit', + title: 'AI Edit — paint mask, then Erase / Replace / Upscale / Expand', on_activate: 'on_activate', attributes: { size: { @@ -133,12 +116,6 @@ config.TOOLS = [ }, }, }, - { - name: 'ai_replace_selection', - title: 'AI Replace Selection - Use any selection tool first', - on_activate: 'on_activate', - attributes: {}, - }, { name: 'magic_wand', title: 'Magic Wand (Color Select)', diff --git a/frontend/src/js/tools/ai_edit.js b/frontend/src/js/tools/ai_edit.js new file mode 100644 index 0000000..14105da --- /dev/null +++ b/frontend/src/js/tools/ai_edit.js @@ -0,0 +1,578 @@ +/** + * AI Edit — unified smart selection + inpainting tool. + * + * Workflow: + * 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 + * + * 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'; +import config from './../config.js'; +import Base_layers_class from './../core/base-layers.js'; +import alertify from './../../../node_modules/alertifyjs/build/alertify.min.js'; + +var instance = null; + +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'; + // 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 ──────────────────────────────────────────────────────── + + on_activate() { + if (!config.layer || config.layer.type !== 'image') { + alertify.error('Select an image layer first.'); + return; + } + this._initMask(); + this._mountOverlay(); + this._mountPanel(); + } + + on_leave() { + this._removeOverlay(); + this._removePanel(); + this._painting = false; + } + + // ── Input routing ───────────────────────────────────────────────────────── + + mousedown(e) { + if (!config.layer || config.layer.type !== 'image') return; + if (this._mode === 'sam') { + this._handleSamClick(e); + } else { + this._painting = true; + this._brushPaint(e); + } + } + + mousemove(e) { + if (this._mode !== 'sam' && this._painting) this._brushPaint(e); + } + + mouseup() { + if (this._painting) { + this._painting = false; + if (this._hasMask) this._showActions(); + } + } + + // ── Coordinate mapping ──────────────────────────────────────────────────── + + _screenToImage(e) { + const canvasEl = document.getElementById('canvas_minipaint') || document.querySelector('canvas'); + 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 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; + + // 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(ix, iy, r, 0, Math.PI * 2); + this._maskCtx.fill(); + this._maskCtx.globalCompositeOperation = 'source-over'; + + // Mirror on overlay + const oc = this._overlayEl; + if (!oc) return; + const oct = oc.getContext('2d'); + 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 ─────────────────────────────────────────────────────────────── + + _mountOverlay() { + this._removeOverlay(); + const base = document.getElementById('canvas_minipaint') || document.querySelector('canvas'); + if (!base) return; + const oc = document.createElement('canvas'); + oc.id = 'ai_edit_overlay'; + 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', + }); + 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; } + } + + // ── Panel ───────────────────────────────────────────────────────────────── + + _mountPanel() { + this._removePanel(); + const panel = document.createElement('div'); + panel.id = 'ai_edit_panel'; + Object.assign(panel.style, { + 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); + this._panel = panel; + this._wirePanel(); + } + + _panelHTML() { + return ` + + + +
+ Select: + + + +
+ Click an object to select it. Alt+click to deselect. +
+ + +
+ Then: + +
+ + + +
+ + + +
`; + } + + _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.', + }; + + // Mode buttons + this._panel.querySelectorAll('[data-mode]').forEach(btn => { + btn.addEventListener('click', () => { + _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'); + }); + }); + + // 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 => { + if (e.key === 'Enter') _this._doReplace(promptEl.value.trim()); + }); + } + } + + _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(), 40); + } + + _removePanel() { + if (this._panel) { this._panel.remove(); this._panel = null; } + } + + // ── 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('Select an area first — click an object or use Brush.'); + return false; + } + return true; + } + + _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', label, [ + new app.Actions.Update_layer_image_action(rc) + ]) + ); + alertify.dismissAll(); + alertify.success(label + ' applied.'); + this._isRunning = false; + this._doClear(); + }; + img.onerror = () => { + alertify.dismissAll(); + alertify.error('Failed to load result image.'); + this._isRunning = false; + }; + img.src = 'data:image/png;base64,' + resultB64; + } + + // ── Actions ─────────────────────────────────────────────────────────────── + + async _doErase() { + if (!this._requireMask() || this._isRunning) return; + this._isRunning = true; + alertify.message('Erasing…', 0); + try { + 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: this._getLayerB64(), mask: maskB64 }), + }); + 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)); + this._isRunning = false; + } + } + + async _doReplace(prompt) { + if (!this._requireMask() || this._isRunning) return; + if (!prompt) { alertify.error('Describe what you want to put there.'); return; } + this._isRunning = true; + alertify.message(`Replacing: "${prompt}"…`, 0); + try { + 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: this._getLayerB64(), mask: maskB64, 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)); + this._isRunning = false; + } + } + + _doUpscale() { + import('./../modules/image/upscale.js').then(m => new m.default().upscale()); + } + + _doExpand() { + import('./../modules/generate/outpaint.js').then(m => new m.default().outpaint()); + } + + _doClear() { + 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'; + } + this._hasMask = false; + this._samPoints = []; + this._samLabels = []; + } +} + +export default Tools_ai_edit_class;