Skip to main content

otf_pixels_core/
pool.rs

1//! The work-stealing thread pool.
2//!
3//! Per ADR-0008 the lock-free deque comes from `crossbeam-deque`; the pool,
4//! the parking policy and every scheduling decision are ours.
5//!
6//! # Shape
7//!
8//! Each worker owns a LIFO deque and pushes new work onto it. LIFO on the
9//! owner's end is deliberate: the most recently produced tile is the one most
10//! likely still in cache, so depth-first execution on each worker keeps the
11//! working set small. Idle workers steal from the *other* end of a victim's
12//! deque, taking the oldest and coldest work — which is also the work least
13//! likely to be stolen back immediately.
14//!
15//! # Panics are contained, not propagated
16//!
17//! A panicking task must not poison the pool or abort the process
18//! (ARCHITECTURE §Failure model). Tasks are caught, and the panic is reported
19//! to the submitter as a [`PixelsError`] rather than resumed. Ops are written
20//! to return errors, so a panic here means a defect — but a defect in one tile
21//! still must not take the host process down with it.
22
23use crate::{PixelsError, Result};
24use crossbeam_deque::{Injector, Stealer, Worker};
25use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering};
26use std::sync::{Arc, Condvar, Mutex};
27
28/// A unit of work for the pool.
29type Task = Box<dyn FnOnce() + Send>;
30
31/// How long an idle worker sleeps when no task is pending anywhere.
32///
33/// Correctness does not depend on it: a worker parks only after seeing
34/// `pending == 0` under the parking lock, and a submitter raises `pending`
35/// before signalling under that same lock, so the wakeup cannot be missed. It
36/// is a backstop, long enough that an idle pool living as long as the process
37/// costs nothing measurable.
38const IDLE_PARK: std::time::Duration = std::time::Duration::from_secs(1);
39
40/// How long an idle worker sleeps while tasks are pending elsewhere.
41///
42/// Work another worker has pulled into its own deque raises no signal, so a
43/// worker that might steal it polls briefly instead.
44const BUSY_PARK: std::time::Duration = std::time::Duration::from_millis(1);
45
46thread_local! {
47    /// Whether this thread is a worker of some [`ThreadPool`].
48    static ON_WORKER: std::cell::Cell<bool> = const { std::cell::Cell::new(false) };
49}
50
51/// Shared state every worker sees.
52struct Shared {
53    /// Global queue: where non-worker threads submit work.
54    injector: Injector<Task>,
55    /// One stealer per worker, for cross-worker theft.
56    stealers: Vec<Stealer<Task>>,
57    /// Set once to tell workers to wind down.
58    shutdown: AtomicBool,
59    /// Tasks queued but not yet finished; drives idle parking.
60    pending: AtomicUsize,
61    /// Parking lot for idle workers.
62    idle: Mutex<()>,
63    wake: Condvar,
64}
65
66impl Shared {
67    /// Wake one parked worker, if any.
68    fn signal_one(&self) {
69        // The lock is taken so a worker cannot check `pending`, decide to
70        // park, and miss this notification in between.
71        let _guard = self
72            .idle
73            .lock()
74            .unwrap_or_else(std::sync::PoisonError::into_inner);
75        self.wake.notify_one();
76    }
77
78    /// Wake every parked worker.
79    fn signal_all(&self) {
80        let _guard = self
81            .idle
82            .lock()
83            .unwrap_or_else(std::sync::PoisonError::into_inner);
84        self.wake.notify_all();
85    }
86
87    /// Take one task, preferring local work, then global, then theft.
88    fn find_task(&self, local: &Worker<Task>) -> Option<Task> {
89        // Local LIFO first: hottest tile, no synchronisation.
90        if let Some(task) = local.pop() {
91            return Some(task);
92        }
93        loop {
94            // Global queue next, in batches so a submitter burst spreads out.
95            match self.injector.steal_batch_and_pop(local) {
96                crossbeam_deque::Steal::Success(task) => return Some(task),
97                crossbeam_deque::Steal::Retry => continue,
98                crossbeam_deque::Steal::Empty => break,
99            }
100        }
101        // Finally steal from a peer's cold end.
102        for stealer in &self.stealers {
103            loop {
104                match stealer.steal_batch_and_pop(local) {
105                    crossbeam_deque::Steal::Success(task) => return Some(task),
106                    crossbeam_deque::Steal::Retry => continue,
107                    crossbeam_deque::Steal::Empty => break,
108                }
109            }
110        }
111        None
112    }
113}
114
115/// A work-stealing pool of worker threads.
116///
117/// Dropping the pool signals shutdown and joins every worker, so no task
118/// outlives it.
119#[derive(Debug)]
120pub struct ThreadPool {
121    shared: Arc<Shared>,
122    workers: Vec<std::thread::JoinHandle<()>>,
123    threads: usize,
124}
125
126impl std::fmt::Debug for Shared {
127    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
128        f.debug_struct("Shared")
129            .field("workers", &self.stealers.len())
130            .field("pending", &self.pending.load(Ordering::Relaxed))
131            .finish_non_exhaustive()
132    }
133}
134
135impl ThreadPool {
136    /// Build a pool with `threads` workers.
137    ///
138    /// `threads` is clamped to at least one. Use [`ThreadPool::default_threads`]
139    /// for the machine's parallelism.
140    ///
141    /// # Errors
142    ///
143    /// Returns [`PixelsError::Io`] if the operating system refuses to spawn a
144    /// worker thread.
145    pub fn new(threads: usize) -> Result<Self> {
146        let threads = threads.max(1);
147        let mut locals = Vec::with_capacity(threads);
148        let mut stealers = Vec::with_capacity(threads);
149        for _ in 0..threads {
150            let worker = Worker::new_lifo();
151            stealers.push(worker.stealer());
152            locals.push(worker);
153        }
154        let shared = Arc::new(Shared {
155            injector: Injector::new(),
156            stealers,
157            shutdown: AtomicBool::new(false),
158            pending: AtomicUsize::new(0),
159            idle: Mutex::new(()),
160            wake: Condvar::new(),
161        });
162
163        let mut workers = Vec::with_capacity(threads);
164        for (index, local) in locals.into_iter().enumerate() {
165            let shared = Arc::clone(&shared);
166            let handle = std::thread::Builder::new()
167                .name(format!("otf-pixels-worker-{index}"))
168                .spawn(move || worker_loop(&shared, &local))
169                .map_err(|e| PixelsError::io("spawning a scheduler worker thread", e))?;
170            workers.push(handle);
171        }
172        Ok(Self {
173            shared,
174            workers,
175            threads,
176        })
177    }
178
179    /// The parallelism to use when the caller has no preference.
180    ///
181    /// Falls back to one thread where the platform cannot report it, which
182    /// yields a correct if serial engine rather than a failure.
183    #[must_use]
184    pub fn default_threads() -> usize {
185        std::thread::available_parallelism().map_or(1, std::num::NonZeroUsize::get)
186    }
187
188    /// A pool sized to [`ThreadPool::default_threads`].
189    ///
190    /// # Errors
191    ///
192    /// As [`ThreadPool::new`].
193    pub fn with_default_threads() -> Result<Self> {
194        Self::new(Self::default_threads())
195    }
196
197    /// Whether the calling thread is a worker of any pool.
198    ///
199    /// A run blocks its caller until its tiles are done, so a run started
200    /// from inside a task would hold a worker of the pool it is waiting on.
201    /// Callers that would otherwise use a shared pool check this first.
202    #[must_use]
203    pub fn on_worker_thread() -> bool {
204        ON_WORKER.with(std::cell::Cell::get)
205    }
206
207    /// How many workers this pool runs.
208    #[must_use]
209    pub const fn threads(&self) -> usize {
210        self.threads
211    }
212
213    /// Queue `task` for execution.
214    ///
215    /// Returns immediately. Panics inside `task` are contained by the worker
216    /// and do not propagate here.
217    pub fn spawn(&self, task: impl FnOnce() + Send + 'static) {
218        self.shared.pending.fetch_add(1, Ordering::SeqCst);
219        self.shared.injector.push(Box::new(task));
220        self.shared.signal_one();
221    }
222
223    /// Run `tasks` on the pool and return once every one has finished.
224    ///
225    /// # Why `'static`
226    ///
227    /// Tasks must be `'static` because they run on long-lived worker threads
228    /// that outlive this call. Erasing a shorter lifetime is what `rayon` uses
229    /// `unsafe` for, and ADR-0008 keeps `unsafe_code = "forbid"` in our
230    /// crates. This costs the scheduler nothing: tiles are already
231    /// `Arc<TileBuf>` and graph nodes `Arc<Node>`, so a task moves cheap
232    /// handles rather than borrowing.
233    ///
234    /// The calling thread blocks. Calling `run_all` from *inside* a pool task
235    /// is not supported and can deadlock; the scheduler submits one flat batch
236    /// per wave instead of nesting.
237    ///
238    /// # Errors
239    ///
240    /// Returns the failing task's error, choosing the **lowest-indexed**
241    /// failure when several fail. Which worker happens to fail first is a race;
242    /// which task index is lowest is not, so the reported error is
243    /// deterministic (SPEC §Guarantees 2).
244    pub fn run_all<F>(&self, tasks: Vec<F>) -> Result<()>
245    where
246        F: FnOnce() -> Result<()> + Send + 'static,
247    {
248        if tasks.is_empty() {
249            return Ok(());
250        }
251        let batch = Arc::new(Batch::new(tasks.len()));
252        for (index, task) in tasks.into_iter().enumerate() {
253            let batch = Arc::clone(&batch);
254            self.spawn(move || {
255                let outcome = catch(task);
256                batch.finish(index, outcome);
257            });
258        }
259        batch.wait();
260        batch.first_error()
261    }
262}
263
264/// Completion tracking for one [`ThreadPool::run_all`] batch.
265#[derive(Debug)]
266struct Batch {
267    /// One slot per task, so failures can be ranked by task index rather than
268    /// by which worker got there first.
269    slots: Vec<Mutex<Option<PixelsError>>>,
270    remaining: AtomicUsize,
271    finished: Mutex<bool>,
272    complete: Condvar,
273}
274
275impl Batch {
276    fn new(count: usize) -> Self {
277        Self {
278            slots: (0..count).map(|_| Mutex::new(None)).collect(),
279            remaining: AtomicUsize::new(count),
280            finished: Mutex::new(false),
281            complete: Condvar::new(),
282        }
283    }
284
285    /// Record one task's outcome and wake the waiter if it was the last.
286    fn finish(&self, index: usize, outcome: Result<()>) {
287        if let Err(error) = outcome {
288            if let Some(slot) = self.slots.get(index) {
289                *slot
290                    .lock()
291                    .unwrap_or_else(std::sync::PoisonError::into_inner) = Some(error);
292            }
293        }
294        if self.remaining.fetch_sub(1, Ordering::SeqCst) == 1 {
295            let mut finished = self
296                .finished
297                .lock()
298                .unwrap_or_else(std::sync::PoisonError::into_inner);
299            *finished = true;
300            self.complete.notify_all();
301        }
302    }
303
304    /// Block until every task in the batch has finished.
305    fn wait(&self) {
306        let mut finished = self
307            .finished
308            .lock()
309            .unwrap_or_else(std::sync::PoisonError::into_inner);
310        while !*finished {
311            finished = self
312                .complete
313                .wait(finished)
314                .unwrap_or_else(std::sync::PoisonError::into_inner);
315        }
316    }
317
318    /// The lowest-indexed failure, if any task failed.
319    fn first_error(&self) -> Result<()> {
320        for slot in &self.slots {
321            let mut slot = slot
322                .lock()
323                .unwrap_or_else(std::sync::PoisonError::into_inner);
324            if let Some(error) = slot.take() {
325                return Err(error);
326            }
327        }
328        Ok(())
329    }
330}
331
332/// Run `task`, converting a panic into an error.
333fn catch<F: FnOnce() -> Result<()>>(task: F) -> Result<()> {
334    match std::panic::catch_unwind(std::panic::AssertUnwindSafe(task)) {
335        Ok(result) => result,
336        Err(payload) => {
337            let detail = panic_message(payload.as_ref());
338            Err(PixelsError::graph(format!(
339                "a scheduler task panicked: {detail}"
340            )))
341        }
342    }
343}
344
345/// Best-effort text of a panic payload.
346fn panic_message(payload: &(dyn std::any::Any + Send)) -> String {
347    if let Some(text) = payload.downcast_ref::<&str>() {
348        return (*text).to_owned();
349    }
350    if let Some(text) = payload.downcast_ref::<String>() {
351        return text.clone();
352    }
353    "non-string panic payload".to_owned()
354}
355
356/// The body of each worker thread.
357fn worker_loop(shared: &Arc<Shared>, local: &Worker<Task>) {
358    ON_WORKER.with(|flag| flag.set(true));
359    loop {
360        if shared.shutdown.load(Ordering::SeqCst) && shared.pending.load(Ordering::SeqCst) == 0 {
361            return;
362        }
363        if let Some(task) = shared.find_task(local) {
364            // A panicking task must not unwind out of the worker thread.
365            let _ = std::panic::catch_unwind(std::panic::AssertUnwindSafe(task));
366            shared.pending.fetch_sub(1, Ordering::SeqCst);
367            continue;
368        }
369        // Nothing to do: park until woken. With nothing pending anywhere the
370        // next submission is guaranteed to signal, so the sleep is long;
371        // with work pending elsewhere, poll briefly in case it can be stolen.
372        let guard = shared
373            .idle
374            .lock()
375            .unwrap_or_else(std::sync::PoisonError::into_inner);
376        if shared.shutdown.load(Ordering::SeqCst) {
377            return;
378        }
379        let park = if shared.pending.load(Ordering::SeqCst) == 0 {
380            IDLE_PARK
381        } else {
382            BUSY_PARK
383        };
384        let _unused = shared
385            .wake
386            .wait_timeout(guard, park)
387            .unwrap_or_else(std::sync::PoisonError::into_inner);
388    }
389}
390
391impl Drop for ThreadPool {
392    fn drop(&mut self) {
393        self.shared.shutdown.store(true, Ordering::SeqCst);
394        self.shared.signal_all();
395        for handle in self.workers.drain(..) {
396            // A worker that panicked has already been contained; joining it
397            // is still correct and must not panic the dropping thread.
398            let _ = handle.join();
399        }
400    }
401}
402
403#[cfg(test)]
404#[allow(
405    clippy::unwrap_used,
406    clippy::indexing_slicing,
407    clippy::panic,
408    reason = "tests operate on known-good values and assert shapes directly"
409)]
410mod tests {
411    use super::*;
412
413    /// A shared counter, the `'static` shape every pool task uses.
414    fn counter() -> Arc<AtomicUsize> {
415        Arc::new(AtomicUsize::new(0))
416    }
417
418    #[test]
419    fn workers_know_they_are_workers() {
420        assert!(!ThreadPool::on_worker_thread());
421        let pool = ThreadPool::new(2).unwrap();
422        let seen = Arc::new(AtomicBool::new(false));
423        let flag = Arc::clone(&seen);
424        pool.run_all(vec![move || {
425            flag.store(ThreadPool::on_worker_thread(), Ordering::SeqCst);
426            Ok(())
427        }])
428        .unwrap();
429        assert!(seen.load(Ordering::SeqCst));
430        assert!(!ThreadPool::on_worker_thread());
431    }
432
433    #[test]
434    fn an_idle_pool_wakes_for_new_work_rather_than_polling_for_it() {
435        // Long enough for every worker to park on `IDLE_PARK`. If submitting
436        // did not wake one, the task would wait out the full park.
437        let pool = ThreadPool::new(2).unwrap();
438        std::thread::sleep(std::time::Duration::from_millis(50));
439        let started = std::time::Instant::now();
440        pool.run_all(vec![|| Ok(())]).unwrap();
441        assert!(
442            started.elapsed() < IDLE_PARK / 2,
443            "an idle pool took {:?} to run one task",
444            started.elapsed()
445        );
446    }
447
448    #[test]
449    fn every_task_runs_exactly_once() {
450        let pool = ThreadPool::new(4).unwrap();
451        let count = counter();
452        let tasks: Vec<_> = (0..1000)
453            .map(|_| {
454                let count = Arc::clone(&count);
455                move || {
456                    count.fetch_add(1, Ordering::Relaxed);
457                    Ok(())
458                }
459            })
460            .collect();
461        pool.run_all(tasks).unwrap();
462        assert_eq!(count.load(Ordering::Relaxed), 1000);
463    }
464
465    #[test]
466    fn tasks_share_state_through_arcs() {
467        // The shape the scheduler uses: tiles are Arc<TileBuf>, so a task
468        // moves handles rather than borrowing from the caller's frame.
469        let pool = ThreadPool::new(4).unwrap();
470        let data: Arc<Vec<usize>> = Arc::new((0..100).collect());
471        let total = counter();
472        let tasks: Vec<_> = (0..10)
473            .map(|chunk| {
474                let (data, total) = (Arc::clone(&data), Arc::clone(&total));
475                move || {
476                    let sum: usize = data[chunk * 10..(chunk + 1) * 10].iter().sum();
477                    total.fetch_add(sum, Ordering::Relaxed);
478                    Ok(())
479                }
480            })
481            .collect();
482        pool.run_all(tasks).unwrap();
483        assert_eq!(total.load(Ordering::Relaxed), (0..100).sum::<usize>());
484    }
485
486    #[test]
487    fn the_lowest_indexed_failure_is_reported() {
488        // Determinism (SPEC §Guarantees 2): whichever worker fails first, the
489        // reported error is always the lowest-indexed failing task.
490        let pool = ThreadPool::new(8).unwrap();
491        for attempt in 0..25 {
492            let tasks: Vec<_> = (0..64)
493                .map(|i| {
494                    move || {
495                        if i == 5 || i == 40 {
496                            return Err(PixelsError::malformed("test", format!("task {i}")));
497                        }
498                        Ok(())
499                    }
500                })
501                .collect();
502            let err = pool.run_all(tasks).unwrap_err();
503            assert!(
504                err.to_string().contains("task 5"),
505                "attempt {attempt}: {err}"
506            );
507        }
508    }
509
510    #[test]
511    fn a_panicking_task_becomes_an_error_not_an_abort() {
512        let pool = ThreadPool::new(4).unwrap();
513        let tasks: Vec<_> = (0..8)
514            .map(|i| {
515                move || {
516                    assert!(i != 3, "kernel defect");
517                    Ok(())
518                }
519            })
520            .collect();
521        let err = pool.run_all(tasks).unwrap_err();
522        assert_eq!(err.code(), crate::ErrorCode::Graph);
523        assert!(err.to_string().contains("panicked"), "got: {err}");
524        assert!(err.to_string().contains("kernel defect"), "got: {err}");
525
526        // The pool is still usable afterwards: one bad tile does not kill it.
527        let count = counter();
528        let c = Arc::clone(&count);
529        pool.run_all(vec![move || {
530            c.fetch_add(1, Ordering::Relaxed);
531            Ok(())
532        }])
533        .unwrap();
534        assert_eq!(count.load(Ordering::Relaxed), 1);
535    }
536
537    #[test]
538    fn a_single_threaded_pool_still_completes() {
539        // The caller blocks, so a one-worker pool must not deadlock.
540        let pool = ThreadPool::new(1).unwrap();
541        let count = counter();
542        let tasks: Vec<_> = (0..100)
543            .map(|_| {
544                let count = Arc::clone(&count);
545                move || {
546                    count.fetch_add(1, Ordering::Relaxed);
547                    Ok(())
548                }
549            })
550            .collect();
551        pool.run_all(tasks).unwrap();
552        assert_eq!(count.load(Ordering::Relaxed), 100);
553        assert_eq!(pool.threads(), 1);
554    }
555
556    #[test]
557    fn zero_threads_is_clamped_to_one() {
558        assert_eq!(ThreadPool::new(0).unwrap().threads(), 1);
559    }
560
561    #[test]
562    fn an_empty_batch_is_a_no_op() {
563        let pool = ThreadPool::new(2).unwrap();
564        let tasks: Vec<fn() -> Result<()>> = Vec::new();
565        pool.run_all(tasks).unwrap();
566    }
567
568    #[test]
569    fn repeated_batches_reuse_the_same_workers() {
570        // Worker threads are long-lived; a pipeline runs many waves.
571        let pool = ThreadPool::new(4).unwrap();
572        let count = counter();
573        for _ in 0..50 {
574            let tasks: Vec<_> = (0..20)
575                .map(|_| {
576                    let count = Arc::clone(&count);
577                    move || {
578                        count.fetch_add(1, Ordering::Relaxed);
579                        Ok(())
580                    }
581                })
582                .collect();
583            pool.run_all(tasks).unwrap();
584        }
585        assert_eq!(count.load(Ordering::Relaxed), 1000);
586    }
587
588    #[test]
589    fn outstanding_spawned_work_completes_before_drop() {
590        let done = counter();
591        {
592            let pool = ThreadPool::new(4).unwrap();
593            for _ in 0..200 {
594                let done = Arc::clone(&done);
595                pool.spawn(move || {
596                    done.fetch_add(1, Ordering::Relaxed);
597                });
598            }
599            // Dropping drains outstanding work, then joins.
600        }
601        assert_eq!(done.load(Ordering::Relaxed), 200);
602    }
603
604    #[test]
605    fn default_threads_is_at_least_one() {
606        assert!(ThreadPool::default_threads() >= 1);
607        assert!(ThreadPool::with_default_threads().unwrap().threads() >= 1);
608    }
609
610    #[test]
611    fn work_is_actually_distributed_across_workers() {
612        // Not merely "it completes": prove more than one thread ran tasks.
613        // The M2 scaling benchmark is meaningless if this does not hold.
614        let pool = ThreadPool::new(4).unwrap();
615        let seen: Arc<Mutex<std::collections::HashSet<std::thread::ThreadId>>> =
616            Arc::new(Mutex::new(std::collections::HashSet::new()));
617        let tasks: Vec<_> = (0..2000)
618            .map(|_| {
619                let seen = Arc::clone(&seen);
620                move || {
621                    // Enough work that tasks overlap rather than draining
622                    // faster than they are queued.
623                    std::hint::black_box((0..500_u64).sum::<u64>());
624                    seen.lock().unwrap().insert(std::thread::current().id());
625                    Ok(())
626                }
627            })
628            .collect();
629        pool.run_all(tasks).unwrap();
630        let count = seen.lock().unwrap().len();
631        assert!(
632            count > 1,
633            "all work ran on one thread; stealing is not happening"
634        );
635    }
636
637    #[test]
638    fn nested_arcs_keep_results_alive_across_batches() {
639        // Results produced in one wave feed the next, as tiles do.
640        let pool = ThreadPool::new(4).unwrap();
641        let stage1: Arc<Mutex<Vec<u64>>> = Arc::new(Mutex::new(vec![0; 16]));
642        let tasks: Vec<_> = (0..16_u64)
643            .map(|i| {
644                let out = Arc::clone(&stage1);
645                move || {
646                    out.lock().unwrap()[i as usize] = i * 2;
647                    Ok(())
648                }
649            })
650            .collect();
651        pool.run_all(tasks).unwrap();
652
653        let total = Arc::new(AtomicUsize::new(0));
654        let tasks: Vec<_> = (0..16_usize)
655            .map(|i| {
656                let (input, total) = (Arc::clone(&stage1), Arc::clone(&total));
657                move || {
658                    let v = input.lock().unwrap()[i];
659                    total.fetch_add(v as usize, Ordering::Relaxed);
660                    Ok(())
661                }
662            })
663            .collect();
664        pool.run_all(tasks).unwrap();
665        assert_eq!(
666            total.load(Ordering::Relaxed),
667            (0..16).map(|i| i * 2).sum::<usize>()
668        );
669    }
670}