Skip to main content

ntex_util/services/
inflight.rs

1//! Middleware for limiting concurrent service calls.
2use std::cell::Cell;
3
4use ntex_service::{Ctx, Middleware, Service};
5
6use super::counter::Counter;
7
8/// Middleware that limits the number of concurrent calls to a service.
9///
10/// Readiness remains pending while every slot is in use. The default limit is
11/// 15 concurrent calls.
12#[derive(Copy, Clone, Debug)]
13pub struct InFlight {
14    max_inflight: usize,
15}
16
17impl InFlight {
18    /// Creates middleware with the specified concurrency limit.
19    ///
20    /// A limit of zero keeps the service permanently unavailable.
21    pub fn new(max: usize) -> Self {
22        Self { max_inflight: max }
23    }
24}
25
26impl Default for InFlight {
27    fn default() -> Self {
28        Self::new(15)
29    }
30}
31
32impl<S, St> Middleware<S, St> for InFlight {
33    type Service = InFlightService<S>;
34
35    fn create(&self, _: &St, service: S) -> Self::Service {
36        InFlightService::new(self.max_inflight, service)
37    }
38}
39
40#[derive(Debug)]
41/// Service wrapper that enforces a concurrent-call limit.
42pub struct InFlightService<S> {
43    count: Counter,
44    service: S,
45    ready: Cell<bool>,
46    entered: Cell<u32>,
47}
48
49impl<S> InFlightService<S> {
50    /// Wraps `service` with the specified concurrency limit.
51    ///
52    /// A limit of zero keeps the service permanently unavailable.
53    pub fn new(max: usize, service: S) -> Self {
54        Self {
55            service,
56            count: Counter::new(max),
57            ready: Cell::new(false),
58            entered: Cell::new(0),
59        }
60    }
61}
62
63impl<S, St, Req> Service<St, Req> for InFlightService<S>
64where
65    S: Service<St, Req>,
66{
67    type Res = S::Res;
68    type Error = S::Error;
69
70    #[inline]
71    async fn ready(&self, ctx: Ctx<'_, Self, St>) -> Result<(), S::Error> {
72        let entered = self.entered.get();
73        let result = if self.count.is_available() {
74            ctx.ready(&self.service).await
75        } else {
76            crate::future::join(self.count.available(), ctx.ready(&self.service))
77                .await
78                .1?;
79            // the inner readiness can be stale after waiting for a free slot
80            ctx.ready(&self.service).await
81        };
82        // valid only if no call entered the service during the check
83        self.ready
84            .set(result.is_ok() && entered == self.entered.get());
85        result
86    }
87
88    #[inline]
89    async fn call(&self, req: Req, ctx: Ctx<'_, Self, St>) -> Result<S::Res, S::Error> {
90        if !self.ready.get() {
91            ctx.ready(self).await?;
92        }
93        self.ready.set(false);
94        self.entered.set(self.entered.get().wrapping_add(1));
95        let _guard = self.count.get();
96        ctx.call_nowait(&self.service, req).await
97    }
98
99    ntex_service::forward_shutdown!(St, service);
100}
101
102#[cfg(test)]
103mod tests {
104    use std::{cell::Cell, cell::RefCell, rc::Rc, task::Poll, time::Duration};
105
106    use async_channel as mpmc;
107    use ntex_service::{Pipeline, apply, fn_factory};
108
109    use super::*;
110    use crate::{channel::oneshot, future::lazy};
111
112    struct SleepService(mpmc::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_service() {
126        let (tx, rx) = mpmc::unbounded();
127        let counter = Rc::new(Cell::new(0));
128
129        let srv = Pipeline::new((), InFlightService::new(1, SleepService(rx)));
130        assert_eq!(lazy(|cx| srv.poll_ready(cx)).await, Poll::Ready(Ok(())));
131
132        let counter2 = counter.clone();
133        let fut = srv.call_static(());
134        ntex::rt::spawn(async move {
135            let _ = fut.await;
136            counter2.set(counter2.get() + 1);
137        });
138        crate::time::sleep(Duration::from_millis(25)).await;
139        assert_eq!(lazy(|cx| srv.poll_ready(cx)).await, Poll::Pending);
140
141        let counter2 = counter.clone();
142        let fut = srv.call_static(());
143        ntex::rt::spawn(async move {
144            let _ = fut.await;
145            counter2.set(counter2.get() + 1);
146        });
147        crate::time::sleep(Duration::from_millis(25)).await;
148        assert_eq!(lazy(|cx| srv.poll_ready(cx)).await, Poll::Pending);
149
150        let counter2 = counter.clone();
151        let fut = srv.call_static(());
152        let (stx, srx) = oneshot::channel::<()>();
153        ntex::rt::spawn(async move {
154            let _ = fut.await;
155            counter2.set(counter2.get() + 1);
156            let _ = stx.send(());
157        });
158        crate::time::sleep(Duration::from_millis(25)).await;
159        assert_eq!(lazy(|cx| srv.poll_ready(cx)).await, Poll::Pending);
160
161        let _ = tx.send(()).await;
162        crate::time::sleep(Duration::from_millis(25)).await;
163        assert_eq!(lazy(|cx| srv.poll_ready(cx)).await, Poll::Ready(Ok(())));
164
165        let _ = tx.send(()).await;
166        crate::time::sleep(Duration::from_millis(25)).await;
167        assert_eq!(lazy(|cx| srv.poll_ready(cx)).await, Poll::Pending);
168
169        let _ = tx.send(()).await;
170        let _ = srx.recv().await;
171        assert_eq!(counter.get(), 3);
172        srv.shutdown().await;
173    }
174
175    #[ntex::test]
176    async fn test_middleware() {
177        assert_eq!(InFlight::default().max_inflight, 15);
178        assert_eq!(
179            format!("{:?}", InFlight::new(1)),
180            "InFlight { max_inflight: 1 }"
181        );
182
183        let (tx, rx) = mpmc::unbounded();
184        let rx = RefCell::new(Some(rx));
185        let sf = apply(
186            InFlight::new(1),
187            fn_factory(move |(): &()| {
188                let rx = rx.borrow_mut().take().unwrap();
189                async move { Ok::<_, ()>(SleepService(rx)) }
190            }),
191        );
192
193        let srv = sf.pipeline(()).await.unwrap();
194        assert_eq!(lazy(|cx| srv.poll_ready(cx)).await, Poll::Ready(Ok(())));
195
196        let srv2 = srv.bind();
197        ntex::rt::spawn(async move {
198            let _ = srv2.call(()).await;
199        });
200        crate::time::sleep(Duration::from_millis(25)).await;
201        assert_eq!(lazy(|cx| srv.poll_ready(cx)).await, Poll::Pending);
202
203        let _ = tx.send(()).await;
204        crate::time::sleep(Duration::from_millis(25)).await;
205        assert_eq!(lazy(|cx| srv.poll_ready(cx)).await, Poll::Ready(Ok(())));
206    }
207
208    #[ntex::test]
209    async fn test_middleware2() {
210        assert_eq!(InFlight::default().max_inflight, 15);
211        assert_eq!(
212            format!("{:?}", InFlight::new(1)),
213            "InFlight { max_inflight: 1 }"
214        );
215
216        let (tx, rx) = mpmc::unbounded();
217        let rx = RefCell::new(Some(rx));
218        let sf = apply(
219            InFlight::new(1),
220            fn_factory(move |(): &()| {
221                let rx = rx.borrow_mut().take().unwrap();
222                async move { Ok::<_, ()>(SleepService(rx)) }
223            }),
224        );
225
226        let srv = sf.pipeline(()).await.unwrap();
227        assert_eq!(lazy(|cx| srv.poll_ready(cx)).await, Poll::Ready(Ok(())));
228
229        let srv2 = srv.bind();
230        ntex::rt::spawn(async move {
231            let _ = srv2.call(()).await;
232        });
233        crate::time::sleep(Duration::from_millis(25)).await;
234        assert_eq!(lazy(|cx| srv.poll_ready(cx)).await, Poll::Pending);
235
236        let _ = tx.send(()).await;
237        crate::time::sleep(Duration::from_millis(25)).await;
238        assert_eq!(lazy(|cx| srv.poll_ready(cx)).await, Poll::Ready(Ok(())));
239    }
240
241    #[derive(Default)]
242    struct ProbeState {
243        active: Cell<usize>,
244        max: Cell<usize>,
245        checks: Cell<usize>,
246        waker: crate::task::LocalWaker,
247    }
248
249    /// Inner service with its own concurrency limit
250    struct Probe(Rc<ProbeState>, usize, mpmc::Receiver<()>);
251
252    impl Service<(), ()> for Probe {
253        type Res = ();
254        type Error = ();
255
256        async fn ready(&self, _: Ctx<'_, Self>) -> Result<(), ()> {
257            self.0.checks.set(self.0.checks.get() + 1);
258            std::future::poll_fn(|cx| {
259                if self.0.active.get() < self.1 {
260                    Poll::Ready(Ok(()))
261                } else {
262                    self.0.waker.register(cx.waker());
263                    Poll::Pending
264                }
265            })
266            .await
267        }
268
269        async fn call(&self, (): (), _: Ctx<'_, Self>) -> Result<(), ()> {
270            let st = &self.0;
271            st.active.set(st.active.get() + 1);
272            st.max.set(st.max.get().max(st.active.get()));
273            let _ = self.2.recv().await;
274            st.active.set(st.active.get() - 1);
275            st.waker.wake();
276            Ok(())
277        }
278    }
279
280    #[ntex::test]
281    async fn test_inner_readiness_checked_once() {
282        let (tx, rx) = mpmc::unbounded();
283        let st = Rc::new(ProbeState::default());
284        let srv = Pipeline::new((), InFlightService::new(4, Probe(st.clone(), 4, rx)));
285        for _ in 0..3 {
286            let _ = tx.send(()).await;
287            srv.call(()).await.unwrap();
288        }
289        assert_eq!(st.checks.get(), 3);
290
291        st.checks.set(0);
292        for _ in 0..3 {
293            let _ = tx.send(()).await;
294            srv.ready().await.unwrap();
295            srv.call(()).await.unwrap();
296        }
297        assert_eq!(st.checks.get(), 3);
298    }
299
300    /// Calls the inner service without checking its readiness, the inner
301    /// service has to enforce its limits
302    struct NoWait<S>(S);
303
304    impl<S: Service<(), (), Res = (), Error = ()>> Service<(), ()> for NoWait<S> {
305        type Res = ();
306        type Error = ();
307
308        async fn call(&self, req: (), ctx: Ctx<'_, Self>) -> Result<(), ()> {
309            ctx.call_nowait(&self.0, req).await
310        }
311    }
312
313    /// Returns the max number of concurrent calls of the inner service
314    async fn run_concurrent(max: usize, inner: usize) -> usize {
315        let (tx, rx) = mpmc::unbounded();
316        let st = Rc::new(ProbeState::default());
317        let srv = Pipeline::new(
318            (),
319            NoWait(InFlightService::new(max, Probe(st.clone(), inner, rx))),
320        );
321        let mut futs = Vec::new();
322        for _ in 0..4 {
323            futs.push(ntex::rt::spawn(srv.call_static(())));
324        }
325        crate::time::sleep(Duration::from_millis(20)).await;
326        for _ in 0..4 {
327            let _ = tx.send(()).await;
328        }
329        for f in futs {
330            let _ = f.await;
331        }
332        st.max.get()
333    }
334
335    #[ntex::test]
336    async fn test_limits_without_readiness_check() {
337        assert_eq!(run_concurrent(1, 4).await, 1, "inflight limit");
338        assert_eq!(run_concurrent(2, 1).await, 1, "inner limit");
339        assert_eq!(run_concurrent(2, 4).await, 2, "inflight limit 2");
340    }
341}