use crate::distance::DistanceMetric;
use parking_lot::RwLock;
use std::collections::HashSet;
use std::sync::atomic::{AtomicU8, Ordering};
const INACTIVE: u8 = 0;
const ACTIVE: u8 = 1;
const DRAINING: u8 = 2;
#[derive(Debug, Clone, Copy, PartialEq, Eq, thiserror::Error)]
#[non_exhaustive]
pub enum ActivateError {
#[error("delta buffer is already active or draining; double-activation rejected")]
AlreadyActive,
}
pub struct DeltaBuffer {
points: RwLock<Vec<(u64, Vec<f32>)>>,
state: AtomicU8,
}
impl DeltaBuffer {
#[must_use]
pub fn new() -> Self {
Self {
points: RwLock::new(Vec::new()),
state: AtomicU8::new(INACTIVE),
}
}
#[must_use]
pub fn is_active(&self) -> bool {
self.state.load(Ordering::Acquire) == ACTIVE
}
#[must_use]
pub fn is_searchable(&self) -> bool {
let s = self.state.load(Ordering::Acquire);
s == ACTIVE || s == DRAINING
}
pub fn activate(&self) {
self.state.store(ACTIVE, Ordering::Release);
}
pub fn try_activate(&self) -> Result<(), ActivateError> {
self.state
.compare_exchange(INACTIVE, ACTIVE, Ordering::AcqRel, Ordering::Acquire)
.map(|_| ())
.map_err(|_| ActivateError::AlreadyActive)
}
pub fn deactivate_and_drain(&self) -> Vec<(u64, Vec<f32>)> {
self.state.store(DRAINING, Ordering::Release);
let mut points = self.points.write();
let drained = std::mem::take(&mut *points);
self.state.store(INACTIVE, Ordering::Release);
drop(points);
drained
}
pub fn push(&self, id: u64, vector: Vec<f32>) {
let mut points = self.points.write();
if self.state.load(Ordering::Acquire) == ACTIVE {
points.retain(|(existing_id, _)| *existing_id != id);
points.push((id, vector));
}
}
pub fn extend(&self, entries: impl IntoIterator<Item = (u64, Vec<f32>)>) {
let new_entries: Vec<(u64, Vec<f32>)> = entries.into_iter().collect();
if new_entries.is_empty() {
return;
}
let new_ids: HashSet<u64> = new_entries.iter().map(|(id, _)| *id).collect();
let mut points = self.points.write();
if self.state.load(Ordering::Acquire) == ACTIVE {
points.retain(|(existing_id, _)| !new_ids.contains(existing_id));
points.extend(new_entries);
}
}
pub fn remove(&self, id: u64) {
self.points.write().retain(|(eid, _)| *eid != id);
}
#[must_use]
pub fn len(&self) -> usize {
self.points.read().len()
}
#[must_use]
pub fn is_empty(&self) -> bool {
self.len() == 0
}
#[must_use]
pub fn stats(&self) -> (usize, bool) {
let len = self.points.read().len();
(len, len == 0)
}
#[must_use]
pub fn search(&self, query: &[f32], k: usize, metric: DistanceMetric) -> Vec<(u64, f32)> {
let current_state = self.state.load(Ordering::Acquire);
if current_state != ACTIVE && current_state != DRAINING {
return Vec::new();
}
let mut results: Vec<(u64, f32)> = {
let points = self.points.read();
if points.is_empty() {
return Vec::new();
}
points
.iter()
.map(|(id, vec)| (*id, metric.calculate(query, vec)))
.collect()
};
metric.sort_results(&mut results);
results.truncate(k);
results
}
}
impl Default for DeltaBuffer {
fn default() -> Self {
Self::new()
}
}
#[cfg(test)]
#[path = "delta_unit_tests.rs"]
mod tests;