use std::num::NonZeroUsize;
use diskann_utils::future::SendFuture;
use crate::{
ANNResult,
error::{ErrorExt, IntoANNResult},
flat::DistancesUnordered,
graph::{SearchOutputBuffer, glue::SearchPostProcess},
neighbor::{Neighbor, NeighborPriorityQueue},
};
#[derive(Debug, Clone, Copy, Default, PartialEq, Eq)]
pub struct SearchStats {
pub cmps: u32,
pub result_count: u32,
}
pub fn knn_search<V, T, O, PP, OB>(
visitor: &mut V,
k: NonZeroUsize,
processor: PP,
query: T,
output: &mut OB,
) -> impl SendFuture<ANNResult<SearchStats>>
where
V: DistancesUnordered,
T: Copy + Send + Sync,
O: Send,
PP: SearchPostProcess<V, T, O> + Send + Sync,
OB: SearchOutputBuffer<O> + Send + ?Sized,
{
async move {
let k = k.get();
let mut queue = NeighborPriorityQueue::new(k);
let mut cmps: u32 = 0;
visitor
.distances_unordered(|id, dist| {
cmps += 1;
queue.insert(Neighbor::new(id, dist));
})
.await
.escalate("flat scan must complete to produce correct k-NN results")?;
let result_count = processor
.post_process(visitor, query, queue.iter().take(k), output)
.await
.into_ann_result()? as u32;
Ok(SearchStats { cmps, result_count })
}
}
#[cfg(test)]
mod tests {
use crate::flat::test::{
harness::{CopyIdsOracle, EvenIdsOnlyOracle, KnnOracleRun, OracleProcessor},
provider::{self as flat_provider},
};
use crate::graph::test::synthetic::Grid;
fn fixture(grid: Grid, size: usize) -> (flat_provider::Provider, usize) {
let provider = flat_provider::Provider::grid(grid, size).unwrap();
let len = provider.len();
(provider, len)
}
#[tokio::test(flavor = "multi_thread", worker_threads = 4)]
async fn multithreaded_knn_search() {
use std::sync::Arc;
let (provider, len) = fixture(Grid::Two, 4);
let provider = Arc::new(provider);
let cases: &[(&[f32], usize)] = &[
(&[-1.0, -1.0], 1),
(&[1.0, 1.0], len),
(&[-1.0, 1.0], len / 2),
(&[1.0, -1.0], len - 1),
(&[0.0, 0.0], 3),
(&[3.0, 3.0], len),
(&[-2.0, 0.5], 2),
(&[0.5, -0.5], len),
];
fn spawn_cases<O>(
set: &mut tokio::task::JoinSet<(Vec<f32>, usize, KnnOracleRun)>,
provider: &Arc<flat_provider::Provider>,
oracle: O,
cases: &[(&[f32], usize)],
) where
O: OracleProcessor + Copy + Send + Sync + 'static,
{
for (query, k) in cases {
let provider = Arc::clone(provider);
let query: Vec<f32> = query.to_vec();
let k = *k;
set.spawn(async move {
let outcome = KnnOracleRun::run(&provider, &oracle, &query, k)
.await
.expect("knn_search failed");
(query, k, outcome)
});
}
}
let mut set = tokio::task::JoinSet::new();
spawn_cases(&mut set, &provider, CopyIdsOracle, cases);
spawn_cases(&mut set, &provider, EvenIdsOnlyOracle, cases);
while let Some(joined) = set.join_next().await {
let (query, k, outcome) = joined.expect("task panicked");
assert_eq!(
outcome.top_k, outcome.ground_truth,
"query = {query:?}, k = {k}: output must match its oracle",
);
assert_eq!(outcome.stats.cmps as usize, len);
assert_eq!(
outcome.stats.result_count as usize,
outcome.ground_truth.len(),
);
}
}
#[test]
fn transient_scan_error() {
for transient_ids in [&[0u32][..], &[3][..], &[1, 2, 5][..]] {
let (provider, _) = fixture(Grid::Two, 3);
let query = &[1.0, 0.0];
let visitor =
flat_provider::Visitor::flaky(&provider, query, transient_ids.iter().copied())
.unwrap();
let err =
KnnOracleRun::run_sync_with_visitor(&provider, visitor, &CopyIdsOracle, query, 4)
.expect_err("transient error during full scan must escalate");
let msg = format!("{err}");
assert!(
transient_ids
.iter()
.any(|id| msg.contains(&format!("id {id}"))),
"transients = {transient_ids:?}: expected error to name one of the \
transient ids, got: {msg}",
);
}
}
fn assert_visitor_error(query: &[f32], expected_msg: &str) {
let (provider, _) = fixture(Grid::Two, 3);
let err = flat_provider::Visitor::new(&provider, query)
.expect_err("expected visitor construction to fail");
let msg = format!("{err}");
assert!(
msg.contains(expected_msg),
"expected error containing {expected_msg:?}, got: {msg}",
);
}
#[test]
fn visitor_constructor_errors() {
assert_visitor_error(&[0.0, 0.0, 0.0], "dimension mismatch");
}
}