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