use std::collections::HashSet;
use std::sync::Mutex;
use crate::cyclotomic::{IsRing, Units};
pub fn reachable_points<ZZ, F, R>(
n: usize,
num_threads: usize,
fold: F,
mut on_round: R,
) -> Vec<Vec<ZZ>>
where
ZZ: IsRing + Units + Send + Sync,
F: Fn(&ZZ) -> ZZ + Sync,
R: FnMut(usize, usize),
{
let start: ZZ = ZZ::zero();
let visited: Mutex<HashSet<ZZ>> = Mutex::new(HashSet::from([start]));
let mut round_pts: Vec<Vec<ZZ>> = vec![vec![start]];
let fold = &fold;
let num_threads = num_threads.max(1);
for k in 1..=n {
let last = round_pts.last().unwrap();
let per_thread = num_threads.max(last.len() / num_threads).max(1);
let curr: Mutex<Vec<ZZ>> = Mutex::new(Vec::new());
let visited_ref = &visited;
let curr_ref = &curr;
std::thread::scope(|s| {
for chunk in last.chunks(per_thread) {
s.spawn(move || {
for p in chunk {
for d in 0..ZZ::turn() {
let raw: ZZ = *p + <ZZ as Units>::unit(d);
let dest = fold(&raw);
let is_new = visited_ref.lock().unwrap().insert(dest);
if is_new {
curr_ref.lock().unwrap().push(dest);
}
}
}
});
}
});
let curr = curr.into_inner().unwrap();
on_round(k, curr.len());
round_pts.push(curr);
}
round_pts
}
#[cfg(test)]
mod tests {
use super::*;
use crate::cyclotomic::ZZ4;
use crate::cyclotomic::Zero;
use crate::cyclotomic::geometry::point_mod_rect;
#[test]
fn zz4_free_reaches_the_l1_ball() {
let rounds = reachable_points::<ZZ4, _, _>(2, 1, |p| *p, |_, _| {});
assert_eq!(rounds[0].len(), 1);
assert_eq!(rounds[1].len(), 4);
assert_eq!(rounds[2].len(), 8);
}
#[test]
fn zz4_unit_cell_fold_collapses_cardinal_steps() {
let anchor = ZZ4::zero();
let rounds =
reachable_points::<ZZ4, _, _>(1, 1, |p| point_mod_rect(p, &anchor, (1, 1)), |_, _| {});
assert_eq!(rounds[1].len(), 0);
}
#[test]
fn on_round_reports_each_round_count() {
let mut counts = Vec::new();
let rounds = reachable_points::<ZZ4, _, _>(2, 1, |p| *p, |_, c| counts.push(c));
assert_eq!(counts, vec![rounds[1].len(), rounds[2].len()]);
}
}