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