use std::sync::Arc;
use diskann_vector::distance::Metric;
use crate::{
graph::{
self, DiskANNIndex,
index::SearchStats,
search::Range,
test::{provider as test_provider, synthetic::Grid},
},
neighbor::Neighbor,
test::{
TestRoot,
cmp::{assert_eq_verbose, verbose_eq},
get_or_save_test_results,
tokio::current_thread_runtime,
},
};
#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
pub(super) struct RangeSearchBaseline {
pub(super) description: String,
pub(super) grid_dims: u8,
pub(super) grid_size: usize,
pub(super) query: Vec<f32>,
pub(super) radius: f32,
pub(super) inner_radius: Option<f32>,
pub(super) starting_l: usize,
pub(super) results: Vec<(u32, f32)>,
pub(super) comparisons: usize,
pub(super) hops: usize,
pub(super) result_count: usize,
pub(super) range_search_second_round: bool,
}
impl RangeSearchBaseline {
pub(super) fn new(
range: &Range,
results: &[Neighbor<u32>],
stats: SearchStats,
grid_dims: Grid,
grid_size: usize,
description: impl Into<String>,
query: Vec<f32>,
) -> Self {
Self {
description: description.into(),
grid_dims: grid_dims.dim(),
grid_size,
query,
radius: range.radius(),
inner_radius: range.inner_radius(),
starting_l: range.starting_l().get(),
results: results.iter().map(|n| (*n.id(), *n.distance())).collect(),
comparisons: stats.cmps as usize,
hops: stats.hops as usize,
result_count: stats.result_count as usize,
range_search_second_round: stats.range_search_second_round,
}
}
}
verbose_eq!(RangeSearchBaseline {
description,
grid_dims,
grid_size,
query,
radius,
inner_radius,
starting_l,
results,
comparisons,
hops,
result_count,
range_search_second_round,
});
fn root() -> TestRoot {
TestRoot::new("graph/test/cases/range_search")
}
pub(super) fn setup_grid_index(
grid_size: usize,
dims: Grid,
) -> Arc<DiskANNIndex<test_provider::Provider>> {
let provider = test_provider::Provider::grid(dims, grid_size).unwrap();
let index_config = graph::config::Builder::new(
provider.max_degree(),
graph::config::MaxDegree::same(),
100,
Metric::L2.into(),
)
.build()
.unwrap();
Arc::new(DiskANNIndex::new(index_config, provider, None))
}
pub(super) fn setup_grid_index_and_default_query(
grid_size: usize,
dims: Grid,
) -> (Arc<DiskANNIndex<test_provider::Provider>>, Vec<f32>) {
let index = setup_grid_index(grid_size, dims);
let query = vec![grid_size as f32; dims.dim().into()];
(index, query)
}
pub(super) fn assert_no_duplicates(results: &[Neighbor<u32>]) {
let mut seen = std::collections::HashSet::new();
for n in results {
assert!(seen.insert(*n.id()), "duplicate result id {}", n.id());
}
}
pub(super) fn assert_range_invariants(
results: &[Neighbor<u32>],
radius: f32,
inner_radius: Option<f32>,
) {
for n in results {
assert!(
*n.distance() <= radius,
"result {} distance {} exceeds radius {}",
n.id(),
n.distance(),
radius
);
if let Some(inner) = inner_radius {
assert!(
*n.distance() > inner,
"result {} distance {} is within inner radius {}",
n.id(),
n.distance(),
inner
);
}
}
}
#[test]
fn basic_range_search() {
let description = "Basic range search test to validate that the range \
search returns results within the specified radius and that there are \
no duplicate results.";
let rt = current_thread_runtime();
let mut test_root = root();
let mut path = test_root.path();
let name = path.push("basic_range_search");
let grid_size = 5;
let (index, query) = setup_grid_index_and_default_query(grid_size, Grid::Three);
let radius = 12.0;
let starting_l = 32;
let range_search = Range::new(starting_l, radius).unwrap();
let mut results: Vec<Neighbor<u32>> = Vec::new();
let stats = rt
.block_on(index.search(
range_search,
&test_provider::Strategy::new(),
&test_provider::Context::new(),
query.as_slice(),
&mut results,
))
.unwrap();
let baseline = RangeSearchBaseline {
description: description.to_string(),
grid_dims: Grid::Three.dim(),
grid_size,
query: query.clone(),
radius,
inner_radius: None,
starting_l,
results: results.iter().map(|n| n.as_tuple()).collect(),
comparisons: stats.cmps as usize,
hops: stats.hops as usize,
result_count: results.len(),
range_search_second_round: stats.range_search_second_round,
};
let expected = get_or_save_test_results(&name, &baseline);
assert_eq_verbose!(expected, baseline);
assert_range_invariants(&results, radius, None);
assert_no_duplicates(&results);
}
#[test]
fn inner_radius_filtering() {
let description = "Inner radius filtering test to validate that the \
range search correctly excludes neighbors within the inner radius.";
let rt = current_thread_runtime();
let mut test_root = root();
let mut path = test_root.path();
let name = path.push("inner_radius_filtering");
let grid_size = 5;
let (index, query) = setup_grid_index_and_default_query(grid_size, Grid::Three);
let radius = 20.0;
let inner_radius = 6.0; let starting_l = 32;
let range_search = Range::builder(starting_l, radius)
.inner_radius(Some(inner_radius))
.build()
.unwrap();
let mut results: Vec<Neighbor<u32>> = Vec::new();
let stats = rt
.block_on(index.search(
range_search,
&test_provider::Strategy::new(),
&test_provider::Context::new(),
query.as_slice(),
&mut results,
))
.unwrap();
let baseline = RangeSearchBaseline {
description: description.to_string(),
grid_dims: Grid::Three.dim(),
grid_size,
query: query.clone(),
radius,
inner_radius: Some(inner_radius),
starting_l,
results: results.iter().map(|n| n.as_tuple()).collect(),
comparisons: stats.cmps as usize,
hops: stats.hops as usize,
result_count: results.len(),
range_search_second_round: stats.range_search_second_round,
};
let expected = get_or_save_test_results(&name, &baseline);
assert_eq_verbose!(expected, baseline);
assert_range_invariants(&results, radius, Some(inner_radius));
assert_no_duplicates(&results);
}
#[test]
fn two_round_search() {
let description = "Two round search test to validate that a \
low starting L with a large radius triggers a second round \
of range search.";
let rt = current_thread_runtime();
let mut test_root = root();
let mut path = test_root.path();
let name = path.push("two_round_search");
let grid_size = 5;
let (index, query) = setup_grid_index_and_default_query(grid_size, Grid::Three);
let radius = 50.0; let starting_l = 4;
let range_search = Range::new(starting_l, radius).unwrap();
let mut results: Vec<Neighbor<u32>> = Vec::new();
let stats = rt
.block_on(index.search(
range_search,
&test_provider::Strategy::new(),
&test_provider::Context::new(),
query.as_slice(),
&mut results,
))
.unwrap();
let baseline = RangeSearchBaseline {
description: description.to_string(),
grid_dims: Grid::Three.dim(),
grid_size,
query: query.clone(),
radius,
inner_radius: None,
starting_l,
results: results.iter().map(|n| n.as_tuple()).collect(),
comparisons: stats.cmps as usize,
hops: stats.hops as usize,
result_count: results.len(),
range_search_second_round: stats.range_search_second_round,
};
let expected = get_or_save_test_results(&name, &baseline);
assert_eq_verbose!(expected, baseline);
assert!(
stats.range_search_second_round,
"low starting_l with large radius should trigger a second round"
);
assert_range_invariants(&results, radius, None);
assert_no_duplicates(&results);
}
#[test]
fn empty_results() {
let rt = current_thread_runtime();
let grid_size = 5;
let (index, query) = setup_grid_index_and_default_query(grid_size, Grid::Three);
let radius = 0.01; let starting_l = 32;
let range_search = Range::new(starting_l, radius).unwrap();
let mut results: Vec<Neighbor<u32>> = Vec::new();
let stats = rt
.block_on(index.search(
range_search,
&test_provider::Strategy::new(),
&test_provider::Context::new(),
query.as_slice(),
&mut results,
))
.unwrap();
assert!(
results.is_empty(),
"no points should be within the radius {}",
radius
);
assert!(
!stats.range_search_second_round,
"empty results shouldn't trigger a second round"
);
}
#[test]
fn max_results_respected_means_no_second_round() {
let description = "Two round search test to validate that max_results = \
starting_l means no second round is triggered.";
let rt = current_thread_runtime();
let mut test_root = root();
let mut path = test_root.path();
let name = path.push("max_results_respected_means_no_second_round");
let grid_size = 5;
let (index, query) = setup_grid_index_and_default_query(grid_size, Grid::Three);
let radius = 1.0e9; let starting_l = 4; let max_results = 4;
let range_search = Range::builder(starting_l, radius)
.max_returned(Some(max_results))
.build()
.unwrap();
let mut results: Vec<Neighbor<u32>> = Vec::new();
let stats = rt
.block_on(index.search(
range_search,
&test_provider::Strategy::new(),
&test_provider::Context::new(),
query.as_slice(),
&mut results,
))
.unwrap();
let baseline = RangeSearchBaseline {
description: description.to_string(),
grid_dims: Grid::Three.dim(),
grid_size,
query: query.clone(),
radius,
inner_radius: None,
starting_l,
results: results.iter().map(|n| n.as_tuple()).collect(),
comparisons: stats.cmps as usize,
hops: stats.hops as usize,
result_count: results.len(),
range_search_second_round: stats.range_search_second_round,
};
let expected = get_or_save_test_results(&name, &baseline);
assert_eq_verbose!(expected, baseline);
assert!(
results.len() <= max_results,
"result count {} exceeds max_results {}",
results.len(),
max_results
);
assert!(
!stats.range_search_second_round,
"If max_results is respected, a second round should not be triggered"
);
assert_range_invariants(&results, radius, None);
assert_no_duplicates(&results);
}
#[test]
fn max_results_respected_and_second_round_triggered() {
let description = "Two round search test to validate that max_results > \
starting_l means a second round is triggered.";
let rt = current_thread_runtime();
let mut test_root = root();
let mut path = test_root.path();
let name = path.push("max_results_respected_and_second_round_triggered");
let grid_size = 5;
let (index, query) = setup_grid_index_and_default_query(grid_size, Grid::Three);
let radius = 1.0e9; let starting_l = 4; let max_results = 5;
let range_search = Range::builder(starting_l, radius)
.max_returned(Some(max_results))
.build()
.unwrap();
let mut results: Vec<Neighbor<u32>> = Vec::new();
let stats = rt
.block_on(index.search(
range_search,
&test_provider::Strategy::new(),
&test_provider::Context::new(),
query.as_slice(),
&mut results,
))
.unwrap();
let baseline = RangeSearchBaseline {
description: description.to_string(),
grid_dims: Grid::Three.dim(),
grid_size,
query: query.clone(),
radius,
inner_radius: None,
starting_l,
results: results.iter().map(|n| n.as_tuple()).collect(),
comparisons: stats.cmps as usize,
hops: stats.hops as usize,
result_count: results.len(),
range_search_second_round: stats.range_search_second_round,
};
let expected = get_or_save_test_results(&name, &baseline);
assert_eq_verbose!(expected, baseline);
assert!(
results.len() <= max_results,
"result count {} exceeds max_results {}",
results.len(),
max_results
);
assert!(
stats.range_search_second_round,
"If max_results is respected, a second round should be triggered"
);
assert_range_invariants(&results, radius, None);
assert_no_duplicates(&results);
}