use parking_lot::{
Condvar,
Mutex,
MutexGuard,
};
use std::sync::Arc;
use std::sync::atomic::{
AtomicU64,
Ordering,
};
use std::time::{
Duration,
Instant,
};
#[cfg(feature = "tokio")]
use crate::sleep::AsyncSleepFuture;
#[cfg(feature = "tokio")]
use tokio::sync::watch;
use crate::{
MockInstant,
MockTimeError,
MockWaiterKind,
};
static NEXT_MOCK_TIMELINE_ID: AtomicU64 = AtomicU64::new(1);
#[derive(Clone, Debug)]
pub struct MockTimeline {
id: u64,
shared: Arc<MockTimelineShared>,
#[cfg(feature = "tokio")]
async_event_sender: watch::Sender<u64>,
}
#[derive(Debug)]
struct MockTimelineShared {
state: Mutex<MockTimelineState>,
event_changed: Condvar,
waiters_changed: Condvar,
}
#[derive(Debug)]
struct MockTimelineState {
elapsed_nanos: u128,
time_epoch: u64,
event_epoch: u64,
sleep_waiters: usize,
deadline_waiters: usize,
}
#[cfg(feature = "tokio")]
#[derive(Debug)]
struct MockTimelineWaiterRegistration {
timeline: MockTimeline,
kind: MockWaiterKind,
}
#[cfg(feature = "tokio")]
impl MockTimelineWaiterRegistration {
fn new(timeline: MockTimeline, kind: MockWaiterKind) -> Self {
{
let mut state = timeline.lock_state();
MockTimeline::increment_waiter(&mut state, kind);
}
timeline.shared.waiters_changed.notify_all();
Self { timeline, kind }
}
}
#[cfg(feature = "tokio")]
impl Drop for MockTimelineWaiterRegistration {
fn drop(&mut self) {
{
let mut state = self.timeline.lock_state();
MockTimeline::decrement_waiter(&mut state, self.kind);
}
self.timeline.shared.waiters_changed.notify_all();
}
}
impl MockTimeline {
#[must_use]
pub fn new() -> Self {
#[cfg(feature = "tokio")]
let (async_event_sender, _) = watch::channel(0);
Self {
id: next_mock_timeline_id(),
shared: Arc::new(MockTimelineShared {
state: Mutex::new(MockTimelineState {
elapsed_nanos: 0,
time_epoch: 0,
event_epoch: 0,
sleep_waiters: 0,
deadline_waiters: 0,
}),
event_changed: Condvar::new(),
waiters_changed: Condvar::new(),
}),
#[cfg(feature = "tokio")]
async_event_sender,
}
}
#[inline]
pub const fn id(&self) -> u64 {
self.id
}
#[inline]
pub fn elapsed(&self) -> Duration {
duration_from_nanos_saturating(self.elapsed_nanos())
}
#[inline]
pub fn elapsed_nanos(&self) -> u128 {
self.lock_state().elapsed_nanos
}
#[inline]
pub fn now(&self) -> MockInstant {
MockInstant::from_nanos_since_origin(self.id, self.elapsed_nanos())
}
#[inline]
pub fn event_epoch(&self) -> u64 {
self.lock_state().event_epoch
}
pub fn advance(&self, duration: Duration) {
let event_epoch = {
let mut state = self.lock_state();
state.elapsed_nanos = state.elapsed_nanos.saturating_add(duration.as_nanos());
state.time_epoch = state.time_epoch.wrapping_add(1);
state.event_epoch = state.event_epoch.wrapping_add(1);
state.event_epoch
};
self.notify_waiters(event_epoch);
}
pub fn reset(&self) -> Result<(), MockTimeError> {
let event_epoch = {
let mut state = self.lock_state();
if state.sleep_waiters != 0 || state.deadline_waiters != 0 {
return Err(MockTimeError::ActiveWaiters);
}
state.elapsed_nanos = 0;
state.time_epoch = state.time_epoch.wrapping_add(1);
state.event_epoch = state.event_epoch.wrapping_add(1);
state.event_epoch
};
self.notify_waiters(event_epoch);
Ok(())
}
pub fn notify_external_change(&self) {
let event_epoch = {
let mut state = self.lock_state();
state.event_epoch = state.event_epoch.wrapping_add(1);
state.event_epoch
};
self.notify_waiters(event_epoch);
}
#[inline]
pub fn wait_until(&self, deadline: MockInstant) -> Result<(), MockTimeError> {
self.wait_until_with_kind(deadline, MockWaiterKind::Deadline)
}
#[inline]
pub fn wait_for(&self, duration: Duration) {
self.wait_until(self.now().saturating_add(duration))
.expect("relative waits should create deadlines on the same timeline");
}
pub fn wait_for_event_after(&self, observed_epoch: u64) {
let mut state = self.lock_state();
while state.event_epoch == observed_epoch {
self.shared.event_changed.wait(&mut state);
}
}
pub fn wait_for_blocked_waiters(&self, kind: MockWaiterKind, count: usize, real_timeout: Duration) -> bool {
let Some(deadline) = Instant::now().checked_add(real_timeout) else {
return false;
};
let mut state = self.lock_state();
while Self::waiter_count(&state, kind) < count {
let Some(remaining) = deadline.checked_duration_since(Instant::now()) else {
return false;
};
let wait_result = self.shared.waiters_changed.wait_for(&mut state, remaining);
if wait_result.timed_out() && Self::waiter_count(&state, kind) < count {
return false;
}
}
true
}
pub(crate) fn wait_until_with_kind(
&self,
deadline: MockInstant,
kind: MockWaiterKind,
) -> Result<(), MockTimeError> {
self.ensure_own_instant(deadline)?;
let mut state = self.lock_state();
if state.elapsed_nanos >= deadline.nanos_since_origin() {
return Ok(());
}
Self::increment_waiter(&mut state, kind);
self.shared.waiters_changed.notify_all();
while state.elapsed_nanos < deadline.nanos_since_origin() {
self.shared.event_changed.wait(&mut state);
}
Self::decrement_waiter(&mut state, kind);
self.shared.waiters_changed.notify_all();
Ok(())
}
#[cfg(feature = "tokio")]
pub(crate) fn wait_until_async_with_kind<'a>(
&'a self,
deadline: MockInstant,
kind: MockWaiterKind,
) -> Result<AsyncSleepFuture<'a>, MockTimeError> {
self.ensure_own_instant(deadline)?;
if self.elapsed_nanos() >= deadline.nanos_since_origin() {
return Ok(Box::pin(async {}));
}
let registration = MockTimelineWaiterRegistration::new(self.clone(), kind);
let mut event_receiver = self.async_event_sender.subscribe();
Ok(Box::pin(async move {
let _registration = registration;
loop {
if self.elapsed_nanos() >= deadline.nanos_since_origin() {
return;
}
event_receiver
.changed()
.await
.expect("mock timeline sender should live while timeline is borrowed");
}
}))
}
fn ensure_own_instant(&self, instant: MockInstant) -> Result<(), MockTimeError> {
if instant.timeline_id() == self.id {
Ok(())
} else {
Err(MockTimeError::MismatchedTimeline {
expected: self.id,
actual: instant.timeline_id(),
})
}
}
#[inline]
fn lock_state(&self) -> MutexGuard<'_, MockTimelineState> {
self.shared.state.lock()
}
fn notify_waiters(&self, event_epoch: u64) {
self.shared.event_changed.notify_all();
self.shared.waiters_changed.notify_all();
self.notify_async_waiters(event_epoch);
}
#[cfg(feature = "tokio")]
#[inline]
fn notify_async_waiters(&self, event_epoch: u64) {
let _ = self.async_event_sender.send(event_epoch);
}
#[cfg(not(feature = "tokio"))]
#[inline]
fn notify_async_waiters(&self, _event_epoch: u64) {}
fn increment_waiter(state: &mut MockTimelineState, kind: MockWaiterKind) {
match kind {
MockWaiterKind::Sleep => {
state.sleep_waiters = state.sleep_waiters.saturating_add(1);
}
MockWaiterKind::Deadline => {
state.deadline_waiters = state.deadline_waiters.saturating_add(1);
}
}
}
fn decrement_waiter(state: &mut MockTimelineState, kind: MockWaiterKind) {
match kind {
MockWaiterKind::Sleep => {
state.sleep_waiters = state.sleep_waiters.saturating_sub(1);
}
MockWaiterKind::Deadline => {
state.deadline_waiters = state.deadline_waiters.saturating_sub(1);
}
}
}
#[inline]
fn waiter_count(state: &MockTimelineState, kind: MockWaiterKind) -> usize {
match kind {
MockWaiterKind::Sleep => state.sleep_waiters,
MockWaiterKind::Deadline => state.deadline_waiters,
}
}
}
impl Default for MockTimeline {
#[inline]
fn default() -> Self {
Self::new()
}
}
fn duration_from_nanos_saturating(nanos: u128) -> Duration {
let secs = nanos / 1_000_000_000;
let sub_nanos = (nanos % 1_000_000_000) as u32;
let secs = match u64::try_from(secs) {
Ok(secs) => secs,
Err(_) => return Duration::MAX,
};
Duration::new(secs, sub_nanos)
}
fn next_mock_timeline_id() -> u64 {
loop {
let current = NEXT_MOCK_TIMELINE_ID.load(Ordering::Relaxed);
assert_ne!(current, u64::MAX, "mock timeline id space exhausted");
if NEXT_MOCK_TIMELINE_ID
.compare_exchange_weak(current, current + 1, Ordering::Relaxed, Ordering::Relaxed)
.is_ok()
{
return current;
}
}
}