1use crate::{PixelsError, Result};
24use crossbeam_deque::{Injector, Stealer, Worker};
25use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering};
26use std::sync::{Arc, Condvar, Mutex};
27
28type Task = Box<dyn FnOnce() + Send>;
30
31const IDLE_PARK: std::time::Duration = std::time::Duration::from_secs(1);
39
40const BUSY_PARK: std::time::Duration = std::time::Duration::from_millis(1);
45
46thread_local! {
47 static ON_WORKER: std::cell::Cell<bool> = const { std::cell::Cell::new(false) };
49}
50
51struct Shared {
53 injector: Injector<Task>,
55 stealers: Vec<Stealer<Task>>,
57 shutdown: AtomicBool,
59 pending: AtomicUsize,
61 idle: Mutex<()>,
63 wake: Condvar,
64}
65
66impl Shared {
67 fn signal_one(&self) {
69 let _guard = self
72 .idle
73 .lock()
74 .unwrap_or_else(std::sync::PoisonError::into_inner);
75 self.wake.notify_one();
76 }
77
78 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 fn find_task(&self, local: &Worker<Task>) -> Option<Task> {
89 if let Some(task) = local.pop() {
91 return Some(task);
92 }
93 loop {
94 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 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#[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 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 #[must_use]
184 pub fn default_threads() -> usize {
185 std::thread::available_parallelism().map_or(1, std::num::NonZeroUsize::get)
186 }
187
188 pub fn with_default_threads() -> Result<Self> {
194 Self::new(Self::default_threads())
195 }
196
197 #[must_use]
203 pub fn on_worker_thread() -> bool {
204 ON_WORKER.with(std::cell::Cell::get)
205 }
206
207 #[must_use]
209 pub const fn threads(&self) -> usize {
210 self.threads
211 }
212
213 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 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#[derive(Debug)]
266struct Batch {
267 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 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 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 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
332fn 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
345fn 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
356fn 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 let _ = std::panic::catch_unwind(std::panic::AssertUnwindSafe(task));
366 shared.pending.fetch_sub(1, Ordering::SeqCst);
367 continue;
368 }
369 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 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 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 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 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 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 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 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 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 }
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 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 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 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}