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/// Buffer - service factory for service that can buffer incoming request.
12///
13/// Default number of buffered requests is 16
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    /// Set size of the buffer.
22    ///
23    /// Default is set to 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 flushed during `poll_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#[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
88/// Buffer service - service that can buffer incoming requests.
89///
90/// Default number of buffered requests is 16
91pub 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        // hold advancement until the last released task either makes a call or is dropped
152        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            // handle inner service readiness
161            if self.service.poll_ready(cx, ctx.st())?.is_pending() {
162                if buffer.len() < self.size {
163                    // buffer next request
164                    self.ready.set(false);
165                    Poll::Ready(Ok(()))
166                } else {
167                    log::trace!("Buffer limit exceeded");
168                    // service is not ready
169                    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                        // the task is gone
177                        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        // hold advancement until the last released task either makes a call or is dropped
193        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                        // the task is gone
214                        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            // release
239            let _task_guard = rx.recv().await.map_err(|_| {
240                log::trace!("Buffered service request canceled");
241                BufferServiceError::RequestCanceled
242            })?;
243
244            // call service
245            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}