From df1749780841b218660c7f2e85a2e31f18615ad9 Mon Sep 17 00:00:00 2001 From: LyAhn Date: Tue, 2 Jun 2026 19:39:31 +0100 Subject: [PATCH] feat: add region-based similarity search Adds a "Search within image" button to the lightbox that lets the user draw a crop region on an image and find visually similar results using that crop's embedding. The crop is embedded in-memory without a temp file via a new embed_image_crop method on ClipImageEmbedder. Introduces a folder-scoped cosine search (search_image_ids_by_embedding_in_folder) to support the current-folder scope option, and wires up the new find_similar_by_region Tauri command end-to-end from Rust through to the Zustand store and Lightbox UI. --- src-tauri/src/commands.rs | 73 ++++++++++ src-tauri/src/embedder.rs | 45 ++++++ src-tauri/src/lib.rs | 1 + src-tauri/src/vector.rs | 34 +++++ src/components/Lightbox.tsx | 281 +++++++++++++++++++++++++++++++++--- src/store.ts | 67 +++++++++ 6 files changed, 477 insertions(+), 24 deletions(-) diff --git a/src-tauri/src/commands.rs b/src-tauri/src/commands.rs index d23eeb4..3c1a8ee 100644 --- a/src-tauri/src/commands.rs +++ b/src-tauri/src/commands.rs @@ -59,6 +59,19 @@ pub struct FindSimilarImagesParams { pub threshold: Option, } +#[derive(Deserialize)] +pub struct FindSimilarByRegionParams { + pub image_id: i64, + /// Normalized crop rect (0.0–1.0). + pub crop_x: f32, + pub crop_y: f32, + pub crop_w: f32, + pub crop_h: f32, + pub folder_id: Option, + pub offset: Option, + pub limit: Option, +} + #[derive(Deserialize)] pub struct DebugSimilarImagesParams { pub image_id: i64, @@ -354,6 +367,66 @@ pub async fn find_similar_images( }) } +#[tauri::command] +pub async fn find_similar_by_region( + db: State<'_, DbState>, + params: FindSimilarByRegionParams, +) -> Result { + let conn = db.get().map_err(|e| e.to_string())?; + let limit = params.limit.unwrap_or(32); + let offset = params.offset.unwrap_or(0); + + // Look up the source image path + let image = db::get_image_by_id(&conn, params.image_id).map_err(|e| e.to_string())?; + let image_path = std::path::Path::new(&image.path); + + // Embed the cropped region in-memory (no temp file needed) + let embedder = embedder::ClipImageEmbedder::new().map_err(|e| e.to_string())?; + let embedding = embedder + .embed_image_crop( + image_path, + params.crop_x, + params.crop_y, + params.crop_w, + params.crop_h, + ) + .map_err(|e| e.to_string())?; + + // Search for similar images using the crop embedding + let image_ids = match params.folder_id { + Some(folder_id) => vector::search_image_ids_by_embedding_in_folder( + &conn, + &embedding, + folder_id, + Some(params.image_id), + offset + limit + 1, + ) + .map_err(|e| e.to_string())?, + None => { + let mut ids = vector::search_image_ids_by_embedding(&conn, &embedding, offset + limit + 1) + .map_err(|e| e.to_string())?; + // Exclude the source image from global results + ids.retain(|&id| id != params.image_id); + ids + } + }; + + let has_more = image_ids.len() > offset + limit; + let page_ids = image_ids + .into_iter() + .skip(offset) + .take(limit) + .collect::>(); + + let images = db::get_images_by_ids(&conn, &page_ids).map_err(|e| e.to_string())?; + Ok(SimilarImagesPage { + images, + offset, + limit, + has_more, + }) +} + #[derive(Serialize)] pub struct SimilarImagesDebug { pub image_id: i64, diff --git a/src-tauri/src/embedder.rs b/src-tauri/src/embedder.rs index 532417d..3a977c7 100644 --- a/src-tauri/src/embedder.rs +++ b/src-tauri/src/embedder.rs @@ -74,6 +74,51 @@ impl ClipImageEmbedder { Ok(self.embed_images(&[path.to_path_buf()])?.remove(0)) } + /// Embed a cropped region of an image without writing a temp file to disk. + /// `crop_x`, `crop_y`, `crop_w`, `crop_h` are normalized 0.0–1.0 coordinates. + pub fn embed_image_crop( + &self, + path: &Path, + crop_x: f32, + crop_y: f32, + crop_w: f32, + crop_h: f32, + ) -> Result> { + let img = image::ImageReader::open(path)? + .with_guessed_format()? + .decode()?; + + let img_w = img.width() as f32; + let img_h = img.height() as f32; + + let x = ((crop_x * img_w) as u32).min(img.width().saturating_sub(1)); + let y = ((crop_y * img_h) as u32).min(img.height().saturating_sub(1)); + let w = ((crop_w * img_w) as u32).max(1).min(img.width() - x); + let h = ((crop_h * img_h) as u32).max(1).min(img.height() - y); + + let cropped = img.crop_imm(x, y, w, h); + let resized = cropped.resize_to_fill( + self.image_size as u32, + self.image_size as u32, + image::imageops::FilterType::Triangle, + ); + + let raw = resized.to_rgb8().into_raw(); + let tensor = candle_core::Tensor::from_vec( + raw, + (self.image_size, self.image_size, 3), + &candle_core::Device::Cpu, + )? + .permute((2, 0, 1))? + .to_dtype(candle_core::DType::F32)? + .affine(2.0 / 255.0, -1.0)?; + + let batch = tensor.unsqueeze(0)?.to_device(&self.device)?; + let features = self.model.get_image_features(&batch)?; + let normalized = candle_transformers::models::clip::div_l2_norm(&features)?; + Ok(normalized.get(0)?.flatten_all()?.to_vec1::()?) + } + pub fn embed_images(&self, paths: &[PathBuf]) -> Result>> { let images = load_images(paths, self.image_size)?.to_device(&self.device)?; let features = self.model.get_image_features(&images)?; diff --git a/src-tauri/src/lib.rs b/src-tauri/src/lib.rs index 39c9922..0ef2241 100644 --- a/src-tauri/src/lib.rs +++ b/src-tauri/src/lib.rs @@ -87,6 +87,7 @@ pub fn run() { commands::reindex_folder, commands::update_image_details, commands::find_similar_images, + commands::find_similar_by_region, commands::debug_similar_images, commands::retry_failed_embeddings, commands::semantic_search_images, diff --git a/src-tauri/src/vector.rs b/src-tauri/src/vector.rs index 0fa753c..644ece8 100644 --- a/src-tauri/src/vector.rs +++ b/src-tauri/src/vector.rs @@ -310,6 +310,40 @@ pub fn search_image_ids_by_embedding( Ok(ids) } +/// Brute-force cosine search scoped to a single folder, ordered by ascending distance. +/// Used for region-based similarity search where we want folder-scoped results. +pub fn search_image_ids_by_embedding_in_folder( + conn: &Connection, + embedding: &[f32], + folder_id: i64, + exclude_image_id: Option, + limit: usize, +) -> Result> { + if embedding.len() != CLIP_VECTOR_DIM { + return Err(anyhow!( + "expected {}-dimensional embedding, got {}", + CLIP_VECTOR_DIM, + embedding.len() + )); + } + + let packed = pack_f32(embedding); + let exclude_id = exclude_image_id.unwrap_or(-1); + let mut stmt = conn.prepare( + "SELECT v.image_id + FROM image_vec v + JOIN images i ON i.id = v.image_id + WHERE i.folder_id = ?2 + AND v.image_id != ?3 + ORDER BY vec_distance_cosine(v.embedding, vec_f32(?1)) ASC + LIMIT ?4", + )?; + let rows = stmt.query_map((&packed, folder_id, exclude_id, limit as i64), |row| { + row.get::<_, i64>(0) + })?; + Ok(rows.collect::>>()?) +} + #[allow(dead_code)] pub fn search_caption_ids_by_embedding( conn: &Connection, diff --git a/src/components/Lightbox.tsx b/src/components/Lightbox.tsx index b357d0f..0ee42f5 100644 --- a/src/components/Lightbox.tsx +++ b/src/components/Lightbox.tsx @@ -58,12 +58,69 @@ function ratingPill(rating: AiRating): { label: string; className: string } { } } +interface DragRect { + startX: number; + startY: number; + endX: number; + endY: number; +} + +/** Compute a CSS-pixel rect (relative to the viewport container) from a DragRect. */ +function normaliseRect(r: DragRect): { left: number; top: number; width: number; height: number } { + return { + left: Math.min(r.startX, r.endX), + top: Math.min(r.startY, r.endY), + width: Math.abs(r.endX - r.startX), + height: Math.abs(r.endY - r.startY), + }; +} + +/** Convert a CSS-pixel drag rect (relative to viewport container) to normalised 0–1 crop coords + * relative to the actual rendered element bounds. */ +function rectToNormalisedCrop( + rect: DragRect, + imgEl: HTMLImageElement, +): { x: number; y: number; w: number; h: number } | null { + const imgBounds = imgEl.getBoundingClientRect(); + if (imgBounds.width === 0 || imgBounds.height === 0) return null; + + // rect coords are already in viewport space (client coords) + const rawX = Math.min(rect.startX, rect.endX); + const rawY = Math.min(rect.startY, rect.endY); + const rawW = Math.abs(rect.endX - rect.startX); + const rawH = Math.abs(rect.endY - rect.startY); + + // Clamp to image bounds + const clampedX = Math.max(rawX, imgBounds.left); + const clampedY = Math.max(rawY, imgBounds.top); + const clampedRight = Math.min(rawX + rawW, imgBounds.right); + const clampedBottom = Math.min(rawY + rawH, imgBounds.bottom); + + const croppedW = clampedRight - clampedX; + const croppedH = clampedBottom - clampedY; + + if (croppedW <= 0 || croppedH <= 0) return null; + + // Normalize by the CSS transform scale — getBoundingClientRect already returns + // the scaled (on-screen) size, so we normalize directly against that. + return { + x: (clampedX - imgBounds.left) / imgBounds.width, + y: (clampedY - imgBounds.top) / imgBounds.height, + w: croppedW / imgBounds.width, + h: croppedH / imgBounds.height, + }; +} + +/** Minimum selection size as a fraction of the viewport container dimension. */ +const MIN_SELECTION_FRACTION = 0.02; + export function Lightbox() { const selectedImage = useGalleryStore((state) => state.selectedImage); const closeImage = useGalleryStore((state) => state.closeImage); const images = useGalleryStore((state) => state.images); const openImage = useGalleryStore((state) => state.openImage); const loadSimilarImages = useGalleryStore((state) => state.loadSimilarImages); + const loadSimilarByRegion = useGalleryStore((state) => state.loadSimilarByRegion); const similarScope = useGalleryStore((state) => state.similarScope); const updateImageDetails = useGalleryStore((state) => state.updateImageDetails); const getImageTags = useGalleryStore((state) => state.getImageTags); @@ -71,16 +128,26 @@ export function Lightbox() { const removeTag = useGalleryStore((state) => state.removeTag); const taggerModelStatus = useGalleryStore((state) => state.taggerModelStatus); const queueTaggingForImage = useGalleryStore((state) => state.queueTaggingForImage); + const [zoom, setZoom] = useState(1); const [imageTags, setImageTags] = useState([]); const [tagInput, setTagInput] = useState(""); const [tagAdding, setTagAdding] = useState(false); const [tagsExpanded, setTagsExpanded] = useState(false); const [taggingQueued, setTaggingQueued] = useState(false); + + // Region selection state + const [regionSelectMode, setRegionSelectMode] = useState(false); + const [isDragging, setIsDragging] = useState(false); + const [dragRect, setDragRect] = useState(null); + const [regionSearching, setRegionSearching] = useState(false); + const imageViewportRef = useRef(null); + const imgRef = useRef(null); const currentIndex = selectedImage ? images.findIndex((image) => image.id === selectedImage.id) : -1; const canFindSimilar = selectedImage?.embedding_status === "ready"; + const canSearchRegion = canFindSimilar && selectedImage?.media_kind === "image"; const goPrev = useCallback(() => { if (currentIndex > 0) openImage(images[currentIndex - 1]); @@ -90,13 +157,21 @@ export function Lightbox() { if (currentIndex >= 0 && currentIndex < images.length - 1) openImage(images[currentIndex + 1]); }, [currentIndex, images, openImage]); + const exitRegionMode = useCallback(() => { + setRegionSelectMode(false); + setIsDragging(false); + setDragRect(null); + }, []); + useEffect(() => { setZoom(1); setImageTags([]); setTagInput(""); setTagsExpanded(false); setTaggingQueued(false); - }, [selectedImage?.id]); + exitRegionMode(); + setRegionSearching(false); + }, [selectedImage?.id, exitRegionMode]); useEffect(() => { if (!selectedImage) return; @@ -115,6 +190,7 @@ export function Lightbox() { if (!viewport || !selectedImage || selectedImage.media_kind !== "image") return; const handleWheel = (event: WheelEvent) => { + if (regionSelectMode) return; // don't zoom during selection if (!event.ctrlKey && Math.abs(event.deltaY) < Math.abs(event.deltaX)) return; event.preventDefault(); setZoom((value) => { @@ -125,12 +201,19 @@ export function Lightbox() { viewport.addEventListener("wheel", handleWheel, { passive: false }); return () => viewport.removeEventListener("wheel", handleWheel); - }, [selectedImage]); + }, [selectedImage, regionSelectMode]); useEffect(() => { const handler = (event: KeyboardEvent) => { if (!selectedImage) return; - if (event.key === "Escape") closeImage(); + if (event.key === "Escape") { + if (regionSelectMode) { + exitRegionMode(); + } else { + closeImage(); + } + } + if (regionSelectMode) return; // block nav keys during selection if (event.key === "ArrowLeft") goPrev(); if (event.key === "ArrowRight") goNext(); if (event.key === "+" || event.key === "=") setZoom((value) => Math.min(3, value + 0.25)); @@ -139,7 +222,75 @@ export function Lightbox() { window.addEventListener("keydown", handler); return () => window.removeEventListener("keydown", handler); - }, [selectedImage, closeImage, goPrev, goNext]); + }, [selectedImage, closeImage, goPrev, goNext, regionSelectMode, exitRegionMode]); + + // ── Region selection pointer handlers ─────────────────────────────────────── + const handleRegionPointerDown = useCallback( + (event: React.PointerEvent) => { + if (!regionSelectMode) return; + event.preventDefault(); + event.currentTarget.setPointerCapture(event.pointerId); + setIsDragging(true); + setDragRect({ + startX: event.clientX, + startY: event.clientY, + endX: event.clientX, + endY: event.clientY, + }); + }, + [regionSelectMode], + ); + + const handleRegionPointerMove = useCallback( + (event: React.PointerEvent) => { + if (!isDragging) return; + setDragRect((prev) => + prev ? { ...prev, endX: event.clientX, endY: event.clientY } : null, + ); + }, + [isDragging], + ); + + const handleRegionPointerUp = useCallback( + (event: React.PointerEvent) => { + if (!isDragging || !dragRect || !selectedImage || !imgRef.current) { + setIsDragging(false); + return; + } + event.currentTarget.releasePointerCapture(event.pointerId); + + const finalRect: DragRect = { ...dragRect, endX: event.clientX, endY: event.clientY }; + const crop = rectToNormalisedCrop(finalRect, imgRef.current); + + setIsDragging(false); + setDragRect(null); + + // Ignore tiny accidental clicks + const containerBounds = imageViewportRef.current?.getBoundingClientRect(); + const containerSize = containerBounds + ? Math.min(containerBounds.width, containerBounds.height) + : 500; + const selW = Math.abs(finalRect.endX - finalRect.startX); + const selH = Math.abs(finalRect.endY - finalRect.startY); + if (!crop || selW < containerSize * MIN_SELECTION_FRACTION || selH < containerSize * MIN_SELECTION_FRACTION) { + exitRegionMode(); + return; + } + + exitRegionMode(); + setRegionSearching(true); + + const folderId = + similarScope === "current_folder" ? selectedImage.folder_id : null; + void loadSimilarByRegion(selectedImage.id, crop, folderId, selectedImage.folder_id) + .finally(() => setRegionSearching(false)); + }, + [isDragging, dragRect, selectedImage, similarScope, loadSimilarByRegion, exitRegionMode], + ); + + // Build the CSS rect for the selection overlay (viewport-relative) + const selectionOverlay = + isDragging && dragRect ? normaliseRect(dragRect) : null; return ( @@ -151,11 +302,11 @@ export function Lightbox() { animate={{ opacity: 1 }} exit={{ opacity: 0 }} transition={{ duration: 0.15 }} - onClick={closeImage} + onClick={regionSelectMode ? undefined : closeImage} > - {Math.round(zoom * 100)}% - + {!regionSelectMode && ( +
+
+ + {Math.round(zoom * 100)}% + +
- + )} )} @@ -261,6 +447,53 @@ export function Lightbox() { + {/* Search region button row */} + {canSearchRegion && ( +
+ +
+ )} +

Rating

@@ -472,7 +705,7 @@ export function Lightbox() {