Skip to main content

a3s_memory/vector/
in_memory.rs

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/// Exact, session-ephemeral vector index backed by immutable partition blocks.
15#[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(&current, 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(&current, 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}