Skip to main content

ntex_util/services/
buffer.rs

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