Skip to main content

behest_store/memory/
embedding.rs

1//! In-memory embedding store with brute-force cosine similarity search
2//! over all stored vectors.
3
4use std::collections::HashMap;
5
6use async_trait::async_trait;
7use tokio::sync::RwLock;
8use uuid::Uuid;
9
10use crate::{EmbeddingRecord, EmbeddingStore, ScoredEmbedding, StoreResult};
11
12/// In-memory embedding store for testing, development, and prototyping.
13///
14/// Uses brute-force O(n) cosine similarity scan for nearest-neighbor search.
15/// Data is lost when the process exits. Implements [`EmbeddingStore`].
16#[derive(Default)]
17pub struct MemoryEmbeddingStore {
18    records: RwLock<HashMap<Uuid, EmbeddingRecord>>,
19}
20
21impl MemoryEmbeddingStore {
22    /// Creates an empty in-memory embedding store.
23    #[must_use]
24    pub fn new() -> Self {
25        Self::default()
26    }
27}
28
29/// Computes the cosine similarity between two vectors of equal length.
30///
31/// Returns a value in `[-1.0, 1.0]` where `1.0` indicates identical direction,
32/// `0.0` indicates orthogonality, and `-1.0` indicates opposite direction.
33/// Returns `0.0` if the vectors have different lengths or either is zero-vector.
34fn cosine_similarity(a: &[f32], b: &[f32]) -> f32 {
35    if a.len() != b.len() || a.is_empty() {
36        return 0.0;
37    }
38
39    let mut dot = 0.0_f32;
40    let mut norm_a = 0.0_f32;
41    let mut norm_b = 0.0_f32;
42
43    for (x, y) in a.iter().zip(b.iter()) {
44        dot += x * y;
45        norm_a += x * x;
46        norm_b += y * y;
47    }
48
49    let denom = (norm_a.sqrt()) * (norm_b.sqrt());
50    if denom == 0.0 { 0.0 } else { dot / denom }
51}
52
53#[async_trait]
54impl EmbeddingStore for MemoryEmbeddingStore {
55    async fn upsert(&self, record: EmbeddingRecord) -> StoreResult<EmbeddingRecord> {
56        let mut records = self.records.write().await;
57        let id = record.id;
58        records.insert(id, record.clone());
59        Ok(record)
60    }
61
62    async fn search(&self, query: &[f32], limit: usize) -> StoreResult<Vec<ScoredEmbedding>> {
63        let records = self.records.read().await;
64
65        let mut scored: Vec<ScoredEmbedding> = records
66            .values()
67            .map(|record| ScoredEmbedding {
68                score: cosine_similarity(query, &record.vector),
69                record: record.clone(),
70            })
71            .collect();
72
73        scored.sort_by(|a, b| {
74            b.score
75                .partial_cmp(&a.score)
76                .unwrap_or(std::cmp::Ordering::Equal)
77        });
78        scored.truncate(limit);
79
80        Ok(scored)
81    }
82
83    async fn delete(&self, id: &Uuid) -> StoreResult<()> {
84        self.records.write().await.remove(id);
85        Ok(())
86    }
87
88    async fn delete_by_session(&self, session_id: &Uuid) -> StoreResult<u64> {
89        let mut records = self.records.write().await;
90        let before = records.len();
91        records.retain(|_, record| record.session_id != Some(*session_id));
92        let deleted = before - records.len();
93        Ok(deleted as u64)
94    }
95}
96
97#[cfg(test)]
98#[allow(clippy::unwrap_used)]
99mod tests {
100    use super::*;
101    use serde_json::json;
102
103    fn test_record(vector: Vec<f32>) -> EmbeddingRecord {
104        EmbeddingRecord::new("test-model", vector)
105    }
106
107    #[tokio::test]
108    async fn memory_embedding_store_should_upsert_and_search() {
109        let store = MemoryEmbeddingStore::new();
110
111        store
112            .upsert(test_record(vec![1.0, 0.0, 0.0]))
113            .await
114            .unwrap();
115        store
116            .upsert(test_record(vec![0.0, 1.0, 0.0]))
117            .await
118            .unwrap();
119        store
120            .upsert(test_record(vec![0.0, 0.0, 1.0]))
121            .await
122            .unwrap();
123
124        let results = store.search(&[1.0, 0.0, 0.0], 2).await.unwrap();
125        assert_eq!(results.len(), 2);
126        assert!((results[0].score - 1.0).abs() < f32::EPSILON);
127        assert!(results[0].score >= results[1].score);
128    }
129
130    #[tokio::test]
131    async fn memory_embedding_store_should_upsert_existing_record() {
132        let store = MemoryEmbeddingStore::new();
133
134        let record = test_record(vec![1.0, 0.0]);
135        let id = record.id;
136        store.upsert(record).await.unwrap();
137
138        let updated = EmbeddingRecord {
139            id,
140            session_id: None,
141            model: "updated-model".to_owned(),
142            vector: vec![0.0, 1.0],
143            metadata: json!({"updated": true}),
144            created_at: chrono::Utc::now(),
145        };
146        store.upsert(updated).await.unwrap();
147
148        let results = store.search(&[0.0, 1.0], 1).await.unwrap();
149        assert_eq!(results[0].record.model, "updated-model");
150    }
151
152    #[tokio::test]
153    async fn memory_embedding_store_should_delete_by_id() {
154        let store = MemoryEmbeddingStore::new();
155
156        let record = test_record(vec![1.0, 0.0]);
157        let id = record.id;
158        store.upsert(record).await.unwrap();
159        store.delete(&id).await.unwrap();
160
161        let results = store.search(&[1.0, 0.0], 10).await.unwrap();
162        assert!(results.is_empty());
163    }
164
165    #[tokio::test]
166    async fn memory_embedding_store_should_delete_by_session() {
167        let store = MemoryEmbeddingStore::new();
168        let session_id = Uuid::now_v7();
169
170        store
171            .upsert(test_record(vec![1.0]).with_session(session_id))
172            .await
173            .unwrap();
174        store
175            .upsert(test_record(vec![0.0, 1.0]).with_session(session_id))
176            .await
177            .unwrap();
178        store
179            .upsert(test_record(vec![0.0, 0.0, 1.0]))
180            .await
181            .unwrap();
182
183        let deleted = store.delete_by_session(&session_id).await.unwrap();
184        assert_eq!(deleted, 2);
185
186        let remaining = store.search(&[1.0], 10).await.unwrap();
187        assert_eq!(remaining.len(), 1);
188    }
189
190    #[tokio::test]
191    async fn memory_embedding_store_should_handle_empty_search() {
192        let store = MemoryEmbeddingStore::new();
193        let results = store.search(&[1.0, 0.0], 5).await.unwrap();
194        assert!(results.is_empty());
195    }
196
197    #[test]
198    fn cosine_similarity_should_return_one_for_identical_vectors() {
199        let v = vec![1.0, 2.0, 3.0];
200        assert!((cosine_similarity(&v, &v) - 1.0).abs() < f32::EPSILON);
201    }
202
203    #[test]
204    fn cosine_similarity_should_return_zero_for_orthogonal_vectors() {
205        let a = vec![1.0, 0.0];
206        let b = vec![0.0, 1.0];
207        assert!(cosine_similarity(&a, &b).abs() < f32::EPSILON);
208    }
209
210    #[test]
211    fn cosine_similarity_should_handle_different_lengths() {
212        assert!((cosine_similarity(&[1.0], &[1.0, 2.0])).abs() < f32::EPSILON);
213    }
214
215    #[test]
216    fn cosine_similarity_should_handle_zero_vectors() {
217        assert!((cosine_similarity(&[0.0, 0.0], &[1.0, 0.0])).abs() < f32::EPSILON);
218    }
219}