use axiolid_core::{Point3, Scalar};
#[derive(Debug, Clone, Copy, PartialEq)]
pub struct PointHit {
pub index: usize,
pub distance: Scalar,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
#[non_exhaustive]
pub enum PointQueryError {
NonFiniteQuery,
InvalidRadius,
}
impl core::fmt::Display for PointQueryError {
fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
match self {
Self::NonFiniteQuery => f.write_str("query position is not finite"),
Self::InvalidRadius => f.write_str("radius must be finite and non-negative"),
}
}
}
impl core::error::Error for PointQueryError {}
#[derive(Debug, Clone)]
pub struct PointIndex {
points: Vec<Point3>,
cell: Scalar,
origin: Point3,
dims: [usize; 3],
starts: Vec<u32>,
ordered: Vec<u32>,
}
impl PointIndex {
pub fn build(points: &[Point3]) -> Self {
let finite: Vec<u32> = points
.iter()
.enumerate()
.filter(|(_, p)| p.is_finite())
.map(|(index, _)| index as u32)
.collect();
if finite.is_empty() {
return Self {
points: points.to_vec(),
cell: 1.0,
origin: Point3::ZERO,
dims: [1, 1, 1],
starts: vec![0, 0],
ordered: Vec::new(),
};
}
let mut min = points[finite[0] as usize];
let mut max = min;
for &i in &finite {
let p = points[i as usize];
min = Point3::new(min.x.min(p.x), min.y.min(p.y), min.z.min(p.z));
max = Point3::new(max.x.max(p.x), max.y.max(p.y), max.z.max(p.z));
}
let span = max - min;
let target = (finite.len() as Scalar).cbrt().max(1.0);
let longest = span.x.max(span.y).max(span.z);
let cell = if longest > 0.0 {
(longest / target).max(Scalar::MIN_POSITIVE)
} else {
1.0
};
let dims = [
((span.x / cell).ceil() as usize + 1).max(1),
((span.y / cell).ceil() as usize + 1).max(1),
((span.z / cell).ceil() as usize + 1).max(1),
];
let cell_count = dims[0] * dims[1] * dims[2];
let mut counts = vec![0u32; cell_count + 1];
let locate = |p: Point3| -> usize {
let ix = (((p.x - min.x) / cell) as usize).min(dims[0] - 1);
let iy = (((p.y - min.y) / cell) as usize).min(dims[1] - 1);
let iz = (((p.z - min.z) / cell) as usize).min(dims[2] - 1);
(iz * dims[1] + iy) * dims[0] + ix
};
for &i in &finite {
counts[locate(points[i as usize]) + 1] += 1;
}
for k in 1..counts.len() {
counts[k] += counts[k - 1];
}
let starts = counts.clone();
let mut cursor = counts;
let mut ordered = vec![0u32; finite.len()];
for &i in &finite {
let cell_index = locate(points[i as usize]);
ordered[cursor[cell_index] as usize] = i;
cursor[cell_index] += 1;
}
Self {
points: points.to_vec(),
cell,
origin: min,
dims,
starts,
ordered,
}
}
pub fn len(&self) -> usize {
self.points.len()
}
pub fn is_empty(&self) -> bool {
self.points.is_empty()
}
pub fn rejected(&self) -> usize {
self.points.len() - self.ordered.len()
}
pub fn for_each_within(
&self,
query: Point3,
radius: Scalar,
mut visit: impl FnMut(PointHit),
) -> Result<(), PointQueryError> {
if !query.is_finite() {
return Err(PointQueryError::NonFiniteQuery);
}
if !radius.is_finite() || radius < 0.0 {
return Err(PointQueryError::InvalidRadius);
}
if self.ordered.is_empty() {
return Ok(());
}
let radius_squared = radius * radius;
let lo = self.cell_of(query - Point3::splat(radius));
let hi = self.cell_of(query + Point3::splat(radius));
for iz in lo[2]..=hi[2] {
for iy in lo[1]..=hi[1] {
for ix in lo[0]..=hi[0] {
let cell_index = (iz * self.dims[1] + iy) * self.dims[0] + ix;
let from = self.starts[cell_index] as usize;
let to = self.starts[cell_index + 1] as usize;
for &point_index in &self.ordered[from..to] {
let point = self.points[point_index as usize];
let squared = (point - query).length_squared();
if squared <= radius_squared {
visit(PointHit {
index: point_index as usize,
distance: squared.sqrt(),
});
}
}
}
}
}
Ok(())
}
pub fn radius_into(
&self,
query: Point3,
radius: Scalar,
out: &mut Vec<PointHit>,
) -> Result<(), PointQueryError> {
out.clear();
self.for_each_within(query, radius, |hit| out.push(hit))?;
sort_hits(out);
Ok(())
}
pub fn nearest_into(
&self,
query: Point3,
k: usize,
out: &mut Vec<PointHit>,
) -> Result<(), PointQueryError> {
out.clear();
if !query.is_finite() {
return Err(PointQueryError::NonFiniteQuery);
}
if k == 0 || self.ordered.is_empty() {
return Ok(());
}
let mut radius = self.cell * (k as Scalar).cbrt().max(1.0);
let ceiling = self.diagonal();
loop {
self.radius_into(query, radius.min(ceiling), out)?;
if out.len() >= k || radius >= ceiling {
break;
}
radius *= 2.0;
}
out.truncate(k);
Ok(())
}
pub fn nearest(&self, query: Point3) -> Result<Option<PointHit>, PointQueryError> {
let mut out = Vec::new();
self.nearest_into(query, 1, &mut out)?;
Ok(out.into_iter().next())
}
fn cell_of(&self, p: Point3) -> [usize; 3] {
let axis = |value: Scalar, origin: Scalar, dim: usize| -> usize {
let raw = (value - origin) / self.cell;
if raw < 0.0 {
0
} else {
(raw as usize).min(dim - 1)
}
};
[
axis(p.x, self.origin.x, self.dims[0]),
axis(p.y, self.origin.y, self.dims[1]),
axis(p.z, self.origin.z, self.dims[2]),
]
}
fn diagonal(&self) -> Scalar {
let span = Point3::new(
self.dims[0] as Scalar * self.cell,
self.dims[1] as Scalar * self.cell,
self.dims[2] as Scalar * self.cell,
);
(span.x * span.x + span.y * span.y + span.z * span.z).sqrt()
}
}
fn sort_hits(hits: &mut [PointHit]) {
hits.sort_by(|a, b| {
a.distance
.partial_cmp(&b.distance)
.unwrap_or(core::cmp::Ordering::Equal)
.then(a.index.cmp(&b.index))
});
}