1use crate::hnsw::query_cache::{QueryCache, QueryCacheConfig};
4use crate::hnsw::{HnswConfig, HnswPerformanceStats, Node};
5use crate::{Vector, VectorIndex};
6use anyhow::Result;
7use std::collections::HashMap;
8use std::sync::atomic::AtomicU64;
9#[cfg(feature = "gpu")]
10use std::sync::Arc;
11
12#[cfg(feature = "gpu")]
13use crate::gpu::GpuAccelerator;
14
15pub struct HnswIndex {
17 config: HnswConfig,
18 nodes: Vec<Node>,
19 uri_to_id: HashMap<String, usize>,
20 entry_point: Option<usize>,
21 level_multiplier: f64,
22 rng_state: u64,
23 stats: HnswPerformanceStats,
25 distance_calculations: AtomicU64,
27 query_cache: Option<QueryCache>,
29 #[cfg(feature = "gpu")]
31 gpu_accelerator: Option<Arc<GpuAccelerator>>,
32 #[cfg(feature = "gpu")]
34 multi_gpu_accelerators: Vec<Arc<GpuAccelerator>>,
35}
36
37impl HnswIndex {
38 pub fn new(config: HnswConfig) -> Result<Self> {
39 #[cfg(feature = "gpu")]
41 let (gpu_accelerator, multi_gpu_accelerators) = if config.enable_gpu {
42 let gpu_config = config.gpu_config.clone().unwrap_or_default();
43
44 if config.enable_multi_gpu && gpu_config.preferred_gpu_ids.len() > 1 {
45 let mut accelerators = Vec::new();
47 for &gpu_id in &gpu_config.preferred_gpu_ids {
48 let mut gpu_conf = gpu_config.clone();
49 gpu_conf.device_id = gpu_id;
50 let accelerator = GpuAccelerator::new(gpu_conf)?;
51 accelerators.push(Arc::new(accelerator));
52 }
53 (None, accelerators)
54 } else {
55 let accelerator = GpuAccelerator::new(gpu_config)?;
57 (Some(Arc::new(accelerator)), Vec::new())
58 }
59 } else {
60 (None, Vec::new())
61 };
62
63 let query_cache = Some(QueryCache::new(QueryCacheConfig::default()));
65
66 Ok(Self {
67 config,
68 nodes: Vec::new(),
69 uri_to_id: HashMap::new(),
70 entry_point: None,
71 level_multiplier: 1.0 / (2.0_f64).ln(),
72 rng_state: 42, stats: HnswPerformanceStats::default(),
74 distance_calculations: AtomicU64::new(0),
75 query_cache,
76 #[cfg(feature = "gpu")]
77 gpu_accelerator,
78 #[cfg(feature = "gpu")]
79 multi_gpu_accelerators,
80 })
81 }
82
83 pub fn new_cpu_only(config: HnswConfig) -> Self {
85 let mut cpu_config = config;
86 cpu_config.enable_gpu = false;
87 cpu_config.enable_multi_gpu = false;
88
89 let query_cache = Some(QueryCache::new(QueryCacheConfig::default()));
91
92 Self {
93 config: cpu_config,
94 nodes: Vec::new(),
95 uri_to_id: HashMap::new(),
96 entry_point: None,
97 level_multiplier: 1.0 / (2.0_f64).ln(),
98 rng_state: 42,
99 stats: HnswPerformanceStats::default(),
100 distance_calculations: AtomicU64::new(0),
101 query_cache,
102 #[cfg(feature = "gpu")]
103 gpu_accelerator: None,
104 #[cfg(feature = "gpu")]
105 multi_gpu_accelerators: Vec::new(),
106 }
107 }
108
109 pub fn enable_query_cache(&mut self, config: QueryCacheConfig) {
111 self.query_cache = Some(QueryCache::new(config));
112 }
113
114 pub fn disable_query_cache(&mut self) {
116 self.query_cache = None;
117 }
118
119 pub fn get_query_cache_stats(&self) -> Option<crate::hnsw::query_cache::QueryCacheStats> {
121 self.query_cache.as_ref().map(|cache| cache.get_stats())
122 }
123
124 pub fn clear_query_cache(&self) {
126 if let Some(ref cache) = self.query_cache {
127 cache.clear();
128 }
129 }
130
131 pub(crate) fn query_cache(&self) -> &Option<QueryCache> {
133 &self.query_cache
134 }
135
136 pub fn uri_to_id(&self) -> &HashMap<String, usize> {
138 &self.uri_to_id
139 }
140
141 pub fn uri_to_id_mut(&mut self) -> &mut HashMap<String, usize> {
143 &mut self.uri_to_id
144 }
145
146 pub fn nodes(&self) -> &Vec<Node> {
148 &self.nodes
149 }
150
151 pub fn nodes_mut(&mut self) -> &mut Vec<Node> {
153 &mut self.nodes
154 }
155
156 pub fn entry_point(&self) -> Option<usize> {
158 self.entry_point
159 }
160
161 pub fn set_entry_point(&mut self, entry_point: Option<usize>) {
163 self.entry_point = entry_point;
164 }
165
166 pub fn config(&self) -> &HnswConfig {
168 &self.config
169 }
170
171 pub fn get_stats(&self) -> &HnswPerformanceStats {
173 &self.stats
174 }
175
176 #[cfg(feature = "gpu")]
178 pub fn is_gpu_available(&self) -> bool {
179 self.config.enable_gpu
180 && (self.gpu_accelerator.is_some() || !self.multi_gpu_accelerators.is_empty())
181 }
182
183 #[cfg(not(feature = "gpu"))]
184 pub fn is_gpu_available(&self) -> bool {
185 false
186 }
187
188 #[cfg(feature = "gpu")]
190 pub fn get_gpu_stats(&self) -> Option<crate::gpu::GpuPerformanceStats> {
191 if let Some(ref _accelerator) = self.gpu_accelerator {
192 None } else {
195 None
196 }
197 }
198
199 #[cfg(feature = "gpu")]
201 pub fn gpu_accelerator(&self) -> Option<&Arc<GpuAccelerator>> {
202 self.gpu_accelerator.as_ref()
203 }
204
205 #[cfg(feature = "gpu")]
207 pub fn multi_gpu_accelerators(&self) -> &Vec<Arc<GpuAccelerator>> {
208 &self.multi_gpu_accelerators
209 }
210
211 pub fn len(&self) -> usize {
213 self.nodes.len()
214 }
215
216 pub fn is_empty(&self) -> bool {
218 self.nodes.is_empty()
219 }
220
221 pub fn stats_mut(&mut self) -> &mut HnswPerformanceStats {
225 &mut self.stats
226 }
227
228 pub fn level_multiplier(&self) -> f64 {
230 self.level_multiplier
231 }
232
233 pub fn rng_state_mut(&mut self) -> &mut u64 {
235 &mut self.rng_state
236 }
237
238 pub fn rng_state(&self) -> u64 {
240 self.rng_state
241 }
242}
243
244impl VectorIndex for HnswIndex {
245 fn insert(&mut self, uri: String, vector: Vector) -> Result<()> {
246 self.add_vector(uri, vector)
248 }
249
250 fn search_knn(&self, query: &Vector, k: usize) -> Result<Vec<(String, f32)>> {
251 let raw = HnswIndex::search_knn(self, query, k)?;
257 Ok(raw
258 .into_iter()
259 .map(|(uri, distance)| (uri, 1.0 / (1.0 + distance)))
260 .collect())
261 }
262
263 fn search_threshold(&self, query: &Vector, threshold: f32) -> Result<Vec<(String, f32)>> {
264 if threshold > 1.0 {
269 return Ok(Vec::new());
271 }
272 let radius = if threshold <= 0.0 {
273 f32::MAX
274 } else {
275 1.0 / threshold - 1.0
276 };
277 let raw = HnswIndex::range_search(self, query, radius)?;
278 Ok(raw
279 .into_iter()
280 .map(|(uri, distance)| (uri, 1.0 / (1.0 + distance)))
281 .collect())
282 }
283
284 fn get_vector(&self, uri: &str) -> Option<&Vector> {
285 self.uri_to_id
286 .get(uri)
287 .and_then(|&id| self.nodes.get(id))
288 .map(|node| &node.vector)
289 }
290
291 fn iter_vectors(&self) -> Vec<(String, Vector)> {
292 self.uri_to_id
296 .iter()
297 .filter_map(|(uri, &id)| {
298 self.nodes
299 .get(id)
300 .map(|node| (uri.clone(), node.vector.clone()))
301 })
302 .collect()
303 }
304
305 fn supports_enumeration(&self) -> bool {
306 true
307 }
308}
309
310impl HnswIndex {
311 pub fn remove(&mut self, uri: &str) -> Result<()> {
313 let node_id = if let Some(&id) = self.uri_to_id.get(uri) {
317 id
318 } else {
319 return Err(anyhow::anyhow!("URI not found: {}", uri));
320 };
321
322 if let Some(node) = self.nodes.get(node_id) {
324 let node_connections = node.connections.clone();
325
326 for (level, connections) in node_connections.iter().enumerate() {
328 for &connected_id in connections {
329 if let Some(connected_node) = self.nodes.get_mut(connected_id) {
330 connected_node.remove_connection(level, node_id);
331 }
332 }
333 }
334 }
335
336 if self.entry_point == Some(node_id) {
338 self.entry_point = None;
339
340 let mut highest_level = 0;
342 let mut new_entry_point = None;
343
344 for (id, node) in self.nodes.iter().enumerate() {
345 if id != node_id && node.level() >= highest_level {
346 highest_level = node.level();
347 new_entry_point = Some(id);
348 }
349 }
350
351 self.entry_point = new_entry_point;
352 }
353
354 self.uri_to_id.remove(uri);
356
357 if let Some(node) = self.nodes.get_mut(node_id) {
361 node.connections.clear();
362 }
364
365 self.stats
367 .total_deletions
368 .fetch_add(1, std::sync::atomic::Ordering::Relaxed);
369
370 Ok(())
371 }
372
373 pub fn update(&mut self, uri: String, vector: Vector) -> Result<()> {
375 if !self.uri_to_id.contains_key(&uri) {
380 return Err(anyhow::anyhow!("URI not found: {}", uri));
381 }
382
383 let node_id = self.uri_to_id[&uri];
385 let _old_connections = self.nodes.get(node_id).map(|node| node.connections.clone());
386
387 self.remove(&uri)?;
389
390 self.insert(uri.clone(), vector)?;
392
393 self.stats
395 .total_updates
396 .fetch_add(1, std::sync::atomic::Ordering::Relaxed);
397
398 Ok(())
404 }
405
406 pub fn clear(&mut self) -> Result<()> {
408 self.nodes.clear();
409 self.uri_to_id.clear();
410 self.entry_point = None;
411 Ok(())
412 }
413
414 pub fn size(&self) -> usize {
416 self.nodes.len()
417 }
418}