use std::any::Any;
use std::cell::UnsafeCell;
use std::panic::{AssertUnwindSafe, catch_unwind};
use std::sync::atomic::{AtomicBool, AtomicU64, AtomicUsize, Ordering};
use std::sync::{Arc, Condvar, Mutex};
use std::thread::JoinHandle;
const SPIN: u32 = 1 << 14;
const BITS: u32 = 20;
const ROUND: usize = 1 << BITS;
const MASK: u64 = (1 << BITS) - 1;
type Job = dyn Fn(usize, usize) + Sync;
fn unpack(s: u64) -> (u64, usize, usize) {
(s >> (2 * BITS), ((s >> BITS) & MASK) as usize, (s & MASK) as usize)
}
struct Inner {
state: AtomicU64,
jobs: [UnsafeCell<*const Job>; 2],
done: AtomicUsize,
sleepers: AtomicUsize,
panic: Mutex<Option<Box<dyn Any + Send>>>,
stop: AtomicBool,
lock: Mutex<()>,
wake: Condvar,
}
unsafe impl Sync for Inner {}
unsafe impl Send for Inner {}
impl Inner {
fn work(&self, worker: usize) {
let mut s = self.state.load(Ordering::Acquire);
loop {
let (generation, tasks, next) = unpack(s);
if next >= tasks {
return;
}
if let Err(now) =
self.state.compare_exchange_weak(s, s + 1, Ordering::Acquire, Ordering::Acquire)
{
s = now;
continue;
}
let job = unsafe { &**self.jobs[(generation & 1) as usize].get() };
if let Err(e) = catch_unwind(AssertUnwindSafe(|| job(next, worker))) {
self.panic
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner)
.get_or_insert(e);
}
self.done.fetch_add(1, Ordering::Release);
s = self.state.load(Ordering::Acquire);
}
}
fn pending(&self, order: Ordering) -> bool {
let (_, tasks, next) = unpack(self.state.load(order));
next < tasks
}
}
pub struct Pool {
inner: Arc<Inner>,
handles: Vec<JoinHandle<()>>,
run: Mutex<()>,
}
impl std::fmt::Debug for Pool {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("Pool").field("threads", &self.threads()).finish()
}
}
impl Pool {
#[must_use]
pub fn new(threads: usize) -> Self {
let noop: &'static Job = &|_, _| {};
let inner = Arc::new(Inner {
state: AtomicU64::new(0),
jobs: [UnsafeCell::new(noop as *const Job), UnsafeCell::new(noop as *const Job)],
done: AtomicUsize::new(0),
sleepers: AtomicUsize::new(0),
panic: Mutex::new(None),
stop: AtomicBool::new(false),
lock: Mutex::new(()),
wake: Condvar::new(),
});
let handles = (1..threads.max(1))
.map(|worker| {
let inner = Arc::clone(&inner);
std::thread::Builder::new()
.name(format!("kime-cpu-{worker}"))
.spawn(move || worker_loop(&inner, worker))
.expect("spawn a worker thread")
})
.collect();
Self { inner, handles, run: Mutex::new(()) }
}
#[must_use]
pub fn threads(&self) -> usize {
self.handles.len() + 1
}
pub fn run(&self, n: usize, f: &(dyn Fn(usize, usize) + Sync)) {
if n == 0 {
return;
}
if n == 1 || self.handles.is_empty() {
(0..n).for_each(|i| f(i, 0));
return;
}
let guard = self.run.lock().unwrap_or_else(std::sync::PoisonError::into_inner);
for base in (0..n).step_by(ROUND) {
let tasks = ROUND.min(n - base);
self.round(tasks, &|i, worker| f(base + i, worker));
}
drop(guard);
}
fn round(&self, tasks: usize, f: &(dyn Fn(usize, usize) + Sync)) {
let inner = &*self.inner;
let (last, _, _) = unpack(inner.state.load(Ordering::Relaxed));
let generation = (last + 1) & ((1 << (64 - 2 * BITS)) - 1);
unsafe {
*inner.jobs[(generation & 1) as usize].get() =
std::mem::transmute::<&(dyn Fn(usize, usize) + Sync + '_), &'static Job>(f)
as *const Job;
}
inner.done.store(0, Ordering::Relaxed);
inner.state.store(generation << (2 * BITS) | (tasks as u64) << BITS, Ordering::SeqCst);
if inner.sleepers.load(Ordering::SeqCst) > 0 {
let _l = inner.lock.lock().unwrap_or_else(std::sync::PoisonError::into_inner);
inner.wake.notify_all();
}
inner.work(0);
while inner.done.load(Ordering::Acquire) < tasks {
std::hint::spin_loop();
}
let panic = inner.panic.lock().unwrap_or_else(std::sync::PoisonError::into_inner).take();
if let Some(e) = panic {
std::panic::resume_unwind(e);
}
}
}
fn worker_loop(inner: &Inner, worker: usize) {
loop {
let mut spins = 0u32;
while !inner.pending(Ordering::Acquire) {
if inner.stop.load(Ordering::Relaxed) {
return;
}
if spins < SPIN {
spins += 1;
std::hint::spin_loop();
continue;
}
let l = inner.lock.lock().unwrap_or_else(std::sync::PoisonError::into_inner);
inner.sleepers.fetch_add(1, Ordering::SeqCst);
let l = if !inner.pending(Ordering::SeqCst) && !inner.stop.load(Ordering::SeqCst) {
inner.wake.wait(l).unwrap_or_else(std::sync::PoisonError::into_inner)
} else {
l
};
inner.sleepers.fetch_sub(1, Ordering::SeqCst);
drop(l);
spins = 0;
}
inner.work(worker);
}
}
impl Drop for Pool {
fn drop(&mut self) {
{
let _l = self.inner.lock.lock().unwrap_or_else(std::sync::PoisonError::into_inner);
self.inner.stop.store(true, Ordering::SeqCst);
self.inner.wake.notify_all();
}
for h in self.handles.drain(..) {
let _ = h.join();
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::sync::atomic::AtomicU32;
#[test]
fn every_task_once_on_a_valid_worker() {
for threads in [1, 2, 5] {
let pool = Pool::new(threads);
for n in [0, 1, 3, 1000] {
let hits: Vec<AtomicU32> = (0..n).map(|_| AtomicU32::new(0)).collect();
pool.run(n, &|i, w| {
assert!(w < threads);
hits[i].fetch_add(1, Ordering::Relaxed);
});
assert!(hits.iter().all(|h| h.load(Ordering::Relaxed) == 1));
}
}
}
#[test]
fn wakes_after_sleeping() {
let pool = Pool::new(3);
let count = AtomicU32::new(0);
for _ in 0..3 {
std::thread::sleep(std::time::Duration::from_millis(30));
pool.run(64, &|_, _| {
count.fetch_add(1, Ordering::Relaxed);
});
}
assert_eq!(count.load(Ordering::Relaxed), 192);
}
#[test]
fn no_worker_runs_two_tasks_at_once_over_many_small_jobs() {
let pool = Pool::new(6);
let busy: Vec<AtomicBool> = (0..6).map(|_| AtomicBool::new(false)).collect();
for job in 0..20_000usize {
let n = 1 + job % 9;
let hits: Vec<AtomicU32> = (0..n).map(|_| AtomicU32::new(0)).collect();
pool.run(n, &|i, w| {
assert!(!busy[w].swap(true, Ordering::AcqRel), "worker {w} ran two tasks at once");
hits[i].fetch_add(1, Ordering::Relaxed);
busy[w].store(false, Ordering::Release);
});
assert!(hits.iter().all(|h| h.load(Ordering::Relaxed) == 1), "job {job}");
}
}
#[test]
fn a_panic_reaches_the_caller_and_the_pool_survives() {
let pool = Pool::new(4);
let r = catch_unwind(AssertUnwindSafe(|| {
pool.run(100, &|i, _| assert!(i != 57, "task 57"));
}));
assert!(r.is_err());
let count = AtomicU32::new(0);
pool.run(10, &|_, _| {
count.fetch_add(1, Ordering::Relaxed);
});
assert_eq!(count.load(Ordering::Relaxed), 10);
}
}