1use 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
13struct 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 use crate::simd_intrinsics;
31 match self.metric {
32 DistanceMetric::Euclidean => simd_intrinsics::euclidean_distance_simd(a, b),
33 DistanceMetric::Cosine => {
34 (1.0_f32 - simd_intrinsics::cosine_similarity_simd(a, b)).max(0.0)
37 }
38 DistanceMetric::DotProduct => {
39 (-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
47pub 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#[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 pub fn new(dimensions: usize, metric: DistanceMetric, config: HnswConfig) -> Result<Self> {
116 let distance_fn = DistanceFn::new(metric);
117
118 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 pub fn config(&self) -> &HnswConfig {
143 &self.config
144 }
145
146 pub fn set_ef_search(&mut self, ef_search: usize) {
151 self.config.ef_search = ef_search;
152 }
153
154 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 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 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 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 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 if inner.vectors.is_empty() {
288 return Ok(vec![]);
289 }
290
291 let effective_ef = ef_search.max(k);
293
294 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 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 inner.hnsw.insert_data(&vector, idx);
335
336 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 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 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 inner.next_idx += entries.len();
370
371 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 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 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 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 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 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 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 let bytes = index.serialize()?;
527
528 let restored_index = HnswIndex::deserialize(&bytes)?;
530
531 assert_eq!(restored_index.len(), 50);
532
533 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}