use std::cell::UnsafeCell;
use std::panic::AssertUnwindSafe;
use std::panic::RefUnwindSafe;
use std::panic::UnwindSafe;
use std::panic::catch_unwind;
use std::panic::resume_unwind;
use std::sync::atomic::AtomicUsize;
use std::sync::atomic::Ordering;
use std::task::Waker;
const WAITING: usize = 0;
const REGISTERING: usize = 0b01;
const WAKING: usize = 0b10;
pub struct AtomicWaker {
state: AtomicUsize,
waker: UnsafeCell<Option<Waker>>,
}
unsafe impl Sync for AtomicWaker {}
impl RefUnwindSafe for AtomicWaker {}
impl UnwindSafe for AtomicWaker {}
impl AtomicWaker {
#[inline]
pub const fn new() -> Self {
Self {
state: AtomicUsize::new(WAITING),
waker: UnsafeCell::new(None),
}
}
#[inline]
pub fn register(&self, waker: &Waker) {
match self
.state
.compare_exchange(WAITING, REGISTERING, Ordering::Acquire, Ordering::Acquire)
.unwrap_or_else(|state| state)
{
WAITING => {
unsafe { self.register_locked(waker) }
}
WAKING => {
waker.wake_by_ref();
}
state => {
debug_assert!(state == REGISTERING || state == REGISTERING | WAKING);
}
}
}
#[inline]
unsafe fn register_locked(&self, waker: &Waker) {
let needs_replacement = match unsafe { &*self.waker.get() } {
Some(current) => !current.will_wake(waker),
None => true,
};
let mut clone_panic = None;
let old_waker = if needs_replacement {
match catch_unwind(AssertUnwindSafe(|| waker.clone())) {
Ok(new_waker) => unsafe { (*self.waker.get()).replace(new_waker) },
Err(payload) => {
clone_panic = Some(payload);
None
}
}
} else {
None
};
let concurrent_wake = match self.state.compare_exchange(
REGISTERING,
WAITING,
Ordering::AcqRel,
Ordering::Acquire,
) {
Ok(_) => None,
Err(state) => {
debug_assert_eq!(state, REGISTERING | WAKING);
let registered = unsafe { (*self.waker.get()).take() };
self.state.swap(WAITING, Ordering::AcqRel);
registered
}
};
if let Some(payload) = clone_panic {
if let Some(waker) = concurrent_wake {
let _ = catch_unwind(AssertUnwindSafe(|| waker.wake()));
}
resume_unwind(payload);
}
if let Some(waker) = concurrent_wake {
if let Some(old_waker) = old_waker {
let _ = catch_unwind(AssertUnwindSafe(|| old_waker.wake()));
}
waker.wake();
} else {
drop(old_waker);
}
}
#[inline]
pub fn wake(&self) {
if let Some(waker) = self.take() {
waker.wake();
}
}
#[inline]
fn take(&self) -> Option<Waker> {
match self.state.fetch_or(WAKING, Ordering::AcqRel) {
WAITING => {
let waker = unsafe { (*self.waker.get()).take() };
let old_state = self.state.swap(WAITING, Ordering::Release);
debug_assert_eq!(old_state, WAKING);
waker
}
state => {
debug_assert!(
state == REGISTERING || state == REGISTERING | WAKING || state == WAKING
);
None
}
}
}
}
#[cfg(test)]
mod tests {
use std::ptr;
use std::sync::Arc;
use std::sync::atomic::AtomicBool;
use std::sync::atomic::AtomicUsize;
use std::sync::atomic::Ordering;
use std::task::RawWaker;
use std::task::RawWakerVTable;
use std::task::Wake;
use super::*;
struct WakeCounter(AtomicUsize);
impl Wake for WakeCounter {
fn wake(self: Arc<Self>) {
self.0.fetch_add(1, Ordering::Relaxed);
}
}
#[cfg(panic = "unwind")]
fn clone_panicking_waker() -> Waker {
static VTABLE: RawWakerVTable = RawWakerVTable::new(
|_| panic!("clone failed"),
|_| unreachable!(),
|_| unreachable!(),
|_| {},
);
unsafe { Waker::from_raw(RawWaker::new(ptr::null(), &VTABLE)) }
}
#[test]
fn wake_notifies_once() {
let counter = Arc::new(WakeCounter(AtomicUsize::new(0)));
let waker = Waker::from(counter.clone());
let atomic_waker = AtomicWaker::new();
atomic_waker.register(&waker);
atomic_waker.wake();
atomic_waker.wake();
assert_eq!(counter.0.load(Ordering::Relaxed), 1);
}
#[test]
fn reregistering_same_task_does_not_clone_waker() {
let counter = Arc::new(WakeCounter(AtomicUsize::new(0)));
let waker = Waker::from(counter.clone());
let atomic_waker = AtomicWaker::new();
atomic_waker.register(&waker);
let registered_refs = Arc::strong_count(&counter);
atomic_waker.register(&waker);
assert_eq!(Arc::strong_count(&counter), registered_refs);
}
#[test]
fn wake_before_register_is_not_remembered() {
let counter = Arc::new(WakeCounter(AtomicUsize::new(0)));
let waker = Waker::from(counter.clone());
let atomic_waker = AtomicWaker::new();
atomic_waker.wake();
atomic_waker.register(&waker);
assert_eq!(counter.0.load(Ordering::Relaxed), 0);
atomic_waker.wake();
assert_eq!(counter.0.load(Ordering::Relaxed), 1);
}
#[test]
fn wake_during_replacement_notifies_old_and_new_tasks() {
let old_counter = Arc::new(WakeCounter(AtomicUsize::new(0)));
let old_waker = Waker::from(old_counter.clone());
let new_counter = Arc::new(WakeCounter(AtomicUsize::new(0)));
let new_waker = Waker::from(new_counter.clone());
let atomic_waker = AtomicWaker::new();
atomic_waker.register(&old_waker);
assert_eq!(
atomic_waker.state.compare_exchange(
WAITING,
REGISTERING,
Ordering::AcqRel,
Ordering::Acquire,
),
Ok(WAITING)
);
std::thread::scope(|scope| scope.spawn(|| atomic_waker.wake()).join().unwrap());
unsafe { atomic_waker.register_locked(&new_waker) };
assert_eq!(old_counter.0.load(Ordering::Relaxed), 1);
assert_eq!(new_counter.0.load(Ordering::Relaxed), 1);
}
#[test]
fn failed_wake_synchronizes_with_next_registration() {
for _ in 0..1_000 {
let did_publish = AtomicBool::new(false);
let atomic_waker = AtomicWaker::new();
atomic_waker.register(Waker::noop());
std::thread::scope(|scope| {
let wake = scope.spawn(|| {
did_publish.store(true, Ordering::Relaxed);
atomic_waker.take()
});
let local_waker = atomic_waker.take();
atomic_waker.register(Waker::noop());
let publication_is_visible = did_publish.load(Ordering::Relaxed);
let concurrent_thread_took_waker = wake.join().unwrap().is_some();
assert!(publication_is_visible || concurrent_thread_took_waker);
drop(local_waker);
});
}
}
#[cfg(panic = "unwind")]
#[test]
fn clone_panic_does_not_poison_state() {
let atomic_waker = AtomicWaker::new();
assert!(
catch_unwind(|| {
atomic_waker.register(&clone_panicking_waker());
})
.is_err()
);
let counter = Arc::new(WakeCounter(AtomicUsize::new(0)));
atomic_waker.register(&Waker::from(counter.clone()));
atomic_waker.wake();
assert_eq!(counter.0.load(Ordering::Relaxed), 1);
}
#[cfg(panic = "unwind")]
#[test]
fn clone_panic_completes_concurrent_wake() {
let counter = Arc::new(WakeCounter(AtomicUsize::new(0)));
let atomic_waker = AtomicWaker::new();
atomic_waker.register(&Waker::from(counter.clone()));
assert_eq!(
atomic_waker.state.compare_exchange(
WAITING,
REGISTERING,
Ordering::Acquire,
Ordering::Acquire,
),
Ok(WAITING)
);
std::thread::scope(|scope| scope.spawn(|| atomic_waker.wake()).join().unwrap());
assert!(
catch_unwind(|| unsafe {
atomic_waker.register_locked(&clone_panicking_waker());
})
.is_err()
);
assert_eq!(counter.0.load(Ordering::Relaxed), 1);
let next_counter = Arc::new(WakeCounter(AtomicUsize::new(0)));
atomic_waker.register(&Waker::from(next_counter.clone()));
atomic_waker.wake();
assert_eq!(next_counter.0.load(Ordering::Relaxed), 1);
}
#[cfg(panic = "unwind")]
#[test]
fn drop_panic_does_not_poison_state() {
unsafe fn clone_drop_panicker(data: *const ()) -> RawWaker {
RawWaker::new(data, &DROP_PANICKING_VTABLE)
}
unsafe fn wake_drop_panicker(_: *const ()) {}
unsafe fn drop_drop_panicker(data: *const ()) {
let should_panic = unsafe { &*data.cast::<AtomicBool>() };
if should_panic.swap(false, Ordering::Relaxed) {
panic!("drop failed");
}
}
static DROP_PANICKING_VTABLE: RawWakerVTable = RawWakerVTable::new(
clone_drop_panicker,
wake_drop_panicker,
wake_drop_panicker,
drop_drop_panicker,
);
let should_panic = AtomicBool::new(true);
let old_waker = unsafe {
Waker::from_raw(RawWaker::new(
ptr::from_ref(&should_panic).cast(),
&DROP_PANICKING_VTABLE,
))
};
let counter = Arc::new(WakeCounter(AtomicUsize::new(0)));
let new_waker = Waker::from(counter.clone());
let atomic_waker = AtomicWaker::new();
atomic_waker.register(&old_waker);
assert!(catch_unwind(AssertUnwindSafe(|| atomic_waker.register(&new_waker))).is_err());
atomic_waker.wake();
assert_eq!(counter.0.load(Ordering::Relaxed), 1);
}
}