1use std::{fmt, marker::PhantomData};
2
3use super::{Ctx, Service, ServiceFactory};
4
5pub 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 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
78pub 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};
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}