Skip to main content

ntex_service/
map_err.rs

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