#![allow(unsafe_code)]
use crate::cell::UnsafeCell;
use std::task::Waker;
#[cfg(loom)]
use loom::sync::atomic::AtomicUsize;
#[cfg(not(loom))]
use std::sync::atomic::AtomicUsize;
use std::sync::atomic::Ordering::{AcqRel, Acquire, Release};
const WAITING: usize = 0;
const REGISTERING: usize = 0b01;
const WAKING: usize = 0b10;
pub struct WakerSlot {
state: AtomicUsize,
waker: UnsafeCell<Option<Waker>>,
}
unsafe impl Send for WakerSlot {}
unsafe impl Sync for WakerSlot {}
impl WakerSlot {
pub fn new() -> Self {
Self { state: AtomicUsize::new(WAITING), waker: UnsafeCell::new(None) }
}
pub fn register(&self, waker: &Waker) {
match self.state.compare_exchange(WAITING, REGISTERING, Acquire, Acquire) {
Ok(_) => {}
Err(WAKING) => {
waker.wake_by_ref();
return;
}
Err(_actual) => {
debug_assert!(
false,
"concurrent registration: a waker slot has exactly one registrar"
);
waker.wake_by_ref();
return;
}
}
let previous = unsafe {
self.waker.with_mut(|slot| {
let previous = (*slot).take();
match previous {
Some(old) if old.will_wake(waker) => {
*slot = Some(old);
None
}
other => {
*slot = Some(waker.clone());
other
}
}
})
};
match self.state.compare_exchange(REGISTERING, WAITING, AcqRel, Acquire) {
Ok(_) => drop(previous),
Err(actual) => {
debug_assert_eq!(actual, REGISTERING | WAKING);
let pending = unsafe { self.waker.with_mut(|slot| (*slot).take()) };
self.state.swap(WAITING, AcqRel);
drop(previous);
if let Some(pending) = pending {
pending.wake();
}
}
}
}
pub fn take(&self) -> Option<Waker> {
match self.state.fetch_or(WAKING, AcqRel) {
WAITING => {
let waker = unsafe { self.waker.with_mut(|slot| (*slot).take()) };
self.state.fetch_and(!WAKING, Release);
waker
}
actual => {
debug_assert!(
actual == REGISTERING || actual == WAKING || actual == REGISTERING | WAKING
);
None
}
}
}
#[inline]
pub fn wake(&self) {
if let Some(waker) = self.take() {
waker.wake();
}
}
}
impl Default for WakerSlot {
fn default() -> Self {
Self::new()
}
}
impl std::fmt::Debug for WakerSlot {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("WakerSlot").finish_non_exhaustive()
}
}
#[cfg(all(test, not(loom)))]
mod tests {
use super::*;
use std::sync::Arc;
use std::sync::atomic::{AtomicUsize, Ordering};
use std::task::Wake;
struct Counter(AtomicUsize);
impl Counter {
fn waker() -> (Arc<Self>, Waker) {
let counter = Arc::new(Self(AtomicUsize::new(0)));
(counter.clone(), Waker::from(counter))
}
fn count(&self) -> usize {
self.0.load(Ordering::Relaxed)
}
}
impl Wake for Counter {
fn wake(self: Arc<Self>) {
self.0.fetch_add(1, Ordering::Relaxed);
}
fn wake_by_ref(self: &Arc<Self>) {
self.0.fetch_add(1, Ordering::Relaxed);
}
}
#[test]
fn a_registered_waker_is_woken_once_and_then_the_slot_is_empty() {
let slot = WakerSlot::new();
let (counter, waker) = Counter::waker();
slot.register(&waker);
slot.wake();
assert_eq!(counter.count(), 1);
slot.wake();
assert_eq!(counter.count(), 1);
}
#[test]
fn waking_an_empty_slot_does_nothing() {
let slot = WakerSlot::new();
assert!(slot.take().is_none());
slot.wake();
}
#[test]
fn registering_replaces_the_previous_waker_and_drops_it() {
let slot = WakerSlot::new();
let (first, first_waker) = Counter::waker();
let (second, second_waker) = Counter::waker();
slot.register(&first_waker);
assert_eq!(Arc::strong_count(&first), 3, "the Arc, the Waker, and the slot's clone");
slot.register(&second_waker);
assert_eq!(Arc::strong_count(&first), 2, "the replaced clone must be dropped");
slot.wake();
assert_eq!(first.count(), 0);
assert_eq!(second.count(), 1);
}
#[test]
fn re_registering_the_same_waker_does_not_clone_it() {
let slot = WakerSlot::new();
let (counter, waker) = Counter::waker();
slot.register(&waker);
let after_first = Arc::strong_count(&counter);
for _ in 0..8 {
slot.register(&waker);
}
assert_eq!(Arc::strong_count(&counter), after_first);
slot.wake();
assert_eq!(counter.count(), 1);
}
#[test]
fn taking_hands_the_waker_to_the_caller_rather_than_waking_it() {
let slot = WakerSlot::new();
let (counter, waker) = Counter::waker();
slot.register(&waker);
let taken = slot.take().expect("a waker was registered");
assert_eq!(counter.count(), 0, "take must not wake on the caller's behalf");
assert!(slot.take().is_none(), "the slot is empty after a take");
taken.wake();
assert_eq!(counter.count(), 1);
}
#[test]
fn concurrent_wakes_racing_registration_are_never_lost() {
let rounds = if cfg!(miri) { 24 } else { 2_000 };
for _ in 0..rounds {
let slot = Arc::new(WakerSlot::new());
let work = Arc::new(AtomicUsize::new(0));
let (counter, waker) = Counter::waker();
let producer = {
let slot = Arc::clone(&slot);
let work = Arc::clone(&work);
std::thread::spawn(move || {
work.store(1, Ordering::Release);
slot.wake();
})
};
slot.register(&waker);
let observed = work.load(Ordering::Acquire) == 1;
producer.join().unwrap();
assert!(
observed || counter.count() > 0,
"the owner neither observed the work nor was woken for it"
);
}
}
#[test]
fn the_debug_rendering_does_not_reach_into_the_cell() {
let slot = WakerSlot::new();
assert!(format!("{slot:?}").contains("WakerSlot"));
}
}
#[cfg(all(test, loom))]
mod loom_tests {
use super::*;
use loom::sync::atomic::{AtomicBool, Ordering};
use std::sync::Arc;
use std::task::Wake;
struct Flag(AtomicBool);
impl Flag {
fn waker() -> (Arc<Self>, Waker) {
let flag = Arc::new(Self(AtomicBool::new(false)));
(flag.clone(), Waker::from(flag))
}
fn woken(&self) -> bool {
self.0.load(Ordering::Acquire)
}
}
impl Wake for Flag {
fn wake(self: Arc<Self>) {
self.0.store(true, Ordering::Release);
}
fn wake_by_ref(self: &Arc<Self>) {
self.0.store(true, Ordering::Release);
}
}
#[test]
fn loom_a_wake_racing_a_registration_is_deferred_not_lost() {
loom::model(|| {
let slot = Arc::new(WakerSlot::new());
let work = Arc::new(AtomicBool::new(false));
let (flag, waker) = Flag::waker();
let producer = {
let slot = Arc::clone(&slot);
let work = Arc::clone(&work);
loom::thread::spawn(move || {
work.store(true, Ordering::Release);
slot.wake();
})
};
slot.register(&waker);
let observed = work.load(Ordering::Acquire);
producer.join().unwrap();
assert!(observed || flag.woken(), "a wake was lost across a registration");
});
}
#[test]
fn loom_concurrent_wakes_deliver_exactly_one_of_them() {
loom::model(|| {
let slot = Arc::new(WakerSlot::new());
let (flag, waker) = Flag::waker();
slot.register(&waker);
let left = {
let slot = Arc::clone(&slot);
loom::thread::spawn(move || slot.take().is_some())
};
let right = {
let slot = Arc::clone(&slot);
loom::thread::spawn(move || slot.take().is_some())
};
let took = usize::from(left.join().unwrap()) + usize::from(right.join().unwrap());
assert_eq!(took, 1, "a registered waker must be handed out exactly once");
assert!(!flag.woken(), "take must not wake on the caller's behalf");
});
}
#[test]
fn loom_a_wake_racing_a_replacement_still_schedules_somebody() {
loom::model(|| {
let slot = Arc::new(WakerSlot::new());
let (first, first_waker) = Flag::waker();
let (second, second_waker) = Flag::waker();
slot.register(&first_waker);
let notifier = {
let slot = Arc::clone(&slot);
loom::thread::spawn(move || slot.wake())
};
slot.register(&second_waker);
notifier.join().unwrap();
assert!(first.woken() || second.woken(), "a wake vanished between two registrations");
});
}
}