1use std::{fmt::Debug, future::Future, num::NonZeroUsize};
7
8use crate::storage::{StorageReadProvider, StorageWriteProvider};
9#[cfg(test)]
10use diskann::neighbor::Neighbor;
11use diskann::{
12 ANNError, ANNResult,
13 graph::AdjacencyList,
14 provider::{
15 DataProvider, DefaultAccessor, DefaultContext, Delete, ElementStatus, ExecutionContext,
16 NeighborAccessor, NeighborAccessorMut, NoopGuard, SetElement,
17 },
18 utils::{IntoUsize, ONE, VectorRepr},
19};
20use diskann_utils::future::AsyncFriendly;
21use diskann_vector::distance::Metric;
22
23use crate::{
24 model::graph::provider::async_::{
25 SimpleNeighborProviderAsync, StartPoints, TableDeleteProviderAsync,
26 common::{
27 CreateDeleteProvider, CreateVectorStore, NoDeletes, NoStore, PrefetchCacheLineLevel,
28 SetElementHelper, VectorStore,
29 },
30 },
31 storage::{AsyncIndexMetadata, AsyncQuantLoadContext, DiskGraphOnly, LoadWith, SaveWith},
32};
33
34pub struct DefaultProvider<U, V = NoStore, D = NoDeletes, Ctx = DefaultContext> {
230 pub base_vectors: U,
232
233 pub aux_vectors: V,
235
236 pub(crate) neighbor_provider: SimpleNeighborProviderAsync<u32>,
238
239 pub(super) deleted: D,
243
244 pub(super) metric: Metric,
246
247 pub(super) start_points: StartPoints,
248
249 context: std::marker::PhantomData<Ctx>,
250}
251
252#[derive(Debug, Clone)]
253pub struct DefaultProviderParameters {
254 pub max_points: usize,
256
257 pub frozen_points: NonZeroUsize,
260
261 pub dim: usize,
263
264 pub metric: Metric,
266
267 pub prefetch_lookahead: Option<usize>,
272
273 pub prefetch_cache_line_level: Option<PrefetchCacheLineLevel>,
274
275 pub max_degree: u32,
277}
278
279impl DefaultProviderParameters {
280 pub fn simple(max_points: usize, dim: usize, metric: Metric, max_degree: u32) -> Self {
281 Self {
282 max_points,
283 frozen_points: ONE,
284 metric,
285 dim,
286 prefetch_lookahead: None,
287 prefetch_cache_line_level: None,
288 max_degree,
289 }
290 }
291}
292
293impl<U, V, D, Ctx> DefaultProvider<U, V, D, Ctx> {
294 pub fn new_empty<CU, CV, CD>(
304 params: DefaultProviderParameters,
305 base_precursor: CU,
306 aux_precursor: CV,
307 delete_precursor: CD,
308 ) -> ANNResult<Self>
309 where
310 CU: CreateVectorStore<Target = U>,
311 CV: CreateVectorStore<Target = V>,
312 CD: CreateDeleteProvider<Target = D>,
313 {
314 let npts = params.max_points + params.frozen_points.get();
315 Ok(Self {
316 base_vectors: base_precursor.create(npts, params.metric, params.prefetch_lookahead),
317 aux_vectors: aux_precursor.create(npts, params.metric, params.prefetch_lookahead),
318 neighbor_provider: SimpleNeighborProviderAsync::new(npts, 1, params.max_degree, 1.0),
319 deleted: delete_precursor.create(npts),
320 metric: params.metric,
321 start_points: StartPoints::new(params.max_points as u32, params.frozen_points)?,
322 context: std::marker::PhantomData,
323 })
324 }
325
326 #[cfg(test)]
329 pub(crate) fn is_not_start_point(&self) -> impl Fn(&Neighbor<u32>) -> bool {
330 let range = self.start_points.range();
331 move |neighbor| !range.contains(&neighbor.id)
332 }
333
334 pub fn starting_points(&self) -> ANNResult<Vec<u32>> {
336 Ok(self.start_points.range().collect())
337 }
338
339 pub fn iter(&self) -> std::ops::Range<u32> {
341 0..self.start_points.end()
342 }
343
344 pub fn neighbors(&self) -> &SimpleNeighborProviderAsync<u32> {
346 &self.neighbor_provider
347 }
348
349 pub fn num_start_points(&self) -> usize {
350 self.start_points.len()
351 }
352
353 pub fn capacity(&self) -> usize {
355 self.start_points.start().into_usize()
356 }
357
358 pub fn total_points(&self) -> usize {
360 self.start_points.end().into_usize()
361 }
362}
363
364impl<U, V, D, Ctx> IntoIterator for &DefaultProvider<U, V, D, Ctx> {
366 type Item = u32;
367 type IntoIter = std::ops::Range<u32>;
368 fn into_iter(self) -> Self::IntoIter {
369 self.iter()
370 }
371}
372
373impl<U, V, Ctx> DefaultProvider<U, V, TableDeleteProviderAsync, Ctx> {
374 pub fn clear_delete_set(&self) {
376 self.deleted.clear();
377 }
378}
379
380impl<U, V, D, Ctx> DefaultProvider<U, V, D, Ctx>
381where
382 U: VectorStore,
383 V: VectorStore,
384{
385 pub fn counts_for_get_vector(&self) -> (usize, usize) {
387 (
388 self.base_vectors.count_for_get_vector(),
389 self.aux_vectors.count_for_get_vector(),
390 )
391 }
392}
393
394pub trait SetStartPoints<T>
395where
396 T: ?Sized + 'static,
397{
398 fn set_start_points<'a, Itr>(&self, itr: Itr) -> ANNResult<()>
399 where
400 Itr: ExactSizeIterator<Item = &'a T> + 'a;
401}
402
403impl<T, U, V, D> SetStartPoints<[T]> for DefaultProvider<U, V, D>
404where
405 U: SetElementHelper<T>,
406 V: SetElementHelper<T>,
407 T: std::fmt::Debug + 'static,
408{
409 fn set_start_points<'a, Itr>(&self, itr: Itr) -> ANNResult<()>
410 where
411 Itr: ExactSizeIterator<Item = &'a [T]> + 'a,
412 {
413 let start_points = self.start_points.range();
414 if itr.len() != start_points.len() {
415 return Err(ANNError::log_async_index_error(format!(
416 "expected `itr` to contain `{}` items, instead it has {}",
417 start_points.len(),
418 itr.len(),
419 )));
420 }
421
422 for (i, v) in std::iter::zip(start_points, itr) {
423 self.aux_vectors.set_element(&i, v)?;
424 self.base_vectors.set_element(&i, v)?;
425 }
426
427 Ok(())
428 }
429}
430
431impl<U, V, D, Ctx> SaveWith<(u32, AsyncIndexMetadata)> for DefaultProvider<U, V, D, Ctx>
436where
437 U: AsyncFriendly + SaveWith<AsyncIndexMetadata>,
438 V: AsyncFriendly + SaveWith<AsyncIndexMetadata>,
439 D: AsyncFriendly,
440 ANNError: From<U::Error> + From<V::Error>,
441 Ctx: ExecutionContext,
442{
443 type Ok = ();
444 type Error = ANNError;
445
446 async fn save_with<P>(
447 &self,
448 provider: &P,
449 auxiliary: &(u32, AsyncIndexMetadata),
450 ) -> Result<Self::Ok, Self::Error>
451 where
452 P: StorageWriteProvider,
453 {
454 self.base_vectors.save_with(provider, &auxiliary.1).await?;
455 self.aux_vectors.save_with(provider, &auxiliary.1).await?;
456 self.neighbor_provider
457 .save_with(provider, auxiliary)
458 .await?;
459 Ok(())
460 }
461}
462
463impl<U, V, D, Ctx> SaveWith<(u32, u32, DiskGraphOnly)> for DefaultProvider<U, V, D, Ctx>
464where
465 U: AsyncFriendly,
466 V: AsyncFriendly,
467 D: AsyncFriendly,
468 Ctx: ExecutionContext,
469{
470 type Ok = ();
471 type Error = ANNError;
472
473 async fn save_with<P>(
474 &self,
475 provider: &P,
476 auxiliary: &(u32, u32, DiskGraphOnly),
477 ) -> Result<Self::Ok, Self::Error>
478 where
479 P: StorageWriteProvider,
480 {
481 self.neighbor_provider
482 .save_with(provider, auxiliary)
483 .await?;
484 Ok(())
485 }
486}
487
488impl<U, V, D, Ctx> LoadWith<AsyncQuantLoadContext> for DefaultProvider<U, V, D, Ctx>
493where
494 U: VectorStore + LoadWith<AsyncQuantLoadContext>,
495 V: VectorStore + AsyncFriendly + LoadWith<AsyncQuantLoadContext>,
496 D: AsyncFriendly + LoadWith<usize>,
497 ANNError: From<U::Error> + From<V::Error> + From<D::Error>,
498 Ctx: ExecutionContext,
499{
500 type Error = ANNError;
501
502 async fn load_with<P>(provider: &P, ctx: &AsyncQuantLoadContext) -> ANNResult<Self>
503 where
504 P: StorageReadProvider,
505 {
506 let base_vectors = U::load_with(provider, ctx).await?;
507 let aux_vectors = V::load_with(provider, ctx).await?;
508 let deleted = D::load_with(provider, &base_vectors.total()).await?;
509
510 let npts = std::cmp::max(base_vectors.total(), aux_vectors.total());
513
514 let valid_points = npts
515 .checked_sub(ctx.num_frozen_points.get())
516 .ok_or_else(|| {
517 ANNError::log_index_error(format_args!(
518 "Expected {} start points but the stored index only has {} total points",
519 ctx.num_frozen_points.get(),
520 base_vectors.total(),
521 ))
522 })?;
523 let start_points = StartPoints::new(valid_points as u32, ctx.num_frozen_points)?;
524 Ok(Self {
525 base_vectors,
526 aux_vectors,
527 neighbor_provider: SimpleNeighborProviderAsync::load_with(provider, ctx).await?,
528 deleted,
529 metric: ctx.metric,
530 start_points,
531 context: std::marker::PhantomData,
532 })
533 }
534}
535
536impl LoadWith<usize> for NoDeletes {
537 type Error = ANNError;
538
539 async fn load_with<P>(_: &P, _num_points: &usize) -> ANNResult<Self>
540 where
541 P: StorageReadProvider,
542 {
543 Ok(NoDeletes)
544 }
545}
546
547impl LoadWith<usize> for TableDeleteProviderAsync {
548 type Error = ANNError;
549
550 async fn load_with<P>(_: &P, num_points: &usize) -> ANNResult<Self>
551 where
552 P: StorageReadProvider,
553 {
554 Ok(TableDeleteProviderAsync::new(*num_points))
555 }
556}
557
558impl<U, V, D, Ctx> DataProvider for DefaultProvider<U, V, D, Ctx>
563where
564 U: AsyncFriendly,
565 V: AsyncFriendly,
566 D: AsyncFriendly,
567 Ctx: ExecutionContext,
568{
569 type Context = Ctx;
570 type InternalId = u32;
572 type ExternalId = u32;
574 type Error = ANNError;
576 type Guard = NoopGuard<u32>;
578
579 fn to_internal_id(
581 &self,
582 _context: &Self::Context,
583 gid: &Self::ExternalId,
584 ) -> Result<Self::InternalId, Self::Error> {
585 Ok(*gid)
586 }
587
588 fn to_external_id(
590 &self,
591 _context: &Self::Context,
592 id: Self::InternalId,
593 ) -> Result<Self::ExternalId, Self::Error> {
594 Ok(id)
595 }
596}
597
598impl<U, V, Ctx> Delete for DefaultProvider<U, V, TableDeleteProviderAsync, Ctx>
600where
601 U: AsyncFriendly,
602 V: AsyncFriendly,
603 Ctx: ExecutionContext,
604{
605 fn release(
606 &self,
607 _context: &Ctx,
608 id: Self::InternalId,
609 ) -> impl Future<Output = Result<(), Self::Error>> + Send {
610 self.deleted.undelete(id.into_usize());
611 let res = self
612 .neighbor_provider
613 .set_neighbors_sync(id.into_usize(), &[])
614 .map_err(|err| err.context(format!("resetting neighbors for undeleted id {}", id)));
615 std::future::ready(res)
616 }
617
618 #[inline]
620 fn delete(
621 &self,
622 _context: &Ctx,
623 gid: &Self::ExternalId,
624 ) -> impl Future<Output = Result<(), Self::Error>> + Send {
625 self.deleted.delete(gid.into_usize());
626 std::future::ready(Ok(()))
627 }
628
629 #[inline]
631 fn status_by_external_id(
632 &self,
633 context: &Ctx,
634 gid: &Self::ExternalId,
635 ) -> impl Future<Output = Result<ElementStatus, Self::Error>> + Send {
636 self.status_by_internal_id(context, *gid)
638 }
639
640 #[inline]
642 fn status_by_internal_id(
643 &self,
644 _context: &Ctx,
645 id: Self::InternalId,
646 ) -> impl Future<Output = Result<ElementStatus, Self::Error>> + Send {
647 let status = if self.deleted.is_deleted(id.into_usize()) {
648 ElementStatus::Deleted
649 } else {
650 ElementStatus::Valid
651 };
652 std::future::ready(Ok(status))
653 }
654}
655
656impl NeighborAccessor for &SimpleNeighborProviderAsync<u32> {
657 async fn get_neighbors(
658 self,
659 id: Self::Id,
660 neighbors: &mut AdjacencyList<Self::Id>,
661 ) -> ANNResult<Self> {
662 self.get_neighbors_sync(id.into_usize(), neighbors)?;
663 Ok(self)
664 }
665}
666
667impl NeighborAccessorMut for &SimpleNeighborProviderAsync<u32> {
668 async fn set_neighbors(self, id: u32, neighbors: &[u32]) -> ANNResult<Self> {
669 self.set_neighbors_sync(id.into_usize(), neighbors)?;
670 Ok(self)
671 }
672
673 async fn append_vector(self, id: u32, new_neighbor_ids: &[u32]) -> ANNResult<Self> {
674 self.append_vector_sync(id.into_usize(), new_neighbor_ids)?;
675 Ok(self)
676 }
677}
678
679impl<U, V, D, Ctx> DefaultAccessor for DefaultProvider<U, V, D, Ctx>
680where
681 U: AsyncFriendly,
682 V: AsyncFriendly,
683 D: AsyncFriendly,
684 Ctx: ExecutionContext,
685{
686 type Accessor<'a> = &'a SimpleNeighborProviderAsync<u32>;
687 fn default_accessor(&self) -> Self::Accessor<'_> {
688 self.neighbors()
689 }
690}
691
692impl<U, V, D, Ctx, T> SetElement<&[T]> for DefaultProvider<U, V, D, Ctx>
698where
699 T: VectorRepr,
700 U: AsyncFriendly + SetElementHelper<T>,
701 V: AsyncFriendly + SetElementHelper<T>,
702 D: AsyncFriendly,
703 Ctx: ExecutionContext,
704{
705 type SetError = ANNError;
706
707 fn set_element(
709 &self,
710 _context: &Self::Context,
711 id: &u32,
712 element: &[T],
713 ) -> impl Future<Output = Result<Self::Guard, Self::SetError>> + Send {
714 if let Err(err) = self.aux_vectors.set_element(id, element) {
716 return std::future::ready(Err(err));
717 }
718
719 if let Err(err) = self.base_vectors.set_element(id, element) {
721 return std::future::ready(Err(err));
722 }
723
724 std::future::ready(Ok(NoopGuard::new(*id)))
726 }
727}
728
729#[cfg(test)]
734mod tests {
735 use super::*;
736 use crate::model::graph::provider::async_::{
737 common::{NoStore, TableBasedDeletes},
738 inmem::CreateFullPrecision,
739 };
740
741 #[tokio::test]
742 async fn test_data_provider_and_delete_interface() {
743 let ctx = &DefaultContext;
744 let provider = DefaultProvider::new_empty(
745 DefaultProviderParameters {
746 max_points: 10,
747 frozen_points: NonZeroUsize::new(2).unwrap(),
748 dim: 5,
749 metric: Metric::L2,
750 prefetch_lookahead: None,
751 max_degree: (64.0 * 1.2) as u32,
752 prefetch_cache_line_level: None,
753 },
754 CreateFullPrecision::<f32>::new(5, None),
755 NoStore,
756 TableBasedDeletes,
757 )
758 .unwrap();
759
760 assert_eq!((&provider).into_iter(), 0..(10 + 2));
762
763 let iter = provider.iter();
764 for i in iter.clone() {
765 assert_eq!(provider.to_external_id(ctx, i).unwrap(), i);
766 assert_eq!(provider.to_internal_id(ctx, &i).unwrap(), i);
767 assert_eq!(
768 provider.status_by_internal_id(ctx, i).await.unwrap(),
769 ElementStatus::Valid
770 );
771 assert_eq!(
772 provider.status_by_external_id(ctx, &i).await.unwrap(),
773 ElementStatus::Valid
774 );
775
776 provider.delete(ctx, &i).await.unwrap();
778 assert_eq!(
779 provider.status_by_internal_id(ctx, i).await.unwrap(),
780 ElementStatus::Deleted
781 );
782 assert_eq!(
783 provider.status_by_external_id(ctx, &i).await.unwrap(),
784 ElementStatus::Deleted
785 );
786 }
787
788 for i in iter.clone() {
790 provider
792 .neighbor_provider
793 .set_neighbors(i, &[1, 2])
794 .await
795 .unwrap();
796 provider.release(ctx, i).await.unwrap();
797 assert_eq!(
798 provider.status_by_internal_id(ctx, i).await.unwrap(),
799 ElementStatus::Valid
800 );
801 assert_eq!(
802 provider.status_by_external_id(ctx, &i).await.unwrap(),
803 ElementStatus::Valid
804 );
805 let mut neighbors = AdjacencyList::new();
807 provider
808 .neighbor_provider
809 .get_neighbors(i, &mut neighbors)
810 .await
811 .unwrap();
812 assert!(neighbors.to_vec().is_empty());
813
814 provider.delete(ctx, &i).await.unwrap();
816 }
817
818 provider.clear_delete_set();
819 for i in iter.clone() {
820 assert_eq!(
821 provider.status_by_internal_id(ctx, i).await.unwrap(),
822 ElementStatus::Valid
823 );
824 assert_eq!(
825 provider.status_by_external_id(ctx, &i).await.unwrap(),
826 ElementStatus::Valid
827 );
828 }
829
830 assert!(
832 provider
833 .set_element(ctx, &100, &[1.0, 2.0, 3.0, 4.0])
834 .await
835 .is_err()
836 );
837 }
838}