1#![allow(clippy::type_complexity)]
3use std::cell::{Cell, RefCell};
4use std::{collections::VecDeque, fmt, marker, 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> {
17 buf_size: usize,
18 cancel_on_shutdown: bool,
19 st: marker::PhantomData<fn(St, Req) -> Result<Res, Err>>,
20}
21
22impl<St: Clone, Req, Res, Err> Buffer<St, Req, Res, Err> {
23 #[must_use]
27 pub fn buf_size(mut self, size: usize) -> Self {
28 self.buf_size = size;
29 self
30 }
31
32 #[must_use]
36 pub fn cancel_on_shutdown(mut self) -> Self {
37 self.cancel_on_shutdown = true;
38 self
39 }
40}
41
42impl<St: Clone, Req, Res, Err> Default for Buffer<St, Req, Res, Err> {
43 fn default() -> Self {
44 Self {
45 buf_size: 16,
46 cancel_on_shutdown: false,
47 st: marker::PhantomData,
48 }
49 }
50}
51
52impl<S, St, Req, Res, Err> Middleware<S, St> for Buffer<St, Req, Res, Err>
53where
54 S: Service<St, Req, Res = Res, Error = Err> + 'static,
55 St: Clone + 'static,
56 Req: 'static,
57 Res: 'static,
58 Err: 'static,
59{
60 type Service = BufferService<St, Req, Res, Err>;
61
62 fn create(&self, _: &St, service: S) -> Self::Service {
63 let service = BufferService::new(self.buf_size, PipelineState::new(service));
64 if self.cancel_on_shutdown {
65 service.cancel_on_shutdown()
66 } else {
67 service
68 }
69 }
70}
71
72#[derive(Clone, Copy, Debug, PartialEq, Eq)]
74pub enum BufferServiceError<E> {
75 Service(E),
77 RequestCanceled,
79}
80
81impl<E> From<E> for BufferServiceError<E> {
82 fn from(err: E) -> Self {
83 BufferServiceError::Service(err)
84 }
85}
86
87impl<E: fmt::Display> fmt::Display for BufferServiceError<E> {
88 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
89 match self {
90 BufferServiceError::Service(e) => fmt::Display::fmt(e, f),
91 BufferServiceError::RequestCanceled => f.write_str("buffer service request canceled"),
92 }
93 }
94}
95
96impl<E: fmt::Display + fmt::Debug> std::error::Error for BufferServiceError<E> {}
97
98pub struct BufferService<St, Req, Res, Err> {
108 size: usize,
109 ready: Cell<bool>,
110 service: PipelineState<St, Req, Res, Err>,
111 buf: RefCell<VecDeque<oneshot::Sender<oneshot::Sender<()>>>>,
112 next_call: RefCell<Option<oneshot::Receiver<()>>>,
113 cancel_on_shutdown: bool,
114}
115
116impl<St, Req, Res, Err> BufferService<St, Req, Res, Err>
117where
118 St: Clone + 'static,
119{
120 #[must_use]
121 pub fn new(size: usize, service: PipelineState<St, Req, Res, Err>) -> Self {
123 Self {
124 size,
125 service,
126 ready: Cell::new(false),
127 buf: RefCell::new(VecDeque::with_capacity(size)),
128 next_call: RefCell::default(),
129 cancel_on_shutdown: false,
130 }
131 }
132
133 #[must_use]
134 pub fn cancel_on_shutdown(self) -> Self {
136 Self {
137 cancel_on_shutdown: true,
138 ..self
139 }
140 }
141}
142
143impl<St, Req, Res, Err> fmt::Debug for BufferService<St, Req, Res, Err> {
144 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
145 f.debug_struct("BufferService")
146 .field("size", &self.size)
147 .field("cancel_on_shutdown", &self.cancel_on_shutdown)
148 .field("ready", &self.ready)
149 .field("service", &self.service)
150 .field("buf", &self.buf)
151 .field("next_call", &self.next_call)
152 .finish()
153 }
154}
155
156impl<St, Req, Res, Err> Service<St, Req> for BufferService<St, Req, Res, Err>
157where
158 St: Clone + 'static,
159 Req: 'static,
160 Res: 'static,
161 Err: 'static,
162{
163 type Res = Res;
164 type Error = BufferServiceError<Err>;
165
166 async fn ready(&self, ctx: Ctx<'_, Self, St>) -> Result<(), Self::Error> {
167 let next_call = self.next_call.borrow_mut().take();
170 if let Some(next_call) = next_call {
171 let _ = next_call.recv().await;
172 }
173
174 ctx.poll_fn(|cx| {
175 let mut buffer = self.buf.borrow_mut();
176
177 if self.service.poll_ready(cx, ctx.st())?.is_pending() {
179 if buffer.len() < self.size {
180 self.ready.set(false);
182 Poll::Ready(Ok(()))
183 } else {
184 log::trace!("Buffer limit exceeded");
185 Poll::Pending
187 }
188 } else {
189 while let Some(sender) = buffer.pop_front() {
190 let (next_call_tx, next_call_rx) = oneshot::channel();
191 if sender.send(next_call_tx).is_err() || next_call_rx.poll_recv(cx).is_ready() {
192 continue;
194 }
195 self.next_call.borrow_mut().replace(next_call_rx);
196 self.ready.set(false);
197 return Poll::Ready(Ok(()));
198 }
199
200 self.ready.set(true);
201 Poll::Ready(Ok(()))
202 }
203 })
204 .await
205 }
206
207 async fn shutdown(&self, ctx: Ctx<'_, Self, St>) {
208 loop {
210 let next_call = self.next_call.borrow_mut().take();
213 if let Some(next_call) = next_call {
214 let _ = next_call.recv().await;
215 }
216
217 if self.cancel_on_shutdown {
218 self.buf.borrow_mut().clear();
219 }
220 if self.buf.borrow().is_empty() {
221 break;
222 }
223
224 if self.service.ready(ctx.st()).await.is_err() {
225 log::error!("Buffered inner service failed while buffer flushing on shutdown");
226 break;
227 }
228
229 let mut buffer = self.buf.borrow_mut();
230 while let Some(sender) = buffer.pop_front() {
231 let (next_call_tx, next_call_rx) = oneshot::channel();
232 if sender.send(next_call_tx).is_ok() {
233 self.next_call.borrow_mut().replace(next_call_rx);
234 break;
235 }
236 }
238 }
239
240 self.service.shutdown(ctx.st()).await;
241 }
242
243 async fn call(&self, req: Req, ctx: Ctx<'_, Self, St>) -> Result<Res, Self::Error> {
244 if self.ready.get() {
245 self.ready.set(false);
246 Ok(self.service.call(req, ctx.st()).await?)
248 } else {
249 let (tx, rx) = oneshot::channel();
250 self.buf.borrow_mut().push_back(tx);
251
252 let _task_guard = rx.recv().await.map_err(|_| {
254 log::trace!("Buffered service request canceled");
255 BufferServiceError::RequestCanceled
256 })?;
257
258 Ok(self.service.call(req, ctx.st()).await?)
260 }
261 }
262}
263
264#[cfg(test)]
265mod tests {
266 #![allow(clippy::unused_async_trait_impl)]
267 use ntex_service::{Pipeline, apply, fn_factory};
268 use std::{cell::RefCell, pin::Pin, rc::Rc, time::Duration};
269
270 use super::*;
271 use crate::{future::lazy, task::LocalWaker};
272
273 #[derive(Debug, Clone)]
274 struct TestService(Rc<Inner>);
275
276 #[derive(Debug)]
277 struct Inner {
278 ready: Cell<bool>,
279 waker: LocalWaker,
280 count: Cell<usize>,
281 }
282
283 impl Service<(), ()> for TestService {
284 type Res = ();
285 type Error = ();
286
287 async fn ready(&self, ctx: Ctx<'_, Self, ()>) -> Result<(), Self::Error> {
288 ctx.poll_fn(|cx| {
289 self.0.waker.register(cx.waker());
290 if self.0.ready.get() {
291 Poll::Ready(Ok(()))
292 } else {
293 Poll::Pending
294 }
295 })
296 .await
297 }
298
299 async fn call(&self, _r: (), _: Ctx<'_, Self, ()>) -> Result<(), ()> {
300 self.0.ready.set(false);
301 self.0.count.set(self.0.count.get() + 1);
302 Ok(())
303 }
304 }
305
306 struct PendingService {
307 inner: Rc<PendingInner>,
308 release: RefCell<Option<oneshot::Receiver<()>>>,
309 }
310
311 struct PendingInner {
312 started: Cell<bool>,
313 completed: Cell<bool>,
314 shutdown: Cell<bool>,
315 }
316
317 impl Service<(), ()> for PendingService {
318 type Res = ();
319 type Error = ();
320
321 async fn call(&self, (): (), _: Ctx<'_, Self, ()>) -> Result<(), ()> {
322 self.inner.started.set(true);
323 let release = self.release.borrow_mut().take().unwrap();
324 release.recv().await.unwrap();
325 self.inner.completed.set(true);
326 Ok(())
327 }
328
329 async fn shutdown(&self, _: Ctx<'_, Self, ()>) {
330 self.inner.shutdown.set(true);
331 }
332 }
333
334 #[ntex::test]
335 async fn test_service() {
336 let inner = Rc::new(Inner {
337 ready: Cell::new(false),
338 waker: LocalWaker::default(),
339 count: Cell::new(0),
340 });
341
342 let svc = BufferService::new(2, PipelineState::new(TestService(inner.clone())));
343 assert!(format!("{svc:?}").contains("BufferService"));
344
345 let srv = Pipeline::new((), svc);
346 assert_eq!(lazy(|cx| srv.poll_ready(cx)).await, Poll::Ready(Ok(())));
347
348 let srv1 = srv.bind();
349 ntex::rt::spawn(async move {
350 let _ = srv1.call(()).await;
351 });
352 crate::time::sleep(Duration::from_millis(25)).await;
353 assert_eq!(inner.count.get(), 0);
354 assert_eq!(lazy(|cx| srv.poll_ready(cx)).await, Poll::Ready(Ok(())));
355
356 let srv1 = srv.bind();
357 ntex::rt::spawn(async move {
358 let _ = srv1.call(()).await;
359 });
360 crate::time::sleep(Duration::from_millis(25)).await;
361 assert_eq!(inner.count.get(), 0);
362 assert_eq!(lazy(|cx| srv.poll_ready(cx)).await, Poll::Pending);
363
364 inner.ready.set(true);
365 inner.waker.wake();
366 assert_eq!(lazy(|cx| srv.poll_ready(cx)).await, Poll::Ready(Ok(())));
367
368 crate::time::sleep(Duration::from_millis(25)).await;
369 assert_eq!(inner.count.get(), 1);
370
371 inner.ready.set(true);
372 inner.waker.wake();
373 assert_eq!(lazy(|cx| srv.poll_ready(cx)).await, Poll::Ready(Ok(())));
374
375 crate::time::sleep(Duration::from_millis(25)).await;
376 assert_eq!(inner.count.get(), 2);
377
378 let inner = Rc::new(Inner {
379 ready: Cell::new(true),
380 waker: LocalWaker::default(),
381 count: Cell::new(0),
382 });
383
384 let srv = Pipeline::new(
385 (),
386 BufferService::new(2, PipelineState::new(TestService(inner.clone()))),
387 );
388 assert_eq!(lazy(|cx| srv.poll_ready(cx)).await, Poll::Ready(Ok(())));
389
390 let _ = srv.call(()).await;
391 assert_eq!(inner.count.get(), 1);
392 assert_eq!(lazy(|cx| srv.poll_ready(cx)).await, Poll::Ready(Ok(())));
393 assert!(lazy(|cx| srv.poll_shutdown(cx)).await.is_ready());
394
395 let err = BufferServiceError::from("test");
396 assert!(format!("{err}").contains("test"));
397 assert!(format!("{:?}", Buffer::<(), (), (), ()>::default()).contains("Buffer"));
398 }
399
400 #[ntex::test]
401 #[allow(clippy::redundant_clone)]
402 async fn test_middleware() {
403 let inner = Rc::new(Inner {
404 ready: Cell::new(false),
405 waker: LocalWaker::default(),
406 count: Cell::new(0),
407 });
408 let inner2 = inner.clone();
409
410 let srv = apply(
411 Buffer::default().buf_size(2),
412 fn_factory(async move |(): &()| Ok::<_, ()>(TestService(inner2.clone()))),
413 );
414
415 let srv = srv.pipeline(()).await.unwrap();
416 assert_eq!(lazy(|cx| srv.poll_ready(cx)).await, Poll::Ready(Ok(())));
417
418 let srv1 = srv.bind();
419 ntex::rt::spawn(async move {
420 let _ = srv1.call(()).await;
421 });
422 crate::time::sleep(Duration::from_millis(25)).await;
423 assert_eq!(inner.count.get(), 0);
424 assert_eq!(lazy(|cx| srv.poll_ready(cx)).await, Poll::Ready(Ok(())));
425
426 let srv1 = srv.bind();
427 ntex::rt::spawn(async move {
428 let _ = srv1.call(()).await;
429 });
430 crate::time::sleep(Duration::from_millis(25)).await;
431 assert_eq!(inner.count.get(), 0);
432 assert_eq!(lazy(|cx| srv.poll_ready(cx)).await, Poll::Pending);
433
434 inner.ready.set(true);
435 inner.waker.wake();
436 assert_eq!(lazy(|cx| srv.poll_ready(cx)).await, Poll::Ready(Ok(())));
437
438 crate::time::sleep(Duration::from_millis(25)).await;
439 assert_eq!(inner.count.get(), 1);
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(), 2);
447 }
448
449 #[ntex::test]
450 #[allow(clippy::redundant_clone)]
451 async fn test_middleware2() {
452 let inner = Rc::new(Inner {
453 ready: Cell::new(false),
454 waker: LocalWaker::default(),
455 count: Cell::new(0),
456 });
457 let inner2 = inner.clone();
458
459 let srv = apply(
460 Buffer::default().buf_size(2),
461 fn_factory(async move |(): &()| Ok::<_, ()>(TestService(inner2.clone()))),
462 );
463
464 let srv = srv.pipeline(()).await.unwrap();
465 assert_eq!(lazy(|cx| srv.poll_ready(cx)).await, Poll::Ready(Ok(())));
466
467 let srv1 = srv.bind();
468 ntex::rt::spawn(async move {
469 let _ = srv1.call(()).await;
470 });
471 crate::time::sleep(Duration::from_millis(25)).await;
472 assert_eq!(inner.count.get(), 0);
473 assert_eq!(lazy(|cx| srv.poll_ready(cx)).await, Poll::Ready(Ok(())));
474
475 let srv1 = srv.bind();
476 ntex::rt::spawn(async move {
477 let _ = srv1.call(()).await;
478 });
479 crate::time::sleep(Duration::from_millis(25)).await;
480 assert_eq!(inner.count.get(), 0);
481 assert_eq!(lazy(|cx| srv.poll_ready(cx)).await, Poll::Pending);
482
483 inner.ready.set(true);
484 inner.waker.wake();
485 assert_eq!(lazy(|cx| srv.poll_ready(cx)).await, Poll::Ready(Ok(())));
486
487 crate::time::sleep(Duration::from_millis(25)).await;
488 assert_eq!(inner.count.get(), 1);
489
490 inner.ready.set(true);
491 inner.waker.wake();
492 assert_eq!(lazy(|cx| srv.poll_ready(cx)).await, Poll::Ready(Ok(())));
493
494 crate::time::sleep(Duration::from_millis(25)).await;
495 assert_eq!(inner.count.get(), 2);
496 }
497
498 #[ntex::test]
499 async fn middleware_cancels_buffered_requests_on_shutdown() {
500 let inner = Rc::new(Inner {
501 ready: Cell::new(false),
502 waker: LocalWaker::default(),
503 count: Cell::new(0),
504 });
505 let inner2 = inner.clone();
506
507 let srv = apply(
508 Buffer::default().buf_size(1).cancel_on_shutdown(),
509 fn_factory(async move |(): &()| Ok::<_, ()>(TestService(inner2.clone()))),
510 )
511 .pipeline(())
512 .await
513 .unwrap();
514
515 assert_eq!(lazy(|cx| srv.poll_ready(cx)).await, Poll::Ready(Ok(())));
516
517 let canceled = Rc::new(Cell::new(false));
518 let canceled2 = canceled.clone();
519 let srv2 = srv.bind();
520 ntex::rt::spawn(async move {
521 canceled2.set(matches!(
522 srv2.call(()).await,
523 Err(BufferServiceError::RequestCanceled)
524 ));
525 });
526
527 crate::time::sleep(Duration::from_millis(25)).await;
528 srv.shutdown().await;
529 crate::time::sleep(Duration::from_millis(25)).await;
530
531 assert!(canceled.get());
532 assert_eq!(inner.count.get(), 0);
533 }
534
535 #[ntex::test]
536 async fn shutdown_with_in_flight_request() {
537 let inner = Rc::new(PendingInner {
538 started: Cell::new(false),
539 completed: Cell::new(false),
540 shutdown: Cell::new(false),
541 });
542 let (release_tx, release_rx) = oneshot::channel();
543 let service = PendingService {
544 inner: inner.clone(),
545 release: RefCell::new(Some(release_rx)),
546 };
547 let srv = Pipeline::new((), BufferService::new(1, PipelineState::new(service)));
548
549 assert_eq!(lazy(|cx| srv.poll_ready(cx)).await, Poll::Ready(Ok(())));
550
551 let srv2 = srv.bind();
552 ntex::rt::spawn(async move {
553 srv2.call(()).await.unwrap();
554 });
555 crate::time::sleep(Duration::from_millis(25)).await;
556
557 assert!(inner.started.get());
558 assert!(!inner.completed.get());
559 assert!(lazy(|cx| srv.poll_shutdown(cx)).await.is_ready());
560 assert!(inner.shutdown.get());
561 assert!(!inner.completed.get());
562
563 release_tx.send(()).unwrap();
564 crate::time::sleep(Duration::from_millis(25)).await;
565
566 assert!(inner.completed.get());
567 }
568
569 #[derive(Default)]
570 struct DrainInner {
571 ready: Cell<bool>,
572 waker: LocalWaker,
573 active: Cell<usize>,
574 max: Cell<usize>,
575 count: Cell<usize>,
576 active_on_shutdown: Cell<Option<usize>>,
577 }
578
579 struct DrainService(Rc<DrainInner>);
580
581 impl Service<(), ()> for DrainService {
582 type Res = ();
583 type Error = ();
584
585 async fn ready(&self, ctx: Ctx<'_, Self, ()>) -> Result<(), ()> {
586 ctx.poll_fn(|cx| {
587 self.0.waker.register(cx.waker());
588 if self.0.ready.get() {
589 Poll::Ready(Ok(()))
590 } else {
591 Poll::Pending
592 }
593 })
594 .await
595 }
596
597 async fn call(&self, (): (), _: Ctx<'_, Self, ()>) -> Result<(), ()> {
598 let inner = &self.0;
599 inner.active.set(inner.active.get() + 1);
600 inner.max.set(inner.max.get().max(inner.active.get()));
601 crate::time::sleep(Duration::from_millis(25)).await;
602 inner.active.set(inner.active.get() - 1);
603 inner.count.set(inner.count.get() + 1);
604 Ok(())
605 }
606
607 async fn shutdown(&self, _: Ctx<'_, Self, ()>) {
608 self.0.active_on_shutdown.set(Some(self.0.active.get()));
609 }
610 }
611
612 #[ntex::test]
613 async fn shutdown_drains_buffer_one_at_a_time() {
614 let inner = Rc::new(DrainInner::default());
615 let srv = Pipeline::new(
616 (),
617 BufferService::new(4, PipelineState::new(DrainService(inner.clone()))),
618 );
619
620 for _ in 0..2 {
621 srv.ready().await.unwrap();
622 let srv = srv.bind();
623 ntex::rt::spawn(async move {
624 srv.call(()).await.unwrap();
625 });
626 }
627 crate::time::sleep(Duration::from_millis(10)).await;
628 assert_eq!(inner.count.get(), 0);
629
630 inner.ready.set(true);
631 inner.waker.wake();
632 srv.shutdown().await;
633
634 assert_eq!(inner.count.get(), 2);
635 assert_eq!(inner.max.get(), 1);
636 assert_eq!(inner.active_on_shutdown.get(), Some(0));
637 }
638
639 struct FailService(Rc<Cell<bool>>);
640
641 impl Service<(), ()> for FailService {
642 type Res = ();
643 type Error = ();
644
645 async fn ready(&self, _: Ctx<'_, Self, ()>) -> Result<(), ()> {
646 if self.0.get() {
647 Err(())
648 } else {
649 std::future::pending().await
650 }
651 }
652
653 async fn call(&self, (): (), _: Ctx<'_, Self, ()>) -> Result<(), ()> {
654 unreachable!("FailService is never ready")
655 }
656 }
657
658 #[test]
659 fn request_canceled_display() {
660 let err = BufferServiceError::<&str>::RequestCanceled;
661 assert_eq!(err.to_string(), "buffer service request canceled");
662 }
663
664 #[ntex::test]
665 async fn dropped_buffered_request_is_skipped() {
666 let inner = Rc::new(Inner {
667 ready: Cell::new(false),
668 waker: LocalWaker::default(),
669 count: Cell::new(0),
670 });
671 let srv = Pipeline::new(
672 (),
673 BufferService::new(2, PipelineState::new(TestService(inner.clone()))),
674 );
675 assert_eq!(lazy(|cx| srv.poll_ready(cx)).await, Poll::Ready(Ok(())));
676
677 let mut fut = srv.call_static(());
679 assert!(lazy(|cx| Pin::new(&mut fut).poll(cx)).await.is_pending());
680 drop(fut);
681
682 inner.ready.set(true);
683 inner.waker.wake();
684 assert_eq!(lazy(|cx| srv.poll_ready(cx)).await, Poll::Ready(Ok(())));
685
686 srv.call_static(()).await.unwrap();
688 assert_eq!(inner.count.get(), 1);
689 }
690
691 #[ntex::test]
692 async fn shutdown_skips_dropped_buffered_request() {
693 let inner = Rc::new(Inner {
694 ready: Cell::new(false),
695 waker: LocalWaker::default(),
696 count: Cell::new(0),
697 });
698 let srv = Pipeline::new(
699 (),
700 BufferService::new(2, PipelineState::new(TestService(inner.clone()))),
701 );
702 assert_eq!(lazy(|cx| srv.poll_ready(cx)).await, Poll::Ready(Ok(())));
703
704 let mut fut = srv.call_static(());
705 assert!(lazy(|cx| Pin::new(&mut fut).poll(cx)).await.is_pending());
706 drop(fut);
707
708 inner.ready.set(true);
709 crate::time::timeout(Duration::from_secs(1), srv.shutdown())
710 .await
711 .unwrap();
712 assert_eq!(inner.count.get(), 0);
713 }
714
715 #[ntex::test]
716 async fn shutdown_stops_on_inner_ready_error() {
717 let fail = Rc::new(Cell::new(false));
718 let srv = Pipeline::new(
719 (),
720 BufferService::new(2, PipelineState::new(FailService(fail.clone()))),
721 );
722 assert_eq!(lazy(|cx| srv.poll_ready(cx)).await, Poll::Ready(Ok(())));
723
724 let mut fut = srv.call_static(());
725 assert!(lazy(|cx| Pin::new(&mut fut).poll(cx)).await.is_pending());
726
727 fail.set(true);
728 crate::time::timeout(Duration::from_secs(1), srv.shutdown())
729 .await
730 .unwrap();
731
732 assert!(lazy(|cx| Pin::new(&mut fut).poll(cx)).await.is_pending());
734 }
735}