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
31struct Shared {
33 injector: Injector<Task>,
35 stealers: Vec<Stealer<Task>>,
37 shutdown: AtomicBool,
39 pending: AtomicUsize,
41 idle: Mutex<()>,
43 wake: Condvar,
44}
45
46impl Shared {
47 fn signal_one(&self) {
49 let _guard = self
52 .idle
53 .lock()
54 .unwrap_or_else(std::sync::PoisonError::into_inner);
55 self.wake.notify_one();
56 }
57
58 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 fn find_task(&self, local: &Worker<Task>) -> Option<Task> {
69 if let Some(task) = local.pop() {
71 return Some(task);
72 }
73 loop {
74 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 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#[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 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 #[must_use]
164 pub fn default_threads() -> usize {
165 std::thread::available_parallelism().map_or(1, std::num::NonZeroUsize::get)
166 }
167
168 pub fn with_default_threads() -> Result<Self> {
174 Self::new(Self::default_threads())
175 }
176
177 #[must_use]
179 pub const fn threads(&self) -> usize {
180 self.threads
181 }
182
183 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 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#[derive(Debug)]
236struct Batch {
237 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 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 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 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
302fn 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
315fn 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
326fn 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 let _ = std::panic::catch_unwind(std::panic::AssertUnwindSafe(task));
335 shared.pending.fetch_sub(1, Ordering::SeqCst);
336 continue;
337 }
338 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 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 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 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 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 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 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 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 }
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 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 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 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}