Skip to main content

agp_service/
timer.rs

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
26    timer_id: u32,
27
28    /// timer type
29    timer_type: TimerType,
30
31    /// constant timer: timer duration
32    /// exponential timer: min timer duration. at every new timer the duration is computers as last_duration * 2
33    duration: Duration,
34
35    /// constant timer: None
36    /// exponential timer: maximum timer duration. once the duration reaches this time it will not be encreased anymore
37    max_duration: Option<Duration>,
38
39    /// if not None, it indicates the maximum number of retryes before call on_failure
40    /// if set to None the timer will go on forever unless cancelled
41    max_retries: Option<u32>,
42
43    /// token used to cancel the timer
44    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// tests
173#[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        // check logs to validate the test
222        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        // check logs to validate the test
304        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        // timeouts after 100ms
356        let expected_msg = "timeout number 1 for timer id 1, retry";
357        assert!(logs_contain(expected_msg));
358
359        // timeouts after 200ms
360        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        // timeouts after 300ms
368        let expected_msg = "timeout number 3 for timer id 1, retry";
369        assert!(logs_contain(expected_msg));
370
371        // timeouts after 400ms
372        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        // timeouts after 500ms
380        let expected_msg = "timeout number 4 for timer id 1, retry";
381        assert!(logs_contain(expected_msg));
382
383        // timeouts after 600ms
384        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        // timeouts after 700ms
392        let expected_msg = "timeout number 6 for timer id 1, stop retry";
393        assert!(logs_contain(expected_msg));
394
395        // stop timer 2 and 3
396        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}