1use ahash::AHashMap;
12use lmdb::{Cursor, Transaction};
13use uuid::Uuid;
14use wm_core::{CoreError, Galaxy, Result};
15
16use crate::MemoryStore;
17
18#[derive(Debug, Clone)]
20pub struct VectorSearchResult {
21 pub memory_id: Uuid,
23 pub galaxy: Galaxy,
25 pub score: f32,
27}
28
29pub trait VectorSearchEngine: Send + Sync {
35 fn add_vector(&mut self, memory_id: Uuid, galaxy: Galaxy, embedding: Vec<f32>);
37
38 fn remove_vector(&mut self, memory_id: Uuid) -> bool;
40
41 fn search_vectors(
46 &self,
47 query: &[f32],
48 limit: usize,
49 galaxy_filter: Option<Galaxy>,
50 ) -> Vec<VectorSearchResult>;
51
52 fn search_similar_vectors(&self, memory_id: Uuid, limit: usize) -> Vec<VectorSearchResult>;
56
57 fn vector_count(&self) -> usize;
59
60 fn is_index_empty(&self) -> bool {
62 self.vector_count() == 0
63 }
64
65 fn load_vectors(&mut self, store: &MemoryStore) -> Result<()>;
67
68 fn clear_vectors(&mut self);
70}
71
72pub struct VectorStore {
77 vectors: AHashMap<Uuid, (Galaxy, Vec<f32>)>,
79 loaded: bool,
81}
82
83impl VectorStore {
84 #[must_use]
86 pub fn new() -> Self {
87 Self {
88 vectors: AHashMap::new(),
89 loaded: false,
90 }
91 }
92
93 pub fn load(&mut self, store: &MemoryStore) -> Result<()> {
99 let db = store.galaxy_db(Galaxy::Embeddings)?;
100
101 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 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 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 pub fn add(&mut self, memory_id: Uuid, galaxy: Galaxy, embedding: Vec<f32>) {
164 self.vectors.insert(memory_id, (galaxy, embedding));
165 }
166
167 pub fn remove(&mut self, memory_id: Uuid) -> bool {
169 self.vectors.remove(&memory_id).is_some()
170 }
171
172 #[must_use]
174 pub fn len(&self) -> usize {
175 self.vectors.len()
176 }
177
178 #[must_use]
180 pub fn is_empty(&self) -> bool {
181 self.vectors.is_empty()
182 }
183
184 #[must_use]
186 pub const fn is_loaded(&self) -> bool {
187 self.loaded
188 }
189
190 #[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 #[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; 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 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
328fn vector_norm(v: &[f32]) -> f32 {
330 v.iter().map(|x| x * x).sum::<f32>().sqrt()
331}
332
333fn 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); assert_eq!(results[0].memory_id, id1);
378 assert!((results[0].score - 1.0).abs() < 0.001); }
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 assert_eq!(results.len(), 1);
417 assert!(results.iter().all(|r| r.memory_id != id1));
418 assert_eq!(results[0].memory_id, id2); }
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()); }
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 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 let mut vs = VectorStore::new();
491 vs.load(&store).unwrap();
492 assert_eq!(vs.len(), 1);
493 assert!(vs.is_loaded());
494
495 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 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}