Skip to main content

ntex_service/
map_err.rs

1use std::{fmt, marker::PhantomData};
2
3use super::{Ctx, Service, ServiceFactory};
4
5/// Service for the `map_err` combinator, changing the type of a service's
6/// error.
7///
8/// This is created by the `ServiceExt::map_err` method.
9pub struct MapErr<F, S, E> {
10    f: F,
11    svc: S,
12    e: PhantomData<E>,
13}
14
15impl<F, S, E> MapErr<F, S, E> {
16    /// Create new `MapErr` combinator
17    pub(crate) fn new<St, Req>(f: F, svc: S) -> Self
18    where
19        S: Service<St, Req>,
20        F: Fn(S::Error) -> E,
21    {
22        Self {
23            f,
24            svc,
25            e: PhantomData,
26        }
27    }
28}
29
30impl<F, S, E> Clone for MapErr<F, S, E>
31where
32    F: Clone,
33    S: Clone,
34{
35    #[inline]
36    fn clone(&self) -> Self {
37        MapErr {
38            f: self.f.clone(),
39            svc: self.svc.clone(),
40            e: PhantomData,
41        }
42    }
43}
44
45impl<F, S, E> fmt::Debug for MapErr<F, S, E>
46where
47    S: fmt::Debug,
48{
49    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
50        f.debug_struct("MapErr")
51            .field("svc", &self.svc)
52            .field("map", &std::any::type_name::<F>())
53            .finish()
54    }
55}
56
57impl<F, S, St, Req, E> Service<St, Req> for MapErr<F, S, E>
58where
59    S: Service<St, Req>,
60    F: Fn(S::Error) -> E,
61{
62    type Res = S::Res;
63    type Error = E;
64
65    #[inline]
66    async fn call(&self, req: Req, ctx: Ctx<'_, Self, St>) -> Result<S::Res, E> {
67        ctx.call(&self.svc, req).await.map_err(|e| (self.f)(e))
68    }
69
70    #[inline]
71    async fn ready(&self, ctx: Ctx<'_, Self, St>) -> Result<(), E> {
72        ctx.ready(&self.svc).await.map_err(&self.f)
73    }
74
75    crate::forward_shutdown!(St, svc);
76}
77
78/// Factory for the `map_err` combinator, changing the type of a new
79/// service's error.
80///
81/// This is created by the `ServiceFactory::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    /// Create new `MapErr` new service instance
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};
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((), Srv(true, cnt_sht.clone()).map_err(|()| "error"));
177        let res = srv.ready().await;
178        assert_eq!(res, Err("error"));
179
180        srv.shutdown().await;
181        assert_eq!(cnt_sht.get(), 1);
182    }
183
184    #[ntex::test]
185    async fn test_service() {
186        let srv = Pipeline::new(
187            (),
188            Srv(false, Rc::new(Cell::new(0)))
189                .map_err(|()| "error")
190                .clone(),
191        );
192        let res = srv.call(()).await;
193        assert!(res.is_err());
194        assert_eq!(res.err().unwrap(), "error");
195
196        let _ = format!("{srv:?}");
197    }
198
199    #[ntex::test]
200    async fn test_pipeline() {
201        let srv = Pipeline::new(
202            (),
203            crate::service(Srv(false, Rc::new(Cell::new(0))))
204                .map_err(|()| "error")
205                .clone(),
206        );
207        let res = srv.call(()).await;
208        assert!(res.is_err());
209        assert_eq!(res.err().unwrap(), "error");
210
211        let _ = format!("{srv:?}");
212    }
213
214    #[ntex::test]
215    async fn test_factory() {
216        let new_srv =
217            crate::fn_factory(|(): &()| async { Ok::<_, ()>(Srv(false, Rc::new(Cell::new(0)))) })
218                .map_err(|()| "error");
219        let srv = Pipeline::new((), new_srv.create(&()).await.unwrap());
220        let res = srv.call(()).await;
221        assert!(res.is_err());
222        assert_eq!(res.err().unwrap(), "error");
223        let _ = format!("{new_srv:?}");
224    }
225
226    #[ntex::test]
227    async fn test_pipeline_factory() {
228        let new_srv =
229            fn_factory(|(): &()| async move { Ok::<Srv, ()>(Srv(false, Rc::new(Cell::new(0)))) })
230                .map_err(|()| "error")
231                .clone();
232        let srv = Pipeline::new((), new_srv.create(&()).await.unwrap());
233        let res = srv.call(()).await;
234        assert!(res.is_err());
235        assert_eq!(res.err().unwrap(), "error");
236        let _ = format!("{new_srv:?}");
237    }
238}