1use std::{cell::Cell, future::poll_fn, task::Poll};
3
4use ntex_service::{Ctx, Middleware, Service};
5
6use crate::channel::condition::Condition;
7
8#[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#[derive(Debug)]
29pub struct OneRequestService<S> {
30 waiters: Condition,
31 service: S,
32 ready: Cell<bool>,
33}
34
35impl<S> OneRequestService<S> {
36 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 async fn acquire(&self) {
52 if self.ready.get() {
53 return;
54 }
55
56 let waiter = self.waiters.wait();
57 poll_fn(|cx| {
58 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
71struct 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 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 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}