use std::collections::VecDeque;
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::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 finite and between 0 and 1.0")]
StartingListSlackValueError,
#[error("range_search_slack must be finite and greater than or equal to 1.0")]
RangeSearchSlackValueError,
#[error("radius must be finite")]
RadiusValueError,
#[error("inner_radius must be finite and 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,
}
}
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 !radius.is_finite() {
return Err(RangeSearchError::RadiusValueError);
}
if !initial_slack.is_finite() || !(0.0..=1.0).contains(&initial_slack) {
return Err(RangeSearchError::StartingListSlackValueError);
}
if !range_slack.is_finite() || range_slack < 1.0 {
return Err(RangeSearchError::RangeSearchSlackValueError);
}
if let Some(inner) = inner_radius
&& (!inner.is_finite() || 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 effective_max_returned(&self, inc: usize) -> usize {
if let Some(max) = self.max_returned {
max.saturating_add(inc)
} else {
usize::MAX
}
}
#[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 starting_l = self.starting_l().get();
let mut scratch = index.search_scratch(starting_l, num_start_ids);
let initial_stats = index
.search_internal(
Some(self.beam_width().get()),
&mut accessor,
&mut scratch,
&mut NoopSearchRecord::new(),
)
.await?;
let in_outer_range = InRange::new(self.radius(), None, starting_l, scratch.best.iter());
let max_returned = self.effective_max_returned(num_start_ids);
let mut in_range = InRange::new(
self.radius(),
self.inner_radius(),
max_returned,
in_outer_range.iter(),
);
let stats = if in_outer_range.len()
>= ((starting_l as f32) * self.initial_slack()) as usize
&& in_outer_range.len() <= max_returned
{
scratch.visited.clear();
scratch
.visited
.extend(in_outer_range.iter().map(|n| *n.id()));
let mut range_frontier: VecDeque<_> =
in_outer_range.take().into_iter().map(|n| *n.id()).collect();
let range_stats = range_search_internal(
index.max_degree_with_slack(),
&self,
&mut accessor,
&mut scratch,
&mut range_frontier,
&mut in_range,
)
.await?;
InternalSearchStats {
cmps: range_stats.cmps,
hops: range_stats.hops,
range_search_second_round: true,
}
} else {
initial_stats
};
let result_count = processor
.post_process(&mut accessor, query, in_range.iter(), output)
.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 InRange<I> {
neighbors: Vec<Neighbor<I>>,
radius: f32,
inner_radius: Option<f32>,
max_returned: usize,
}
impl<I> InRange<I> {
#[must_use]
pub(super) fn push(&mut self, neighbor: Neighbor<I>) -> bool {
let d = *neighbor.distance();
if self.neighbors.len() < self.max_returned && self.check(d) {
self.neighbors.push(neighbor);
true
} else {
false
}
}
pub(super) fn sort_and_dedup(&mut self)
where
I: Ord,
{
self.neighbors
.sort_unstable_by(crate::neighbor::ord::fast_distance_total);
self.neighbors
.dedup_by(|left, right| left.id() == right.id());
}
pub(super) fn len(&self) -> usize {
self.neighbors.len()
}
pub(super) fn is_full(&self) -> bool {
self.len() == self.max_returned
}
pub(super) fn take(self) -> Vec<Neighbor<I>> {
self.neighbors
}
pub(super) fn iter(&self) -> impl ExactSizeIterator<Item = Neighbor<I>>
where
I: Copy,
{
self.neighbors.iter().copied()
}
#[must_use]
pub(super) fn check(&self, distance: f32) -> bool {
distance <= self.radius && self.inner_radius.is_none_or(|inner| distance > inner)
}
pub(super) fn new<Itr>(
radius: f32,
inner_radius: Option<f32>,
max_returned: usize,
candidates: Itr,
) -> Self
where
Itr: IntoIterator<Item = Neighbor<I>>,
{
Self {
neighbors: candidates
.into_iter()
.filter(|n| {
let dist = *n.distance();
if dist > radius {
return false;
}
if let Some(inner) = inner_radius
&& dist <= inner
{
return false;
}
true
})
.take(max_returned)
.collect(),
radius,
inner_radius,
max_returned,
}
}
}
pub(crate) async fn range_search_internal<A>(
max_degree_with_slack: usize,
search_params: &Range,
accessor: &mut A,
scratch: &mut SearchScratch<A::Id>,
range_frontier: &mut VecDeque<A::Id>,
in_range: &mut InRange<A::Id>,
) -> ANNResult<InternalSearchStats>
where
A: SearchAccessor,
{
let beam_width = search_params.beam_width().get();
let mut neighbors = Vec::with_capacity(max_degree_with_slack);
while !range_frontier.is_empty() && !in_range.is_full() {
scratch.beam_nodes.clear();
while !range_frontier.is_empty() && scratch.beam_nodes.len() < beam_width {
let next = 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?;
let navigation_radius = search_params.radius() * search_params.range_slack();
for neighbor in neighbors.iter() {
if in_range.is_full() {
break;
}
if in_range.push(*neighbor) || *neighbor.distance() <= navigation_radius {
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::*;
fn neighbor(id: u32, distance: f32) -> Neighbor<u32> {
Neighbor::new(id, distance)
}
#[test]
fn in_range_applies_inner_and_outer_radius_boundaries() {
let candidates = [
neighbor(0, 0.1),
neighbor(1, 0.2),
neighbor(2, 0.3),
neighbor(3, 0.5),
neighbor(4, 0.6),
];
let in_range = InRange::new(0.5, Some(0.2), usize::MAX, candidates);
let ids: Vec<_> = in_range.iter().map(|candidate| *candidate.id()).collect();
assert_eq!(ids, [2, 3]);
assert!(!in_range.check(0.2));
assert!(in_range.check(0.5));
}
#[test]
fn in_range_filters_before_truncating_and_preserves_order() {
let candidates = [
neighbor(0, 0.8),
neighbor(1, 0.3),
neighbor(2, 0.7),
neighbor(3, 0.2),
neighbor(4, 0.1),
];
let in_range = InRange::new(0.5, None, 2, candidates);
let ids: Vec<_> = in_range.iter().map(|candidate| *candidate.id()).collect();
assert_eq!(ids, [1, 3]);
assert_eq!(in_range.len(), 2);
assert!(in_range.is_full());
}
#[test]
fn in_range_push_rejects_out_of_range_and_over_capacity() {
let mut in_range = InRange::new(0.5, Some(0.1), 2, []);
assert!(!in_range.push(neighbor(0, 0.1)));
assert!(!in_range.push(neighbor(1, 0.6)));
assert!(in_range.push(neighbor(2, 0.2)));
assert!(!in_range.is_full());
assert!(in_range.push(neighbor(3, 0.5)));
assert!(in_range.is_full());
assert!(!in_range.push(neighbor(4, 0.3)));
}
#[test]
fn in_range_sort_and_dedup_sorts_by_distance_and_removes_repeated_ids() {
let mut in_range = InRange::new(
1.0,
None,
usize::MAX,
[
neighbor(2, 0.4),
neighbor(1, 0.2),
neighbor(1, 0.2),
neighbor(3, 0.3),
],
);
in_range.sort_and_dedup();
let neighbors = in_range.take();
let ids: Vec<_> = neighbors.iter().map(|candidate| *candidate.id()).collect();
let distances: Vec<_> = neighbors
.iter()
.map(|candidate| *candidate.distance())
.collect();
assert_eq!(ids, [1, 3, 2]);
assert_eq!(distances, [0.2, 0.3, 0.4]);
}
#[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 range_builder_rejects_non_finite_float_values() {
for value in [f32::NAN, f32::INFINITY, f32::NEG_INFINITY] {
assert!(matches!(
Range::builder(100, value).build(),
Err(RangeSearchError::RadiusValueError)
));
assert!(matches!(
Range::builder(100, 0.5).inner_radius(Some(value)).build(),
Err(RangeSearchError::InnerRadiusValueError)
));
assert!(matches!(
Range::builder(100, 0.5).initial_slack(value).build(),
Err(RangeSearchError::StartingListSlackValueError)
));
assert!(matches!(
Range::builder(100, 0.5).range_slack(value).build(),
Err(RangeSearchError::RangeSearchSlackValueError)
));
}
}
#[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()
);
}
}