1use std::{fmt, marker::PhantomData};
2
3use super::{Ctx, Service, ServiceFactory};
4
5pub 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 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(&self.svc, req).await.map_err(|e| (self.f)(e))
67 }
68
69 #[inline]
70 async fn ready(&self, ctx: Ctx<'_, Self, St>) -> Result<(), E> {
71 ctx.ready(&self.svc).await.map_err(&self.f)
72 }
73
74 crate::forward_shutdown!(St, svc);
75}
76
77pub struct MapErrFactory<F, Sf, E> {
81 f: F,
82 sf: Sf,
83 e: PhantomData<fn(Sf) -> E>,
84}
85
86impl<F, Sf, E> MapErrFactory<F, Sf, E> {
87 pub(crate) fn new<St, Req>(f: F, sf: Sf) -> Self
89 where
90 Sf: ServiceFactory<St, Req>,
91 F: Fn(Sf::Error) -> E + Clone,
92 {
93 Self {
94 f,
95 sf,
96 e: PhantomData,
97 }
98 }
99}
100
101impl<F: Clone, Sf: Clone, E> Clone for MapErrFactory<F, Sf, E> {
102 fn clone(&self) -> Self {
103 Self {
104 f: self.f.clone(),
105 sf: self.sf.clone(),
106 e: PhantomData,
107 }
108 }
109}
110
111impl<F, Sf, E> fmt::Debug for MapErrFactory<F, Sf, E>
112where
113 Sf: fmt::Debug,
114{
115 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
116 f.debug_struct("MapErrFactory")
117 .field("sf", &self.sf)
118 .field("map", &std::any::type_name::<F>())
119 .finish()
120 }
121}
122
123impl<F, Sf, St, Req, E> ServiceFactory<St, Req> for MapErrFactory<F, Sf, E>
124where
125 Sf: ServiceFactory<St, Req>,
126 F: Fn(Sf::Error) -> E + Clone,
127{
128 type Res = Sf::Res;
129 type Error = E;
130
131 type Service = MapErr<F, Sf::Service, E>;
132 type InitError = Sf::InitError;
133
134 #[inline]
135 async fn create(&self, st: &St) -> Result<Self::Service, Self::InitError> {
136 self.sf.create(st).await.map(|svc| MapErr {
137 svc,
138 f: self.f.clone(),
139 e: PhantomData,
140 })
141 }
142}
143
144#[cfg(test)]
145mod tests {
146 use std::{cell::Cell, rc::Rc};
147
148 use super::*;
149 use crate::{Pipeline, fn_factory, service};
150
151 #[derive(Debug, Clone)]
152 struct Srv(bool, Rc<Cell<usize>>);
153
154 impl Service<(), ()> for Srv {
155 type Res = ();
156 type Error = ();
157
158 async fn ready(&self, _: Ctx<'_, Self>) -> Result<(), Self::Error> {
159 if self.0 { Err(()) } else { Ok(()) }
160 }
161
162 async fn call(&self, _m: (), _: Ctx<'_, Self>) -> Result<(), ()> {
163 Err(())
164 }
165
166 async fn shutdown(&self, _: Ctx<'_, Self, ()>) {
167 self.1.set(self.1.get() + 1);
168 }
169 }
170
171 #[ntex::test]
172 async fn test_ready() {
173 let cnt_sht = Rc::new(Cell::new(0));
174 let srv = Pipeline::new(
175 (),
176 service(Srv(true, cnt_sht.clone())).map_err(|()| "error"),
177 );
178 let res = srv.ready().await;
179 assert_eq!(res, Err("error"));
180
181 srv.shutdown().await;
182 assert_eq!(cnt_sht.get(), 1);
183 }
184
185 #[ntex::test]
186 async fn test_service() {
187 let srv = Pipeline::new(
188 (),
189 Srv(false, Rc::new(Cell::new(0)))
190 .map_err(|()| "error")
191 .clone(),
192 );
193 let res = srv.call(()).await;
194 assert!(res.is_err());
195 assert_eq!(res.err().unwrap(), "error");
196
197 let _ = format!("{srv:?}");
198 }
199
200 #[ntex::test]
201 async fn test_pipeline() {
202 let srv = Pipeline::new(
203 (),
204 crate::service(Srv(false, Rc::new(Cell::new(0))))
205 .map_err(|()| "error")
206 .clone(),
207 );
208 let res = srv.call(()).await;
209 assert!(res.is_err());
210 assert_eq!(res.err().unwrap(), "error");
211
212 let _ = format!("{srv:?}");
213 }
214
215 #[ntex::test]
216 async fn test_factory() {
217 let new_srv =
218 crate::fn_factory(|(): &()| async { Ok::<_, ()>(Srv(false, Rc::new(Cell::new(0)))) })
219 .map_err(|()| "error");
220 let srv = Pipeline::new((), new_srv.create(&()).await.unwrap());
221 let res = srv.call(()).await;
222 assert!(res.is_err());
223 assert_eq!(res.err().unwrap(), "error");
224 let _ = format!("{new_srv:?}");
225 }
226
227 #[ntex::test]
228 async fn test_pipeline_factory() {
229 let new_srv =
230 fn_factory(|(): &()| async move { Ok::<Srv, ()>(Srv(false, Rc::new(Cell::new(0)))) })
231 .map_err(|()| "error")
232 .clone();
233 let srv = Pipeline::new((), new_srv.create(&()).await.unwrap());
234 let res = srv.call(()).await;
235 assert!(res.is_err());
236 assert_eq!(res.err().unwrap(), "error");
237 let _ = format!("{new_srv:?}");
238 }
239}