1use super::search::search_snapshot;
2use super::{
3 VectorBudgetResource, VectorIndex, VectorIndexChangeToken, VectorIndexDescriptor,
4 VectorIndexError, VectorIndexObservation, VectorIndexStatus, VectorMutationConsistency,
5 VectorNormalization, VectorRecord, VectorResult, VectorRevision, VectorSearchRequest,
6 VectorSearchResult,
7};
8use sha2::{Digest, Sha256};
9use std::collections::{BTreeMap, BTreeSet};
10use std::sync::{Arc, RwLock, RwLockReadGuard, RwLockWriteGuard};
11
12const HISTORY_DIGEST_DOMAIN: &str = "a3s.memory.vector-index-history.v1";
13
14#[derive(Clone)]
16pub struct InMemoryVectorIndex {
17 inner: Arc<IndexInner>,
18}
19
20impl std::fmt::Debug for InMemoryVectorIndex {
21 fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
22 formatter
23 .debug_struct("InMemoryVectorIndex")
24 .field("descriptor", &self.inner.descriptor)
25 .field("status", &self.status())
26 .finish()
27 }
28}
29
30struct IndexInner {
31 descriptor: VectorIndexDescriptor,
32 initial_change_token: VectorIndexChangeToken,
33 snapshot: RwLock<Arc<IndexSnapshot>>,
34}
35
36#[derive(Default)]
37pub(super) struct IndexSnapshot {
38 pub(super) revision: VectorRevision,
39 pub(super) partitions: BTreeMap<String, Arc<PartitionBlock>>,
40 pub(super) record_count: usize,
41 pub(super) byte_count: usize,
42}
43
44pub(super) struct PartitionBlock {
45 pub(super) name: String,
46 pub(super) ids: Vec<String>,
47 pub(super) labels: Vec<BTreeMap<String, String>>,
48 pub(super) vectors: Vec<f32>,
49 pub(super) byte_count: usize,
50}
51
52impl PartitionBlock {
53 pub(super) fn record_count(&self) -> usize {
54 self.ids.len()
55 }
56}
57
58impl IndexSnapshot {
59 pub(super) fn status(&self) -> VectorIndexStatus {
60 VectorIndexStatus {
61 revision: self.revision,
62 partition_count: self.partitions.len(),
63 record_count: self.record_count,
64 byte_count: self.byte_count,
65 }
66 }
67}
68
69impl InMemoryVectorIndex {
70 pub fn new(descriptor: VectorIndexDescriptor) -> VectorResult<Self> {
71 descriptor.validate()?;
72 let initial_change_token =
73 VectorIndexChangeToken::try_new(new_history_digest(), VectorRevision::default())?;
74 Ok(Self {
75 inner: Arc::new(IndexInner {
76 descriptor,
77 initial_change_token,
78 snapshot: RwLock::new(Arc::new(IndexSnapshot::default())),
79 }),
80 })
81 }
82
83 fn snapshot(&self) -> Arc<IndexSnapshot> {
84 read_unpoisoned(&self.inner.snapshot).clone()
85 }
86}
87
88pub(super) fn new_history_digest() -> String {
89 let mut hasher = Sha256::new();
90 hasher.update(HISTORY_DIGEST_DOMAIN.as_bytes());
91 hasher.update([0]);
92 hasher.update(uuid::Uuid::new_v4().as_bytes());
93 format!("sha256:{:x}", hasher.finalize())
94}
95
96#[async_trait::async_trait]
97impl VectorIndex for InMemoryVectorIndex {
98 fn descriptor(&self) -> &VectorIndexDescriptor {
99 &self.inner.descriptor
100 }
101
102 fn status(&self) -> VectorIndexStatus {
103 self.snapshot().status()
104 }
105
106 fn change_token(&self) -> Option<VectorIndexChangeToken> {
107 let snapshot = self.snapshot();
108 Some(
109 self.inner
110 .initial_change_token
111 .with_revision(snapshot.revision),
112 )
113 }
114
115 async fn observe(&self) -> VectorResult<VectorIndexObservation> {
116 let snapshot = self.snapshot();
117 let observation = VectorIndexObservation {
118 status: snapshot.status(),
119 change_token: Some(
120 self.inner
121 .initial_change_token
122 .with_revision(snapshot.revision),
123 ),
124 };
125 observation.verify()?;
126 Ok(observation)
127 }
128
129 fn mutation_consistency(&self) -> VectorMutationConsistency {
130 VectorMutationConsistency::IndexRevisionCas
131 }
132
133 async fn replace_partition(
134 &self,
135 partition: &str,
136 records: Vec<VectorRecord>,
137 ) -> VectorResult<VectorIndexStatus> {
138 let partition = validate_partition(partition)?.to_string();
139 let inner = Arc::clone(&self.inner);
140 run_blocking(move || {
141 let block = build_partition(&inner.descriptor, partition, records)?;
142 publish_partition(&inner, block, None)
143 })
144 .await
145 }
146
147 async fn replace_partition_if_revision(
148 &self,
149 partition: &str,
150 expected_revision: VectorRevision,
151 records: Vec<VectorRecord>,
152 ) -> VectorResult<VectorIndexStatus> {
153 let partition = validate_partition(partition)?.to_string();
154 let inner = Arc::clone(&self.inner);
155 run_blocking(move || {
156 let block = build_partition(&inner.descriptor, partition, records)?;
157 publish_partition(&inner, block, Some(expected_revision))
158 })
159 .await
160 }
161
162 async fn remove_partition(&self, partition: &str) -> VectorResult<VectorIndexStatus> {
163 let partition = validate_partition(partition)?.to_string();
164 let inner = Arc::clone(&self.inner);
165 run_blocking(move || remove_partition(&inner, &partition, None)).await
166 }
167
168 async fn remove_partition_if_revision(
169 &self,
170 partition: &str,
171 expected_revision: VectorRevision,
172 ) -> VectorResult<VectorIndexStatus> {
173 let partition = validate_partition(partition)?.to_string();
174 let inner = Arc::clone(&self.inner);
175 run_blocking(move || remove_partition(&inner, &partition, Some(expected_revision))).await
176 }
177
178 async fn search(&self, mut request: VectorSearchRequest) -> VectorResult<VectorSearchResult> {
179 validate_request_filters(&request)?;
180 if request.limit == 0 {
181 return Err(VectorIndexError::InvalidRequest(
182 "limit must be greater than zero".to_string(),
183 ));
184 }
185 let query = prepare_vector(
186 std::mem::take(&mut request.embedding),
187 &self.inner.descriptor,
188 "query".to_string(),
189 )?;
190 let descriptor = self.inner.descriptor.clone();
191 let snapshot = self.snapshot();
192 run_blocking(move || search_snapshot(snapshot, &descriptor, query, request)).await
193 }
194
195 async fn clear(&self) -> VectorResult<VectorIndexStatus> {
196 let inner = Arc::clone(&self.inner);
197 run_blocking(move || clear_index(&inner)).await
198 }
199}
200
201async fn run_blocking<T, F>(operation: F) -> VectorResult<T>
202where
203 T: Send + 'static,
204 F: FnOnce() -> VectorResult<T> + Send + 'static,
205{
206 tokio::task::spawn_blocking(operation)
207 .await
208 .map_err(|error| VectorIndexError::WorkerFailed(error.to_string()))?
209}
210
211pub(super) fn validate_partition(partition: &str) -> VectorResult<&str> {
212 let partition = partition.trim();
213 if partition.is_empty() {
214 Err(VectorIndexError::InvalidPartition)
215 } else {
216 Ok(partition)
217 }
218}
219
220pub(super) fn validate_request_filters(request: &VectorSearchRequest) -> VectorResult<()> {
221 if request
222 .partitions
223 .iter()
224 .any(|partition| partition.trim().is_empty())
225 {
226 return Err(VectorIndexError::InvalidPartition);
227 }
228 if request.labels.keys().any(|key| key.trim().is_empty()) {
229 return Err(VectorIndexError::InvalidLabel {
230 context: "query filter".to_string(),
231 });
232 }
233 Ok(())
234}
235
236pub(super) fn build_partition(
237 descriptor: &VectorIndexDescriptor,
238 name: String,
239 records: Vec<VectorRecord>,
240) -> VectorResult<Arc<PartitionBlock>> {
241 if records.len() > descriptor.max_records {
242 return Err(VectorIndexError::BudgetExceeded {
243 resource: VectorBudgetResource::Records,
244 limit: descriptor.max_records,
245 required: records.len(),
246 });
247 }
248 let minimum_vector_bytes = records
249 .len()
250 .checked_mul(descriptor.dimension)
251 .and_then(|elements| elements.checked_mul(std::mem::size_of::<f32>()))
252 .ok_or(VectorIndexError::SizeOverflow)?;
253 if minimum_vector_bytes > descriptor.max_bytes {
254 return Err(VectorIndexError::BudgetExceeded {
255 resource: VectorBudgetResource::Bytes,
256 limit: descriptor.max_bytes,
257 required: minimum_vector_bytes,
258 });
259 }
260 let mut seen = BTreeSet::new();
261 let mut byte_count = name.len();
262
263 for (record_index, record) in records.iter().enumerate() {
264 if record.id.trim().is_empty() {
265 return Err(VectorIndexError::InvalidRecordId {
266 partition: name.clone(),
267 record_index,
268 });
269 }
270 if !seen.insert(record.id.clone()) {
271 return Err(VectorIndexError::DuplicateRecordId {
272 partition: name.clone(),
273 id: record.id.clone(),
274 });
275 }
276 if record.labels.keys().any(|key| key.trim().is_empty()) {
277 return Err(VectorIndexError::InvalidLabel {
278 context: format!("record '{}' in partition '{name}'", record.id),
279 });
280 }
281 let context = format!("record '{}' in partition '{name}'", record.id);
282 validate_vector(&record.embedding, descriptor, context)?;
283 byte_count = accounted_record_bytes(byte_count, &record.id, &record.labels, descriptor)?;
284 if byte_count > descriptor.max_bytes {
285 return Err(VectorIndexError::BudgetExceeded {
286 resource: VectorBudgetResource::Bytes,
287 limit: descriptor.max_bytes,
288 required: byte_count,
289 });
290 }
291 }
292
293 let vector_capacity = records
294 .len()
295 .checked_mul(descriptor.dimension)
296 .ok_or(VectorIndexError::SizeOverflow)?;
297 let mut ids = Vec::with_capacity(records.len());
298 let mut labels = Vec::with_capacity(records.len());
299 let mut vectors = Vec::with_capacity(vector_capacity);
300 for record in records {
301 let context = format!("record '{}' in partition '{name}'", record.id);
302 let embedding = prepare_vector(record.embedding, descriptor, context)?;
303 ids.push(record.id);
304 labels.push(record.labels);
305 vectors.extend(embedding);
306 }
307
308 Ok(Arc::new(PartitionBlock {
309 name,
310 ids,
311 labels,
312 vectors,
313 byte_count,
314 }))
315}
316
317fn accounted_record_bytes(
318 current: usize,
319 id: &str,
320 labels: &BTreeMap<String, String>,
321 descriptor: &VectorIndexDescriptor,
322) -> VectorResult<usize> {
323 let label_bytes = labels.iter().try_fold(0usize, |total, (key, value)| {
324 total
325 .checked_add(key.len())
326 .and_then(|total| total.checked_add(value.len()))
327 .ok_or(VectorIndexError::SizeOverflow)
328 })?;
329 let vector_bytes = descriptor
330 .dimension
331 .checked_mul(std::mem::size_of::<f32>())
332 .ok_or(VectorIndexError::SizeOverflow)?;
333 current
334 .checked_add(id.len())
335 .and_then(|value| value.checked_add(label_bytes))
336 .and_then(|value| value.checked_add(vector_bytes))
337 .ok_or(VectorIndexError::SizeOverflow)
338}
339
340pub(super) fn prepare_vector(
341 mut vector: Vec<f32>,
342 descriptor: &VectorIndexDescriptor,
343 context: String,
344) -> VectorResult<Vec<f32>> {
345 validate_vector(&vector, descriptor, context.clone())?;
346 if descriptor.normalization == VectorNormalization::Unit {
347 normalize_unit(&mut vector);
348 }
349 Ok(vector)
350}
351
352fn validate_vector(
353 vector: &[f32],
354 descriptor: &VectorIndexDescriptor,
355 context: String,
356) -> VectorResult<()> {
357 if vector.len() != descriptor.dimension {
358 return Err(VectorIndexError::DimensionMismatch {
359 context,
360 expected: descriptor.dimension,
361 actual: vector.len(),
362 });
363 }
364 if let Some(element_index) = vector.iter().position(|value| !value.is_finite()) {
365 return Err(VectorIndexError::NonFiniteVector {
366 context,
367 element_index,
368 });
369 }
370 if descriptor.normalization == VectorNormalization::Unit {
371 let squared_norm = vector.iter().fold(0.0f64, |sum, value| {
372 let value = f64::from(*value);
373 sum + value * value
374 });
375 if squared_norm == 0.0 {
376 return Err(VectorIndexError::ZeroVector { context });
377 }
378 }
379 Ok(())
380}
381
382fn normalize_unit(vector: &mut [f32]) {
383 let norm = vector
384 .iter()
385 .fold(0.0f64, |sum, value| {
386 let value = f64::from(*value);
387 sum + value * value
388 })
389 .sqrt();
390 for value in vector {
391 *value = (f64::from(*value) / norm) as f32;
392 }
393}
394
395fn publish_partition(
396 inner: &IndexInner,
397 block: Arc<PartitionBlock>,
398 expected_revision: Option<VectorRevision>,
399) -> VectorResult<VectorIndexStatus> {
400 let mut published = write_unpoisoned(&inner.snapshot);
401 let current = Arc::clone(&published);
402 verify_expected_revision(¤t, expected_revision)?;
403 let existing = current.partitions.get(&block.name);
404
405 if block.record_count() == 0 && existing.is_none() {
406 return Ok(current.status());
407 }
408
409 let old_records = existing.map_or(0, |partition| partition.record_count());
410 let old_bytes = existing.map_or(0, |partition| partition.byte_count);
411 let record_count = current
412 .record_count
413 .checked_sub(old_records)
414 .and_then(|count| count.checked_add(block.record_count()))
415 .ok_or(VectorIndexError::SizeOverflow)?;
416 let retained_bytes = current
417 .byte_count
418 .checked_sub(old_bytes)
419 .ok_or(VectorIndexError::SizeOverflow)?;
420 let byte_count = if block.record_count() == 0 {
421 retained_bytes
422 } else {
423 retained_bytes
424 .checked_add(block.byte_count)
425 .ok_or(VectorIndexError::SizeOverflow)?
426 };
427 enforce_budgets(&inner.descriptor, record_count, byte_count)?;
428
429 let mut partitions = current.partitions.clone();
430 if block.record_count() == 0 {
431 partitions.remove(&block.name);
432 } else {
433 partitions.insert(block.name.clone(), block);
434 }
435 let next = Arc::new(IndexSnapshot {
436 revision: current.revision.next()?,
437 partitions,
438 record_count,
439 byte_count,
440 });
441 let status = next.status();
442 *published = next;
443 Ok(status)
444}
445
446fn remove_partition(
447 inner: &IndexInner,
448 partition: &str,
449 expected_revision: Option<VectorRevision>,
450) -> VectorResult<VectorIndexStatus> {
451 let mut published = write_unpoisoned(&inner.snapshot);
452 let current = Arc::clone(&published);
453 verify_expected_revision(¤t, expected_revision)?;
454 let Some(existing) = current.partitions.get(partition) else {
455 return Ok(current.status());
456 };
457 let mut partitions = current.partitions.clone();
458 partitions.remove(partition);
459 let next = Arc::new(IndexSnapshot {
460 revision: current.revision.next()?,
461 partitions,
462 record_count: current
463 .record_count
464 .checked_sub(existing.record_count())
465 .ok_or(VectorIndexError::SizeOverflow)?,
466 byte_count: current
467 .byte_count
468 .checked_sub(existing.byte_count)
469 .ok_or(VectorIndexError::SizeOverflow)?,
470 });
471 let status = next.status();
472 *published = next;
473 Ok(status)
474}
475
476fn verify_expected_revision(
477 current: &IndexSnapshot,
478 expected_revision: Option<VectorRevision>,
479) -> VectorResult<()> {
480 if let Some(expected) = expected_revision {
481 if current.revision != expected {
482 return Err(VectorIndexError::RevisionConflict {
483 expected,
484 actual: current.revision,
485 });
486 }
487 }
488 Ok(())
489}
490
491fn clear_index(inner: &IndexInner) -> VectorResult<VectorIndexStatus> {
492 let mut published = write_unpoisoned(&inner.snapshot);
493 let current = Arc::clone(&published);
494 if current.partitions.is_empty() {
495 return Ok(current.status());
496 }
497 let next = Arc::new(IndexSnapshot {
498 revision: current.revision.next()?,
499 ..IndexSnapshot::default()
500 });
501 let status = next.status();
502 *published = next;
503 Ok(status)
504}
505
506pub(super) fn enforce_budgets(
507 descriptor: &VectorIndexDescriptor,
508 record_count: usize,
509 byte_count: usize,
510) -> VectorResult<()> {
511 if record_count > descriptor.max_records {
512 return Err(VectorIndexError::BudgetExceeded {
513 resource: VectorBudgetResource::Records,
514 limit: descriptor.max_records,
515 required: record_count,
516 });
517 }
518 if byte_count > descriptor.max_bytes {
519 return Err(VectorIndexError::BudgetExceeded {
520 resource: VectorBudgetResource::Bytes,
521 limit: descriptor.max_bytes,
522 required: byte_count,
523 });
524 }
525 Ok(())
526}
527
528fn read_unpoisoned<T>(lock: &RwLock<T>) -> RwLockReadGuard<'_, T> {
529 lock.read()
530 .unwrap_or_else(std::sync::PoisonError::into_inner)
531}
532
533fn write_unpoisoned<T>(lock: &RwLock<T>) -> RwLockWriteGuard<'_, T> {
534 lock.write()
535 .unwrap_or_else(std::sync::PoisonError::into_inner)
536}
537
538#[cfg(test)]
539mod lifetime_tests {
540 use super::*;
541
542 #[test]
543 fn last_index_handle_releases_the_complete_index_graph() {
544 let index = InMemoryVectorIndex::new(VectorIndexDescriptor::new(3)).unwrap();
545 let clone = index.clone();
546 let weak = Arc::downgrade(&index.inner);
547
548 drop(index);
549 assert!(weak.upgrade().is_some());
550 drop(clone);
551 assert!(weak.upgrade().is_none());
552 }
553}