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_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
79pub 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 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}