use core::f64::consts::PI;
use std::collections::{HashMap, HashSet};
type Key = (i32, i32, i32);
#[derive(Debug)]
pub(crate) struct Region<'a, T> {
pub votes: usize,
pub members: &'a [T],
}
pub(crate) struct SkyVotes<T> {
step: f64,
log_step: f64,
buckets: HashMap<Key, Vec<T>>,
}
impl<T> SkyVotes<T> {
pub(crate) fn new(step: f64, log_step: f64) -> Self {
Self {
step,
log_step,
buckets: HashMap::new(),
}
}
pub(crate) fn len(&self) -> usize {
self.buckets.len()
}
fn ra_bins(&self, band: i32) -> i32 {
let dec_mid = (f64::from(band) + 0.5) * self.step - PI / 2.0;
let width = self.step / dec_mid.cos().max(1e-6);
((2.0 * PI / width).floor() as i32).max(1)
}
fn key(&self, ra: f64, dec: f64, scale: f64) -> Key {
let n_bands = (PI / self.step).ceil() as i32;
let band = (((dec + PI / 2.0) / self.step).floor() as i32).clamp(0, n_bands - 1);
let n_ra = self.ra_bins(band);
let ra = ra.rem_euclid(2.0 * PI);
let ra_bin = ((ra / (2.0 * PI) * f64::from(n_ra)).floor() as i32).rem_euclid(n_ra);
let s_bin = (scale.max(1e-12).ln() / self.log_step).floor() as i32;
(band, ra_bin, s_bin)
}
fn centre(&self, key: Key) -> (f64, f64) {
let dec = (f64::from(key.0) + 0.5) * self.step - PI / 2.0;
let n_ra = self.ra_bins(key.0);
let ra = (f64::from(key.1) + 0.5) * 2.0 * PI / f64::from(n_ra);
(ra, dec)
}
fn neighbours(&self, key: Key, reach: i32, s_reach: i32) -> Vec<Key> {
let (ra, dec) = self.centre(key);
let mut out = Vec::with_capacity(((2 * reach + 1).pow(2) * (2 * s_reach + 1)) as usize);
for dd in -reach..=reach {
let d = dec + f64::from(dd) * self.step;
if !(-PI / 2.0..=PI / 2.0).contains(&d) {
continue;
}
let cos_d = d.cos().max(1e-6);
for dr in -reach..=reach {
let r = ra + f64::from(dr) * self.step / cos_d;
let (band, ra_bin, _) = self.key(r, d, 1.0);
for ds in -s_reach..=s_reach {
let k = (band, ra_bin, key.2 + ds);
if !out.contains(&k) {
out.push(k);
}
}
}
}
out
}
pub(crate) fn add(&mut self, ra: f64, dec: f64, scale: f64, item: T) {
let key = self.key(ra, dec, scale);
self.buckets.entry(key).or_default().push(item);
}
pub(crate) fn regions(&self, limit: usize) -> Vec<Region<'_, T>> {
let mut smoothed: Vec<(usize, Key, Key)> = self
.buckets
.keys()
.map(|&key| {
let mut sum = 0usize;
let mut best: Option<(usize, Key)> = None;
for n in self.neighbours(key, 1, 1) {
if let Some(v) = self.buckets.get(&n) {
sum += v.len();
let better =
best.is_none_or(|(bn, bk)| v.len() > bn || (v.len() == bn && n < bk));
if better {
best = Some((v.len(), n));
}
}
}
let (_, rep) = best.unwrap_or((0, key));
(sum, key, rep)
})
.collect();
smoothed.sort_by(|a, b| b.0.cmp(&a.0).then_with(|| a.1.cmp(&b.1)));
let mut taken: HashSet<Key> = HashSet::new();
let mut out = Vec::new();
for (votes, key, rep) in smoothed {
if out.len() >= limit {
break;
}
if taken.contains(&key) || taken.contains(&rep) {
continue;
}
taken.extend(self.neighbours(key, 2, 1));
taken.extend(self.neighbours(rep, 1, 1));
if let Some(members) = self.buckets.get(&rep) {
out.push(Region {
votes,
members: members.as_slice(),
});
}
}
out
}
}
pub(crate) fn medoid<T>(members: &[T], pos: impl Fn(&T) -> (f64, f64)) -> usize {
let m = &members[..members.len().min(64)];
if m.len() <= 2 {
return 0;
}
let pts: Vec<[f64; 3]> = m
.iter()
.map(|t| {
let (ra, dec) = pos(t);
[dec.cos() * ra.cos(), dec.cos() * ra.sin(), dec.sin()]
})
.collect();
let mut best = (f64::INFINITY, 0usize);
for (i, p) in pts.iter().enumerate() {
let total: f64 = pts
.iter()
.map(|q| {
let d = [p[0] - q[0], p[1] - q[1], p[2] - q[2]];
(d[0] * d[0] + d[1] * d[1] + d[2] * d[2]).sqrt()
})
.sum();
if total < best.0 {
best = (total, i);
}
}
best.1
}
#[cfg(test)]
mod tests {
use super::*;
fn deg(d: f64) -> f64 {
d.to_radians()
}
#[test]
fn a_split_peak_recombines_and_outranks_a_single_dense_bucket() {
let mut v = SkyVotes::new(deg(0.1), 0.05);
for i in 0..6 {
let ra = 10.0 + if i % 2 == 0 { -0.01 } else { 0.01 };
v.add(deg(ra), deg(20.05), 2.0, i);
}
for i in 0..4 {
v.add(deg(200.0), deg(-30.05), 2.0, 100 + i);
}
let r = v.regions(10);
assert_eq!(r[0].votes, 6);
assert!(r[0].members.iter().all(|&m| m < 100));
assert_eq!(r[1].votes, 4);
assert_eq!(r.len(), 2);
}
#[test]
fn ra_bins_are_square_near_the_pole() {
let mut v = SkyVotes::new(deg(0.1), 0.05);
for i in 0..10 {
v.add(deg(100.0 + f64::from(i) * 0.5), deg(89.55), 1.0, i);
}
assert!(v.len() <= 2, "{} buckets", v.len());
assert_eq!(v.regions(5)[0].votes, 10);
}
#[test]
fn ra_wraps_at_zero() {
let mut v = SkyVotes::new(deg(0.1), 0.05);
v.add(deg(359.99), deg(0.0), 1.0, 0);
v.add(deg(0.01), deg(0.0), 1.0, 1);
v.add(deg(0.03), deg(0.0), 1.0, 2);
assert_eq!(v.regions(5)[0].votes, 3);
}
#[test]
fn scale_separates_hypotheses_at_one_position() {
let mut v = SkyVotes::new(deg(0.1), 0.05);
for i in 0..3 {
v.add(deg(50.0), deg(10.0), 1.0, i);
}
for i in 0..3 {
v.add(deg(50.0), deg(10.0), 2.0, 10 + i);
}
let r = v.regions(5);
assert_eq!(r.len(), 2);
assert_eq!(r[0].votes, 3);
}
#[test]
fn the_medoid_ignores_a_stray_member() {
let pts = [(0.0, 0.0), (deg(5.0), 0.0), (0.0001, 0.0), (0.0002, 0.0)];
let i = medoid(&pts, |&p| p);
assert!(i == 2 || i == 3 || i == 0);
assert_ne!(i, 1);
}
}