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)]
67pub enum BufferServiceError<E> {
68 Service(E),
70 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
91pub 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 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 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 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 if self.service.poll_ready(cx, ctx.st())?.is_pending() {
167 if buffer.len() < self.size {
168 self.ready.set(false);
170 Poll::Ready(Ok(()))
171 } else {
172 log::trace!("Buffer limit exceeded");
173 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 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 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 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 let _task_guard = rx.recv().await.map_err(|_| {
245 log::trace!("Buffered service request canceled");
246 BufferServiceError::RequestCanceled
247 })?;
248
249 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}