use distances::Number;
use crate::{Cluster, Dataset, Tree};
use super::linear;
pub fn search<T, U, D>(tree: &Tree<T, U, D>, query: T, radius: U) -> Vec<(usize, U)>
where
T: Send + Sync + Copy,
U: Number,
D: Dataset<T, U>,
{
let [confirmed, straddlers] = tree_search(tree.data(), &tree.root, query, radius);
leaf_search(tree.data(), confirmed, straddlers, query, radius)
}
pub fn tree_search<'a, T, U, D>(
data: &D,
root: &'a Cluster<T, U>,
query: T,
radius: U,
) -> [Vec<(&'a Cluster<T, U>, U)>; 2]
where
T: Send + Sync + Copy,
U: Number,
D: Dataset<T, U>,
{
let mut confirmed = Vec::new();
let mut straddlers = Vec::new();
let mut candidates = vec![root];
let (mut terminal, mut non_terminal): (Vec<_>, Vec<_>);
while !candidates.is_empty() {
(terminal, non_terminal) = candidates
.into_iter()
.map(|c| (c, c.distance_to_instance(data, query)))
.filter(|&(c, d)| d <= (c.radius + radius))
.partition(|&(c, d)| (c.radius + d) <= radius);
confirmed.append(&mut terminal);
(terminal, non_terminal) = non_terminal.into_iter().partition(|&(c, _)| c.is_leaf());
straddlers.append(&mut terminal);
candidates = non_terminal
.into_iter()
.flat_map(|(c, d)| {
if d < c.radius {
c.overlapping_children(data, query, radius)
} else {
c.children()
.map_or_else(|| unreachable!("Non-leaf cluster without children"), |v| v.to_vec())
}
})
.collect();
}
[confirmed, straddlers]
}
pub fn leaf_search<T, U, D>(
data: &D,
confirmed: Vec<(&Cluster<T, U>, U)>,
straddlers: Vec<(&Cluster<T, U>, U)>,
query: T,
radius: U,
) -> Vec<(usize, U)>
where
T: Send + Sync + Copy,
U: Number,
D: Dataset<T, U>,
{
let hits = confirmed.into_iter().flat_map(|(c, d)| {
let distances = if c.is_singleton() {
vec![d; c.cardinality]
} else {
data.query_to_many(query, c.indices(data))
};
c.indices(data).iter().copied().zip(distances)
});
let indices = straddlers
.into_iter()
.flat_map(|(c, _)| c.indices(data))
.copied()
.collect::<Vec<_>>();
hits.chain(linear::search(data, query, radius, &indices)).collect()
}