use std::collections::hash_map::DefaultHasher;
use std::collections::{HashSet, VecDeque};
use std::hash::{Hash, Hasher};
use std::ops::Range;
use std::sync::Arc;
use ahash::HashMap;
use anyhow::{Result, bail};
use diskann::graph::DiskANNIndex as RawDiskAnnIndex;
use diskann::graph::config::{Builder, MaxDegree, PruneKind};
use diskann::graph::search::Knn;
use diskann::graph::search_output_buffer::IdDistance;
use diskann::provider::{Delete, Guard, SetElement};
use diskann_vector::Half;
use diskann_vector::distance::Metric;
use reblessive::tree::Stk;
use roaring::RoaringTreemap;
use tokio::sync::RwLock;
use crate::catalog::{DiskAnnParams, Distance, TableId, VectorType};
use crate::ctx::{Context, FrozenContext};
use crate::dbs::Options;
use crate::err::Error;
use crate::expr::Cond;
use crate::idx::planner::ScanDirection;
use crate::idx::planner::iterators::KnnIteratorResult;
use crate::idx::trees::diskann::cache::DiskAnnCache;
use crate::idx::trees::diskann::docs::{DiskAnnDocs, DiskAnnVecDocs};
use crate::idx::trees::diskann::filter::DiskAnnTruthyDocumentFilter;
use crate::idx::trees::diskann::provider::{
DiskAnnProvider, DiskAnnProviderContext, DiskAnnStrategy, DiskAnnVectorElement,
};
use crate::idx::trees::diskann::{
DISKANN_PENDING_STATE_SHARDS, DiskAnnPendingState, DiskAnnPendingStateKind,
DiskAnnRecordPendingUpdate, ElementId,
};
use crate::idx::trees::hnsw::VectorId;
use crate::idx::trees::knn::KnnResultBuilder;
use crate::idx::trees::vector::{SerializedVector, Vector};
use crate::idx::{
IndexKeyBase, bump_compaction_generation, is_transaction_condition_not_met,
read_compaction_generation,
};
use crate::key::index::dr::DiskAnnRecordPending;
use crate::key::index::dw::DiskAnnRecordPendingShard;
use crate::kvs::{KVKey, KVValue, Key, Transaction, Val};
use crate::val::{Number, RecordId, RecordIdKey, Value};
const DISKANN_COMPACTION_MAX_PENDING_KEYS: usize = 1024;
const DISKANN_COMPACTION_MAX_PENDING_BYTES: usize = 16 * 1024 * 1024;
struct CapturedPendingKey {
key: Key,
value: Val,
}
#[derive(Clone)]
struct PendingOperation {
id: VectorId,
old_vectors: Vec<SerializedVector>,
new_vectors: Vec<SerializedVector>,
}
type PendingStateSnapshot = Vec<Option<DiskAnnPendingState>>;
pub(crate) struct DiskAnnCompactionPlan {
generation: Option<u64>,
pending_state: PendingStateSnapshot,
captured_keys: Vec<CapturedPendingKey>,
pending: Vec<PendingOperation>,
cleared_shards: Vec<u16>,
has_more: bool,
}
impl DiskAnnCompactionPlan {
pub(crate) fn has_work(&self) -> bool {
!self.captured_keys.is_empty()
}
pub(crate) fn requires_apply(&self) -> bool {
self.has_work()
|| self.pending_state.iter().any(|state| {
state.as_ref().is_some_and(|state| state.kind != DiskAnnPendingStateKind::Empty)
})
}
pub(crate) fn has_more(&self) -> bool {
self.has_more
}
}
struct PendingPlanBuilder {
generation: Option<u64>,
pending_state: PendingStateSnapshot,
captured_keys: Vec<CapturedPendingKey>,
pending: Vec<PendingOperation>,
pending_by_id: HashMap<VectorId, usize>,
encoded_bytes: usize,
has_more: bool,
}
impl PendingPlanBuilder {
fn new(generation: Option<u64>, pending_state: PendingStateSnapshot) -> Self {
Self {
generation,
pending_state,
captured_keys: Vec::new(),
pending: Vec::new(),
pending_by_id: HashMap::default(),
encoded_bytes: 0,
has_more: false,
}
}
fn add(&mut self, key: Key, value: Val, pending: PendingOperation) -> bool {
if self.captured_keys.len() >= DISKANN_COMPACTION_MAX_PENDING_KEYS
|| (!self.captured_keys.is_empty()
&& self.encoded_bytes + key.len() + value.len()
> DISKANN_COMPACTION_MAX_PENDING_BYTES)
{
self.has_more = true;
return false;
}
self.encoded_bytes += key.len() + value.len();
self.captured_keys.push(CapturedPendingKey {
key,
value,
});
self.add_pending(pending);
if self.captured_keys.len() >= DISKANN_COMPACTION_MAX_PENDING_KEYS
|| self.encoded_bytes >= DISKANN_COMPACTION_MAX_PENDING_BYTES
{
self.has_more = true;
}
true
}
fn has_room_for(&self, keys: usize, bytes: usize) -> bool {
if self.captured_keys.len() + keys > DISKANN_COMPACTION_MAX_PENDING_KEYS {
return false;
}
self.captured_keys.is_empty()
|| self.encoded_bytes + bytes <= DISKANN_COMPACTION_MAX_PENDING_BYTES
}
fn add_authorized(&mut self, key: Key, value: Val, pending: PendingOperation) {
self.encoded_bytes += key.len() + value.len();
self.captured_keys.push(CapturedPendingKey {
key,
value,
});
self.add_pending(pending);
if self.captured_keys.len() >= DISKANN_COMPACTION_MAX_PENDING_KEYS
|| self.encoded_bytes >= DISKANN_COMPACTION_MAX_PENDING_BYTES
{
self.has_more = true;
}
}
fn add_pending(&mut self, pending: PendingOperation) {
if let Some(&pos) = self.pending_by_id.get(&pending.id) {
let existing = &mut self.pending[pos];
let existing_precedes = existing.new_vectors == pending.old_vectors;
let pending_precedes = pending.new_vectors == existing.old_vectors;
if pending_precedes {
existing.old_vectors = pending.old_vectors;
} else if existing_precedes {
existing.new_vectors = pending.new_vectors;
} else {
debug_assert!(
existing_precedes || pending_precedes,
"DiskANN pending coalesce: non-chaining entries for {:?} (existing {:?} -> {:?}, incoming {:?} -> {:?})",
existing.id,
existing.old_vectors,
existing.new_vectors,
pending.old_vectors,
pending.new_vectors,
);
existing.new_vectors = pending.new_vectors;
}
return;
}
let pos = self.pending.len();
self.pending_by_id.insert(pending.id.clone(), pos);
self.pending.push(pending);
}
fn into_plan(self, cleared_shards: Vec<u16>) -> DiskAnnCompactionPlan {
DiskAnnCompactionPlan {
generation: self.generation,
pending_state: self.pending_state,
captured_keys: self.captured_keys,
pending: self.pending,
cleared_shards,
has_more: self.has_more,
}
}
}
pub(crate) struct DiskAnnIndex {
dim: usize,
distance: Distance,
ikb: IndexKeyBase,
table_id: TableId,
vector_type: VectorType,
cache: DiskAnnCache,
graph: RwLock<DiskAnnGraph>,
vec_docs: DiskAnnVecDocs,
}
pub(super) struct DiskAnnContext<'a> {
pub(super) ctx: &'a FrozenContext,
pub(super) tx: Arc<Transaction>,
pub(super) ikb: IndexKeyBase,
pub(super) provider_context: DiskAnnProviderContext,
}
impl<'a> DiskAnnContext<'a> {
fn new(
ctx: &'a FrozenContext,
ikb: IndexKeyBase,
provider_context: DiskAnnProviderContext,
) -> Self {
let tx = ctx.tx();
Self {
ctx,
tx,
ikb,
provider_context,
}
}
}
pub(super) struct DiskAnnGraph {
index: RawDiskAnnIndex<DiskAnnProvider>,
}
type DiskAnnSearchResult = (ElementId, f64);
impl DiskAnnGraph {
fn new(ikb: IndexKeyBase, tb: TableId, p: &DiskAnnParams, cache: DiskAnnCache) -> Result<Self> {
let metric = distance_to_metric(&p.distance)?;
let alpha = p.alpha.to_float() as f32;
if !alpha.is_finite() || alpha <= 0.0 {
bail!("DISKANN ALPHA must be finite and greater than 0")
}
let mut builder = Builder::new(
p.degree as usize,
MaxDegree::default_slack(),
p.l_build as usize,
PruneKind::from_metric(metric),
);
builder.alpha(alpha);
let config = builder.build()?;
let provider = DiskAnnProvider::new(ikb, tb, cache, p.dimension as usize, metric);
Ok(Self {
index: RawDiskAnnIndex::new(config, provider, None),
})
}
pub(super) async fn insert(
&mut self,
ctx: &DiskAnnContext<'_>,
vector: Vector,
) -> Result<ElementId> {
match vector {
Vector::F32(values) => {
let Some(values) = values.as_slice() else {
bail!("DISKANN vector storage must be contiguous")
};
self.insert_typed(ctx, values).await
}
Vector::F16(values) => {
let Some(values) = values.as_slice() else {
bail!("DISKANN vector storage must be contiguous")
};
self.insert_typed(ctx, values).await
}
Vector::I8(values) => {
let Some(values) = values.as_slice() else {
bail!("DISKANN vector storage must be contiguous")
};
self.insert_typed(ctx, values).await
}
Vector::U8(values) => {
let Some(values) = values.as_slice() else {
bail!("DISKANN vector storage must be contiguous")
};
self.insert_typed(ctx, values).await
}
_ => bail!("DISKANN supports TYPE F32, F16, I8, and U8"),
}
}
async fn insert_typed<T>(&mut self, ctx: &DiskAnnContext<'_>, values: &[T]) -> Result<ElementId>
where
T: DiskAnnVectorElement,
for<'a> DiskAnnProvider: SetElement<&'a [T], SetError = diskann::ANNError>,
{
let provider = self.index.provider();
let element_id = provider.allocate_element_id(&ctx.provider_context).await?;
if provider.valid_starting_points(&ctx.provider_context).await?.is_empty() {
let guard = provider.set_element(&ctx.provider_context, &element_id, values).await?;
guard.complete().await;
let node: crate::idx::trees::diskann::DiskAnnNode = Default::default();
ctx.tx.set(&ctx.ikb.new_dn_key(element_id), &node).await?;
provider.set_entry_point(&ctx.provider_context, Some(element_id)).await?;
} else {
let strategy = DiskAnnStrategy::<T>::default();
self.index.insert(&strategy, &ctx.provider_context, &element_id, values).await?;
provider.ensure_entry_point(&ctx.provider_context, element_id).await?;
}
Ok(element_id)
}
pub(super) async fn remove(
&mut self,
ctx: &DiskAnnContext<'_>,
element_id: ElementId,
) -> Result<()> {
let provider = self.index.provider();
provider.delete(&ctx.provider_context, &element_id).await?;
let next = provider.valid_starting_points(&ctx.provider_context).await?.into_iter().next();
provider.set_entry_point(&ctx.provider_context, next).await?;
Ok(())
}
async fn search(
&self,
ctx: &DiskAnnContext<'_>,
query: &DiskAnnQuery,
k: usize,
l: usize,
) -> Result<Vec<DiskAnnSearchResult>> {
match query {
DiskAnnQuery::F32(query) => self.search_typed(ctx, query, k, l).await,
DiskAnnQuery::F16(query) => self.search_typed(ctx, query, k, l).await,
DiskAnnQuery::I8(query) => self.search_typed(ctx, query, k, l).await,
DiskAnnQuery::U8(query) => self.search_typed(ctx, query, k, l).await,
}
}
async fn search_typed<T>(
&self,
ctx: &DiskAnnContext<'_>,
query: &[T],
k: usize,
l: usize,
) -> Result<Vec<DiskAnnSearchResult>>
where
T: DiskAnnVectorElement,
{
if self.index.provider().valid_starting_points(&ctx.provider_context).await?.is_empty() {
return Ok(Vec::new());
}
let limit = l.max(k).max(1);
let params = Knn::new_default(limit, limit)?;
let mut ids = vec![0; limit];
let mut distances = vec![0.0; limit];
let mut output = IdDistance::new(&mut ids, &mut distances);
let strategy = DiskAnnStrategy::<T>::default();
let stats =
self.index.search(params, &strategy, &ctx.provider_context, query, &mut output).await?;
let result_count = stats.result_count as usize;
Ok(ids
.into_iter()
.zip(distances)
.take(result_count)
.map(|(id, distance)| (id, distance as f64))
.collect())
}
}
async fn cancel_silently(tx: &Transaction) {
let _ = tx.cancel().await;
}
fn distance_to_metric(distance: &Distance) -> Result<Metric> {
match distance {
Distance::Euclidean => Ok(Metric::L2),
Distance::Cosine => Ok(Metric::Cosine),
Distance::InnerProduct => Ok(Metric::InnerProduct),
Distance::CosineNormalized => Ok(Metric::CosineNormalized),
_ => bail!(
"DISKANN supports EUCLIDEAN, COSINE, INNER_PRODUCT, and COSINE_NORMALIZED distances"
),
}
}
enum DiskAnnQuery {
F32(Vec<f32>),
F16(Vec<Half>),
I8(Vec<i8>),
U8(Vec<u8>),
}
struct DiskAnnSearch {
pt: Vector,
query: DiskAnnQuery,
k: usize,
l: usize,
}
impl DiskAnnSearch {
fn new(pt: Vector, k: usize, l: usize) -> Result<Self> {
let query = match &pt {
Vector::F32(values) => DiskAnnQuery::F32(values.to_vec()),
Vector::F16(values) => DiskAnnQuery::F16(values.to_vec()),
Vector::I8(values) => DiskAnnQuery::I8(values.to_vec()),
Vector::U8(values) => DiskAnnQuery::U8(values.to_vec()),
_ => bail!("DISKANN supports TYPE F32, F16, I8, and U8"),
};
Ok(Self {
query,
pt,
k,
l,
})
}
}
struct DiskAnnGraphSearch<'a, 'b> {
graph: &'a DiskAnnGraph,
search: &'a DiskAnnSearch,
pending_docs: Option<RoaringTreemap>,
filter: &'a mut Option<DiskAnnTruthyDocumentFilter<'b>>,
builder: &'a mut KnnResultBuilder,
}
impl DiskAnnIndex {
pub(crate) async fn new(
ikb: IndexKeyBase,
tb: TableId,
p: &DiskAnnParams,
cache: DiskAnnCache,
) -> Result<Self> {
if !matches!(
p.vector_type,
VectorType::F32 | VectorType::F16 | VectorType::I8 | VectorType::U8
) {
bail!("DISKANN supports TYPE F32, F16, I8, and U8")
}
if matches!(p.distance, Distance::CosineNormalized)
&& matches!(p.vector_type, VectorType::I8 | VectorType::U8)
{
bail!("DISKANN COSINE_NORMALIZED supports TYPE F32 and F16 only")
}
distance_to_metric(&p.distance)?;
Ok(Self {
dim: p.dimension as usize,
vector_type: p.vector_type,
distance: p.distance.clone(),
table_id: tb,
cache: cache.clone(),
graph: RwLock::new(DiskAnnGraph::new(ikb.clone(), tb, p, cache.clone())?),
vec_docs: DiskAnnVecDocs::new(ikb.clone(), tb, cache, p.use_hashed_vector),
ikb,
})
}
fn graph_distance(&self, distance: f64) -> f64 {
match self.distance {
Distance::Euclidean => distance.sqrt(),
_ => distance,
}
}
fn content_to_vectors(&self, content: Vec<Value>) -> Result<Vec<SerializedVector>> {
let mut vectors = Vec::with_capacity(content.len());
for value in content.into_iter().filter(|v| !v.is_nullish()) {
let vector = SerializedVector::try_from_value(self.vector_type, self.dim, value)?;
Vector::check_expected_dimension(vector.dimension(), self.dim)?;
vectors.push(vector);
}
Ok(vectors)
}
fn pending_state_shard(id: &RecordIdKey) -> u16 {
if let RecordIdKey::Number(id) = id {
return id.rem_euclid(i64::from(DISKANN_PENDING_STATE_SHARDS)) as u16;
}
let mut hasher = DefaultHasher::new();
id.hash(&mut hasher);
(hasher.finish() % u64::from(DISKANN_PENDING_STATE_SHARDS)) as u16
}
async fn read_pending_state(
tx: &Transaction,
ikb: &IndexKeyBase,
) -> Result<PendingStateSnapshot> {
let keys: Vec<_> =
(0..DISKANN_PENDING_STATE_SHARDS).map(|shard| ikb.new_dy_key(shard)).collect();
tx.getm(keys, None).await
}
async fn mark_pending_non_empty(
tx: &Transaction,
ikb: &IndexKeyBase,
id: &RecordIdKey,
) -> Result<()> {
let key = ikb.new_dy_key(Self::pending_state_shard(id));
let current: Option<DiskAnnPendingState> = tx.get(&key, None).await?;
if current.as_ref().is_some_and(|state| state.kind == DiskAnnPendingStateKind::NonEmpty) {
return Ok(());
}
let next = DiskAnnPendingState {
kind: DiskAnnPendingStateKind::NonEmpty,
generation: current.as_ref().map_or(0, |state| state.generation).saturating_add(1),
};
tx.putc(&key, &next, current.as_ref()).await
}
async fn clear_pending_state_if_current(
tx: &Transaction,
ikb: &IndexKeyBase,
current: &[Option<DiskAnnPendingState>],
shards: &[u16],
) -> Result<bool> {
let mut changed = false;
for &shard in shards {
let current = current.get(shard as usize).and_then(|state| state.as_ref());
if current.is_some_and(|state| state.kind == DiskAnnPendingStateKind::Empty) {
continue;
}
let key = ikb.new_dy_key(shard);
let kind = match current.map(|state| state.kind) {
Some(DiskAnnPendingStateKind::NonEmpty) => DiskAnnPendingStateKind::MaybeEmpty,
Some(DiskAnnPendingStateKind::MaybeEmpty) | None => DiskAnnPendingStateKind::Empty,
Some(DiskAnnPendingStateKind::Empty) => continue,
};
let next = DiskAnnPendingState {
kind,
generation: current.map_or(0, |state| state.generation.saturating_add(1)),
};
match tx.putc(&key, &next, current).await {
Ok(()) => changed = true,
Err(e) if is_transaction_condition_not_met(&e) => return Ok(false),
Err(e) => return Err(e),
}
}
Ok(changed)
}
async fn range_empty(ctx: &FrozenContext, tx: &Transaction, rng: Range<Key>) -> Result<bool> {
let mut cursor = tx.open_vals_cursor(rng, ScanDirection::Forward, 0, None).await?;
let batch = cursor.next_batch(1).await?;
if !batch.is_empty() {
return Ok(false);
}
drop(cursor);
if ctx.is_done(None).await? {
bail!(Error::QueryCancelled)
}
Ok(true)
}
async fn pending_shard_ranges_empty(
ctx: &FrozenContext,
tx: &Transaction,
ikb: &IndexKeyBase,
shards: &[u16],
) -> Result<bool> {
for &shard in shards {
if !Self::range_empty(ctx, tx, ikb.new_dw_shard_range(shard)?).await? {
return Ok(false);
}
}
Ok(true)
}
pub(crate) async fn index(
&self,
ctx: &Context,
id: &RecordIdKey,
old_values: Option<Vec<Value>>,
new_values: Option<Vec<Value>>,
) -> Result<()> {
if old_values.is_none() && new_values.is_none() {
return Ok(());
}
let old_vectors = if let Some(v) = old_values {
self.content_to_vectors(v)?
} else {
vec![]
};
let new_vectors = if let Some(v) = new_values {
self.content_to_vectors(v)?
} else {
vec![]
};
let tx = ctx.tx();
let shard = Self::pending_state_shard(id);
let key = self.ikb.new_dw_key(shard, id);
let legacy = Self::take_legacy_pending(&tx, &self.ikb, id).await?;
let pending = match (tx.get(&key, None).await?, legacy) {
(Some(mut sharded), None) => {
sharded.new_vectors = new_vectors;
sharded
}
(None, Some(mut legacy)) => {
legacy.new_vectors = new_vectors;
legacy
}
(Some(sharded), Some(legacy)) => {
let sharded_precedes = sharded.new_vectors == legacy.old_vectors;
let legacy_precedes = legacy.new_vectors == sharded.old_vectors;
debug_assert!(
sharded_precedes || legacy_precedes,
"DiskANN write fold: non-chaining dual entries for {id:?} (sharded {:?} -> {:?}, legacy {:?} -> {:?})",
sharded.old_vectors,
sharded.new_vectors,
legacy.old_vectors,
legacy.new_vectors,
);
DiskAnnRecordPendingUpdate {
doc_id: sharded.doc_id.or(legacy.doc_id),
old_vectors: if sharded_precedes {
sharded.old_vectors
} else {
legacy.old_vectors
},
new_vectors,
}
}
(None, None) => DiskAnnRecordPendingUpdate {
doc_id: DiskAnnDocs::get_doc_id(&self.ikb, &tx, id).await?,
old_vectors,
new_vectors,
},
};
tx.set(&key, &pending).await?;
Self::mark_pending_non_empty(&tx, &self.ikb, id).await?;
Ok(())
}
async fn take_legacy_pending(
tx: &Transaction,
ikb: &IndexKeyBase,
id: &RecordIdKey,
) -> Result<Option<DiskAnnRecordPendingUpdate>> {
let legacy_key = ikb.new_dr_key(id);
let Some(legacy) = tx.get(&legacy_key, None).await? else {
return Ok(None);
};
tx.del(&legacy_key).await?;
Ok(Some(legacy))
}
fn record_pending_to_operation(
id: RecordIdKey,
pending: DiskAnnRecordPendingUpdate,
) -> PendingOperation {
let id = if let Some(doc_id) = pending.doc_id {
VectorId::DocId(doc_id)
} else {
VectorId::RecordKey(Arc::new(id))
};
PendingOperation {
id,
old_vectors: pending.old_vectors,
new_vectors: pending.new_vectors,
}
}
fn new_diskann_context<'a>(
&'a self,
ctx: &'a FrozenContext,
provider_context: DiskAnnProviderContext,
) -> DiskAnnContext<'a> {
DiskAnnContext::new(ctx, self.ikb.clone(), provider_context)
}
pub(in crate::idx) async fn prepare_compaction(
ctx: &FrozenContext,
ikb: &IndexKeyBase,
) -> Result<DiskAnnCompactionPlan> {
let tx = ctx.tx();
let generation = read_compaction_generation(&tx, &ikb.new_dg_key()).await?;
let pending_state = Self::read_pending_state(&tx, ikb).await?;
let mut builder = PendingPlanBuilder::new(generation, pending_state.clone());
let mut count = 0;
let mut folded_shard_keys: HashSet<Key> = HashSet::new();
let legacy_drained = Self::capture_legacy_range(
ctx,
&tx,
ikb,
&mut count,
&mut builder,
&mut folded_shard_keys,
)
.await?;
if !legacy_drained {
builder.has_more = true;
return Ok(builder.into_plan(Vec::new()));
}
let mut cleared_shards = Vec::new();
for (shard, state) in pending_state.iter().enumerate() {
if state.as_ref().is_none_or(|s| s.kind == DiskAnnPendingStateKind::Empty) {
continue;
}
let shard = shard as u16;
let drained = Self::capture_shard_range(
ctx,
&tx,
ikb.new_dw_shard_range(shard)?,
&mut count,
&mut builder,
&folded_shard_keys,
)
.await?;
if !drained {
break;
}
cleared_shards.push(shard);
}
Ok(builder.into_plan(cleared_shards))
}
async fn capture_legacy_range(
ctx: &FrozenContext,
tx: &Transaction,
ikb: &IndexKeyBase,
count: &mut usize,
builder: &mut PendingPlanBuilder,
folded_shard_keys: &mut HashSet<Key>,
) -> Result<bool> {
let mut cursor =
tx.open_vals_cursor(ikb.new_dr_range()?, ScanDirection::Forward, 0, None).await?;
loop {
let batch = cursor.next_batch(crate::kvs::NORMAL_BATCH_SIZE).await?;
if batch.is_empty() {
return Ok(true);
}
let owned: Vec<(Vec<u8>, Vec<u8>)> =
batch.iter().map(|(k, v)| (k.to_vec(), v.to_vec())).collect();
for (legacy_key, legacy_value) in owned {
if ctx.is_done(Some(*count)).await? {
bail!(Error::QueryCancelled)
}
let id = DiskAnnRecordPending::decode_key(&legacy_key)?.id.into_owned();
let legacy_update = DiskAnnRecordPendingUpdate::kv_decode_value(&legacy_value, ())?;
let legacy_op = Self::record_pending_to_operation(id.clone(), legacy_update);
let shard_key = ikb.new_dw_key(Self::pending_state_shard(&id), &id);
let shard_entry = match tx.get_raw(&shard_key, None).await? {
Some(shard_value) => {
let update = DiskAnnRecordPendingUpdate::kv_decode_value(&shard_value, ())?;
let op = Self::record_pending_to_operation(id.clone(), update);
Some((shard_key.encode_key()?, shard_value, op))
}
None => None,
};
let pair_keys = 1 + usize::from(shard_entry.is_some());
let pair_bytes = legacy_key.len()
+ legacy_value.len()
+ shard_entry.as_ref().map_or(0, |(k, v, _)| k.len() + v.len());
if !builder.has_room_for(pair_keys, pair_bytes) {
return Ok(false);
}
builder.add_authorized(legacy_key, legacy_value, legacy_op);
if let Some((shard_key_bytes, shard_value, shard_op)) = shard_entry {
folded_shard_keys.insert(shard_key_bytes.clone());
builder.add_authorized(shard_key_bytes, shard_value, shard_op);
}
*count += 1;
if builder.has_more {
return Ok(false);
}
}
}
}
async fn capture_shard_range(
ctx: &FrozenContext,
tx: &Transaction,
rng: Range<Key>,
count: &mut usize,
builder: &mut PendingPlanBuilder,
folded_shard_keys: &HashSet<Key>,
) -> Result<bool> {
let mut cursor = tx.open_vals_cursor(rng, ScanDirection::Forward, 0, None).await?;
loop {
let batch = cursor.next_batch(crate::kvs::NORMAL_BATCH_SIZE).await?;
if batch.is_empty() {
return Ok(true);
}
let owned: Vec<(Vec<u8>, Vec<u8>)> =
batch.iter().map(|(k, v)| (k.to_vec(), v.to_vec())).collect();
for (key, value) in owned {
if ctx.is_done(Some(*count)).await? {
bail!(Error::QueryCancelled)
}
if folded_shard_keys.contains(&key) {
continue;
}
let id = DiskAnnRecordPendingShard::decode_key(&key)?.id.into_owned();
let pending = DiskAnnRecordPendingUpdate::kv_decode_value(&value, ())?;
let pending = Self::record_pending_to_operation(id, pending);
if !builder.add(key, value, pending) {
return Ok(false);
}
*count += 1;
if builder.has_more {
return Ok(false);
}
}
}
}
pub(in crate::idx) async fn apply_compaction(
&self,
ctx: &FrozenContext,
plan: DiskAnnCompactionPlan,
) -> Result<bool> {
let DiskAnnCompactionPlan {
generation,
pending_state,
captured_keys,
pending,
cleared_shards,
has_more: _,
} = plan;
let tx = ctx.tx();
if captured_keys.is_empty() {
if !cleared_shards.is_empty()
&& Self::pending_shard_ranges_empty(ctx, &tx, &self.ikb, &cleared_shards).await?
&& Self::clear_pending_state_if_current(
&tx,
&self.ikb,
&pending_state,
&cleared_shards,
)
.await?
{
return tx.commit().await.map(|()| true);
}
cancel_silently(&tx).await;
return Ok(false);
}
if !bump_compaction_generation(&tx, &self.ikb.new_dg_key(), generation).await? {
cancel_silently(&tx).await;
return Ok(false);
}
for captured in &captured_keys {
match tx.delc(&captured.key, Some(&captured.value)).await {
Ok(()) => {}
Err(e) if is_transaction_condition_not_met(&e) => {
cancel_silently(&tx).await;
return Ok(false);
}
Err(e) => {
cancel_silently(&tx).await;
return Err(e);
}
}
}
let mut graph = self.graph.write().await;
let apply_result: Result<()> = async {
let mut docs = DiskAnnDocs::new(&tx, self.ikb.clone()).await?;
let provider_context = graph.index.provider().context(Arc::clone(&tx));
let diskann_ctx = self.new_diskann_context(ctx, provider_context);
for pending in pending {
self.apply_pending_operation(&diskann_ctx, &mut docs, &mut graph, pending).await?;
}
docs.finish(&tx).await?;
if !cleared_shards.is_empty()
&& Self::pending_shard_ranges_empty(ctx, &tx, &self.ikb, &cleared_shards).await?
{
Self::clear_pending_state_if_current(
&tx,
&self.ikb,
&pending_state,
&cleared_shards,
)
.await?;
}
Ok(())
}
.await;
if let Err(e) = apply_result {
cancel_silently(&tx).await;
self.clear_local_cache().await;
return Err(e);
}
if let Err(e) = tx.commit().await {
self.clear_local_cache().await;
return Err(e);
}
Ok(true)
}
async fn clear_local_cache(&self) {
self.cache
.remove_index(self.ikb.ns(), self.ikb.db(), self.table_id, self.ikb.index())
.await;
}
async fn apply_pending_operation(
&self,
ctx: &DiskAnnContext<'_>,
docs: &mut DiskAnnDocs,
graph: &mut DiskAnnGraph,
pending: PendingOperation,
) -> Result<()> {
match pending.id {
VectorId::DocId(doc_id) => {
for vector in pending.old_vectors {
let vector = Vector::from(vector);
self.vec_docs.remove(ctx, &vector, doc_id, graph).await?;
}
if pending.new_vectors.is_empty() {
docs.remove(&ctx.tx, doc_id, self.table_id, &self.cache).await?;
} else {
for vector in pending.new_vectors {
self.vec_docs.insert(ctx, Vector::from(vector), doc_id, graph).await?;
}
}
}
VectorId::RecordKey(id) => {
if !pending.new_vectors.is_empty() {
let doc_id = docs.resolve(&ctx.tx, &id).await?;
for vector in pending.new_vectors {
self.vec_docs.insert(ctx, Vector::from(vector), doc_id, graph).await?;
}
}
}
}
Ok(())
}
pub(crate) async fn check_state(&self) -> Result<()> {
Ok(())
}
pub(crate) async fn knn_search(
&self,
ctx: &FrozenContext,
stk: &mut Stk,
pt: &[Number],
k: usize,
ef: usize,
cond_filter: Option<(&Options, Arc<Cond>)>,
) -> Result<VecDeque<KnnIteratorResult>> {
let pending_state = Self::read_pending_state(&ctx.tx(), &self.ikb).await?;
let compaction_generation =
read_compaction_generation(&ctx.tx(), &self.ikb.new_dg_key()).await?;
let mut filter = cond_filter.map(|(opt, cond)| {
DiskAnnTruthyDocumentFilter::new(
opt,
self.ikb.clone(),
self.table_id,
self.cache.clone(),
compaction_generation,
cond,
)
});
let vector = Vector::try_from_vector(self.vector_type, pt)?;
vector.check_dimension(self.dim)?;
let search = DiskAnnSearch::new(vector, k, ef)?;
let graph = self.graph.read().await;
let provider_context = graph.index.provider().context(ctx.tx());
let ctx = self.new_diskann_context(ctx, provider_context);
let mut builder = KnnResultBuilder::new(k);
let pending_docs = self
.search_pendings(&ctx, stk, &search, &mut filter, &mut builder, &pending_state)
.await?;
self.search_graph(
&ctx,
stk,
DiskAnnGraphSearch {
graph: &graph,
search: &search,
pending_docs,
filter: &mut filter,
builder: &mut builder,
},
)
.await?;
let result = builder.collect();
let cache = filter.map(DiskAnnTruthyDocumentFilter::release);
let doc_ids: Vec<_> = result
.iter()
.filter_map(|(_, id)| match id {
VectorId::DocId(doc_id) => Some(*doc_id),
VectorId::RecordKey(_) => None,
})
.collect();
let mut doc_rids = DiskAnnDocs::get_things_batch(
&ctx.ikb,
self.table_id,
&self.cache,
&ctx.tx,
&doc_ids,
compaction_generation,
)
.await?
.into_iter();
let mut res = VecDeque::with_capacity(result.len());
for (dist, id) in result {
let dist: f64 = dist.into();
let cached = cache.as_ref().and_then(|cache| cache.get(&id)).cloned();
match id {
VectorId::DocId(_) => {
let rid = doc_rids.next().unwrap_or(None);
if let Some(Some((rid, record))) = cached {
res.push_back((rid, dist, Some(record)));
} else if let Some(rid) = rid {
res.push_back((rid, dist, None));
}
}
VectorId::RecordKey(key) => {
if let Some(Some((rid, record))) = cached {
res.push_back((rid, dist, Some(record)));
continue;
}
let rid = RecordId::new(self.ikb.table().clone(), key.as_ref().clone());
res.push_back((Arc::new(rid), dist, None));
}
}
}
Ok(res)
}
async fn search_graph(
&self,
ctx: &DiskAnnContext<'_>,
stk: &mut Stk,
state: DiskAnnGraphSearch<'_, '_>,
) -> Result<()> {
let results =
state.graph.search(ctx, &state.search.query, state.search.k, state.search.l).await?;
let candidates: Vec<_> = results
.into_iter()
.map(|(element_id, distance)| (element_id, self.graph_distance(distance)))
.filter(|(_, distance)| state.builder.check_add(*distance))
.collect();
if candidates.is_empty() {
return Ok(());
}
let mut docs = self.vec_docs.get_docs_by_element_batch(&ctx.tx, &candidates).await?;
docs.sort_by(|a, b| a.1.total_cmp(&b.1));
let mut idx = 0usize;
let mut window =
(*crate::cnf::DISKANN_FILTER_PREFETCH_MIN_CHUNK).max(state.search.k).max(1);
'windows: while idx < docs.len() {
let end = idx.saturating_add(window).min(docs.len());
let slice = &docs[idx..end];
if let Some(filter) = state.filter.as_mut() {
let mut prefetch_ids: Vec<VectorId> = Vec::new();
for (_, distance, docs) in slice {
if !state.builder.check_add(*distance) {
break;
}
let Some(docs) = docs else {
continue;
};
for doc_id in docs.iter() {
if state
.pending_docs
.as_ref()
.is_some_and(|pending| pending.contains(doc_id))
{
continue;
}
prefetch_ids.push(VectorId::DocId(doc_id));
}
}
filter.prefetch_records(ctx, &prefetch_ids).await?;
}
for (_, distance, docs) in slice {
if !state.builder.check_add(*distance) {
break 'windows;
}
let Some(docs) = docs else {
continue;
};
for doc_id in docs.iter() {
if state.pending_docs.as_ref().is_some_and(|pending| pending.contains(doc_id)) {
continue;
}
let id = VectorId::DocId(doc_id);
if let Some(filter) = state.filter.as_mut()
&& !filter.check_vector_id_truthy(ctx, stk, id.clone()).await?
{
continue;
}
if let Some(evicted_id) = state.builder.add_vector_id_result(*distance, id)
&& let Some(filter) = state.filter.as_mut()
{
filter.expire(&evicted_id);
}
}
}
idx = end;
window = window.saturating_mul(2).min(*crate::cnf::DISKANN_FILTER_PREFETCH_MAX_CHUNK);
}
Ok(())
}
async fn search_pendings(
&self,
ctx: &DiskAnnContext<'_>,
stk: &mut Stk,
search: &DiskAnnSearch,
filter: &mut Option<DiskAnnTruthyDocumentFilter<'_>>,
builder: &mut KnnResultBuilder,
pending_state: &[Option<DiskAnnPendingState>],
) -> Result<Option<RoaringTreemap>> {
let mut all_existing_docs = RoaringTreemap::new();
let mut non_deleted_docs = HashMap::default();
self.collect_pending(ctx.ctx, &ctx.tx, pending_state, |pending| {
if let VectorId::DocId(doc_id) = &pending.id {
all_existing_docs.insert(*doc_id);
};
if pending.new_vectors.is_empty() {
non_deleted_docs.remove(&pending.id);
} else {
non_deleted_docs.insert(pending.id, pending.new_vectors);
}
})
.await?;
if all_existing_docs.is_empty() && non_deleted_docs.is_empty() {
return Ok(None);
}
if let Some(filter) = filter.as_mut() {
let ids: Vec<VectorId> = non_deleted_docs.keys().cloned().collect();
filter.prefetch_records(ctx, &ids).await?;
}
for (id, vectors) in non_deleted_docs {
if let Some(filter) = filter
&& !filter.check_vector_id_truthy(ctx, stk, id.clone()).await?
{
continue;
}
for vector in vectors {
let vector = Vector::from(vector);
let d = self.distance.calculate(&search.pt, &vector);
if builder.check_add(d)
&& let Some(evicted_id) = builder.add_vector_id_result(d, id.clone())
&& let Some(filter) = filter
{
filter.expire(&evicted_id);
}
}
}
if all_existing_docs.is_empty() {
return Ok(None);
}
Ok(Some(all_existing_docs))
}
async fn collect_pending<F>(
&self,
ctx: &Context,
tx: &Transaction,
pending_state: &[Option<DiskAnnPendingState>],
mut collector: F,
) -> Result<()>
where
F: FnMut(PendingOperation),
{
let mut count = 0;
for (shard, state) in pending_state.iter().enumerate() {
if state.as_ref().is_none_or(|s| s.kind == DiskAnnPendingStateKind::Empty) {
continue;
}
let rng = self.ikb.new_dw_shard_range(shard as u16)?;
Self::scan_pending_range(ctx, tx, rng, true, &mut count, &mut collector).await?;
}
let rng = self.ikb.new_dr_range()?;
Self::scan_pending_range(ctx, tx, rng, false, &mut count, &mut collector).await?;
Ok(())
}
async fn scan_pending_range<F>(
ctx: &Context,
tx: &Transaction,
rng: Range<Key>,
sharded: bool,
count: &mut usize,
collector: &mut F,
) -> Result<()>
where
F: FnMut(PendingOperation),
{
let mut cursor = tx.open_vals_cursor(rng, ScanDirection::Forward, 0, None).await?;
loop {
let batch = cursor.next_batch(crate::kvs::NORMAL_BATCH_SIZE).await?;
if batch.is_empty() {
break;
}
for (key, value) in &batch {
if ctx.is_done(Some(*count)).await? {
bail!(Error::QueryCancelled)
}
let id = if sharded {
DiskAnnRecordPendingShard::decode_key(key)?.id.into_owned()
} else {
DiskAnnRecordPending::decode_key(key)?.id.into_owned()
};
let pending = DiskAnnRecordPendingUpdate::kv_decode_value(value, ())?;
collector(Self::record_pending_to_operation(id, pending));
*count += 1;
}
}
Ok(())
}
}
#[cfg(test)]
mod tests {
#[cfg(feature = "kv-rocksdb")]
use temp_dir::TempDir;
use super::*;
use crate::catalog::{DatabaseId, IndexId, NamespaceId};
use crate::idx::trees::diskann::cache::DiskAnnCache;
use crate::kvs::{Datastore, LockType, TransactionType};
fn ikb() -> IndexKeyBase {
IndexKeyBase::new(NamespaceId(1), DatabaseId(2), "tb".into(), IndexId(3))
}
fn cache() -> DiskAnnCache {
DiskAnnCache::new(1024 * 1024)
}
fn params(vector_type: VectorType, distance: Distance) -> DiskAnnParams {
DiskAnnParams {
dimension: 4,
distance,
vector_type,
degree: 16,
l_build: 32,
alpha: 1.2.into(),
use_hashed_vector: false,
}
}
fn diskann_pending_state(kind: DiskAnnPendingStateKind) -> DiskAnnPendingState {
DiskAnnPendingState {
kind,
generation: 0,
}
}
fn diskann_empty_pending_states() -> PendingStateSnapshot {
(0..DISKANN_PENDING_STATE_SHARDS)
.map(|_| Some(diskann_pending_state(DiskAnnPendingStateKind::Empty)))
.collect()
}
fn diskann_compaction_plan(
pending_state: PendingStateSnapshot,
captured_keys: Vec<CapturedPendingKey>,
) -> DiskAnnCompactionPlan {
DiskAnnCompactionPlan {
generation: None,
pending_state,
captured_keys,
pending: Vec::new(),
cleared_shards: Vec::new(),
has_more: false,
}
}
async fn new_ctx(ds: &Datastore, tt: TransactionType) -> FrozenContext {
let tx = Arc::new(ds.transaction(tt, LockType::Optimistic).await.unwrap());
let mut ctx = Context::new_test();
ctx.set_transaction(tx);
ctx.freeze()
}
async fn diskann_pending_states(
tx: &Transaction,
ikb: &IndexKeyBase,
) -> Result<Vec<Option<DiskAnnPendingState>>> {
let keys: Vec<_> =
(0..DISKANN_PENDING_STATE_SHARDS).map(|shard| ikb.new_dy_key(shard)).collect();
tx.getm(keys, None).await
}
fn diskann_any_pending_state_non_empty(states: &[Option<DiskAnnPendingState>]) -> bool {
states.iter().flatten().any(|state| state.kind == DiskAnnPendingStateKind::NonEmpty)
}
fn diskann_any_pending_state_maybe_empty(states: &[Option<DiskAnnPendingState>]) -> bool {
states.iter().flatten().any(|state| state.kind == DiskAnnPendingStateKind::MaybeEmpty)
}
fn diskann_pending_states_require_scan(states: &[Option<DiskAnnPendingState>]) -> bool {
states.iter().any(|state| {
state.as_ref().is_none_or(|state| state.kind != DiskAnnPendingStateKind::Empty)
})
}
fn diskann_all_pending_states_empty(states: &[Option<DiskAnnPendingState>]) -> bool {
states.iter().flatten().all(|state| state.kind == DiskAnnPendingStateKind::Empty)
}
fn f32_value(values: &[f32]) -> Value {
Value::from(values.iter().map(|v| Value::from(*v as f64)).collect::<Vec<_>>())
}
fn f32_content(values: &[f32]) -> Vec<Value> {
vec![f32_value(values)]
}
fn f32_pending(values: &[f32]) -> DiskAnnRecordPendingUpdate {
DiskAnnRecordPendingUpdate {
doc_id: None,
old_vectors: vec![],
new_vectors: vec![SerializedVector::F32(values.to_vec())],
}
}
fn dw_key<'a>(ikb: &'a IndexKeyBase, id: &'a RecordIdKey) -> DiskAnnRecordPendingShard<'a> {
ikb.new_dw_key(DiskAnnIndex::pending_state_shard(id), id)
}
fn f32_query(values: &[f32]) -> Vec<Number> {
values.iter().map(|v| Number::from(*v)).collect()
}
async fn knn_len_with_k(
index: &DiskAnnIndex,
ds: &Datastore,
values: &[f32],
k: usize,
) -> Result<usize> {
let ctx = new_ctx(ds, TransactionType::Read).await;
let query = f32_query(values);
let mut stack = reblessive::tree::TreeStack::new();
let res = stack
.enter(|stk| async { index.knn_search(&ctx, stk, &query, k, 8, None).await })
.finish()
.await?;
ctx.tx().cancel().await?;
Ok(res.len())
}
async fn knn_len(index: &DiskAnnIndex, ds: &Datastore, values: &[f32]) -> Result<usize> {
knn_len_with_k(index, ds, values, 1).await
}
async fn knn_nearest(
index: &DiskAnnIndex,
ds: &Datastore,
values: &[f32],
) -> Result<Option<f64>> {
let ctx = new_ctx(ds, TransactionType::Read).await;
let query = f32_query(values);
let mut stack = reblessive::tree::TreeStack::new();
let res = stack
.enter(|stk| async { index.knn_search(&ctx, stk, &query, 1, 8, None).await })
.finish()
.await?;
ctx.tx().cancel().await?;
Ok(res.front().map(|(_, dist, _)| *dist))
}
async fn compact_once(
index: &DiskAnnIndex,
ds: &Datastore,
ikb: &IndexKeyBase,
) -> Result<bool> {
let plan = {
let ctx = new_ctx(ds, TransactionType::Read).await;
let plan = DiskAnnIndex::prepare_compaction(&ctx, ikb).await?;
ctx.tx().cancel().await?;
plan
};
let ctx = new_ctx(ds, TransactionType::Write).await;
let applied = index.apply_compaction(&ctx, plan).await?;
Ok(applied)
}
fn cached_doc_ids(
cache: &DiskAnnCache,
ikb: &IndexKeyBase,
element_id: ElementId,
) -> Option<Vec<u64>> {
cache
.get_doc_set((ikb.ns(), ikb.db(), TableId(4), ikb.index()), element_id)
.map(|docs| docs.iter().collect())
}
#[tokio::test(flavor = "multi_thread")]
async fn test_diskann_filtered_knn_batches_record_fetches() -> Result<()> {
use crate::catalog::providers::CatalogProvider;
use crate::dbs::{NewPlannerStrategy, Session};
let ds = Arc::new(Datastore::new("memory").await?);
{
let tx = ds.transaction(TransactionType::Write, LockType::Optimistic).await?;
tx.ensure_ns_db(None, "test", "test").await?;
tx.commit().await?;
}
let session = Session::owner()
.with_ns("test")
.with_db("test")
.new_planner_strategy(NewPlannerStrategy::AllReadOnlyStatements);
let n = 500u32;
let cats = 20u32;
let mut setup = String::from(
"DEFINE INDEX emb ON pts FIELDS vec DISKANN DIMENSION 8 DIST EUCLIDEAN TYPE F32;\n",
);
for i in 0..n {
let mut v = String::new();
for j in 0..8u32 {
if j > 0 {
v.push_str(", ");
}
let f =
((i.wrapping_mul(7).wrapping_add(j.wrapping_mul(131))) % 1000) as f32 / 1000.0;
v.push_str(&format!("{f}f"));
}
setup.push_str(&format!("CREATE pts:{i} SET vec = [{v}], category = {};\n", i % cats));
}
for response in ds.execute(&setup, &session, None).await? {
response.result?;
}
async fn run(
ds: &Arc<Datastore>,
session: &Session,
query: &str,
) -> Result<(usize, crate::observe::TransactionMetricsSnapshot)> {
let tx = Arc::new(ds.transaction(TransactionType::Read, LockType::Optimistic).await?);
let mut response =
ds.execute_with_transaction(query, session, None, Arc::clone(&tx)).await?;
let len = match response.remove(0).result? {
surrealdb_types::Value::Array(a) => a.len(),
_ => 0,
};
Ok((len, tx.metrics_snapshot_for_test()))
}
let selective = "SELECT id FROM pts \
WHERE vec <|10,400|> [0.5f,0.5f,0.5f,0.5f,0.5f,0.5f,0.5f,0.5f] AND category = 7;";
let (pending_len, pending_m) = run(&ds, &session, selective).await?;
eprintln!(
"PENDING ops_get={} keys_read={} value_bytes_read={} results={pending_len}",
pending_m.ops_get, pending_m.keys_read, pending_m.value_bytes_read
);
Datastore::index_compaction(
Arc::clone(&ds),
std::time::Duration::from_secs(1),
tokio_util::sync::CancellationToken::new(),
)
.await?;
let (committed_len, committed_m) = run(&ds, &session, selective).await?;
eprintln!(
"COMMITTED ops_get={} keys_read={} value_bytes_read={} results={committed_len}",
committed_m.ops_get, committed_m.keys_read, committed_m.value_bytes_read
);
let nonselective = "SELECT id FROM pts \
WHERE vec <|5,400|> [0.5f,0.5f,0.5f,0.5f,0.5f,0.5f,0.5f,0.5f] AND category < 10;";
let (ns_len, ns_m) = run(&ds, &session, nonselective).await?;
eprintln!(
"NONSELECTIVE ops_get={} keys_read={} value_bytes_read={} results={ns_len}",
ns_m.ops_get, ns_m.keys_read, ns_m.value_bytes_read
);
assert_eq!(pending_len, 10, "pending filtered KNN should return K matches");
assert_eq!(committed_len, 10, "committed filtered KNN should return K matches");
assert!(
u64::from(pending_m.ops_get) * 4 < pending_m.keys_read * 3,
"pending path should batch: ops_get={} keys_read={}",
pending_m.ops_get,
pending_m.keys_read
);
assert!(
u64::from(committed_m.ops_get) * 4 < committed_m.keys_read * 3,
"committed path should batch: ops_get={} keys_read={}",
committed_m.ops_get,
committed_m.keys_read
);
assert_eq!(ns_len, 5, "non-selective filtered KNN should return K matches");
assert!(
ns_m.keys_read * 4 < committed_m.keys_read,
"windowed prefetch should bound non-selective over-fetch: \
non-selective keys_read={} vs selective keys_read={}",
ns_m.keys_read,
committed_m.keys_read
);
Ok(())
}
#[tokio::test(flavor = "multi_thread")]
async fn test_diskann_filtered_knn_skips_missing_record() -> Result<()> {
use crate::catalog::providers::CatalogProvider;
use crate::dbs::{NewPlannerStrategy, Session};
let ds = Arc::new(Datastore::new("memory").await?);
let db_def = {
let tx = ds.transaction(TransactionType::Write, LockType::Optimistic).await?;
let db = tx.ensure_ns_db(None, "test", "test").await?;
tx.commit().await?;
db
};
let session = Session::owner()
.with_ns("test")
.with_db("test")
.new_planner_strategy(NewPlannerStrategy::AllReadOnlyStatements);
let mut setup = String::from(
"DEFINE INDEX pt ON pts FIELDS point DISKANN DIMENSION 1 DIST EUCLIDEAN TYPE F32;\n",
);
for i in 1..=12u32 {
let cat = if i % 2 == 1 {
"a"
} else {
"b"
};
setup.push_str(&format!(
"CREATE pts:{i} SET point = [{}f], category = '{cat}';\n",
i * 10
));
}
for response in ds.execute(&setup, &session, None).await? {
response.result?;
}
Datastore::index_compaction(
Arc::clone(&ds),
std::time::Duration::from_secs(1),
tokio_util::sync::CancellationToken::new(),
)
.await?;
{
let tx = ds.transaction(TransactionType::Write, LockType::Optimistic).await?;
let tb = crate::val::TableName::from("pts");
let key = crate::key::record::new(
db_def.namespace_id,
db_def.database_id,
&tb,
&RecordIdKey::Number(1),
);
tx.del(&key).await?;
tx.commit().await?;
}
let query = "SELECT VALUE vector::distance::knn() FROM pts \
WHERE point <|2,40|> [0f] AND category = 'a';";
let mut dists: Vec<f64> =
ds.execute(query, &session, None).await?.remove(0).result?.into_t::<Vec<f64>>()?;
dists.sort_by(f64::total_cmp);
assert_eq!(
dists,
vec![30.0, 50.0],
"missing pts:1 (dist 10) must be excluded and backfilled, got {dists:?}"
);
Ok(())
}
#[test]
fn diskann_compaction_plan_requires_apply_for_captured_keys() {
let plan = diskann_compaction_plan(
diskann_empty_pending_states(),
vec![CapturedPendingKey {
key: vec![0],
value: vec![1],
}],
);
assert!(plan.has_work());
assert!(plan.requires_apply());
}
#[test]
fn diskann_compaction_plan_skips_apply_when_empty_confirmed() {
let plan = diskann_compaction_plan(diskann_empty_pending_states(), Vec::new());
assert!(!plan.has_work());
assert!(!plan.requires_apply());
}
#[test]
fn diskann_compaction_plan_requires_apply_only_for_non_empty_shards() {
let mut maybe_empty = diskann_empty_pending_states();
maybe_empty[0] = Some(diskann_pending_state(DiskAnnPendingStateKind::MaybeEmpty));
let mut non_empty = diskann_empty_pending_states();
non_empty[0] = Some(diskann_pending_state(DiskAnnPendingStateKind::NonEmpty));
for pending_state in [maybe_empty, non_empty] {
let plan = diskann_compaction_plan(pending_state, Vec::new());
assert!(!plan.has_work());
assert!(plan.requires_apply());
}
let mut missing = diskann_empty_pending_states();
missing[0] = None;
let all_none: PendingStateSnapshot =
(0..DISKANN_PENDING_STATE_SHARDS).map(|_| None).collect();
for pending_state in [missing, all_none] {
let plan = diskann_compaction_plan(pending_state, Vec::new());
assert!(!plan.has_work());
assert!(!plan.requires_apply());
}
}
#[tokio::test]
async fn diskann_accepts_supported_vector_types_and_distances() -> Result<()> {
for (vector_type, distance) in [
(VectorType::F32, Distance::Euclidean),
(VectorType::F16, Distance::CosineNormalized),
(VectorType::U8, Distance::InnerProduct),
(VectorType::I8, Distance::Euclidean),
] {
DiskAnnIndex::new(ikb(), TableId(4), ¶ms(vector_type, distance), cache()).await?;
}
Ok(())
}
#[tokio::test]
async fn diskann_rejects_unsupported_type_metric_combinations() -> Result<()> {
assert!(
DiskAnnIndex::new(
ikb(),
TableId(4),
¶ms(VectorType::I16, Distance::Euclidean),
cache()
)
.await
.is_err()
);
assert!(
DiskAnnIndex::new(
ikb(),
TableId(4),
¶ms(VectorType::U8, Distance::CosineNormalized),
cache()
)
.await
.is_err()
);
assert!(
DiskAnnIndex::new(
ikb(),
TableId(4),
¶ms(VectorType::I8, Distance::CosineNormalized),
cache()
)
.await
.is_err()
);
Ok(())
}
#[tokio::test]
async fn diskann_graph_distance_matches_public_euclidean_distance() -> Result<()> {
let index = DiskAnnIndex::new(
ikb(),
TableId(4),
¶ms(VectorType::F32, Distance::Euclidean),
cache(),
)
.await?;
assert_eq!(index.graph_distance(9.0), 3.0);
Ok(())
}
#[tokio::test]
async fn diskann_doc_set_cache_evicted_and_refilled_for_duplicate_vector_updates() -> Result<()>
{
let ds = Datastore::new("memory").await?;
let ikb = ikb();
let cache = cache();
let index = DiskAnnIndex::new(
ikb.clone(),
TableId(4),
¶ms(VectorType::F32, Distance::Euclidean),
cache.clone(),
)
.await?;
let first_id = RecordIdKey::Number(1);
let second_id = RecordIdKey::Number(2);
let vector = [1.0, 2.0, 3.0, 4.0];
{
let ctx = new_ctx(&ds, TransactionType::Write).await;
index.index(&ctx, &first_id, None, Some(f32_content(&vector))).await?;
index.index(&ctx, &second_id, None, Some(f32_content(&vector))).await?;
ctx.tx().commit().await?;
}
assert!(compact_once(&index, &ds, &ikb).await?);
assert!(cached_doc_ids(&cache, &ikb, 0).is_none());
assert_eq!(knn_len(&index, &ds, &vector).await?, 1);
assert_eq!(cached_doc_ids(&cache, &ikb, 0), Some(vec![0, 1]));
{
let ctx = new_ctx(&ds, TransactionType::Write).await;
index.index(&ctx, &first_id, Some(f32_content(&vector)), None).await?;
ctx.tx().commit().await?;
}
assert!(compact_once(&index, &ds, &ikb).await?);
assert!(cached_doc_ids(&cache, &ikb, 0).is_none());
assert_eq!(knn_len(&index, &ds, &vector).await?, 1);
assert_eq!(cached_doc_ids(&cache, &ikb, 0), Some(vec![1]));
{
let ctx = new_ctx(&ds, TransactionType::Write).await;
index.index(&ctx, &second_id, Some(f32_content(&vector)), None).await?;
ctx.tx().commit().await?;
}
assert!(compact_once(&index, &ds, &ikb).await?);
assert!(cached_doc_ids(&cache, &ikb, 0).is_none());
assert_eq!(knn_len(&index, &ds, &vector).await?, 0);
Ok(())
}
#[tokio::test]
async fn diskann_index_write_marks_pending_state_non_empty() -> Result<()> {
let ds = Datastore::new("memory").await?;
let ikb = ikb();
let index = DiskAnnIndex::new(
ikb.clone(),
TableId(4),
¶ms(VectorType::F32, Distance::Euclidean),
cache(),
)
.await?;
let ctx = new_ctx(&ds, TransactionType::Write).await;
let tx = ctx.tx();
let id = RecordIdKey::Number(1);
index.index(&ctx, &id, None, Some(f32_content(&[1.0, 2.0, 3.0, 4.0]))).await?;
let pending: DiskAnnRecordPendingUpdate = tx.get(&dw_key(&ikb, &id), None).await?.unwrap();
let states = diskann_pending_states(&tx, &ikb).await?;
let state = states
.iter()
.flatten()
.find(|state| state.kind == DiskAnnPendingStateKind::NonEmpty)
.unwrap();
assert!(pending.old_vectors.is_empty());
assert_eq!(pending.new_vectors, vec![SerializedVector::F32(vec![1.0, 2.0, 3.0, 4.0])]);
assert_eq!(state.kind, DiskAnnPendingStateKind::NonEmpty);
assert_eq!(state.generation, 1);
index
.index(
&ctx,
&id,
Some(f32_content(&[1.0, 2.0, 3.0, 4.0])),
Some(f32_content(&[4.0, 3.0, 2.0, 1.0])),
)
.await?;
let pending: DiskAnnRecordPendingUpdate = tx.get(&dw_key(&ikb, &id), None).await?.unwrap();
let updated_states = diskann_pending_states(&tx, &ikb).await?;
let updated_state = updated_states
.iter()
.flatten()
.find(|state| state.kind == DiskAnnPendingStateKind::NonEmpty)
.unwrap();
assert_eq!(pending.new_vectors, vec![SerializedVector::F32(vec![4.0, 3.0, 2.0, 1.0])]);
assert_eq!(updated_state.kind, DiskAnnPendingStateKind::NonEmpty);
assert!(updated_state.generation >= state.generation);
tx.cancel().await?;
Ok(())
}
#[tokio::test]
async fn diskann_lookup_skips_sharded_pendings_only_when_guard_is_empty() -> Result<()> {
let ds = Datastore::new("memory").await?;
let ikb = ikb();
let index = DiskAnnIndex::new(
ikb.clone(),
TableId(4),
¶ms(VectorType::F32, Distance::Euclidean),
cache(),
)
.await?;
let id = RecordIdKey::Number(1);
let shard = DiskAnnIndex::pending_state_shard(&id);
{
let ctx = new_ctx(&ds, TransactionType::Write).await;
let tx = ctx.tx();
tx.set(&ikb.new_dw_key(shard, &id), &f32_pending(&[1.0, 2.0, 3.0, 4.0])).await?;
tx.set(
&ikb.new_dy_key(shard),
&DiskAnnPendingState {
kind: DiskAnnPendingStateKind::NonEmpty,
generation: 1,
},
)
.await?;
tx.commit().await?;
}
assert_eq!(knn_len(&index, &ds, &[1.0, 2.0, 3.0, 4.0]).await?, 1);
{
let ctx = new_ctx(&ds, TransactionType::Write).await;
ctx.tx()
.set(
&ikb.new_dy_key(shard),
&DiskAnnPendingState {
kind: DiskAnnPendingStateKind::MaybeEmpty,
generation: 2,
},
)
.await?;
ctx.tx().commit().await?;
}
assert_eq!(knn_len(&index, &ds, &[1.0, 2.0, 3.0, 4.0]).await?, 1);
{
let ctx = new_ctx(&ds, TransactionType::Write).await;
ctx.tx()
.set(
&ikb.new_dy_key(shard),
&DiskAnnPendingState {
kind: DiskAnnPendingStateKind::Empty,
generation: 3,
},
)
.await?;
ctx.tx().commit().await?;
}
assert_eq!(knn_len(&index, &ds, &[1.0, 2.0, 3.0, 4.0]).await?, 0);
Ok(())
}
#[tokio::test]
async fn diskann_old_compactor_clearing_legacy_guard_keeps_sharded_visible() -> Result<()> {
let ds = Datastore::new("memory").await?;
let ikb = ikb();
let index = DiskAnnIndex::new(
ikb.clone(),
TableId(4),
¶ms(VectorType::F32, Distance::Euclidean),
cache(),
)
.await?;
let id = RecordIdKey::Number(1);
{
let ctx = new_ctx(&ds, TransactionType::Write).await;
index.index(&ctx, &id, None, Some(f32_content(&[1.0, 2.0, 3.0, 4.0]))).await?;
ctx.tx().commit().await?;
}
assert_eq!(knn_len(&index, &ds, &[1.0, 2.0, 3.0, 4.0]).await?, 1);
{
let ctx = new_ctx(&ds, TransactionType::Write).await;
for s in 0..DISKANN_PENDING_STATE_SHARDS {
ctx.tx()
.set(
&ikb.new_dp_key(s),
&DiskAnnPendingState {
kind: DiskAnnPendingStateKind::Empty,
generation: 1,
},
)
.await?;
}
ctx.tx().commit().await?;
}
assert_eq!(knn_len(&index, &ds, &[1.0, 2.0, 3.0, 4.0]).await?, 1);
Ok(())
}
#[tokio::test]
async fn diskann_compaction_clears_pending_state_after_empty_confirmation() -> Result<()> {
let ds = Datastore::new("memory").await?;
let ikb = ikb();
let index = DiskAnnIndex::new(
ikb.clone(),
TableId(4),
¶ms(VectorType::F32, Distance::Euclidean),
cache(),
)
.await?;
let id = RecordIdKey::Number(1);
{
let ctx = new_ctx(&ds, TransactionType::Write).await;
index.index(&ctx, &id, None, Some(f32_content(&[1.0, 2.0, 3.0, 4.0]))).await?;
ctx.tx().commit().await?;
}
let plan = {
let ctx = new_ctx(&ds, TransactionType::Read).await;
let plan = DiskAnnIndex::prepare_compaction(&ctx, &ikb).await?;
ctx.tx().cancel().await?;
plan
};
assert!(plan.has_work());
assert!(!plan.has_more());
{
let ctx = new_ctx(&ds, TransactionType::Write).await;
assert!(index.apply_compaction(&ctx, plan).await?);
}
{
let ctx = new_ctx(&ds, TransactionType::Read).await;
let states = diskann_pending_states(&ctx.tx(), &ikb).await?;
assert!(diskann_pending_states_require_scan(&states));
assert!(diskann_any_pending_state_maybe_empty(&states));
assert!(ctx.tx().get::<_>(&dw_key(&ikb, &id), None).await?.is_none());
ctx.tx().cancel().await?;
}
assert!(compact_once(&index, &ds, &ikb).await?);
{
let ctx = new_ctx(&ds, TransactionType::Read).await;
let states = diskann_pending_states(&ctx.tx(), &ikb).await?;
assert!(diskann_all_pending_states_empty(&states));
assert!(ctx.tx().get::<_>(&dw_key(&ikb, &id), None).await?.is_none());
ctx.tx().cancel().await?;
}
Ok(())
}
#[cfg(feature = "kv-rocksdb")]
#[tokio::test]
async fn diskann_rocksdb_clear_race_keeps_concurrent_writer_visible() -> Result<()> {
let dir = TempDir::new()?;
let path = format!("rocksdb:{}", dir.path().to_string_lossy());
let ds = Datastore::new(&path).await?;
let ikb = ikb();
let index = DiskAnnIndex::new(
ikb.clone(),
TableId(4),
¶ms(VectorType::F32, Distance::Euclidean),
cache(),
)
.await?;
let first_id = RecordIdKey::Number(1);
let second_id = RecordIdKey::Number(1 + i64::from(DISKANN_PENDING_STATE_SHARDS));
assert_eq!(
DiskAnnIndex::pending_state_shard(&first_id),
DiskAnnIndex::pending_state_shard(&second_id)
);
{
let ctx = new_ctx(&ds, TransactionType::Write).await;
index.index(&ctx, &first_id, None, Some(f32_content(&[1.0, 2.0, 3.0, 4.0]))).await?;
ctx.tx().commit().await?;
}
let plan = {
let ctx = new_ctx(&ds, TransactionType::Read).await;
let plan = DiskAnnIndex::prepare_compaction(&ctx, &ikb).await?;
ctx.tx().cancel().await?;
plan
};
assert!(plan.has_work());
assert!(!plan.has_more());
let apply_ctx = new_ctx(&ds, TransactionType::Write).await;
{
let ctx = new_ctx(&ds, TransactionType::Write).await;
index.index(&ctx, &second_id, None, Some(f32_content(&[4.0, 3.0, 2.0, 1.0]))).await?;
ctx.tx().commit().await?;
}
assert!(index.apply_compaction(&apply_ctx, plan).await?);
{
let ctx = new_ctx(&ds, TransactionType::Read).await;
let states = diskann_pending_states(&ctx.tx(), &ikb).await?;
let shard = DiskAnnIndex::pending_state_shard(&second_id) as usize;
assert_eq!(
states[shard].as_ref().map(|state| state.kind),
Some(DiskAnnPendingStateKind::MaybeEmpty)
);
assert!(ctx.tx().get::<_>(&dw_key(&ikb, &second_id), None).await?.is_some());
ctx.tx().cancel().await?;
}
assert_eq!(knn_len_with_k(&index, &ds, &[4.0, 3.0, 2.0, 1.0], 2).await?, 2);
Ok(())
}
#[tokio::test]
async fn diskann_empty_compaction_plan_does_not_clear_concurrent_pending_write() -> Result<()> {
let ds = Datastore::new("memory").await?;
let ikb = ikb();
let index = DiskAnnIndex::new(
ikb.clone(),
TableId(4),
¶ms(VectorType::F32, Distance::Euclidean),
cache(),
)
.await?;
let id = RecordIdKey::Number(1);
let plan = {
let ctx = new_ctx(&ds, TransactionType::Read).await;
let plan = DiskAnnIndex::prepare_compaction(&ctx, &ikb).await?;
ctx.tx().cancel().await?;
plan
};
assert!(!plan.has_work());
{
let ctx = new_ctx(&ds, TransactionType::Write).await;
index.index(&ctx, &id, None, Some(f32_content(&[1.0, 2.0, 3.0, 4.0]))).await?;
ctx.tx().commit().await?;
}
{
let ctx = new_ctx(&ds, TransactionType::Write).await;
assert!(!index.apply_compaction(&ctx, plan).await?);
}
{
let ctx = new_ctx(&ds, TransactionType::Read).await;
let states = diskann_pending_states(&ctx.tx(), &ikb).await?;
assert!(diskann_any_pending_state_non_empty(&states));
assert!(ctx.tx().get::<_>(&dw_key(&ikb, &id), None).await?.is_some());
ctx.tx().cancel().await?;
}
Ok(())
}
#[tokio::test]
async fn diskann_final_compaction_plan_preserves_concurrent_pending_write() -> Result<()> {
let ds = Datastore::new("memory").await?;
let ikb = ikb();
let index = DiskAnnIndex::new(
ikb.clone(),
TableId(4),
¶ms(VectorType::F32, Distance::Euclidean),
cache(),
)
.await?;
let first_id = RecordIdKey::Number(1);
let second_id = RecordIdKey::Number(2);
{
let ctx = new_ctx(&ds, TransactionType::Write).await;
index.index(&ctx, &first_id, None, Some(f32_content(&[1.0, 2.0, 3.0, 4.0]))).await?;
ctx.tx().commit().await?;
}
let plan = {
let ctx = new_ctx(&ds, TransactionType::Read).await;
let plan = DiskAnnIndex::prepare_compaction(&ctx, &ikb).await?;
ctx.tx().cancel().await?;
plan
};
assert!(plan.has_work());
assert!(!plan.has_more());
{
let ctx = new_ctx(&ds, TransactionType::Write).await;
index.index(&ctx, &second_id, None, Some(f32_content(&[4.0, 3.0, 2.0, 1.0]))).await?;
ctx.tx().commit().await?;
}
{
let ctx = new_ctx(&ds, TransactionType::Write).await;
assert!(index.apply_compaction(&ctx, plan).await?);
}
{
let ctx = new_ctx(&ds, TransactionType::Read).await;
let states = diskann_pending_states(&ctx.tx(), &ikb).await?;
assert!(diskann_any_pending_state_non_empty(&states));
assert!(ctx.tx().get::<_>(&dw_key(&ikb, &first_id), None).await?.is_none());
assert!(ctx.tx().get::<_>(&dw_key(&ikb, &second_id), None).await?.is_some());
ctx.tx().cancel().await?;
}
Ok(())
}
#[tokio::test]
async fn diskann_dual_read_drains_legacy_pending_then_uses_shards() -> Result<()> {
let ds = Datastore::new("memory").await?;
let ikb = ikb();
let index = DiskAnnIndex::new(
ikb.clone(),
TableId(4),
¶ms(VectorType::F32, Distance::Euclidean),
cache(),
)
.await?;
let legacy_id = RecordIdKey::Number(1);
{
let ctx = new_ctx(&ds, TransactionType::Write).await;
let tx = ctx.tx();
tx.set(&ikb.new_dr_key(&legacy_id), &f32_pending(&[1.0, 2.0, 3.0, 4.0])).await?;
tx.commit().await?;
}
assert_eq!(knn_len(&index, &ds, &[1.0, 2.0, 3.0, 4.0]).await?, 1);
assert!(compact_once(&index, &ds, &ikb).await?);
{
let ctx = new_ctx(&ds, TransactionType::Read).await;
assert!(ctx.tx().get::<_>(&ikb.new_dr_key(&legacy_id), None).await?.is_none());
ctx.tx().cancel().await?;
}
compact_once(&index, &ds, &ikb).await?;
{
let ctx = new_ctx(&ds, TransactionType::Read).await;
let states = diskann_pending_states(&ctx.tx(), &ikb).await?;
assert!(diskann_all_pending_states_empty(&states));
ctx.tx().cancel().await?;
}
assert_eq!(knn_len(&index, &ds, &[1.0, 2.0, 3.0, 4.0]).await?, 1);
let new_id = RecordIdKey::Number(2);
{
let ctx = new_ctx(&ds, TransactionType::Write).await;
index.index(&ctx, &new_id, None, Some(f32_content(&[4.0, 3.0, 2.0, 1.0]))).await?;
ctx.tx().commit().await?;
}
{
let ctx = new_ctx(&ds, TransactionType::Read).await;
assert!(ctx.tx().get::<_>(&dw_key(&ikb, &new_id), None).await?.is_some());
assert!(ctx.tx().get::<_>(&ikb.new_dr_key(&new_id), None).await?.is_none());
ctx.tx().cancel().await?;
}
assert_eq!(knn_len_with_k(&index, &ds, &[4.0, 3.0, 2.0, 1.0], 2).await?, 2);
Ok(())
}
#[tokio::test]
async fn diskann_write_folds_legacy_pending_into_shard() -> Result<()> {
let ds = Datastore::new("memory").await?;
let ikb = ikb();
let index = DiskAnnIndex::new(
ikb.clone(),
TableId(4),
¶ms(VectorType::F32, Distance::Euclidean),
cache(),
)
.await?;
let id = RecordIdKey::Number(1);
{
let ctx = new_ctx(&ds, TransactionType::Write).await;
let tx = ctx.tx();
tx.set(&ikb.new_dr_key(&id), &f32_pending(&[1.0, 2.0, 3.0, 4.0])).await?;
tx.commit().await?;
}
{
let ctx = new_ctx(&ds, TransactionType::Write).await;
index
.index(
&ctx,
&id,
Some(f32_content(&[1.0, 2.0, 3.0, 4.0])),
Some(f32_content(&[9.0, 9.0, 9.0, 9.0])),
)
.await?;
ctx.tx().commit().await?;
}
{
let ctx = new_ctx(&ds, TransactionType::Read).await;
assert!(ctx.tx().get::<_>(&ikb.new_dr_key(&id), None).await?.is_none());
let folded: DiskAnnRecordPendingUpdate =
ctx.tx().get(&dw_key(&ikb, &id), None).await?.unwrap();
assert_eq!(folded.new_vectors, vec![SerializedVector::F32(vec![9.0, 9.0, 9.0, 9.0])]);
ctx.tx().cancel().await?;
}
assert_eq!(knn_nearest(&index, &ds, &[9.0, 9.0, 9.0, 9.0]).await?, Some(0.0));
Ok(())
}
#[tokio::test]
async fn diskann_write_folds_legacy_even_when_sharded_exists() -> Result<()> {
let ds = Datastore::new("memory").await?;
let ikb = ikb();
let index = DiskAnnIndex::new(
ikb.clone(),
TableId(4),
¶ms(VectorType::F32, Distance::Euclidean),
cache(),
)
.await?;
let id = RecordIdKey::Number(1);
let shard = DiskAnnIndex::pending_state_shard(&id);
{
let ctx = new_ctx(&ds, TransactionType::Write).await;
let tx = ctx.tx();
tx.set(
&ikb.new_dw_key(shard, &id),
&DiskAnnRecordPendingUpdate {
doc_id: None,
old_vectors: vec![SerializedVector::F32(vec![1.0, 1.0, 1.0, 1.0])],
new_vectors: vec![SerializedVector::F32(vec![2.0, 2.0, 2.0, 2.0])],
},
)
.await?;
tx.set(
&ikb.new_dr_key(&id),
&DiskAnnRecordPendingUpdate {
doc_id: None,
old_vectors: vec![SerializedVector::F32(vec![2.0, 2.0, 2.0, 2.0])],
new_vectors: vec![SerializedVector::F32(vec![3.0, 3.0, 3.0, 3.0])],
},
)
.await?;
tx.set(
&ikb.new_dy_key(shard),
&DiskAnnPendingState {
kind: DiskAnnPendingStateKind::NonEmpty,
generation: 1,
},
)
.await?;
tx.commit().await?;
}
{
let ctx = new_ctx(&ds, TransactionType::Write).await;
index
.index(
&ctx,
&id,
Some(f32_content(&[3.0, 3.0, 3.0, 3.0])),
Some(f32_content(&[4.0, 4.0, 4.0, 4.0])),
)
.await?;
ctx.tx().commit().await?;
}
{
let ctx = new_ctx(&ds, TransactionType::Read).await;
assert!(ctx.tx().get::<_>(&ikb.new_dr_key(&id), None).await?.is_none());
let folded: DiskAnnRecordPendingUpdate =
ctx.tx().get(&dw_key(&ikb, &id), None).await?.unwrap();
assert_eq!(folded.old_vectors, vec![SerializedVector::F32(vec![1.0, 1.0, 1.0, 1.0])]);
assert_eq!(folded.new_vectors, vec![SerializedVector::F32(vec![4.0, 4.0, 4.0, 4.0])]);
ctx.tx().cancel().await?;
}
assert_eq!(knn_nearest(&index, &ds, &[4.0, 4.0, 4.0, 4.0]).await?, Some(0.0));
Ok(())
}
#[tokio::test]
async fn diskann_compaction_orders_cross_layout_pending_by_vector_chain() -> Result<()> {
let ds = Datastore::new("memory").await?;
let ikb = ikb();
let index = DiskAnnIndex::new(
ikb.clone(),
TableId(4),
¶ms(VectorType::F32, Distance::Euclidean),
cache(),
)
.await?;
let id = RecordIdKey::Number(1);
let shard = DiskAnnIndex::pending_state_shard(&id);
{
let ctx = new_ctx(&ds, TransactionType::Write).await;
let tx = ctx.tx();
tx.set(
&ikb.new_dw_key(shard, &id),
&DiskAnnRecordPendingUpdate {
doc_id: None,
old_vectors: vec![SerializedVector::F32(vec![10.0, 10.0, 10.0, 10.0])],
new_vectors: vec![SerializedVector::F32(vec![20.0, 20.0, 20.0, 20.0])],
},
)
.await?;
tx.set(
&ikb.new_dr_key(&id),
&DiskAnnRecordPendingUpdate {
doc_id: None,
old_vectors: vec![SerializedVector::F32(vec![20.0, 20.0, 20.0, 20.0])],
new_vectors: vec![SerializedVector::F32(vec![30.0, 30.0, 30.0, 30.0])],
},
)
.await?;
tx.set(
&ikb.new_dy_key(shard),
&DiskAnnPendingState {
kind: DiskAnnPendingStateKind::NonEmpty,
generation: 1,
},
)
.await?;
tx.commit().await?;
}
while compact_once(&index, &ds, &ikb).await? {}
assert_eq!(knn_nearest(&index, &ds, &[30.0, 30.0, 30.0, 30.0]).await?, Some(0.0));
assert_eq!(knn_nearest(&index, &ds, &[20.0, 20.0, 20.0, 20.0]).await?, Some(20.0));
Ok(())
}
#[tokio::test]
async fn diskann_revert_across_layouts_leaves_no_phantom() -> Result<()> {
let ds = Datastore::new("memory").await?;
let ikb = ikb();
let index = DiskAnnIndex::new(
ikb.clone(),
TableId(4),
¶ms(VectorType::F32, Distance::Euclidean),
cache(),
)
.await?;
let id = RecordIdKey::Number(1);
let a = [1.0, 1.0, 1.0, 1.0];
let b = [2.0, 2.0, 2.0, 2.0];
{
let ctx = new_ctx(&ds, TransactionType::Write).await;
index.index(&ctx, &id, None, Some(f32_content(&a))).await?;
ctx.tx().commit().await?;
}
while compact_once(&index, &ds, &ikb).await? {}
{
let ctx = new_ctx(&ds, TransactionType::Write).await;
index.index(&ctx, &id, Some(f32_content(&a)), Some(f32_content(&b))).await?;
ctx.tx().commit().await?;
}
let doc_id = {
let ctx = new_ctx(&ds, TransactionType::Read).await;
let dw: DiskAnnRecordPendingUpdate =
ctx.tx().get(&dw_key(&ikb, &id), None).await?.unwrap();
ctx.tx().cancel().await?;
dw.doc_id
};
assert!(doc_id.is_some());
{
let ctx = new_ctx(&ds, TransactionType::Write).await;
ctx.tx()
.set(
&ikb.new_dr_key(&id),
&DiskAnnRecordPendingUpdate {
doc_id,
old_vectors: vec![SerializedVector::F32(b.to_vec())],
new_vectors: vec![SerializedVector::F32(a.to_vec())],
},
)
.await?;
ctx.tx().commit().await?;
}
while compact_once(&index, &ds, &ikb).await? {}
assert_eq!(knn_nearest(&index, &ds, &a).await?, Some(0.0));
assert_eq!(knn_nearest(&index, &ds, &b).await?, Some(2.0));
Ok(())
}
#[tokio::test]
async fn diskann_write_fold_inverse_pair_keeps_chain_head() -> Result<()> {
let ds = Datastore::new("memory").await?;
let ikb = ikb();
let index = DiskAnnIndex::new(
ikb.clone(),
TableId(4),
¶ms(VectorType::F32, Distance::Euclidean),
cache(),
)
.await?;
let id = RecordIdKey::Number(1);
let shard = DiskAnnIndex::pending_state_shard(&id);
let a = [1.0, 1.0, 1.0, 1.0];
let b = [2.0, 2.0, 2.0, 2.0];
let c = [3.0, 3.0, 3.0, 3.0];
{
let ctx = new_ctx(&ds, TransactionType::Write).await;
index.index(&ctx, &id, None, Some(f32_content(&a))).await?;
ctx.tx().commit().await?;
}
while compact_once(&index, &ds, &ikb).await? {}
{
let ctx = new_ctx(&ds, TransactionType::Write).await;
index.index(&ctx, &id, Some(f32_content(&a)), Some(f32_content(&b))).await?;
ctx.tx().commit().await?;
}
let doc_id = {
let ctx = new_ctx(&ds, TransactionType::Read).await;
let dw: DiskAnnRecordPendingUpdate =
ctx.tx().get(&dw_key(&ikb, &id), None).await?.unwrap();
ctx.tx().cancel().await?;
dw.doc_id
};
{
let ctx = new_ctx(&ds, TransactionType::Write).await;
ctx.tx()
.set(
&ikb.new_dr_key(&id),
&DiskAnnRecordPendingUpdate {
doc_id,
old_vectors: vec![SerializedVector::F32(b.to_vec())],
new_vectors: vec![SerializedVector::F32(a.to_vec())],
},
)
.await?;
ctx.tx().commit().await?;
}
{
let ctx = new_ctx(&ds, TransactionType::Write).await;
index.index(&ctx, &id, Some(f32_content(&a)), Some(f32_content(&c))).await?;
ctx.tx().commit().await?;
}
{
let ctx = new_ctx(&ds, TransactionType::Read).await;
assert!(ctx.tx().get::<_>(&ikb.new_dr_key(&id), None).await?.is_none());
let folded: DiskAnnRecordPendingUpdate =
ctx.tx().get(&ikb.new_dw_key(shard, &id), None).await?.unwrap();
assert_eq!(folded.old_vectors, vec![SerializedVector::F32(a.to_vec())]);
assert_eq!(folded.new_vectors, vec![SerializedVector::F32(c.to_vec())]);
ctx.tx().cancel().await?;
}
while compact_once(&index, &ds, &ikb).await? {}
assert_eq!(knn_nearest(&index, &ds, &c).await?, Some(0.0));
assert_eq!(knn_nearest(&index, &ds, &a).await?, Some(4.0));
Ok(())
}
#[test]
fn diskann_builder_authorized_pair_survives_byte_budget() {
fn op(n: i64) -> PendingOperation {
PendingOperation {
id: VectorId::RecordKey(Arc::new(RecordIdKey::Number(n))),
old_vectors: vec![],
new_vectors: vec![],
}
}
let half = DISKANN_COMPACTION_MAX_PENDING_BYTES / 2 + 1;
let empty_state = vec![None; DISKANN_PENDING_STATE_SHARDS as usize];
let mut builder = PendingPlanBuilder::new(None, empty_state.clone());
assert!(builder.has_room_for(2, 2 + half + half), "empty batch admits the oversized pair");
builder.add_authorized(vec![0u8; 1], vec![0u8; half], op(1));
builder.add_authorized(vec![1u8; 1], vec![0u8; half], op(2));
assert_eq!(builder.captured_keys.len(), 2, "both halves captured");
let mut naive = PendingPlanBuilder::new(None, empty_state);
assert!(naive.add(vec![0u8; 1], vec![0u8; half], op(1)));
assert!(!naive.add(vec![1u8; 1], vec![0u8; half], op(2)));
assert_eq!(naive.captured_keys.len(), 1, "second half rejected by the byte guard");
}
#[tokio::test]
async fn diskann_compaction_split_batch_does_not_leave_phantom_vector() -> Result<()> {
let ds = Datastore::new("memory").await?;
let ikb = ikb();
let index = DiskAnnIndex::new(
ikb.clone(),
TableId(4),
¶ms(VectorType::F32, Distance::Euclidean),
cache(),
)
.await?;
let id = RecordIdKey::Number(1);
{
let ctx = new_ctx(&ds, TransactionType::Write).await;
index.index(&ctx, &id, None, Some(f32_content(&[2.0, 2.0, 2.0, 2.0]))).await?;
ctx.tx().commit().await?;
}
{
let ctx = new_ctx(&ds, TransactionType::Write).await;
let tx = ctx.tx();
tx.set(
&ikb.new_dr_key(&id),
&DiskAnnRecordPendingUpdate {
doc_id: None,
old_vectors: vec![SerializedVector::F32(vec![2.0, 2.0, 2.0, 2.0])],
new_vectors: vec![SerializedVector::F32(vec![3.0, 3.0, 3.0, 3.0])],
},
)
.await?;
for i in 0..DISKANN_COMPACTION_MAX_PENDING_KEYS {
let filler = RecordIdKey::Number(1000 + i as i64);
tx.set(
&ikb.new_dr_key(&filler),
&DiskAnnRecordPendingUpdate {
doc_id: None,
old_vectors: vec![],
new_vectors: vec![],
},
)
.await?;
}
tx.commit().await?;
}
while compact_once(&index, &ds, &ikb).await? {}
assert_eq!(knn_nearest(&index, &ds, &[3.0, 3.0, 3.0, 3.0]).await?, Some(0.0));
assert_eq!(
knn_nearest(&index, &ds, &[2.0, 2.0, 2.0, 2.0]).await?,
Some(2.0),
"phantom vector: the superseded [2,2,2,2] is still indexed for record {id:?}; the \
legacy/sharded pending pair was applied uncoalesced across separate compaction batches",
);
Ok(())
}
#[tokio::test]
async fn diskann_legacy_overflow_reports_has_more() -> Result<()> {
let ds = Datastore::new("memory").await?;
let ikb = ikb();
let index = DiskAnnIndex::new(
ikb.clone(),
TableId(4),
¶ms(VectorType::F32, Distance::Euclidean),
cache(),
)
.await?;
let dual = RecordIdKey::Number(1_000_000);
{
let ctx = new_ctx(&ds, TransactionType::Write).await;
index.index(&ctx, &dual, None, Some(f32_content(&[1.0, 1.0, 1.0, 1.0]))).await?;
ctx.tx().commit().await?;
}
{
let ctx = new_ctx(&ds, TransactionType::Write).await;
let tx = ctx.tx();
for i in 0..(DISKANN_COMPACTION_MAX_PENDING_KEYS - 1) {
let filler = RecordIdKey::Number(i as i64);
tx.set(
&ikb.new_dr_key(&filler),
&DiskAnnRecordPendingUpdate {
doc_id: None,
old_vectors: vec![],
new_vectors: vec![],
},
)
.await?;
}
tx.set(
&ikb.new_dr_key(&dual),
&DiskAnnRecordPendingUpdate {
doc_id: None,
old_vectors: vec![SerializedVector::F32(vec![1.0, 1.0, 1.0, 1.0])],
new_vectors: vec![SerializedVector::F32(vec![2.0, 2.0, 2.0, 2.0])],
},
)
.await?;
tx.commit().await?;
}
let ctx = new_ctx(&ds, TransactionType::Read).await;
let plan = DiskAnnIndex::prepare_compaction(&ctx, &ikb).await?;
ctx.tx().cancel().await?;
assert!(plan.has_work(), "the batch captured the legacy fillers");
assert!(
plan.has_more(),
"legacy `!dr` backlog exceeded one batch but the plan reported no more work; \
process_diskann_compaction would strand the remaining legacy entries until the next \
write re-enqueues the index",
);
Ok(())
}
#[tokio::test]
async fn diskann_failed_compaction_clears_cache_and_keeps_knn_working() -> Result<()> {
let ds = Datastore::new("memory").await?;
let ikb = ikb();
let cache = cache();
let index = DiskAnnIndex::new(
ikb.clone(),
TableId(4),
¶ms(VectorType::F32, Distance::Euclidean),
cache.clone(),
)
.await?;
for i in 0..4_i64 {
let ctx = new_ctx(&ds, TransactionType::Write).await;
let v = [i as f32, 0.0, 0.0, 0.0];
index.index(&ctx, &RecordIdKey::Number(i), None, Some(f32_content(&v))).await?;
ctx.tx().commit().await?;
}
let plan_a = {
let ctx = new_ctx(&ds, TransactionType::Read).await;
let plan = DiskAnnIndex::prepare_compaction(&ctx, &ikb).await?;
ctx.tx().cancel().await?;
plan
};
let plan_b = {
let ctx = new_ctx(&ds, TransactionType::Read).await;
let plan = DiskAnnIndex::prepare_compaction(&ctx, &ikb).await?;
ctx.tx().cancel().await?;
plan
};
assert!(plan_a.has_work());
assert!(plan_b.has_work());
let ctx_a = new_ctx(&ds, TransactionType::Write).await;
let ctx_b = new_ctx(&ds, TransactionType::Write).await;
assert!(index.apply_compaction(&ctx_a, plan_a).await?);
let res = index.apply_compaction(&ctx_b, plan_b).await;
assert!(res.is_err(), "expected commit failure, got {res:?}");
let cache_index = (ikb.ns(), ikb.db(), TableId(4), ikb.index());
assert!(cache.get_state(cache_index).is_none(), "state cache should be empty");
for id in 0..4 {
assert!(cache.get_element(cache_index, id).is_none(), "element {id} cached");
assert!(cache.get_node(cache_index, id).is_none(), "node {id} cached");
}
assert_eq!(knn_len_with_k(&index, &ds, &[2.0, 0.0, 0.0, 0.0], 4).await?, 4);
Ok(())
}
#[tokio::test]
async fn diskann_hashed_vector_compaction_and_knn() -> Result<()> {
let ds = Datastore::new("memory").await?;
let ikb = ikb();
let cache = cache();
let params = DiskAnnParams {
use_hashed_vector: true,
..params(VectorType::F32, Distance::Euclidean)
};
let index = DiskAnnIndex::new(ikb.clone(), TableId(4), ¶ms, cache.clone()).await?;
let v0 = [1.0_f32, 0.0, 0.0, 0.0];
let v1 = [0.0_f32, 1.0, 0.0, 0.0];
let v2 = [0.0_f32, 0.0, 1.0, 0.0];
for (id, v) in [(0, &v0), (1, &v1), (2, &v2)] {
let ctx = new_ctx(&ds, TransactionType::Write).await;
index.index(&ctx, &RecordIdKey::Number(id), None, Some(f32_content(v))).await?;
ctx.tx().commit().await?;
}
assert!(compact_once(&index, &ds, &ikb).await?);
assert_eq!(knn_len(&index, &ds, &v0).await?, 1);
assert_eq!(knn_len(&index, &ds, &v1).await?, 1);
assert_eq!(knn_len(&index, &ds, &v2).await?, 1);
let ctx = new_ctx(&ds, TransactionType::Write).await;
index.index(&ctx, &RecordIdKey::Number(3), None, Some(f32_content(&v0))).await?;
ctx.tx().commit().await?;
assert!(compact_once(&index, &ds, &ikb).await?);
assert_eq!(knn_len(&index, &ds, &v0).await?, 1);
let ctx = new_ctx(&ds, TransactionType::Write).await;
index.index(&ctx, &RecordIdKey::Number(0), Some(f32_content(&v0)), None).await?;
ctx.tx().commit().await?;
assert!(compact_once(&index, &ds, &ikb).await?);
assert_eq!(knn_len(&index, &ds, &v0).await?, 1);
let ctx = new_ctx(&ds, TransactionType::Write).await;
index.index(&ctx, &RecordIdKey::Number(3), Some(f32_content(&v0)), None).await?;
ctx.tx().commit().await?;
assert!(compact_once(&index, &ds, &ikb).await?);
assert_eq!(knn_len(&index, &ds, &v1).await?, 1);
Ok(())
}
}