1use std::collections::HashMap;
4use std::sync::Mutex;
5
6use async_trait::async_trait;
7use gonzalo_core::{CoreError, KeyPrefix, RecordKey, Result};
8
9use crate::{Match, VectorIndex};
10
11pub(crate) fn cosine(a: &[f32], b: &[f32]) -> f32 {
16 let dot: f32 = a.iter().zip(b.iter()).map(|(x, y)| x * y).sum();
17 let norm_a: f32 = a.iter().map(|x| x * x).sum::<f32>().sqrt();
18 let norm_b: f32 = b.iter().map(|x| x * x).sum::<f32>().sqrt();
19 if norm_a == 0.0
24 || norm_b == 0.0
25 || !dot.is_finite()
26 || !norm_a.is_finite()
27 || !norm_b.is_finite()
28 {
29 0.0
30 } else {
31 let score = dot / (norm_a * norm_b);
32 if score.is_finite() { score } else { 0.0 }
33 }
34}
35
36#[derive(Default)]
47pub struct MemoryVectorIndex {
48 store: Mutex<HashMap<RecordKey, Vec<f32>>>,
49}
50
51impl MemoryVectorIndex {
52 pub fn new() -> Self {
54 Self::default()
55 }
56
57 fn stored_dim(map: &HashMap<RecordKey, Vec<f32>>) -> Option<usize> {
60 map.values().next().map(|v| v.len())
61 }
62
63 fn check_dim(map: &HashMap<RecordKey, Vec<f32>>, incoming: usize) -> Result<()> {
64 if let Some(expected) = Self::stored_dim(map)
65 && incoming != expected
66 {
67 return Err(CoreError::Backend(format!(
68 "vector dimension mismatch: expected {expected}, got {incoming}"
69 )));
70 }
71 Ok(())
72 }
73}
74
75#[async_trait]
76impl VectorIndex for MemoryVectorIndex {
77 async fn upsert(&self, key: RecordKey, vector: Vec<f32>) -> Result<()> {
78 let mut map = self.store.lock().expect("mutex poisoned");
79 Self::check_dim(&map, vector.len())?;
80 map.insert(key, vector);
81 Ok(())
82 }
83
84 async fn remove(&self, key: &RecordKey) -> Result<()> {
85 let mut map = self.store.lock().expect("mutex poisoned");
86 map.remove(key);
87 Ok(())
88 }
89
90 async fn query(&self, query: &[f32], k: usize, filter: &KeyPrefix) -> Result<Vec<Match>> {
91 let map = self.store.lock().expect("mutex poisoned");
92 if !map.is_empty() {
93 Self::check_dim(&map, query.len())?;
94 }
95
96 let mut matches: Vec<Match> = map
97 .iter()
98 .filter(|(key, _)| filter.matches(key))
99 .map(|(key, vec)| Match {
100 key: key.clone(),
101 score: cosine(query, vec),
102 })
103 .collect();
104
105 matches.sort_by(|a, b| b.score.total_cmp(&a.score).then_with(|| a.key.cmp(&b.key)));
110
111 matches.truncate(k);
112 Ok(matches)
113 }
114}
115
116#[cfg(test)]
121mod tests {
122 use super::*;
123 use crate::Embedder;
124
125 struct BucketEmbedder;
128
129 #[async_trait]
130 impl Embedder for BucketEmbedder {
131 async fn embed(&self, text: &str) -> Result<Vec<f32>> {
132 let mut out = vec![0.0f32; 4];
133 for (i, b) in text.bytes().enumerate() {
134 out[i % 4] += b as f32;
135 }
136 Ok(out)
137 }
138 }
139
140 #[tokio::test]
144 async fn query_returns_nearest_first() {
145 let idx = MemoryVectorIndex::new();
146
147 idx.upsert(RecordKey::new("ns", "col", "east"), vec![1.0, 0.0])
149 .await
150 .unwrap();
151 idx.upsert(RecordKey::new("ns", "col", "north"), vec![0.0, 1.0])
152 .await
153 .unwrap();
154 idx.upsert(RecordKey::new("ns", "col", "neg"), vec![-1.0, 0.0])
155 .await
156 .unwrap();
157
158 let results = idx
160 .query(&[0.99, 0.01], 3, &KeyPrefix::default())
161 .await
162 .unwrap();
163
164 assert_eq!(results[0].key.id, "east");
165 }
166
167 #[tokio::test]
171 async fn k_limits_and_oversized_k() {
172 let idx = MemoryVectorIndex::new();
173 for i in 0..5u8 {
174 let v = vec![i as f32, 0.0];
175 idx.upsert(RecordKey::new("ns", "col", format!("{i}")), v)
176 .await
177 .unwrap();
178 }
179
180 let q = vec![1.0f32, 0.0];
182
183 let limited = idx.query(&q, 2, &KeyPrefix::default()).await.unwrap();
184 assert_eq!(limited.len(), 2);
185
186 let all = idx.query(&q, 100, &KeyPrefix::default()).await.unwrap();
187 assert_eq!(all.len(), 5);
188 }
189
190 #[tokio::test]
194 async fn filter_restricts_to_namespace() {
195 let idx = MemoryVectorIndex::new();
196
197 idx.upsert(RecordKey::new("alpha", "col", "a1"), vec![1.0, 0.0])
198 .await
199 .unwrap();
200 idx.upsert(RecordKey::new("alpha", "col", "a2"), vec![1.0, 0.1])
201 .await
202 .unwrap();
203 idx.upsert(RecordKey::new("beta", "col", "b1"), vec![0.0, 1.0])
204 .await
205 .unwrap();
206
207 let filter = KeyPrefix {
208 namespace: Some("alpha".into()),
209 collection: None,
210 };
211 let results = idx.query(&[1.0, 0.0], 10, &filter).await.unwrap();
212
213 assert_eq!(results.len(), 2);
214 assert!(results.iter().all(|m| m.key.namespace == "alpha"));
215 }
216
217 #[tokio::test]
221 async fn upsert_dimension_mismatch_is_error() {
222 let idx = MemoryVectorIndex::new();
223 idx.upsert(RecordKey::new("ns", "col", "a"), vec![1.0, 0.0])
224 .await
225 .unwrap();
226
227 let err = idx
228 .upsert(RecordKey::new("ns", "col", "b"), vec![1.0, 0.0, 0.0])
229 .await
230 .unwrap_err();
231
232 assert!(
233 matches!(err, CoreError::Backend(ref msg) if msg.contains("dimension mismatch")),
234 "unexpected error: {err}"
235 );
236 }
237
238 #[tokio::test]
242 async fn remove_drops_key() {
243 let idx = MemoryVectorIndex::new();
244 let key = RecordKey::new("ns", "col", "target");
245 idx.upsert(key.clone(), vec![1.0, 0.0]).await.unwrap();
246 idx.upsert(RecordKey::new("ns", "col", "other"), vec![0.0, 1.0])
247 .await
248 .unwrap();
249
250 idx.remove(&key).await.unwrap();
251
252 let results = idx
253 .query(&[1.0, 0.0], 10, &KeyPrefix::default())
254 .await
255 .unwrap();
256 assert!(results.iter().all(|m| m.key != key));
257 }
258
259 #[tokio::test]
263 async fn cosine_identical_direction_is_one() {
264 let idx = MemoryVectorIndex::new();
265 let v = vec![3.0f32, 4.0]; idx.upsert(RecordKey::new("ns", "col", "a"), v.clone())
267 .await
268 .unwrap();
269
270 let results = idx.query(&v, 1, &KeyPrefix::default()).await.unwrap();
271 assert!(
272 (results[0].score - 1.0).abs() < 1e-6,
273 "score={}",
274 results[0].score
275 );
276 }
277
278 #[tokio::test]
282 async fn embedder_trait_integration() {
283 let embedder = BucketEmbedder;
284 let idx = MemoryVectorIndex::new();
285
286 let texts = ["hello", "world", "rust"];
287 for text in &texts {
288 let vec = embedder.embed(text).await.unwrap();
289 idx.upsert(RecordKey::new("ns", "col", *text), vec)
290 .await
291 .unwrap();
292 }
293
294 let query_vec = embedder.embed("hello").await.unwrap();
296 let results = idx
297 .query(&query_vec, 3, &KeyPrefix::default())
298 .await
299 .unwrap();
300
301 assert!(!results.is_empty());
302 assert_eq!(results[0].key.id, "hello");
303 assert!(
305 (results[0].score - 1.0).abs() < 1e-6,
306 "score={}",
307 results[0].score
308 );
309 }
310
311 #[tokio::test]
315 async fn remove_absent_key_is_ok() {
316 let idx = MemoryVectorIndex::new();
317 let key = RecordKey::new("ns", "col", "ghost");
318 assert!(idx.remove(&key).await.is_ok());
319 }
320
321 #[test]
325 fn cosine_nan_component_scores_zero() {
326 assert_eq!(cosine(&[f32::NAN, 1.0], &[1.0, 0.0]), 0.0);
328 assert_eq!(cosine(&[1.0, 0.0], &[f32::NAN, 1.0]), 0.0);
329 assert_eq!(cosine(&[f32::INFINITY, 1.0], &[1.0, 0.0]), 0.0);
331 assert_eq!(cosine(&[1.0, 0.0], &[f32::NEG_INFINITY, 1.0]), 0.0);
332 assert_eq!(cosine(&[0.0, 0.0], &[1.0, 0.0]), 0.0);
334 assert!((cosine(&[3.0, 4.0], &[3.0, 4.0]) - 1.0).abs() < 1e-6);
336 }
337
338 #[tokio::test]
342 async fn nan_vector_never_occupies_top_k_ahead_of_match() {
343 let idx = MemoryVectorIndex::new();
344
345 idx.upsert(RecordKey::new("ns", "col", "good"), vec![1.0, 0.0])
347 .await
348 .unwrap();
349 idx.upsert(RecordKey::new("ns", "col", "poison"), vec![f32::NAN, 1.0])
351 .await
352 .unwrap();
353
354 let results = idx
355 .query(&[1.0, 0.0], 2, &KeyPrefix::default())
356 .await
357 .unwrap();
358
359 assert_eq!(results.len(), 2);
360 assert_eq!(results[0].key.id, "good");
362 assert_eq!(results[1].key.id, "poison");
363 assert!(results.iter().all(|m| m.score.is_finite()));
364 assert_eq!(results[1].score, 0.0);
365 }
366
367 #[tokio::test]
371 async fn ranking_is_deterministic_with_poisoned_vectors() {
372 let idx = MemoryVectorIndex::new();
373 idx.upsert(RecordKey::new("ns", "col", "a"), vec![1.0, 0.0])
374 .await
375 .unwrap();
376 idx.upsert(RecordKey::new("ns", "col", "b"), vec![f32::NAN, 1.0])
377 .await
378 .unwrap();
379 idx.upsert(RecordKey::new("ns", "col", "c"), vec![f32::INFINITY, 0.0])
380 .await
381 .unwrap();
382 idx.upsert(RecordKey::new("ns", "col", "d"), vec![0.9, 0.1])
383 .await
384 .unwrap();
385
386 let first = idx
387 .query(&[1.0, 0.0], 4, &KeyPrefix::default())
388 .await
389 .unwrap();
390 for _ in 0..10 {
391 let again = idx
392 .query(&[1.0, 0.0], 4, &KeyPrefix::default())
393 .await
394 .unwrap();
395 let ids_first: Vec<_> = first.iter().map(|m| m.key.id.clone()).collect();
396 let ids_again: Vec<_> = again.iter().map(|m| m.key.id.clone()).collect();
397 assert_eq!(ids_first, ids_again, "ordering must be deterministic");
398 }
399 }
400
401 #[tokio::test]
405 async fn query_dimension_mismatch_is_error() {
406 let idx = MemoryVectorIndex::new();
407 idx.upsert(RecordKey::new("ns", "col", "a"), vec![1.0, 0.0])
408 .await
409 .unwrap();
410
411 let err = idx
412 .query(&[1.0, 0.0, 0.0], 1, &KeyPrefix::default())
413 .await
414 .unwrap_err();
415
416 assert!(
417 matches!(err, CoreError::Backend(ref msg) if msg.contains("dimension mismatch")),
418 "unexpected error: {err}"
419 );
420 }
421}