use std::cell::Cell;
use std::cell::RefCell;
use std::marker::PhantomData;
use std::sync::Once;
use std::sync::atomic::AtomicU32;
use std::sync::atomic::Ordering;
const LAUNCHING: u32 = 1 << 31;
pub const OPEN_WAIT_LIMIT: std::time::Duration = std::time::Duration::from_millis(100);
static STATE: AtomicU32 = AtomicU32::new(0);
static CHILD_RESET: Once = Once::new();
thread_local! {
static OPENS_HELD: Cell<u32> = const { Cell::new(0) };
static LAUNCH_DEPTH: Cell<u32> = const { Cell::new(0) };
static OPENS_IN_LAUNCH: Cell<u32> = const { Cell::new(0) };
static AFTER_GUARDS: RefCell<Vec<Box<dyn FnOnce()>>> = const { RefCell::new(Vec::new()) };
static DEFERRED: Cell<bool> = const { Cell::new(false) };
}
fn holds_a_guard() -> bool {
OPENS_HELD.with(Cell::get) > 0
|| LAUNCH_DEPTH.with(Cell::get) > 0
|| OPENS_IN_LAUNCH.with(Cell::get) > 0
}
pub fn run_after_guards(work: impl FnOnce() + 'static) {
if holds_a_guard() {
AFTER_GUARDS.with(|queue| queue.borrow_mut().push(Box::new(work)));
DEFERRED.with(|deferred| deferred.set(true));
} else {
work();
}
}
fn run_deferred_work() {
while DEFERRED.with(Cell::get) && !holds_a_guard() {
DEFERRED.with(|deferred| deferred.set(false));
let queue = AFTER_GUARDS.with(|queue| queue.take());
if queue.is_empty() {
return;
}
for work in queue {
work();
}
}
}
extern "C" fn reset_in_child() {
let held = OPENS_HELD.try_with(Cell::get).unwrap_or(0)
+ OPENS_IN_LAUNCH
.try_with(|opens| opens.replace(0))
.unwrap_or(0);
let _ = OPENS_HELD.try_with(|opens| opens.set(held));
let _ = LAUNCH_DEPTH.try_with(|depth| depth.set(0));
if matches!(
DEFERRED.try_with(|deferred| deferred.replace(false)),
Ok(true)
) {
let _ = AFTER_GUARDS.try_with(|queue| {
if let Ok(mut queue) = queue.try_borrow_mut() {
std::mem::forget(std::mem::take(&mut *queue));
}
});
}
STATE.store(held, Ordering::Relaxed);
}
fn register_child_reset() {
CHILD_RESET.call_once(|| {
let rc = unsafe { libc::pthread_atfork(None, None, Some(reset_in_child)) };
assert_eq!(rc, 0, "pthread_atfork failed: {rc}");
});
}
pub(crate) fn reset_after_raw_clone() {
reset_in_child();
}
const RECHECK_AFTER: libc::timespec = libc::timespec {
tv_sec: 0,
tv_nsec: 10_000_000,
};
fn wait_while(observed: u32) {
unsafe {
libc::syscall(
libc::SYS_futex,
STATE.as_ptr(),
libc::FUTEX_WAIT | libc::FUTEX_PRIVATE_FLAG,
observed,
&RECHECK_AFTER as *const libc::timespec,
)
};
}
fn wake_all() {
let _woken = unsafe {
libc::syscall(
libc::SYS_futex,
STATE.as_ptr(),
libc::FUTEX_WAKE | libc::FUTEX_PRIVATE_FLAG,
i32::MAX,
)
};
#[cfg(test)]
LAST_WAKE.with(|last| {
last.set(Some(match _woken {
-1 => Err(std::io::Error::last_os_error().raw_os_error().unwrap_or(0)),
woken => Ok(woken),
}))
});
}
#[cfg(test)]
thread_local! {
static LAST_WAKE: Cell<Option<Result<libc::c_long, i32>>> = const { Cell::new(None) };
}
pub fn held_by_this_thread() -> bool {
LAUNCH_DEPTH.try_with(Cell::get).unwrap_or(0) > 0
}
pub fn launching() -> bool {
STATE.load(Ordering::Relaxed) & LAUNCHING != 0
}
#[must_use]
pub struct Launch {
_not_send: PhantomData<*const ()>,
}
impl Launch {
pub fn begin() -> Self {
register_child_reset();
let depth = LAUNCH_DEPTH.with(Cell::get);
if depth == 0 {
let own = OPENS_HELD.with(Cell::get);
loop {
match STATE.compare_exchange_weak(
own,
LAUNCHING,
Ordering::Acquire,
Ordering::Relaxed,
) {
Ok(_) => break,
Err(observed) if observed == own => {}
Err(observed) => wait_while(observed),
}
}
OPENS_HELD.with(|opens| opens.set(0));
OPENS_IN_LAUNCH.with(|opens| opens.set(own));
}
LAUNCH_DEPTH.with(|launches| launches.set(depth + 1));
Launch {
_not_send: PhantomData,
}
}
}
impl Drop for Launch {
fn drop(&mut self) {
let depth = LAUNCH_DEPTH.with(Cell::get);
if depth == 0 {
return;
}
LAUNCH_DEPTH.with(|launches| launches.set(depth - 1));
if depth > 1 {
return;
}
let own = OPENS_IN_LAUNCH.with(|opens| opens.replace(0));
OPENS_HELD.with(|opens| opens.set(own));
STATE.fetch_sub(LAUNCHING - own, Ordering::Release);
wake_all();
run_deferred_work();
}
}
#[must_use]
pub struct TransientOpen {
_not_send: PhantomData<*const ()>,
}
impl TransientOpen {
pub fn begin() -> Self {
register_child_reset();
if LAUNCH_DEPTH.with(Cell::get) > 0 {
OPENS_IN_LAUNCH.with(|opens| opens.set(opens.get() + 1));
return TransientOpen {
_not_send: PhantomData,
};
}
let mut observed = STATE.load(Ordering::Relaxed);
let mut give_up_at = None;
loop {
if observed & LAUNCHING != 0
&& std::time::Instant::now()
< *give_up_at.get_or_insert_with(|| std::time::Instant::now() + OPEN_WAIT_LIMIT)
{
wait_while(observed);
observed = STATE.load(Ordering::Relaxed);
continue;
}
match STATE.compare_exchange_weak(
observed,
observed + 1,
Ordering::Acquire,
Ordering::Relaxed,
) {
Ok(_) => break,
Err(now) => observed = now,
}
}
OPENS_HELD.with(|held| held.set(held.get() + 1));
TransientOpen {
_not_send: PhantomData,
}
}
}
impl Drop for TransientOpen {
fn drop(&mut self) {
if LAUNCH_DEPTH.with(Cell::get) > 0 {
OPENS_IN_LAUNCH.with(|opens| opens.set(opens.get() - 1));
return;
}
OPENS_HELD.with(|held| held.set(held.get() - 1));
if STATE.fetch_sub(1, Ordering::Release) == 1 {
wake_all();
}
run_deferred_work();
}
}
pub fn read_to_string<P: AsRef<std::path::Path>>(path: P) -> std::io::Result<String> {
let _open = TransientOpen::begin();
std::fs::read_to_string(path)
}
pub fn read<P: AsRef<std::path::Path>>(path: P) -> std::io::Result<Vec<u8>> {
let _open = TransientOpen::begin();
std::fs::read(path)
}
#[cfg(test)]
mod tests {
use std::sync::Arc;
use std::sync::atomic::AtomicBool;
use std::sync::mpsc;
use std::time::Duration;
use std::time::Instant;
use super::*;
fn gettid() -> libc::pid_t {
unsafe { libc::gettid() }
}
fn wait_until_asleep_on_the_lock(tid: libc::pid_t) {
wait_until_asleep_on_the_lock_or_proceeded(tid, &AtomicBool::new(false));
}
fn wait_until_asleep_on_the_lock_or_proceeded(tid: libc::pid_t, proceeded: &AtomicBool) {
let expected = format!(
"{} {:#x} {:#x} ",
libc::SYS_futex,
STATE.as_ptr() as usize,
libc::FUTEX_WAIT | libc::FUTEX_PRIVATE_FLAG
);
let deadline = Instant::now() + Duration::from_secs(10);
loop {
let syscall = std::fs::read_to_string(format!("/proc/self/task/{tid}/syscall"))
.unwrap_or_default();
if syscall.starts_with(&expected) || proceeded.load(Ordering::SeqCst) {
return;
}
assert!(
Instant::now() < deadline,
"thread {tid} was not seen waiting on the lock within 10 s (last: {syscall:?})"
);
std::thread::sleep(Duration::from_millis(1));
}
}
#[test]
fn transient_opens_nest_on_one_thread() {
let outer = TransientOpen::begin();
let inner = TransientOpen::begin();
assert_eq!(OPENS_HELD.with(Cell::get), 2);
drop(inner);
drop(outer);
assert_eq!(OPENS_HELD.with(Cell::get), 0);
}
#[test]
fn a_launch_waits_for_a_transient_open_to_close() {
let open = TransientOpen::begin();
let launched = Arc::new(AtomicBool::new(false));
let (tid_tx, tid_rx) = mpsc::channel();
let launcher = {
let launched = Arc::clone(&launched);
std::thread::spawn(move || {
tid_tx.send(gettid()).unwrap();
let _launch = Launch::begin();
launched.store(true, Ordering::SeqCst);
})
};
wait_until_asleep_on_the_lock(tid_rx.recv().unwrap());
assert!(
!launched.load(Ordering::SeqCst),
"launch ran during a transient open"
);
drop(open);
launcher.join().unwrap();
assert!(launched.load(Ordering::SeqCst));
}
#[test]
fn a_transient_open_waits_for_a_launch_to_finish() {
let launch = Launch::begin();
let opened = Arc::new(AtomicBool::new(false));
let (tid_tx, tid_rx) = mpsc::channel();
let reader = {
let opened = Arc::clone(&opened);
std::thread::spawn(move || {
tid_tx.send(gettid()).unwrap();
let began = Instant::now();
let _open = TransientOpen::begin();
let during = launching();
opened.store(true, Ordering::SeqCst);
(began.elapsed(), during)
})
};
wait_until_asleep_on_the_lock_or_proceeded(tid_rx.recv().unwrap(), &opened);
drop(launch);
let (waited, during) = reader.join().unwrap();
assert!(opened.load(Ordering::SeqCst));
assert!(
!during || waited >= OPEN_WAIT_LIMIT,
"transient open ran during a launch after {waited:?}, before {OPEN_WAIT_LIMIT:?}"
);
}
#[test]
fn a_transient_open_proceeds_after_the_limit_and_a_later_launch_waits_for_it() {
let launch = Launch::begin();
let (tid_tx, tid_rx) = mpsc::channel();
let (opened_tx, opened_rx) = mpsc::channel();
let (close_tx, close_rx) = mpsc::channel::<()>();
let proceeded = Arc::new(AtomicBool::new(false));
let reader = {
let proceeded = Arc::clone(&proceeded);
std::thread::spawn(move || {
tid_tx.send(gettid()).unwrap();
let began = Instant::now();
let open = TransientOpen::begin();
proceeded.store(true, Ordering::SeqCst);
opened_tx.send(began.elapsed()).unwrap();
close_rx.recv().unwrap();
drop(open);
})
};
wait_until_asleep_on_the_lock_or_proceeded(tid_rx.recv().unwrap(), &proceeded);
let waited = opened_rx
.recv_timeout(Duration::from_secs(10))
.expect("a transient open still waited for a launch 10 s later");
assert!(launching(), "the launch ended before the open proceeded");
assert!(
waited >= OPEN_WAIT_LIMIT,
"the open proceeded after {waited:?}, before {OPEN_WAIT_LIMIT:?}"
);
drop(launch);
let launched = Arc::new(AtomicBool::new(false));
let (launcher_tx, launcher_rx) = mpsc::channel();
let launcher = {
let launched = Arc::clone(&launched);
std::thread::spawn(move || {
launcher_tx.send(gettid()).unwrap();
let _launch = Launch::begin();
launched.store(true, Ordering::SeqCst);
})
};
wait_until_asleep_on_the_lock(launcher_rx.recv().unwrap());
assert!(
!launched.load(Ordering::SeqCst),
"a launch ran while an open that began during the previous launch was held"
);
close_tx.send(()).unwrap();
reader.join().unwrap();
launcher.join().unwrap();
assert!(launched.load(Ordering::SeqCst));
}
fn deny_wake_on_this_thread() {
const fn stmt(code: u32, k: u32) -> libc::sock_filter {
libc::sock_filter {
code: code as u16,
jt: 0,
jf: 0,
k,
}
}
const fn jeq_else_allow(k: u32, to_allow: u8) -> libc::sock_filter {
libc::sock_filter {
code: (libc::BPF_JMP | libc::BPF_JEQ | libc::BPF_K) as u16,
jt: 0,
jf: to_allow,
k,
}
}
let load = libc::BPF_LD | libc::BPF_W | libc::BPF_ABS;
let filter = [
stmt(load, 0),
jeq_else_allow(libc::SYS_futex as u32, 5),
stmt(load, 24),
jeq_else_allow((libc::FUTEX_WAKE | libc::FUTEX_PRIVATE_FLAG) as u32, 3),
stmt(load, 32),
jeq_else_allow(i32::MAX as u32, 1),
stmt(
libc::BPF_RET | libc::BPF_K,
libc::SECCOMP_RET_ERRNO | libc::EPERM as u32,
),
stmt(libc::BPF_RET | libc::BPF_K, libc::SECCOMP_RET_ALLOW),
];
let prog = libc::sock_fprog {
len: filter.len() as u16,
filter: filter.as_ptr() as *mut libc::sock_filter,
};
unsafe {
assert_eq!(libc::prctl(libc::PR_SET_NO_NEW_PRIVS, 1, 0, 0, 0), 0);
assert_eq!(
libc::prctl(
libc::PR_SET_SECCOMP,
libc::SECCOMP_MODE_FILTER,
&prog as *const libc::sock_fprog,
),
0,
"installing the test filter failed: {}",
std::io::Error::last_os_error()
);
}
}
const ALONE: &str = "REVERIE_LAUNCH_WINDOW_TEST_ALONE";
fn running_alone(test: &str) -> bool {
let (_crate, module) = module_path!().split_once("::").unwrap();
let name = format!("{module}::{test}");
if std::env::var_os(ALONE).is_some_and(|alone| alone == name.as_str()) {
return true;
}
let output = std::process::Command::new(std::env::current_exe().unwrap())
.args([name.as_str(), "--exact", "--test-threads=1", "--nocapture"])
.env(ALONE, &name)
.output()
.unwrap();
let stdout = String::from_utf8_lossy(&output.stdout);
assert!(
output.status.success() && stdout.contains("test result: ok. 1 passed"),
"{name} failed when run alone ({}):\n{stdout}\n{}",
output.status,
String::from_utf8_lossy(&output.stderr)
);
false
}
#[test]
fn a_waiting_launch_proceeds_when_the_release_wake_is_denied() {
if !running_alone("a_waiting_launch_proceeds_when_the_release_wake_is_denied") {
return;
}
let opened = Arc::new(AtomicBool::new(false));
let release = Arc::new(AtomicBool::new(false));
let holder = {
let (opened, release) = (Arc::clone(&opened), Arc::clone(&release));
std::thread::spawn(move || {
let open = TransientOpen::begin();
opened.store(true, Ordering::SeqCst);
while !release.load(Ordering::SeqCst) {
std::thread::sleep(Duration::from_millis(5));
}
deny_wake_on_this_thread();
let before = STATE.load(Ordering::SeqCst);
LAST_WAKE.with(|last| last.set(None));
drop(open);
(before, LAST_WAKE.with(Cell::get))
})
};
while !opened.load(Ordering::SeqCst) {
std::thread::sleep(Duration::from_millis(5));
}
let launched = Arc::new(AtomicBool::new(false));
let (tid_tx, tid_rx) = mpsc::channel();
let launcher = {
let launched = Arc::clone(&launched);
std::thread::spawn(move || {
tid_tx.send(gettid()).unwrap();
let _launch = Launch::begin();
launched.store(true, Ordering::SeqCst);
})
};
wait_until_asleep_on_the_lock(tid_rx.recv().unwrap());
assert!(
!launched.load(Ordering::SeqCst),
"launch ran during a transient open"
);
release.store(true, Ordering::SeqCst);
let (before, wake) = holder.join().unwrap();
assert_eq!(before, 1, "the holder's open was not the only one");
assert_eq!(
wake,
Some(Err(libc::EPERM)),
"the release did not issue a wake that the filter denied"
);
let deadline = Instant::now() + Duration::from_secs(5);
while !launched.load(Ordering::SeqCst) && Instant::now() < deadline {
std::thread::sleep(Duration::from_millis(5));
}
assert!(
launched.load(Ordering::SeqCst),
"the launch was still waiting 5 s after a release whose wake was denied"
);
launcher.join().unwrap();
}
#[test]
fn work_handed_over_under_a_guard_runs_after_the_last_guard() {
use std::cell::RefCell;
use std::rc::Rc;
let ran = Rc::new(RefCell::new(Vec::new()));
let note = |name: &'static str| {
let ran = Rc::clone(&ran);
move || ran.borrow_mut().push(name)
};
run_after_guards(note("unguarded"));
assert_eq!(*ran.borrow(), ["unguarded"]);
let outer = TransientOpen::begin();
let inner = TransientOpen::begin();
run_after_guards(note("first"));
let (second, third) = (note("second"), note("third"));
run_after_guards(move || {
second();
run_after_guards(third);
});
drop(inner);
assert_eq!(*ran.borrow(), ["unguarded"], "ran under the outer open");
drop(outer);
assert_eq!(*ran.borrow(), ["unguarded", "first", "second", "third"]);
let launch = Launch::begin();
run_after_guards(note("launch"));
assert_eq!(ran.borrow().len(), 4, "ran under the launch");
drop(launch);
assert_eq!(ran.borrow().last(), Some(&"launch"));
}
#[test]
fn a_thread_never_waits_for_its_own_guards() {
let open = TransientOpen::begin();
let launch = Launch::begin();
let inner_open = TransientOpen::begin();
let inner_launch = Launch::begin();
drop(inner_launch);
drop(inner_open);
assert!(
launching() && held_by_this_thread(),
"the outer launch must stay held after the nested one ends"
);
assert_eq!(OPENS_IN_LAUNCH.with(Cell::get), 1);
let opened = Arc::new(AtomicBool::new(false));
let (tid_tx, tid_rx) = mpsc::channel();
let reader = {
let opened = Arc::clone(&opened);
std::thread::spawn(move || {
tid_tx.send(gettid()).unwrap();
let began = Instant::now();
let _open = TransientOpen::begin();
let during = launching();
opened.store(true, Ordering::SeqCst);
(began.elapsed(), during)
})
};
wait_until_asleep_on_the_lock_or_proceeded(tid_rx.recv().unwrap(), &opened);
drop(launch);
let (waited, during) = reader.join().unwrap();
assert!(opened.load(Ordering::SeqCst));
assert!(
!during || waited >= OPEN_WAIT_LIMIT,
"another thread's open ran during the outer launch after {waited:?}"
);
assert_eq!(OPENS_HELD.with(Cell::get), 1);
drop(open);
assert_eq!(OPENS_HELD.with(Cell::get), 0);
}
}