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;
use crate::snapshot::SnapshotPublisher;
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,
serialized: bool,
draining: 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> Sink<T> {
fn drain_queues(&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);
}
}
}
impl<T: Send + 'static> Drain for Sink<T> {
fn drain(&self) {
if !self.serialized {
self.drain_queues();
return;
}
if self.draining.swap(true, Ordering::Acquire) {
return;
}
loop {
self.drain_queues();
self.draining.store(false, Ordering::Release);
fence(Ordering::SeqCst);
if !self.scheduled.load(Ordering::SeqCst) {
return;
}
if self.draining.swap(true, Ordering::Acquire) {
return;
}
}
}
}
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);
match thread::Builder::new()
.name("truce-task-pool".into())
.spawn(move || worker_loop(&shared))
{
Ok(handle) => workers.push(handle.thread().clone()),
Err(e) => {
eprintln!("[truce] task-pool worker spawn failed: {e}");
break;
}
}
}
Pool { shared, workers }
})
}
pub fn warm_pool() {
let _ = pool();
}
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.workers.is_empty() {
return false;
}
if pool.shared.injector.push(sink).is_err() {
return false;
}
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::with_mode(run, false)
}
pub fn new_serialized(run: impl Fn(T) + Send + Sync + 'static) -> Self {
Self::with_mode(run, true)
}
fn with_mode(run: impl Fn(T) + Send + Sync + 'static, serialized: bool) -> Self {
Self {
sink: Arc::new(Sink {
queue: ArrayQueue::new(TASK_QUEUE_PREALLOC),
coalesced: ArrayQueue::new(1),
scheduled: AtomicBool::new(false),
serialized,
draining: 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);
}
}
}
}
type ErasedLane = Arc<dyn std::any::Any + Send + Sync>;
#[derive(Clone)]
pub struct AnyTaskSpawner(Arc<[ErasedLane]>);
impl AnyTaskSpawner {
#[must_use]
pub fn new<T: Send + 'static>(spawner: &TaskSpawner<T>) -> Self {
Self(Arc::from(vec![Arc::new(spawner.clone()) as ErasedLane]))
}
#[must_use]
pub fn from_lanes(lanes: Vec<ErasedLane>) -> Self {
Self(Arc::from(lanes))
}
#[must_use]
pub fn downcast<T: Send + 'static>(&self) -> Option<TaskSpawner<T>> {
self.0
.iter()
.find_map(|lane| lane.downcast_ref::<TaskSpawner<T>>().cloned())
}
}
#[derive(Default)]
pub struct TaskSpawnerBundle(Vec<ErasedLane>);
impl TaskSpawnerBundle {
#[must_use]
pub fn new() -> Self {
Self(Vec::new())
}
pub fn push<T: Send + 'static>(&mut self, spawner: TaskSpawner<T>) {
self.0.push(Arc::new(spawner) as ErasedLane);
}
#[must_use]
pub fn into_any(self) -> Option<AnyTaskSpawner> {
if self.0.is_empty() {
None
} else {
Some(AnyTaskSpawner::from_lanes(self.0))
}
}
}
pub struct InitContext {
tasks: Option<AnyTaskSpawner>,
snapshot: Option<SnapshotPublisher>,
}
impl InitContext {
#[must_use]
pub fn new(tasks: Option<AnyTaskSpawner>) -> Self {
Self {
tasks,
snapshot: None,
}
}
#[must_use]
pub fn with_snapshot(mut self, snapshot: SnapshotPublisher) -> Self {
self.snapshot = Some(snapshot);
self
}
#[must_use]
pub fn tasks<T: Send + 'static>(&self) -> Option<TaskSpawner<T>> {
self.tasks.as_ref().and_then(AnyTaskSpawner::downcast::<T>)
}
#[must_use]
pub fn snapshot_publisher(&self) -> Option<SnapshotPublisher> {
self.snapshot.clone()
}
}
#[cfg(test)]
mod tests {
#![allow(clippy::cast_possible_truncation)]
use super::*;
use crate::snapshot::{SnapshotPublisher, SnapshotSlot};
use std::sync::atomic::AtomicU32;
use std::sync::{Condvar, Mutex};
use std::time::Instant;
#[test]
fn init_context_exposes_snapshot_publisher() {
let slot = SnapshotSlot::new();
let cx = InitContext::new(None).with_snapshot(SnapshotPublisher::new(&slot));
cx.snapshot_publisher()
.expect("publisher present")
.publish(vec![1, 2, 3]);
assert_eq!(slot.read(), Some(vec![1, 2, 3]));
assert!(InitContext::new(None).snapshot_publisher().is_none());
}
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 warm_pool_starts_workers_and_is_idempotent() {
warm_pool();
warm_pool();
assert!(
!pool().workers.is_empty(),
"warming spawns at least one worker"
);
}
#[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 serialized_runs_one_at_a_time_and_drops_nothing() {
const N: u32 = 64;
let in_flight = Arc::new(AtomicU32::new(0));
let peak = Arc::new(AtomicU32::new(0));
let latch = Arc::new(Latch::default());
let (inf, pk, l) = (
Arc::clone(&in_flight),
Arc::clone(&peak),
Arc::clone(&latch),
);
let spawner = TaskSpawner::<u32>::new_serialized(move |_| {
let now = inf.fetch_add(1, Ordering::AcqRel) + 1;
pk.fetch_max(now, Ordering::AcqRel);
thread::sleep(Duration::from_millis(1));
inf.fetch_sub(1, Ordering::AcqRel);
l.bump();
});
for n in 0..N {
while spawner.try_spawn(n).is_err() {
thread::sleep(Duration::from_millis(1));
}
}
latch.wait_for(N);
assert_eq!(
peak.load(Ordering::Acquire),
1,
"serialized: at most one handler in flight at a time"
);
}
#[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"
);
});
}
#[test]
fn serialized_drain_never_strands_a_task() {
loom::model(|| {
let scheduled = Arc::new(AtomicBool::new(true));
let item = Arc::new(AtomicBool::new(true));
let draining = Arc::new(AtomicBool::new(false));
let worker =
|scheduled: Arc<AtomicBool>, item: Arc<AtomicBool>, draining: Arc<AtomicBool>| {
if draining.swap(true, Ordering::Acquire) {
return; }
for _ in 0..2 {
scheduled.store(false, Ordering::SeqCst);
fence(Ordering::SeqCst);
let _ = item.swap(false, Ordering::AcqRel);
draining.store(false, Ordering::Release);
fence(Ordering::SeqCst);
if !scheduled.load(Ordering::SeqCst) {
return;
}
if draining.swap(true, Ordering::Acquire) {
return;
}
}
};
let (s1, i1, d1) = (scheduled.clone(), item.clone(), draining.clone());
let w1 = thread::spawn(move || worker(s1, i1, d1));
let (s2, i2, d2) = (scheduled.clone(), item.clone(), draining.clone());
let w2 = thread::spawn(move || worker(s2, i2, d2));
item.store(true, Ordering::Release);
fence(Ordering::SeqCst);
let _ = scheduled.swap(true, Ordering::SeqCst);
w1.join().unwrap();
w2.join().unwrap();
let pending = item.load(Ordering::SeqCst);
let is_scheduled = scheduled.load(Ordering::SeqCst);
assert!(
!pending || is_scheduled,
"serialized task stranded: pending with scheduled == false"
);
});
}
}