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::err::Error;
use crate::idx::planner::ScanDirection;
use crate::idx::planner::iterators::KnnIteratorResult;
use crate::idx::trees::KnnCondFilter;
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::pending::{
PendingBacklogReport, PendingScan, PendingScanStats, take_pending_record_id,
};
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, ValsBatch};
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,
pending_backlog_reported: PendingBacklogReport,
pending_scan_stats: PendingScanStats,
}
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)?;
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,
})
}
}
#[derive(Clone, Copy)]
enum PendingLayout {
Sharded,
Legacy,
}
fn pending_record_id(
key: &[u8],
layout: PendingLayout,
pending: &mut DiskAnnRecordPendingUpdate,
) -> Result<RecordIdKey> {
take_pending_record_id(&mut pending.id, || {
Ok(match layout {
PendingLayout::Sharded => DiskAnnRecordPendingShard::decode_key(key)?.id.into_owned(),
PendingLayout::Legacy => DiskAnnRecordPending::decode_key(key)?.id.into_owned(),
})
})
}
struct DiskAnnPendingScan<'a, 'b> {
search: &'a DiskAnnSearch,
filter: &'a mut Option<DiskAnnTruthyDocumentFilter<'b>>,
builder: &'a mut KnnResultBuilder,
pending: PendingScan<'a>,
suppressed: RoaringTreemap,
legacy_shards: u32,
}
const _: () = assert!(DISKANN_PENDING_STATE_SHARDS as u32 <= u32::BITS);
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,
pending_backlog_reported: PendingBacklogReport::default(),
pending_scan_stats: PendingScanStats::default(),
})
}
#[cfg(test)]
pub(crate) fn pending_scan_stats(&self) -> &PendingScanStats {
&self.pending_scan_stats
}
#[cfg(test)]
pub(crate) fn pending_backlog_reported(&self) -> bool {
self.pending_backlog_reported.is_reported()
}
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> = if tx.shared_locked_reads() {
tx.getu(&key).await?
} else {
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 mut 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,
id: None,
}
}
(None, None) => DiskAnnRecordPendingUpdate {
doc_id: DiskAnnDocs::get_doc_id(&self.ikb, &tx, id).await?,
old_vectors,
new_vectors,
id: None,
},
};
pending.id = Some(id.clone());
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 mut legacy_update =
DiskAnnRecordPendingUpdate::kv_decode_value(&legacy_value, ())?;
let mut id =
pending_record_id(&legacy_key, PendingLayout::Legacy, &mut legacy_update)?;
let (shard_key, shard_value) = {
let shard_key = ikb.new_dw_key(Self::pending_state_shard(&id), &id);
let value = tx.get_raw(&shard_key, None).await?;
let encoded = value.is_some().then(|| shard_key.encode_key()).transpose()?;
(encoded, value)
};
let mut shard_update = shard_value
.as_ref()
.map(|v| DiskAnnRecordPendingUpdate::kv_decode_value(v, ()))
.transpose()?;
if let Some(exact) = shard_update.as_mut().and_then(|u| u.id.take()) {
id = exact;
}
let legacy_op = Self::record_pending_to_operation(id.clone(), legacy_update);
let shard_entry = match (shard_key, shard_value, shard_update) {
(Some(shard_key), Some(shard_value), Some(update)) => {
let op = Self::record_pending_to_operation(id.clone(), update);
Some((shard_key, shard_value, op))
}
_ => 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 mut pending = DiskAnnRecordPendingUpdate::kv_decode_value(&value, ())?;
let id = pending_record_id(&key, PendingLayout::Sharded, &mut pending)?;
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<KnnCondFilter<'_>>,
) -> 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(|f| {
DiskAnnTruthyDocumentFilter::new(
f.opt,
self.ikb.clone(),
self.table_id,
self.cache.clone(),
compaction_generation,
f.cond,
f.select_gate,
)
});
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 scan = DiskAnnPendingScan {
search,
filter,
builder,
pending: PendingScan::new(
"diskann",
&self.ikb,
&self.pending_backlog_reported,
&self.pending_scan_stats,
),
suppressed: RoaringTreemap::new(),
legacy_shards: 0,
};
let rng = self.ikb.new_dr_range()?;
self.scan_pending_range(ctx, stk, &mut scan, rng, PendingLayout::Legacy, false).await?;
for shard in 0..DISKANN_PENDING_STATE_SHARDS {
let state = pending_state.get(usize::from(shard)).and_then(|state| state.as_ref());
if state.is_none_or(|state| state.kind == DiskAnnPendingStateKind::Empty) {
continue;
}
let probe_legacy = scan.legacy_shards & (1u32 << u32::from(shard)) != 0;
let rng = self.ikb.new_dw_shard_range(shard)?;
self.scan_pending_range(ctx, stk, &mut scan, rng, PendingLayout::Sharded, probe_legacy)
.await?;
}
self.score_pending_batch(ctx, stk, &mut scan).await?;
scan.pending.finish();
if scan.suppressed.is_empty() {
return Ok(None);
}
Ok(Some(scan.suppressed))
}
async fn scan_pending_range(
&self,
ctx: &DiskAnnContext<'_>,
stk: &mut Stk,
scan: &mut DiskAnnPendingScan<'_, '_>,
rng: Range<Key>,
layout: PendingLayout,
probe_legacy: bool,
) -> Result<()> {
let mut cursor = ctx.tx.open_vals_cursor(rng, ScanDirection::Forward, 0, None).await?;
loop {
let read = cursor.next_batch(scan.pending.rows()).await?;
scan.pending.observe_page(&read);
if read.is_empty() {
break;
}
let legacy_owned = if probe_legacy {
self.score_pending_batch(ctx, stk, scan).await?;
self.legacy_owned_rows(ctx, &read, &scan.pending).await?
} else {
Vec::new()
};
for (row, (key, value)) in read.iter().enumerate() {
let entry_bytes = key.len() + value.len();
scan.pending.charge_entry(entry_bytes);
if ctx.ctx.is_done(Some(scan.pending.entries())).await? {
bail!(Error::QueryCancelled)
}
let mut pending = DiskAnnRecordPendingUpdate::kv_decode_value(value, ())?;
let id = pending_record_id(key, layout, &mut pending)?;
match layout {
PendingLayout::Legacy => {
scan.legacy_shards |= 1u32 << u32::from(Self::pending_state_shard(&id));
}
PendingLayout::Sharded if probe_legacy => {
if legacy_owned[row] {
continue;
}
}
PendingLayout::Sharded => {}
}
let pending = Self::record_pending_to_operation(id, pending);
if let VectorId::DocId(doc_id) = &pending.id {
scan.suppressed.insert(*doc_id);
}
if pending.new_vectors.is_empty() {
continue;
}
if scan.pending.rollover_required(entry_bytes) {
self.score_pending_batch(ctx, stk, scan).await?;
}
scan.pending.push(pending.id, pending.new_vectors, entry_bytes);
}
}
Ok(())
}
async fn legacy_owned_rows(
&self,
ctx: &DiskAnnContext<'_>,
read: &ValsBatch<'_>,
pending: &PendingScan<'_>,
) -> Result<Vec<bool>> {
debug_assert!(
pending.batch.is_empty(),
"the pending batch is scored before an ownership reply is read, so only the page, \
request, and reply are live while the reply is"
);
let mut ids = Vec::with_capacity(read.len());
for (key, _) in read {
ids.push(DiskAnnRecordPendingShard::decode_key(key)?.id.into_owned());
}
let keys: Vec<_> = ids.iter().map(|id| self.ikb.new_dr_key(id)).collect();
let found = ctx.tx.getm_raw(keys, None).await?;
let reply_bytes: usize = found.iter().flatten().map(|value| value.len()).sum();
pending.record_side_read(
read.len(),
pending.batch.bytes()
+ pending.batch.page_bytes()
+ read.key_bytes as usize
+ reply_bytes,
);
Ok(found.into_iter().map(|value| value.is_some()).collect())
}
async fn score_pending_batch(
&self,
ctx: &DiskAnnContext<'_>,
stk: &mut Stk,
scan: &mut DiskAnnPendingScan<'_, '_>,
) -> Result<()> {
if scan.pending.batch.is_empty() {
return Ok(());
}
scan.pending.begin_scoring();
if let Some(filter) = scan.filter.as_mut() {
filter.prefetch_records(ctx, &scan.pending.batch.ids).await?;
}
let batch = &mut scan.pending.batch;
for (id, vectors) in batch.ids.drain(..).zip(batch.vectors.drain(..)) {
let truthy = match scan.filter.as_mut() {
Some(filter) => filter.check_vector_id_truthy(ctx, stk, id.clone()).await?,
None => true,
};
if truthy {
for vector in vectors {
let vector = Vector::from(vector);
let d = self.distance.calculate(&scan.search.pt, &vector);
if scan.builder.check_add(d)
&& let Some(evicted_id) = scan.builder.add_vector_id_result(d, id.clone())
&& let Some(filter) = scan.filter.as_mut()
{
filter.expire(&evicted_id);
}
}
}
if let Some(filter) = scan.filter.as_mut()
&& !scan.builder.contains(&id)
{
filter.expire(&id);
}
}
Ok(())
}
}
#[cfg(test)]
mod tests {
#[cfg(feature = "kv-rocksdb")]
use temp_dir::TempDir;
use super::*;
use crate::catalog::{DatabaseId, IndexId, NamespaceId};
use crate::idx::DocId;
use crate::idx::trees::diskann::cache::DiskAnnCache;
use crate::idx::trees::pending::{
PENDING_MAX_BATCH_KEYS, PENDING_MAX_BYTES, PENDING_MAX_PAGE_BYTES, PENDING_MAX_ROWS,
PENDING_PROBE_ROWS,
};
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 {
wide_params(4, vector_type, distance)
}
fn wide_params(dimension: u16, vector_type: VectorType, distance: Distance) -> DiskAnnParams {
DiskAnnParams {
dimension,
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())],
id: None,
}
}
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(())
}
async fn knn_keys(
index: &DiskAnnIndex,
ds: &Datastore,
values: &[f32],
k: usize,
) -> Result<Vec<i64>> {
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?;
let mut keys: Vec<i64> = res
.iter()
.map(|(rid, _, _)| match &rid.key {
RecordIdKey::Number(n) => *n,
other => panic!("unexpected record key: {other:?}"),
})
.collect();
keys.sort();
Ok(keys)
}
async fn knn_cancelled_at(
index: &DiskAnnIndex,
ds: &Datastore,
values: &[f32],
k: usize,
at: usize,
) -> Result<anyhow::Error> {
let tx = Arc::new(ds.transaction(TransactionType::Read, LockType::Optimistic).await?);
let mut ctx = Context::new_test();
ctx.set_transaction(tx);
let canceller = ctx.add_cancel();
let ctx = ctx.freeze();
index.pending_scan_stats().interrupt_at(at, canceller);
let query = f32_query(values);
let mut stack = reblessive::tree::TreeStack::new();
let err = stack
.enter(|stk| async { index.knn_search(&ctx, stk, &query, k, 8, None).await })
.finish()
.await
.expect_err("the scan is cancelled before it finishes reading the queue");
ctx.tx().cancel().await?;
Ok(err)
}
async fn new_pending_backlog_fixture(
ds: &Datastore,
ikb: &IndexKeyBase,
n: i64,
) -> Result<DiskAnnIndex> {
let index = DiskAnnIndex::new(
ikb.clone(),
TableId(4),
¶ms(VectorType::F32, Distance::Euclidean),
cache(),
)
.await?;
let ctx = new_ctx(ds, TransactionType::Write).await;
for i in 1..=n {
index
.index(
&ctx,
&RecordIdKey::Number(i),
None,
Some(f32_content(&[i as f32, 0.0, 0.0, 0.0])),
)
.await?;
}
ctx.tx().commit().await?;
Ok(index)
}
async fn new_large_key_pending_backlog_fixture(
ds: &Datastore,
ikb: &IndexKeyBase,
n: i64,
key_payload_bytes: usize,
) -> Result<DiskAnnIndex> {
let index = DiskAnnIndex::new(
ikb.clone(),
TableId(4),
¶ms(VectorType::F32, Distance::Euclidean),
cache(),
)
.await?;
let ctx = new_ctx(ds, TransactionType::Write).await;
index
.index(&ctx, &RecordIdKey::Number(0), None, Some(f32_content(&[0.0, 0.0, 0.0, 0.0])))
.await?;
let suffix = "x".repeat(key_payload_bytes);
for i in 1..=n {
let id = RecordIdKey::String(format!("{i:04}-{suffix}").into());
index.index(&ctx, &id, None, Some(f32_content(&[i as f32, 0.0, 0.0, 0.0]))).await?;
}
ctx.tx().commit().await?;
Ok(index)
}
#[tokio::test]
async fn diskann_pending_scan_accounts_for_record_key_bytes() -> Result<()> {
const PENDING: i64 = 600;
const KEY_PAYLOAD_BYTES: usize = 8 * 1024;
const _: () = assert!(PENDING as usize + 1 < PENDING_MAX_BATCH_KEYS);
const _: () = assert!(PENDING as usize * KEY_PAYLOAD_BYTES > PENDING_MAX_BYTES);
let ds = Datastore::new("memory").await?;
let ikb = ikb();
let index =
new_large_key_pending_backlog_fixture(&ds, &ikb, PENDING, KEY_PAYLOAD_BYTES).await?;
assert_eq!(knn_keys(&index, &ds, &[0.0, 0.0, 0.0, 0.0], 1).await?, vec![0]);
let stats = index.pending_scan_stats();
assert!(
stats.batches() > 1,
"record key bytes alone must split the queue into several scoring batches"
);
assert!(
stats.peak_batch_bytes() <= PENDING_MAX_BYTES,
"retained record-key bytes stay inside the scoring-batch budget: {} bytes",
stats.peak_batch_bytes(),
);
Ok(())
}
async fn new_wide_pending_backlog_fixture(
ds: &Datastore,
ikb: &IndexKeyBase,
dimension: u16,
n: i64,
) -> Result<DiskAnnIndex> {
let index = DiskAnnIndex::new(
ikb.clone(),
TableId(4),
&wide_params(dimension, VectorType::F32, Distance::Euclidean),
cache(),
)
.await?;
let ctx = new_ctx(ds, TransactionType::Write).await;
for i in 1..=n {
let mut values = vec![0.0f32; usize::from(dimension)];
values[0] = i as f32;
index
.index(
&ctx,
&RecordIdKey::Number(i * i64::from(DISKANN_PENDING_STATE_SHARDS)),
None,
Some(vec![f32_value(&values)]),
)
.await?;
}
ctx.tx().commit().await?;
Ok(index)
}
async fn new_widening_pending_backlog_fixture(
ds: &Datastore,
ikb: &IndexKeyBase,
dimension: u16,
narrow: i64,
wide: i64,
vectors_per_wide_record: usize,
) -> Result<DiskAnnIndex> {
let index = DiskAnnIndex::new(
ikb.clone(),
TableId(4),
&wide_params(dimension, VectorType::F32, Distance::Euclidean),
cache(),
)
.await?;
let ctx = new_ctx(ds, TransactionType::Write).await;
for i in 1..=(narrow + wide) {
let mut values = vec![0.0f32; usize::from(dimension)];
values[0] = i as f32;
let copies = if i <= narrow {
1
} else {
vectors_per_wide_record
};
index
.index(
&ctx,
&RecordIdKey::Number(i * i64::from(DISKANN_PENDING_STATE_SHARDS)),
None,
Some(vec![f32_value(&values); copies]),
)
.await?;
}
ctx.tx().commit().await?;
Ok(index)
}
#[tokio::test]
async fn diskann_pending_scan_materialises_bounded_batches() -> Result<()> {
const PENDING: i64 = PENDING_MAX_BATCH_KEYS as i64 * 2 + 7;
const EXPECTED_BATCHES: usize = 3;
const _: () = assert!(PENDING_MAX_BYTES > PENDING_MAX_BATCH_KEYS * 1024);
let ds = Datastore::new("memory").await?;
let ikb = ikb();
let index = new_pending_backlog_fixture(&ds, &ikb, PENDING).await?;
assert_eq!(knn_keys(&index, &ds, &[0.0, 0.0, 0.0, 0.0], 4).await?, vec![1, 2, 3, 4]);
let stats = index.pending_scan_stats();
assert_eq!(
stats.batches(),
EXPECTED_BATCHES,
"a {PENDING}-entry queue scores in {EXPECTED_BATCHES} batches"
);
assert_eq!(
stats.peak_batch_entries(),
PENDING_MAX_BATCH_KEYS,
"no batch holds more than the entry budget"
);
assert!(
stats.peak_batch_bytes() < PENDING_MAX_BYTES,
"the byte budget must not bind for this fixture: {} bytes",
stats.peak_batch_bytes()
);
assert!(
stats.peak_page_entries() <= PENDING_MAX_ROWS as usize,
"no page holds more than the row cap: {} entries",
stats.peak_page_entries()
);
assert!(
stats.peak_page_bytes() <= PENDING_MAX_PAGE_BYTES,
"no page holds more than its byte share: {} bytes",
stats.peak_page_bytes()
);
assert!(
stats.peak_resident_bytes() <= PENDING_MAX_BYTES,
"page and batch together stay inside the residency budget: {} bytes",
stats.peak_resident_bytes()
);
assert_eq!(stats.side_reads(), 0, "an empty legacy range costs no ownership resolution");
Ok(())
}
#[tokio::test]
async fn diskann_pending_scan_pages_shrink_for_wide_vectors() -> Result<()> {
const DIMENSION: u16 = 16_384;
const PENDING: i64 = 64;
let ds = Datastore::new("memory").await?;
let ikb = ikb();
let index = new_wide_pending_backlog_fixture(&ds, &ikb, DIMENSION, PENDING).await?;
let query = vec![0.0f32; usize::from(DIMENSION)];
assert_eq!(
knn_keys(&index, &ds, &query, 2).await?,
vec![
i64::from(DISKANN_PENDING_STATE_SHARDS),
2 * i64::from(DISKANN_PENDING_STATE_SHARDS)
]
);
let stats = index.pending_scan_stats();
let entry_bytes = stats.peak_page_bytes().div_ceil(stats.peak_page_entries());
assert!(
entry_bytes * PENDING_MAX_ROWS as usize > PENDING_MAX_PAGE_BYTES,
"the fixture must be wide enough for a full-size page to overrun the budget: \
{entry_bytes} bytes per entry"
);
assert!(
stats.peak_page_entries() < PENDING_MAX_ROWS as usize,
"the byte share, not the row cap, sized the pages: {} entries",
stats.peak_page_entries()
);
assert!(
stats.peak_page_bytes() <= PENDING_MAX_PAGE_BYTES,
"no page holds more than its byte share: {} bytes",
stats.peak_page_bytes()
);
assert!(
stats.peak_batch_entries() < PENDING_MAX_BATCH_KEYS,
"the byte budget, not the entry budget, ended a batch: {} entries",
stats.peak_batch_entries()
);
assert!(
stats.peak_resident_bytes() <= PENDING_MAX_BYTES,
"page and batch together stay inside the residency budget: {} bytes",
stats.peak_resident_bytes()
);
let after_probe = PENDING as usize - PENDING_PROBE_ROWS as usize;
let shard_pages = 1 + after_probe.div_ceil(stats.peak_page_entries()) + 1;
assert_eq!(
stats.pages(),
1 + shard_pages,
"every page after the probe is the size the byte share allows"
);
Ok(())
}
#[tokio::test]
async fn diskann_pending_scan_caps_pages_when_entries_widen() -> Result<()> {
const DIMENSION: u16 = 256;
const NARROW: i64 = PENDING_PROBE_ROWS as i64;
const VECTORS_PER_WIDE_RECORD: usize = 10;
const WIDE: i64 = PENDING_MAX_ROWS as i64 * 3 + 5;
let ds = Datastore::new("memory").await?;
let ikb = ikb();
let index = new_widening_pending_backlog_fixture(
&ds,
&ikb,
DIMENSION,
NARROW,
WIDE,
VECTORS_PER_WIDE_RECORD,
)
.await?;
let shard = i64::from(DISKANN_PENDING_STATE_SHARDS);
let query = vec![0.0f32; usize::from(DIMENSION)];
assert_eq!(
knn_keys(&index, &ds, &query, 4).await?,
vec![shard, 2 * shard, 3 * shard, 4 * shard],
"the conservative page cap changes residency, not results"
);
let stats = index.pending_scan_stats();
assert_eq!(
stats.peak_page_entries(),
PENDING_MAX_ROWS as usize,
"the narrow prefix sized the next page at the row cap"
);
let wide_entry_bytes = stats.peak_page_bytes().div_ceil(PENDING_MAX_ROWS as usize);
assert!(
wide_entry_bytes * crate::kvs::NORMAL_BATCH_SIZE as usize > PENDING_MAX_BYTES,
"the ordinary row cap must recreate the allocation this regression guards: \
{wide_entry_bytes} bytes per entry"
);
assert!(
stats.peak_page_bytes() <= PENDING_MAX_PAGE_BYTES,
"the pending-specific row cap keeps the widening page inside its share: {} bytes",
stats.peak_page_bytes()
);
assert!(
stats.peak_resident_bytes() <= PENDING_MAX_BYTES,
"page and batch stay inside the residency target: {} bytes",
stats.peak_resident_bytes()
);
assert_eq!(
stats.pages(),
1 + 1 + (WIDE as usize).div_ceil(PENDING_MAX_ROWS as usize) + 1,
"the widening tail stays split by the conservative row cap"
);
Ok(())
}
#[tokio::test]
async fn diskann_pending_backlog_report_arms_on_a_scoring_rollover() -> Result<()> {
const DIMENSION: u16 = 12_000;
const PENDING: i64 = 80;
const _: () = assert!(PENDING as usize <= PENDING_MAX_BATCH_KEYS);
let ds = Datastore::new("memory").await?;
let ikb = ikb();
let index = new_wide_pending_backlog_fixture(&ds, &ikb, DIMENSION, PENDING).await?;
assert!(!index.pending_backlog_reported(), "no scan has read the queue yet");
let query = vec![0.0f32; usize::from(DIMENSION)];
assert_eq!(
knn_keys(&index, &ds, &query, 1).await?,
vec![i64::from(DISKANN_PENDING_STATE_SHARDS)]
);
let stats = index.pending_scan_stats();
assert!(
stats.batches() > 1,
"the fixture must cost more than one batch to score: {} batches",
stats.batches()
);
let entry_bytes = stats.peak_page_bytes().div_ceil(stats.peak_page_entries());
assert!(
stats.entries_read() * entry_bytes < PENDING_MAX_BYTES,
"the queue's raw bytes must stay inside the budget: {} entries of {entry_bytes} bytes",
stats.entries_read()
);
assert!(
index.pending_backlog_reported(),
"a queue that costs several materialisation batches arms the report"
);
{
let ctx = new_ctx(&ds, TransactionType::Write).await;
ctx.tx().delr(ikb.new_dw_shard_range(0)?).await?;
ctx.tx().commit().await?;
}
assert!(knn_keys(&index, &ds, &query, 1).await?.is_empty());
assert!(
!index.pending_backlog_reported(),
"a scan that completes inside both budgets re-arms the report"
);
Ok(())
}
#[tokio::test]
async fn diskann_pending_scan_reports_the_entry_it_is_cancelled_on() -> Result<()> {
const CROSSING: usize = PENDING_MAX_BATCH_KEYS + 1;
const PENDING: i64 = PENDING_MAX_BATCH_KEYS as i64 * 2;
let ds = Datastore::new("memory").await?;
let ikb = ikb();
let index = new_pending_backlog_fixture(&ds, &ikb, PENDING).await?;
assert!(!index.pending_backlog_reported(), "no scan has read the queue yet");
let err = knn_cancelled_at(&index, &ds, &[0.0, 0.0, 0.0, 0.0], 4, CROSSING).await?;
assert!(
matches!(err.downcast_ref::<Error>(), Some(Error::QueryCancelled)),
"unexpected error: {err}"
);
let stats = index.pending_scan_stats();
assert_eq!(
stats.entries_read(),
CROSSING,
"the entry the scan is cancelled on is charged before the checkpoint"
);
assert!(
index.pending_backlog_reported(),
"the crossing is reported before the deadline bails out of the scan"
);
Ok(())
}
#[tokio::test]
async fn diskann_pending_scan_reports_byte_rollover_before_cancellation() -> Result<()> {
const DIMENSION: u16 = 256;
const PENDING: i64 = 2;
const VECTORS_PER_RECORD: usize = 1536;
const CANCEL_AT: usize = 2;
const _: () = assert!((PENDING as usize) < PENDING_MAX_BATCH_KEYS);
let ds = Datastore::new("memory").await?;
let ikb = ikb();
let index = new_widening_pending_backlog_fixture(
&ds,
&ikb,
DIMENSION,
0,
PENDING,
VECTORS_PER_RECORD,
)
.await?;
assert!(!index.pending_backlog_reported(), "no scan has read the queue yet");
let query = vec![0.0f32; usize::from(DIMENSION)];
let err = knn_cancelled_at(&index, &ds, &query, 1, CANCEL_AT).await?;
assert!(
matches!(err.downcast_ref::<Error>(), Some(Error::QueryCancelled)),
"unexpected error: {err}"
);
assert_eq!(index.pending_scan_stats().entries_read(), CANCEL_AT);
assert!(
index.pending_backlog_reported(),
"the byte rollover must be reported before its cancellation checkpoint"
);
Ok(())
}
#[tokio::test]
async fn diskann_legacy_ownership_resolves_one_page_at_a_time() -> Result<()> {
const PENDING: i64 = PENDING_MAX_ROWS as i64 + 7;
let shard = i64::from(DISKANN_PENDING_STATE_SHARDS);
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;
for i in 1..=PENDING {
index
.index(
&ctx,
&RecordIdKey::Number(i * shard),
None,
Some(f32_content(&[i as f32, 0.0, 0.0, 0.0])),
)
.await?;
}
ctx.tx().commit().await?;
}
{
let ctx = new_ctx(&ds, TransactionType::Write).await;
ctx.tx()
.set(
&ikb.new_dr_key(&RecordIdKey::Number(shard)),
&f32_pending(&[9999.0, 0.0, 0.0, 0.0]),
)
.await?;
ctx.tx().commit().await?;
}
assert_eq!(
knn_keys(&index, &ds, &[0.0, 0.0, 0.0, 0.0], 2).await?,
vec![2 * shard, 3 * shard],
"the legacy entry owns its record, so the sharded entry for it is not scored"
);
let stats = index.pending_scan_stats();
assert_eq!(
stats.side_read_keys(),
PENDING as usize,
"every record of the masked shard is resolved"
);
let expected_batches = (PENDING as usize).div_ceil(PENDING_MAX_ROWS as usize);
assert_eq!(
stats.side_reads(),
expected_batches,
"resolving them costs one round trip per page, not one per record"
);
let first_request_bytes: usize = (1..=PENDING_MAX_ROWS as i64)
.map(|i| ikb.new_dr_key(&RecordIdKey::Number(i * shard)).encode_key().unwrap().len())
.sum();
assert!(
stats.peak_resident_bytes() >= stats.peak_page_bytes() + first_request_bytes,
"ownership residency includes the encoded lookup request: page={} request={} peak={}",
stats.peak_page_bytes(),
first_request_bytes,
stats.peak_resident_bytes(),
);
assert!(
stats.peak_resident_bytes() <= PENDING_MAX_BYTES,
"the ownership request and reply land inside the residency budget: {} bytes",
stats.peak_resident_bytes()
);
Ok(())
}
#[tokio::test]
async fn diskann_pending_scan_batching_preserves_top_k() -> Result<()> {
const PENDING: i64 = PENDING_MAX_BATCH_KEYS as i64 * 2 + 7;
let ds = Datastore::new("memory").await?;
let ikb = ikb();
let index = new_pending_backlog_fixture(&ds, &ikb, PENDING).await?;
assert_eq!(
knn_keys(&index, &ds, &[2047.5, 0.0, 0.0, 0.0], 4).await?,
vec![2046, 2047, 2048, 2049]
);
assert_eq!(index.pending_scan_stats().batches(), 3);
Ok(())
}
#[tokio::test]
async fn diskann_pending_scan_below_budget_scores_in_one_batch() -> Result<()> {
const SHALLOW: i64 = PENDING_MAX_BATCH_KEYS as i64 - 1;
const DEEP: i64 = PENDING_MAX_BATCH_KEYS as i64 * 2 + 7;
const QUERY: [f32; 4] = [511.5, 0.0, 0.0, 0.0];
const EXPECTED: [i64; 4] = [510, 511, 512, 513];
let ds = Datastore::new("memory").await?;
let shallow_ikb = IndexKeyBase::new(NamespaceId(1), DatabaseId(2), "tb".into(), IndexId(3));
let deep_ikb = IndexKeyBase::new(NamespaceId(1), DatabaseId(2), "tb".into(), IndexId(4));
let shallow = new_pending_backlog_fixture(&ds, &shallow_ikb, SHALLOW).await?;
let deep = new_pending_backlog_fixture(&ds, &deep_ikb, DEEP).await?;
assert_eq!(knn_keys(&shallow, &ds, &QUERY, 4).await?, EXPECTED.to_vec());
assert_eq!(
shallow.pending_scan_stats().batches(),
1,
"a queue one entry short of the budget scores in a single batch"
);
assert_eq!(
knn_keys(&deep, &ds, &QUERY, 4).await?,
EXPECTED.to_vec(),
"batching the same records across three batches returns the same neighbours"
);
assert_eq!(deep.pending_scan_stats().batches(), 3);
Ok(())
}
#[tokio::test]
async fn diskann_pending_scan_masks_stale_graph_entries_across_batches() -> Result<()> {
const FILLERS: i64 = PENDING_MAX_BATCH_KEYS as i64 * 2;
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 near = RecordIdKey::Number(9023);
let winner = RecordIdKey::Number(9024);
{
let ctx = new_ctx(&ds, TransactionType::Write).await;
index.index(&ctx, &near, None, Some(f32_content(&[1.0, 0.0, 0.0, 0.0]))).await?;
index.index(&ctx, &winner, None, Some(f32_content(&[5.0, 0.0, 0.0, 0.0]))).await?;
ctx.tx().commit().await?;
}
for _ in 0..8 {
if !compact_once(&index, &ds, &ikb).await? {
break;
}
}
{
let ctx = new_ctx(&ds, TransactionType::Read).await;
assert!(ctx.tx().get(&dw_key(&ikb, &near), None).await?.is_none());
assert!(ctx.tx().get(&dw_key(&ikb, &winner), None).await?.is_none());
ctx.tx().cancel().await?;
}
{
let ctx = new_ctx(&ds, TransactionType::Write).await;
for i in 1..=FILLERS {
index
.index(
&ctx,
&RecordIdKey::Number(i),
None,
Some(f32_content(&[10000.0 + i as f32, 0.0, 0.0, 0.0])),
)
.await?;
}
ctx.tx().commit().await?;
}
{
let ctx = new_ctx(&ds, TransactionType::Write).await;
index
.index(
&ctx,
&near,
Some(f32_content(&[1.0, 0.0, 0.0, 0.0])),
Some(f32_content(&[30000.0, 0.0, 0.0, 0.0])),
)
.await?;
ctx.tx().commit().await?;
}
assert_eq!(
knn_keys(&index, &ds, &[0.0, 0.0, 0.0, 0.0], 1).await?,
vec![9024],
"the moved record's stale graph entry must stay masked"
);
assert_eq!(
index.pending_scan_stats().batches(),
3,
"the moved record is scored alone in the third batch"
);
Ok(())
}
#[tokio::test]
async fn diskann_legacy_pending_supersedes_sharded_entry() -> 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 = RecordIdKey::Number(1);
let second = RecordIdKey::Number(2);
{
let ctx = new_ctx(&ds, TransactionType::Write).await;
let tx = ctx.tx();
tx.set(&dw_key(&ikb, &first), &f32_pending(&[1.0, 0.0, 0.0, 0.0])).await?;
tx.set(
&ikb.new_dy_key(DiskAnnIndex::pending_state_shard(&first)),
&DiskAnnPendingState {
kind: DiskAnnPendingStateKind::NonEmpty,
generation: 1,
},
)
.await?;
tx.set(&ikb.new_dr_key(&first), &f32_pending(&[9.0, 0.0, 0.0, 0.0])).await?;
tx.commit().await?;
}
{
let ctx = new_ctx(&ds, TransactionType::Write).await;
index.index(&ctx, &second, None, Some(f32_content(&[5.0, 0.0, 0.0, 0.0]))).await?;
ctx.tx().commit().await?;
}
assert_eq!(knn_keys(&index, &ds, &[0.0, 0.0, 0.0, 0.0], 1).await?, vec![2]);
assert_eq!(knn_keys(&index, &ds, &[0.0, 0.0, 0.0, 0.0], 3).await?, vec![1, 2]);
assert_eq!(
index.pending_scan_stats().side_reads(),
2,
"one round trip per search, for the single shard a legacy record hashes to"
);
Ok(())
}
#[tokio::test]
async fn diskann_legacy_pending_delete_cancels_sharded_entry() -> 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 = RecordIdKey::Number(1);
let second = RecordIdKey::Number(2);
{
let ctx = new_ctx(&ds, TransactionType::Write).await;
let tx = ctx.tx();
tx.set(&dw_key(&ikb, &first), &f32_pending(&[1.0, 0.0, 0.0, 0.0])).await?;
tx.set(
&ikb.new_dy_key(DiskAnnIndex::pending_state_shard(&first)),
&DiskAnnPendingState {
kind: DiskAnnPendingStateKind::NonEmpty,
generation: 1,
},
)
.await?;
tx.set(
&ikb.new_dr_key(&first),
&DiskAnnRecordPendingUpdate {
doc_id: Some(7),
old_vectors: vec![SerializedVector::F32(vec![1.0, 0.0, 0.0, 0.0])],
new_vectors: vec![],
id: None,
},
)
.await?;
tx.commit().await?;
}
{
let ctx = new_ctx(&ds, TransactionType::Write).await;
index.index(&ctx, &second, None, Some(f32_content(&[5.0, 0.0, 0.0, 0.0]))).await?;
ctx.tx().commit().await?;
}
assert_eq!(
knn_keys(&index, &ds, &[0.0, 0.0, 0.0, 0.0], 3).await?,
vec![2],
"the deleted record's sharded entry must not resurface"
);
Ok(())
}
#[tokio::test]
async fn diskann_pending_scan_probes_only_shards_a_legacy_record_touches() -> 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;
for i in 1..=3 {
index
.index(
&ctx,
&RecordIdKey::Number(i),
None,
Some(f32_content(&[i as f32, 0.0, 0.0, 0.0])),
)
.await?;
}
ctx.tx().commit().await?;
}
assert_eq!(knn_keys(&index, &ds, &[0.0, 0.0, 0.0, 0.0], 4).await?, vec![1, 2, 3]);
assert_eq!(
index.pending_scan_stats().side_read_keys(),
0,
"an empty legacy range costs no ownership resolution"
);
{
let ctx = new_ctx(&ds, TransactionType::Write).await;
ctx.tx()
.set(&ikb.new_dr_key(&RecordIdKey::Number(3)), &f32_pending(&[7.0, 0.0, 0.0, 0.0]))
.await?;
ctx.tx().commit().await?;
}
assert_eq!(knn_keys(&index, &ds, &[0.0, 0.0, 0.0, 0.0], 4).await?, vec![1, 2, 3]);
assert_eq!(
index.pending_scan_stats().side_read_keys(),
1,
"only the one shard a legacy record hashes to is resolved"
);
assert_eq!(
index.pending_scan_stats().side_reads(),
1,
"and that shard's single page costs one round trip"
);
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(())
}
#[cfg(feature = "kv-rocksdb")]
#[tokio::test]
async fn diskann_pending_written_while_its_shard_empties_stays_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 shard = DiskAnnIndex::pending_state_shard(&RecordIdKey::Number(1));
assert_eq!(shard, DiskAnnIndex::pending_state_shard(&RecordIdKey::Number(33)));
{
let ctx = new_ctx(&ds, TransactionType::Write).await;
index
.index(
&ctx,
&RecordIdKey::Number(1),
None,
Some(f32_content(&[1.0, 0.0, 0.0, 0.0])),
)
.await?;
ctx.tx().commit().await?;
}
let writer = new_ctx(&ds, TransactionType::Write).await;
assert!(
writer.tx().shared_locked_reads(),
"the store must share locked reads for the writer to take the locked path"
);
index
.index(
&writer,
&RecordIdKey::Number(33),
None,
Some(f32_content(&[33.0, 0.0, 0.0, 0.0])),
)
.await?;
for _ in 0..3 {
compact_once(&index, &ds, &ikb).await?;
}
{
let ctx = new_ctx(&ds, TransactionType::Read).await;
let state = DiskAnnIndex::read_pending_state(&ctx.tx(), &ikb).await?;
ctx.tx().cancel().await?;
assert_eq!(
state[shard as usize].as_ref().map(|state| state.kind),
Some(DiskAnnPendingStateKind::Empty),
"the shard must reach Empty while the writer is open for this to test anything"
);
}
match writer.tx().commit().await {
Ok(()) => {}
Err(e) if crate::kvs::is_retryable_transaction_conflict(&e) => {
let ctx = new_ctx(&ds, TransactionType::Write).await;
index
.index(
&ctx,
&RecordIdKey::Number(33),
None,
Some(f32_content(&[33.0, 0.0, 0.0, 0.0])),
)
.await?;
ctx.tx().commit().await?;
}
Err(e) => return Err(e),
}
let query = [33.0, 0.0, 0.0, 0.0];
assert_eq!(
knn_keys(&index, &ds, &query, 2).await?,
vec![1, 33],
"search must see the committed pending"
);
while compact_once(&index, &ds, &ikb).await? {}
assert_eq!(
knn_keys(&index, &ds, &query, 2).await?,
vec![1, 33],
"compaction must fold the committed pending into the graph"
);
Ok(())
}
#[cfg(feature = "kv-rocksdb")]
#[tokio::test]
async fn diskann_writer_committing_inside_the_emptying_pass_stays_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 shard = DiskAnnIndex::pending_state_shard(&RecordIdKey::Number(1));
assert_eq!(shard, DiskAnnIndex::pending_state_shard(&RecordIdKey::Number(33)));
{
let ctx = new_ctx(&ds, TransactionType::Write).await;
index
.index(
&ctx,
&RecordIdKey::Number(1),
None,
Some(f32_content(&[1.0, 0.0, 0.0, 0.0])),
)
.await?;
ctx.tx().commit().await?;
}
let writer = new_ctx(&ds, TransactionType::Write).await;
assert!(
writer.tx().shared_locked_reads(),
"the store must share locked reads for the writer to take the locked path"
);
index
.index(
&writer,
&RecordIdKey::Number(33),
None,
Some(f32_content(&[33.0, 0.0, 0.0, 0.0])),
)
.await?;
assert!(compact_once(&index, &ds, &ikb).await?);
let plan = {
let ctx = new_ctx(&ds, TransactionType::Read).await;
let plan = DiskAnnIndex::prepare_compaction(&ctx, &ikb).await?;
ctx.tx().cancel().await?;
plan
};
let emptying = new_ctx(&ds, TransactionType::Write).await;
match writer.tx().commit().await {
Ok(()) => {}
Err(e) if crate::kvs::is_retryable_transaction_conflict(&e) => {
let ctx = new_ctx(&ds, TransactionType::Write).await;
index
.index(
&ctx,
&RecordIdKey::Number(33),
None,
Some(f32_content(&[33.0, 0.0, 0.0, 0.0])),
)
.await?;
ctx.tx().commit().await?;
}
Err(e) => return Err(e),
}
match index.apply_compaction(&emptying, plan).await {
Ok(_) => {}
Err(e) if crate::kvs::is_retryable_transaction_conflict(&e) => {}
Err(e) => return Err(e),
}
let query = [33.0, 0.0, 0.0, 0.0];
assert_eq!(
knn_keys(&index, &ds, &query, 2).await?,
vec![1, 33],
"search must see the committed pending"
);
while compact_once(&index, &ds, &ikb).await? {}
assert_eq!(
knn_keys(&index, &ds, &query, 2).await?,
vec![1, 33],
"compaction must fold the committed pending into the graph"
);
Ok(())
}
#[cfg(feature = "kv-rocksdb")]
#[tokio::test]
async fn diskann_overlapping_writers_into_a_non_empty_shard_all_commit() -> 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 shard = DiskAnnIndex::pending_state_shard(&RecordIdKey::Number(1));
for id in [33, 65] {
assert_eq!(shard, DiskAnnIndex::pending_state_shard(&RecordIdKey::Number(id)));
}
{
let ctx = new_ctx(&ds, TransactionType::Write).await;
index
.index(
&ctx,
&RecordIdKey::Number(1),
None,
Some(f32_content(&[1.0, 0.0, 0.0, 0.0])),
)
.await?;
ctx.tx().commit().await?;
}
let guard = {
let ctx = new_ctx(&ds, TransactionType::Read).await;
let state = DiskAnnIndex::read_pending_state(&ctx.tx(), &ikb).await?;
ctx.tx().cancel().await?;
state[shard as usize].clone().expect("record 1 sets its shard's guard")
};
assert_eq!(guard.kind, DiskAnnPendingStateKind::NonEmpty);
let first = new_ctx(&ds, TransactionType::Write).await;
let second = new_ctx(&ds, TransactionType::Write).await;
assert!(
first.tx().shared_locked_reads(),
"the store must share locked reads for the writers to take the locked path"
);
index
.index(
&first,
&RecordIdKey::Number(33),
None,
Some(f32_content(&[33.0, 0.0, 0.0, 0.0])),
)
.await?;
index
.index(
&second,
&RecordIdKey::Number(65),
None,
Some(f32_content(&[65.0, 0.0, 0.0, 0.0])),
)
.await?;
first.tx().commit().await?;
second.tx().commit().await?;
{
let ctx = new_ctx(&ds, TransactionType::Read).await;
let state = DiskAnnIndex::read_pending_state(&ctx.tx(), &ikb).await?;
ctx.tx().cancel().await?;
assert_eq!(state[shard as usize], Some(guard), "neither writer rewrites the guard");
}
assert_eq!(
knn_keys(&index, &ds, &[33.0, 0.0, 0.0, 0.0], 3).await?,
vec![1, 33, 65],
"search must see both writers' pendings"
);
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_stamps_the_id_on_an_entry_that_predates_it() -> Result<()> {
use crate::val::{Number, Value};
let arr = |n: i64| RecordIdKey::Array(vec![Value::Number(Number::Int(n))].into());
for (layout, sharded) in [("sharded", true), ("legacy", false)] {
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 = arr(1);
{
let ctx = new_ctx(&ds, TransactionType::Write).await;
let tx = ctx.tx();
let mut pending = f32_pending(&[1.0, 2.0, 3.0, 4.0]);
pending.id = None;
if sharded {
let shard = DiskAnnIndex::pending_state_shard(&id);
tx.set(&ikb.new_dw_key(shard, &id), &pending).await?;
} else {
tx.set(&ikb.new_dr_key(&id), &pending).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;
let shard = DiskAnnIndex::pending_state_shard(&id);
let stored: DiskAnnRecordPendingUpdate =
ctx.tx().get(&ikb.new_dw_key(shard, &id), None).await?.expect("an entry");
ctx.tx().cancel().await?;
assert_eq!(
stored.id.as_ref().map(storekey::encode_vec).transpose().unwrap(),
Some(storekey::encode_vec(&id).unwrap()),
"the {layout} arm must stamp the record id on the entry it reuses"
);
}
Ok(())
}
#[tokio::test]
async fn diskann_dual_layout_folds_under_the_id_the_value_carries() -> Result<()> {
use crate::val::{Number, Value};
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::Array(vec![Value::Number(Number::Int(1))].into());
{
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;
let tx = ctx.tx();
tx.set(
&ikb.new_dr_key(&id),
&DiskAnnRecordPendingUpdate {
doc_id: None,
old_vectors: vec![SerializedVector::F32(vec![1.0, 2.0, 3.0, 4.0])],
new_vectors: vec![SerializedVector::F32(vec![5.0, 6.0, 7.0, 8.0])],
id: None,
},
)
.await?;
tx.commit().await?;
}
while compact_once(&index, &ds, &ikb).await? {}
let ctx = new_ctx(&ds, TransactionType::Read).await;
let tx = ctx.tx();
let doc_id: DocId =
tx.get(&ikb.new_di_key(&id), None).await?.expect("mapped after compaction");
let mapped: RecordIdKey =
tx.get(&ikb.new_dd_key(doc_id), None).await?.expect("reverse mapping");
tx.cancel().await?;
assert_eq!(
storekey::encode_vec(&mapped).unwrap(),
storekey::encode_vec(&id).unwrap(),
"folded under {mapped:?}, which is not the key the record is stored under"
);
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])],
id: None,
},
)
.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])],
id: None,
},
)
.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])],
id: None,
},
)
.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])],
id: None,
},
)
.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())],
id: None,
},
)
.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())],
id: None,
},
)
.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])],
id: None,
},
)
.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![],
id: None,
},
)
.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![],
id: None,
},
)
.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])],
id: None,
},
)
.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(())
}
}