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
7pub 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 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
73pub 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 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 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 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 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 assert!(!svc.sleep.is_elapsed());
237 }
238}