use std::num::NonZeroU16;
use diskann::{ANNResult, neighbor::Neighbor};
use thiserror::Error;
use crate::{
counters::LocalCounters,
num::{Capacity, IdLimit, MaxDegree},
};
mod internal;
pub mod full;
pub use full::{Full, FullPrecision};
#[cfg(feature = "quantization")]
#[cfg_attr(docsrs, doc(cfg(feature = "quantization")))]
pub mod spherical;
#[cfg(feature = "quantization")]
#[cfg_attr(docsrs, doc(cfg(feature = "quantization")))]
pub use spherical::Spherical;
#[cfg(test)]
mod test;
pub trait RepresentationConfig {
type Representation: Representation;
fn build(self) -> ANNResult<Self::Representation>;
}
pub trait Representation: Send + Sync + 'static {
fn max_degree(&self) -> MaxDegree;
fn id_limit(&self) -> IdLimit;
fn capacity(&self) -> Capacity;
fn retire(&self, i: u32) -> ANNResult<()>;
fn is_readable(&self, i: u32) -> Option<bool>;
}
pub trait Set<T>: Representation {
type Guard<'a>: Guard;
fn set(&self, element: T) -> ANNResult<Self::Guard<'_>>;
}
pub trait Guard {
fn id(&self) -> u32;
fn publish(self);
}
pub trait Search: Send + Sync + 'static {
type Query<'a>;
#[doc(hidden)]
fn search_accessor<'a>(
&'a self,
query: Self::Query<'a>,
provider: &'a (dyn std::any::Any + Send + Sync),
counters: LocalCounters<'a>,
) -> ANNResult<crate::provider::SearchAccessor<'a>>;
}
pub trait Insert: Search + for<'a> Set<Self::Query<'a>> {
#[doc(hidden)]
fn insert_search_accessor<'a>(
&'a self,
query: Self::Query<'a>,
provider: &'a (dyn std::any::Any + Send + Sync),
counters: LocalCounters<'a>,
) -> ANNResult<crate::provider::SearchAccessor<'a>> {
self.search_accessor(query, provider, counters)
}
#[doc(hidden)]
fn prune_accessor<'a>(
&'a self,
counters: LocalCounters<'a>,
) -> ANNResult<crate::provider::PruneAccessor<'a>>;
}
pub(crate) unsafe trait ExpandBeam: Send + Sync + std::fmt::Debug {
fn evaluate(&self, i: u32) -> ANNResult<Option<f32>>;
fn id_limit(&self) -> IdLimit;
unsafe fn expand_beam(&self, list: &[u32], buffer: &mut [Neighbor<u32>]) -> ANNResult<usize>;
}
#[cfg(test)]
fn safe_expand_beam(
expand_beam: &dyn ExpandBeam,
list: &[u32],
buffer: &mut [Neighbor<u32>],
) -> ANNResult<usize> {
assert!(list.len() <= buffer.len());
let limit = expand_beam.id_limit();
for id in list.iter() {
if !limit.is_in_bounds(*id) {
panic!("id {} is not within {} -- {:?}", id, limit, expand_beam);
}
}
unsafe { expand_beam.expand_beam(list, buffer) }
}
pub(crate) trait PostProcess: Send + Sync + std::fmt::Debug {
fn post_process(&mut self, buffer: &mut Vec<Neighbor<u32>>) -> ANNResult<()>;
}
pub(crate) trait Prune: Send + Sync + std::fmt::Debug {
fn prepare(
&mut self,
items: hashbrown::hash_map::IterMut<'_, u32, Option<PruneKey>>,
) -> ANNResult<usize>;
fn evaluate(&self, a: PruneKey, b: PruneKey) -> f32;
}
#[derive(Debug, Clone, Copy)]
pub(crate) struct PruneKey(NonZeroU16);
impl PruneKey {
const ONE: Self = Self(NonZeroU16::new(1).unwrap());
pub(crate) fn counter() -> Self {
Self::ONE
}
pub(crate) fn increment(self) -> Result<Self, Overflow> {
match self.0.checked_add(1) {
Some(v) => Ok(Self(v)),
None => Err(Overflow),
}
}
pub(crate) fn index(self) -> usize {
usize::from(self.0.get()) - 1
}
}
#[derive(Debug, Error)]
#[error("prune list exceeded u16::MAX")]
pub(crate) struct Overflow;
diskann::convert_error!(Overflow);