#![cfg(loom)]
use core::future::Future;
use loom::thread;
use rivet::sync::atomic::Ordering;
use rivet::sync::{Channel, Semaphore, Signal};
use rivet::waker;
fn model<F: Fn() + Send + Sync + 'static>(f: F) {
loom::model(move || {
rivet::waker::reset();
f();
});
}
loom::lazy_static! {
static ref CHAN: Channel<u32, 4> = Channel::new();
}
#[test]
fn waker_no_lost_wakeups_single_producer() {
model(|| {
const N: usize = 2;
let mut handles = Vec::new();
for i in 0..N {
let h = thread::spawn(move || {
waker::mark_ready(rivet::task::TaskId::new(2, i as u8));
});
handles.push(h);
}
for h in handles {
h.join().unwrap();
}
let mut seen = [false; N];
while let Some(id) = waker::next_ready() {
assert_eq!(id.priority(), 2, "unexpected priority");
let idx = id.index() as usize;
assert!(!seen[idx], "task {idx} dequeued twice");
seen[idx] = true;
}
assert!(seen.iter().all(|&s| s), "lost wakeups: {seen:?}");
});
}
#[test]
fn semaphore_single_waiter_no_lost_wakeup() {
model(|| {
let sem: loom::sync::Arc<Semaphore<1>> = loom::sync::Arc::new(Semaphore::new(0));
let sem_w = sem.clone();
let w = thread::spawn(move || {
rivet::executor::set_current_for_test(1, 0);
let waker = rivet::waker::task_waker(rivet::task::TaskId::new(1, 0));
let mut cx = core::task::Context::from_waker(&waker);
let mut fut = sem_w.acquire();
let pinned = unsafe { core::pin::Pin::new_unchecked(&mut fut) };
pinned.poll(&mut cx).is_ready()
});
let sem_s = sem.clone();
let s = thread::spawn(move || {
sem_s.release();
});
let poll_ready = w.join().unwrap();
s.join().unwrap();
assert!(
poll_ready || waker::has_pending() || sem.try_acquire(),
"lost wakeup: waiter never woken, semaphore never released"
);
});
}
#[test]
fn semaphore_two_waiters_no_lost_wakeup() {
model(|| {
let sem: loom::sync::Arc<Semaphore<1>> = loom::sync::Arc::new(Semaphore::new(0));
let mut poll1 = false;
let mut poll2 = false;
{
rivet::executor::set_current_for_test(1, 0);
let waker = rivet::waker::task_waker(rivet::task::TaskId::new(1, 0));
let mut cx = core::task::Context::from_waker(&waker);
let mut fut = sem.acquire();
let pinned = unsafe { core::pin::Pin::new_unchecked(&mut fut) };
poll1 = pinned.poll(&mut cx).is_ready();
core::mem::forget(fut);
}
{
rivet::executor::set_current_for_test(2, 0);
let waker = rivet::waker::task_waker(rivet::task::TaskId::new(2, 0));
let mut cx = core::task::Context::from_waker(&waker);
let mut fut = sem.acquire();
let pinned = unsafe { core::pin::Pin::new_unchecked(&mut fut) };
poll2 = pinned.poll(&mut cx).is_ready();
core::mem::forget(fut);
}
rivet::executor::clear_current_for_test();
let sem_s1 = sem.clone();
let s1 = thread::spawn(move || {
sem_s1.release();
});
let sem_s2 = sem.clone();
let s2 = thread::spawn(move || {
sem_s2.release();
});
s1.join().unwrap();
s2.join().unwrap();
let mut marked: Vec<(u8, u8)> = Vec::new();
while let Some(id) = waker::next_ready() {
marked.push((id.priority(), id.index()));
}
let w1_ok = poll1 || marked.contains(&(1, 0));
let w2_ok = poll2 || marked.contains(&(2, 0));
assert!(
w1_ok && w2_ok,
"[B9] lost wakeup: poll1={poll1} poll2={poll2} marked={marked:?}"
);
});
}
#[test]
fn channel_spsc_every_value_received_exactly_once() {
model(|| {
let (tx, rx) = CHAN.split().expect("split once");
let producer = thread::spawn(move || {
for v in 1..=3u32 {
while tx.try_send(v).is_err() {
loom::hint::spin_loop();
}
}
});
let consumer = thread::spawn(move || {
let mut got = Vec::new();
while got.len() < 3 {
if let Some(v) = rx.try_recv() {
got.push(v);
} else {
loom::hint::spin_loop();
}
}
got
});
producer.join().unwrap();
let got = consumer.join().unwrap();
assert_eq!(
got,
vec![1, 2, 3],
"SPSC values lost, duplicated, or reordered"
);
});
}
#[test]
fn signal_no_lost_wakeup() {
model(|| {
let sig: loom::sync::Arc<Signal> = loom::sync::Arc::new(Signal::new());
let sig_w = sig.clone();
let w = thread::spawn(move || {
rivet::executor::set_current_for_test(1, 0);
let waker = rivet::waker::task_waker(rivet::task::TaskId::new(1, 0));
let mut cx = core::task::Context::from_waker(&waker);
let mut fut = sig_w.wait();
let pinned = unsafe { core::pin::Pin::new_unchecked(&mut fut) };
pinned.poll(&mut cx).is_ready()
});
let sig_s = sig.clone();
let s = thread::spawn(move || {
sig_s.signal();
});
let poll_ready = w.join().unwrap();
s.join().unwrap();
assert!(
poll_ready || waker::has_pending() || sig.try_take(),
"lost wakeup: waiter never woken, signal never observed"
);
});
}