use std::num::NonZeroUsize;
use std::sync::Mutex;
use std::sync::atomic::{AtomicUsize, Ordering};
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum Jobs {
Serial,
Threads(NonZeroUsize),
}
impl Default for Jobs {
fn default() -> Jobs {
Jobs::available()
}
}
impl Jobs {
#[must_use]
pub fn available() -> Jobs {
match std::thread::available_parallelism() {
Ok(n) if n.get() > 1 => Jobs::Threads(n),
_ => Jobs::Serial,
}
}
#[must_use]
pub fn count(self) -> usize {
match self {
Jobs::Serial => 1,
Jobs::Threads(n) => n.get(),
}
}
pub fn parse(arg: &str) -> Result<Jobs, String> {
if arg.is_empty() {
return Ok(Jobs::available());
}
match arg.parse::<NonZeroUsize>() {
Ok(n) if n.get() == 1 => Ok(Jobs::Serial),
Ok(n) => Ok(Jobs::Threads(n)),
Err(_) if arg == "0" => Err("-j0 asks for no workers at all".to_owned()),
Err(_) => Err(format!("`{arg}` is not a job count")),
}
}
}
pub fn run<T, R, F>(jobs: Jobs, items: &[T], work: F) -> Vec<R>
where
T: Sync,
R: Send,
F: Fn(usize, &T) -> R + Sync,
{
if items.is_empty() {
return Vec::new();
}
let workers = jobs.count().min(items.len());
if workers <= 1 {
return items.iter().enumerate().map(|(i, t)| work(i, t)).collect();
}
let slots: Vec<Mutex<Option<R>>> = items.iter().map(|_| Mutex::new(None)).collect();
let next = AtomicUsize::new(0);
std::thread::scope(|scope| {
for _ in 0..workers {
scope.spawn(|| {
loop {
let i = next.fetch_add(1, Ordering::Relaxed);
let Some(item) = items.get(i) else { break };
let result = work(i, item);
*slots[i].lock().expect("a slot lock is only held to store one result") =
Some(result);
}
});
}
});
slots
.into_iter()
.map(|slot| {
slot.into_inner()
.expect("a slot lock is only held to store one result")
.expect("every index was claimed exactly once")
})
.collect()
}
#[cfg(test)]
mod tests {
use std::sync::atomic::AtomicUsize;
use super::*;
#[test]
fn results_come_back_in_input_order_however_they_finish() {
let items: Vec<u64> = (0..8).collect();
let out = run(Jobs::Threads(NonZeroUsize::new(8).unwrap()), &items, |i, x| {
std::thread::sleep(std::time::Duration::from_millis((8 - i as u64) * 4));
x * 10
});
assert_eq!(out, vec![0, 10, 20, 30, 40, 50, 60, 70]);
}
#[test]
fn serial_and_parallel_give_the_same_answer() {
let items: Vec<usize> = (0..64).collect();
let serial = run(Jobs::Serial, &items, |i, x| i + x);
let parallel = run(Jobs::Threads(NonZeroUsize::new(4).unwrap()), &items, |i, x| i + x);
assert_eq!(serial, parallel);
}
#[test]
fn every_item_runs_exactly_once() {
let items: Vec<usize> = (0..500).collect();
let calls = AtomicUsize::new(0);
let out = run(Jobs::Threads(NonZeroUsize::new(16).unwrap()), &items, |_, x| {
calls.fetch_add(1, Ordering::Relaxed);
*x
});
assert_eq!(calls.load(Ordering::Relaxed), 500);
assert_eq!(out, items);
}
#[test]
fn more_workers_than_items_is_fine() {
let items = [1, 2];
let out = run(Jobs::Threads(NonZeroUsize::new(64).unwrap()), &items, |_, x| *x);
assert_eq!(out, vec![1, 2]);
}
#[test]
fn no_items_is_no_threads_and_no_results() {
let items: [u8; 0] = [];
let out: Vec<u8> = run(Jobs::available(), &items, |_, x| *x);
assert!(out.is_empty());
}
#[test]
fn dash_j_reads_the_way_make_reads_it() {
assert_eq!(Jobs::parse("1").unwrap(), Jobs::Serial);
assert_eq!(Jobs::parse("4").unwrap(), Jobs::Threads(NonZeroUsize::new(4).unwrap()));
assert_eq!(Jobs::parse("").unwrap(), Jobs::available());
assert!(Jobs::parse("0").is_err());
assert!(Jobs::parse("many").is_err());
}
#[test]
fn a_panicking_job_is_not_swallowed() {
let items = [0, 1, 2];
let r = std::panic::catch_unwind(|| {
run(Jobs::Threads(NonZeroUsize::new(3).unwrap()), &items, |_, x| {
assert!(*x != 1, "planted failure");
*x
})
});
assert!(r.is_err());
}
}