Skip to main content

ntex_util/services/
onerequest.rs

1//! Service that limits number of in-flight async requests to 1.
2use std::{cell::Cell, future::poll_fn, task::Poll};
3
4use ntex_service::{Ctx, Middleware, Service};
5
6use crate::task::LocalWaker;
7
8/// `OneRequest` - service factory for service that can limit number of in-flight
9/// async requests to 1.
10#[derive(Copy, Clone, Default, Debug)]
11pub struct OneRequest;
12
13impl<S, St> Middleware<S, St> for OneRequest {
14    type Service = OneRequestService<S>;
15
16    fn create(&self, _: &St, service: S) -> Self::Service {
17        OneRequestService {
18            service,
19            ready: Cell::new(true),
20            waker: LocalWaker::new(),
21        }
22    }
23}
24
25#[derive(Clone, Debug)]
26pub struct OneRequestService<S> {
27    waker: LocalWaker,
28    service: S,
29    ready: Cell<bool>,
30}
31
32impl<S> OneRequestService<S> {
33    pub fn new<St, Req>(service: S) -> Self
34    where
35        S: Service<St, Req>,
36    {
37        Self {
38            service,
39            ready: Cell::new(true),
40            waker: LocalWaker::new(),
41        }
42    }
43}
44
45impl<S: Service<St, Req>, St, Req> Service<St, Req> for OneRequestService<S> {
46    type Res = S::Res;
47    type Error = S::Error;
48
49    #[inline]
50    async fn ready(&self, ctx: Ctx<'_, Self, St>) -> Result<(), S::Error> {
51        if !self.ready.get() {
52            poll_fn(|cx| {
53                self.waker.register(cx.waker());
54                if self.ready.get() {
55                    Poll::Ready(())
56                } else {
57                    Poll::Pending
58                }
59            })
60            .await;
61        }
62        ctx.ready(&self.service).await
63    }
64
65    #[inline]
66    async fn call(&self, req: Req, ctx: Ctx<'_, Self, St>) -> Result<S::Res, S::Error> {
67        self.ready.set(false);
68
69        let result = ctx.call(&self.service, req).await;
70        self.ready.set(true);
71        self.waker.wake();
72        result
73    }
74
75    ntex_service::forward_shutdown!(St, service);
76}
77
78#[cfg(test)]
79mod tests {
80    use ntex_service::{Pipeline, apply, fn_factory};
81    use std::{cell::RefCell, time::Duration};
82
83    use super::*;
84    use crate::{channel::oneshot, future::lazy};
85
86    struct SleepService(oneshot::Receiver<()>);
87
88    impl Service<(), ()> for SleepService {
89        type Res = ();
90        type Error = ();
91
92        async fn call(&self, _r: (), _: Ctx<'_, Self>) -> Result<(), ()> {
93            let _ = self.0.recv().await;
94            Ok::<_, ()>(())
95        }
96    }
97
98    #[ntex::test]
99    async fn test_oneshot() {
100        let (tx, rx) = oneshot::channel();
101
102        let srv = Pipeline::new((), OneRequestService::new(SleepService(rx)));
103        assert_eq!(lazy(|cx| srv.poll_ready(cx)).await, Poll::Ready(Ok(())));
104
105        let srv2 = srv.bind();
106        ntex::rt::spawn(async move {
107            let _ = srv2.call(()).await;
108        });
109        crate::time::sleep(Duration::from_millis(25)).await;
110        assert_eq!(lazy(|cx| srv.poll_ready(cx)).await, Poll::Pending);
111
112        let _ = tx.send(());
113        crate::time::sleep(Duration::from_millis(25)).await;
114        assert_eq!(lazy(|cx| srv.poll_ready(cx)).await, Poll::Ready(Ok(())));
115        srv.shutdown().await;
116    }
117
118    #[ntex::test]
119    async fn test_middleware() {
120        assert_eq!(format!("{OneRequest:?}"), "OneRequest");
121
122        let (tx, rx) = oneshot::channel();
123        let rx = RefCell::new(Some(rx));
124        let sf = apply(
125            OneRequest,
126            fn_factory(move |(): &()| {
127                let rx = rx.borrow_mut().take().unwrap();
128                async move { Ok::<_, ()>(SleepService(rx)) }
129            }),
130        );
131
132        let srv = sf.pipeline(()).await.unwrap();
133        assert_eq!(lazy(|cx| srv.poll_ready(cx)).await, Poll::Ready(Ok(())));
134
135        let srv1 = srv.bind();
136        ntex::rt::spawn(async move {
137            let _ = srv1.call(()).await;
138        });
139        crate::time::sleep(Duration::from_millis(25)).await;
140        assert_eq!(lazy(|cx| srv.poll_ready(cx)).await, Poll::Pending);
141
142        let _ = tx.send(());
143        crate::time::sleep(Duration::from_millis(25)).await;
144        assert_eq!(lazy(|cx| srv.poll_ready(cx)).await, Poll::Ready(Ok(())));
145    }
146
147    #[ntex::test]
148    async fn test_middleware2() {
149        assert_eq!(format!("{OneRequest:?}"), "OneRequest");
150
151        let (tx, rx) = oneshot::channel();
152        let rx = RefCell::new(Some(rx));
153        let sf = apply(
154            OneRequest,
155            fn_factory(move |(): &()| {
156                let rx = rx.borrow_mut().take().unwrap();
157                async move { Ok::<_, ()>(SleepService(rx)) }
158            }),
159        );
160
161        let srv = sf.pipeline(()).await.unwrap();
162        assert_eq!(lazy(|cx| srv.poll_ready(cx)).await, Poll::Ready(Ok(())));
163
164        let srv1 = srv.bind();
165        ntex::rt::spawn(async move {
166            let _ = srv1.call(()).await;
167        });
168        crate::time::sleep(Duration::from_millis(25)).await;
169        assert_eq!(lazy(|cx| srv.poll_ready(cx)).await, Poll::Pending);
170
171        let _ = tx.send(());
172        crate::time::sleep(Duration::from_millis(25)).await;
173        assert_eq!(lazy(|cx| srv.poll_ready(cx)).await, Poll::Ready(Ok(())));
174    }
175}