use std::{marker::PhantomData, pin::Pin, sync::Arc};
use diskann::{
graph::{
glue::{InsertStrategy, PruneStrategy},
Config, DiskANNIndex,
},
provider::DefaultContext,
utils::VectorRepr,
ANNError, ANNResult,
};
use diskann_providers::storage::{DynWriteProvider, StorageReadProvider, WriteProviderWrapper};
use diskann_providers::{
index::diskann_async,
model::{
graph::provider::async_::{
common::{FullPrecision, NoDeletes, NoStore, Quantized, SetElementHelper, VectorStore},
inmem::{
DefaultProvider, DefaultProviderParameters, DefaultQuant, FullPrecisionProvider,
SQStore, SetStartPoints,
},
},
IndexConfiguration,
},
storage::{
index_storage::load_index, load_fp_index, AsyncIndexMetadata, DiskGraphOnly, SaveWith,
},
};
use diskann_utils::future::{AsyncFriendly, SendFuture};
use super::quantizer::BuildQuantizer;
pub(super) trait InmemIndexBuilder<T: Sized>: Send + Sync {
fn capacity(&self) -> usize;
fn total_points(&self) -> usize;
fn set_start_point(&self, start_point: &[T]) -> ANNResult<()>;
fn insert_vector<'a>(
&'a self,
id: u32,
vector: &'a [T],
) -> Pin<Box<dyn SendFuture<ANNResult<()>> + 'a>>;
fn final_prune(
&self,
range: core::ops::Range<u32>,
) -> Pin<Box<dyn SendFuture<ANNResult<()>> + '_>>;
fn save_index<'a>(
&'a self,
storage_provider: &'a dyn DynWriteProvider,
metadata: &'a AsyncIndexMetadata,
) -> Pin<Box<dyn SendFuture<ANNResult<()>> + 'a>>;
fn save_graph<'a>(
&'a self,
storage_provider: &'a dyn DynWriteProvider,
start_point_and_path: &'a (u32, DiskGraphOnly),
) -> Pin<Box<dyn SendFuture<ANNResult<()>> + 'a>>;
#[cfg(debug_assertions)]
fn counts_for_get_vector(&self) -> (usize, usize);
#[cfg(debug_assertions)]
fn count_reachable_nodes(&self) -> Pin<Box<dyn SendFuture<ANNResult<usize>> + '_>>;
}
impl<T> InmemIndexBuilder<T> for DiskANNIndex<FullPrecisionProvider<T>>
where
T: VectorRepr,
{
fn capacity(&self) -> usize {
self.provider().capacity()
}
fn total_points(&self) -> usize {
self.provider().total_points()
}
fn set_start_point(&self, start_point: &[T]) -> ANNResult<()> {
self.provider()
.set_start_points(std::iter::once(start_point))
}
fn insert_vector<'a>(
&'a self,
id: u32,
vector: &'a [T],
) -> Pin<Box<dyn SendFuture<ANNResult<()>> + 'a>> {
Box::pin(async move {
self.insert(&FullPrecision, &DefaultContext, &id, vector)
.await
})
}
fn final_prune(
&self,
range: core::ops::Range<u32>,
) -> Pin<Box<dyn SendFuture<ANNResult<()>> + '_>> {
Box::pin(async move {
self.prune_range(&FullPrecision, &DefaultContext, range)
.await
})
}
fn save_index<'a>(
&'a self,
storage_provider: &'a dyn DynWriteProvider,
metadata: &'a AsyncIndexMetadata,
) -> Pin<Box<dyn SendFuture<ANNResult<()>> + 'a>> {
Box::pin(async move {
let wrapper = WriteProviderWrapper::new(storage_provider);
self.save_with(&wrapper, metadata).await
})
}
fn save_graph<'a>(
&'a self,
storage_provider: &'a dyn DynWriteProvider,
start_point_and_path: &'a (u32, DiskGraphOnly),
) -> Pin<Box<dyn SendFuture<ANNResult<()>> + 'a>> {
Box::pin(async move {
let wrapper = WriteProviderWrapper::new(storage_provider);
self.save_with(&wrapper, start_point_and_path).await
})
}
#[cfg(debug_assertions)]
fn counts_for_get_vector(&self) -> (usize, usize) {
self.provider().counts_for_get_vector()
}
#[cfg(debug_assertions)]
fn count_reachable_nodes(&self) -> Pin<Box<dyn SendFuture<ANNResult<usize>> + '_>> {
Box::pin(async move {
let provider = self.provider();
let start_points = provider.starting_points()?;
let mut neighbor_accessor = provider.neighbors();
self.count_reachable_nodes(&start_points, &mut neighbor_accessor)
.await
})
}
}
pub(super) struct QuantInMemBuilder<T, Q>
where
Q: AsyncFriendly,
{
index: DiskANNIndex<DefaultProvider<NoStore, Q>>,
_vector_data_type: PhantomData<T>,
}
impl<T, Q> QuantInMemBuilder<T, Q>
where
Q: AsyncFriendly,
{
pub fn new(index: DiskANNIndex<DefaultProvider<NoStore, Q>>) -> Self {
Self {
index,
_vector_data_type: PhantomData,
}
}
fn index(&self) -> &DiskANNIndex<DefaultProvider<NoStore, Q>> {
&self.index
}
}
impl<T, Q> InmemIndexBuilder<T> for QuantInMemBuilder<T, Q>
where
T: VectorRepr,
Q: AsyncFriendly + VectorStore + SetElementHelper<T>,
Quantized: for<'a> InsertStrategy<'a, DefaultProvider<NoStore, Q>, &'a [T]>
+ PruneStrategy<DefaultProvider<NoStore, Q>>,
DefaultProvider<NoStore, Q>: SaveWith<(u32, AsyncIndexMetadata), Error = ANNError>,
{
fn capacity(&self) -> usize {
self.index().provider().capacity()
}
fn total_points(&self) -> usize {
self.index().provider().total_points()
}
fn set_start_point(&self, start_point: &[T]) -> ANNResult<()> {
self.index()
.provider()
.set_start_points(std::iter::once(start_point))
}
fn insert_vector<'a>(
&'a self,
id: u32,
vector: &'a [T],
) -> Pin<Box<dyn SendFuture<ANNResult<()>> + 'a>> {
Box::pin(async move {
self.index()
.insert(&Quantized, &DefaultContext, &id, vector)
.await
})
}
fn final_prune(
&self,
range: core::ops::Range<u32>,
) -> Pin<Box<dyn SendFuture<ANNResult<()>> + '_>> {
Box::pin(async move {
self.index()
.prune_range(&Quantized, &DefaultContext, range)
.await
})
}
fn save_index<'a>(
&'a self,
storage_provider: &'a dyn DynWriteProvider,
metadata: &'a AsyncIndexMetadata,
) -> Pin<Box<dyn SendFuture<ANNResult<()>> + 'a>> {
Box::pin(async move {
let wrapper = WriteProviderWrapper::new(storage_provider);
self.index().save_with(&wrapper, metadata).await
})
}
fn save_graph<'a>(
&'a self,
storage_provider: &'a dyn DynWriteProvider,
start_point_and_path: &'a (u32, DiskGraphOnly),
) -> Pin<Box<dyn SendFuture<ANNResult<()>> + 'a>> {
Box::pin(async move {
let wrapper = WriteProviderWrapper::new(storage_provider);
self.index().save_with(&wrapper, start_point_and_path).await
})
}
#[cfg(debug_assertions)]
fn counts_for_get_vector(&self) -> (usize, usize) {
self.index().provider().counts_for_get_vector()
}
#[cfg(debug_assertions)]
fn count_reachable_nodes(&self) -> Pin<Box<dyn SendFuture<ANNResult<usize>> + '_>> {
Box::pin(async move {
let provider = self.index().provider();
let start_points = provider.starting_points()?;
let mut neighbor_accessor = provider.neighbors();
self.index()
.count_reachable_nodes(&start_points, &mut neighbor_accessor)
.await
})
}
}
pub(super) fn new_inmem_index_builder<T>(
config: Config,
params: DefaultProviderParameters,
build_quantizer: &BuildQuantizer,
) -> ANNResult<Arc<dyn InmemIndexBuilder<T>>>
where
T: VectorRepr,
{
match &build_quantizer {
BuildQuantizer::NoQuant(_) => diskann_async::new_index::<T, _>(config, params, NoDeletes)
.map(|index| index as Arc<dyn InmemIndexBuilder<T>>),
BuildQuantizer::Scalar1Bit(q) => {
let index = diskann_async::new_quant_only_index(config, params, q.clone(), NoDeletes)?;
Ok(Arc::new(QuantInMemBuilder::<T, _>::new(index)))
}
BuildQuantizer::PQ(table) => {
let index =
diskann_async::new_quant_only_index(config, params, table.clone(), NoDeletes)?;
Ok(Arc::new(QuantInMemBuilder::<T, _>::new(index)))
}
}
}
pub(super) async fn load_inmem_index_builder<T, P>(
storage_provider: &P,
build_quantizer: &BuildQuantizer,
config: IndexConfiguration,
index_path_prefix: &str,
) -> ANNResult<Arc<dyn InmemIndexBuilder<T>>>
where
P: StorageReadProvider,
T: VectorRepr,
{
match build_quantizer {
BuildQuantizer::NoQuant(_) => {
load_fp_index::<T, _, NoStore>(storage_provider, index_path_prefix, config)
.await
.map(|index| Arc::new(index) as Arc<dyn InmemIndexBuilder<T>>)
}
BuildQuantizer::Scalar1Bit(_) => {
let index =
load_index::<_, NoStore, SQStore<1>>(storage_provider, index_path_prefix, config)
.await?;
Ok(Arc::new(QuantInMemBuilder::<T, _>::new(index)))
}
BuildQuantizer::PQ(_) => {
let index =
load_index::<_, NoStore, DefaultQuant>(storage_provider, index_path_prefix, config)
.await?;
Ok(Arc::new(QuantInMemBuilder::<T, _>::new(index)))
}
}
}