1use std::sync::Arc;
2
3use tonic::async_trait;
4
5use tokio::time::{self, Duration};
6use tokio_util::sync::CancellationToken;
7
8use tracing::trace;
9
10#[async_trait]
11pub trait TimerObserver {
12 async fn on_timeout(&self, timer_id: u32, timeouts: u32);
13 async fn on_failure(&self, timer_id: u32, timeouts: u32);
14 async fn on_stop(&self, timer_id: u32);
15}
16
17#[derive(Debug, Clone)]
18pub enum TimerType {
19 Constant = 0,
20 Exponential = 1,
21}
22
23#[derive(Debug)]
24pub struct Timer {
25 timer_id: u32,
27
28 timer_type: TimerType,
30
31 duration: Duration,
34
35 max_duration: Option<Duration>,
38
39 max_retries: Option<u32>,
42
43 cancellation_token: CancellationToken,
45}
46
47impl Timer {
48 pub fn new(
49 timer_id: u32,
50 timer_type: TimerType,
51 duration: Duration,
52 max_duration: Option<Duration>,
53 max_retries: Option<u32>,
54 ) -> Self {
55 Timer {
56 timer_id,
57 timer_type,
58 duration,
59 max_duration,
60 max_retries,
61 cancellation_token: CancellationToken::new(),
62 }
63 }
64
65 pub fn start<T: TimerObserver + Send + Sync + 'static>(&self, observer: Arc<T>) {
66 let timer_id = self.timer_id;
67 let timer_type = self.timer_type.clone();
68 let duration = self.duration;
69 let max_retries = self.max_retries;
70 let max_duration = self.max_duration;
71 let cancellation_token = self.cancellation_token.clone();
72
73 tokio::spawn(async move {
74 let mut retry = 0;
75 let mut timeouts = 0;
76 let mut last_duration = duration;
77
78 trace!("timer {} started", timer_id);
79 loop {
80 let timer_duration = match timer_type {
81 TimerType::Constant => {
82 trace!(
83 "constant timer {}, next in {} ms",
84 timer_id,
85 duration.as_millis()
86 );
87 duration
88 }
89 TimerType::Exponential => {
90 let mut d = duration;
91 if timeouts != 0 {
92 d = last_duration * 2;
93 }
94 match max_duration {
95 None => {
96 trace!(
97 "exponential timer {}, next in {} ms",
98 timer_id,
99 d.as_millis()
100 );
101 last_duration = d;
102 d
103 }
104 Some(max_d) => {
105 if d > max_d {
106 trace!(
107 "exponential timer {}, next in {} ms (use max duration)",
108 timer_id,
109 max_d.as_millis()
110 );
111 last_duration = max_d;
112 max_d
113 } else {
114 trace!(
115 "exponential timer {}, next in {} ms",
116 timer_id,
117 d.as_millis()
118 );
119 last_duration = d;
120 d
121 }
122 }
123 }
124 }
125 };
126
127 let timer = time::sleep(timer_duration);
128 tokio::pin!(timer);
129
130 tokio::select! {
131 _ = timer.as_mut() => {
132 timeouts += 1;
133 match max_retries {
134 Some(max) => {
135 if retry < max {
136 observer.on_timeout(timer_id, timeouts).await
137 } else {
138 observer.on_failure(timer_id, timeouts).await;
139 break;
140 }
141 }
142 None => observer.on_timeout(timer_id, timeouts).await
143 }
144 retry += 1;
145 },
146 _ = cancellation_token.cancelled() => {
147 observer.on_stop(timer_id).await;
148 break;
149 },
150 }
151 }
152 });
153 }
154
155 pub fn stop(&mut self) {
156 self.cancellation_token.cancel();
157 self.cancellation_token = CancellationToken::new();
158 }
159
160 pub fn reset<T: TimerObserver + Send + Sync + 'static>(&mut self, observer: Arc<T>) {
161 self.stop();
162 self.start(observer);
163 }
164}
165
166impl Drop for Timer {
167 fn drop(&mut self) {
168 self.cancellation_token.cancel();
169 }
170}
171
172#[cfg(test)]
174mod tests {
175 use tracing::debug;
176 use tracing_test::traced_test;
177
178 use super::*;
179
180 struct Observer {
181 id: u32,
182 }
183
184 #[async_trait]
185 impl TimerObserver for Observer {
186 async fn on_timeout(&self, timer_id: u32, timeouts: u32) {
187 debug!(
188 "timeout number {} for timer id {}, retry",
189 timeouts, timer_id
190 );
191 }
192
193 async fn on_failure(&self, timer_id: u32, timeouts: u32) {
194 debug!(
195 "timeout number {} for timer id {}, stop retry",
196 timeouts, timer_id
197 );
198 }
199
200 async fn on_stop(&self, timer_id: u32) {
201 debug!("timer id {} cancelled", timer_id);
202 }
203 }
204
205 #[tokio::test]
206 #[traced_test]
207 async fn test_timer() {
208 let o = Arc::new(Observer { id: 10 });
209 let t = Timer::new(
210 o.id,
211 TimerType::Constant,
212 Duration::from_millis(100),
213 None,
214 Some(3),
215 );
216
217 t.start(o);
218
219 time::sleep(Duration::from_millis(500)).await;
220
221 let expected_msg = "timeout number 1 for timer id 10, retry";
223 assert!(logs_contain(expected_msg));
224 let expected_msg = "timeout number 2 for timer id 10, retry";
225 assert!(logs_contain(expected_msg));
226 let expected_msg = "timeout number 3 for timer id 10, retry";
227 assert!(logs_contain(expected_msg));
228 let expected_msg = "timeout number 4 for timer id 10, stop retry";
229 assert!(logs_contain(expected_msg));
230
231 let o = Arc::new(Observer { id: 20 });
232 let t = Timer::new(
233 o.id,
234 TimerType::Exponential,
235 Duration::from_millis(100),
236 Some(Duration::from_millis(400)),
237 Some(3),
238 );
239
240 t.start(o);
241 time::sleep(Duration::from_millis(1200)).await;
242
243 let expected_msg = "exponential timer 20, next in 100 ms";
244 assert!(logs_contain(expected_msg));
245 let expected_msg = "exponential timer 20, next in 200 ms";
246 assert!(logs_contain(expected_msg));
247 let expected_msg = "exponential timer 20, next in 400 ms";
248 assert!(logs_contain(expected_msg));
249 let expected_msg = "exponential timer 20, next in 400 ms (use max duration)";
250 assert!(logs_contain(expected_msg));
251 let expected_msg = "timeout number 4 for timer id 20, stop retry";
252 assert!(logs_contain(expected_msg));
253
254 let o = Arc::new(Observer { id: 30 });
255 let mut t = Timer::new(
256 o.id,
257 TimerType::Exponential,
258 Duration::from_millis(100),
259 None,
260 None,
261 );
262
263 t.start(o);
264
265 time::sleep(Duration::from_millis(2000)).await;
266 t.stop();
267 time::sleep(Duration::from_millis(500)).await;
268 let expected_msg = "exponential timer 30, next in 100 ms";
269 assert!(logs_contain(expected_msg));
270 let expected_msg = "exponential timer 30, next in 200 ms";
271 assert!(logs_contain(expected_msg));
272 let expected_msg = "exponential timer 30, next in 400 ms";
273 assert!(logs_contain(expected_msg));
274 let expected_msg = "exponential timer 30, next in 800 ms";
275 assert!(logs_contain(expected_msg));
276 let expected_msg = "exponential timer 30, next in 1600 ms";
277 assert!(logs_contain(expected_msg));
278 let expected_msg = "timer id 30 cancelled";
279 assert!(logs_contain(expected_msg))
280 }
281
282 #[tokio::test]
283 #[traced_test]
284 async fn test_timer_stop() {
285 let o = Arc::new(Observer { id: 10 });
286
287 let mut t = Timer::new(
288 o.id,
289 TimerType::Constant,
290 Duration::from_millis(100),
291 None,
292 Some(5),
293 );
294
295 t.start(o);
296
297 time::sleep(Duration::from_millis(350)).await;
298
299 t.stop();
300
301 time::sleep(Duration::from_millis(500)).await;
302
303 let expected_msg = "timeout number 1 for timer id 10, retry";
305 assert!(logs_contain(expected_msg));
306 let expected_msg = "timeout number 2 for timer id 10, retry";
307 assert!(logs_contain(expected_msg));
308 let expected_msg = "timeout number 3 for timer id 10, retry";
309 assert!(logs_contain(expected_msg));
310 let expected_msg = "timer id 10 cancelled";
311 assert!(logs_contain(expected_msg));
312 }
313
314 #[tokio::test]
315 #[traced_test]
316 async fn test_multiple_timers() {
317 let o1 = Arc::new(Observer { id: 1 });
318 let o2 = Arc::new(Observer { id: 2 });
319 let o3 = Arc::new(Observer { id: 3 });
320
321 let mut t1 = Timer::new(
322 o1.id,
323 TimerType::Constant,
324 Duration::from_millis(100),
325 None,
326 Some(5),
327 );
328 let mut t2 = Timer::new(
329 o2.id,
330 TimerType::Constant,
331 Duration::from_millis(200),
332 None,
333 Some(5),
334 );
335 let mut t3 = Timer::new(
336 o3.id,
337 TimerType::Constant,
338 Duration::from_millis(200),
339 None,
340 Some(5),
341 );
342
343 t1.start(o1);
344 t2.start(o2);
345 t3.start(o3);
346
347 time::sleep(Duration::from_millis(700)).await;
348
349 t1.stop();
350 t2.stop();
351 t3.stop();
352
353 time::sleep(Duration::from_millis(500)).await;
354
355 let expected_msg = "timeout number 1 for timer id 1, retry";
357 assert!(logs_contain(expected_msg));
358
359 let expected_msg = "timeout number 1 for timer id 2, retry";
361 assert!(logs_contain(expected_msg));
362 let expected_msg = "timeout number 1 for timer id 3, retry";
363 assert!(logs_contain(expected_msg));
364 let expected_msg = "timeout number 2 for timer id 1, retry";
365 assert!(logs_contain(expected_msg));
366
367 let expected_msg = "timeout number 3 for timer id 1, retry";
369 assert!(logs_contain(expected_msg));
370
371 let expected_msg = "timeout number 2 for timer id 2, retry";
373 assert!(logs_contain(expected_msg));
374 let expected_msg = "timeout number 2 for timer id 3, retry";
375 assert!(logs_contain(expected_msg));
376 let expected_msg = "timeout number 4 for timer id 1, retry";
377 assert!(logs_contain(expected_msg));
378
379 let expected_msg = "timeout number 4 for timer id 1, retry";
381 assert!(logs_contain(expected_msg));
382
383 let expected_msg = "timeout number 3 for timer id 2, retry";
385 assert!(logs_contain(expected_msg));
386 let expected_msg = "timeout number 3 for timer id 3, retry";
387 assert!(logs_contain(expected_msg));
388 let expected_msg = "timeout number 5 for timer id 1, retry";
389 assert!(logs_contain(expected_msg));
390
391 let expected_msg = "timeout number 6 for timer id 1, stop retry";
393 assert!(logs_contain(expected_msg));
394
395 let expected_msg = "timer id 2 cancelled";
397 assert!(logs_contain(expected_msg));
398 let expected_msg = "timer id 3 cancelled";
399 assert!(logs_contain(expected_msg));
400 }
401
402 #[tokio::test]
403 #[traced_test]
404 async fn test_timer_reset() {
405 let o = Arc::new(Observer { id: 10 });
406
407 let mut t = Timer::new(
408 o.id,
409 TimerType::Constant,
410 Duration::from_millis(100),
411 None,
412 Some(5),
413 );
414
415 t.start(o.clone());
416
417 time::sleep(Duration::from_millis(350)).await;
418
419 let expected_msg = "timeout number 3 for timer id 10, retry";
420 assert!(logs_contain(expected_msg));
421
422 t.reset(o.clone());
423
424 time::sleep(Duration::from_millis(250)).await;
425
426 let expected_msg = "timeout number 2 for timer id 10, retry";
427 assert!(logs_contain(expected_msg));
428
429 t.reset(o.clone());
430
431 time::sleep(Duration::from_millis(700)).await;
432
433 let expected_msg = "timeout number 6 for timer id 10, stop retry";
434 assert!(logs_contain(expected_msg));
435
436 t.reset(o);
437
438 time::sleep(Duration::from_millis(700)).await;
439
440 let expected_msg = "timeout number 6 for timer id 10, stop retry";
441 assert!(logs_contain(expected_msg));
442 }
443
444 #[tokio::test]
445 #[traced_test]
446 async fn test_timer_reset_without_start() {
447 let o = Arc::new(Observer { id: 10 });
448
449 let mut t = Timer::new(
450 o.id,
451 TimerType::Constant,
452 Duration::from_millis(100),
453 None,
454 Some(5),
455 );
456
457 t.reset(o);
458
459 time::sleep(Duration::from_millis(350)).await;
460
461 let expected_msg = "timeout number 3 for timer id 10, retry";
462 assert!(logs_contain(expected_msg));
463 }
464}