1use std::collections::BTreeMap;
14use std::sync::Arc;
15
16use uqa_core::{DocId, Payload, PostingEntry, PostingList};
17
18use crate::{StorageBackendError, StorageBackendResult};
19
20mod config;
21
22pub use config::{HNSWIndexParams, IVFIndexParams, VectorIndexOpenMode, VectorIndexSpec};
23
24pub fn validate_vector_values(dimensions: u32, vector: &[f32]) -> StorageBackendResult<()> {
25 let dimensions = usize::try_from(dimensions).map_err(|_| {
26 StorageBackendError::Other(format!(
27 "vector dimension {dimensions} exceeds the platform usize range"
28 ))
29 })?;
30 if vector.len() != dimensions {
31 return Err(StorageBackendError::Other(format!(
32 "vector dimension mismatch: expected {dimensions}, got {}",
33 vector.len()
34 )));
35 }
36 if let Some((index, value)) = vector
37 .iter()
38 .copied()
39 .enumerate()
40 .find(|(_, value)| !value.is_finite())
41 {
42 return Err(StorageBackendError::Other(format!(
43 "vector component {index} must be finite, got {value}"
44 )));
45 }
46 Ok(())
47}
48
49fn checked_vector_count(counts: impl IntoIterator<Item = usize>) -> StorageBackendResult<usize> {
50 counts.into_iter().try_fold(0_usize, |total, count| {
51 total
52 .checked_add(count)
53 .ok_or_else(|| StorageBackendError::Other("vector count overflow".into()))
54 })
55}
56
57fn validate_threshold(threshold: f32) -> StorageBackendResult<()> {
58 if threshold.is_finite() {
59 Ok(())
60 } else {
61 Err(StorageBackendError::Other(format!(
62 "vector similarity threshold must be finite, got {threshold}"
63 )))
64 }
65}
66
67pub fn select_top_k_scored(scored: &mut Vec<(DocId, f32)>, k: usize) {
68 if scored.len() > k {
69 scored.select_nth_unstable_by(k, |a, b| b.1.total_cmp(&a.1).then_with(|| a.0.cmp(&b.0)));
70 scored.truncate(k);
71 }
72}
73
74pub(crate) fn deduplicate_scored_by_doc(scored: &mut Vec<(DocId, f32)>) {
78 if scored.len() < 2 {
79 return;
80 }
81 scored.sort_unstable_by_key(|(doc_id, _)| *doc_id);
82 let mut write = 1;
83 for read in 1..scored.len() {
84 let (doc_id, score) = scored[read];
85 if scored[write - 1].0 == doc_id {
86 scored[write - 1].1 = scored[write - 1].1.max(score);
87 } else {
88 scored[write] = (doc_id, score);
89 write += 1;
90 }
91 }
92 scored.truncate(write);
93}
94
95pub(crate) fn vector_norm(vector: &[f32]) -> f32 {
96 let mut squared_norm = 0.0_f32;
97 for value in vector {
98 squared_norm += value * value;
99 }
100 squared_norm.sqrt()
101}
102
103pub(crate) fn cosine_similarity_with_norms(a: &[f32], b: &[f32], norm_a: f32, norm_b: f32) -> f32 {
107 if a.len() != b.len() || a.is_empty() || norm_a == 0.0 || norm_b == 0.0 {
108 return 0.0;
109 }
110 let mut dot = 0.0_f32;
111 for (x, y) in a.iter().zip(b) {
112 dot += x * y;
113 }
114 dot / (norm_a * norm_b)
115}
116
117pub fn cosine_similarity(a: &[f32], b: &[f32]) -> f32 {
124 if a.len() != b.len() || a.is_empty() {
125 return 0.0;
126 }
127 let mut dot = 0.0f32;
128 let mut norm_a = 0.0f32;
129 let mut norm_b = 0.0f32;
130 for (x, y) in a.iter().zip(b.iter()) {
131 dot += x * y;
132 norm_a += x * x;
133 norm_b += y * y;
134 }
135 if norm_a == 0.0 || norm_b == 0.0 {
136 return 0.0;
137 }
138 dot / (norm_a.sqrt() * norm_b.sqrt())
139}
140
141pub trait VectorIndex: Send + Sync {
142 fn dimensions(&self) -> u32;
143 fn index_kind(&self) -> &'static str {
144 "vector"
145 }
146 fn add(&mut self, doc_id: DocId, vector: Vec<f32>) -> StorageBackendResult<()>;
147 fn add_many(&mut self, doc_id: DocId, vectors: Vec<Vec<f32>>) -> StorageBackendResult<()>;
148 fn delete(&mut self, doc_id: DocId) -> StorageBackendResult<()>;
149 fn clear(&mut self) -> StorageBackendResult<()>;
150 fn search_knn(&self, query: &[f32], k: usize) -> StorageBackendResult<PostingList>;
151 fn search_threshold(&self, query: &[f32], threshold: f32) -> StorageBackendResult<PostingList>;
152 fn count(&self) -> StorageBackendResult<usize>;
153
154 fn initialize(&mut self) -> StorageBackendResult<()> {
159 Ok(())
160 }
161
162 fn snapshot(&self) -> StorageBackendResult<Arc<dyn VectorIndex>>;
164
165 fn writable_snapshot(&self) -> StorageBackendResult<Box<dyn VectorIndex>> {
169 Err(StorageBackendError::Other(
170 "writable vector-index snapshots are not supported by this backend".into(),
171 ))
172 }
173}
174
175#[derive(Debug, Clone)]
176pub struct MemoryVectorIndex {
177 dimensions: u32,
178 vectors: BTreeMap<DocId, Vec<Vec<f32>>>,
179}
180
181impl MemoryVectorIndex {
182 pub fn new(dimensions: u32) -> Self {
183 Self {
184 dimensions,
185 vectors: BTreeMap::new(),
186 }
187 }
188
189 pub fn vectors(&self) -> &BTreeMap<DocId, Vec<Vec<f32>>> {
190 &self.vectors
191 }
192}
193
194impl VectorIndex for MemoryVectorIndex {
195 fn dimensions(&self) -> u32 {
196 self.dimensions
197 }
198
199 fn index_kind(&self) -> &'static str {
200 "memory-bruteforce"
201 }
202
203 fn add(&mut self, doc_id: DocId, vector: Vec<f32>) -> StorageBackendResult<()> {
204 self.validate_dimensions(&vector)?;
205 self.vectors.insert(doc_id, vec![vector]);
206 Ok(())
207 }
208
209 fn add_many(&mut self, doc_id: DocId, vectors: Vec<Vec<f32>>) -> StorageBackendResult<()> {
210 for vector in &vectors {
211 self.validate_dimensions(vector)?;
212 }
213 if vectors.is_empty() {
214 self.vectors.remove(&doc_id);
215 } else {
216 self.vectors.insert(doc_id, vectors);
217 }
218 Ok(())
219 }
220
221 fn delete(&mut self, doc_id: DocId) -> StorageBackendResult<()> {
222 self.vectors.remove(&doc_id);
223 Ok(())
224 }
225
226 fn clear(&mut self) -> StorageBackendResult<()> {
227 self.vectors.clear();
228 Ok(())
229 }
230
231 fn search_knn(&self, query: &[f32], k: usize) -> StorageBackendResult<PostingList> {
233 self.validate_dimensions(query)?;
234 if k == 0 || self.vectors.is_empty() {
235 return Ok(PostingList::new());
236 }
237 let mut scored: Vec<(DocId, f32)> = self
238 .vectors
239 .iter()
240 .filter_map(|(&doc_id, vectors)| best_vector_score(query, vectors).map(|s| (doc_id, s)))
241 .collect();
242 select_top_k_scored(&mut scored, k);
243 scored.sort_by_key(|(id, _)| *id);
246 let entries = scored
247 .into_iter()
248 .map(|(doc_id, sim)| PostingEntry::new(doc_id, Payload::with_score(f64::from(sim))))
249 .collect::<Vec<_>>();
250 Ok(PostingList::from_sorted_unchecked(entries))
251 }
252
253 fn search_threshold(&self, query: &[f32], threshold: f32) -> StorageBackendResult<PostingList> {
255 self.validate_dimensions(query)?;
256 validate_threshold(threshold)?;
257 let mut entries: Vec<PostingEntry> = self
258 .vectors
259 .iter()
260 .filter_map(|(&doc_id, vectors)| {
261 let sim = best_vector_score(query, vectors)?;
262 if sim >= threshold {
263 Some(PostingEntry::new(
264 doc_id,
265 Payload::with_score(f64::from(sim)),
266 ))
267 } else {
268 None
269 }
270 })
271 .collect();
272 entries.sort_by_key(|e| e.doc_id);
275 Ok(PostingList::from_sorted_unchecked(entries))
276 }
277
278 fn count(&self) -> StorageBackendResult<usize> {
279 checked_vector_count(self.vectors.values().map(Vec::len))
280 }
281
282 fn snapshot(&self) -> StorageBackendResult<Arc<dyn VectorIndex>> {
283 Ok(Arc::new(self.clone()))
284 }
285
286 fn writable_snapshot(&self) -> StorageBackendResult<Box<dyn VectorIndex>> {
287 Ok(Box::new(self.clone()))
288 }
289}
290
291impl MemoryVectorIndex {
292 fn validate_dimensions(&self, vector: &[f32]) -> StorageBackendResult<()> {
293 validate_vector_values(self.dimensions, vector)
294 }
295}
296
297fn best_vector_score(query: &[f32], vectors: &[Vec<f32>]) -> Option<f32> {
298 vectors
299 .iter()
300 .map(|vector| cosine_similarity(query, vector))
301 .max_by(f32::total_cmp)
302}
303
304#[cfg(test)]
305mod tests {
306 use super::*;
307
308 fn approx_eq(a: f32, b: f32, eps: f32) {
309 assert!((a - b).abs() < eps, "expected {a} ~ {b} within {eps}");
310 }
311
312 #[test]
313 fn cosine_identity_is_one() {
314 let v = vec![1.0, 2.0, 3.0];
315 approx_eq(cosine_similarity(&v, &v), 1.0, 1e-6);
316 }
317
318 #[test]
319 fn vector_count_overflow_is_reported() {
320 let error = checked_vector_count([usize::MAX, 1]).unwrap_err();
321 assert!(error.to_string().contains("vector count overflow"));
322 }
323
324 #[test]
325 fn cosine_orthogonal_is_zero() {
326 let a = vec![1.0, 0.0];
327 let b = vec![0.0, 1.0];
328 approx_eq(cosine_similarity(&a, &b), 0.0, 1e-6);
329 }
330
331 #[test]
332 fn cosine_zero_norm_is_zero() {
333 let a = vec![0.0, 0.0];
334 let b = vec![1.0, 1.0];
335 approx_eq(cosine_similarity(&a, &b), 0.0, 1e-6);
336 }
337
338 #[test]
339 fn knn_orders_by_similarity_descending_then_doc_id() {
340 let mut idx = MemoryVectorIndex::new(2);
341 idx.add(1, vec![1.0, 0.0]).unwrap();
342 idx.add(2, vec![0.5, 0.5]).unwrap();
343 idx.add(3, vec![0.0, 1.0]).unwrap();
344 let pl = idx.search_knn(&[1.0, 0.0], 2).unwrap();
345 let docs: Vec<_> = pl.iter().map(|e| e.doc_id).collect();
346 assert_eq!(docs, vec![1, 2]);
348 let entry1 = pl.get_entry(1).unwrap();
349 let entry2 = pl.get_entry(2).unwrap();
350 assert!(entry1.payload.score > entry2.payload.score);
351 }
352
353 #[test]
354 fn partial_top_k_keeps_deterministic_doc_id_ties() {
355 let mut scored = vec![(10, 0.5), (3, 0.9), (1, 0.9), (8, 0.7), (2, 0.1)];
356 select_top_k_scored(&mut scored, 2);
357 scored.sort_by_key(|(doc_id, _)| *doc_id);
358 assert_eq!(scored, vec![(1, 0.9), (3, 0.9)]);
359 }
360
361 #[test]
362 fn precomputed_norm_cosine_matches_reference_bits() {
363 let a = [0.25, -3.0, 1.5, 8.0];
364 let b = [2.0, 0.75, -4.0, 0.5];
365 let expected = cosine_similarity(&a, &b);
366 let actual = cosine_similarity_with_norms(&a, &b, vector_norm(&a), vector_norm(&b));
367 assert_eq!(actual.to_bits(), expected.to_bits());
368 }
369
370 #[test]
371 fn score_deduplication_keeps_best_tensor_vector() {
372 let mut scored = vec![(7, 0.3), (2, 0.8), (7, 0.9), (2, 0.4), (9, -0.2)];
373 deduplicate_scored_by_doc(&mut scored);
374 assert_eq!(scored, vec![(2, 0.8), (7, 0.9), (9, -0.2)]);
375 }
376
377 #[test]
378 fn threshold_filters_below_cutoff() {
379 let mut idx = MemoryVectorIndex::new(2);
380 idx.add(1, vec![1.0, 0.0]).unwrap();
381 idx.add(2, vec![0.5, 0.5]).unwrap();
382 idx.add(3, vec![0.0, 1.0]).unwrap();
383 let pl = idx.search_threshold(&[1.0, 0.0], 0.5).unwrap();
384 let docs: Vec<_> = pl.iter().map(|e| e.doc_id).collect();
385 assert_eq!(docs, vec![1, 2]);
386 }
387
388 #[test]
389 fn delete_removes_vector() {
390 let mut idx = MemoryVectorIndex::new(2);
391 idx.add(1, vec![1.0, 0.0]).unwrap();
392 idx.delete(1).unwrap();
393 assert_eq!(idx.count().unwrap(), 0);
394 }
395
396 #[test]
397 fn non_finite_vectors_queries_and_thresholds_are_errors() {
398 let mut idx = MemoryVectorIndex::new(2);
399 assert!(idx.add(1, vec![f32::NAN, 0.0]).is_err());
400 idx.add(1, vec![1.0, 0.0]).unwrap();
401 assert!(idx.search_knn(&[f32::INFINITY, 0.0], 1).is_err());
402 assert!(idx
403 .search_threshold(&[1.0, 0.0], f32::NEG_INFINITY)
404 .is_err());
405 }
406}