Skip to main content

ntex_service/
fn_shutdown.rs

1use std::{cell::Cell, convert::Infallible, fmt, marker::PhantomData};
2
3use crate::{Ctx, Service, ServiceFactory};
4
5/// A pass-through service that invokes an asynchronous shutdown callback once.
6pub struct FnShutdown<F, Err> {
7    f_shutdown: Cell<Option<F>>,
8    err: PhantomData<Err>,
9}
10
11impl<F, Err> FnShutdown<F, Err> {
12    /// Creates a service with the supplied shutdown callback.
13    pub fn new<St>(f: F) -> Self
14    where
15        F: AsyncFnOnce(&St),
16    {
17        Self {
18            f_shutdown: Cell::new(Some(f)),
19            err: PhantomData,
20        }
21    }
22}
23
24impl<F, Err> Clone for FnShutdown<F, Err>
25where
26    F: Clone,
27{
28    #[inline]
29    fn clone(&self) -> Self {
30        let f = self.f_shutdown.take();
31        self.f_shutdown.set(f.clone());
32        Self {
33            f_shutdown: Cell::new(f),
34            err: PhantomData,
35        }
36    }
37}
38
39impl<F, Err> fmt::Debug for FnShutdown<F, Err> {
40    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
41        f.debug_struct("FnShutdown")
42            .field("fn", &std::any::type_name::<F>())
43            .finish()
44    }
45}
46
47impl<F, St, Req, Err> Service<St, Req> for FnShutdown<F, Err>
48where
49    F: AsyncFnOnce(&St),
50{
51    type Res = Req;
52    type Error = Err;
53
54    #[inline]
55    async fn shutdown(&self, ctx: Ctx<'_, Self, St>) {
56        if let Some(f) = self.f_shutdown.take() {
57            (f)(ctx.st()).await;
58        }
59    }
60
61    #[inline]
62    async fn call(&self, req: Req, _: Ctx<'_, Self, St>) -> Result<Req, Err> {
63        Ok(req)
64    }
65}
66
67impl<F, St, Req, Err> ServiceFactory<St, Req> for FnShutdown<F, Err>
68where
69    F: AsyncFnOnce(&St) + Clone,
70{
71    type Res = Req;
72    type Error = Err;
73
74    type Service = FnShutdown<F, Err>;
75    type InitError = Infallible;
76
77    #[inline]
78    async fn create(&self, _: &St) -> Result<Self::Service, Self::InitError> {
79        if let Some(f) = self.f_shutdown.take() {
80            self.f_shutdown.set(Some(f.clone()));
81            Ok(FnShutdown {
82                f_shutdown: Cell::new(Some(f)),
83                err: PhantomData,
84            })
85        } else {
86            panic!("FnShutdown was used already");
87        }
88    }
89}
90
91#[cfg(test)]
92mod tests {
93    use std::{future::poll_fn, rc::Rc};
94
95    use crate::{Pipeline, factory, service};
96
97    use super::*;
98
99    #[ntex::test]
100    async fn test_fn_shutdown() {
101        // Service factory
102        let is_called = Rc::new(Cell::new(false));
103        let is_called2 = is_called.clone();
104        let fac = factory(|()| async { Ok::<_, ()>("pipe") }).shutdown(async move |()| {
105            is_called2.set(true);
106        });
107        let _ = format!("{fac:?}");
108
109        let pipe = Pipeline::new((), fac.clone().create(&()).await.unwrap());
110
111        let res = pipe.call(()).await;
112        assert_eq!(pipe.ready().await, Ok(()));
113        assert!(res.is_ok());
114        assert_eq!(res.unwrap(), "pipe");
115        assert!(!pipe.is_shutdown());
116        pipe.shutdown().await;
117        assert!(is_called.get());
118        assert!(pipe.is_shutdown());
119
120        poll_fn(|cx| pipe.poll_shutdown(cx)).await;
121        assert!(pipe.is_shutdown());
122
123        // Service
124        let is_called = Rc::new(Cell::new(false));
125        let is_called2 = is_called.clone();
126        let svc = service(|()| async { Ok::<_, ()>("pipe") }).shutdown(async move |()| {
127            is_called2.set(true);
128        });
129        let _ = format!("{fac:?}");
130
131        let pipe = Pipeline::new((), svc);
132
133        let res = pipe.call(()).await;
134        assert_eq!(pipe.ready().await, Ok(()));
135        assert!(res.is_ok());
136        assert_eq!(res.unwrap(), "pipe");
137        assert!(!pipe.is_shutdown());
138        pipe.shutdown().await;
139        assert!(is_called.get());
140        assert!(pipe.is_shutdown());
141
142        poll_fn(|cx| pipe.poll_shutdown(cx)).await;
143        assert!(pipe.is_shutdown());
144    }
145}