use std::num::NonZeroUsize;
use diskann_utils::future::SendFuture;
use thiserror::Error;
use crate::{
ANNResult, convert_error,
error::IntoANNResult,
graph::{
glue::{self, SearchAccessor, SearchStrategy},
index::{DiskANNIndex, InternalSearchStats, SearchStats},
search::{
Knn, Search, filtered_range_search::FilteredRange, record::NoopSearchRecord,
scratch::SearchScratch,
},
search_output_buffer::{self, SearchOutputBuffer},
},
neighbor::Neighbor,
provider::DataProvider,
};
#[derive(Debug, Error)]
pub enum RangeSearchError {
#[error("beam width cannot be zero")]
BeamWidthZero,
#[error("l_value cannot be zero")]
LZero,
#[error("initial_search_slack must be between 0 and 1.0")]
StartingListSlackValueError,
#[error("range_search_slack must be greater than or equal to 1.0")]
RangeSearchSlackValueError,
#[error("inner_radius must be less than or equal to radius")]
InnerRadiusValueError,
#[error("max_returned must be greater than or equal to starting_l")]
MaxReturnedLessThanInitialL,
}
convert_error!(RangeSearchError);
#[derive(Debug, Clone, Copy)]
pub struct Range {
max_returned: Option<usize>,
starting_l: NonZeroUsize,
beam_width: NonZeroUsize,
radius: f32,
inner_radius: Option<f32>,
initial_slack: f32,
range_slack: f32,
}
impl Range {
pub fn new(starting_l: usize, radius: f32) -> Result<Self, RangeSearchError> {
Self::builder(starting_l, radius).build()
}
pub fn builder(starting_l: usize, radius: f32) -> RangeBuilder {
RangeBuilder {
max_returned: None,
starting_l,
beam_width: None,
radius,
inner_radius: None,
initial_slack: 1.0,
range_slack: 1.0,
}
}
#[allow(clippy::too_many_arguments)]
fn validate_and_create(
max_returned: Option<usize>,
starting_l: usize,
beam_width: Option<usize>,
radius: f32,
inner_radius: Option<f32>,
initial_slack: f32,
range_slack: f32,
) -> Result<Self, RangeSearchError> {
let beam_width = match NonZeroUsize::new(beam_width.unwrap_or(1)) {
Some(bw) => bw,
None => return Err(RangeSearchError::BeamWidthZero),
};
let starting_l = match NonZeroUsize::new(starting_l) {
Some(l) => l,
None => return Err(RangeSearchError::LZero),
};
if let Some(max) = max_returned
&& max < starting_l.get()
{
return Err(RangeSearchError::MaxReturnedLessThanInitialL);
}
if !(0.0..=1.0).contains(&initial_slack) {
return Err(RangeSearchError::StartingListSlackValueError);
}
if range_slack < 1.0 {
return Err(RangeSearchError::RangeSearchSlackValueError);
}
if let Some(inner) = inner_radius
&& inner > radius
{
return Err(RangeSearchError::InnerRadiusValueError);
}
Ok(Self {
max_returned,
starting_l,
beam_width,
radius,
inner_radius,
initial_slack,
range_slack,
})
}
#[inline]
pub fn max_returned(&self) -> Option<usize> {
self.max_returned
}
#[inline]
pub fn starting_l(&self) -> NonZeroUsize {
self.starting_l
}
#[inline]
pub fn beam_width(&self) -> NonZeroUsize {
self.beam_width
}
#[inline]
pub fn radius(&self) -> f32 {
self.radius
}
#[inline]
pub fn inner_radius(&self) -> Option<f32> {
self.inner_radius
}
#[inline]
pub fn initial_slack(&self) -> f32 {
self.initial_slack
}
#[inline]
pub fn range_slack(&self) -> f32 {
self.range_slack
}
pub(super) fn to_knn(self) -> Knn {
Knn::new_infallible(self.starting_l, self.beam_width)
}
}
#[derive(Debug, Clone, Copy)]
pub struct RangeBuilder {
max_returned: Option<usize>,
starting_l: usize,
beam_width: Option<usize>,
radius: f32,
inner_radius: Option<f32>,
initial_slack: f32,
range_slack: f32,
}
impl RangeBuilder {
pub fn build_filtered(self) -> Result<FilteredRange, RangeSearchError> {
let range_params = self.build()?;
Ok(FilteredRange::from_range_params(range_params))
}
pub fn max_returned(mut self, value: Option<usize>) -> Self {
self.max_returned = value;
self
}
pub fn beam_width(mut self, value: Option<usize>) -> Self {
self.beam_width = value;
self
}
pub fn inner_radius(mut self, value: Option<f32>) -> Self {
self.inner_radius = value;
self
}
pub fn initial_slack(mut self, value: f32) -> Self {
self.initial_slack = value;
self
}
pub fn range_slack(mut self, value: f32) -> Self {
self.range_slack = value;
self
}
pub fn build(self) -> Result<Range, RangeSearchError> {
Range::validate_and_create(
self.max_returned,
self.starting_l,
self.beam_width,
self.radius,
self.inner_radius,
self.initial_slack,
self.range_slack,
)
}
}
impl<'a, DP, S, T> Search<'a, DP, S, T> for Range
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: glue::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.starting_l().get(), num_start_ids);
let initial_stats = index
.search_internal(
Some(self.beam_width().get()),
&mut accessor,
&mut scratch,
&mut NoopSearchRecord::new(),
)
.await?;
let mut in_range = Vec::with_capacity(self.starting_l().get());
let starting_l = self.starting_l().get();
let max_returned = self.max_returned().unwrap_or(usize::MAX);
for neighbor in scratch.best.iter().take(starting_l) {
if *neighbor.distance() <= self.radius() {
in_range.push(neighbor);
}
}
scratch.visited.clear();
for neighbor in in_range.iter() {
scratch.visited.insert(*neighbor.id());
}
scratch.in_range = in_range;
let stats = if scratch.in_range.len()
>= ((starting_l as f32) * self.initial_slack()) as usize
&& scratch.in_range.len() < max_returned
{
let range_stats = range_search_internal(
index.max_degree_with_slack(),
&self,
&mut accessor,
&mut scratch,
)
.await?;
InternalSearchStats {
cmps: initial_stats.cmps,
hops: initial_stats.hops + range_stats.hops,
range_search_second_round: true,
}
} else {
initial_stats
};
let radius = self.radius();
let inner_radius = self.inner_radius();
let mut filtered = DistanceFiltered::new(output, |dist| {
if let Some(ir) = inner_radius
&& dist <= ir
{
return false;
}
dist <= radius
});
let result_count = processor
.post_process(
&mut accessor,
query,
scratch.in_range.iter().copied(),
&mut filtered,
)
.await
.into_ann_result()?;
Ok(SearchStats {
cmps: stats.cmps,
hops: stats.hops,
result_count: result_count as u32,
range_search_second_round: stats.range_search_second_round,
})
}
}
}
pub(super) struct DistanceFiltered<'a, F, B: ?Sized> {
predicate: F,
inner: &'a mut B,
}
impl<'a, F, B: ?Sized> DistanceFiltered<'a, F, B> {
pub(super) fn new(inner: &'a mut B, predicate: F) -> Self {
Self { predicate, inner }
}
}
impl<I, F, B> SearchOutputBuffer<I> for DistanceFiltered<'_, F, B>
where
F: FnMut(f32) -> bool,
B: SearchOutputBuffer<I> + ?Sized,
{
fn size_hint(&self) -> Option<usize> {
self.inner.size_hint()
}
fn push(&mut self, neighbor: Neighbor<I>) -> search_output_buffer::BufferState {
if (self.predicate)(*neighbor.distance()) {
self.inner.push(neighbor)
} else {
match self.inner.size_hint() {
Some(0) => search_output_buffer::BufferState::Full,
_ => search_output_buffer::BufferState::Available,
}
}
}
fn current_len(&self) -> usize {
self.inner.current_len()
}
fn extend<Itr>(&mut self, itr: Itr) -> usize
where
Itr: IntoIterator<Item = Neighbor<I>>,
{
self.inner
.extend(itr.into_iter().filter(|n| (self.predicate)(*n.distance())))
}
}
pub(crate) async fn range_search_internal<A>(
max_degree_with_slack: usize,
search_params: &Range,
accessor: &mut A,
scratch: &mut SearchScratch<A::Id>,
) -> ANNResult<InternalSearchStats>
where
A: SearchAccessor,
{
let beam_width = search_params.beam_width().get();
for neighbor in &scratch.in_range {
scratch.range_frontier.push_back(*neighbor.id());
}
let mut neighbors = Vec::with_capacity(max_degree_with_slack);
let max_returned = search_params.max_returned().unwrap_or(usize::MAX);
while !scratch.range_frontier.is_empty() && scratch.in_range.len() < max_returned {
scratch.beam_nodes.clear();
while !scratch.range_frontier.is_empty() && scratch.beam_nodes.len() < beam_width {
let next = scratch.range_frontier.pop_front();
if let Some(next_node) = next {
scratch.beam_nodes.push(next_node);
}
}
neighbors.clear();
accessor
.expand_beam(
scratch.beam_nodes.iter().copied(),
glue::NotInMut::new(&mut scratch.visited),
|id, distance| neighbors.push(Neighbor::new(id, distance)),
)
.await?;
for neighbor in neighbors.iter() {
if *neighbor.distance() <= search_params.radius() * search_params.range_slack()
&& scratch.in_range.len() < max_returned
{
scratch.in_range.push(*neighbor);
scratch.range_frontier.push_back(*neighbor.id());
}
}
scratch.cmps += neighbors.len() as u32;
scratch.hops += scratch.beam_nodes.len() as u32;
}
Ok(InternalSearchStats {
cmps: scratch.cmps,
hops: scratch.hops,
range_search_second_round: true,
})
}
#[cfg(test)]
mod tests {
use super::*;
use crate::graph::search_output_buffer::BufferState;
use crate::neighbor::Neighbor;
#[test]
fn range_builder_defaults_match_new() {
let from_new = Range::new(100, 0.5).unwrap();
let from_builder = Range::builder(100, 0.5).build().unwrap();
assert_eq!(from_builder.max_returned(), from_new.max_returned());
assert_eq!(from_builder.starting_l(), from_new.starting_l());
assert_eq!(from_builder.beam_width(), from_new.beam_width());
assert_eq!(from_builder.radius(), from_new.radius());
assert_eq!(from_builder.inner_radius(), from_new.inner_radius());
assert_eq!(from_builder.initial_slack(), from_new.initial_slack());
assert_eq!(from_builder.range_slack(), from_new.range_slack());
}
#[test]
fn range_builder_custom_options_match_expected_values() {
let built = Range::builder(100, 0.8)
.max_returned(Some(101))
.beam_width(Some(8))
.inner_radius(Some(0.3))
.initial_slack(0.9)
.range_slack(1.2)
.build()
.unwrap();
assert_eq!(built.max_returned(), Some(101));
assert_eq!(built.starting_l().get(), 100);
assert_eq!(built.beam_width().get(), 8);
assert_eq!(built.radius(), 0.8);
assert_eq!(built.inner_radius(), Some(0.3));
assert_eq!(built.initial_slack(), 0.9);
assert_eq!(built.range_slack(), 1.2);
}
#[test]
fn range_builder_validation_error() {
let err = Range::builder(100, 0.5)
.beam_width(Some(0))
.build()
.unwrap_err();
assert!(matches!(err, RangeSearchError::BeamWidthZero));
}
#[test]
fn test_range_search_validation() {
assert!(Range::new(100, 0.5).is_ok());
assert!(Range::new(0, 0.5).is_err());
assert!(Range::builder(100, 0.5).initial_slack(1.5).build().is_err());
assert!(Range::builder(100, 0.5).range_slack(0.5).build().is_err());
assert!(
Range::builder(100, 0.5)
.inner_radius(Some(1.0))
.build()
.is_err()
);
assert!(
Range::builder(100, 0.5)
.max_returned(Some(1))
.build()
.is_err()
);
}
#[test]
fn distance_filtered_push_accepts_passing_items() {
let mut inner: Vec<Neighbor<u32>> = Vec::new();
let mut filtered = DistanceFiltered::new(&mut inner, |d| d < 1.0);
assert_eq!(filtered.push(Neighbor::new(1, 0.5)), BufferState::Available);
assert_eq!(filtered.current_len(), 1);
assert_eq!(*inner[0].id(), 1);
assert_eq!(*inner[0].distance(), 0.5);
}
#[test]
fn distance_filtered_push_rejects_failing_items() {
let mut inner: Vec<Neighbor<u32>> = Vec::new();
let mut filtered = DistanceFiltered::new(&mut inner, |d| d < 1.0);
assert_eq!(filtered.push(Neighbor::new(1, 1.5)), BufferState::Available);
assert_eq!(filtered.current_len(), 0);
}
#[test]
fn distance_filtered_extend_filters_correctly() {
let mut inner: Vec<Neighbor<u32>> = Vec::new();
let mut filtered = DistanceFiltered::new(&mut inner, |d| d < 1.0);
assert!(filtered.size_hint().is_none());
let items = [(1u32, 0.3), (2, 1.5), (3, 0.7), (4, 2.0), (5, 0.9)].map(Neighbor::from_tuple);
let count = filtered.extend(items);
assert_eq!(count, 3);
assert_eq!(inner.len(), 3);
assert_eq!(*inner[0].id(), 1);
assert_eq!(*inner[1].id(), 3);
assert_eq!(*inner[2].id(), 5);
}
#[test]
fn distance_filtered_respects_inner_capacity() {
let mut ids = [0u32; 2];
let mut dists = [0.0f32; 2];
let mut inner = search_output_buffer::IdDistance::new(&mut ids, &mut dists);
let mut filtered = DistanceFiltered::new(&mut inner, |d| d < 1.0);
assert_eq!(filtered.size_hint(), Some(2));
let items = [(1u32, 0.1), (2, 0.2), (3, 0.3)].map(Neighbor::from_tuple);
let count = filtered.extend(items);
assert_eq!(count, 2);
assert_eq!(ids, [1, 2]);
}
#[test]
fn distance_filtered_inner_radius_pattern() {
let mut inner: Vec<Neighbor<u32>> = Vec::new();
let radius = 1.0f32;
let inner_radius = Some(0.3f32);
let mut filtered = DistanceFiltered::new(&mut inner, |dist| {
if let Some(ir) = inner_radius
&& dist <= ir
{
return false;
}
dist < radius
});
let items = [(1u32, 0.1), (2, 0.5), (3, 0.3), (4, 1.0), (5, 0.8)].map(Neighbor::from_tuple);
let count = filtered.extend(items);
assert_eq!(count, 2);
assert_eq!(*inner[0].id(), 2);
assert_eq!(*inner[1].id(), 5);
}
}