Skip to main content

ntex_service/
map.rs

1use std::{fmt, marker::PhantomData};
2
3use super::{Ctx, Service, ServiceFactory};
4
5/// Service for the `map` combinator, changing the type of a service's response.
6///
7/// This is created by the `ServiceExt::map` method.
8pub struct Map<F, S, Res> {
9    f: F,
10    svc: S,
11    _t: PhantomData<fn() -> Res>,
12}
13
14impl<F, S, Res> Map<F, S, Res> {
15    /// Create new `Map` combinator
16    pub(crate) fn new<St, Req>(f: F, svc: S) -> Self
17    where
18        F: Fn(S::Res) -> Res,
19        S: Service<St, Req>,
20    {
21        Self {
22            f,
23            svc,
24            _t: PhantomData,
25        }
26    }
27}
28
29impl<F, S, Res> Clone for Map<F, S, Res>
30where
31    F: Clone,
32    S: Clone,
33{
34    #[inline]
35    fn clone(&self) -> Self {
36        Map {
37            f: self.f.clone(),
38            svc: self.svc.clone(),
39            _t: PhantomData,
40        }
41    }
42}
43
44impl<F, S, Res> fmt::Debug for Map<F, S, Res>
45where
46    S: fmt::Debug,
47{
48    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
49        f.debug_struct("Map")
50            .field("svc", &self.svc)
51            .field("map", &std::any::type_name::<F>())
52            .finish()
53    }
54}
55
56impl<F, S, St, Req, Res> Service<St, Req> for Map<F, S, Res>
57where
58    S: Service<St, Req>,
59    F: Fn(S::Res) -> Res,
60{
61    type Res = Res;
62    type Error = S::Error;
63
64    #[inline]
65    async fn call(&self, req: Req, ctx: Ctx<'_, Self, St>) -> Result<Res, S::Error> {
66        ctx.call(&self.svc, req).await.map(|r| (self.f)(r))
67    }
68
69    crate::forward_ready!(St, svc);
70    crate::forward_shutdown!(St, svc);
71}
72
73/// `MapNewService` new service combinator
74pub struct MapFactory<F, Sf, Res> {
75    f: F,
76    sf: Sf,
77    r: PhantomData<fn() -> Res>,
78}
79
80impl<F, Sf, Res> MapFactory<F, Sf, Res> {
81    /// Create new `Map` new service instance
82    pub(crate) fn new<St, Req>(f: F, sf: Sf) -> Self
83    where
84        F: Fn(Sf::Res) -> Res,
85        Sf: ServiceFactory<St, Req>,
86    {
87        Self {
88            f,
89            sf,
90            r: PhantomData,
91        }
92    }
93}
94
95impl<F, Sf, Res> Clone for MapFactory<F, Sf, Res>
96where
97    F: Clone,
98    Sf: Clone,
99{
100    #[inline]
101    fn clone(&self) -> Self {
102        Self {
103            sf: self.sf.clone(),
104            f: self.f.clone(),
105            r: PhantomData,
106        }
107    }
108}
109
110impl<F, Sf, Res> fmt::Debug for MapFactory<F, Sf, Res>
111where
112    Sf: fmt::Debug,
113{
114    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
115        f.debug_struct("MapFactory")
116            .field("factory", &self.sf)
117            .field("map", &std::any::type_name::<F>())
118            .finish()
119    }
120}
121
122impl<F, Sf, St, Req, Res> ServiceFactory<St, Req> for MapFactory<F, Sf, Res>
123where
124    F: Fn(Sf::Res) -> Res + Clone,
125    Sf: ServiceFactory<St, Req>,
126{
127    type Res = Res;
128    type Error = Sf::Error;
129
130    type Service = Map<F, Sf::Service, Res>;
131    type InitError = Sf::InitError;
132
133    #[inline]
134    async fn create(&self, st: &St) -> Result<Self::Service, Self::InitError> {
135        Ok(Map {
136            svc: self.sf.create(st).await?,
137            f: self.f.clone(),
138            _t: PhantomData,
139        })
140    }
141}
142
143#[cfg(test)]
144mod tests {
145    use std::{cell::Cell, rc::Rc};
146
147    use crate::{Ctx, Pipeline, Service, ServiceFactory, fn_factory};
148
149    #[derive(Debug, Default, Clone)]
150    struct Srv(Rc<Cell<usize>>);
151
152    impl Service<(), ()> for Srv {
153        type Res = ();
154        type Error = ();
155
156        async fn ready(&self, _: Ctx<'_, Self>) -> Result<(), Self::Error> {
157            Ok(())
158        }
159
160        async fn call(&self, _r: (), _: Ctx<'_, Self>) -> Result<(), ()> {
161            Ok(())
162        }
163
164        async fn shutdown(&self, _: Ctx<'_, Self, ()>) {
165            self.0.set(self.0.get() + 1);
166        }
167    }
168
169    #[ntex::test]
170    async fn test_service() {
171        let cnt_sht = Rc::new(Cell::new(0));
172        let srv = Pipeline::new((), Srv(cnt_sht.clone()).map(|()| "ok").clone());
173        let res = srv.call(()).await;
174        assert!(res.is_ok());
175        assert_eq!(res.unwrap(), "ok");
176
177        let res = srv.ready().await;
178        assert_eq!(res, Ok(()));
179
180        srv.shutdown().await;
181        assert_eq!(cnt_sht.get(), 1);
182        let _ = format!("{srv:?}");
183
184        let cnt_sht = Rc::new(Cell::new(0));
185        let svc = Srv(cnt_sht.clone()).map(|()| "ok");
186        let srv = Pipeline::new((), svc);
187        let res = srv.call(()).await;
188        assert!(res.is_ok());
189        assert_eq!(res.unwrap(), "ok");
190
191        let res = srv.ready().await;
192        assert_eq!(res, Ok(()));
193
194        srv.shutdown().await;
195        assert_eq!(cnt_sht.get(), 1);
196        let _ = format!("{srv:?}");
197    }
198
199    #[ntex::test]
200    async fn test_pipeline() {
201        let srv = Pipeline::new((), crate::service(Srv::default()).map(|()| "ok").clone());
202        let res = srv.call(()).await;
203        assert!(res.is_ok());
204        assert_eq!(res.unwrap(), "ok");
205
206        let res = srv.ready().await;
207        assert_eq!(res, Ok(()));
208    }
209
210    #[ntex::test]
211    async fn test_factory() {
212        let new_srv = fn_factory(|(): &()| async { Ok::<_, ()>(Srv::default()) })
213            .map(|()| "ok")
214            .clone();
215        let srv = Pipeline::new((), new_srv.create(&()).await.unwrap());
216        let res = srv.call(()).await;
217        assert!(res.is_ok());
218        assert_eq!(res.unwrap(), ("ok"));
219
220        let _ = format!("{new_srv:?}");
221    }
222
223    #[ntex::test]
224    async fn test_pipeline_factory() {
225        let new_srv = crate::fn_factory(|(): &()| async { Ok::<_, ()>(Srv::default()) })
226            .map(|()| "ok")
227            .clone();
228        let srv = Pipeline::new((), new_srv.create(&()).await.unwrap());
229        let res = srv.call(()).await;
230        assert!(res.is_ok());
231        assert_eq!(res.unwrap(), ("ok"));
232
233        let _ = format!("{new_srv:?}");
234    }
235}