use crate::{PixelsError, Result};
use crossbeam_deque::{Injector, Stealer, Worker};
use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering};
use std::sync::{Arc, Condvar, Mutex};
type Task = Box<dyn FnOnce() + Send>;
struct Shared {
injector: Injector<Task>,
stealers: Vec<Stealer<Task>>,
shutdown: AtomicBool,
pending: AtomicUsize,
idle: Mutex<()>,
wake: Condvar,
}
impl Shared {
fn signal_one(&self) {
let _guard = self
.idle
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
self.wake.notify_one();
}
fn signal_all(&self) {
let _guard = self
.idle
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
self.wake.notify_all();
}
fn find_task(&self, local: &Worker<Task>) -> Option<Task> {
if let Some(task) = local.pop() {
return Some(task);
}
loop {
match self.injector.steal_batch_and_pop(local) {
crossbeam_deque::Steal::Success(task) => return Some(task),
crossbeam_deque::Steal::Retry => continue,
crossbeam_deque::Steal::Empty => break,
}
}
for stealer in &self.stealers {
loop {
match stealer.steal_batch_and_pop(local) {
crossbeam_deque::Steal::Success(task) => return Some(task),
crossbeam_deque::Steal::Retry => continue,
crossbeam_deque::Steal::Empty => break,
}
}
}
None
}
}
#[derive(Debug)]
pub struct ThreadPool {
shared: Arc<Shared>,
workers: Vec<std::thread::JoinHandle<()>>,
threads: usize,
}
impl std::fmt::Debug for Shared {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("Shared")
.field("workers", &self.stealers.len())
.field("pending", &self.pending.load(Ordering::Relaxed))
.finish_non_exhaustive()
}
}
impl ThreadPool {
pub fn new(threads: usize) -> Result<Self> {
let threads = threads.max(1);
let mut locals = Vec::with_capacity(threads);
let mut stealers = Vec::with_capacity(threads);
for _ in 0..threads {
let worker = Worker::new_lifo();
stealers.push(worker.stealer());
locals.push(worker);
}
let shared = Arc::new(Shared {
injector: Injector::new(),
stealers,
shutdown: AtomicBool::new(false),
pending: AtomicUsize::new(0),
idle: Mutex::new(()),
wake: Condvar::new(),
});
let mut workers = Vec::with_capacity(threads);
for (index, local) in locals.into_iter().enumerate() {
let shared = Arc::clone(&shared);
let handle = std::thread::Builder::new()
.name(format!("otf-pixels-worker-{index}"))
.spawn(move || worker_loop(&shared, &local))
.map_err(|e| PixelsError::io("spawning a scheduler worker thread", e))?;
workers.push(handle);
}
Ok(Self {
shared,
workers,
threads,
})
}
#[must_use]
pub fn default_threads() -> usize {
std::thread::available_parallelism().map_or(1, std::num::NonZeroUsize::get)
}
pub fn with_default_threads() -> Result<Self> {
Self::new(Self::default_threads())
}
#[must_use]
pub const fn threads(&self) -> usize {
self.threads
}
pub fn spawn(&self, task: impl FnOnce() + Send + 'static) {
self.shared.pending.fetch_add(1, Ordering::SeqCst);
self.shared.injector.push(Box::new(task));
self.shared.signal_one();
}
pub fn run_all<F>(&self, tasks: Vec<F>) -> Result<()>
where
F: FnOnce() -> Result<()> + Send + 'static,
{
if tasks.is_empty() {
return Ok(());
}
let batch = Arc::new(Batch::new(tasks.len()));
for (index, task) in tasks.into_iter().enumerate() {
let batch = Arc::clone(&batch);
self.spawn(move || {
let outcome = catch(task);
batch.finish(index, outcome);
});
}
batch.wait();
batch.first_error()
}
}
#[derive(Debug)]
struct Batch {
slots: Vec<Mutex<Option<PixelsError>>>,
remaining: AtomicUsize,
finished: Mutex<bool>,
complete: Condvar,
}
impl Batch {
fn new(count: usize) -> Self {
Self {
slots: (0..count).map(|_| Mutex::new(None)).collect(),
remaining: AtomicUsize::new(count),
finished: Mutex::new(false),
complete: Condvar::new(),
}
}
fn finish(&self, index: usize, outcome: Result<()>) {
if let Err(error) = outcome {
if let Some(slot) = self.slots.get(index) {
*slot
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner) = Some(error);
}
}
if self.remaining.fetch_sub(1, Ordering::SeqCst) == 1 {
let mut finished = self
.finished
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
*finished = true;
self.complete.notify_all();
}
}
fn wait(&self) {
let mut finished = self
.finished
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
while !*finished {
finished = self
.complete
.wait(finished)
.unwrap_or_else(std::sync::PoisonError::into_inner);
}
}
fn first_error(&self) -> Result<()> {
for slot in &self.slots {
let mut slot = slot
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
if let Some(error) = slot.take() {
return Err(error);
}
}
Ok(())
}
}
fn catch<F: FnOnce() -> Result<()>>(task: F) -> Result<()> {
match std::panic::catch_unwind(std::panic::AssertUnwindSafe(task)) {
Ok(result) => result,
Err(payload) => {
let detail = panic_message(payload.as_ref());
Err(PixelsError::graph(format!(
"a scheduler task panicked: {detail}"
)))
}
}
}
fn panic_message(payload: &(dyn std::any::Any + Send)) -> String {
if let Some(text) = payload.downcast_ref::<&str>() {
return (*text).to_owned();
}
if let Some(text) = payload.downcast_ref::<String>() {
return text.clone();
}
"non-string panic payload".to_owned()
}
fn worker_loop(shared: &Arc<Shared>, local: &Worker<Task>) {
loop {
if shared.shutdown.load(Ordering::SeqCst) && shared.pending.load(Ordering::SeqCst) == 0 {
return;
}
if let Some(task) = shared.find_task(local) {
let _ = std::panic::catch_unwind(std::panic::AssertUnwindSafe(task));
shared.pending.fetch_sub(1, Ordering::SeqCst);
continue;
}
let guard = shared
.idle
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
if shared.shutdown.load(Ordering::SeqCst) {
return;
}
let _unused = shared
.wake
.wait_timeout(guard, std::time::Duration::from_millis(1))
.unwrap_or_else(std::sync::PoisonError::into_inner);
}
}
impl Drop for ThreadPool {
fn drop(&mut self) {
self.shared.shutdown.store(true, Ordering::SeqCst);
self.shared.signal_all();
for handle in self.workers.drain(..) {
let _ = handle.join();
}
}
}
#[cfg(test)]
#[allow(
clippy::unwrap_used,
clippy::indexing_slicing,
clippy::panic,
reason = "tests operate on known-good values and assert shapes directly"
)]
mod tests {
use super::*;
fn counter() -> Arc<AtomicUsize> {
Arc::new(AtomicUsize::new(0))
}
#[test]
fn every_task_runs_exactly_once() {
let pool = ThreadPool::new(4).unwrap();
let count = counter();
let tasks: Vec<_> = (0..1000)
.map(|_| {
let count = Arc::clone(&count);
move || {
count.fetch_add(1, Ordering::Relaxed);
Ok(())
}
})
.collect();
pool.run_all(tasks).unwrap();
assert_eq!(count.load(Ordering::Relaxed), 1000);
}
#[test]
fn tasks_share_state_through_arcs() {
let pool = ThreadPool::new(4).unwrap();
let data: Arc<Vec<usize>> = Arc::new((0..100).collect());
let total = counter();
let tasks: Vec<_> = (0..10)
.map(|chunk| {
let (data, total) = (Arc::clone(&data), Arc::clone(&total));
move || {
let sum: usize = data[chunk * 10..(chunk + 1) * 10].iter().sum();
total.fetch_add(sum, Ordering::Relaxed);
Ok(())
}
})
.collect();
pool.run_all(tasks).unwrap();
assert_eq!(total.load(Ordering::Relaxed), (0..100).sum::<usize>());
}
#[test]
fn the_lowest_indexed_failure_is_reported() {
let pool = ThreadPool::new(8).unwrap();
for attempt in 0..25 {
let tasks: Vec<_> = (0..64)
.map(|i| {
move || {
if i == 5 || i == 40 {
return Err(PixelsError::malformed("test", format!("task {i}")));
}
Ok(())
}
})
.collect();
let err = pool.run_all(tasks).unwrap_err();
assert!(
err.to_string().contains("task 5"),
"attempt {attempt}: {err}"
);
}
}
#[test]
fn a_panicking_task_becomes_an_error_not_an_abort() {
let pool = ThreadPool::new(4).unwrap();
let tasks: Vec<_> = (0..8)
.map(|i| {
move || {
assert!(i != 3, "kernel defect");
Ok(())
}
})
.collect();
let err = pool.run_all(tasks).unwrap_err();
assert_eq!(err.code(), crate::ErrorCode::Graph);
assert!(err.to_string().contains("panicked"), "got: {err}");
assert!(err.to_string().contains("kernel defect"), "got: {err}");
let count = counter();
let c = Arc::clone(&count);
pool.run_all(vec![move || {
c.fetch_add(1, Ordering::Relaxed);
Ok(())
}])
.unwrap();
assert_eq!(count.load(Ordering::Relaxed), 1);
}
#[test]
fn a_single_threaded_pool_still_completes() {
let pool = ThreadPool::new(1).unwrap();
let count = counter();
let tasks: Vec<_> = (0..100)
.map(|_| {
let count = Arc::clone(&count);
move || {
count.fetch_add(1, Ordering::Relaxed);
Ok(())
}
})
.collect();
pool.run_all(tasks).unwrap();
assert_eq!(count.load(Ordering::Relaxed), 100);
assert_eq!(pool.threads(), 1);
}
#[test]
fn zero_threads_is_clamped_to_one() {
assert_eq!(ThreadPool::new(0).unwrap().threads(), 1);
}
#[test]
fn an_empty_batch_is_a_no_op() {
let pool = ThreadPool::new(2).unwrap();
let tasks: Vec<fn() -> Result<()>> = Vec::new();
pool.run_all(tasks).unwrap();
}
#[test]
fn repeated_batches_reuse_the_same_workers() {
let pool = ThreadPool::new(4).unwrap();
let count = counter();
for _ in 0..50 {
let tasks: Vec<_> = (0..20)
.map(|_| {
let count = Arc::clone(&count);
move || {
count.fetch_add(1, Ordering::Relaxed);
Ok(())
}
})
.collect();
pool.run_all(tasks).unwrap();
}
assert_eq!(count.load(Ordering::Relaxed), 1000);
}
#[test]
fn outstanding_spawned_work_completes_before_drop() {
let done = counter();
{
let pool = ThreadPool::new(4).unwrap();
for _ in 0..200 {
let done = Arc::clone(&done);
pool.spawn(move || {
done.fetch_add(1, Ordering::Relaxed);
});
}
}
assert_eq!(done.load(Ordering::Relaxed), 200);
}
#[test]
fn default_threads_is_at_least_one() {
assert!(ThreadPool::default_threads() >= 1);
assert!(ThreadPool::with_default_threads().unwrap().threads() >= 1);
}
#[test]
fn work_is_actually_distributed_across_workers() {
let pool = ThreadPool::new(4).unwrap();
let seen: Arc<Mutex<std::collections::HashSet<std::thread::ThreadId>>> =
Arc::new(Mutex::new(std::collections::HashSet::new()));
let tasks: Vec<_> = (0..2000)
.map(|_| {
let seen = Arc::clone(&seen);
move || {
std::hint::black_box((0..500_u64).sum::<u64>());
seen.lock().unwrap().insert(std::thread::current().id());
Ok(())
}
})
.collect();
pool.run_all(tasks).unwrap();
let count = seen.lock().unwrap().len();
assert!(
count > 1,
"all work ran on one thread; stealing is not happening"
);
}
#[test]
fn nested_arcs_keep_results_alive_across_batches() {
let pool = ThreadPool::new(4).unwrap();
let stage1: Arc<Mutex<Vec<u64>>> = Arc::new(Mutex::new(vec![0; 16]));
let tasks: Vec<_> = (0..16_u64)
.map(|i| {
let out = Arc::clone(&stage1);
move || {
out.lock().unwrap()[i as usize] = i * 2;
Ok(())
}
})
.collect();
pool.run_all(tasks).unwrap();
let total = Arc::new(AtomicUsize::new(0));
let tasks: Vec<_> = (0..16_usize)
.map(|i| {
let (input, total) = (Arc::clone(&stage1), Arc::clone(&total));
move || {
let v = input.lock().unwrap()[i];
total.fetch_add(v as usize, Ordering::Relaxed);
Ok(())
}
})
.collect();
pool.run_all(tasks).unwrap();
assert_eq!(
total.load(Ordering::Relaxed),
(0..16).map(|i| i * 2).sum::<usize>()
);
}
}