use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering, fence};
use std::sync::{Arc, OnceLock};
use std::thread::{self, Thread};
use std::time::Duration;
use crossbeam_queue::ArrayQueue;
pub const TASK_QUEUE_PREALLOC: usize = 256;
const INJECTOR_CAP: usize = 4096;
const PARK_TIMEOUT: Duration = Duration::from_secs(1);
trait Drain: Send + Sync {
fn drain(&self);
}
struct Sink<T: Send + 'static> {
queue: ArrayQueue<T>,
coalesced: ArrayQueue<T>,
scheduled: AtomicBool,
run: Box<dyn Fn(T) + Send + Sync>,
}
impl<T: Send + 'static> Sink<T> {
fn run_one(&self, task: T) {
let run = &self.run;
let _ = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| run(task)));
}
}
impl<T: Send + 'static> Drain for Sink<T> {
fn drain(&self) {
self.scheduled.store(false, Ordering::SeqCst);
fence(Ordering::SeqCst);
if let Some(task) = self.coalesced.pop() {
self.run_one(task);
}
while let Some(task) = self.queue.pop() {
self.run_one(task);
}
}
}
struct Shared {
injector: ArrayQueue<Arc<dyn Drain>>,
next: AtomicUsize,
}
struct Pool {
shared: Arc<Shared>,
workers: Vec<Thread>,
}
static POOL: OnceLock<Pool> = OnceLock::new();
fn pool() -> &'static Pool {
POOL.get_or_init(|| {
let shared = Arc::new(Shared {
injector: ArrayQueue::new(INJECTOR_CAP),
next: AtomicUsize::new(0),
});
let n = thread::available_parallelism().map_or(1, |p| p.get().saturating_sub(1).max(1));
let mut workers = Vec::with_capacity(n);
for _ in 0..n {
let shared = Arc::clone(&shared);
let handle = thread::Builder::new()
.name("truce-task-pool".into())
.spawn(move || worker_loop(&shared))
.expect("spawn truce task-pool worker");
workers.push(handle.thread().clone());
}
Pool { shared, workers }
})
}
fn worker_loop(shared: &Shared) -> ! {
loop {
while let Some(sink) = shared.injector.pop() {
sink.drain();
}
thread::park_timeout(PARK_TIMEOUT);
}
}
fn schedule(sink: Arc<dyn Drain>) -> bool {
let pool = pool();
if pool.shared.injector.push(sink).is_err() {
return false;
}
if !pool.workers.is_empty() {
let i = pool.shared.next.fetch_add(1, Ordering::Relaxed) % pool.workers.len();
pool.workers[i].unpark();
}
true
}
pub struct TaskSpawner<T: Send + 'static> {
sink: Arc<Sink<T>>,
}
impl<T: Send + 'static> Clone for TaskSpawner<T> {
fn clone(&self) -> Self {
Self {
sink: Arc::clone(&self.sink),
}
}
}
impl<T: Send + 'static> TaskSpawner<T> {
pub fn new(run: impl Fn(T) + Send + Sync + 'static) -> Self {
Self {
sink: Arc::new(Sink {
queue: ArrayQueue::new(TASK_QUEUE_PREALLOC),
coalesced: ArrayQueue::new(1),
scheduled: AtomicBool::new(false),
run: Box::new(run),
}),
}
}
pub fn try_spawn(&self, task: T) -> Result<(), T> {
self.sink.queue.push(task)?;
self.arm();
Ok(())
}
pub fn spawn_coalescing(&self, task: T) {
let _ = self.sink.coalesced.force_push(task);
self.arm();
}
fn arm(&self) {
fence(Ordering::SeqCst);
if !self.sink.scheduled.swap(true, Ordering::SeqCst) {
let sink: Arc<dyn Drain> = Arc::clone(&self.sink) as Arc<dyn Drain>;
if !schedule(sink) {
self.sink.scheduled.store(false, Ordering::SeqCst);
}
}
}
}
#[derive(Clone)]
pub struct AnyTaskSpawner(Arc<dyn std::any::Any + Send + Sync>);
impl AnyTaskSpawner {
#[must_use]
pub fn new<T: Send + 'static>(spawner: &TaskSpawner<T>) -> Self {
Self(Arc::new(spawner.clone()) as Arc<dyn std::any::Any + Send + Sync>)
}
#[must_use]
pub fn downcast<T: Send + 'static>(&self) -> Option<TaskSpawner<T>> {
self.0.downcast_ref::<TaskSpawner<T>>().cloned()
}
}
pub struct InitContext {
tasks: Option<AnyTaskSpawner>,
}
impl InitContext {
#[must_use]
pub fn new(tasks: Option<AnyTaskSpawner>) -> Self {
Self { tasks }
}
#[must_use]
pub fn tasks<T: Send + 'static>(&self) -> Option<TaskSpawner<T>> {
self.tasks.as_ref().and_then(AnyTaskSpawner::downcast::<T>)
}
}
#[cfg(test)]
mod tests {
#![allow(clippy::cast_possible_truncation)]
use super::*;
use std::sync::atomic::AtomicU32;
use std::sync::{Condvar, Mutex};
use std::time::Instant;
fn wait_until(deadline: Duration, mut done: impl FnMut() -> bool) -> bool {
let start = Instant::now();
while start.elapsed() < deadline {
if done() {
return true;
}
thread::sleep(Duration::from_millis(1));
}
done()
}
#[derive(Default)]
struct Latch {
ran: Mutex<u32>,
woke: Condvar,
}
impl Latch {
fn bump(&self) {
*self.ran.lock().unwrap() += 1;
self.woke.notify_all();
}
fn wait_for(&self, target: u32) {
let mut ran = self.ran.lock().unwrap();
while *ran < target {
ran = self.woke.wait(ran).unwrap();
}
}
}
#[test]
fn runs_scheduled_tasks_off_thread() {
let latch = Arc::new(Latch::default());
let sum = Arc::new(AtomicU32::new(0));
let (l, s) = (Arc::clone(&latch), Arc::clone(&sum));
let spawner = TaskSpawner::<u32>::new(move |n| {
s.fetch_add(n, Ordering::Relaxed);
l.bump();
});
for n in 1..=10 {
spawner.try_spawn(n).expect("queue has room");
}
latch.wait_for(10);
assert_eq!(sum.load(Ordering::Relaxed), 55, "all ten tasks ran");
}
#[test]
fn full_queue_returns_the_task() {
let gate = Arc::new(AtomicBool::new(false));
let g = Arc::clone(&gate);
let spawner = TaskSpawner::<u32>::new(move |_| {
while !g.load(Ordering::Acquire) {
thread::sleep(Duration::from_millis(1));
}
});
let mut rejected = 0u32;
for n in 0..(TASK_QUEUE_PREALLOC as u32 + 64) {
if spawner.try_spawn(n).is_err() {
rejected += 1;
}
}
assert!(rejected > 0, "a full inbound queue rejects further tasks");
gate.store(true, Ordering::Release);
}
#[test]
fn panicking_task_does_not_kill_the_worker() {
let latch = Arc::new(Latch::default());
let l = Arc::clone(&latch);
let spawner = TaskSpawner::<bool>::new(move |should_panic| {
assert!(!should_panic, "intentional panic, caught by the pool");
l.bump();
});
spawner.try_spawn(true).expect("queue has room"); spawner.try_spawn(false).expect("queue has room"); latch.wait_for(1);
}
#[test]
fn coalescing_never_rejects() {
let last = Arc::new(AtomicU32::new(0));
let l = Arc::clone(&last);
let spawner = TaskSpawner::<u32>::new(move |n| {
l.store(n, Ordering::Relaxed);
});
for n in 0..(TASK_QUEUE_PREALLOC as u32 * 4) {
spawner.spawn_coalescing(n); }
let target = TASK_QUEUE_PREALLOC as u32 * 4 - 1;
assert!(
wait_until(Duration::from_secs(2), || last.load(Ordering::Relaxed)
== target),
"the newest task always runs"
);
}
#[test]
fn coalescing_collapses_to_the_newest() {
let runs = Arc::new(AtomicU32::new(0));
let last = Arc::new(AtomicU32::new(0));
let (r, l) = (Arc::clone(&runs), Arc::clone(&last));
let spawner = TaskSpawner::<u32>::new(move |n| {
r.fetch_add(1, Ordering::Relaxed);
l.store(n, Ordering::Relaxed);
});
for n in 1..=1000 {
let _ = spawner.sink.coalesced.force_push(n);
}
spawner.sink.drain();
assert_eq!(runs.load(Ordering::Relaxed), 1, "the burst ran once");
assert_eq!(last.load(Ordering::Relaxed), 1000, "and it was the newest");
}
}
#[cfg(all(test, feature = "loom"))]
mod loom_tests {
use loom::sync::Arc;
use loom::sync::atomic::{AtomicBool, Ordering, fence};
use loom::thread;
#[test]
fn schedule_drain_never_strands_a_task() {
loom::model(|| {
let flag = Arc::new(AtomicBool::new(true));
let item = Arc::new(AtomicBool::new(false));
let (f, i) = (flag.clone(), item.clone());
let worker = thread::spawn(move || {
f.store(false, Ordering::SeqCst);
fence(Ordering::SeqCst);
if i.load(Ordering::Acquire) {
i.store(false, Ordering::Release); }
});
item.store(true, Ordering::Release);
fence(Ordering::SeqCst);
let was_scheduled = flag.swap(true, Ordering::SeqCst);
let _ = was_scheduled;
worker.join().unwrap();
let pending = item.load(Ordering::SeqCst);
let scheduled = flag.load(Ordering::SeqCst);
assert!(
!pending || scheduled,
"task stranded: pending with scheduled == false"
);
});
}
}