use crate::metadata::{Metadata, MetadataFilter, MetadataIndex};
use crate::simd::DistanceKernel;
use crate::vector::inverse_norm;
use crate::{IndexConfig, SearchError, SearchHit, VectorIndex};
use std::collections::{BTreeMap, BTreeSet};
use std::path::Path;
use std::sync::{Arc, RwLock};
#[derive(Debug, Clone, PartialEq)]
pub struct VectorRecord {
pub key: u64,
pub vector: Vec<f32>,
pub metadata: Metadata,
}
impl VectorRecord {
#[must_use]
pub fn new(key: u64, vector: Vec<f32>) -> Self {
Self {
key,
vector,
metadata: Metadata::new(),
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum MutationOutcome {
Inserted,
Updated,
}
#[derive(Debug)]
pub struct MutableVectorIndex {
config: IndexConfig,
state: RwLock<MutableState>,
distance_kernel: DistanceKernel,
}
#[derive(Debug)]
struct MutableState {
base: Arc<VectorIndex>,
pending: BTreeMap<u64, Vec<f32>>,
deleted: BTreeSet<u64>,
metadata: MetadataIndex,
generation: u64,
}
impl MutableVectorIndex {
pub fn build(config: IndexConfig, records: &[VectorRecord]) -> Result<Self, SearchError> {
let vectors = records
.iter()
.map(|record| (record.key, record.vector.as_slice()))
.collect::<Vec<_>>();
let base = Arc::new(VectorIndex::build(config.clone(), &vectors)?);
let mut metadata = MetadataIndex::new();
for record in records {
if !record.metadata.is_empty() {
metadata.insert(record.key, record.metadata.clone());
}
}
Ok(Self {
config,
state: RwLock::new(MutableState {
base,
pending: BTreeMap::new(),
deleted: BTreeSet::new(),
metadata,
generation: 0,
}),
distance_kernel: DistanceKernel::detect(),
})
}
#[must_use]
pub fn from_index(index: VectorIndex) -> Self {
let config = index.config().clone();
Self {
config,
state: RwLock::new(MutableState {
base: Arc::new(index),
pending: BTreeMap::new(),
deleted: BTreeSet::new(),
metadata: MetadataIndex::new(),
generation: 0,
}),
distance_kernel: DistanceKernel::detect(),
}
}
pub fn load(path: impl AsRef<Path>) -> Result<Self, SearchError> {
Ok(Self::from_index(VectorIndex::load(path)?))
}
#[must_use]
pub const fn config(&self) -> &IndexConfig {
&self.config
}
#[must_use]
pub fn len(&self) -> usize {
let state = self
.state
.read()
.unwrap_or_else(std::sync::PoisonError::into_inner);
current_len(&state)
}
#[must_use]
pub fn is_empty(&self) -> bool {
self.len() == 0
}
#[must_use]
pub fn delta_len(&self) -> usize {
let state = self
.state
.read()
.unwrap_or_else(std::sync::PoisonError::into_inner);
state.pending.len().saturating_add(state.deleted.len())
}
#[must_use]
pub fn should_compact(&self, maximum_delta: usize) -> bool {
self.delta_len() >= maximum_delta
}
pub fn insert(&self, key: u64, vector: &[f32], metadata: Metadata) -> Result<(), SearchError> {
let normalized = normalize(self.config.dimensions, vector)?;
let mut state = self
.state
.write()
.unwrap_or_else(std::sync::PoisonError::into_inner);
if key_exists(&state, key) {
return Err(SearchError::DuplicateKey(key));
}
state.deleted.remove(&key);
state.pending.insert(key, normalized);
if metadata.is_empty() {
state.metadata.remove(key);
} else {
state.metadata.insert(key, metadata);
}
state.generation = state.generation.wrapping_add(1);
Ok(())
}
pub fn upsert(
&self,
key: u64,
vector: &[f32],
metadata: Metadata,
) -> Result<MutationOutcome, SearchError> {
let normalized = normalize(self.config.dimensions, vector)?;
let mut state = self
.state
.write()
.unwrap_or_else(std::sync::PoisonError::into_inner);
let outcome = if key_exists(&state, key) {
MutationOutcome::Updated
} else {
MutationOutcome::Inserted
};
state.deleted.remove(&key);
state.pending.insert(key, normalized);
if metadata.is_empty() {
state.metadata.remove(key);
} else {
state.metadata.insert(key, metadata);
}
state.generation = state.generation.wrapping_add(1);
Ok(outcome)
}
pub fn delete(&self, key: u64) -> bool {
let mut state = self
.state
.write()
.unwrap_or_else(std::sync::PoisonError::into_inner);
if !key_exists(&state, key) {
return false;
}
state.pending.remove(&key);
if state.base.vector(key).is_some() {
state.deleted.insert(key);
}
state.metadata.remove(key);
state.generation = state.generation.wrapping_add(1);
true
}
pub fn set_metadata(&self, key: u64, metadata: Metadata) -> Result<(), SearchError> {
let mut state = self
.state
.write()
.unwrap_or_else(std::sync::PoisonError::into_inner);
if !key_exists(&state, key) {
return Err(SearchError::MissingKey(key));
}
if metadata.is_empty() {
state.metadata.remove(key);
} else {
state.metadata.insert(key, metadata);
}
state.generation = state.generation.wrapping_add(1);
Ok(())
}
#[must_use]
pub fn metadata(&self, key: u64) -> Option<Metadata> {
let state = self
.state
.read()
.unwrap_or_else(std::sync::PoisonError::into_inner);
state.metadata.get(key).cloned()
}
pub fn search(&self, query: &[f32], count: usize) -> Result<Vec<SearchHit>, SearchError> {
self.search_where(query, count, |_| true)
}
pub fn search_filtered(
&self,
query: &[f32],
count: usize,
filter: &MetadataFilter,
) -> Result<Vec<SearchHit>, SearchError> {
let state = self
.state
.read()
.unwrap_or_else(std::sync::PoisonError::into_inner);
self.search_locked(&state, query, count, |key| {
state.metadata.matches(key, filter)
})
}
pub fn compact(&self) -> Result<(), SearchError> {
for _ in 0..3 {
let (generation, base, pending, deleted) = {
let state = self
.state
.read()
.unwrap_or_else(std::sync::PoisonError::into_inner);
(
state.generation,
Arc::clone(&state.base),
state.pending.clone(),
state.deleted.clone(),
)
};
let mut owned = Vec::new();
owned
.try_reserve_exact(
base.len()
.saturating_sub(deleted.len())
.saturating_add(pending.len()),
)
.map_err(|_| SearchError::AllocationFailed)?;
for key in base.keys() {
if deleted.contains(&key) || pending.contains_key(&key) {
continue;
}
let vector = base.vector(key).ok_or(SearchError::MissingKey(key))?;
owned.push((key, vector.to_vec()));
}
owned.extend(pending.iter().map(|(key, vector)| (*key, vector.clone())));
owned.sort_unstable_by_key(|record| record.0);
let borrowed = owned
.iter()
.map(|(key, vector)| (*key, vector.as_slice()))
.collect::<Vec<_>>();
let rebuilt = Arc::new(VectorIndex::build(self.config.clone(), &borrowed)?);
let mut state = self
.state
.write()
.unwrap_or_else(std::sync::PoisonError::into_inner);
if state.generation != generation {
continue;
}
state.base = rebuilt;
state.pending.clear();
state.deleted.clear();
return Ok(());
}
Err(SearchError::MutationConflict)
}
pub fn save(&self, path: impl AsRef<Path>) -> Result<(), SearchError> {
self.compact()?;
let state = self
.state
.read()
.unwrap_or_else(std::sync::PoisonError::into_inner);
state.base.save(path)
}
fn search_where<F>(
&self,
query: &[f32],
count: usize,
accepts: F,
) -> Result<Vec<SearchHit>, SearchError>
where
F: Fn(u64) -> bool,
{
let state = self
.state
.read()
.unwrap_or_else(std::sync::PoisonError::into_inner);
self.search_locked(&state, query, count, accepts)
}
fn search_locked<F>(
&self,
state: &MutableState,
query: &[f32],
count: usize,
accepts: F,
) -> Result<Vec<SearchHit>, SearchError>
where
F: Fn(u64) -> bool,
{
if query.len() != self.config.dimensions {
return Err(SearchError::DimensionMismatch {
expected: self.config.dimensions,
actual: query.len(),
vector: None,
});
}
let query_inverse_norm = inverse_norm(query, None)?;
let live_count = current_len(state);
let limit = count.min(live_count);
if limit == 0 {
return Ok(Vec::new());
}
let mut hits = state.base.search_filtered(query, limit, |key| {
!state.deleted.contains(&key) && !state.pending.contains_key(&key) && accepts(key)
})?;
hits.try_reserve(state.pending.len())
.map_err(|_| SearchError::AllocationFailed)?;
hits.extend(
state
.pending
.iter()
.filter(|(key, _)| accepts(**key))
.map(|(key, vector)| SearchHit {
key: *key,
distance: self.distance_kernel.cosine_distance(
vector,
query,
query_inverse_norm,
),
}),
);
hits.sort_unstable_by(|left, right| {
left.distance
.total_cmp(&right.distance)
.then_with(|| left.key.cmp(&right.key))
});
hits.dedup_by_key(|hit| hit.key);
hits.truncate(limit);
Ok(hits)
}
}
fn normalize(dimensions: usize, vector: &[f32]) -> Result<Vec<f32>, SearchError> {
if vector.len() != dimensions {
return Err(SearchError::DimensionMismatch {
expected: dimensions,
actual: vector.len(),
vector: Some(0),
});
}
let inverse = inverse_norm(vector, Some(0))?;
let mut normalized = Vec::new();
normalized
.try_reserve_exact(dimensions)
.map_err(|_| SearchError::AllocationFailed)?;
normalized.extend(vector.iter().map(|value| value * inverse));
Ok(normalized)
}
fn key_exists(state: &MutableState, key: u64) -> bool {
state.pending.contains_key(&key)
|| (!state.deleted.contains(&key) && state.base.vector(key).is_some())
}
fn current_len(state: &MutableState) -> usize {
let retained_base = state
.base
.len()
.saturating_sub(state.deleted.len())
.saturating_sub(
state
.pending
.keys()
.filter(|key| state.base.vector(**key).is_some())
.count(),
);
retained_base.saturating_add(state.pending.len())
}