#![doc = include_str!("../README.md")]
#![cfg_attr(docsrs, feature(doc_cfg))]
#![cfg_attr(coverage_nightly, feature(coverage_attribute))]
#![forbid(unsafe_code)]
#[cfg(any(test, feature = "test-clock"))]
mod paused;
use std::fmt;
use std::sync::{Arc, Condvar, Mutex, MutexGuard};
use std::time::Instant;
#[derive(Clone)]
#[cfg_attr(not(any(test, feature = "test-clock")), derive(PartialEq, Eq))]
pub struct Clock {
#[cfg(any(test, feature = "test-clock"))]
paused: Option<Arc<paused::Paused>>,
}
impl Clock {
pub fn real() -> Self {
Self {
#[cfg(any(test, feature = "test-clock"))]
paused: None,
}
}
pub fn now(&self) -> Instant {
#[cfg(any(test, feature = "test-clock"))]
if let Some(paused) = &self.paused {
return paused.now();
}
Instant::now()
}
pub fn waiter(&self) -> Waiter {
let signal = Arc::new(Signal::default());
#[cfg(any(test, feature = "test-clock"))]
if let Some(paused) = &self.paused {
paused.register(&signal);
}
Waiter {
clock: self.clone(),
signal,
}
}
}
impl fmt::Debug for Clock {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
let mut clock = f.debug_struct("Clock");
#[cfg(any(test, feature = "test-clock"))]
if let Some(paused) = &self.paused {
return clock
.field("paused", &true)
.field("advanced", &paused.advanced())
.finish();
}
clock.field("paused", &false).finish()
}
}
#[derive(Clone)]
pub struct Waiter {
clock: Clock,
signal: Arc<Signal>,
}
impl Waiter {
pub fn notify_all(&self) {
self.signal.notify_all();
}
pub fn wait_until<T>(
&self,
deadline: Option<Instant>,
mut ready: impl FnMut() -> Option<T>,
) -> Option<T> {
#[cfg(any(test, feature = "test-clock"))]
let timer = deadline.filter(|_| self.clock.paused.is_none());
#[cfg(not(any(test, feature = "test-clock")))]
let timer = deadline;
loop {
let seen = self.signal.generation();
if let Some(value) = ready() {
return Some(value);
}
if deadline.is_some_and(|deadline| self.clock.now() >= deadline) {
return None;
}
self.signal.park(seen, timer);
}
}
}
impl fmt::Debug for Waiter {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("Waiter")
.field("clock", &self.clock)
.finish_non_exhaustive()
}
}
#[derive(Default)]
struct Signal {
state: Mutex<SignalState>,
changed: Condvar,
}
#[derive(Default)]
struct SignalState {
generation: u64,
#[cfg(test)]
parks: usize,
}
impl Signal {
fn lock(&self) -> MutexGuard<'_, SignalState> {
self.state.lock().expect("waiter signal not poisoned")
}
fn generation(&self) -> u64 {
self.lock().generation
}
fn notify_all(&self) {
self.lock().generation += 1;
self.changed.notify_all();
}
fn park(&self, seen: u64, timer: Option<Instant>) {
let mut state = self.lock();
#[cfg(test)]
if state.generation == seen {
state.parks += 1;
self.changed.notify_all();
}
while state.generation == seen {
state = match timer {
None => self
.changed
.wait(state)
.expect("waiter signal not poisoned"),
Some(timer) => {
let now = Instant::now();
if now >= timer {
break;
}
self.changed
.wait_timeout(state, timer - now)
.expect("waiter signal not poisoned")
.0
}
};
}
}
}
#[cfg(test)]
#[cfg_attr(coverage_nightly, coverage(off))]
mod tests {
use super::*;
use std::panic::{self, AssertUnwindSafe};
use std::sync::atomic::{AtomicBool, Ordering};
use std::thread;
use std::time::Duration;
fn parked(waiter: &Waiter, count: usize) {
let mut state = waiter.signal.lock();
while state.parks < count {
state = waiter.signal.changed.wait(state).unwrap();
}
}
#[test]
fn test_paused_clock_sharing() {
let other = Clock::paused();
let clock = Clock::paused();
let clone = clock.clone();
let other_start = other.now();
let start = clock.now();
clone.advance(Duration::from_secs(1));
assert_eq!(clock.now(), start + Duration::from_secs(1));
clock.advance_to(start + Duration::from_secs(3));
assert_eq!(clone.now(), start + Duration::from_secs(3));
clock.advance(Duration::ZERO);
clock.advance_to(start + Duration::from_secs(3));
assert_eq!(clock.now(), start + Duration::from_secs(3));
assert_eq!(other.now(), other_start);
other.advance_to(clock.now());
assert_eq!(other.now(), clock.now());
assert_eq!(clock, clone);
assert_ne!(clock, other);
assert_ne!(clock, Clock::real());
assert_eq!(Clock::real(), Clock::real());
}
#[test]
fn test_failed_advance_keeps_time() {
let clock = Clock::paused();
let real = Clock::real();
clock.advance(Duration::from_secs(1));
let now = clock.now();
let misuses: [(&str, &dyn Fn()); 3] = [
("real", &|| real.advance(Duration::from_secs(1))),
("backwards", &|| {
clock.advance_to(now - Duration::from_secs(1))
}),
("overflow", &|| clock.advance(Duration::MAX)),
];
for (case, misuse) in misuses {
assert!(
panic::catch_unwind(AssertUnwindSafe(misuse)).is_err(),
"{case}"
);
assert_eq!(clock.now(), now, "{case}");
}
clock.advance(Duration::from_secs(1));
assert_eq!(clock.now(), now + Duration::from_secs(1));
}
#[test]
fn test_wait_checks_ready_first() {
for clock in [Clock::real(), Clock::paused()] {
let waiter = clock.waiter();
let now = clock.now();
assert_eq!(
waiter.wait_until(Some(now), || Some(1)),
Some(1),
"{clock:?}"
);
assert_eq!(
waiter.wait_until(Some(now), || None::<u8>),
None,
"{clock:?}"
);
}
}
#[test]
fn test_wait_sees_changes_after_its_check() {
let clock = Clock::paused();
let waiter = clock.waiter();
let deadline = clock.now() + Duration::from_secs(2);
let mut checks = 0;
let result = waiter.wait_until(Some(deadline), || {
checks += 1;
if checks == 1 {
waiter.notify_all();
} else if clock.now() < deadline {
clock.advance(Duration::from_secs(1));
}
None::<()>
});
assert_eq!(result, None);
assert_eq!(clock.now(), deadline);
}
#[test]
fn test_advance_wakes_parked_waiters() {
let clock = Clock::paused();
let deadline = clock.now() + Duration::from_secs(60);
let mut waiters: Vec<_> = (0..4).map(|_| clock.waiter()).collect();
waiters.truncate(1);
waiters.push(clock.clone().waiter());
drop(clock.waiter());
let threads: Vec<_> = waiters
.iter()
.map(|waiter| {
let waiter = waiter.clone();
thread::spawn(move || waiter.wait_until(Some(deadline), || None::<()>))
})
.collect();
for waiter in &waiters {
parked(waiter, 1);
}
clock.advance(Duration::from_secs(30));
for waiter in &waiters {
parked(waiter, 2);
}
clock.advance(Duration::from_secs(30));
for thread in threads {
assert_eq!(thread.join().unwrap(), None);
}
}
#[test]
fn test_ready_value_survives_advance() {
let clock = Clock::paused();
let waiter = clock.waiter();
let deadline = clock.now() + Duration::from_secs(1);
let ready = Arc::new(AtomicBool::new(false));
let waiting = thread::spawn({
let (waiter, ready) = (waiter.clone(), ready.clone());
move || {
waiter.wait_until(Some(deadline), || {
ready.load(Ordering::SeqCst).then_some(())
})
}
});
parked(&waiter, 1);
ready.store(true, Ordering::SeqCst);
clock.advance(Duration::from_secs(1));
assert_eq!(waiting.join().unwrap(), Some(()));
}
#[test]
fn test_notify_wakes_parked_waiters() {
for clock in [Clock::real(), Clock::paused()] {
let waiter = clock.waiter();
let open = Arc::new(AtomicBool::new(false));
let threads: Vec<_> = (0..2)
.map(|_| {
let (waiter, open) = (waiter.clone(), open.clone());
thread::spawn(move || {
waiter.wait_until(None, || open.load(Ordering::SeqCst).then_some(()))
})
})
.collect();
parked(&waiter, 2);
open.store(true, Ordering::SeqCst);
waiter.notify_all();
for thread in threads {
assert_eq!(thread.join().unwrap(), Some(()), "{clock:?}");
}
}
}
#[test]
fn test_real_deadline_expires() {
let waiter = Clock::real().waiter();
let deadline = Instant::now() + Duration::from_millis(1);
assert_eq!(waiter.wait_until(Some(deadline), || None::<()>), None);
assert!(Instant::now() >= deadline);
}
}