1use std::cell::Cell;
3
4use ntex_service::{Ctx, Middleware, Service};
5
6use super::counter::Counter;
7
8#[derive(Copy, Clone, Debug)]
13pub struct InFlight {
14 max_inflight: usize,
15}
16
17impl InFlight {
18 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)]
41pub struct InFlightService<S> {
43 count: Counter,
44 service: S,
45 ready: Cell<bool>,
46 entered: Cell<u32>,
47}
48
49impl<S> InFlightService<S> {
50 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 ctx.ready(&self.service).await
81 };
82 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 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 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 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}