Skip to main content

ntex_util/services/
keepalive.rs

1use std::{cell::Cell, convert::Infallible, fmt, task::Poll, time};
2
3use ntex_service::{Ctx, Service, ServiceFactory};
4
5use crate::time::{Millis, Sleep, now, sleep};
6
7/// `KeepAlive` service factory
8///
9/// Controls min time between requests.
10pub struct KeepAlive<F, E>
11where
12    F: Fn() -> E + Clone,
13{
14    f: F,
15    ka: Millis,
16}
17
18impl<F, E> KeepAlive<F, E>
19where
20    F: Fn() -> E + Clone,
21{
22    /// Construct `KeepAlive` service factory.
23    ///
24    /// ka - keep-alive timeout
25    /// err - error factory function
26    pub fn new(ka: Millis, f: F) -> Self {
27        KeepAlive { f, ka }
28    }
29}
30
31impl<F, E> Clone for KeepAlive<F, E>
32where
33    F: Fn() -> E + Clone,
34{
35    fn clone(&self) -> Self {
36        KeepAlive {
37            f: self.f.clone(),
38            ka: self.ka,
39        }
40    }
41}
42
43impl<F, E> fmt::Debug for KeepAlive<F, E>
44where
45    F: Fn() -> E + Clone,
46{
47    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
48        f.debug_struct("KeepAlive")
49            .field("ka", &self.ka)
50            .field("f", &std::any::type_name::<F>())
51            .finish()
52    }
53}
54
55impl<F, E, St, Req> ServiceFactory<St, Req> for KeepAlive<F, E>
56where
57    F: Fn() -> E + Clone,
58{
59    type Res = Req;
60    type Error = E;
61
62    type Service = KeepAliveService<F, E>;
63    type InitError = Infallible;
64
65    #[inline]
66    async fn create(&self, _: &St) -> Result<Self::Service, Self::InitError> {
67        Ok(KeepAliveService::new(self.ka, self.f.clone()))
68    }
69}
70
71pub struct KeepAliveService<F, E>
72where
73    F: Fn() -> E,
74{
75    f: F,
76    dur: time::Duration,
77    sleep: Sleep,
78    expire: Cell<time::Instant>,
79}
80
81impl<F, E> KeepAliveService<F, E>
82where
83    F: Fn() -> E,
84{
85    pub fn new(dur: Millis, f: F) -> Self {
86        let expire = Cell::new(now());
87
88        KeepAliveService {
89            f,
90            expire,
91            sleep: sleep(dur),
92            dur: time::Duration::from(dur),
93        }
94    }
95}
96
97impl<F, E> fmt::Debug for KeepAliveService<F, E>
98where
99    F: Fn() -> E,
100{
101    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
102        f.debug_struct("KeepAliveService")
103            .field("dur", &self.dur)
104            .field("expire", &self.expire)
105            .field("f", &std::any::type_name::<F>())
106            .finish()
107    }
108}
109
110impl<F, E, St, Req> Service<St, Req> for KeepAliveService<F, E>
111where
112    F: Fn() -> E,
113{
114    type Res = Req;
115    type Error = E;
116
117    async fn ready(&self, ctx: Ctx<'_, Self, St>) -> Result<(), Self::Error> {
118        let expire = self.expire.get() + self.dur;
119        if expire <= now() {
120            Err((self.f)())
121        } else {
122            ctx.poll_once(|cx| {
123                loop {
124                    match self.sleep.poll_elapsed(cx) {
125                        Poll::Ready(()) => {
126                            let now = now();
127                            let expire = self.expire.get() + self.dur;
128                            if expire <= now {
129                                return Err((self.f)());
130                            }
131                            let expire = expire - now;
132
133                            // sleep must be reset to non zero duration,
134                            // otherwise it stays in elapsed state and waker
135                            // never gets registered
136                            let expire: u32 = expire.as_millis().try_into().unwrap_or(u32::MAX);
137                            self.sleep.reset(Millis(expire.max(1)));
138                        }
139                        Poll::Pending => return Ok(()),
140                    }
141                }
142            })
143        }
144    }
145
146    #[inline]
147    async fn call(&self, req: Req, _: Ctx<'_, Self, St>) -> Result<Req, E> {
148        self.expire.set(now());
149        Ok(req)
150    }
151}
152
153#[cfg(test)]
154mod tests {
155    use std::{pin::Pin, task::Context, task::Poll, task::ready};
156
157    use ntex_service::{Pipeline, boxed, factory};
158
159    use super::*;
160    use crate::{channel::oneshot, spawn};
161
162    #[derive(Debug, PartialEq)]
163    struct TestErr;
164
165    struct Dispatcher {
166        p: Pipeline<usize, usize, TestErr>,
167        tx: Option<oneshot::Sender<()>>,
168    }
169
170    impl Future for Dispatcher {
171        type Output = ();
172
173        fn poll(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
174            let mut this = self.as_mut();
175            if ready!(this.p.poll_ready(cx)).is_err() {
176                if let Some(tx) = this.tx.take() {
177                    let _ = tx.send(());
178                }
179                Poll::Ready(())
180            } else {
181                Poll::Pending
182            }
183        }
184    }
185
186    #[ntex::test]
187    async fn test_ka() {
188        let factory = factory::<_, (), usize>(KeepAlive::new(Millis(100), || TestErr));
189        assert!(format!("{factory:?}").contains("KeepAlive"));
190        let _ = factory.clone();
191
192        let svc = factory.create(&()).await.unwrap();
193        assert!(format!("{svc:?}").contains("KeepAliveService"));
194
195        let p = Pipeline::new((), boxed::service(svc));
196        assert_eq!(p.call(1usize).await, Ok(1usize));
197        let svc = p.bind();
198
199        let (tx, rx) = oneshot::channel();
200        spawn(Dispatcher { p, tx: Some(tx) }).detach();
201
202        sleep(Millis(25)).await;
203        assert_eq!(svc.call(1usize).await, Ok(1usize));
204        sleep(Millis(100)).await;
205
206        let res = rx.await;
207        assert_eq!(res, Ok(()));
208        assert_eq!(svc.ready().await, Err(TestErr));
209    }
210
211    #[ntex::test]
212    async fn test_ka_sub_millis() {
213        let svc = std::rc::Rc::new(KeepAliveService::new(Millis(100), || TestErr));
214
215        // less than millisecond is left before expiration
216        svc.expire
217            .set(now().checked_sub(svc.dur).unwrap() + time::Duration::from_micros(500));
218        svc.sleep.elapse();
219
220        let p = Pipeline::<usize, usize, TestErr>::new((), svc.clone()).bind();
221        assert_eq!(p.ready().await, Ok(()));
222
223        // timer has to be re-armed, otherwise waker never gets registered
224        // and service readiness never resolves
225        assert!(!svc.sleep.is_elapsed());
226    }
227}