use std::{fmt::Debug, num::NonZeroUsize};
use diskann_utils::future::SendFuture;
use thiserror::Error;
use super::Search;
use crate::{
ANNResult, convert_error,
error::IntoANNResult,
graph::{
glue::{SearchAccessor, SearchPostProcess, SearchStrategy},
index::{DiskANNIndex, SearchStats},
search::record::NoopSearchRecord,
search_output_buffer::SearchOutputBuffer,
},
provider::DataProvider,
};
#[derive(Debug, Error)]
pub enum KnnSearchError {
#[error("beam width cannot be zero")]
BeamWidthZero,
#[error("l_value cannot be zero")]
LZero,
}
convert_error!(KnnSearchError);
#[derive(Debug, Clone, Copy)]
pub struct Knn {
l_value: NonZeroUsize,
beam_width: NonZeroUsize,
}
impl Knn {
pub fn new(l_value: usize, beam_width: Option<usize>) -> Result<Self, KnnSearchError> {
let l_value = NonZeroUsize::new(l_value).ok_or(KnnSearchError::LZero)?;
const ONE: NonZeroUsize = NonZeroUsize::new(1).unwrap();
let beam_width = match beam_width {
Some(bw) => NonZeroUsize::new(bw).ok_or(KnnSearchError::BeamWidthZero)?,
None => ONE,
};
Ok(Self {
l_value,
beam_width,
})
}
pub fn new_default(l_value: usize) -> Result<Self, KnnSearchError> {
Self::new(l_value, None)
}
#[inline]
pub fn l_value(&self) -> NonZeroUsize {
self.l_value
}
#[inline]
pub fn beam_width(&self) -> NonZeroUsize {
self.beam_width
}
pub(crate) fn new_infallible(l_value: NonZeroUsize, beam_width: NonZeroUsize) -> Self {
Self {
l_value,
beam_width,
}
}
}
impl<'a, DP, S, T> Search<'a, DP, S, T> for Knn
where
DP: DataProvider,
S: SearchStrategy<'a, DP, T, SearchAccessor: SearchAccessor>,
T: Copy + Send + Sync,
{
type Output = SearchStats;
fn search<O, PP, OB>(
self,
index: &'a DiskANNIndex<DP>,
strategy: &'a S,
processor: PP,
context: &'a DP::Context,
query: T,
output: &mut OB,
) -> impl SendFuture<ANNResult<Self::Output>>
where
O: Send,
PP: SearchPostProcess<S::SearchAccessor, T, O> + Send + Sync,
OB: SearchOutputBuffer<O> + Send + ?Sized,
{
async move {
let mut accessor = strategy
.search_accessor(&index.data_provider, context, query)
.into_ann_result()?;
let num_start_ids = accessor.num_starting_points().await?;
let mut scratch = index.search_scratch(self.l_value.get(), num_start_ids);
let stats = index
.search_internal(
Some(self.beam_width.get()),
&mut accessor,
&mut scratch,
&mut NoopSearchRecord::new(),
)
.await?;
let result_count = processor
.post_process(&mut accessor, query, scratch.best.iter(), output)
.await
.into_ann_result()?;
Ok(stats.finish(result_count as u32))
}
}
}
#[derive(Debug)]
pub struct RecordedKnn<'r, SR: ?Sized> {
pub inner: Knn,
pub recorder: &'r mut SR,
}
impl<'r, SR: ?Sized> RecordedKnn<'r, SR> {
pub fn new(inner: Knn, recorder: &'r mut SR) -> Self {
Self { inner, recorder }
}
}
impl<'a, DP, S, T, SR> Search<'a, DP, S, T> for RecordedKnn<'a, SR>
where
DP: DataProvider,
S: SearchStrategy<'a, DP, T, SearchAccessor: SearchAccessor>,
T: Copy + Send + Sync,
SR: super::record::SearchRecord<DP::InternalId> + ?Sized,
{
type Output = SearchStats;
fn search<O, PP, OB>(
self,
index: &'a DiskANNIndex<DP>,
strategy: &'a S,
processor: PP,
context: &'a DP::Context,
query: T,
output: &mut OB,
) -> impl SendFuture<ANNResult<Self::Output>>
where
O: Send,
PP: SearchPostProcess<S::SearchAccessor, T, O> + Send + Sync,
OB: SearchOutputBuffer<O> + Send + ?Sized,
{
async move {
let mut accessor = strategy
.search_accessor(&index.data_provider, context, query)
.into_ann_result()?;
let num_start_ids = accessor.num_starting_points().await?;
let mut scratch = index.search_scratch(self.inner.l_value.get(), num_start_ids);
let stats = index
.search_internal(
Some(self.inner.beam_width.get()),
&mut accessor,
&mut scratch,
self.recorder,
)
.await?;
let result_count = processor
.post_process(&mut accessor, query, scratch.best.iter(), output)
.await
.into_ann_result()?;
Ok(stats.finish(result_count as u32))
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_knn_search_validation() {
assert!(Knn::new(100, None).is_ok());
assert!(Knn::new(100, Some(4)).is_ok());
assert!(matches!(Knn::new(0, None), Err(KnnSearchError::LZero)));
assert!(matches!(
Knn::new(100, Some(0)),
Err(KnnSearchError::BeamWidthZero)
));
}
}