use std::any::Any;
use std::cell::Cell;
use std::panic::{catch_unwind, resume_unwind, AssertUnwindSafe};
use std::sync::atomic::{AtomicBool, AtomicU64, AtomicUsize, Ordering};
use std::sync::{Arc, Condvar, Mutex};
use std::time::{Duration, Instant};
fn spin_window() -> Duration {
use std::sync::OnceLock;
static US: OnceLock<u64> = OnceLock::new();
Duration::from_micros(*US.get_or_init(|| {
std::env::var("FERROX_CPU_POOL_SPIN_US")
.ok()
.and_then(|v| v.trim().parse::<u64>().ok())
.unwrap_or(100)
}))
}
const PARK_TIMEOUT: Duration = Duration::from_millis(20);
#[derive(Clone, Copy)]
struct Job {
ptr: *const (),
call: unsafe fn(*const (), usize),
n_tasks: usize,
}
impl Default for Job {
fn default() -> Self {
unsafe fn never(_: *const (), _: usize) {}
Self {
ptr: std::ptr::null(),
call: never,
n_tasks: 0,
}
}
}
unsafe fn call_shim<F: Fn(usize) + Sync>(ptr: *const (), task: usize) {
let f = unsafe { &*(ptr as *const F) };
f(task)
}
struct Shared {
epoch: AtomicU64,
job: std::cell::UnsafeCell<Job>,
next: AtomicUsize,
aborted: AtomicBool,
active: AtomicUsize,
shutdown: AtomicBool,
parked: Mutex<usize>,
cv: Condvar,
panic: Mutex<Option<Box<dyn Any + Send + 'static>>>,
}
unsafe impl Sync for Shared {}
unsafe impl Send for Shared {}
impl Shared {
fn wait_for_job(&self, seen: u64) -> u64 {
let deadline = Instant::now() + spin_window();
loop {
let epoch = self.epoch.load(Ordering::Acquire);
if epoch != seen {
return epoch;
}
for _ in 0..64 {
std::hint::spin_loop();
}
if Instant::now() >= deadline {
break;
}
}
let mut parked = self.parked.lock().unwrap_or_else(|e| e.into_inner());
loop {
let epoch = self.epoch.load(Ordering::Acquire);
if epoch != seen {
return epoch;
}
*parked += 1;
let (guard, _) = self
.cv
.wait_timeout(parked, PARK_TIMEOUT)
.unwrap_or_else(|e| e.into_inner());
parked = guard;
*parked -= 1;
}
}
fn drain(&self, job: Job) {
let outcome = catch_unwind(AssertUnwindSafe(|| {
loop {
if self.aborted.load(Ordering::Relaxed) {
break;
}
let task = self.next.fetch_add(1, Ordering::Relaxed);
if task >= job.n_tasks {
break;
}
unsafe { (job.call)(job.ptr, task) };
}
}));
if let Err(payload) = outcome {
self.aborted.store(true, Ordering::Relaxed);
let mut slot = self.panic.lock().unwrap_or_else(|e| e.into_inner());
if slot.is_none() {
*slot = Some(payload);
}
}
}
}
thread_local! {
static IN_REGION: Cell<bool> = const { Cell::new(false) };
}
pub fn in_region() -> bool {
IN_REGION.with(|c| c.get())
}
pub struct CpuPool {
shared: Arc<Shared>,
workers: Vec<std::thread::JoinHandle<()>>,
submit: Mutex<()>,
}
impl CpuPool {
pub fn new(threads: usize) -> Self {
let threads = threads.max(1);
let shared = Arc::new(Shared {
epoch: AtomicU64::new(0),
job: std::cell::UnsafeCell::new(Job::default()),
next: AtomicUsize::new(0),
aborted: AtomicBool::new(false),
active: AtomicUsize::new(0),
shutdown: AtomicBool::new(false),
parked: Mutex::new(0),
cv: Condvar::new(),
panic: Mutex::new(None),
});
let mut workers = Vec::with_capacity(threads - 1);
for idx in 0..threads - 1 {
let shared = Arc::clone(&shared);
let handle = std::thread::Builder::new()
.name(format!("ferrox-cpu-{idx}"))
.spawn(move || worker_loop(&shared))
.expect("ferrox: cannot spawn CPU pool worker");
workers.push(handle);
}
Self {
shared,
workers,
submit: Mutex::new(()),
}
}
pub fn num_threads(&self) -> usize {
self.workers.len() + 1
}
pub fn run<F: Fn(usize) + Sync>(&self, n_tasks: usize, f: &F) -> bool {
if n_tasks == 0 {
return true;
}
if in_region() {
for task in 0..n_tasks {
f(task);
}
return true;
}
let guard = match self.submit.try_lock() {
Ok(guard) => guard,
Err(std::sync::TryLockError::Poisoned(guard)) => guard.into_inner(),
Err(std::sync::TryLockError::WouldBlock) => return false,
};
let job = Job {
ptr: std::ptr::from_ref(f) as *const (),
call: call_shim::<F>,
n_tasks,
};
unsafe { *self.shared.job.get() = job };
self.shared.next.store(0, Ordering::Relaxed);
self.shared.aborted.store(false, Ordering::Relaxed);
self.shared
.active
.store(self.workers.len(), Ordering::Relaxed);
let parked = {
let parked = self.shared.parked.lock().unwrap_or_else(|e| e.into_inner());
self.shared.epoch.fetch_add(1, Ordering::Release);
*parked
};
if parked > 0 {
self.shared.cv.notify_all();
}
IN_REGION.with(|c| c.set(true));
self.shared.drain(job);
IN_REGION.with(|c| c.set(false));
let mut spins = 0u32;
while self.shared.active.load(Ordering::Acquire) != 0 {
spins += 1;
if spins.is_multiple_of(512) {
std::thread::yield_now();
} else {
std::hint::spin_loop();
}
}
let payload = self
.shared
.panic
.lock()
.unwrap_or_else(|e| e.into_inner())
.take();
drop(guard);
if let Some(payload) = payload {
resume_unwind(payload);
}
true
}
}
impl Drop for CpuPool {
fn drop(&mut self) {
{
let _parked = self.shared.parked.lock().unwrap_or_else(|e| e.into_inner());
self.shared.shutdown.store(true, Ordering::Release);
self.shared.epoch.fetch_add(1, Ordering::Release);
}
self.shared.cv.notify_all();
for handle in self.workers.drain(..) {
let _ = handle.join();
}
}
}
fn worker_loop(shared: &Shared) {
crate::threads::set_user_interactive_qos();
let mut seen = 0u64;
loop {
seen = shared.wait_for_job(seen);
if shared.shutdown.load(Ordering::Acquire) {
return;
}
let job = unsafe { *shared.job.get() };
IN_REGION.with(|c| c.set(true));
shared.drain(job);
IN_REGION.with(|c| c.set(false));
shared.active.fetch_sub(1, Ordering::Release);
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::sync::atomic::AtomicU32;
#[test]
fn every_task_runs_exactly_once() {
for threads in [1usize, 2, 4, 8] {
let pool = CpuPool::new(threads);
for n in [1usize, 3, 17, 1000] {
let counts: Vec<AtomicU32> = (0..n).map(|_| AtomicU32::new(0)).collect();
assert!(pool.run(n, &|task: usize| {
counts[task].fetch_add(1, Ordering::Relaxed);
}));
for (task, count) in counts.iter().enumerate() {
assert_eq!(
count.load(Ordering::Relaxed),
1,
"task {task} of {n} on {threads} threads"
);
}
}
}
}
#[test]
fn regions_still_complete_when_every_worker_has_to_be_woken_from_a_park() {
let pool = CpuPool::new(4);
for round in 0..20 {
std::thread::sleep(spin_window() * 3);
let seen = AtomicU32::new(0);
assert!(pool.run(64, &|_| {
seen.fetch_add(1, Ordering::Relaxed);
}));
assert_eq!(seen.load(Ordering::Relaxed), 64, "round {round}");
}
}
#[test]
fn no_worker_touches_the_closure_after_run_returns() {
let pool = CpuPool::new(4);
for _ in 0..50 {
let cells: Vec<AtomicU32> = (0..256).map(|_| AtomicU32::new(0)).collect();
let local = 7u32;
assert!(pool.run(256, &|task: usize| {
cells[task].store(local + task as u32, Ordering::Relaxed);
}));
for (task, cell) in cells.iter().enumerate() {
assert_eq!(cell.load(Ordering::Relaxed), 7 + task as u32);
}
}
}
#[test]
fn a_nested_region_runs_inline_instead_of_deadlocking() {
let pool = CpuPool::new(4);
let inner_total = AtomicU32::new(0);
assert!(pool.run(8, &|_outer: usize| {
assert!(in_region());
assert!(pool.run(5, &|_inner: usize| {
inner_total.fetch_add(1, Ordering::Relaxed);
}));
}));
assert_eq!(inner_total.load(Ordering::Relaxed), 40);
}
#[test]
fn a_concurrent_submitter_is_refused_rather_than_serialized() {
let pool = Arc::new(CpuPool::new(2));
let refused = AtomicU32::new(0);
let started = Arc::new(AtomicBool::new(false));
std::thread::scope(|scope| {
let held = Arc::clone(&pool);
let started_w = Arc::clone(&started);
scope.spawn(move || {
held.run(1, &|_| {
started_w.store(true, Ordering::Release);
std::thread::sleep(Duration::from_millis(150));
});
});
while !started.load(Ordering::Acquire) {
std::hint::spin_loop();
}
if !pool.run(4, &|_| {}) {
refused.fetch_add(1, Ordering::Relaxed);
}
});
assert_eq!(
refused.load(Ordering::Relaxed),
1,
"the pool must report a busy region rather than block"
);
}
#[test]
fn a_panicking_task_is_re_raised_on_the_submitter() {
let pool = CpuPool::new(4);
let outcome = catch_unwind(AssertUnwindSafe(|| {
pool.run(64, &|task: usize| {
if task == 33 {
panic!("ferrox test panic");
}
});
}));
assert!(outcome.is_err(), "the panic must not be swallowed");
let ran = AtomicU32::new(0);
assert!(pool.run(10, &|_| {
ran.fetch_add(1, Ordering::Relaxed);
}));
assert_eq!(ran.load(Ordering::Relaxed), 10);
}
#[test]
fn dropping_the_pool_joins_every_worker() {
for threads in [1usize, 2, 6] {
let pool = CpuPool::new(threads);
assert!(pool.run(4, &|_| {}));
std::thread::sleep(spin_window() * 2);
drop(pool);
}
}
#[test]
fn a_single_thread_pool_runs_everything_on_the_submitter() {
let pool = CpuPool::new(1);
assert_eq!(pool.num_threads(), 1);
let here = std::thread::current().id();
let same = AtomicU32::new(0);
assert!(pool.run(32, &|_| {
if std::thread::current().id() == here {
same.fetch_add(1, Ordering::Relaxed);
}
}));
assert_eq!(same.load(Ordering::Relaxed), 32);
}
}