use crate::vector; use anyhow::Result; use hnsw_rs::prelude::{DistCosine, Hnsw, Neighbour}; use rusqlite::Connection; use std::collections::HashMap; use std::sync::{OnceLock, RwLock}; const HNSW_MAX_CONNECTIONS: usize = 24; const HNSW_EF_CONSTRUCTION: usize = 300; const HNSW_EF_SEARCH: usize = 96; struct CachedHnswIndex { revision: String, image_ids_by_external: Vec, external_by_image_id: HashMap, hnsw: Hnsw<'static, f32, DistCosine>, } static IMAGE_HNSW_INDEX: OnceLock>> = OnceLock::new(); fn cache() -> &'static RwLock> { IMAGE_HNSW_INDEX.get_or_init(|| RwLock::new(None)) } fn build_index(conn: &Connection) -> Result { // Read the revision *before* fetching embeddings so we can detect any write // that races with the build. If the revision advances while we are building, // the resulting index would be stale — retry until it is stable. loop { let revision_before = vector::get_embedding_revision(conn)?; let embeddings = vector::get_all_image_embeddings_with_ids(conn, None)?; let max_elements = embeddings.len().max(1); let max_layer = 16.min((max_elements as f32).ln().trunc() as usize).max(1); let mut hnsw = Hnsw::::new( HNSW_MAX_CONNECTIONS, max_elements, max_layer, HNSW_EF_CONSTRUCTION, DistCosine {}, ); let image_ids_by_external = embeddings .iter() .map(|(image_id, _)| *image_id) .collect::>(); let external_by_image_id = image_ids_by_external .iter() .enumerate() .map(|(external_id, image_id)| (*image_id, external_id)) .collect::>(); let data_with_id = embeddings .iter() .enumerate() .map(|(external_id, (_, embedding))| (embedding, external_id)) .collect::>(); hnsw.parallel_insert(&data_with_id); hnsw.set_searching_mode(true); // If the revision is unchanged the index reflects a consistent snapshot. let revision_after = vector::get_embedding_revision(conn)?; if revision_before == revision_after { return Ok(CachedHnswIndex { revision: revision_after, image_ids_by_external, external_by_image_id, hnsw, }); } // A concurrent write advanced the revision — discard this build and retry. } } fn ensure_index(conn: &Connection) -> Result<()> { let revision = vector::get_embedding_revision(conn)?; { let guard = cache().read().expect("hnsw cache poisoned"); if guard.as_ref().map(|cached| cached.revision.as_str()) == Some(revision.as_str()) { return Ok(()); } } let next = build_index(conn)?; let mut guard = cache().write().expect("hnsw cache poisoned"); *guard = Some(next); Ok(()) } pub fn find_similar_image_matches( conn: &Connection, image_id: i64, folder_id: Option, threshold: f32, offset: usize, limit: usize, ) -> Result> { ensure_index(conn)?; let query_embedding = match vector::get_image_embedding(conn, image_id)? { Some(embedding) => embedding, None => return Ok(Vec::new()), }; // Fetch folder image IDs *before* acquiring the read lock so we don't hold // the lock across a potentially slow SQLite query, which would delay any // concurrent ensure_index call waiting for a write lock. let folder_image_ids: Option> = if let Some(folder_id) = folder_id { let ids = vector::get_all_image_embeddings_with_ids(conn, Some(folder_id))? .into_iter() .map(|(id, _)| id) .collect(); Some(ids) } else { None }; let guard = cache().read().expect("hnsw cache poisoned"); let Some(cached) = guard.as_ref() else { return Ok(Vec::new()); }; let knbn = (offset + limit).max(limit).saturating_add(32); let neighbours: Vec = if let Some(image_ids) = folder_image_ids { let mut allowed_ids = image_ids .into_iter() .filter_map(|allowed_image_id| { cached.external_by_image_id.get(&allowed_image_id).copied() }) .collect::>(); allowed_ids.sort_unstable(); cached .hnsw .search_filter(&query_embedding, knbn, HNSW_EF_SEARCH, Some(&allowed_ids)) } else { cached.hnsw.search(&query_embedding, knbn, HNSW_EF_SEARCH) }; let matches = neighbours .into_iter() .filter_map(|neighbour| { let image_id_match = cached.image_ids_by_external.get(neighbour.d_id).copied()?; if image_id_match == image_id || neighbour.distance > threshold { return None; } Some((image_id_match, neighbour.distance)) }) .skip(offset) .take(limit) .collect::>(); Ok(matches) }