Skip to main content

ruvector_core/index/
hnsw.rs

1//! HNSW (Hierarchical Navigable Small World) index implementation
2
3use crate::distance::distance;
4use crate::error::{Result, RuvectorError};
5use crate::index::VectorIndex;
6use crate::types::{DistanceMetric, HnswConfig, SearchResult, VectorId};
7use bincode::{Decode, Encode};
8use dashmap::DashMap;
9use hnsw_rs::prelude::*;
10use parking_lot::RwLock;
11use std::sync::Arc;
12
13/// Distance function wrapper for hnsw_rs
14struct DistanceFn {
15    metric: DistanceMetric,
16}
17
18impl DistanceFn {
19    fn new(metric: DistanceMetric) -> Self {
20        Self { metric }
21    }
22}
23
24impl Distance<f32> for DistanceFn {
25    #[inline(always)]
26    fn eval(&self, a: &[f32], b: &[f32]) -> f32 {
27        // Bypass the simsimd/Result-overhead path and call our hand-written
28        // SIMD kernels directly.  hnsw_rs asserts dist >= 0 in its search
29        // loop, so clamp any floating-point rounding below zero.
30        use crate::simd_intrinsics;
31        match self.metric {
32            DistanceMetric::Euclidean => simd_intrinsics::euclidean_distance_simd(a, b),
33            DistanceMetric::Cosine => {
34                // cosine_similarity_simd returns dot/(|a||b|); HNSW needs
35                // cosine DISTANCE = 1 - sim, clamped to 0.
36                (1.0_f32 - simd_intrinsics::cosine_similarity_simd(a, b)).max(0.0)
37            }
38            DistanceMetric::DotProduct => {
39                // Negate for minimization; clamp per hnsw_rs assertion.
40                (-simd_intrinsics::dot_product_simd(a, b)).max(0.0)
41            }
42            DistanceMetric::Manhattan => simd_intrinsics::manhattan_distance_simd(a, b),
43        }
44    }
45}
46
47/// HNSW index wrapper
48pub struct HnswIndex {
49    inner: Arc<RwLock<HnswInner>>,
50    config: HnswConfig,
51    metric: DistanceMetric,
52    dimensions: usize,
53}
54
55struct HnswInner {
56    hnsw: Hnsw<'static, f32, DistanceFn>,
57    vectors: DashMap<VectorId, Vec<f32>>,
58    id_to_idx: DashMap<VectorId, usize>,
59    idx_to_id: DashMap<usize, VectorId>,
60    next_idx: usize,
61}
62
63/// Serializable HNSW index state
64#[derive(Encode, Decode, Clone)]
65pub struct HnswState {
66    vectors: Vec<(String, Vec<f32>)>,
67    id_to_idx: Vec<(String, usize)>,
68    idx_to_id: Vec<(usize, String)>,
69    next_idx: usize,
70    config: SerializableHnswConfig,
71    dimensions: usize,
72    metric: SerializableDistanceMetric,
73}
74
75#[derive(Encode, Decode, Clone)]
76struct SerializableHnswConfig {
77    m: usize,
78    ef_construction: usize,
79    ef_search: usize,
80    max_elements: usize,
81}
82
83#[derive(Encode, Decode, Clone, Copy)]
84enum SerializableDistanceMetric {
85    Euclidean,
86    Cosine,
87    DotProduct,
88    Manhattan,
89}
90
91impl From<DistanceMetric> for SerializableDistanceMetric {
92    fn from(metric: DistanceMetric) -> Self {
93        match metric {
94            DistanceMetric::Euclidean => SerializableDistanceMetric::Euclidean,
95            DistanceMetric::Cosine => SerializableDistanceMetric::Cosine,
96            DistanceMetric::DotProduct => SerializableDistanceMetric::DotProduct,
97            DistanceMetric::Manhattan => SerializableDistanceMetric::Manhattan,
98        }
99    }
100}
101
102impl From<SerializableDistanceMetric> for DistanceMetric {
103    fn from(metric: SerializableDistanceMetric) -> Self {
104        match metric {
105            SerializableDistanceMetric::Euclidean => DistanceMetric::Euclidean,
106            SerializableDistanceMetric::Cosine => DistanceMetric::Cosine,
107            SerializableDistanceMetric::DotProduct => DistanceMetric::DotProduct,
108            SerializableDistanceMetric::Manhattan => DistanceMetric::Manhattan,
109        }
110    }
111}
112
113impl HnswIndex {
114    /// Create a new HNSW index
115    pub fn new(dimensions: usize, metric: DistanceMetric, config: HnswConfig) -> Result<Self> {
116        let distance_fn = DistanceFn::new(metric);
117
118        // Create HNSW with configured parameters
119        let hnsw = Hnsw::<f32, DistanceFn>::new(
120            config.m,
121            config.max_elements,
122            dimensions,
123            config.ef_construction,
124            distance_fn,
125        );
126
127        Ok(Self {
128            inner: Arc::new(RwLock::new(HnswInner {
129                hnsw,
130                vectors: DashMap::new(),
131                id_to_idx: DashMap::new(),
132                idx_to_id: DashMap::new(),
133                next_idx: 0,
134            })),
135            config,
136            metric,
137            dimensions,
138        })
139    }
140
141    /// Get configuration
142    pub fn config(&self) -> &HnswConfig {
143        &self.config
144    }
145
146    /// Set efSearch parameter for query-time accuracy tuning.
147    ///
148    /// Higher values increase recall at the cost of search latency.
149    /// Typical range: 50–500. Must be >= k for meaningful results.
150    pub fn set_ef_search(&mut self, ef_search: usize) {
151        self.config.ef_search = ef_search;
152    }
153
154    /// Serialize the index to bytes using bincode
155    pub fn serialize(&self) -> Result<Vec<u8>> {
156        let inner = self.inner.read();
157
158        let state = HnswState {
159            vectors: inner
160                .vectors
161                .iter()
162                .map(|entry| (entry.key().clone(), entry.value().clone()))
163                .collect(),
164            id_to_idx: inner
165                .id_to_idx
166                .iter()
167                .map(|entry| (entry.key().clone(), *entry.value()))
168                .collect(),
169            idx_to_id: inner
170                .idx_to_id
171                .iter()
172                .map(|entry| (*entry.key(), entry.value().clone()))
173                .collect(),
174            next_idx: inner.next_idx,
175            config: SerializableHnswConfig {
176                m: self.config.m,
177                ef_construction: self.config.ef_construction,
178                ef_search: self.config.ef_search,
179                max_elements: self.config.max_elements,
180            },
181            dimensions: self.dimensions,
182            metric: self.metric.into(),
183        };
184
185        bincode::encode_to_vec(&state, bincode::config::standard()).map_err(|e| {
186            RuvectorError::SerializationError(format!("Failed to serialize HNSW index: {}", e))
187        })
188    }
189
190    /// Deserialize the index from bytes using bincode
191    pub fn deserialize(bytes: &[u8]) -> Result<Self> {
192        let (state, _): (HnswState, usize) =
193            bincode::decode_from_slice(bytes, bincode::config::standard()).map_err(|e| {
194                RuvectorError::SerializationError(format!(
195                    "Failed to deserialize HNSW index: {}",
196                    e
197                ))
198            })?;
199
200        let config = HnswConfig {
201            m: state.config.m,
202            ef_construction: state.config.ef_construction,
203            ef_search: state.config.ef_search,
204            max_elements: state.config.max_elements,
205        };
206
207        let dimensions = state.dimensions;
208        let metric: DistanceMetric = state.metric.into();
209
210        let distance_fn = DistanceFn::new(metric);
211        let mut hnsw = Hnsw::<'static, f32, DistanceFn>::new(
212            config.m,
213            config.max_elements,
214            dimensions,
215            config.ef_construction,
216            distance_fn,
217        );
218
219        // Rebuild the index by inserting all vectors.
220        // Build a HashMap first to avoid O(n^2) linear search in the loop below.
221        let vectors_lookup: std::collections::HashMap<&str, &Vec<f32>> = state
222            .vectors
223            .iter()
224            .map(|(id, v)| (id.as_str(), v))
225            .collect();
226
227        let id_to_idx: DashMap<VectorId, usize> = state.id_to_idx.into_iter().collect();
228        let idx_to_id: DashMap<usize, VectorId> = state.idx_to_id.into_iter().collect();
229
230        // Insert vectors into HNSW in index order for deterministic reconstruction.
231        let mut sorted_entries: Vec<_> = idx_to_id
232            .iter()
233            .map(|e| (*e.key(), e.value().clone()))
234            .collect();
235        sorted_entries.sort_unstable_by_key(|(idx, _)| *idx);
236
237        for (idx, id) in &sorted_entries {
238            if let Some(vector) = vectors_lookup.get(id.as_str()) {
239                hnsw.insert_data(vector, *idx);
240            }
241        }
242
243        let vectors_map: DashMap<VectorId, Vec<f32>> = state.vectors.into_iter().collect();
244
245        Ok(Self {
246            inner: Arc::new(RwLock::new(HnswInner {
247                hnsw,
248                vectors: vectors_map,
249                id_to_idx,
250                idx_to_id,
251                next_idx: state.next_idx,
252            })),
253            config,
254            metric,
255            dimensions,
256        })
257    }
258
259    /// Search with custom efSearch parameter.
260    ///
261    /// `ef_search` must be >= `k`; values smaller than `k` are clamped to `k`
262    /// to avoid silent under-recall.  Results are returned sorted by ascending
263    /// distance (closest first).
264    pub fn search_with_ef(
265        &self,
266        query: &[f32],
267        k: usize,
268        ef_search: usize,
269    ) -> Result<Vec<SearchResult>> {
270        if query.len() != self.dimensions {
271            return Err(RuvectorError::DimensionMismatch {
272                expected: self.dimensions,
273                actual: query.len(),
274            });
275        }
276
277        if k == 0 {
278            return Ok(vec![]);
279        }
280
281        let inner = self.inner.read();
282
283        // hnsw_rs panics in its BinaryHeap traversal when the index is empty
284        // or contains only a single element (the candidate/return-point loop
285        // calls .peek().unwrap() without an emptiness guard).  Return early
286        // to surface a clean error instead of an assertion panic.
287        if inner.vectors.is_empty() {
288            return Ok(vec![]);
289        }
290
291        // ef_search < k causes hnsw_rs to return fewer than k candidates; clamp.
292        let effective_ef = ef_search.max(k);
293
294        // Use HNSW search with custom ef parameter (knbn)
295        let neighbors = inner.hnsw.search(query, k, effective_ef);
296
297        let mut results: Vec<SearchResult> = neighbors
298            .into_iter()
299            .filter_map(|neighbor| {
300                inner.idx_to_id.get(&neighbor.d_id).map(|id| SearchResult {
301                    id: id.clone(),
302                    score: neighbor.distance,
303                    vector: None,
304                    metadata: None,
305                })
306            })
307            .collect();
308
309        // hnsw_rs does not guarantee sort order — ensure ascending distance.
310        results.sort_unstable_by(|a, b| {
311            a.score
312                .partial_cmp(&b.score)
313                .unwrap_or(std::cmp::Ordering::Equal)
314        });
315
316        Ok(results)
317    }
318}
319
320impl VectorIndex for HnswIndex {
321    fn add(&mut self, id: VectorId, vector: Vec<f32>) -> Result<()> {
322        if vector.len() != self.dimensions {
323            return Err(RuvectorError::DimensionMismatch {
324                expected: self.dimensions,
325                actual: vector.len(),
326            });
327        }
328
329        let mut inner = self.inner.write();
330        let idx = inner.next_idx;
331        inner.next_idx += 1;
332
333        // Insert into HNSW graph using insert_data
334        inner.hnsw.insert_data(&vector, idx);
335
336        // Store mappings
337        inner.vectors.insert(id.clone(), vector);
338        inner.id_to_idx.insert(id.clone(), idx);
339        inner.idx_to_id.insert(idx, id);
340
341        Ok(())
342    }
343
344    fn add_batch(&mut self, entries: Vec<(VectorId, Vec<f32>)>) -> Result<()> {
345        // Validate all dimensions first
346        for (_, vector) in &entries {
347            if vector.len() != self.dimensions {
348                return Err(RuvectorError::DimensionMismatch {
349                    expected: self.dimensions,
350                    actual: vector.len(),
351                });
352            }
353        }
354
355        let mut inner = self.inner.write();
356
357        // Prepare batch data for insertion
358        // First, assign indices and collect vector data
359        let data_with_ids: Vec<_> = entries
360            .iter()
361            .enumerate()
362            .map(|(i, (id, vector))| {
363                let idx = inner.next_idx + i;
364                (id.clone(), idx, vector.clone())
365            })
366            .collect();
367
368        // Update next_idx
369        inner.next_idx += entries.len();
370
371        // For large batches (>=PARALLEL_THRESHOLD), use hnsw_rs parallel
372        // insert (rayon-based) to cut build time.  Below this threshold,
373        // sequential insert maintains better graph connectivity — parallel
374        // workers can miss each other's in-flight inserts, producing fewer
375        // optimal neighbors and increasing search latency on small indexes.
376        //
377        // Rule of thumb from hnsw_rs: parallel is efficient only when
378        // n_inserts >= 1000 * num_threads.  We conservatively gate at 10 K.
379        const PARALLEL_THRESHOLD: usize = 10_000;
380        if data_with_ids.len() >= PARALLEL_THRESHOLD {
381            let datas: Vec<(&[f32], usize)> = data_with_ids
382                .iter()
383                .map(|(_id, idx, vector)| (vector.as_slice(), *idx))
384                .collect();
385            inner.hnsw.parallel_insert_slice(&datas);
386        } else {
387            for (_id, idx, vector) in &data_with_ids {
388                inner.hnsw.insert_data(vector, *idx);
389            }
390        }
391
392        // Store mappings
393        for (id, idx, vector) in data_with_ids {
394            inner.vectors.insert(id.clone(), vector);
395            inner.id_to_idx.insert(id.clone(), idx);
396            inner.idx_to_id.insert(idx, id);
397        }
398
399        Ok(())
400    }
401
402    fn search(&self, query: &[f32], k: usize) -> Result<Vec<SearchResult>> {
403        // Use configured ef_search
404        self.search_with_ef(query, k, self.config.ef_search)
405    }
406
407    fn remove(&mut self, id: &VectorId) -> Result<bool> {
408        let inner = self.inner.write();
409
410        // Note: hnsw_rs doesn't support direct deletion
411        // We remove from our mappings but the graph structure remains
412        // This is a known limitation of HNSW
413        let removed = inner.vectors.remove(id).is_some();
414
415        if removed {
416            if let Some((_, idx)) = inner.id_to_idx.remove(id) {
417                inner.idx_to_id.remove(&idx);
418            }
419        }
420
421        Ok(removed)
422    }
423
424    fn len(&self) -> usize {
425        self.inner.read().vectors.len()
426    }
427}
428
429#[cfg(test)]
430mod tests {
431    use super::*;
432
433    fn generate_random_vectors(count: usize, dimensions: usize) -> Vec<Vec<f32>> {
434        use rand::Rng;
435        let mut rng = rand::thread_rng();
436
437        (0..count)
438            .map(|_| (0..dimensions).map(|_| rng.gen::<f32>()).collect())
439            .collect()
440    }
441
442    fn normalize_vector(v: &[f32]) -> Vec<f32> {
443        let norm = v.iter().map(|x| x * x).sum::<f32>().sqrt();
444        if norm > 0.0 {
445            v.iter().map(|x| x / norm).collect()
446        } else {
447            v.to_vec()
448        }
449    }
450
451    #[test]
452    fn test_hnsw_index_creation() -> Result<()> {
453        let config = HnswConfig::default();
454        let index = HnswIndex::new(128, DistanceMetric::Cosine, config)?;
455        assert_eq!(index.len(), 0);
456        Ok(())
457    }
458
459    #[test]
460    fn test_hnsw_insert_and_search() -> Result<()> {
461        let config = HnswConfig {
462            m: 16,
463            ef_construction: 100,
464            ef_search: 50,
465            max_elements: 1000,
466        };
467
468        let mut index = HnswIndex::new(128, DistanceMetric::Cosine, config)?;
469
470        // Insert a few vectors
471        let vectors = generate_random_vectors(100, 128);
472        for (i, vector) in vectors.iter().enumerate() {
473            let normalized = normalize_vector(vector);
474            index.add(format!("vec_{}", i), normalized)?;
475        }
476
477        assert_eq!(index.len(), 100);
478
479        // Search for the first vector
480        let query = normalize_vector(&vectors[0]);
481        let results = index.search(&query, 10)?;
482
483        assert!(!results.is_empty());
484        assert_eq!(results[0].id, "vec_0");
485
486        Ok(())
487    }
488
489    #[test]
490    fn test_hnsw_batch_insert() -> Result<()> {
491        let config = HnswConfig::default();
492        let mut index = HnswIndex::new(128, DistanceMetric::Cosine, config)?;
493
494        let vectors = generate_random_vectors(100, 128);
495        let entries: Vec<_> = vectors
496            .iter()
497            .enumerate()
498            .map(|(i, v)| (format!("vec_{}", i), normalize_vector(v)))
499            .collect();
500
501        index.add_batch(entries)?;
502        assert_eq!(index.len(), 100);
503
504        Ok(())
505    }
506
507    #[test]
508    fn test_hnsw_serialization() -> Result<()> {
509        let config = HnswConfig {
510            m: 16,
511            ef_construction: 100,
512            ef_search: 50,
513            max_elements: 1000,
514        };
515
516        let mut index = HnswIndex::new(128, DistanceMetric::Cosine, config)?;
517
518        // Insert vectors
519        let vectors = generate_random_vectors(50, 128);
520        for (i, vector) in vectors.iter().enumerate() {
521            let normalized = normalize_vector(vector);
522            index.add(format!("vec_{}", i), normalized)?;
523        }
524
525        // Serialize
526        let bytes = index.serialize()?;
527
528        // Deserialize
529        let restored_index = HnswIndex::deserialize(&bytes)?;
530
531        assert_eq!(restored_index.len(), 50);
532
533        // Test search on restored index
534        let query = normalize_vector(&vectors[0]);
535        let results = restored_index.search(&query, 5)?;
536
537        assert!(!results.is_empty());
538
539        Ok(())
540    }
541
542    #[test]
543    fn test_dimension_mismatch() -> Result<()> {
544        let config = HnswConfig::default();
545        let mut index = HnswIndex::new(128, DistanceMetric::Cosine, config)?;
546
547        let result = index.add("test".to_string(), vec![1.0; 64]);
548        assert!(result.is_err());
549
550        Ok(())
551    }
552}