use crate::banding::{Banding, Solution, W};
use crate::hash::{coeff_row, ribbon_hash, start};
use std::thread;
const G: usize = 1 << 14;
pub fn from_keys_parallel_seeded<const R: usize>(
keys: &[u64],
seed: u64,
window_shift: u32,
threads: usize,
) -> Solution<R> {
let mut band = Banding::<R>::new(crate::filter::num_slots_for(keys.len(), R), seed);
let num_slots = band.num_slots();
let num_starts = (num_slots - W + 1) as u64;
let threads = threads.max(1).min(num_slots / W);
if threads <= 1 {
band.add_all(keys);
return band.solve();
}
let window = crate::filter::window_size(window_shift);
let n_windows = num_slots.div_ceil(window);
let per = n_windows.div_ceil(threads);
let mut bounds: Vec<usize> = (0..=threads)
.map(|t| t.saturating_mul(per).saturating_mul(window).min(num_slots))
.collect();
bounds.dedup();
let nt = bounds.len() - 1;
let bucket_capacity = keys.len().div_ceil(nt);
let mut buckets: Vec<Vec<u64>> = (0..nt)
.map(|_| Vec::with_capacity(bucket_capacity))
.collect();
let mut deferred: Vec<u64> = Vec::with_capacity(keys.len() / 4);
for &k in keys {
let s = start(ribbon_hash(k, seed), num_starts) as usize;
let t = bounds.partition_point(|&b| b <= s) - 1;
let hi = bounds[t + 1];
if t + 1 < nt && s.saturating_add(G) >= hi {
deferred.push(k); } else {
buckets[t].push(k);
}
}
let coeff = band.coeff_rows_mut();
let mut slices: Vec<(usize, &mut [u64])> = Vec::with_capacity(nt);
let mut rest = &mut coeff[..];
let mut consumed = 0usize;
for t in 0..nt {
let len = bounds[t + 1] - bounds[t];
let (head, tail) = rest.split_at_mut(len);
slices.push((consumed, head));
rest = tail;
consumed += len;
}
let spilled: Vec<Vec<u64>> = thread::scope(|scope| {
let handles: Vec<_> = slices
.into_iter()
.zip(buckets.iter())
.map(|((base, slice), bkeys)| {
scope.spawn(move || band_slice(slice, base, num_starts, seed, bkeys))
})
.collect();
handles.into_iter().map(|h| h.join().unwrap()).collect()
});
band.add_all(&deferred);
for spill in spilled {
band.add_all(&spill);
}
band.solve()
}
fn band_slice(
slice: &mut [u64],
base: usize,
num_starts: u64,
seed: u64,
keys: &[u64],
) -> Vec<u64> {
let len = slice.len();
let mut spilled = Vec::new();
'key: for &k in keys {
let h = ribbon_hash(k, seed);
let mut i = start(h, num_starts) as usize - base;
let mut cr = coeff_row(h);
loop {
if i >= len {
spilled.push(k); continue 'key;
}
let cr_at_i = slice[i];
if cr_at_i == 0 {
slice[i] = cr;
break;
}
cr ^= cr_at_i;
if cr == 0 {
break;
}
let tz = cr.trailing_zeros() as usize;
i += tz;
cr >>= tz;
}
}
spilled
}
#[cfg(test)]
mod tests {
use crate::banding::solution_fnv;
use crate::filter::RibbonFilter;
fn mix64(mut z: u64) -> u64 {
z = (z ^ (z >> 30)).wrapping_mul(0xbf58_476d_1ce4_e5b9);
z = (z ^ (z >> 27)).wrapping_mul(0x94d0_49bb_1331_11eb);
z ^ (z >> 31)
}
fn keys(n: usize, seed: u64) -> Vec<u64> {
let mut s = seed;
(0..n)
.map(|_| {
s = s.wrapping_add(0x9e37_79b9_7f4a_7c15);
mix64(s)
})
.collect()
}
#[test]
fn parallel_build_is_bit_identical_to_sequential() {
let k = keys(300_000, 0xA11CE);
let seq = RibbonFilter::from_keys(&k);
for t in [2usize, 4, 8] {
let par = super::from_keys_parallel_seeded::<7>(&k, 0, 16, t);
assert_eq!(
solution_fnv(seq_segments(&seq)),
solution_fnv(par.segments()),
"parallel build (t={t}) diverges from sequential"
);
assert!(
k.iter().all(|&x| par.contains(x)),
"parallel: false negative t={t}"
);
}
}
fn seq_segments(f: &RibbonFilter) -> &[u64] {
f.solution_segments()
}
}