use std::sync::atomic::{AtomicUsize, Ordering};
pub(crate) fn parallel_drain<R: Send>(
total: usize,
workers: usize,
init: impl Fn() -> R + Sync,
fold: impl Fn(&mut R, usize) + Sync,
reduce: impl Fn(R, R) -> R,
) -> R {
let workers = workers.max(1);
let next = AtomicUsize::new(0);
let (next, init, fold) = (&next, &init, &fold);
std::thread::scope(|s| {
let handles: Vec<_> = (0..workers)
.map(|_| {
s.spawn(move || {
let mut acc = init();
loop {
let i = next.fetch_add(1, Ordering::Relaxed);
if i >= total {
break;
}
fold(&mut acc, i);
}
acc
})
})
.collect();
handles
.into_iter()
.map(|h| h.join().expect("parallel_drain worker panicked"))
.reduce(reduce)
.unwrap_or_else(init)
})
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn sums_all_indices_work_stealing() {
let total = 10_000usize;
let got = parallel_drain(total, 8, || 0usize, |acc, i| *acc += i, |a, b| a + b);
assert_eq!(got, total * (total - 1) / 2);
}
#[test]
fn every_index_folded_exactly_once() {
let total = 5_000usize;
let hits = parallel_drain(
total,
8,
Vec::<usize>::new,
|acc, i| acc.push(i),
|mut a, b| {
a.extend(b);
a
},
);
let mut hits = hits;
hits.sort_unstable();
assert_eq!(hits.len(), total);
assert!(hits.iter().enumerate().all(|(k, &v)| k == v));
}
#[test]
fn empty_and_single_worker() {
assert_eq!(parallel_drain(0, 4, || 7usize, |_, _| {}, |a, _| a), 7);
assert_eq!(
parallel_drain(100, 1, || 0usize, |a, _| *a += 1, |a, b| a + b),
100
);
}
}