use chrono::{DateTime, Utc};
use parking_lot::Mutex;
use std::{
cmp::Ordering as CmpOrdering,
collections::BinaryHeap,
sync::atomic::{AtomicI64, Ordering},
task::Waker,
time::Duration,
};
fn truncate_to_millis(time: DateTime<Utc>) -> DateTime<Utc> {
DateTime::from_timestamp_millis(time.timestamp_millis()).expect("valid timestamp")
}
fn datetime_from_millis(ms: i64) -> DateTime<Utc> {
DateTime::from_timestamp_millis(ms).unwrap_or(if ms > 0 {
DateTime::<Utc>::MAX_UTC
} else {
DateTime::<Utc>::MIN_UTC
})
}
pub(crate) struct ManualClock {
current_ms: AtomicI64,
pending_wakes: Mutex<BinaryHeap<PendingWake>>,
coalesce_wakes: Mutex<Vec<PendingWake>>,
}
pub(crate) struct PendingWake {
wake_at_ms: i64,
sleep_id: u64,
waker: Waker,
}
impl PartialEq for PendingWake {
fn eq(&self, other: &Self) -> bool {
self.wake_at_ms == other.wake_at_ms && self.sleep_id == other.sleep_id
}
}
impl Eq for PendingWake {}
impl PartialOrd for PendingWake {
fn partial_cmp(&self, other: &Self) -> Option<CmpOrdering> {
Some(self.cmp(other))
}
}
impl Ord for PendingWake {
fn cmp(&self, other: &Self) -> CmpOrdering {
match other.wake_at_ms.cmp(&self.wake_at_ms) {
CmpOrdering::Equal => other.sleep_id.cmp(&self.sleep_id),
ord => ord,
}
}
}
impl ManualClock {
pub fn new() -> Self {
Self::new_at(Utc::now())
}
pub fn new_at(start_at: DateTime<Utc>) -> Self {
Self {
current_ms: AtomicI64::new(truncate_to_millis(start_at).timestamp_millis()),
pending_wakes: Mutex::new(BinaryHeap::new()),
coalesce_wakes: Mutex::new(Vec::new()),
}
}
pub fn now(&self) -> DateTime<Utc> {
datetime_from_millis(self.now_ms())
}
pub fn now_ms(&self) -> i64 {
self.current_ms.load(Ordering::SeqCst)
}
pub fn register_wake(&self, wake_at_ms: i64, sleep_id: u64, waker: Waker) {
let mut pending = self.pending_wakes.lock();
pending.push(PendingWake {
wake_at_ms,
sleep_id,
waker,
});
}
pub fn register_coalesce_wake(&self, wake_at_ms: i64, sleep_id: u64, waker: Waker) {
let mut coalesce = self.coalesce_wakes.lock();
coalesce.push(PendingWake {
wake_at_ms,
sleep_id,
waker,
});
}
pub fn cancel_wake(&self, sleep_id: u64) {
{
let mut pending = self.pending_wakes.lock();
let entries: Vec<_> = pending.drain().filter(|w| w.sleep_id != sleep_id).collect();
pending.extend(entries);
}
{
let mut coalesce = self.coalesce_wakes.lock();
coalesce.retain(|w| w.sleep_id != sleep_id);
}
}
pub fn next_wake_time(&self) -> Option<i64> {
let pending = self.pending_wakes.lock();
pending.peek().map(|w| w.wake_at_ms)
}
pub fn wake_tasks_at(&self, up_to_ms: i64) -> usize {
let wakers: Vec<Waker> = {
let mut pending = self.pending_wakes.lock();
let mut wakers = Vec::new();
while let Some(wake) = pending.peek() {
if wake.wake_at_ms > up_to_ms {
break;
}
let wake = pending.pop().unwrap();
wakers.push(wake.waker);
}
wakers
};
let count = wakers.len();
for waker in wakers {
waker.wake();
}
count
}
pub async fn advance(&self, duration: Duration) -> usize {
let start_ms = self.current_ms.load(Ordering::SeqCst);
let added = i64::try_from(duration.as_millis()).unwrap_or(i64::MAX);
let target_ms = start_ms.saturating_add(added);
let mut total_woken = 0;
loop {
let next_wake_ms = self.next_wake_time();
match next_wake_ms {
Some(wake_ms) if wake_ms <= target_ms => {
self.current_ms.store(wake_ms, Ordering::SeqCst);
let woken = self.wake_tasks_at(wake_ms);
total_woken += woken;
tokio::task::yield_now().await;
}
_ => {
self.current_ms.store(target_ms, Ordering::SeqCst);
break;
}
}
}
let coalesce_woken = self.wake_coalesce_tasks_at(target_ms);
if coalesce_woken > 0 {
total_woken += coalesce_woken;
tokio::task::yield_now().await;
}
total_woken
}
pub async fn advance_to_next_wake(&self) -> Option<DateTime<Utc>> {
let next_regular = self.next_wake_time();
let next_coalesce = self.next_coalesce_wake_time();
let next_wake_ms = match (next_regular, next_coalesce) {
(Some(r), Some(c)) => Some(r.min(c)),
(Some(r), None) => Some(r),
(None, Some(c)) => Some(c),
(None, None) => None,
}?;
self.current_ms.store(next_wake_ms, Ordering::SeqCst);
self.wake_tasks_at(next_wake_ms);
self.wake_coalesce_tasks_at(next_wake_ms);
tokio::task::yield_now().await;
Some(datetime_from_millis(next_wake_ms))
}
pub fn wake_coalesce_tasks_at(&self, up_to_ms: i64) -> usize {
let wakers: Vec<Waker> = {
let mut coalesce = self.coalesce_wakes.lock();
let mut wakers = Vec::new();
let mut remaining = Vec::new();
for wake in coalesce.drain(..) {
if wake.wake_at_ms <= up_to_ms {
wakers.push(wake.waker);
} else {
remaining.push(wake);
}
}
*coalesce = remaining;
wakers
};
let count = wakers.len();
for waker in wakers {
waker.wake();
}
count
}
fn next_coalesce_wake_time(&self) -> Option<i64> {
let coalesce = self.coalesce_wakes.lock();
coalesce.iter().map(|w| w.wake_at_ms).min()
}
pub fn pending_wake_count(&self) -> usize {
self.pending_wakes.lock().len() + self.coalesce_wakes.lock().len()
}
}
#[cfg(test)]
mod tests {
use super::*;
use proptest::prelude::*;
#[test]
fn test_manual_now() {
let clock = ManualClock::new();
let start = clock.now();
std::thread::sleep(Duration::from_millis(10));
assert_eq!(clock.now(), start);
}
#[test]
fn test_pending_wake_ordering() {
let clock = ManualClock::new();
let waker = futures::task::noop_waker();
clock.register_wake(3000, 1, waker.clone());
clock.register_wake(1000, 2, waker.clone());
clock.register_wake(2000, 3, waker);
assert_eq!(clock.next_wake_time(), Some(1000));
clock.wake_tasks_at(1000);
assert_eq!(clock.next_wake_time(), Some(2000));
clock.wake_tasks_at(2000);
assert_eq!(clock.next_wake_time(), Some(3000));
}
proptest! {
#[test]
fn next_wake_time_tracks_min_after_each_insert(
raw in proptest::collection::vec((any::<i64>(), any::<u64>()), 0..30),
) {
let clock = ManualClock::new();
let mut seen = std::collections::HashSet::<u64>::new();
let mut running_min: Option<i64> = None;
let mut count = 0usize;
for (ms, id) in raw {
if !seen.insert(id) {
continue;
}
clock.register_wake(ms, id, futures::task::noop_waker());
running_min = Some(running_min.map_or(ms, |m| m.min(ms)));
count += 1;
prop_assert_eq!(clock.next_wake_time(), running_min);
prop_assert_eq!(clock.pending_wake_count(), count);
}
}
#[test]
fn wakes_drain_in_nondecreasing_order(
raw in proptest::collection::vec((any::<i64>(), any::<u64>()), 0..30),
) {
let clock = ManualClock::new();
let mut seen = std::collections::HashSet::<u64>::new();
let mut inserted = 0usize;
for (ms, id) in raw {
if seen.insert(id) {
clock.register_wake(ms, id, futures::task::noop_waker());
inserted += 1;
}
}
let mut prev: Option<i64> = None;
let mut total = 0usize;
while let Some(cur) = clock.next_wake_time() {
if let Some(p) = prev {
prop_assert!(cur >= p, "wake time decreased: {cur} < {p}");
}
prev = Some(cur);
total += clock.wake_tasks_at(cur);
}
prop_assert_eq!(total, inserted);
prop_assert_eq!(clock.pending_wake_count(), 0);
}
#[test]
fn cancel_wake_removes_exactly_one(
raw in proptest::collection::vec((any::<i64>(), any::<u64>()), 1..20),
) {
let clock = ManualClock::new();
let mut seen = std::collections::HashSet::<u64>::new();
let mut ids: Vec<u64> = Vec::new();
for (ms, id) in raw {
if seen.insert(id) {
clock.register_wake(ms, id, futures::task::noop_waker());
ids.push(id);
}
}
let before = clock.pending_wake_count();
let victim = ids[0];
clock.cancel_wake(victim);
prop_assert_eq!(clock.pending_wake_count(), before - 1);
clock.cancel_wake(victim);
prop_assert_eq!(clock.pending_wake_count(), before - 1);
let mut total = 0usize;
while let Some(cur) = clock.next_wake_time() {
total += clock.wake_tasks_at(cur);
}
prop_assert_eq!(total, before - 1);
}
#[test]
fn now_never_panics_for_any_stored_ms(ms in any::<i64>()) {
let clock = ManualClock::new();
clock
.current_ms
.store(ms, std::sync::atomic::Ordering::SeqCst);
let _ = clock.now();
}
}
#[tokio::test]
async fn advance_huge_duration_does_not_panic_in_now() {
let clock = ManualClock::new();
let start = clock.now();
clock.advance(Duration::MAX).await;
let after = clock.now();
assert!(after >= start);
}
}