behest_store/memory/
embedding.rs1use 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#[derive(Default)]
17pub struct MemoryEmbeddingStore {
18 records: RwLock<HashMap<Uuid, EmbeddingRecord>>,
19}
20
21impl MemoryEmbeddingStore {
22 #[must_use]
24 pub fn new() -> Self {
25 Self::default()
26 }
27}
28
29fn 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}