uqa_storage/hnsw_index/
index.rs1use std::sync::Arc;
10
11use uqa_core::{DocId, Payload, PostingEntry, PostingList};
12
13use super::metric::normalize_with_norm;
14use super::types::HNSWIndex;
15use crate::vector_index::{
16 cosine_similarity_with_norms, deduplicate_scored_by_doc, select_top_k_scored,
17 validate_vector_values, vector_norm, VectorIndex,
18};
19use crate::{StorageBackendError, StorageBackendResult};
20
21impl VectorIndex for HNSWIndex {
22 fn dimensions(&self) -> u32 {
23 self.dimensions
24 }
25
26 fn index_kind(&self) -> &'static str {
27 "hnsw"
28 }
29
30 fn add(&mut self, doc_id: DocId, vector: Vec<f32>) -> StorageBackendResult<()> {
31 self.add_many(doc_id, vec![vector])
32 }
33
34 fn add_many(&mut self, doc_id: DocId, vectors: Vec<Vec<f32>>) -> StorageBackendResult<()> {
35 self.replace_document_vectors(doc_id, vectors)
36 }
37
38 fn delete(&mut self, doc_id: DocId) -> StorageBackendResult<()> {
39 self.mark_document_deleted(doc_id)?;
40 self.maybe_rebuild()
41 }
42
43 fn clear(&mut self) -> StorageBackendResult<()> {
44 self.nodes.clear();
45 self.active.clear();
46 self.entry_point = None;
47 self.max_level = 0;
48 self.next_node_id = 1;
49 self.deleted_count = 0;
50 self.dirty_nodes.clear();
51 self.full_rewrite = true;
52 Ok(())
53 }
54
55 fn search_knn(&self, query: &[f32], k: usize) -> StorageBackendResult<PostingList> {
56 validate_vector_values(self.dimensions, query)?;
57 if k == 0 || self.active.is_empty() {
58 return Ok(PostingList::new());
59 }
60 let (normalized_query, query_norm) = normalize_with_norm(query);
61 let mut ef = self.params.ef_search.max(k).min(self.nodes.len());
62 let mut scored = Vec::<(DocId, f32)>::with_capacity(ef);
63 loop {
64 scored.clear();
65 for candidate in self.query_candidates(&normalized_query, ef) {
66 let Some(node) = self.nodes.get(&candidate.node_id) else {
67 continue;
68 };
69 if node.deleted {
70 continue;
71 }
72 let score =
73 cosine_similarity_with_norms(query, &node.raw_vector, query_norm, node.norm);
74 scored.push((node.doc_id, score));
75 }
76 deduplicate_scored_by_doc(&mut scored);
77 if scored.len() >= k || ef >= self.nodes.len() {
78 break;
79 }
80 ef = ef
81 .checked_mul(2)
82 .unwrap_or(self.nodes.len())
83 .min(self.nodes.len());
84 }
85 select_top_k_scored(&mut scored, k);
86 scored.sort_by_key(|(doc_id, _)| *doc_id);
87 Ok(PostingList::from_sorted_unchecked(
88 scored
89 .into_iter()
90 .map(|(doc_id, score)| {
91 PostingEntry::new(doc_id, Payload::with_score(f64::from(score)))
92 })
93 .collect(),
94 ))
95 }
96
97 fn search_threshold(&self, query: &[f32], threshold: f32) -> StorageBackendResult<PostingList> {
98 validate_vector_values(self.dimensions, query)?;
99 if !threshold.is_finite() {
100 return Err(StorageBackendError::Other(format!(
101 "vector similarity threshold must be finite, got {threshold}"
102 )));
103 }
104 let query_norm = vector_norm(query);
105 let mut scored = Vec::<(DocId, f32)>::new();
106 for node_id in self.active.values() {
107 let node = self.nodes.get(node_id).ok_or_else(|| {
108 StorageBackendError::Other(format!(
109 "HNSW active map references missing node {node_id}"
110 ))
111 })?;
112 let score =
113 cosine_similarity_with_norms(query, &node.raw_vector, query_norm, node.norm);
114 if score >= threshold {
115 scored.push((node.doc_id, score));
116 }
117 }
118 deduplicate_scored_by_doc(&mut scored);
119 Ok(PostingList::from_sorted_unchecked(
120 scored
121 .into_iter()
122 .map(|(doc_id, score)| {
123 PostingEntry::new(doc_id, Payload::with_score(f64::from(score)))
124 })
125 .collect(),
126 ))
127 }
128
129 fn count(&self) -> StorageBackendResult<usize> {
130 Ok(self.active.len())
131 }
132
133 fn snapshot(&self) -> StorageBackendResult<Arc<dyn VectorIndex>> {
134 Ok(Arc::new(self.clone()))
135 }
136
137 fn writable_snapshot(&self) -> StorageBackendResult<Box<dyn VectorIndex>> {
138 Ok(Box::new(self.clone()))
139 }
140}