Skip to main content

sva_engine/
threads.rs

1// Concern: runs independent tasks on at most a thread limit, answers in task order | Non-concern: which tasks are independent | IO: (limit, tasks) -> answers
2
3use std::num::NonZeroUsize;
4use std::sync::Mutex;
5use std::sync::atomic::{AtomicUsize, Ordering};
6
7pub fn default_threads() -> NonZeroUsize {
8    std::thread::available_parallelism().unwrap_or(NonZeroUsize::MIN)
9}
10
11pub(crate) fn each<T: Send, R: Send>(
12    limit: NonZeroUsize,
13    tasks: Vec<T>,
14    run: impl Fn(T) -> R + Sync,
15) -> Vec<R> {
16    let workers = limit.get().min(tasks.len());
17    if workers <= 1 {
18        return tasks.into_iter().map(run).collect();
19    }
20    let answers: Vec<Mutex<Option<R>>> = tasks.iter().map(|_| Mutex::new(None)).collect();
21    let tasks: Vec<Mutex<Option<T>>> = tasks.into_iter().map(|t| Mutex::new(Some(t))).collect();
22    let next = AtomicUsize::new(0);
23    let work = || {
24        loop {
25            let at = next.fetch_add(1, Ordering::Relaxed);
26            let Some(task) = tasks.get(at) else {
27                return;
28            };
29            let task = task.lock().expect("a task").take().expect("taken once");
30            *answers[at].lock().expect("an answer") = Some(run(task));
31        }
32    };
33    rayon::scope(|scope| {
34        for _ in 1..workers {
35            scope.spawn(|_| work());
36        }
37        work();
38    });
39    answers
40        .into_iter()
41        .map(|answer| {
42            answer
43                .into_inner()
44                .expect("an answer")
45                .expect("each task ran")
46        })
47        .collect()
48}
49
50#[cfg(test)]
51mod tests {
52    use std::sync::atomic::{AtomicUsize, Ordering};
53
54    #[test]
55    fn tasks_run_on_at_most_the_limit_and_answer_in_task_order() {
56        let (running, most) = (AtomicUsize::new(0), AtomicUsize::new(0));
57        let limit = std::num::NonZeroUsize::new(3).expect("three");
58        let answers = super::each(limit, (0..64u64).collect(), |task| {
59            let now = running.fetch_add(1, Ordering::SeqCst) + 1;
60            most.fetch_max(now, Ordering::SeqCst);
61            let spun = (0..10_000u64).fold(task, |held, k| std::hint::black_box(held ^ k));
62            running.fetch_sub(1, Ordering::SeqCst);
63            (task, spun)
64        });
65        let tasks: Vec<u64> = answers.iter().map(|(task, _)| *task).collect();
66        assert_eq!(tasks, (0..64).collect::<Vec<_>>());
67        assert!(most.load(Ordering::SeqCst) <= 3, "{most:?} at once");
68    }
69}