zeusdb-vector-database 0.0.1__tar.gz → 0.0.2__tar.gz

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.

Potentially problematic release.


This version of zeusdb-vector-database might be problematic. Click here for more details.

@@ -1,9 +1,10 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: zeusdb-vector-database
3
- Version: 0.0.1
3
+ Version: 0.0.2
4
4
  Classifier: Programming Language :: Rust
5
5
  Classifier: Programming Language :: Python :: Implementation :: CPython
6
- Requires-Dist: maturin>=1.8.7 ; extra == 'dev'
6
+ Requires-Dist: numpy>=2.2.6,<3.0.0
7
+ Requires-Dist: maturin>=1.9.0 ; extra == 'dev'
7
8
  Requires-Dist: pytest>=8.4.0 ; extra == 'dev'
8
9
  Provides-Extra: dev
9
10
  License-File: LICENSE
@@ -141,7 +142,7 @@ query_vec = [0.1, 0.2, 0.3, 0.1, 0.4, 0.2, 0.6, 0.7]
141
142
  # Query with no filter (all documents)
142
143
  print("\n--- Querying without filter (all documents) ---")
143
144
  results = index.query(vector=query_vec, filter=None, top_k=2)
144
- for doc_id, score in results_all:
145
+ for doc_id, score in results:
145
146
  print(f"{doc_id} (score={score:.4f})")
146
147
  ```
147
148
 
@@ -123,7 +123,7 @@ query_vec = [0.1, 0.2, 0.3, 0.1, 0.4, 0.2, 0.6, 0.7]
123
123
  # Query with no filter (all documents)
124
124
  print("\n--- Querying without filter (all documents) ---")
125
125
  results = index.query(vector=query_vec, filter=None, top_k=2)
126
- for doc_id, score in results_all:
126
+ for doc_id, score in results:
127
127
  print(f"{doc_id} (score={score:.4f})")
128
128
  ```
129
129
 
@@ -1,6 +1,6 @@
1
1
  [project]
2
2
  name = "zeusdb-vector-database"
3
- version = "0.0.1"
3
+ version = "0.0.2"
4
4
  description = "Blazing-fast vector DB with real-time similarity search and metadata filtering."
5
5
  readme = "README.md"
