use std::collections::HashMap;
use std::marker::PhantomData;
use std::sync::Arc;
use anyhow::Result;
use diskann::graph::glue::{
CopyIds, DefaultPostProcessor, HybridPredicate, InsertStrategy, PruneAccessor, PruneStrategy,
SearchAccessor, SearchStrategy,
};
use diskann::graph::{AdjacencyList, workingset};
use diskann::provider::{
DataProvider, Delete, ElementStatus, ExecutionContext, HasId, NeighborAccessor,
NeighborAccessorMut, NoopGuard, SetElement,
};
use diskann::utils::VectorRepr;
use diskann::{ANNError, ANNResult, default_post_processor};
use diskann_vector::distance::Metric;
use diskann_vector::{Half, PreprocessedDistanceFunction};
use tracing::warn;
use crate::catalog::{DatabaseId, IndexId, NamespaceId, TableId};
use crate::idx::IndexKeyBase;
use crate::idx::trees::diskann::cache::DiskAnnCache;
use crate::idx::trees::diskann::{DiskAnnElement, DiskAnnNode, DiskAnnState, ElementId};
use crate::idx::trees::vector::SerializedVector;
#[cfg(test)]
use crate::idx::trees::vector::Vector;
use crate::kvs::{KVValue, Transaction};
#[derive(Clone)]
pub(super) struct DiskAnnProviderContext {
pub(super) tx: Arc<Transaction>,
pub(super) ikb: IndexKeyBase,
}
impl ExecutionContext for DiskAnnProviderContext {}
pub(super) trait DiskAnnVectorElement:
VectorRepr + Copy + Default + Send + Sync + 'static
{
fn serialized_from_slice(slice: &[Self]) -> SerializedVector;
fn copy_from_serialized(vector: &SerializedVector, buffer: &mut [Self]) -> ANNResult<()>;
}
fn vector_type_mismatch(expected: &str, vector: &SerializedVector) -> ANNError {
ANNError::log_index_error(format!("DiskANN expected {expected} vector, got {vector:?}"))
}
fn vector_len_mismatch(current: usize, expected: usize) -> ANNError {
ANNError::log_index_error(format!(
"DiskANN vector dimension mismatch: got {current}, expected {expected}"
))
}
impl DiskAnnVectorElement for f32 {
fn serialized_from_slice(slice: &[Self]) -> SerializedVector {
SerializedVector::F32(slice.to_vec())
}
fn copy_from_serialized(vector: &SerializedVector, buffer: &mut [Self]) -> ANNResult<()> {
let SerializedVector::F32(values) = vector else {
return Err(vector_type_mismatch("F32", vector));
};
if values.len() != buffer.len() {
return Err(vector_len_mismatch(values.len(), buffer.len()));
}
buffer.copy_from_slice(values);
Ok(())
}
}
impl DiskAnnVectorElement for Half {
fn serialized_from_slice(slice: &[Self]) -> SerializedVector {
SerializedVector::F16(slice.iter().map(|v| v.to_bits()).collect())
}
fn copy_from_serialized(vector: &SerializedVector, buffer: &mut [Self]) -> ANNResult<()> {
let SerializedVector::F16(values) = vector else {
return Err(vector_type_mismatch("F16", vector));
};
if values.len() != buffer.len() {
return Err(vector_len_mismatch(values.len(), buffer.len()));
}
for (dst, bits) in buffer.iter_mut().zip(values.iter()) {
*dst = Half::from_bits(*bits);
}
Ok(())
}
}
impl DiskAnnVectorElement for i8 {
fn serialized_from_slice(slice: &[Self]) -> SerializedVector {
SerializedVector::I8(slice.to_vec())
}
fn copy_from_serialized(vector: &SerializedVector, buffer: &mut [Self]) -> ANNResult<()> {
let SerializedVector::I8(values) = vector else {
return Err(vector_type_mismatch("I8", vector));
};
if values.len() != buffer.len() {
return Err(vector_len_mismatch(values.len(), buffer.len()));
}
buffer.copy_from_slice(values);
Ok(())
}
}
impl DiskAnnVectorElement for u8 {
fn serialized_from_slice(slice: &[Self]) -> SerializedVector {
SerializedVector::U8(slice.to_vec())
}
fn copy_from_serialized(vector: &SerializedVector, buffer: &mut [Self]) -> ANNResult<()> {
let SerializedVector::U8(values) = vector else {
return Err(vector_type_mismatch("U8", vector));
};
if values.len() != buffer.len() {
return Err(vector_len_mismatch(values.len(), buffer.len()));
}
buffer.copy_from_slice(values);
Ok(())
}
}
pub(super) struct DiskAnnProvider {
pub(super) ikb: IndexKeyBase,
table_id: TableId,
cache: DiskAnnCache,
dim: usize,
metric: Metric,
}
impl DiskAnnProvider {
pub(super) fn new(
ikb: IndexKeyBase,
table_id: TableId,
cache: DiskAnnCache,
dim: usize,
metric: Metric,
) -> Self {
Self {
ikb,
table_id,
cache,
dim,
metric,
}
}
fn cache_index(&self) -> (NamespaceId, DatabaseId, TableId, IndexId) {
(self.ikb.ns(), self.ikb.db(), self.table_id, self.ikb.index())
}
pub(super) fn context(&self, tx: Arc<Transaction>) -> DiskAnnProviderContext {
DiskAnnProviderContext {
tx,
ikb: self.ikb.clone(),
}
}
async fn read_state(&self, context: &DiskAnnProviderContext) -> Result<DiskAnnState> {
let index = self.cache_index();
if let Some(state) = self.cache.get_state(index) {
return Ok(state);
}
let state = context.tx.get(&context.ikb.new_ds_key(), None).await?.unwrap_or_default();
if !context.tx.writeable() {
self.cache.insert_state(index, state.clone());
}
Ok(state)
}
async fn write_state(
&self,
context: &DiskAnnProviderContext,
state: &DiskAnnState,
) -> Result<()> {
context.tx.set(&context.ikb.new_ds_key(), state).await?;
self.cache.insert_state(self.cache_index(), state.clone());
Ok(())
}
async fn read_elements(
&self,
context: &DiskAnnProviderContext,
element_ids: &[ElementId],
) -> Result<Vec<Option<Arc<DiskAnnElement>>>> {
let index = self.cache_index();
let mut elements = vec![None; element_ids.len()];
let mut miss_positions = Vec::new();
let mut miss_keys = Vec::new();
for (pos, element_id) in element_ids.iter().copied().enumerate() {
if let Some(element) = self.cache.get_element(index, element_id) {
elements[pos] = Some(element);
} else {
miss_positions.push((pos, element_id));
miss_keys.push(context.ikb.new_de_key(element_id));
}
}
if !miss_keys.is_empty() {
let fetched: Vec<Option<DiskAnnElement>> = context.tx.getm(miss_keys, None).await?;
let cache_misses = !context.tx.writeable();
for ((pos, element_id), element) in miss_positions.into_iter().zip(fetched) {
if let Some(element) = element {
elements[pos] = Some(if cache_misses {
self.cache.insert_element(index, element_id, element)
} else {
Arc::new(element)
});
}
}
}
Ok(elements)
}
async fn read_nodes(
&self,
context: &DiskAnnProviderContext,
element_ids: &[ElementId],
) -> Result<Vec<Option<Arc<DiskAnnNode>>>> {
let index = self.cache_index();
let mut nodes = vec![None; element_ids.len()];
let mut miss_positions = Vec::new();
let mut miss_keys = Vec::new();
for (pos, element_id) in element_ids.iter().copied().enumerate() {
if let Some(node) = self.cache.get_node(index, element_id) {
nodes[pos] = Some(node);
} else {
miss_positions.push((pos, element_id));
miss_keys.push(context.ikb.new_dn_key(element_id));
}
}
if !miss_keys.is_empty() {
let fetched: Vec<Option<DiskAnnNode>> = context.tx.getm(miss_keys, None).await?;
let cache_misses = !context.tx.writeable();
for ((pos, element_id), node) in miss_positions.into_iter().zip(fetched) {
if let Some(node) = node {
nodes[pos] = Some(if cache_misses {
self.cache.insert_node(index, element_id, node)
} else {
Arc::new(node)
});
}
}
}
Ok(nodes)
}
pub(super) async fn allocate_element_id(
&self,
context: &DiskAnnProviderContext,
) -> Result<ElementId> {
let mut state = self.read_state(context).await?;
let element_id = state.next_element_id;
state.next_element_id = state.next_element_id.saturating_add(1);
self.write_state(context, &state).await?;
Ok(element_id)
}
pub(super) async fn ensure_entry_point(
&self,
context: &DiskAnnProviderContext,
element_id: ElementId,
) -> Result<()> {
let mut state = self.read_state(context).await?;
if state.enter_point.is_none() {
state.enter_point = Some(element_id);
self.write_state(context, &state).await?;
}
Ok(())
}
pub(super) async fn set_entry_point(
&self,
context: &DiskAnnProviderContext,
element_id: Option<ElementId>,
) -> Result<()> {
let mut state = self.read_state(context).await?;
state.enter_point = element_id;
self.write_state(context, &state).await
}
pub(super) async fn valid_starting_points(
&self,
context: &DiskAnnProviderContext,
) -> ANNResult<Vec<ElementId>> {
let state =
self.read_state(context).await.map_err(|e| ANNError::log_index_error(e.to_string()))?;
if let Some(element_id) = state.enter_point
&& self
.status_by_internal_id(context, element_id)
.await
.is_ok_and(ElementStatus::is_valid)
{
return Ok(vec![element_id]);
}
let rng =
context.ikb.new_de_range().map_err(|e| ANNError::log_index_error(e.to_string()))?;
let mut cursor = context
.tx
.open_vals_cursor(rng, crate::idx::planner::ScanDirection::Forward, 0, None)
.await
.map_err(|e| ANNError::log_index_error(e.to_string()))?;
let cache_misses = !context.tx.writeable();
loop {
let batch = cursor
.next_batch(crate::kvs::NORMAL_BATCH_SIZE)
.await
.map_err(|e| ANNError::log_index_error(e.to_string()))?;
if batch.is_empty() {
break;
}
for (key, value) in &batch {
let key: crate::key::index::de::De<'_> = storekey::decode_borrow(key)
.map_err(|e| ANNError::log_index_error(e.to_string()))?;
let element = DiskAnnElement::kv_decode_value(value, ())
.map_err(|e| ANNError::log_index_error(e.to_string()))?;
let deleted = element.deleted;
if !deleted && cache_misses {
self.cache.insert_element(self.cache_index(), key.element_id, element);
}
if !deleted {
if context.tx.writeable()
&& let Err(e) = self.set_entry_point(context, Some(key.element_id)).await
{
warn!(
error = %e,
element_id = key.element_id,
"Failed to persist DiskANN entry-point hint; continuing with in-memory value",
);
}
return Ok(vec![key.element_id]);
}
}
}
Ok(Vec::new())
}
#[cfg(test)]
async fn get_vectors(
&self,
context: &DiskAnnProviderContext,
element_ids: &[ElementId],
) -> Result<Vec<Option<Vector>>> {
let elements = self.read_elements(context, element_ids).await?;
Ok(elements
.into_iter()
.map(|element| {
element.and_then(|element| {
if element.deleted {
None
} else {
Some(Vector::from(element.vector.clone()))
}
})
})
.collect())
}
}
impl DataProvider for DiskAnnProvider {
type Context = DiskAnnProviderContext;
type InternalId = ElementId;
type ExternalId = ElementId;
type Error = ANNError;
type Guard = NoopGuard<ElementId>;
fn to_internal_id(
&self,
_: &Self::Context,
gid: &Self::ExternalId,
) -> Result<Self::InternalId, Self::Error> {
Ok(*gid)
}
fn to_external_id(
&self,
_: &Self::Context,
id: Self::InternalId,
) -> Result<Self::ExternalId, Self::Error> {
Ok(id)
}
}
impl<T> SetElement<&[T]> for DiskAnnProvider
where
T: DiskAnnVectorElement,
{
type SetError = ANNError;
async fn set_element(
&self,
context: &Self::Context,
id: &Self::ExternalId,
element: &[T],
) -> Result<Self::Guard, Self::SetError> {
if element.len() != self.dim {
return Err(vector_len_mismatch(element.len(), self.dim));
}
let key = context.ikb.new_de_key(*id);
let element = DiskAnnElement {
vector: T::serialized_from_slice(element),
deleted: false,
};
context
.tx
.set(&key, &element)
.await
.map_err(|e| ANNError::log_index_error(e.to_string()))?;
self.cache.insert_element(self.cache_index(), *id, element);
Ok(NoopGuard::new(*id))
}
}
impl Delete for DiskAnnProvider {
async fn delete(
&self,
context: &Self::Context,
gid: &Self::ExternalId,
) -> Result<(), Self::Error> {
let key = context.ikb.new_de_key(*gid);
let Some(element) = self
.read_elements(context, std::slice::from_ref(gid))
.await
.map_err(|e| ANNError::log_index_error(e.to_string()))?
.pop()
.flatten()
else {
return Ok(());
};
let mut element = (*element).clone();
element.deleted = true;
context
.tx
.set(&key, &element)
.await
.map_err(|e| ANNError::log_index_error(e.to_string()))?;
self.cache.insert_element(self.cache_index(), *gid, element);
Ok(())
}
async fn release(
&self,
context: &Self::Context,
id: Self::InternalId,
) -> Result<(), Self::Error> {
context
.tx
.del(&context.ikb.new_de_key(id))
.await
.map_err(|e| ANNError::log_index_error(e.to_string()))?;
self.cache.remove_element(self.cache_index(), id);
Ok(())
}
async fn status_by_internal_id(
&self,
context: &Self::Context,
id: Self::InternalId,
) -> Result<ElementStatus, Self::Error> {
let Some(element) = self
.read_elements(context, std::slice::from_ref(&id))
.await
.map_err(|e| ANNError::log_index_error(e.to_string()))?
.pop()
.flatten()
else {
return Err(ANNError::log_index_error(format!("DiskANN element {id} is missing")));
};
if element.deleted {
Ok(ElementStatus::Deleted)
} else {
Ok(ElementStatus::Valid)
}
}
fn status_by_external_id(
&self,
context: &Self::Context,
gid: &Self::ExternalId,
) -> impl std::future::Future<Output = Result<ElementStatus, Self::Error>> + Send {
self.status_by_internal_id(context, *gid)
}
}
pub(super) struct DiskAnnSearchAccessor<'a, T: DiskAnnVectorElement> {
provider: &'a DiskAnnProvider,
context: &'a DiskAnnProviderContext,
query: <T as VectorRepr>::QueryDistance,
buffer: Box<[T]>,
}
impl<'a, T> DiskAnnSearchAccessor<'a, T>
where
T: DiskAnnVectorElement,
{
fn new(
provider: &'a DiskAnnProvider,
context: &'a DiskAnnProviderContext,
query: &[T],
) -> Result<Self, ANNError> {
if query.len() != provider.dim {
return Err(vector_len_mismatch(query.len(), provider.dim));
}
Ok(Self {
provider,
context,
query: T::query_distance(query, provider.metric),
buffer: vec![T::default(); provider.dim].into_boxed_slice(),
})
}
fn distance_to(&mut self, element: &DiskAnnElement) -> ANNResult<f32> {
T::copy_from_serialized(&element.vector, &mut self.buffer)?;
Ok(self.query.evaluate_similarity(&self.buffer))
}
}
impl<T: DiskAnnVectorElement> HasId for DiskAnnSearchAccessor<'_, T> {
type Id = ElementId;
}
impl<T> SearchAccessor for DiskAnnSearchAccessor<'_, T>
where
T: DiskAnnVectorElement,
{
async fn starting_points(&self) -> ANNResult<Vec<ElementId>> {
self.provider.valid_starting_points(self.context).await
}
async fn start_point_distances<F>(&mut self, mut f: F) -> ANNResult<()>
where
F: FnMut(ElementId, f32) + Send,
{
let ids = self.provider.valid_starting_points(self.context).await?;
let elements = self
.provider
.read_elements(self.context, &ids)
.await
.map_err(|e| ANNError::log_index_error(e.to_string()))?;
for (id, element) in ids.into_iter().zip(elements) {
let Some(element) = element else {
return Err(ANNError::log_index_error(format!(
"DiskANN start point {id} is missing"
)));
};
let distance = self.distance_to(&element)?;
f(id, distance);
}
Ok(())
}
async fn expand_beam<Itr, P, F>(
&mut self,
ids: Itr,
mut pred: P,
mut on_neighbors: F,
) -> ANNResult<()>
where
Itr: Iterator<Item = ElementId> + Send,
P: HybridPredicate<ElementId> + Send + Sync,
F: FnMut(ElementId, f32) + Send,
{
let ids: Vec<_> = ids.collect();
let nodes = self
.provider
.read_nodes(self.context, &ids)
.await
.map_err(|e| ANNError::log_index_error(e.to_string()))?;
let mut candidates = Vec::new();
for node in nodes.iter().flatten() {
for &id in &node.neighbors {
if pred.eval_mut(&id) {
candidates.push(id);
}
}
}
let elements = self
.provider
.read_elements(self.context, &candidates)
.await
.map_err(|e| ANNError::log_index_error(e.to_string()))?;
for (id, element) in candidates.into_iter().zip(elements) {
if let Some(element) = element {
let distance = self.distance_to(&element)?;
on_neighbors(id, distance);
}
}
Ok(())
}
}
type DiskAnnWorkingSet<T> = workingset::Map<ElementId, Box<[T]>, workingset::map::Ref<[T]>>;
type DiskAnnWorkingSetView<'a, T> =
workingset::map::View<'a, ElementId, Box<[T]>, workingset::map::Ref<[T]>>;
pub(super) struct DiskAnnPruneAccessor<'a, T: DiskAnnVectorElement> {
provider: &'a DiskAnnProvider,
context: &'a DiskAnnProviderContext,
set: DiskAnnWorkingSet<T>,
}
impl<'a, T> DiskAnnPruneAccessor<'a, T>
where
T: DiskAnnVectorElement,
{
fn new(
provider: &'a DiskAnnProvider,
context: &'a DiskAnnProviderContext,
set: DiskAnnWorkingSet<T>,
) -> Self {
Self {
provider,
context,
set,
}
}
}
impl<T: DiskAnnVectorElement> HasId for DiskAnnPruneAccessor<'_, T> {
type Id = ElementId;
}
impl<T> PruneAccessor for DiskAnnPruneAccessor<'_, T>
where
T: DiskAnnVectorElement,
{
type Neighbors<'b>
= DiskAnnNeighborAccessor<'b>
where
Self: 'b;
type ElementRef<'b> = &'b [T];
type View<'b>
= DiskAnnWorkingSetView<'b, T>
where
Self: 'b;
type Distance<'b>
= <T as VectorRepr>::Distance
where
Self: 'b;
fn neighbors(&mut self) -> Self::Neighbors<'_> {
DiskAnnNeighborAccessor {
provider: self.provider,
context: self.context,
}
}
async fn fill<Itr>(&mut self, itr: Itr) -> ANNResult<(Self::View<'_>, Self::Distance<'_>)>
where
Itr: ExactSizeIterator<Item = ElementId> + Clone + Send + Sync,
{
let missing: Vec<ElementId> = itr.clone().filter(|id| !self.set.contains_key(id)).collect();
let fetched = self
.provider
.read_elements(self.context, &missing)
.await
.map_err(|e| ANNError::log_index_error(e.to_string()))?;
let mut decoded: HashMap<ElementId, Box<[T]>> = HashMap::with_capacity(missing.len());
for (id, element) in missing.into_iter().zip(fetched) {
if let Some(element) = element
&& !element.deleted
{
let mut buffer = vec![T::default(); self.provider.dim].into_boxed_slice();
T::copy_from_serialized(&element.vector, &mut buffer)?;
decoded.insert(id, buffer);
}
}
let distance = T::distance(self.provider.metric, Some(self.provider.dim));
let view =
self.set.fill(itr, |id| -> ANNResult<Option<Box<[T]>>> { Ok(decoded.remove(&id)) })?;
Ok((view, distance))
}
}
#[derive(Clone, Copy)]
pub(super) struct DiskAnnNeighborAccessor<'a> {
provider: &'a DiskAnnProvider,
context: &'a DiskAnnProviderContext,
}
impl HasId for DiskAnnNeighborAccessor<'_> {
type Id = ElementId;
}
impl NeighborAccessor for DiskAnnNeighborAccessor<'_> {
async fn get_neighbors(
&mut self,
id: ElementId,
neighbors: &mut AdjacencyList<ElementId>,
) -> ANNResult<()> {
if let Some(node) = self
.provider
.read_nodes(self.context, &[id])
.await
.map_err(|e| ANNError::log_index_error(e.to_string()))?
.into_iter()
.next()
.flatten()
{
neighbors.overwrite_trusted(&node.neighbors);
} else {
neighbors.clear();
}
Ok(())
}
}
impl NeighborAccessorMut for DiskAnnNeighborAccessor<'_> {
async fn set_neighbors(&mut self, id: ElementId, neighbors: &[ElementId]) -> ANNResult<()> {
let node = DiskAnnNode {
neighbors: neighbors.to_vec(),
};
self.context
.tx
.set(&self.context.ikb.new_dn_key(id), &node)
.await
.map_err(|e| ANNError::log_index_error(e.to_string()))?;
self.provider.cache.insert_node(self.provider.cache_index(), id, node);
Ok(())
}
async fn append_vector(&mut self, id: ElementId, neighbors: &[ElementId]) -> ANNResult<()> {
let key = self.context.ikb.new_dn_key(id);
let mut node: DiskAnnNode = self
.provider
.read_nodes(self.context, &[id])
.await
.map_err(|e| ANNError::log_index_error(e.to_string()))?
.into_iter()
.next()
.flatten()
.map(|node| (*node).clone())
.unwrap_or_default();
for neighbor in neighbors {
if !node.neighbors.contains(neighbor) {
node.neighbors.push(*neighbor);
}
}
self.context
.tx
.set(&key, &node)
.await
.map_err(|e| ANNError::log_index_error(e.to_string()))?;
self.provider.cache.insert_node(self.provider.cache_index(), id, node);
Ok(())
}
}
#[derive(Debug)]
pub(super) struct DiskAnnStrategy<T>(PhantomData<fn() -> T>);
impl<T> Default for DiskAnnStrategy<T> {
fn default() -> Self {
Self(PhantomData)
}
}
impl<T> Clone for DiskAnnStrategy<T> {
fn clone(&self) -> Self {
Self::default()
}
}
impl<'a, T> SearchStrategy<'a, DiskAnnProvider, &'a [T]> for DiskAnnStrategy<T>
where
T: DiskAnnVectorElement,
{
type SearchAccessorError = ANNError;
type SearchAccessor = DiskAnnSearchAccessor<'a, T>;
fn search_accessor(
&'a self,
provider: &'a DiskAnnProvider,
context: &'a DiskAnnProviderContext,
query: &'a [T],
) -> Result<DiskAnnSearchAccessor<'a, T>, ANNError> {
DiskAnnSearchAccessor::new(provider, context, query)
}
}
impl<'a, T> DefaultPostProcessor<'a, DiskAnnProvider, &'a [T]> for DiskAnnStrategy<T>
where
T: DiskAnnVectorElement,
{
default_post_processor!(CopyIds);
}
impl<T> PruneStrategy<DiskAnnProvider> for DiskAnnStrategy<T>
where
T: DiskAnnVectorElement,
{
type PruneAccessor<'a> = DiskAnnPruneAccessor<'a, T>;
type PruneAccessorError = ANNError;
fn prune_accessor<'a>(
&'a self,
provider: &'a DiskAnnProvider,
context: &'a DiskAnnProviderContext,
capacity: usize,
) -> Result<Self::PruneAccessor<'a>, ANNError> {
let set = workingset::map::Builder::new(workingset::map::Capacity::Default).build(capacity);
Ok(DiskAnnPruneAccessor::new(provider, context, set))
}
}
impl<'a, T> InsertStrategy<'a, DiskAnnProvider, &'a [T]> for DiskAnnStrategy<T>
where
T: DiskAnnVectorElement,
{
type PruneStrategy = Self;
fn prune_strategy(&self) -> Self::PruneStrategy {
self.clone()
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::catalog::{DatabaseId, IndexId, NamespaceId, TableId};
use crate::idx::trees::diskann::cache::DiskAnnCache;
use crate::kvs::{Datastore, LockType, TransactionType};
fn ikb() -> IndexKeyBase {
IndexKeyBase::new(NamespaceId(1), DatabaseId(2), "tb".into(), IndexId(3))
}
async fn provider_and_context() -> Result<(DiskAnnProvider, DiskAnnProviderContext)> {
let ds = Datastore::new("memory").await?;
let tx = Arc::new(ds.transaction(TransactionType::Write, LockType::Optimistic).await?);
let ikb = ikb();
let provider =
DiskAnnProvider::new(ikb, TableId(4), DiskAnnCache::new(1024 * 1024), 2, Metric::L2);
let context = provider.context(tx);
Ok((provider, context))
}
#[tokio::test]
async fn diskann_provider_batches_element_reads_and_reports_missing() -> Result<()> {
let (provider, context) = provider_and_context().await?;
context
.tx
.set(
&context.ikb.new_de_key(1),
&DiskAnnElement {
vector: SerializedVector::F32(vec![1.0, 2.0]),
deleted: false,
},
)
.await?;
context
.tx
.set(
&context.ikb.new_de_key(2),
&DiskAnnElement {
vector: SerializedVector::F32(vec![3.0, 4.0]),
deleted: false,
},
)
.await?;
let elements = provider.read_elements(&context, &[1, 2, 3]).await?;
assert_eq!(
elements[0].as_ref().map(|e| e.vector.clone()),
Some(SerializedVector::F32(vec![1.0, 2.0]))
);
assert_eq!(
elements[1].as_ref().map(|e| e.vector.clone()),
Some(SerializedVector::F32(vec![3.0, 4.0]))
);
assert!(elements[2].is_none());
assert!(provider.cache.get_element(provider.cache_index(), 1).is_none());
assert!(provider.cache.get_element(provider.cache_index(), 2).is_none());
Ok(())
}
#[tokio::test]
async fn diskann_provider_batches_nodes_and_updates_cache_on_mutation() -> Result<()> {
let (provider, context) = provider_and_context().await?;
context
.tx
.set(
&context.ikb.new_dn_key(1),
&DiskAnnNode {
neighbors: vec![2, 3],
},
)
.await?;
context
.tx
.set(
&context.ikb.new_dn_key(2),
&DiskAnnNode {
neighbors: vec![4],
},
)
.await?;
let nodes = provider.read_nodes(&context, &[1, 2]).await?;
assert_eq!(nodes[0].as_ref().unwrap().neighbors, vec![2, 3]);
assert_eq!(nodes[1].as_ref().unwrap().neighbors, vec![4]);
let mut neighbor_accessor = DiskAnnNeighborAccessor {
provider: &provider,
context: &context,
};
neighbor_accessor.set_neighbors(1, &[8, 9]).await.unwrap();
let cached = provider.cache.get_node(provider.cache_index(), 1).unwrap();
assert_eq!(cached.neighbors, vec![8, 9]);
Ok(())
}
#[tokio::test]
async fn diskann_provider_batch_vectors_omits_deleted_elements() -> Result<()> {
let (provider, context) = provider_and_context().await?;
context
.tx
.set(
&context.ikb.new_de_key(1),
&DiskAnnElement {
vector: SerializedVector::F32(vec![1.0, 2.0]),
deleted: false,
},
)
.await?;
context
.tx
.set(
&context.ikb.new_de_key(2),
&DiskAnnElement {
vector: SerializedVector::F32(vec![3.0, 4.0]),
deleted: true,
},
)
.await?;
let vectors = provider.get_vectors(&context, &[1, 2, 3]).await?;
assert!(vectors[0].is_some());
assert!(vectors[1].is_none());
assert!(vectors[2].is_none());
Ok(())
}
}