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