Skip to main content

gonzalo_vector/
index.rs

1//! In-memory exact cosine vector index keyed by [`RecordKey`].
2
3use 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
11// ---------------------------------------------------------------------------
12// Cosine helper
13// ---------------------------------------------------------------------------
14
15pub(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    // Guard against a zero norm (undefined direction) and any non-finite input
20    // (a NaN/inf component propagates into `dot`/`norm_*`). A non-finite result
21    // would otherwise scatter arbitrarily during ranking, so score it 0.0 —
22    // ranked last, never NaN.
23    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// ---------------------------------------------------------------------------
37// MemoryVectorIndex
38// ---------------------------------------------------------------------------
39
40/// An exact, brute-force in-memory vector index.
41///
42/// All vectors must share the same dimension once the first entry is inserted.
43/// Cosine similarity is used for scoring; results are ordered by descending
44/// score with ties broken by [`RecordKey`] ordering (lexicographic), making
45/// results deterministic.
46#[derive(Default)]
47pub struct MemoryVectorIndex {
48    store: Mutex<HashMap<RecordKey, Vec<f32>>>,
49}
50
51impl MemoryVectorIndex {
52    /// Create an empty index.
53    pub fn new() -> Self {
54        Self::default()
55    }
56
57    /// Return the current dimension of stored vectors, or `None` if the index
58    /// is empty.
59    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        // Sort descending by score; break ties by RecordKey order (ascending).
106        // Use `f32::total_cmp` so ordering is total and deterministic even if a
107        // non-finite score (NaN/inf) slips through — `partial_cmp` returns
108        // `None` for NaN and would scatter such entries arbitrarily.
109        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// ---------------------------------------------------------------------------
117// Tests
118// ---------------------------------------------------------------------------
119
120#[cfg(test)]
121mod tests {
122    use super::*;
123    use crate::Embedder;
124
125    /// Tiny deterministic embedder: maps a string to a 4-element Vec<f32> by
126    /// bucketing each byte value into one of four bins and summing them.
127    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    // ------------------------------------------------------------------
141    // Test 1: nearest-first ordering
142    // ------------------------------------------------------------------
143    #[tokio::test]
144    async fn query_returns_nearest_first() {
145        let idx = MemoryVectorIndex::new();
146
147        // Three 2-D vectors pointing in clearly different directions.
148        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        // Query pointing almost east.
159        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    // ------------------------------------------------------------------
168    // Test 2: k limits results; k > index size returns all
169    // ------------------------------------------------------------------
170    #[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        // Non-zero query so we don't divide by zero on the [0.0, 0.0] vector.
181        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    // ------------------------------------------------------------------
191    // Test 3: filter by namespace
192    // ------------------------------------------------------------------
193    #[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    // ------------------------------------------------------------------
218    // Test 4: dimension mismatch returns Backend error on upsert
219    // ------------------------------------------------------------------
220    #[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    // ------------------------------------------------------------------
239    // Test 5: remove drops key from results
240    // ------------------------------------------------------------------
241    #[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    // ------------------------------------------------------------------
260    // Test 6: cosine of identical direction is ~1.0
261    // ------------------------------------------------------------------
262    #[tokio::test]
263    async fn cosine_identical_direction_is_one() {
264        let idx = MemoryVectorIndex::new();
265        let v = vec![3.0f32, 4.0]; // magnitude 5
266        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    // ------------------------------------------------------------------
279    // Test 7: Embedder trait integration
280    // ------------------------------------------------------------------
281    #[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        // Query with the same embedding as "hello" — it should be the top hit.
295        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        // Self-similarity must be ~1.0.
304        assert!(
305            (results[0].score - 1.0).abs() < 1e-6,
306            "score={}",
307            results[0].score
308        );
309    }
310
311    // ------------------------------------------------------------------
312    // Bonus: remove absent key does not error
313    // ------------------------------------------------------------------
314    #[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    // ------------------------------------------------------------------
322    // #154: non-finite components score 0.0, never NaN
323    // ------------------------------------------------------------------
324    #[test]
325    fn cosine_nan_component_scores_zero() {
326        // A NaN in either operand must score 0.0, not NaN.
327        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        // An infinity likewise scores 0.0 rather than a non-finite value.
330        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        // The pre-existing zero-norm guard still holds.
333        assert_eq!(cosine(&[0.0, 0.0], &[1.0, 0.0]), 0.0);
334        // A genuine match still scores ~1.0.
335        assert!((cosine(&[3.0, 4.0], &[3.0, 4.0]) - 1.0).abs() < 1e-6);
336    }
337
338    // ------------------------------------------------------------------
339    // #154: a NaN-bearing vector never outranks a genuinely similar one
340    // ------------------------------------------------------------------
341    #[tokio::test]
342    async fn nan_vector_never_occupies_top_k_ahead_of_match() {
343        let idx = MemoryVectorIndex::new();
344
345        // A vector aligned with the query (should rank first) ...
346        idx.upsert(RecordKey::new("ns", "col", "good"), vec![1.0, 0.0])
347            .await
348            .unwrap();
349        // ... and a poisoned vector with a NaN component.
350        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        // The genuine match ranks first; the poisoned vector scores 0.0 last.
361        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    // ------------------------------------------------------------------
368    // #154: ranking is deterministic across repeated queries
369    // ------------------------------------------------------------------
370    #[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    // ------------------------------------------------------------------
402    // Bonus: query dimension mismatch is error
403    // ------------------------------------------------------------------
404    #[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}