use std::time::Duration;
#[cfg(loom)]
use loom::sync::{
Condvar, Mutex, MutexGuard,
atomic::{AtomicBool, AtomicUsize, Ordering},
};
#[cfg(not(loom))]
use std::sync::{
Condvar, Mutex, MutexGuard,
atomic::{AtomicBool, AtomicUsize, Ordering},
};
use crossbeam_deque::{Injector, Steal, Worker};
const CANCELLATION_POLL_INTERVAL: Duration = Duration::from_millis(10);
const SLEEPER_BITS: u32 = usize::BITS - crate::MAX_WORKERS.leading_zeros();
const SLEEPER_MASK: usize = (1 << SLEEPER_BITS) - 1;
const WAKE_EPOCH_STEP: usize = 1 << SLEEPER_BITS;
pub(crate) struct Scheduler<T> {
injector: Injector<T>,
}
impl<T> Scheduler<T> {
pub(crate) fn new() -> Self {
Self {
injector: Injector::new(),
}
}
pub(crate) fn push(&self, task: T) {
self.injector.push(task);
}
pub(crate) fn worker(&self) -> Worker<T> {
Worker::new_fifo()
}
pub(crate) fn is_empty(&self) -> bool {
self.injector.is_empty()
}
pub(crate) fn steal_into(&self, worker: &Worker<T>) -> Option<T> {
loop {
match self.injector.steal_batch_and_pop(worker) {
Steal::Success(task) => return Some(task),
Steal::Empty => return None,
Steal::Retry => continue,
}
}
}
}
pub(crate) struct Coordinator {
pending: CacheLine<AtomicUsize>,
active_workers: CacheLine<AtomicUsize>,
wake_state: AtomicUsize,
lone: AtomicBool,
wake_lock: Mutex<()>,
wake: Condvar,
}
#[cfg_attr(any(target_arch = "aarch64", target_arch = "x86_64"), repr(align(128)))]
#[cfg_attr(
not(any(target_arch = "aarch64", target_arch = "x86_64")),
repr(align(64))
)]
pub(crate) struct CacheLine<T>(pub(crate) T);
impl<T> std::ops::Deref for CacheLine<T> {
type Target = T;
fn deref(&self) -> &T {
&self.0
}
}
impl Coordinator {
pub(crate) fn new() -> Self {
Self {
pending: CacheLine(AtomicUsize::new(0)),
active_workers: CacheLine(AtomicUsize::new(0)),
wake_state: AtomicUsize::new(0),
lone: AtomicBool::new(true),
wake_lock: Mutex::new(()),
wake: Condvar::new(),
}
}
pub(crate) fn begin_task(&self) {
self.pending.fetch_add(1, Ordering::AcqRel);
}
pub(crate) fn schedule(&self, publish: impl FnOnce()) {
self.begin_task();
publish();
if self.lone.load(Ordering::Acquire) {
return;
}
self.announce_one();
}
pub(crate) fn widen(&self) {
self.lone.store(false, Ordering::Release);
}
pub(crate) fn claim_task(&self) -> TaskGuard<'_> {
TaskGuard { coordinator: self }
}
pub(crate) fn pending(&self) -> usize {
self.pending.load(Ordering::Acquire)
}
pub(crate) fn wake_waiters(&self) {
let guard = lock(&self.wake_lock);
self.wake.notify_all();
drop(guard);
}
pub(crate) fn announce_batch_handoff(&self) {
self.announce_one();
}
fn announce_one(&self) {
let previous = self.wake_state.fetch_add(WAKE_EPOCH_STEP, Ordering::SeqCst);
if previous & SLEEPER_MASK == 0 {
return;
}
let guard = lock(&self.wake_lock);
self.wake.notify_one();
drop(guard);
}
pub(crate) fn wait_for_task<T>(
&self,
should_stop: impl Fn() -> bool,
mut try_take: impl FnMut() -> Option<T>,
) -> Option<T> {
loop {
if should_stop() {
return None;
}
let observed_epoch = self.wake_state.load(Ordering::SeqCst) & !SLEEPER_MASK;
if let Some(task) = try_take() {
return Some(task);
}
let guard = lock(&self.wake_lock);
if self.pending.load(Ordering::Acquire) == 0 || should_stop() {
return None;
}
let previous = self.wake_state.fetch_add(1, Ordering::SeqCst);
debug_assert_ne!(previous & SLEEPER_MASK, SLEEPER_MASK);
if previous & !SLEEPER_MASK != observed_epoch {
self.wake_state.fetch_sub(1, Ordering::SeqCst);
drop(guard);
continue;
}
let (guard, _) = self
.wake
.wait_timeout(guard, CANCELLATION_POLL_INTERVAL)
.unwrap_or_else(|poisoned| poisoned.into_inner());
self.wake_state.fetch_sub(1, Ordering::SeqCst);
drop(guard);
}
}
pub(crate) fn claim_caller_slot(&self) -> WorkerSlot<'_> {
self.active_workers.fetch_add(1, Ordering::AcqRel);
WorkerSlot { coordinator: self }
}
pub(crate) fn claim_worker_slot(&self, budget: usize) -> Option<WorkerSlot<'_>> {
let mut active = self.active_workers.load(Ordering::Acquire);
loop {
if active >= budget || self.pending() <= active {
return None;
}
match self.active_workers.compare_exchange_weak(
active,
active + 1,
Ordering::AcqRel,
Ordering::Acquire,
) {
Ok(_) => return Some(WorkerSlot { coordinator: self }),
Err(observed) => active = observed,
}
}
}
#[cfg(loom)]
pub(crate) fn active_workers(&self) -> usize {
self.active_workers.load(Ordering::Acquire)
}
#[cfg(loom)]
pub(crate) fn sleepers(&self) -> usize {
self.wake_state.load(Ordering::SeqCst) & SLEEPER_MASK
}
fn finish_task(&self) {
if self.pending.fetch_sub(1, Ordering::AcqRel) == 1 {
self.wake_waiters();
}
}
}
pub(crate) struct TaskGuard<'a> {
coordinator: &'a Coordinator,
}
impl Drop for TaskGuard<'_> {
fn drop(&mut self) {
self.coordinator.finish_task();
}
}
pub(crate) struct WorkerSlot<'a> {
coordinator: &'a Coordinator,
}
impl Drop for WorkerSlot<'_> {
fn drop(&mut self) {
self.coordinator
.active_workers
.fetch_sub(1, Ordering::AcqRel);
}
}
fn lock<T>(mutex: &Mutex<T>) -> MutexGuard<'_, T> {
mutex
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner())
}
#[cfg(all(test, not(loom)))]
mod tests {
use super::Scheduler;
#[test]
fn injector_distributes_tasks_to_worker_local_queues() {
let scheduler = Scheduler::new();
scheduler.push(1);
scheduler.push(2);
let worker = scheduler.worker();
let first = scheduler.steal_into(&worker).expect("first task");
let second = worker.pop().unwrap_or_else(|| {
scheduler
.steal_into(&worker)
.expect("second task after local batch")
});
assert_ne!(first, second);
assert!(scheduler.steal_into(&worker).is_none());
}
}
#[cfg(loom)]
mod loom_models {
use std::{
collections::VecDeque,
panic::{AssertUnwindSafe, catch_unwind},
sync::atomic::{AtomicUsize, Ordering},
};
use loom::{
sync::{Arc, Mutex},
thread,
};
use super::Coordinator;
fn report(model: &str, executions: &AtomicUsize) {
eprintln!(
"loom explored {} executions of {model}",
executions.load(Ordering::Relaxed)
);
}
fn pop(queue: &Mutex<VecDeque<usize>>) -> Option<usize> {
queue
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner())
.pop_front()
}
fn push(queue: &Mutex<VecDeque<usize>>, task: usize) {
queue
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner())
.push_back(task);
}
#[test]
fn a_queued_task_reaches_a_parking_worker() {
static EXECUTIONS: AtomicUsize = AtomicUsize::new(0);
loom::model(|| {
EXECUTIONS.fetch_add(1, Ordering::Relaxed);
let coordinator = Arc::new(Coordinator::new());
let queue = Arc::new(Mutex::new(VecDeque::new()));
coordinator.widen();
coordinator.begin_task();
let producer_coordinator = Arc::clone(&coordinator);
let producer_queue = Arc::clone(&queue);
let producer = thread::spawn(move || {
let _root = producer_coordinator.claim_task();
producer_coordinator.schedule(|| push(&producer_queue, 7));
});
let consumer_coordinator = Arc::clone(&coordinator);
let consumer_queue = Arc::clone(&queue);
let consumer = thread::spawn(move || {
let task = consumer_coordinator.wait_for_task(|| false, || pop(&consumer_queue));
assert_eq!(task, Some(7), "the queued task must reach the worker");
drop(consumer_coordinator.claim_task());
assert_eq!(
consumer_coordinator.wait_for_task(|| false, || pop(&consumer_queue)),
None,
"a drained walk must let its workers finish"
);
});
producer.join().expect("producer joins");
consumer.join().expect("consumer joins");
assert_eq!(coordinator.pending(), 0);
});
report("a_queued_task_reaches_a_parking_worker", &EXECUTIONS);
}
#[test]
fn zero_sleeper_skip_races_with_parker_registration() {
static EXECUTIONS: AtomicUsize = AtomicUsize::new(0);
loom::model(|| {
EXECUTIONS.fetch_add(1, Ordering::Relaxed);
let coordinator = Arc::new(Coordinator::new());
let queue = Arc::new(Mutex::new(VecDeque::new()));
coordinator.begin_task();
let producer_coordinator = Arc::clone(&coordinator);
let producer_queue = Arc::clone(&queue);
let producer = thread::spawn(move || {
push(&producer_queue, 7);
producer_coordinator.announce_batch_handoff();
});
let worker_coordinator = Arc::clone(&coordinator);
let worker_queue = Arc::clone(&queue);
let worker = thread::spawn(move || {
let task = worker_coordinator.wait_for_task(|| false, || pop(&worker_queue));
assert_eq!(task, Some(7), "the skipped notify must not lose work");
drop(worker_coordinator.claim_task());
});
producer.join().expect("producer joins");
worker.join().expect("worker joins");
assert_eq!(coordinator.pending(), 0);
assert_eq!(coordinator.sleepers(), 0);
});
report(
"zero_sleeper_skip_races_with_parker_registration",
&EXECUTIONS,
);
}
#[test]
fn widening_does_not_lose_a_concurrent_scheduled_task() {
static EXECUTIONS: AtomicUsize = AtomicUsize::new(0);
loom::model(|| {
EXECUTIONS.fetch_add(1, Ordering::Relaxed);
let coordinator = Arc::new(Coordinator::new());
let queue = Arc::new(Mutex::new(VecDeque::new()));
coordinator.begin_task();
let producer_coordinator = Arc::clone(&coordinator);
let producer_queue = Arc::clone(&queue);
let producer = thread::spawn(move || {
let _root = producer_coordinator.claim_task();
producer_coordinator.schedule(|| push(&producer_queue, 7));
});
let helper_coordinator = Arc::clone(&coordinator);
let helper_queue = Arc::clone(&queue);
let helper = thread::spawn(move || {
helper_coordinator.widen();
let task = helper_coordinator.wait_for_task(|| false, || pop(&helper_queue));
assert_eq!(task, Some(7), "widening must not strand published work");
drop(helper_coordinator.claim_task());
});
producer.join().expect("producer joins");
helper.join().expect("helper joins");
assert_eq!(coordinator.pending(), 0);
});
report(
"widening_does_not_lose_a_concurrent_scheduled_task",
&EXECUTIONS,
);
}
#[test]
fn a_batch_handoff_reaches_a_parking_worker() {
static EXECUTIONS: AtomicUsize = AtomicUsize::new(0);
loom::model(|| {
EXECUTIONS.fetch_add(1, Ordering::Relaxed);
let coordinator = Arc::new(Coordinator::new());
let queue = Arc::new(Mutex::new(VecDeque::new()));
coordinator.widen();
coordinator.begin_task();
let thief_coordinator = Arc::clone(&coordinator);
let thief_queue = Arc::clone(&queue);
let thief = thread::spawn(move || {
push(&thief_queue, 7);
thief_coordinator.announce_batch_handoff();
});
let worker_coordinator = Arc::clone(&coordinator);
let worker_queue = Arc::clone(&queue);
let worker = thread::spawn(move || {
let task = worker_coordinator.wait_for_task(|| false, || pop(&worker_queue));
assert_eq!(task, Some(7), "the batch sibling must reach a worker");
drop(worker_coordinator.claim_task());
});
thief.join().expect("batch thief joins");
worker.join().expect("worker joins");
assert_eq!(coordinator.pending(), 0);
});
report("a_batch_handoff_reaches_a_parking_worker", &EXECUTIONS);
}
#[test]
fn a_panicking_worker_releases_its_task() {
static EXECUTIONS: AtomicUsize = AtomicUsize::new(0);
let previous = std::sync::Arc::new(std::panic::take_hook());
let filtered = std::sync::Arc::clone(&previous);
std::panic::set_hook(Box::new(move |info| {
if !info.to_string().contains("injected task panic") {
filtered(info);
}
}));
loom::model(|| {
EXECUTIONS.fetch_add(1, Ordering::Relaxed);
let coordinator = Arc::new(Coordinator::new());
coordinator.begin_task();
let panicking_coordinator = Arc::clone(&coordinator);
let panicking = thread::spawn(move || {
let outcome = catch_unwind(AssertUnwindSafe(|| {
let _task = panicking_coordinator.claim_task();
panic!("injected task panic");
}));
assert!(outcome.is_err(), "the injected panic must be captured");
});
let sibling_coordinator = Arc::clone(&coordinator);
let sibling = thread::spawn(move || {
assert_eq!(
sibling_coordinator.wait_for_task(|| false, || None::<usize>),
None,
"the sibling must not park forever after a panic"
);
});
panicking.join().expect("panicking worker joins");
sibling.join().expect("sibling joins");
assert_eq!(coordinator.pending(), 0);
});
let _ = std::panic::take_hook();
std::panic::set_hook(Box::new(move |info| previous(info)));
report("a_panicking_worker_releases_its_task", &EXECUTIONS);
}
#[test]
fn cancellation_releases_a_parking_worker() {
static EXECUTIONS: AtomicUsize = AtomicUsize::new(0);
loom::model(|| {
EXECUTIONS.fetch_add(1, Ordering::Relaxed);
let coordinator = Arc::new(Coordinator::new());
let cancelled = Arc::new(Mutex::new(false));
coordinator.begin_task();
let canceller_coordinator = Arc::clone(&coordinator);
let canceller_flag = Arc::clone(&cancelled);
let canceller = thread::spawn(move || {
*canceller_flag
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner()) = true;
canceller_coordinator.wake_waiters();
});
let worker_coordinator = Arc::clone(&coordinator);
let worker_flag = Arc::clone(&cancelled);
let worker = thread::spawn(move || {
let stopped = || {
*worker_flag
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner())
};
assert_eq!(
worker_coordinator.wait_for_task(stopped, || None::<usize>),
None,
"a cancelled walk must release its parked workers"
);
});
canceller.join().expect("canceller joins");
worker.join().expect("worker joins");
drop(coordinator.claim_task());
});
report("cancellation_releases_a_parking_worker", &EXECUTIONS);
}
#[test]
fn concurrent_growth_stays_within_the_thread_budget() {
static EXECUTIONS: AtomicUsize = AtomicUsize::new(0);
loom::model(|| {
EXECUTIONS.fetch_add(1, Ordering::Relaxed);
let coordinator = Arc::new(Coordinator::new());
let caller = coordinator.claim_caller_slot();
for _ in 0..3 {
coordinator.begin_task();
}
let contender_coordinator = Arc::clone(&coordinator);
let contender = thread::spawn(move || {
let slot = contender_coordinator.claim_worker_slot(2);
assert!(contender_coordinator.active_workers() <= 2);
drop(slot);
});
let slot = coordinator.claim_worker_slot(2);
assert!(coordinator.active_workers() <= 2);
drop(slot);
contender.join().expect("contender joins");
drop(caller);
assert_eq!(coordinator.active_workers(), 0);
for _ in 0..3 {
drop(coordinator.claim_task());
}
assert_eq!(coordinator.pending(), 0);
});
report(
"concurrent_growth_stays_within_the_thread_budget",
&EXECUTIONS,
);
}
}