use ahash::AHashMap;
use lmdb::{Cursor, Transaction};
use uuid::Uuid;
use wm_core::{CoreError, Galaxy, Result};
use crate::MemoryStore;
#[derive(Debug, Clone)]
pub struct VectorSearchResult {
pub memory_id: Uuid,
pub galaxy: Galaxy,
pub score: f32,
}
pub trait VectorSearchEngine: Send + Sync {
fn add_vector(&mut self, memory_id: Uuid, galaxy: Galaxy, embedding: Vec<f32>);
fn remove_vector(&mut self, memory_id: Uuid) -> bool;
fn search_vectors(
&self,
query: &[f32],
limit: usize,
galaxy_filter: Option<Galaxy>,
) -> Vec<VectorSearchResult>;
fn search_similar_vectors(&self, memory_id: Uuid, limit: usize) -> Vec<VectorSearchResult>;
fn vector_count(&self) -> usize;
fn is_index_empty(&self) -> bool {
self.vector_count() == 0
}
fn load_vectors(&mut self, store: &MemoryStore) -> Result<()>;
fn clear_vectors(&mut self);
}
pub struct VectorStore {
vectors: AHashMap<Uuid, (Galaxy, Vec<f32>)>,
loaded: bool,
}
impl VectorStore {
#[must_use]
pub fn new() -> Self {
Self {
vectors: AHashMap::new(),
loaded: false,
}
}
pub fn load(&mut self, store: &MemoryStore) -> Result<()> {
let db = store.galaxy_db(Galaxy::Embeddings)?;
let mut entries: Vec<(Uuid, Vec<f32>)> = Vec::new();
{
let tx = store
.env()
.begin_ro_txn()
.map_err(|e| CoreError::Memory(format!("LMDB ro_txn failed: {e}")))?;
let mut cursor = tx
.open_ro_cursor(db)
.map_err(|e| CoreError::Memory(format!("LMDB cursor failed: {e}")))?;
for (key, val) in cursor.iter() {
if key.len() == 16 {
let bytes: [u8; 16] = key.try_into().unwrap_or([0u8; 16]);
let id = Uuid::from_bytes(bytes);
let embedding = crate::memory::decode_embedding(val);
entries.push((id, embedding));
}
}
drop(cursor);
tx.commit()
.map_err(|e| CoreError::Memory(format!("LMDB commit failed: {e}")))?;
}
let mut count = 0;
for (id, embedding) in entries {
match self.find_memory_galaxy(store, id) {
Some(galaxy) => {
self.vectors.insert(id, (galaxy, embedding));
count += 1;
}
None => {
tracing::warn!(
"Skipping orphaned embedding (memory not found in any galaxy, id={})",
id
);
}
}
}
self.loaded = true;
tracing::info!("Loaded {count} embedding vectors into VectorStore");
Ok(())
}
fn find_memory_galaxy(&self, store: &MemoryStore, id: Uuid) -> Option<Galaxy> {
for galaxy in Galaxy::all() {
if galaxy == Galaxy::Embeddings {
continue;
}
if store.get(galaxy, id).ok().flatten().is_some() {
return Some(galaxy);
}
}
None
}
pub fn add(&mut self, memory_id: Uuid, galaxy: Galaxy, embedding: Vec<f32>) {
self.vectors.insert(memory_id, (galaxy, embedding));
}
pub fn remove(&mut self, memory_id: Uuid) -> bool {
self.vectors.remove(&memory_id).is_some()
}
#[must_use]
pub fn len(&self) -> usize {
self.vectors.len()
}
#[must_use]
pub fn is_empty(&self) -> bool {
self.vectors.is_empty()
}
#[must_use]
pub const fn is_loaded(&self) -> bool {
self.loaded
}
#[must_use]
pub fn search(
&self,
query: &[f32],
limit: usize,
galaxy_filter: Option<Galaxy>,
) -> Vec<VectorSearchResult> {
if self.vectors.is_empty() || query.is_empty() {
return Vec::new();
}
let query_norm = vector_norm(query);
if query_norm == 0.0 {
return Vec::new();
}
let mut results: Vec<VectorSearchResult> = self
.vectors
.iter()
.filter(|(_, (galaxy, _))| galaxy_filter.is_none_or(|g| g == *galaxy))
.filter_map(|(id, (galaxy, embedding))| {
let score = cosine_similarity(query, embedding, query_norm);
if score > 0.0 {
Some(VectorSearchResult {
memory_id: *id,
galaxy: *galaxy,
score,
})
} else {
None
}
})
.collect();
results.sort_by(|a, b| {
b.score
.partial_cmp(&a.score)
.unwrap_or(std::cmp::Ordering::Equal)
});
results.truncate(limit);
results
}
#[must_use]
pub fn search_similar_to(&self, memory_id: Uuid, limit: usize) -> Vec<VectorSearchResult> {
let (galaxy, embedding) = match self.vectors.get(&memory_id) {
Some(v) => v,
None => return Vec::new(),
};
let query_norm = vector_norm(embedding);
if query_norm == 0.0 {
return Vec::new();
}
let mut results: Vec<VectorSearchResult> = self
.vectors
.iter()
.filter(|(id, _)| **id != memory_id)
.filter_map(|(id, (g, emb))| {
let score = cosine_similarity(embedding, emb, query_norm);
if score > 0.0 {
Some(VectorSearchResult {
memory_id: *id,
galaxy: *g,
score,
})
} else {
None
}
})
.collect();
let _ = galaxy; results.sort_by(|a, b| {
b.score
.partial_cmp(&a.score)
.unwrap_or(std::cmp::Ordering::Equal)
});
results.truncate(limit);
results
}
pub fn clear(&mut self) {
self.vectors.clear();
self.loaded = false;
}
}
impl Default for VectorStore {
fn default() -> Self {
Self::new()
}
}
impl VectorSearchEngine for VectorStore {
fn add_vector(&mut self, memory_id: Uuid, galaxy: Galaxy, embedding: Vec<f32>) {
self.add(memory_id, galaxy, embedding);
}
fn remove_vector(&mut self, memory_id: Uuid) -> bool {
self.remove(memory_id)
}
fn search_vectors(
&self,
query: &[f32],
limit: usize,
galaxy_filter: Option<Galaxy>,
) -> Vec<VectorSearchResult> {
self.search(query, limit, galaxy_filter)
}
fn search_similar_vectors(&self, memory_id: Uuid, limit: usize) -> Vec<VectorSearchResult> {
self.search_similar_to(memory_id, limit)
}
fn vector_count(&self) -> usize {
self.len()
}
fn load_vectors(&mut self, store: &MemoryStore) -> Result<()> {
self.load(store)
}
fn clear_vectors(&mut self) {
self.clear();
}
}
fn vector_norm(v: &[f32]) -> f32 {
v.iter().map(|x| x * x).sum::<f32>().sqrt()
}
fn cosine_similarity(query: &[f32], target: &[f32], query_norm: f32) -> f32 {
if query.len() != target.len() {
return 0.0;
}
let dot: f32 = query.iter().zip(target.iter()).map(|(a, b)| a * b).sum();
let target_norm = vector_norm(target);
if target_norm == 0.0 {
return 0.0;
}
dot / (query_norm * target_norm)
}
#[cfg(test)]
mod tests {
use super::*;
use crate::Memory;
#[test]
fn vector_store_empty_search() {
let vs = VectorStore::new();
let results = vs.search(&[1.0, 0.0, 0.0], 10, None);
assert!(results.is_empty());
}
#[test]
fn vector_store_add_and_search() {
let mut vs = VectorStore::new();
let id1 = Uuid::new_v4();
let id2 = Uuid::new_v4();
let id3 = Uuid::new_v4();
vs.add(id1, Galaxy::Codex, vec![1.0, 0.0, 0.0]);
vs.add(id2, Galaxy::Codex, vec![0.0, 1.0, 0.0]);
vs.add(id3, Galaxy::Codex, vec![1.0, 1.0, 0.0]);
let results = vs.search(&[1.0, 0.0, 0.0], 10, None);
assert_eq!(results.len(), 2); assert_eq!(results[0].memory_id, id1);
assert!((results[0].score - 1.0).abs() < 0.001); }
#[test]
fn vector_store_search_with_limit() {
let mut vs = VectorStore::new();
for _ in 0..10 {
vs.add(Uuid::new_v4(), Galaxy::Codex, vec![1.0, 0.0, 0.0]);
}
let results = vs.search(&[1.0, 0.0, 0.0], 3, None);
assert_eq!(results.len(), 3);
}
#[test]
fn vector_store_galaxy_filter() {
let mut vs = VectorStore::new();
vs.add(Uuid::new_v4(), Galaxy::Codex, vec![1.0, 0.0]);
vs.add(Uuid::new_v4(), Galaxy::Research, vec![1.0, 0.0]);
vs.add(Uuid::new_v4(), Galaxy::Codex, vec![0.9, 0.1]);
let results = vs.search(&[1.0, 0.0], 10, Some(Galaxy::Codex));
assert_eq!(results.len(), 2);
assert!(results.iter().all(|r| r.galaxy == Galaxy::Codex));
}
#[test]
fn vector_store_search_similar_to() {
let mut vs = VectorStore::new();
let id1 = Uuid::new_v4();
let id2 = Uuid::new_v4();
let id3 = Uuid::new_v4();
vs.add(id1, Galaxy::Codex, vec![1.0, 0.0, 0.0]);
vs.add(id2, Galaxy::Codex, vec![0.95, 0.05, 0.0]);
vs.add(id3, Galaxy::Codex, vec![0.0, 1.0, 0.0]);
let results = vs.search_similar_to(id1, 10);
assert_eq!(results.len(), 1);
assert!(results.iter().all(|r| r.memory_id != id1));
assert_eq!(results[0].memory_id, id2); }
#[test]
fn vector_store_remove() {
let mut vs = VectorStore::new();
let id = Uuid::new_v4();
vs.add(id, Galaxy::Codex, vec![1.0, 0.0]);
assert_eq!(vs.len(), 1);
assert!(vs.remove(id));
assert_eq!(vs.len(), 0);
assert!(!vs.remove(id));
}
#[test]
fn vector_store_clear() {
let mut vs = VectorStore::new();
vs.add(Uuid::new_v4(), Galaxy::Codex, vec![1.0]);
vs.add(Uuid::new_v4(), Galaxy::Codex, vec![1.0]);
assert_eq!(vs.len(), 2);
vs.clear();
assert_eq!(vs.len(), 0);
assert!(!vs.is_loaded());
}
#[test]
fn vector_store_zero_query_returns_empty() {
let mut vs = VectorStore::new();
vs.add(Uuid::new_v4(), Galaxy::Codex, vec![1.0, 0.0]);
let results = vs.search(&[0.0, 0.0], 10, None);
assert!(results.is_empty());
}
#[test]
fn vector_store_mismatched_dimensions() {
let mut vs = VectorStore::new();
vs.add(Uuid::new_v4(), Galaxy::Codex, vec![1.0, 0.0, 0.0]);
let results = vs.search(&[1.0, 0.0], 10, None);
assert!(results.is_empty()); }
#[test]
fn cosine_similarity_exact_match() {
let sim = cosine_similarity(&[1.0, 0.0, 0.0], &[1.0, 0.0, 0.0], 1.0);
assert!((sim - 1.0).abs() < 0.001);
}
#[test]
fn cosine_similarity_orthogonal() {
let sim = cosine_similarity(&[1.0, 0.0], &[0.0, 1.0], 1.0);
assert!(sim.abs() < 0.001);
}
#[test]
fn cosine_similarity_45_degrees() {
let sim = cosine_similarity(&[1.0, 0.0], &[1.0, 1.0], 1.0);
assert!((sim - std::f32::consts::FRAC_1_SQRT_2).abs() < 0.01);
}
#[test]
fn vector_store_load_from_lmdb() {
let tmp = tempfile::tempdir().unwrap();
let store = MemoryStore::open_default(tmp.path()).unwrap();
let mem = Memory::new(Galaxy::Codex, "test content".into());
store.put(Galaxy::Codex, &mem).unwrap();
store
.put_embedding(mem.metadata.id, &[0.1, 0.2, 0.3])
.unwrap();
let mut vs = VectorStore::new();
vs.load(&store).unwrap();
assert_eq!(vs.len(), 1);
assert!(vs.is_loaded());
let results = vs.search(&[0.1, 0.2, 0.3], 10, None);
assert_eq!(results.len(), 1);
assert_eq!(results[0].memory_id, mem.metadata.id);
}
#[test]
fn vector_store_load_multiple_embeddings() {
let tmp = tempfile::tempdir().unwrap();
let store = MemoryStore::open_default(tmp.path()).unwrap();
for i in 0..5 {
let mem = Memory::new(Galaxy::Codex, format!("content {i}"));
store.put(Galaxy::Codex, &mem).unwrap();
let embedding = vec![i as f32 * 0.1, (i as f32).mul_add(-0.1, 1.0), 0.5];
store.put_embedding(mem.metadata.id, &embedding).unwrap();
}
let mut vs = VectorStore::new();
vs.load(&store).unwrap();
assert_eq!(vs.len(), 5);
}
#[test]
fn vector_store_load_empty() {
let tmp = tempfile::tempdir().unwrap();
let store = MemoryStore::open_default(tmp.path()).unwrap();
let mut vs = VectorStore::new();
vs.load(&store).unwrap();
assert_eq!(vs.len(), 0);
assert!(vs.is_loaded());
}
}