Skip to main content

wm_memory/
vector.rs

1//! Vector store — cosine similarity search over memory embeddings.
2//!
3//! Provides true vector similarity search using embeddings stored in the
4//! Embeddings galaxy (LMDB). Builds an in-memory index from all stored
5//! embeddings for fast cosine similarity lookups.
6//!
7//! For small-to-medium datasets (<100K memories), brute-force cosine
8//! similarity is fast enough (sub-millisecond for 10K vectors at 384 dims).
9//! For larger datasets, LanceDB can be added as an optional backend.
10
11use ahash::AHashMap;
12use lmdb::{Cursor, Transaction};
13use uuid::Uuid;
14use wm_core::{CoreError, Galaxy, Result};
15
16use crate::MemoryStore;
17
18/// A vector search result.
19#[derive(Debug, Clone)]
20pub struct VectorSearchResult {
21    /// Memory UUID
22    pub memory_id: Uuid,
23    /// Galaxy the memory belongs to
24    pub galaxy: Galaxy,
25    /// Cosine similarity score (0.0 to 1.0, higher = more similar)
26    pub score: f32,
27}
28
29/// Trait for pluggable vector search backends.
30///
31/// Implemented by:
32/// - `VectorStore` — in-memory brute-force cosine similarity (default)
33/// - `LanceVectorStore` — LanceDB-backed ANN search (feature-gated under `lancedb`)
34pub trait VectorSearchEngine: Send + Sync {
35    /// Add a vector to the index.
36    fn add_vector(&mut self, memory_id: Uuid, galaxy: Galaxy, embedding: Vec<f32>);
37
38    /// Remove a vector from the index.
39    fn remove_vector(&mut self, memory_id: Uuid) -> bool;
40
41    /// Search for the most similar vectors to a query embedding.
42    ///
43    /// Returns results sorted by similarity (highest first).
44    /// Optionally filter by galaxy.
45    fn search_vectors(
46        &self,
47        query: &[f32],
48        limit: usize,
49        galaxy_filter: Option<Galaxy>,
50    ) -> Vec<VectorSearchResult>;
51
52    /// Search for similar vectors to a memory's own embedding.
53    ///
54    /// Excludes the memory itself from results.
55    fn search_similar_vectors(&self, memory_id: Uuid, limit: usize) -> Vec<VectorSearchResult>;
56
57    /// Get the number of indexed vectors.
58    fn vector_count(&self) -> usize;
59
60    /// Whether the index is empty.
61    fn is_index_empty(&self) -> bool {
62        self.vector_count() == 0
63    }
64
65    /// Load all embeddings from the LMDB Embeddings galaxy.
66    fn load_vectors(&mut self, store: &MemoryStore) -> Result<()>;
67
68    /// Clear the index.
69    fn clear_vectors(&mut self);
70}
71
72/// In-memory vector index for fast cosine similarity search.
73///
74/// Loads all embeddings from the Embeddings galaxy into memory on first access.
75/// Supports incremental updates (add/remove vectors without full rebuild).
76pub struct VectorStore {
77    /// Indexed vectors: (memory_id, galaxy, embedding)
78    vectors: AHashMap<Uuid, (Galaxy, Vec<f32>)>,
79    /// Whether the index has been loaded from LMDB
80    loaded: bool,
81}
82
83impl VectorStore {
84    /// Create a new empty vector store.
85    #[must_use]
86    pub fn new() -> Self {
87        Self {
88            vectors: AHashMap::new(),
89            loaded: false,
90        }
91    }
92
93    /// Load all embeddings from the LMDB Embeddings galaxy into memory.
94    ///
95    /// This scans the Embeddings galaxy and loads all stored embedding vectors.
96    /// Must be called before `search` if vectors were stored directly via
97    /// `MemoryStore::put_embedding` without going through `add`.
98    pub fn load(&mut self, store: &MemoryStore) -> Result<()> {
99        let db = store.galaxy_db(Galaxy::Embeddings)?;
100
101        // Pass 1: Collect all embeddings from LMDB within a single read txn
102        let mut entries: Vec<(Uuid, Vec<f32>)> = Vec::new();
103        {
104            let tx = store
105                .env()
106                .begin_ro_txn()
107                .map_err(|e| CoreError::Memory(format!("LMDB ro_txn failed: {e}")))?;
108
109            let mut cursor = tx
110                .open_ro_cursor(db)
111                .map_err(|e| CoreError::Memory(format!("LMDB cursor failed: {e}")))?;
112
113            for (key, val) in cursor.iter() {
114                if key.len() == 16 {
115                    let bytes: [u8; 16] = key.try_into().unwrap_or([0u8; 16]);
116                    let id = Uuid::from_bytes(bytes);
117                    let embedding = crate::memory::decode_embedding(val);
118                    entries.push((id, embedding));
119                }
120            }
121
122            drop(cursor);
123            tx.commit()
124                .map_err(|e| CoreError::Memory(format!("LMDB commit failed: {e}")))?;
125        }
126
127        // Pass 2: Look up galaxy for each embedding (separate txn per lookup)
128        let mut count = 0;
129        for (id, embedding) in entries {
130            match self.find_memory_galaxy(store, id) {
131                Some(galaxy) => {
132                    self.vectors.insert(id, (galaxy, embedding));
133                    count += 1;
134                }
135                None => {
136                    tracing::warn!(
137                        "Skipping orphaned embedding (memory not found in any galaxy, id={})",
138                        id
139                    );
140                }
141            }
142        }
143
144        self.loaded = true;
145        tracing::info!("Loaded {count} embedding vectors into VectorStore");
146        Ok(())
147    }
148
149    /// Find which galaxy a memory belongs to by scanning all galaxies.
150    fn find_memory_galaxy(&self, store: &MemoryStore, id: Uuid) -> Option<Galaxy> {
151        for galaxy in Galaxy::all() {
152            if galaxy == Galaxy::Embeddings {
153                continue;
154            }
155            if store.get(galaxy, id).ok().flatten().is_some() {
156                return Some(galaxy);
157            }
158        }
159        None
160    }
161
162    /// Add a vector to the index.
163    pub fn add(&mut self, memory_id: Uuid, galaxy: Galaxy, embedding: Vec<f32>) {
164        self.vectors.insert(memory_id, (galaxy, embedding));
165    }
166
167    /// Remove a vector from the index.
168    pub fn remove(&mut self, memory_id: Uuid) -> bool {
169        self.vectors.remove(&memory_id).is_some()
170    }
171
172    /// Get the number of indexed vectors.
173    #[must_use]
174    pub fn len(&self) -> usize {
175        self.vectors.len()
176    }
177
178    /// Check if the index is empty.
179    #[must_use]
180    pub fn is_empty(&self) -> bool {
181        self.vectors.is_empty()
182    }
183
184    /// Whether the index has been loaded from LMDB.
185    #[must_use]
186    pub const fn is_loaded(&self) -> bool {
187        self.loaded
188    }
189
190    /// Search for the most similar vectors to a query embedding.
191    ///
192    /// Returns results sorted by cosine similarity (highest first).
193    /// Optionally filter by galaxy.
194    #[must_use]
195    pub fn search(
196        &self,
197        query: &[f32],
198        limit: usize,
199        galaxy_filter: Option<Galaxy>,
200    ) -> Vec<VectorSearchResult> {
201        if self.vectors.is_empty() || query.is_empty() {
202            return Vec::new();
203        }
204
205        let query_norm = vector_norm(query);
206        if query_norm == 0.0 {
207            return Vec::new();
208        }
209
210        let mut results: Vec<VectorSearchResult> = self
211            .vectors
212            .iter()
213            .filter(|(_, (galaxy, _))| galaxy_filter.is_none_or(|g| g == *galaxy))
214            .filter_map(|(id, (galaxy, embedding))| {
215                let score = cosine_similarity(query, embedding, query_norm);
216                if score > 0.0 {
217                    Some(VectorSearchResult {
218                        memory_id: *id,
219                        galaxy: *galaxy,
220                        score,
221                    })
222                } else {
223                    None
224                }
225            })
226            .collect();
227
228        results.sort_by(|a, b| {
229            b.score
230                .partial_cmp(&a.score)
231                .unwrap_or(std::cmp::Ordering::Equal)
232        });
233        results.truncate(limit);
234        results
235    }
236
237    /// Search for similar vectors to a memory's own embedding.
238    ///
239    /// Excludes the memory itself from results.
240    #[must_use]
241    pub fn search_similar_to(&self, memory_id: Uuid, limit: usize) -> Vec<VectorSearchResult> {
242        let (galaxy, embedding) = match self.vectors.get(&memory_id) {
243            Some(v) => v,
244            None => return Vec::new(),
245        };
246
247        let query_norm = vector_norm(embedding);
248        if query_norm == 0.0 {
249            return Vec::new();
250        }
251
252        let mut results: Vec<VectorSearchResult> = self
253            .vectors
254            .iter()
255            .filter(|(id, _)| **id != memory_id)
256            .filter_map(|(id, (g, emb))| {
257                let score = cosine_similarity(embedding, emb, query_norm);
258                if score > 0.0 {
259                    Some(VectorSearchResult {
260                        memory_id: *id,
261                        galaxy: *g,
262                        score,
263                    })
264                } else {
265                    None
266                }
267            })
268            .collect();
269
270        let _ = galaxy; // galaxy filter not applied for similar-to
271        results.sort_by(|a, b| {
272            b.score
273                .partial_cmp(&a.score)
274                .unwrap_or(std::cmp::Ordering::Equal)
275        });
276        results.truncate(limit);
277        results
278    }
279
280    /// Clear the index.
281    pub fn clear(&mut self) {
282        self.vectors.clear();
283        self.loaded = false;
284    }
285}
286
287impl Default for VectorStore {
288    fn default() -> Self {
289        Self::new()
290    }
291}
292
293impl VectorSearchEngine for VectorStore {
294    fn add_vector(&mut self, memory_id: Uuid, galaxy: Galaxy, embedding: Vec<f32>) {
295        self.add(memory_id, galaxy, embedding);
296    }
297
298    fn remove_vector(&mut self, memory_id: Uuid) -> bool {
299        self.remove(memory_id)
300    }
301
302    fn search_vectors(
303        &self,
304        query: &[f32],
305        limit: usize,
306        galaxy_filter: Option<Galaxy>,
307    ) -> Vec<VectorSearchResult> {
308        self.search(query, limit, galaxy_filter)
309    }
310
311    fn search_similar_vectors(&self, memory_id: Uuid, limit: usize) -> Vec<VectorSearchResult> {
312        self.search_similar_to(memory_id, limit)
313    }
314
315    fn vector_count(&self) -> usize {
316        self.len()
317    }
318
319    fn load_vectors(&mut self, store: &MemoryStore) -> Result<()> {
320        self.load(store)
321    }
322
323    fn clear_vectors(&mut self) {
324        self.clear();
325    }
326}
327
328/// Compute the L2 norm of a vector.
329fn vector_norm(v: &[f32]) -> f32 {
330    v.iter().map(|x| x * x).sum::<f32>().sqrt()
331}
332
333/// Compute cosine similarity between two vectors.
334///
335/// `query_norm` is pre-computed for the query vector to avoid redundant
336/// calculations when searching against many vectors.
337fn cosine_similarity(query: &[f32], target: &[f32], query_norm: f32) -> f32 {
338    if query.len() != target.len() {
339        return 0.0;
340    }
341
342    let dot: f32 = query.iter().zip(target.iter()).map(|(a, b)| a * b).sum();
343
344    let target_norm = vector_norm(target);
345    if target_norm == 0.0 {
346        return 0.0;
347    }
348
349    dot / (query_norm * target_norm)
350}
351
352#[cfg(test)]
353mod tests {
354    use super::*;
355    use crate::Memory;
356
357    #[test]
358    fn vector_store_empty_search() {
359        let vs = VectorStore::new();
360        let results = vs.search(&[1.0, 0.0, 0.0], 10, None);
361        assert!(results.is_empty());
362    }
363
364    #[test]
365    fn vector_store_add_and_search() {
366        let mut vs = VectorStore::new();
367        let id1 = Uuid::new_v4();
368        let id2 = Uuid::new_v4();
369        let id3 = Uuid::new_v4();
370
371        vs.add(id1, Galaxy::Codex, vec![1.0, 0.0, 0.0]);
372        vs.add(id2, Galaxy::Codex, vec![0.0, 1.0, 0.0]);
373        vs.add(id3, Galaxy::Codex, vec![1.0, 1.0, 0.0]);
374
375        let results = vs.search(&[1.0, 0.0, 0.0], 10, None);
376        assert_eq!(results.len(), 2); // id1 and id3 have positive similarity
377        assert_eq!(results[0].memory_id, id1);
378        assert!((results[0].score - 1.0).abs() < 0.001); // exact match
379    }
380
381    #[test]
382    fn vector_store_search_with_limit() {
383        let mut vs = VectorStore::new();
384        for _ in 0..10 {
385            vs.add(Uuid::new_v4(), Galaxy::Codex, vec![1.0, 0.0, 0.0]);
386        }
387        let results = vs.search(&[1.0, 0.0, 0.0], 3, None);
388        assert_eq!(results.len(), 3);
389    }
390
391    #[test]
392    fn vector_store_galaxy_filter() {
393        let mut vs = VectorStore::new();
394        vs.add(Uuid::new_v4(), Galaxy::Codex, vec![1.0, 0.0]);
395        vs.add(Uuid::new_v4(), Galaxy::Research, vec![1.0, 0.0]);
396        vs.add(Uuid::new_v4(), Galaxy::Codex, vec![0.9, 0.1]);
397
398        let results = vs.search(&[1.0, 0.0], 10, Some(Galaxy::Codex));
399        assert_eq!(results.len(), 2);
400        assert!(results.iter().all(|r| r.galaxy == Galaxy::Codex));
401    }
402
403    #[test]
404    fn vector_store_search_similar_to() {
405        let mut vs = VectorStore::new();
406        let id1 = Uuid::new_v4();
407        let id2 = Uuid::new_v4();
408        let id3 = Uuid::new_v4();
409
410        vs.add(id1, Galaxy::Codex, vec![1.0, 0.0, 0.0]);
411        vs.add(id2, Galaxy::Codex, vec![0.95, 0.05, 0.0]);
412        vs.add(id3, Galaxy::Codex, vec![0.0, 1.0, 0.0]);
413
414        let results = vs.search_similar_to(id1, 10);
415        // id3 is orthogonal (cosine sim = 0.0), so only id2 has positive similarity
416        assert_eq!(results.len(), 1);
417        assert!(results.iter().all(|r| r.memory_id != id1));
418        assert_eq!(results[0].memory_id, id2); // most similar
419    }
420
421    #[test]
422    fn vector_store_remove() {
423        let mut vs = VectorStore::new();
424        let id = Uuid::new_v4();
425        vs.add(id, Galaxy::Codex, vec![1.0, 0.0]);
426        assert_eq!(vs.len(), 1);
427        assert!(vs.remove(id));
428        assert_eq!(vs.len(), 0);
429        assert!(!vs.remove(id));
430    }
431
432    #[test]
433    fn vector_store_clear() {
434        let mut vs = VectorStore::new();
435        vs.add(Uuid::new_v4(), Galaxy::Codex, vec![1.0]);
436        vs.add(Uuid::new_v4(), Galaxy::Codex, vec![1.0]);
437        assert_eq!(vs.len(), 2);
438        vs.clear();
439        assert_eq!(vs.len(), 0);
440        assert!(!vs.is_loaded());
441    }
442
443    #[test]
444    fn vector_store_zero_query_returns_empty() {
445        let mut vs = VectorStore::new();
446        vs.add(Uuid::new_v4(), Galaxy::Codex, vec![1.0, 0.0]);
447        let results = vs.search(&[0.0, 0.0], 10, None);
448        assert!(results.is_empty());
449    }
450
451    #[test]
452    fn vector_store_mismatched_dimensions() {
453        let mut vs = VectorStore::new();
454        vs.add(Uuid::new_v4(), Galaxy::Codex, vec![1.0, 0.0, 0.0]);
455        let results = vs.search(&[1.0, 0.0], 10, None);
456        assert!(results.is_empty()); // dimension mismatch → score 0.0
457    }
458
459    #[test]
460    fn cosine_similarity_exact_match() {
461        let sim = cosine_similarity(&[1.0, 0.0, 0.0], &[1.0, 0.0, 0.0], 1.0);
462        assert!((sim - 1.0).abs() < 0.001);
463    }
464
465    #[test]
466    fn cosine_similarity_orthogonal() {
467        let sim = cosine_similarity(&[1.0, 0.0], &[0.0, 1.0], 1.0);
468        assert!(sim.abs() < 0.001);
469    }
470
471    #[test]
472    fn cosine_similarity_45_degrees() {
473        let sim = cosine_similarity(&[1.0, 0.0], &[1.0, 1.0], 1.0);
474        assert!((sim - std::f32::consts::FRAC_1_SQRT_2).abs() < 0.01);
475    }
476
477    #[test]
478    fn vector_store_load_from_lmdb() {
479        let tmp = tempfile::tempdir().unwrap();
480        let store = MemoryStore::open_default(tmp.path()).unwrap();
481
482        // Create a memory and store its embedding
483        let mem = Memory::new(Galaxy::Codex, "test content".into());
484        store.put(Galaxy::Codex, &mem).unwrap();
485        store
486            .put_embedding(mem.metadata.id, &[0.1, 0.2, 0.3])
487            .unwrap();
488
489        // Load the vector store
490        let mut vs = VectorStore::new();
491        vs.load(&store).unwrap();
492        assert_eq!(vs.len(), 1);
493        assert!(vs.is_loaded());
494
495        // Search should find it
496        let results = vs.search(&[0.1, 0.2, 0.3], 10, None);
497        assert_eq!(results.len(), 1);
498        assert_eq!(results[0].memory_id, mem.metadata.id);
499    }
500
501    #[test]
502    fn vector_store_load_multiple_embeddings() {
503        let tmp = tempfile::tempdir().unwrap();
504        let store = MemoryStore::open_default(tmp.path()).unwrap();
505
506        // Store multiple memories with embeddings
507        for i in 0..5 {
508            let mem = Memory::new(Galaxy::Codex, format!("content {i}"));
509            store.put(Galaxy::Codex, &mem).unwrap();
510            let embedding = vec![i as f32 * 0.1, (i as f32).mul_add(-0.1, 1.0), 0.5];
511            store.put_embedding(mem.metadata.id, &embedding).unwrap();
512        }
513
514        let mut vs = VectorStore::new();
515        vs.load(&store).unwrap();
516        assert_eq!(vs.len(), 5);
517    }
518
519    #[test]
520    fn vector_store_load_empty() {
521        let tmp = tempfile::tempdir().unwrap();
522        let store = MemoryStore::open_default(tmp.path()).unwrap();
523
524        let mut vs = VectorStore::new();
525        vs.load(&store).unwrap();
526        assert_eq!(vs.len(), 0);
527        assert!(vs.is_loaded());
528    }
529}