moirai_async/timer/
wheel.rs1use std::cmp::Ordering;
4use std::collections::{BinaryHeap, HashSet};
5use std::task::Waker;
6use std::time::Instant;
7
8#[derive(Debug)]
10struct TimerEntry {
11 id: u64,
12 deadline: Instant,
13 waker: Option<Waker>,
14}
15
16impl PartialEq for TimerEntry {
17 fn eq(&self, other: &Self) -> bool {
18 self.id == other.id && self.deadline == other.deadline
19 }
20}
21
22impl Eq for TimerEntry {}
23
24impl PartialOrd for TimerEntry {
25 fn partial_cmp(&self, other: &Self) -> Option<Ordering> {
26 Some(self.cmp(other))
27 }
28}
29
30impl Ord for TimerEntry {
31 fn cmp(&self, other: &Self) -> Ordering {
32 other
33 .deadline
34 .cmp(&self.deadline)
35 .then_with(|| other.id.cmp(&self.id))
36 }
37}
38
39pub struct TimerWheel {
41 timers: BinaryHeap<TimerEntry>,
42 active: HashSet<u64>,
46 cancelled: HashSet<u64>,
50 next_id: u64,
51}
52
53pub enum TimerCommand {
55 Schedule {
57 timer_id: u64,
59 deadline: Instant,
61 waker: Waker,
63 },
64 Cancel {
66 timer_id: u64,
68 },
69 Reschedule {
71 timer_id: u64,
73 new_deadline: Instant,
75 },
76}
77
78impl TimerWheel {
79 pub fn new() -> Self {
81 Self {
82 timers: BinaryHeap::new(),
83 active: HashSet::new(),
84 cancelled: HashSet::new(),
85 next_id: 1,
86 }
87 }
88
89 pub fn schedule(&mut self, deadline: Instant, waker: Waker) -> u64 {
91 let timer_id = self.next_id;
92 self.next_id = self.next_id.saturating_add(1);
93
94 self.active.insert(timer_id);
95 self.timers.push(TimerEntry {
96 id: timer_id,
97 deadline,
98 waker: Some(waker),
99 });
100
101 timer_id
102 }
103
104 pub fn cancel(&mut self, timer_id: u64) -> bool {
110 if !self.active.contains(&timer_id) {
111 return false;
112 }
113
114 self.cancelled.insert(timer_id)
115 }
116
117 fn pop_head(&mut self) -> TimerEntry {
119 let entry = self.timers.pop().expect("entry existed after peek");
120 self.active.remove(&entry.id);
121 entry
122 }
123
124 fn drain_cancelled_head(&mut self) {
127 while self
128 .timers
129 .peek()
130 .is_some_and(|entry| self.cancelled.contains(&entry.id))
131 {
132 let entry = self.pop_head();
133 self.cancelled.remove(&entry.id);
134 }
135 }
136
137 pub fn poll_expired(&mut self) -> usize {
139 let now = Instant::now();
140 let mut expired_count = 0;
141
142 self.drain_cancelled_head();
143
144 while self
145 .timers
146 .peek()
147 .is_some_and(|entry| entry.deadline <= now)
148 {
149 let mut expired = self.pop_head();
150 if self.cancelled.remove(&expired.id) {
151 continue;
152 }
153
154 if let Some(waker) = expired.waker.take() {
155 waker.wake();
156 expired_count += 1;
157 }
158 }
159
160 expired_count
161 }
162
163 pub fn next_expiration(&mut self) -> Option<Instant> {
165 self.drain_cancelled_head();
166 self.timers.peek().map(|entry| entry.deadline)
167 }
168
169 pub fn timer_count(&self) -> usize {
171 self.active.len() - self.cancelled.len()
173 }
174
175 #[cfg(test)]
177 fn tombstone_count(&self) -> usize {
178 self.cancelled.len()
179 }
180}
181
182impl Default for TimerWheel {
183 fn default() -> Self {
184 Self::new()
185 }
186}
187
188#[cfg(test)]
189mod tests {
190 use super::TimerWheel;
191 use std::sync::{
192 Arc,
193 atomic::{AtomicUsize, Ordering},
194 };
195 use std::task::{Wake, Waker};
196 use std::time::{Duration, Instant};
197
198 struct CountingWake {
199 count: Arc<AtomicUsize>,
200 }
201
202 impl Wake for CountingWake {
203 fn wake(self: Arc<Self>) {
204 self.count.fetch_add(1, Ordering::Release);
205 }
206
207 fn wake_by_ref(self: &Arc<Self>) {
208 self.count.fetch_add(1, Ordering::Release);
209 }
210 }
211
212 fn counting_waker(count: Arc<AtomicUsize>) -> Waker {
213 Waker::from(Arc::new(CountingWake { count }))
214 }
215
216 #[test]
217 fn timer_wheel_cancelled_timer_does_not_wake() {
218 let mut wheel = TimerWheel::new();
219 let wake_count = Arc::new(AtomicUsize::new(0));
220 let timer_id = wheel.schedule(
221 Instant::now()
222 .checked_sub(Duration::from_millis(1))
223 .expect("invariant: process uptime exceeds 1ms"),
224 counting_waker(Arc::clone(&wake_count)),
225 );
226
227 assert!(wheel.cancel(timer_id));
228 assert_eq!(wheel.poll_expired(), 0);
229 assert_eq!(wake_count.load(Ordering::Acquire), 0);
230 assert_eq!(wheel.timer_count(), 0);
231 }
232
233 #[test]
234 fn timer_wheel_poll_wakes_only_uncancelled_expired_timers() {
235 let mut wheel = TimerWheel::new();
236 let wake_count = Arc::new(AtomicUsize::new(0));
237 let deadline = Instant::now()
238 .checked_sub(Duration::from_millis(1))
239 .expect("invariant: process uptime exceeds 1ms");
240 let cancelled = wheel.schedule(deadline, counting_waker(Arc::clone(&wake_count)));
241 let active = wheel.schedule(deadline, counting_waker(Arc::clone(&wake_count)));
242
243 assert!(wheel.cancel(cancelled));
244 assert_ne!(cancelled, active);
245 assert_eq!(wheel.poll_expired(), 1);
246 assert_eq!(wake_count.load(Ordering::Acquire), 1);
247 assert_eq!(wheel.timer_count(), 0);
248 }
249
250 #[test]
251 fn timer_wheel_cancel_after_expiry_does_not_leak_tombstones() {
252 let mut wheel = TimerWheel::new();
256 let wake_count = Arc::new(AtomicUsize::new(0));
257 let id = wheel.schedule(
258 Instant::now()
259 .checked_sub(Duration::from_millis(1))
260 .expect("invariant: process uptime exceeds 1ms"),
261 counting_waker(Arc::clone(&wake_count)),
262 );
263
264 assert_eq!(wheel.poll_expired(), 1);
265 assert_eq!(wheel.tombstone_count(), 0);
266
267 assert!(!wheel.cancel(id));
270 assert!(!wheel.cancel(9_999));
271 assert_eq!(wheel.tombstone_count(), 0);
272 assert_eq!(wheel.timer_count(), 0);
273 }
274
275 #[test]
276 fn timer_wheel_cancelled_then_fired_reclaims_tombstone() {
277 let mut wheel = TimerWheel::new();
280 let wake_count = Arc::new(AtomicUsize::new(0));
281 let id = wheel.schedule(
282 Instant::now()
283 .checked_sub(Duration::from_millis(1))
284 .expect("invariant: process uptime exceeds 1ms"),
285 counting_waker(Arc::clone(&wake_count)),
286 );
287
288 assert!(wheel.cancel(id));
289 assert_eq!(wheel.tombstone_count(), 1);
290
291 assert_eq!(wheel.poll_expired(), 0);
292 assert_eq!(wake_count.load(Ordering::Acquire), 0);
293 assert_eq!(wheel.tombstone_count(), 0, "tombstone must be reclaimed");
294 assert_eq!(wheel.timer_count(), 0);
295 }
296}