1#![allow(clippy::type_complexity)]
3use std::cell::{Cell, RefCell};
4use std::{collections::VecDeque, fmt, future, marker, task, task::Poll};
5
6use ntex_service::{Ctx, Middleware, Service, pipeline::PipelineState};
7
8use crate::channel::oneshot;
9
10#[derive(Copy, Clone, Debug)]
11pub struct Buffer<St: Clone, Req, Res, Err> {
15 buf_size: usize,
16 cancel_on_shutdown: bool,
17 st: marker::PhantomData<fn(St, Req) -> Result<Res, Err>>,
18}
19
20impl<St: Clone, Req, Res, Err> Buffer<St, Req, Res, Err> {
21 #[must_use]
25 pub fn buf_size(mut self, size: usize) -> Self {
26 self.buf_size = size;
27 self
28 }
29
30 #[must_use]
34 pub fn cancel_on_shutdown(mut self) -> Self {
35 self.cancel_on_shutdown = true;
36 self
37 }
38}
39
40impl<St: Clone, Req, Res, Err> Default for Buffer<St, Req, Res, Err> {
41 fn default() -> Self {
42 Self {
43 buf_size: 16,
44 cancel_on_shutdown: false,
45 st: marker::PhantomData,
46 }
47 }
48}
49
50impl<S, St, Req, Res, Err> Middleware<S, St> for Buffer<St, Req, Res, Err>
51where
52 S: Service<St, Req, Res = Res, Error = Err> + 'static,
53 St: Clone + 'static,
54 Req: 'static,
55 Res: 'static,
56 Err: 'static,
57{
58 type Service = BufferService<St, Req, Res, Err>;
59
60 fn create(&self, _: &St, service: S) -> Self::Service {
61 BufferService::new(self.buf_size, PipelineState::new(service))
62 }
63}
64
65#[derive(Clone, Copy, Debug, PartialEq, Eq)]
66pub enum BufferServiceError<E> {
67 Service(E),
68 RequestCanceled,
69}
70
71impl<E> From<E> for BufferServiceError<E> {
72 fn from(err: E) -> Self {
73 BufferServiceError::Service(err)
74 }
75}
76
77impl<E: fmt::Display> fmt::Display for BufferServiceError<E> {
78 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
79 match self {
80 BufferServiceError::Service(e) => fmt::Display::fmt(e, f),
81 BufferServiceError::RequestCanceled => f.write_str("buffer service request canceled"),
82 }
83 }
84}
85
86impl<E: fmt::Display + fmt::Debug> std::error::Error for BufferServiceError<E> {}
87
88pub struct BufferService<St, Req, Res, Err> {
92 size: usize,
93 ready: Cell<bool>,
94 service: PipelineState<St, Req, Res, Err>,
95 buf: RefCell<VecDeque<oneshot::Sender<oneshot::Sender<()>>>>,
96 next_call: RefCell<Option<oneshot::Receiver<()>>>,
97 cancel_on_shutdown: bool,
98 readiness: Cell<Option<task::Waker>>,
99}
100
101impl<St, Req, Res, Err> BufferService<St, Req, Res, Err>
102where
103 St: Clone + 'static,
104{
105 #[must_use]
106 pub fn new(size: usize, service: PipelineState<St, Req, Res, Err>) -> Self {
107 Self {
108 size,
109 service,
110 ready: Cell::new(false),
111 buf: RefCell::new(VecDeque::with_capacity(size)),
112 next_call: RefCell::default(),
113 cancel_on_shutdown: false,
114 readiness: Cell::new(None),
115 }
116 }
117
118 #[must_use]
119 pub fn cancel_on_shutdown(self) -> Self {
120 Self {
121 cancel_on_shutdown: true,
122 ..self
123 }
124 }
125}
126
127impl<St, Req, Res, Err> fmt::Debug for BufferService<St, Req, Res, Err> {
128 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
129 f.debug_struct("BufferService")
130 .field("size", &self.size)
131 .field("cancel_on_shutdown", &self.cancel_on_shutdown)
132 .field("ready", &self.ready)
133 .field("service", &self.service)
134 .field("buf", &self.buf)
135 .field("next_call", &self.next_call)
136 .finish()
137 }
138}
139
140impl<St, Req, Res, Err> Service<St, Req> for BufferService<St, Req, Res, Err>
141where
142 St: Clone + 'static,
143 Req: 'static,
144 Res: 'static,
145 Err: 'static,
146{
147 type Res = Res;
148 type Error = BufferServiceError<Err>;
149
150 async fn ready(&self, ctx: Ctx<'_, Self, St>) -> Result<(), Self::Error> {
151 let next_call = self.next_call.borrow_mut().take();
153 if let Some(next_call) = next_call {
154 let _ = next_call.recv().await;
155 }
156
157 ctx.poll_fn(|cx| {
158 let mut buffer = self.buf.borrow_mut();
159
160 if self.service.poll_ready(cx, ctx.st())?.is_pending() {
162 if buffer.len() < self.size {
163 self.ready.set(false);
165 Poll::Ready(Ok(()))
166 } else {
167 log::trace!("Buffer limit exceeded");
168 let _ = self.readiness.take().map(task::Waker::wake);
170 Poll::Pending
171 }
172 } else {
173 while let Some(sender) = buffer.pop_front() {
174 let (next_call_tx, next_call_rx) = oneshot::channel();
175 if sender.send(next_call_tx).is_err() || next_call_rx.poll_recv(cx).is_ready() {
176 continue;
178 }
179 self.next_call.borrow_mut().replace(next_call_rx);
180 self.ready.set(false);
181 return Poll::Ready(Ok(()));
182 }
183
184 self.ready.set(true);
185 Poll::Ready(Ok(()))
186 }
187 })
188 .await
189 }
190
191 async fn shutdown(&self, ctx: Ctx<'_, Self, St>) {
192 let next_call = self.next_call.borrow_mut().take();
194 if let Some(next_call) = next_call {
195 let _ = next_call.recv().await;
196 }
197
198 future::poll_fn(|cx| {
199 let mut buffer = self.buf.borrow_mut();
200 if self.cancel_on_shutdown {
201 buffer.clear();
202 }
203
204 if !buffer.is_empty() {
205 if task::ready!(self.service.poll_ready(cx, ctx.st())).is_err() {
206 log::error!("Buffered inner service failed while buffer flushing on shutdown");
207 return Poll::Ready(());
208 }
209
210 while let Some(sender) = buffer.pop_front() {
211 let (next_call_tx, next_call_rx) = oneshot::channel();
212 if sender.send(next_call_tx).is_err() || next_call_rx.poll_recv(cx).is_ready() {
213 continue;
215 }
216 self.next_call.borrow_mut().replace(next_call_rx);
217 if buffer.is_empty() {
218 break;
219 }
220 return Poll::Pending;
221 }
222 }
223 Poll::Ready(())
224 })
225 .await;
226
227 self.service.shutdown(ctx.st()).await;
228 }
229
230 async fn call(&self, req: Req, ctx: Ctx<'_, Self, St>) -> Result<Res, Self::Error> {
231 if self.ready.get() {
232 self.ready.set(false);
233 Ok(self.service.call_nowait(req, ctx.st()).await?)
234 } else {
235 let (tx, rx) = oneshot::channel();
236 self.buf.borrow_mut().push_back(tx);
237
238 let _task_guard = rx.recv().await.map_err(|_| {
240 log::trace!("Buffered service request canceled");
241 BufferServiceError::RequestCanceled
242 })?;
243
244 Ok(self.service.call(req, ctx.st()).await?)
246 }
247 }
248}
249
250#[cfg(test)]
251mod tests {
252 #![allow(clippy::unused_async_trait_impl)]
253 use ntex_service::{Pipeline, apply, fn_factory};
254 use std::{rc::Rc, time::Duration};
255
256 use super::*;
257 use crate::{future::lazy, task::LocalWaker};
258
259 #[derive(Debug, Clone)]
260 struct TestService(Rc<Inner>);
261
262 #[derive(Debug)]
263 struct Inner {
264 ready: Cell<bool>,
265 waker: LocalWaker,
266 count: Cell<usize>,
267 }
268
269 impl Service<(), ()> for TestService {
270 type Res = ();
271 type Error = ();
272
273 async fn ready(&self, ctx: Ctx<'_, Self, ()>) -> Result<(), Self::Error> {
274 ctx.poll_fn(|cx| {
275 self.0.waker.register(cx.waker());
276 if self.0.ready.get() {
277 Poll::Ready(Ok(()))
278 } else {
279 Poll::Pending
280 }
281 })
282 .await
283 }
284
285 async fn call(&self, _r: (), _: Ctx<'_, Self, ()>) -> Result<(), ()> {
286 self.0.ready.set(false);
287 self.0.count.set(self.0.count.get() + 1);
288 Ok(())
289 }
290 }
291
292 #[ntex::test]
293 async fn test_service() {
294 let inner = Rc::new(Inner {
295 ready: Cell::new(false),
296 waker: LocalWaker::default(),
297 count: Cell::new(0),
298 });
299
300 let svc = BufferService::new(2, PipelineState::new(TestService(inner.clone())));
301 assert!(format!("{svc:?}").contains("BufferService"));
302
303 let srv = Pipeline::new((), svc);
304 assert_eq!(lazy(|cx| srv.poll_ready(cx)).await, Poll::Ready(Ok(())));
305
306 let srv1 = srv.bind();
307 ntex::rt::spawn(async move {
308 let _ = srv1.call(()).await;
309 });
310 crate::time::sleep(Duration::from_millis(25)).await;
311 assert_eq!(inner.count.get(), 0);
312 assert_eq!(lazy(|cx| srv.poll_ready(cx)).await, Poll::Ready(Ok(())));
313
314 let srv1 = srv.bind();
315 ntex::rt::spawn(async move {
316 let _ = srv1.call(()).await;
317 });
318 crate::time::sleep(Duration::from_millis(25)).await;
319 assert_eq!(inner.count.get(), 0);
320 assert_eq!(lazy(|cx| srv.poll_ready(cx)).await, Poll::Pending);
321
322 inner.ready.set(true);
323 inner.waker.wake();
324 assert_eq!(lazy(|cx| srv.poll_ready(cx)).await, Poll::Ready(Ok(())));
325
326 crate::time::sleep(Duration::from_millis(25)).await;
327 assert_eq!(inner.count.get(), 1);
328
329 inner.ready.set(true);
330 inner.waker.wake();
331 assert_eq!(lazy(|cx| srv.poll_ready(cx)).await, Poll::Ready(Ok(())));
332
333 crate::time::sleep(Duration::from_millis(25)).await;
334 assert_eq!(inner.count.get(), 2);
335
336 let inner = Rc::new(Inner {
337 ready: Cell::new(true),
338 waker: LocalWaker::default(),
339 count: Cell::new(0),
340 });
341
342 let srv = Pipeline::new(
343 (),
344 BufferService::new(2, PipelineState::new(TestService(inner.clone()))),
345 );
346 assert_eq!(lazy(|cx| srv.poll_ready(cx)).await, Poll::Ready(Ok(())));
347
348 let _ = srv.call(()).await;
349 assert_eq!(inner.count.get(), 1);
350 assert_eq!(lazy(|cx| srv.poll_ready(cx)).await, Poll::Ready(Ok(())));
351 assert!(lazy(|cx| srv.poll_shutdown(cx)).await.is_ready());
352
353 let err = BufferServiceError::from("test");
354 assert!(format!("{err}").contains("test"));
355 assert!(format!("{:?}", Buffer::<(), (), (), ()>::default()).contains("Buffer"));
356 }
357
358 #[ntex::test]
359 #[allow(clippy::redundant_clone)]
360 async fn test_middleware() {
361 let inner = Rc::new(Inner {
362 ready: Cell::new(false),
363 waker: LocalWaker::default(),
364 count: Cell::new(0),
365 });
366 let inner2 = inner.clone();
367
368 let srv = apply(
369 Buffer::default().buf_size(2),
370 fn_factory(async move |(): &()| Ok::<_, ()>(TestService(inner2.clone()))),
371 );
372
373 let srv = srv.pipeline(()).await.unwrap();
374 assert_eq!(lazy(|cx| srv.poll_ready(cx)).await, Poll::Ready(Ok(())));
375
376 let srv1 = srv.bind();
377 ntex::rt::spawn(async move {
378 let _ = srv1.call(()).await;
379 });
380 crate::time::sleep(Duration::from_millis(25)).await;
381 assert_eq!(inner.count.get(), 0);
382 assert_eq!(lazy(|cx| srv.poll_ready(cx)).await, Poll::Ready(Ok(())));
383
384 let srv1 = srv.bind();
385 ntex::rt::spawn(async move {
386 let _ = srv1.call(()).await;
387 });
388 crate::time::sleep(Duration::from_millis(25)).await;
389 assert_eq!(inner.count.get(), 0);
390 assert_eq!(lazy(|cx| srv.poll_ready(cx)).await, Poll::Pending);
391
392 inner.ready.set(true);
393 inner.waker.wake();
394 assert_eq!(lazy(|cx| srv.poll_ready(cx)).await, Poll::Ready(Ok(())));
395
396 crate::time::sleep(Duration::from_millis(25)).await;
397 assert_eq!(inner.count.get(), 1);
398
399 inner.ready.set(true);
400 inner.waker.wake();
401 assert_eq!(lazy(|cx| srv.poll_ready(cx)).await, Poll::Ready(Ok(())));
402
403 crate::time::sleep(Duration::from_millis(25)).await;
404 assert_eq!(inner.count.get(), 2);
405 }
406
407 #[ntex::test]
408 #[allow(clippy::redundant_clone)]
409 async fn test_middleware2() {
410 let inner = Rc::new(Inner {
411 ready: Cell::new(false),
412 waker: LocalWaker::default(),
413 count: Cell::new(0),
414 });
415 let inner2 = inner.clone();
416
417 let srv = apply(
418 Buffer::default().buf_size(2),
419 fn_factory(async move |(): &()| Ok::<_, ()>(TestService(inner2.clone()))),
420 );
421
422 let srv = srv.pipeline(()).await.unwrap();
423 assert_eq!(lazy(|cx| srv.poll_ready(cx)).await, Poll::Ready(Ok(())));
424
425 let srv1 = srv.bind();
426 ntex::rt::spawn(async move {
427 let _ = srv1.call(()).await;
428 });
429 crate::time::sleep(Duration::from_millis(25)).await;
430 assert_eq!(inner.count.get(), 0);
431 assert_eq!(lazy(|cx| srv.poll_ready(cx)).await, Poll::Ready(Ok(())));
432
433 let srv1 = srv.bind();
434 ntex::rt::spawn(async move {
435 let _ = srv1.call(()).await;
436 });
437 crate::time::sleep(Duration::from_millis(25)).await;
438 assert_eq!(inner.count.get(), 0);
439 assert_eq!(lazy(|cx| srv.poll_ready(cx)).await, Poll::Pending);
440
441 inner.ready.set(true);
442 inner.waker.wake();
443 assert_eq!(lazy(|cx| srv.poll_ready(cx)).await, Poll::Ready(Ok(())));
444
445 crate::time::sleep(Duration::from_millis(25)).await;
446 assert_eq!(inner.count.get(), 1);
447
448 inner.ready.set(true);
449 inner.waker.wake();
450 assert_eq!(lazy(|cx| srv.poll_ready(cx)).await, Poll::Ready(Ok(())));
451
452 crate::time::sleep(Duration::from_millis(25)).await;
453 assert_eq!(inner.count.get(), 2);
454 }
455}