Skip to main content

ntex_util/services/
onerequest.rs

1//! Middleware that permits one service call at a time.
2use std::{cell::Cell, future::poll_fn, task::Poll};
3
4use ntex_service::{Ctx, Middleware, Service};
5
6use crate::channel::condition::Condition;
7
8/// Middleware that serializes calls to its wrapped service.
9#[derive(Copy, Clone, Default, Debug)]
10pub struct OneRequest;
11
12impl<S, St> Middleware<S, St> for OneRequest {
13    type Service = OneRequestService<S>;
14
15    fn create(&self, _: &St, service: S) -> Self::Service {
16        OneRequestService {
17            service,
18            ready: Cell::new(true),
19            waiters: Condition::new(),
20        }
21    }
22}
23
24/// Service wrapper that allows only one call to run at a time.
25///
26/// This type is intentionally not cloneable. Cloning its readiness state would
27/// create another independent gate and allow calls to overlap.
28#[derive(Debug)]
29pub struct OneRequestService<S> {
30    waiters: Condition,
31    service: S,
32    ready: Cell<bool>,
33}
34
35impl<S> OneRequestService<S> {
36    /// Wraps a service so that concurrent callers wait for the active call.
37    pub fn new<St, Req>(service: S) -> Self
38    where
39        S: Service<St, Req>,
40    {
41        Self {
42            service,
43            ready: Cell::new(true),
44            waiters: Condition::new(),
45        }
46    }
47}
48
49impl<S> OneRequestService<S> {
50    /// Waits until no call is active.
51    async fn acquire(&self) {
52        if self.ready.get() {
53            return;
54        }
55
56        let waiter = self.waiters.wait();
57        poll_fn(|cx| {
58            // the waiter stays registered after a notification, another task
59            // may have started a call before this one was polled
60            let _ = waiter.poll_ready(cx);
61            if self.ready.get() {
62                Poll::Ready(())
63            } else {
64                Poll::Pending
65            }
66        })
67        .await;
68    }
69}
70
71/// Releases the call slot even if the call future is dropped or panics.
72struct Release<'a, S>(&'a OneRequestService<S>);
73
74impl<S> Drop for Release<'_, S> {
75    fn drop(&mut self) {
76        self.0.ready.set(true);
77        self.0.waiters.notify(());
78    }
79}
80
81impl<S: Service<St, Req>, St, Req> Service<St, Req> for OneRequestService<S> {
82    type Res = S::Res;
83    type Error = S::Error;
84
85    #[inline]
86    async fn ready(&self, ctx: Ctx<'_, Self, St>) -> Result<(), S::Error> {
87        self.acquire().await;
88        ctx.ready(&self.service).await
89    }
90
91    #[inline]
92    async fn call(&self, req: Req, ctx: Ctx<'_, Self, St>) -> Result<S::Res, S::Error> {
93        // several callers can observe readiness before any of them calls
94        self.acquire().await;
95        self.ready.set(false);
96        let _release = Release(self);
97
98        ctx.call(&self.service, req).await
99    }
100
101    ntex_service::forward_shutdown!(St, service);
102}
103
104#[cfg(test)]
105mod tests {
106    use ntex_service::{Pipeline, apply, fn_factory};
107    use std::{cell::RefCell, rc::Rc, time::Duration};
108
109    use super::*;
110    use crate::{channel::oneshot, future::lazy};
111
112    struct SleepService(oneshot::Receiver<()>);
113
114    impl Service<(), ()> for SleepService {
115        type Res = ();
116        type Error = ();
117
118        async fn call(&self, _r: (), _: Ctx<'_, Self>) -> Result<(), ()> {
119            let _ = self.0.recv().await;
120            Ok::<_, ()>(())
121        }
122    }
123
124    #[ntex::test]
125    async fn test_oneshot() {
126        let (tx, rx) = oneshot::channel();
127
128        let srv = Pipeline::new((), OneRequestService::new(SleepService(rx)));
129        assert_eq!(lazy(|cx| srv.poll_ready(cx)).await, Poll::Ready(Ok(())));
130
131        let srv2 = srv.bind();
132        ntex::rt::spawn(async move {
133            let _ = srv2.call(()).await;
134        });
135        crate::time::sleep(Duration::from_millis(25)).await;
136        assert_eq!(lazy(|cx| srv.poll_ready(cx)).await, Poll::Pending);
137
138        let _ = tx.send(());
139        crate::time::sleep(Duration::from_millis(25)).await;
140        assert_eq!(lazy(|cx| srv.poll_ready(cx)).await, Poll::Ready(Ok(())));
141        srv.shutdown().await;
142    }
143
144    #[ntex::test]
145    async fn test_middleware() {
146        assert_eq!(format!("{OneRequest:?}"), "OneRequest");
147
148        let (tx, rx) = oneshot::channel();
149        let rx = RefCell::new(Some(rx));
150        let sf = apply(
151            OneRequest,
152            fn_factory(move |(): &()| {
153                let rx = rx.borrow_mut().take().unwrap();
154                async move { Ok::<_, ()>(SleepService(rx)) }
155            }),
156        );
157
158        let srv = sf.pipeline(()).await.unwrap();
159        assert_eq!(lazy(|cx| srv.poll_ready(cx)).await, Poll::Ready(Ok(())));
160
161        let srv1 = srv.bind();
162        ntex::rt::spawn(async move {
163            let _ = srv1.call(()).await;
164        });
165        crate::time::sleep(Duration::from_millis(25)).await;
166        assert_eq!(lazy(|cx| srv.poll_ready(cx)).await, Poll::Pending);
167
168        let _ = tx.send(());
169        crate::time::sleep(Duration::from_millis(25)).await;
170        assert_eq!(lazy(|cx| srv.poll_ready(cx)).await, Poll::Ready(Ok(())));
171    }
172
173    #[ntex::test]
174    async fn test_middleware2() {
175        assert_eq!(format!("{OneRequest:?}"), "OneRequest");
176
177        let (tx, rx) = oneshot::channel();
178        let rx = RefCell::new(Some(rx));
179        let sf = apply(
180            OneRequest,
181            fn_factory(move |(): &()| {
182                let rx = rx.borrow_mut().take().unwrap();
183                async move { Ok::<_, ()>(SleepService(rx)) }
184            }),
185        );
186
187        let srv = sf.pipeline(()).await.unwrap();
188        assert_eq!(lazy(|cx| srv.poll_ready(cx)).await, Poll::Ready(Ok(())));
189
190        let srv1 = srv.bind();
191        ntex::rt::spawn(async move {
192            let _ = srv1.call(()).await;
193        });
194        crate::time::sleep(Duration::from_millis(25)).await;
195        assert_eq!(lazy(|cx| srv.poll_ready(cx)).await, Poll::Pending);
196
197        let _ = tx.send(());
198        crate::time::sleep(Duration::from_millis(25)).await;
199        assert_eq!(lazy(|cx| srv.poll_ready(cx)).await, Poll::Ready(Ok(())));
200    }
201
202    #[ntex::test]
203    async fn test_cancelled_call_releases() {
204        let (_tx, rx) = oneshot::channel();
205        let srv = Pipeline::new((), OneRequestService::new(SleepService(rx)));
206
207        let res = crate::time::timeout(Duration::from_millis(25), srv.call(())).await;
208        assert!(res.is_err());
209        assert_eq!(lazy(|cx| srv.poll_ready(cx)).await, Poll::Ready(Ok(())));
210    }
211
212    struct CountService {
213        active: Rc<Cell<usize>>,
214        max: Rc<Cell<usize>>,
215    }
216
217    impl Service<(), ()> for CountService {
218        type Res = ();
219        type Error = ();
220
221        async fn call(&self, _r: (), _: Ctx<'_, Self>) -> Result<(), ()> {
222            self.active.set(self.active.get() + 1);
223            self.max.set(self.max.get().max(self.active.get()));
224            crate::time::sleep(Duration::from_millis(10)).await;
225            self.active.set(self.active.get() - 1);
226            Ok(())
227        }
228    }
229
230    fn count_srv() -> (Pipeline<(), (), ()>, Rc<Cell<usize>>) {
231        let max = Rc::new(Cell::new(0));
232        let srv = Pipeline::new(
233            (),
234            OneRequestService::new(CountService {
235                active: Rc::new(Cell::new(0)),
236                max: max.clone(),
237            }),
238        );
239        (srv, max)
240    }
241
242    #[ntex::test]
243    async fn test_many_waiters() {
244        let (srv, max) = count_srv();
245        let done = Rc::new(Cell::new(0));
246        for _ in 0..4 {
247            let (srv, done) = (srv.bind(), done.clone());
248            ntex::rt::spawn(async move {
249                srv.call(()).await.unwrap();
250                done.set(done.get() + 1);
251            });
252        }
253
254        // calls run one after another, wait for all of them with a generous
255        // deadline, timers on loaded CI runners overshoot a lot
256        for _ in 0..100 {
257            if done.get() == 4 {
258                break;
259            }
260            crate::time::sleep(Duration::from_millis(50)).await;
261        }
262        assert_eq!(done.get(), 4);
263        assert_eq!(max.get(), 1);
264    }
265
266    #[ntex::test]
267    async fn test_no_overlap_after_ready() {
268        let (srv, max) = count_srv();
269
270        srv.ready().await.unwrap();
271        srv.ready().await.unwrap();
272        let (r1, r2) = crate::future::join(srv.call_static(()), srv.call_static(())).await;
273        assert!(r1.is_ok() && r2.is_ok());
274        assert_eq!(max.get(), 1);
275    }
276}