6
6
  authors = [
@@ -12,10 +12,12 @@ classifiers = [
12
12
  "Programming Language :: Rust",
13
13
  "Programming Language :: Python :: Implementation :: CPython",
14
14
  ]
15
- dependencies = []
15
+ dependencies = [
16
+ "numpy>=2.2.6,<3.0.0"
17
+ ]
16
18
 
17
19
  [build-system]
18
- requires = ["maturin>=1.8.7,<2.0"]
20
+ requires = ["maturin>=1.9.0,<2.0"]
19
21
  build-backend = "maturin"
20
22
 
21
23
  [tool.maturin]
@@ -28,7 +30,7 @@ include = ["LICENSE", "NOTICE", "README.md", "src/**"]
28
30
 
29
31
  [project.optional-dependencies]
30
32
  dev = [
31
- "maturin >=1.8.7",
33
+ "maturin >=1.9.0",
32
34
  "pytest >=8.4.0",
33
35
  ]
34
36
 
@@ -1,4 +1,7 @@
1
- __version__ = "0.0.1"
1
+ """
2
+ ZeusDB Vector Database Module
3
+ """
4
+ __version__ = "0.0.2"
2
5
 
3
6
  from .vector_database import VectorDatabase # imports the VectorDatabase class from the vector_database.py file
4
7
 
@@ -0,0 +1,50 @@
1
+ """
2
+ vector_database.py
3
+
4
+ Pure factory for creating vector indexes using Rust backend.
5
+ Currently supports HNSW (Hierarchical Navigable Small World).
6
+ """
7
+ from .zeusdb_vector_database import HNSWIndex
8
+
9
+ class VectorDatabase:
10
+ """
11
+ Pure factory for creating vector indexes.
12
+ No state management - just creates and returns indexes.
13
+ """
14
+
15
+ def __init__(self):
16
+ """Initialize VectorDatabase factory."""
17
+ pass
18
+
19
+ def create_index_hnsw(
20
+ self,
21
+ dim: int = 1536,
22
+ space: str = "cosine",
23
+ M: int = 16,
24
+ ef_construction: int = 200,
25
+ expected_size: int = 10000
26
+ ) -> HNSWIndex:
27
+ """
28
+ Creates a new HNSW (Hierarchical Navigable Small World) index.
29
+
30
+ Args:
31
+ dim: Vector dimension (default: 1536)
32
+ space: Distance metric, only 'cosine' supported (default: 'cosine')
33
+ M: Bidirectional links per node (default: 16, max: 256)
34
+ ef_construction: Construction candidate list size (default: 200)
35
+ expected_size: Expected number of vectors (default: 10000)
36
+
37
+ Returns:
38
+ HNSWIndex: Use this index directly for all operations
39
+
40
+ Example:
41
+ vdb = VectorDatabase()
42
+ index = vdb.create_index_hnsw(dim=1536, expected_size=10000)
43
+ index.add_point("doc1", vector, metadata)
44
+ results = index.query(query_vector, k=10)
45
+ """
46
+ try:
47
+ return HNSWIndex(dim, space, M, ef_construction, expected_size)
48
+ except Exception as e:
49
+ raise RuntimeError(f"Failed to create HNSW index: {e}") from e
50
+
@@ -301,9 +301,9 @@ dependencies = [
301
301
 
302
302
  [[package]]
303
303
  name = "indexmap"
304
- version = "2.9.0"
304
+ version = "2.10.0"
305
305
  source = "registry+https://github.com/rust-lang/crates.io-index"
306
- checksum = "cea70ddb795996207ad57735b50c5982d8844f38ba9ee5f1aedcfb708a2aa11e"
306
+ checksum = "fe4cd85333e22411419a0bcae1297d25e58c9443848b11dc6a86fefe8c78a661"
307
307
  dependencies = [
308
308
  "equivalent",
309
309
  "hashbrown",
@@ -375,9 +375,9 @@ checksum = "13dc2df351e3202783a1fe0d44375f7295ffb4049267b0f3018346dc122a1d94"
375
375
 
376
376
  [[package]]
377
377
  name = "mach2"
378
- version = "0.4.2"
378
+ version = "0.4.3"
379
379
  source = "registry+https://github.com/rust-lang/crates.io-index"
380
- checksum = "19b955cdeb2a02b9117f121ce63aa52d08ade45de53e48fe6a38b39c10f6f709"
380
+ checksum = "d640282b302c0bb0a2a8e0233ead9035e3bed871f0b7e81fe4a1ec829765db44"
381
381
  dependencies = [
382
382
  "libc",
383
383
  ]
@@ -736,9 +736,9 @@ checksum = "67b1b7a3b5fe4f1376887184045fcf45c69e92af734b7aaddc05fb777b6fbd03"
736
736
 
737
737
  [[package]]
738
738
  name = "syn"
739
- version = "2.0.103"
739
+ version = "2.0.104"
740
740
  source = "registry+https://github.com/rust-lang/crates.io-index"
741
- checksum = "e4307e30089d6fd6aff212f2da3a1f9e32f3223b1f010fb09b7c95f90f3ca1e8"
741
+ checksum = "17b6f705963418cdb9927482fa304bc562ece2fdd4f616084c50b7023b435a40"
742
742
  dependencies = [
743
743
  "proc-macro2",
744
744
  "quote",
@@ -1029,7 +1029,7 @@ dependencies = [
1029
1029
 
1030
1030
  [[package]]
1031
1031
  name = "zeusdb-vector-database"
1032
- version = "0.0.1"
1032
+ version = "0.0.2"
1033
1033
  dependencies = [
1034
1034
  "hnsw_rs",
1035
1035
  "pyo3",
@@ -1,6 +1,6 @@
1
1
  [package]
2
2
  name = "zeusdb-vector-database"
3
- version = "0.0.1"
3
+ version = "0.0.2"
4
4
  edition = "2021"
5
5
  resolver = "2" # <-- Avoid compiling unnecessary features from dependencies.
6
6
 
@@ -10,7 +10,7 @@ name = "zeusdb_vector_database" # <-- This is the name of the compiled Python mo
10
10
  crate-type = ["cdylib"]
11
11
 
12
12
  [dependencies]
13
- pyo3 = { version = "0.25.0", features = ["extension-module"] }
13
+ pyo3 = { version = "0.25.1", features = ["extension-module"] }
14
14
  hnsw_rs = "0.3.2"
15
15
 
16
16
  [profile.release]
@@ -0,0 +1,344 @@
1
+ use pyo3::prelude::*;
2
+ use std::collections::HashMap;
3
+ use hnsw_rs::prelude::{Hnsw, DistCosine};
4
+
5
+ #[pyclass]
6
+ pub struct HNSWIndex {
7
+ dim: usize,
8
+ space: String,
9
+ m: usize,
10
+ ef_construction: usize,
11
+ expected_size: usize,
12
+
13
+ // Index-level metadata
14
+ metadata: HashMap<String, String>,
15
+
16
+ // Vector store
17
+ vectors: HashMap<String, Vec<f32>>,
18
+ vector_metadata: HashMap<String, HashMap<String, String>>,
19
+
20
+ hnsw: Hnsw<'static, f32, DistCosine>, // Actual graph
21
+ id_map: HashMap<String, usize>, // Maps external ID → usize
22
+ rev_map: HashMap<usize, String>, // Maps usize → external ID
23
+ id_counter: usize,
24
+ }
25
+
26
+ #[pymethods]
27
+ impl HNSWIndex {
28
+ #[new]
29
+ fn new(
30
+ dim: usize,
31
+ space: String,
32
+ m: usize,
33
+ ef_construction: usize,
34
+ expected_size: usize
35
+ ) -> PyResult<Self> { // Return PyResult for validation
36
+ // Validate parameters in Rust
37
+ if dim == 0 {
38
+ return Err(PyErr::new::<pyo3::exceptions::PyValueError, _>(
39
+ "dim must be positive"
40
+ ));
41
+ }
42
+ if ef_construction == 0 {
43
+ return Err(PyErr::new::<pyo3::exceptions::PyValueError, _>(
44
+ "ef_construction must be positive"
45
+ ));
46
+ }
47
+ if expected_size == 0 {
48
+ return Err(PyErr::new::<pyo3::exceptions::PyValueError, _>(
49
+ "expected_size must be positive"
50
+ ));
51
+ }
52
+ if m > 256 {
53
+ return Err(PyErr::new::<pyo3::exceptions::PyValueError, _>(
54
+ "M must be less than or equal to 256"
55
+ ));
56
+ }
57
+ if space != "cosine" {
58
+ return Err(PyErr::new::<pyo3::exceptions::PyValueError, _>(
59
+ format!("Unsupported space: {}. Only 'cosine' is supported", space)
60
+ ));
61
+ }
62
+
63
+ // Calculate max_layer as log2(expected_size).ceil()
64
+ let max_layer = (expected_size as f32).log2().ceil() as usize;
65
+ let hnsw = Hnsw::<f32, DistCosine>::new(
66
+ m, // M
67
+ expected_size, // expected number of vectors
68
+ max_layer, // number of layers
69
+ ef_construction, // ef
70
+ DistCosine {}
71
+ );
72
+
73
+ Ok(HNSWIndex {
74
+ dim,
75
+ space,
76
+ m,
77
+ ef_construction,
78
+ expected_size,
79
+ metadata: HashMap::new(),
80
+ vectors: HashMap::new(),
81
+ vector_metadata: HashMap::new(),
82
+ hnsw,
83
+ id_map: HashMap::new(),
84
+ rev_map: HashMap::new(),
85
+ id_counter: 0,
86
+ })
87
+ }
88
+
89
+ /// Adds a vector to the index. Fails if the ID already exists.
90
+ pub fn add_point(&mut self, id: String, vector: Vec<f32>, metadata: Option<HashMap<String, String>>) -> PyResult<()> {
91
+ // Check vector dimension
92
+ if vector.len() != self.dim {
93
+ return Err(PyErr::new::<pyo3::exceptions::PyValueError, _>(format!(
94
+ "Vector dimension mismatch: expected {}, got {}",
95
+ self.dim, vector.len()
96
+ )));
97
+ }
98
+
99
+ // Check for duplicate ID
100
+ if self.vectors.contains_key(&id) {
101
+ return Err(PyErr::new::<pyo3::exceptions::PyValueError, _>(format!(
102
+ "Duplicate ID: '{}' already exists", id
103
+ )));
104
+ }
105
+
106
+ // Assign internal index
107
+ let internal_id = self.id_counter;
108
+ self.id_counter += 1;
109
+
110
+ // Store the vector and mappings
111
+ self.vectors.insert(id.clone(), vector.clone());
112
+ self.id_map.insert(id.clone(), internal_id);
113
+ self.rev_map.insert(internal_id, id.clone());
114
+
115
+ // Store metadata if provided
116
+ if let Some(meta) = metadata {
117
+ self.vector_metadata.insert(id.clone(), meta);
118
+ }
119
+
120
+ // Insert into HNSW using a reference to the stored vector
121
+ let stored_vec = self.vectors.get(&id).unwrap();
122
+ self.hnsw.insert((stored_vec.as_slice(), internal_id));
123
+
124
+ Ok(())
125
+ }
126
+
127
+ /// Add multiple vectors in batch for better performance
128
+ pub fn add_batch(&mut self,
129
+ points: Vec<(String, Vec<f32>, Option<HashMap<String, String>>)>
130
+ ) -> PyResult<HashMap<String, Vec<String>>> {
131
+ let mut errors = Vec::new();
132
+ let mut success_count = 0;
133
+
134
+ for (id, vector, metadata) in points {
135
+ match self.add_point(id.clone(), vector, metadata) {
136
+ Ok(()) => success_count += 1,
137
+ Err(e) => errors.push(format!("ID '{}': {}", id, e)),
138
+ }
139
+ }
140
+
141
+ let mut result = HashMap::new();
142
+ result.insert("success_count".to_string(), vec![success_count.to_string()]);
143
+ result.insert("error_count".to_string(), vec![errors.len().to_string()]);
144
+ result.insert("errors".to_string(), errors);
145
+
146
+ Ok(result)
147
+ }
148
+
149
+ /// Query the index for the k-nearest neighbors of a vector
150
+ #[pyo3(signature = (vector, filter=None, top_k=10, ef_search=None))]
151
+ pub fn query(
152
+ &self,
153
+ vector: Vec<f32>,
154
+ filter: Option<HashMap<String, String>>,
155
+ top_k: usize,
156
+ ef_search: Option<usize>,
157
+ ) -> PyResult<Vec<(String, f32)>> {
158
+ if vector.len() != self.dim {
159
+ return Err(PyErr::new::<pyo3::exceptions::PyValueError, _>(format!(
160
+ "Query vector dimension mismatch: expected {}, got {}",
161
+ self.dim, vector.len()
162
+ )));
163
+ }
164
+
165
+ // Get results from HNSW graph
166
+ let ef = ef_search.unwrap_or_else(|| std::cmp::max(2 * top_k, 100));
167
+ let results = self.hnsw.search(&vector, top_k, ef);
168
+
169
+ let mut filtered_results = Vec::new();
170
+
171
+ for neighbor in results {
172
+ let score = neighbor.distance;
173
+ let internal_id = neighbor.get_origin_id();
174
+
175
+ if let Some(ext_id) = self.rev_map.get(&internal_id) {
176
+ // Apply filtering if provided
177
+ if let Some(ref filter_map) = filter {
178
+ if let Some(meta) = self.vector_metadata.get(ext_id) {
179
+ let mut matches = true;
180
+ for (k, v) in filter_map {
181
+ if meta.get(k) != Some(v) {
182
+ matches = false;
183
+ break;
184
+ }
185
+ }
186
+ if !matches {
187
+ continue;
188
+ }
189
+ } else {
190
+ // No metadata, but filter required - skip
191
+ continue;
192
+ }
193
+ }
194
+ filtered_results.push((ext_id.clone(), score));
195
+ }
196
+ }
197
+
198
+ Ok(filtered_results)
199
+ }
200
+
201
+ /// Search with metadata included in results
202
+ #[pyo3(signature = (vector, filter=None, top_k=10, ef_search=None, include_metadata=false))]
203
+ pub fn search_with_metadata(
204
+ &self,
205
+ vector: Vec<f32>,
206
+ filter: Option<HashMap<String, String>>,
207
+ top_k: usize,
208
+ ef_search: Option<usize>,
209
+ include_metadata: bool,
210
+ ) -> PyResult<Vec<(String, f32, Option<HashMap<String, String>>)>> {
211
+ if vector.len() != self.dim {
212
+ return Err(PyErr::new::<pyo3::exceptions::PyValueError, _>(format!(
213
+ "Query vector dimension mismatch: expected {}, got {}",
214
+ self.dim, vector.len()
215
+ )));
216
+ }
217
+
218
+ let ef = ef_search.unwrap_or_else(|| std::cmp::max(2 * top_k, 100));
219
+ let results = self.hnsw.search(&vector, top_k, ef);
220
+
221
+ let mut filtered_results = Vec::new();
222
+
223
+ for neighbor in results {
224
+ let score = neighbor.distance;
225
+ let internal_id = neighbor.get_origin_id();
226
+
227
+ if let Some(ext_id) = self.rev_map.get(&internal_id) {
228
+ // Apply filtering if provided
229
+ if let Some(ref filter_map) = filter {
230
+ if let Some(meta) = self.vector_metadata.get(ext_id) {
231
+ let mut matches = true;
232
+ for (k, v) in filter_map {
233
+ if meta.get(k) != Some(v) {
234
+ matches = false;
235
+ break;
236
+ }
237
+ }
238
+ if !matches {
239
+ continue;
240
+ }
241
+ } else {
242
+ continue;
243
+ }
244
+ }
245
+
246
+ let metadata = if include_metadata {
247
+ self.vector_metadata.get(ext_id).cloned()
248
+ } else {
249
+ None
250
+ };
251
+
252
+ filtered_results.push((ext_id.clone(), score, metadata));
253
+ }
254
+ }
255
+
256
+ Ok(filtered_results)
257
+ }
258
+
259
+ /// Get vector by ID
260
+ pub fn get_vector(&self, id: String) -> Option<Vec<f32>> {
261
+ self.vectors.get(&id).cloned()
262
+ }
263
+
264
+ /// Get metadata by ID
265
+ pub fn get_vector_metadata(&self, id: String) -> Option<HashMap<String, String>> {
266
+ self.vector_metadata.get(&id).cloned()
267
+ }
268
+
269
+ /// Get comprehensive statistics
270
+ pub fn get_stats(&self) -> HashMap<String, String> {
271
+ let mut stats = HashMap::new();
272
+ stats.insert("total_vectors".to_string(), self.vectors.len().to_string());
273
+ stats.insert("dimension".to_string(), self.dim.to_string());
274
+ stats.insert("space".to_string(), self.space.clone());
275
+ stats.insert("M".to_string(), self.m.to_string());
276
+ stats.insert("ef_construction".to_string(), self.ef_construction.to_string());
277
+ stats.insert("expected_size".to_string(), self.expected_size.to_string());
278
+ stats.insert("index_type".to_string(), "HNSW".to_string());
279
+ stats
280
+ }
281
+
282
+ /// List the first `number` records in the index (ID and metadata).
283
+ #[pyo3(signature = (number=10))]
284
+ pub fn list(&self, number: usize) -> Vec<(String, Option<HashMap<String, String>>)> {
285
+ self.vectors
286
+ .iter()
287
+ .take(number)
288
+ .map(|(id, _vec)| {
289
+ let meta = self.vector_metadata.get(id).cloned();
290
+ (id.clone(), meta)
291
+ })
292
+ .collect()
293
+ }
294
+
295
+ /// Add multiple key-value pairs to index-level metadata
296
+ pub fn add_metadata(&mut self, metadata: HashMap<String, String>) {
297
+ for (key, value) in metadata {
298
+ self.metadata.insert(key, value);
299
+ }
300
+ }
301
+
302
+ /// Get a single index-level metadata value
303
+ pub fn get_metadata(&self, key: String) -> Option<String> {
304
+ self.metadata.get(&key).cloned()
305
+ }
306
+
307
+ /// Get all index-level metadata
308
+ pub fn get_all_metadata(&self) -> HashMap<String, String> {
309
+ self.metadata.clone()
310
+ }
311
+
312
+ /// Returns basic info about the index
313
+ pub fn info(&self) -> String {
314
+ format!(
315
+ "HNSWIndex(dim={}, space={}, M={}, ef_construction={}, expected_size={}, vectors={})",
316
+ self.dim,
317
+ self.space,
318
+ self.m,
319
+ self.ef_construction,
320
+ self.expected_size,
321
+ self.vectors.len()
322
+ )
323
+ }
324
+
325
+ /// Check if vector ID exists
326
+ pub fn contains(&self, id: String) -> bool {
327
+ self.vectors.contains_key(&id)
328
+ }
329
+
330
+ /// Remove vector by ID
331
+ pub fn remove_point(&mut self, id: String) -> PyResult<bool> {
332
+ if let Some(internal_id) = self.id_map.remove(&id) {
333
+ self.vectors.remove(&id);
334
+ self.vector_metadata.remove(&id);
335
+ self.rev_map.remove(&internal_id);
336
+ // Note: HNSW doesn't support removal, so the graph still contains the point
337
+ // but it won't be accessible via our mappings
338
+ Ok(true)
339
+ } else {
340
+ Ok(false)
341
+ }
342
+ }
343
+ }
344
+
@@ -1,10 +1,10 @@
1
1
  // lib.rs
2
- mod create_index_hnsw;
2
+ mod hnsw_index;
3
3
 
4
4
  use pyo3::prelude::*;
5
5
 
6
6
  #[pymodule]
7
7
  fn zeusdb_vector_database(m: &Bound<'_, PyModule>) -> PyResult<()> {
8
- m.add_class::<create_index_hnsw::HNSWIndex>()?;
8
+ m.add_class::<hnsw_index::HNSWIndex>()?;
9
9
  Ok(())
10
10
  }
@@ -1,24 +0,0 @@
1
- from .zeusdb_vector_database import HNSWIndex
2
-
3
- def create_index_hnsw(dim: int, space: str, M: int, ef_construction: int, expected_size: int) -> HNSWIndex:
4
- """
5
- Create a new HNSW (Hierarchical Navigable Small World) index using the Rust backend with expected capacity.
6
-
7
- Args:
8
- dim (int): Dimension of the vectors to be indexed.
9
- space (str): Distance metric to use. Only 'cosine' is currently supported.
10
- M (int): Number of bi-directional links created for every new element.
11
- ef_construction (int): Size of the dynamic list for the nearest neighbors during index construction.
12
-
13
- Returns:
14
- HNSWIndex: An instance of the HNSWIndex class representing the created index.
15
-
16
- Raises:
17
- ValueError: If an unsupported distance metric is provided.
18
- """
19
- #if space not in {"cosine", "l2", "dot"}:
20
- if space not in {"cosine"}:
21
- raise ValueError(f"Unsupported space: {space}")
22
- if M > 256:
23
- raise ValueError("M (max_nb_connection) must be less than or equal to 256")
24
- return HNSWIndex(dim, space, M, ef_construction, expected_size)
@@ -1,33 +0,0 @@
1
- from .create_index_hnsw import HNSWIndex, create_index_hnsw
2
-
3
- class VectorDatabase:
4
- def __init__(self):
5
- self.index = None
6
-
7
- def create_index_hnsw(
8
- self,
9
- dim: int = 1536,
10
- space: str = "cosine",
11
- M: int = 16,
12
- ef_construction: int = 200,
13
- expected_size: int = 10000 # Default capacity
14
- ) -> HNSWIndex:
15
- """
16
- Creates a new HNSW (Hierarchical Navigable Small World) index using the specified configuration.
17
-
18
- This method initializes the index for approximate nearest neighbor search using the HNSW algorithm.
19
- It supports configuration of vector dimension, distance metric, connectivity, and construction parameters.
20
-
21
- Args:
22
- dim (int): The number of dimensions for each vector in the index (default is 1536).
23
- space (str): The distance metric to use for similarity, currently only 'cosine' is supported.
24
- M (int): The number of bidirectional links each node maintains in the graph (higher = more accuracy).
25
- ef_construction (int): Size of the dynamic candidate list during index construction (higher = better recall).
26
- expected_size (int): Estimated number of vectors to store; used to preallocate internal data structures (default is 10,000).
27
-
28
- Returns:
29
- HNSWIndex: An initialized HNSWIndex object ready for vector insertion and similarity search.
30
- """
31
- return create_index_hnsw(dim, space, M, ef_construction, expected_size)
32
-
33
-
@@ -1,199 +0,0 @@
1
- use pyo3::prelude::*;
2
- use std::collections::HashMap;
3
- use hnsw_rs::prelude::{Hnsw, DistCosine};
4
-
5
- #[pyclass]
6
- pub struct HNSWIndex {
7
- dim: usize,
8
- space: String,
9
- m: usize,
10
- ef_construction: usize,
11
- expected_size: usize,
12
-
13
- // Index-level metadata
14
- metadata: HashMap<String, String>,
15
-
16
- // Vector store
17
- vectors: HashMap<String, Vec<f32>>,
18
- vector_metadata: HashMap<String, HashMap<String, String>>,
19
-
20
- hnsw: Hnsw<'static, f32, DistCosine>, // Actual graph
21
- id_map: HashMap<String, usize>, // Maps external ID → usize
22
- rev_map: HashMap<usize, String>, // Maps usize → external ID
23
- id_counter: usize,
24
- }
25
-
26
- #[pymethods]
27
- impl HNSWIndex {
28
- #[new]
29
- fn new(
30
- dim: usize,
31
- space: String,
32
- m: usize,
33
- ef_construction: usize,
34
- expected_size: usize
35
- ) -> Self {
36
- // Calculate max_layer as log2(expected_size).ceil()
37
- let max_layer = (expected_size as f32).log2().ceil() as usize;
38
- let hnsw = Hnsw::<f32, DistCosine>::new(
39
- m, // M
40
- expected_size, // expected number of vectors
41
- max_layer, // number of layers
42
- ef_construction, // ef
43
- DistCosine {}
44
- );
45
- // Initialize the HNSW index with the given parameters
46
- HNSWIndex {
47
- dim,
48
- space,
49
- m,
50
- ef_construction,
51
- expected_size,
52
- metadata: HashMap::new(),
53
- vectors: HashMap::new(),
54
- vector_metadata: HashMap::new(),
55
- hnsw,
56
- id_map: HashMap::new(),
57
- rev_map: HashMap::new(),
58
- id_counter: 0,
59
- }
60
- }
61
-
62
- /// Adds a vector to the index. Fails if the ID already exists.
63
- pub fn add_point(&mut self, id: String, vector: Vec<f32>, metadata: Option<HashMap<String, String>>) -> PyResult<()> {
64
- // Check vector dimension
65
- if vector.len() != self.dim {
66
- return Err(PyErr::new::<pyo3::exceptions::PyValueError, _>(format!(
67
- "Vector dimension mismatch: expected {}, got {}",
68
- self.dim, vector.len()
69
- )));
70
- }
71
-
72
- // Check for duplicate ID
73
- if self.vectors.contains_key(&id) {
74
- return Err(PyErr::new::<pyo3::exceptions::PyValueError, _>(format!(
75
- "Duplicate ID: '{}' already exists", id
76
- )));
77
- }
78
-
79
- // Assign internal index
80
- let internal_id = self.id_counter;
81
- self.id_counter += 1;
82
-
83
- // Store the vector and mappings
84
- self.vectors.insert(id.clone(), vector.clone());
85
- self.id_map.insert(id.clone(), internal_id);
86
- self.rev_map.insert(internal_id, id.clone());
87
-
88
- // Store metadata if provided
89
- if let Some(meta) = metadata {
90
- self.vector_metadata.insert(id.clone(), meta);
91
- }
92
-
93
- // Debugging output
94
- //println!("Adding vector with external_id = '{}', internal_id = {}", id, internal_id);
95
-
96
- // Insert into HNSW using a reference to the stored vector
97
- let stored_vec = self.vectors.get(&id).unwrap();
98
- self.hnsw.insert((stored_vec.as_slice(), internal_id));
99
-
100
- Ok(())
101
- }
102
-
103
- /// Query the index for the k-nearest neighbors of a vector
104
- #[pyo3(signature = (vector, filter=None, top_k=10, ef_search=None))]
105
- pub fn query(
106
- &self,
107
- vector: Vec<f32>,
108
- filter: Option<HashMap<String, String>>,
109
- top_k: usize,
110
- ef_search: Option<usize>,
111
- ) -> PyResult<Vec<(String, f32)>> {
112
- if vector.len() != self.dim {
113
- return Err(PyErr::new::<pyo3::exceptions::PyValueError, _>(format!(
114
- "Query vector dimension mismatch: expected {}, got {}",
115
- self.dim, vector.len()
116
- )));
117
- }
118
-
119
- // Get results from HNSW graph
120
- let ef = ef_search.unwrap_or_else(|| std::cmp::max(2 * top_k, 100));
121
- let results = self.hnsw.search(&vector, top_k, ef);
122
-
123
- let mut filtered_results = Vec::new();
124
-
125
- for neighbor in results {
126
- let score = neighbor.distance;
127
- let internal_id = neighbor.get_origin_id();
128
-
129
- // Debugging: resolved internal ID to external ID mapping
130
- //println!("Resolved internal_id {} → {:?}", internal_id, self.rev_map.get(&internal_id));
131
-
132
- if let Some(ext_id) = self.rev_map.get(&internal_id) {
133
- if let Some(ref filter_map) = filter {
134
- let meta = self.vector_metadata.get(ext_id);
135
- if meta.is_none() {
136
- continue;
137
- }
138
- let meta = meta.unwrap();
139
- let mut matches = true;
140
- for (k, v) in filter_map {
141
- if meta.get(k) != Some(v) {
142
- matches = false;
143
- break;
144
- }
145
- }
146
- if !matches {
147
- continue;
148
- }
149
- }
150
- filtered_results.push((ext_id.clone(), score));
151
- }
152
- }
153
-
154
- Ok(filtered_results)
155
- }
156
-
157
- /// List the first `number` records in the index (ID and metadata).
158
- #[pyo3(signature = (number=10))]
159
- pub fn list(&self, number: usize) -> Vec<(String, Option<HashMap<String, String>>)> {
160
- self.vectors
161
- .iter()
162
- .take(number)
163
- .map(|(id, _vec)| {
164
- let meta = self.vector_metadata.get(id).cloned();
165
- (id.clone(), meta)
166
- })
167
- .collect()
168
- }
169
-
170
- /// Add multiple key-value pairs to index-level metadata
171
- pub fn add_metadata(&mut self, metadata: HashMap<String, String>) {
172
- for (key, value) in metadata {
173
- self.metadata.insert(key, value);
174
- }
175
- }
176
-
177
- /// Get a single index-level metadata value
178
- pub fn get_metadata(&self, key: String) -> Option<String> {
179
- self.metadata.get(&key).cloned()
180
- }
181
-
182
- /// Get all index-level metadata
183
- pub fn get_all_metadata(&self) -> HashMap<String, String> {
184
- self.metadata.clone()
185
- }
186
-
187
- /// Returns basic info about the index
188
- pub fn info(&self) -> String {
189
- format!(
190
- "HNSWIndex(dim={}, space={}, M={}, ef_construction={}, expected_size={}, vectors={})",
191
- self.dim,
192
- self.space,
193
- self.m,
194
- self.ef_construction,
195
- self.expected_size,
196
- self.vectors.len()
197
- )
198
- }
199
- }