use std::cmp::Ordering;
use std::collections::{BinaryHeap, HashSet};
use std::task::Waker;
use std::time::Instant;
#[derive(Debug)]
struct TimerEntry {
id: u64,
deadline: Instant,
waker: Option<Waker>,
}
impl PartialEq for TimerEntry {
fn eq(&self, other: &Self) -> bool {
self.id == other.id && self.deadline == other.deadline
}
}
impl Eq for TimerEntry {}
impl PartialOrd for TimerEntry {
fn partial_cmp(&self, other: &Self) -> Option<Ordering> {
Some(self.cmp(other))
}
}
impl Ord for TimerEntry {
fn cmp(&self, other: &Self) -> Ordering {
other
.deadline
.cmp(&self.deadline)
.then_with(|| other.id.cmp(&self.id))
}
}
pub struct TimerWheel {
timers: BinaryHeap<TimerEntry>,
active: HashSet<u64>,
cancelled: HashSet<u64>,
next_id: u64,
}
pub enum TimerCommand {
Schedule {
timer_id: u64,
deadline: Instant,
waker: Waker,
},
Cancel {
timer_id: u64,
},
Reschedule {
timer_id: u64,
new_deadline: Instant,
},
}
impl TimerWheel {
pub fn new() -> Self {
Self {
timers: BinaryHeap::new(),
active: HashSet::new(),
cancelled: HashSet::new(),
next_id: 1,
}
}
pub fn schedule(&mut self, deadline: Instant, waker: Waker) -> u64 {
let timer_id = self.next_id;
self.next_id = self.next_id.saturating_add(1);
self.active.insert(timer_id);
self.timers.push(TimerEntry {
id: timer_id,
deadline,
waker: Some(waker),
});
timer_id
}
pub fn cancel(&mut self, timer_id: u64) -> bool {
if !self.active.contains(&timer_id) {
return false;
}
self.cancelled.insert(timer_id)
}
fn pop_head(&mut self) -> TimerEntry {
let entry = self.timers.pop().expect("entry existed after peek");
self.active.remove(&entry.id);
entry
}
fn drain_cancelled_head(&mut self) {
while self
.timers
.peek()
.is_some_and(|entry| self.cancelled.contains(&entry.id))
{
let entry = self.pop_head();
self.cancelled.remove(&entry.id);
}
}
pub fn poll_expired(&mut self) -> usize {
let now = Instant::now();
let mut expired_count = 0;
self.drain_cancelled_head();
while self
.timers
.peek()
.is_some_and(|entry| entry.deadline <= now)
{
let mut expired = self.pop_head();
if self.cancelled.remove(&expired.id) {
continue;
}
if let Some(waker) = expired.waker.take() {
waker.wake();
expired_count += 1;
}
}
expired_count
}
pub fn next_expiration(&mut self) -> Option<Instant> {
self.drain_cancelled_head();
self.timers.peek().map(|entry| entry.deadline)
}
pub fn timer_count(&self) -> usize {
self.active.len() - self.cancelled.len()
}
#[cfg(test)]
fn tombstone_count(&self) -> usize {
self.cancelled.len()
}
}
impl Default for TimerWheel {
fn default() -> Self {
Self::new()
}
}
#[cfg(test)]
mod tests {
use super::TimerWheel;
use std::sync::{
Arc,
atomic::{AtomicUsize, Ordering},
};
use std::task::{Wake, Waker};
use std::time::{Duration, Instant};
struct CountingWake {
count: Arc<AtomicUsize>,
}
impl Wake for CountingWake {
fn wake(self: Arc<Self>) {
self.count.fetch_add(1, Ordering::Release);
}
fn wake_by_ref(self: &Arc<Self>) {
self.count.fetch_add(1, Ordering::Release);
}
}
fn counting_waker(count: Arc<AtomicUsize>) -> Waker {
Waker::from(Arc::new(CountingWake { count }))
}
#[test]
fn timer_wheel_cancelled_timer_does_not_wake() {
let mut wheel = TimerWheel::new();
let wake_count = Arc::new(AtomicUsize::new(0));
let timer_id = wheel.schedule(
Instant::now()
.checked_sub(Duration::from_millis(1))
.expect("invariant: process uptime exceeds 1ms"),
counting_waker(Arc::clone(&wake_count)),
);
assert!(wheel.cancel(timer_id));
assert_eq!(wheel.poll_expired(), 0);
assert_eq!(wake_count.load(Ordering::Acquire), 0);
assert_eq!(wheel.timer_count(), 0);
}
#[test]
fn timer_wheel_poll_wakes_only_uncancelled_expired_timers() {
let mut wheel = TimerWheel::new();
let wake_count = Arc::new(AtomicUsize::new(0));
let deadline = Instant::now()
.checked_sub(Duration::from_millis(1))
.expect("invariant: process uptime exceeds 1ms");
let cancelled = wheel.schedule(deadline, counting_waker(Arc::clone(&wake_count)));
let active = wheel.schedule(deadline, counting_waker(Arc::clone(&wake_count)));
assert!(wheel.cancel(cancelled));
assert_ne!(cancelled, active);
assert_eq!(wheel.poll_expired(), 1);
assert_eq!(wake_count.load(Ordering::Acquire), 1);
assert_eq!(wheel.timer_count(), 0);
}
#[test]
fn timer_wheel_cancel_after_expiry_does_not_leak_tombstones() {
let mut wheel = TimerWheel::new();
let wake_count = Arc::new(AtomicUsize::new(0));
let id = wheel.schedule(
Instant::now()
.checked_sub(Duration::from_millis(1))
.expect("invariant: process uptime exceeds 1ms"),
counting_waker(Arc::clone(&wake_count)),
);
assert_eq!(wheel.poll_expired(), 1);
assert_eq!(wheel.tombstone_count(), 0);
assert!(!wheel.cancel(id));
assert!(!wheel.cancel(9_999));
assert_eq!(wheel.tombstone_count(), 0);
assert_eq!(wheel.timer_count(), 0);
}
#[test]
fn timer_wheel_cancelled_then_fired_reclaims_tombstone() {
let mut wheel = TimerWheel::new();
let wake_count = Arc::new(AtomicUsize::new(0));
let id = wheel.schedule(
Instant::now()
.checked_sub(Duration::from_millis(1))
.expect("invariant: process uptime exceeds 1ms"),
counting_waker(Arc::clone(&wake_count)),
);
assert!(wheel.cancel(id));
assert_eq!(wheel.tombstone_count(), 1);
assert_eq!(wheel.poll_expired(), 0);
assert_eq!(wake_count.load(Ordering::Acquire), 0);
assert_eq!(wheel.tombstone_count(), 0, "tombstone must be reclaimed");
assert_eq!(wheel.timer_count(), 0);
}
}