1use std::{fmt, marker::PhantomData, rc::Rc};
2
3use crate::dev::{Apply, ApplyCtx};
4use crate::{IntoServiceFactory, Service, ServiceChainFactory, ServiceFactory};
5
6pub fn apply<Sf, St, Req, M>(
8 mw: M,
9 factory: impl IntoServiceFactory<Sf, St, Req>,
10) -> ServiceChainFactory<ApplyMiddleware<M, Sf>, St, Req>
11where
12 Sf: ServiceFactory<St, Req>,
13 M: Middleware<Sf::Service, St>,
14{
15 ServiceChainFactory {
16 factory: ApplyMiddleware::new(mw, factory.into_factory()),
17 _t: PhantomData,
18 }
19}
20
21pub trait Middleware<S, St> {
91 type Service;
93
94 fn create(&self, st: &St, service: S) -> Self::Service;
96
97 fn apply_to<Sf, Req>(
102 self,
103 factory: Sf,
104 ) -> ServiceChainFactory<ApplyMiddleware<Self, Sf>, St, Req>
105 where
106 Sf: ServiceFactory<St, Req, Service = S>,
107 Self: Sized,
108 Self::Service: Service<St, Req>,
109 {
110 crate::factory(ApplyMiddleware::new(self, factory))
111 }
112}
113
114impl<M, S, St> Middleware<S, St> for Rc<M>
115where
116 M: Middleware<S, St>,
117{
118 type Service = M::Service;
119
120 fn create(&self, st: &St, service: S) -> M::Service {
121 self.as_ref().create(st, service)
122 }
123}
124
125pub struct ApplyMiddleware<M, Sf>(Rc<(M, Sf)>);
127
128impl<M, Sf> ApplyMiddleware<M, Sf> {
129 pub(crate) fn new(mw: M, sf: Sf) -> Self {
131 Self(Rc::new((mw, sf)))
132 }
133}
134
135impl<M, Sf> Clone for ApplyMiddleware<M, Sf> {
136 fn clone(&self) -> Self {
137 Self(self.0.clone())
138 }
139}
140
141impl<M, Sf> fmt::Debug for ApplyMiddleware<M, Sf>
142where
143 M: fmt::Debug,
144 Sf: fmt::Debug,
145{
146 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
147 f.debug_struct("ApplyMiddleware")
148 .field("factory", &self.0.1)
149 .field("middleware", &self.0.0)
150 .finish()
151 }
152}
153
154impl<M, Sf, St, Req> ServiceFactory<St, Req> for ApplyMiddleware<M, Sf>
155where
156 Sf: ServiceFactory<St, Req>,
157 M: Middleware<Sf::Service, St>,
158 M::Service: Service<St, Req>,
159{
160 type Res = <M::Service as Service<St, Req>>::Res;
161 type Error = <M::Service as Service<St, Req>>::Error;
162
163 type Service = M::Service;
164 type InitError = Sf::InitError;
165
166 #[inline]
167 async fn create(&self, st: &St) -> Result<Self::Service, Self::InitError> {
168 Ok(self.0.0.create(st, self.0.1.create(st).await?))
169 }
170}
171
172#[derive(Debug, Clone, Copy)]
176pub struct Identity;
177
178impl<S, St> Middleware<S, St> for Identity {
179 type Service = S;
180
181 #[inline]
182 fn create(&self, _: &St, service: S) -> Self::Service {
183 service
184 }
185}
186
187#[derive(Debug, Clone)]
189pub struct Stack<Inner, Outer> {
190 inner: Inner,
191 outer: Outer,
192}
193
194impl<Inner, Outer> Stack<Inner, Outer> {
195 pub fn new(inner: Inner, outer: Outer) -> Self {
196 Stack { inner, outer }
197 }
198}
199
200impl<S, St, Inner, Outer> Middleware<S, St> for Stack<Inner, Outer>
201where
202 Inner: Middleware<S, St>,
203 Outer: Middleware<Inner::Service, St>,
204{
205 type Service = Outer::Service;
206
207 fn create(&self, st: &St, service: S) -> Self::Service {
208 self.outer.create(st, self.inner.create(st, service))
209 }
210}
211
212#[doc(hidden)]
213pub fn fn_layer<F, S, St, Req, In, Out, Err>(f: F) -> FnMiddleware<F, S, St, Req, In, Out, Err>
215where
216 F: AsyncFn(In, &ApplyCtx<'_, S, St, Req>) -> Result<Out, Err> + Clone,
217 S: Service<St, Req>,
218{
219 FnMiddleware { f, r: PhantomData }
220}
221
222pub struct FnMiddleware<F, S, St, Req, In, Out, Err> {
224 f: F,
225 r: PhantomData<fn(S, St, Req) -> (In, Out, Err)>,
226}
227
228impl<F, S, St, Req, In, Out, Err> Clone for FnMiddleware<F, S, St, Req, In, Out, Err>
229where
230 F: Clone,
231{
232 fn clone(&self) -> Self {
233 FnMiddleware {
234 f: self.f.clone(),
235 r: PhantomData,
236 }
237 }
238}
239
240impl<F, S, St, Req, In, Out, Err> fmt::Debug for FnMiddleware<F, S, St, Req, In, Out, Err> {
241 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
242 f.debug_struct("FnMiddleware")
243 .field("layer", &std::any::type_name::<F>())
244 .finish()
245 }
246}
247
248impl<F, S, St, Req, In, Out, Err> Middleware<S, St> for FnMiddleware<F, S, St, Req, In, Out, Err>
249where
250 S: Service<St, Req>,
251 F: AsyncFn(In, &ApplyCtx<'_, S, St, Req>) -> Result<Out, Err> + Clone,
252 Err: From<S::Error>,
253{
254 type Service = Apply<S, St, Req, F, In, Out, Err>;
255
256 fn create(&self, _: &St, service: S) -> Self::Service {
257 Apply::new(service, self.f.clone())
258 }
259}
260
261#[cfg(test)]
262#[allow(clippy::redundant_clone)]
263mod tests {
264 use std::{cell::Cell, rc::Rc};
265
266 use super::*;
267 use crate::{Ctx, Pipeline, factory, fn_service};
268
269 #[derive(Debug, Clone)]
270 struct Mw(Rc<Cell<usize>>);
271
272 impl<S, St> Middleware<S, St> for Mw {
273 type Service = Srv<S>;
274
275 fn create(&self, _: &St, service: S) -> Self::Service {
276 self.0.set(self.0.get() + 1);
277 Srv(service, self.0.clone())
278 }
279 }
280
281 #[derive(Debug, Clone)]
282 struct Srv<S>(S, Rc<Cell<usize>>);
283
284 impl<S: Service<(), R>, R> Service<(), R> for Srv<S> {
285 type Res = S::Res;
286 type Error = S::Error;
287
288 async fn ready(&self, ctx: Ctx<'_, Self>) -> Result<(), Self::Error> {
289 ctx.ready(&self.0).await
290 }
291
292 async fn call(&self, req: R, ctx: Ctx<'_, Self>) -> Result<S::Res, S::Error> {
293 ctx.call(&self.0, req).await
294 }
295
296 async fn shutdown(&self, _: Ctx<'_, Self, ()>) {
297 self.1.set(self.1.get() + 1);
298 }
299 }
300
301 #[ntex::test]
302 async fn middleware() {
303 let cnt_sht = Rc::new(Cell::new(0));
304 let fac = apply(
305 Rc::new(Mw(cnt_sht.clone()).clone()),
306 fn_service(|i: usize| async move { Ok::<_, ()>(i * 2) }),
307 )
308 .clone();
309
310 let srv = Pipeline::new((), fac.create(&()).await.unwrap().clone());
311 let res = srv.call(10).await;
312 assert!(res.is_ok());
313 assert_eq!(res.unwrap(), 20);
314 let _ = format!("{fac:?} {srv:?}");
315
316 assert_eq!(srv.ready().await, Ok(()));
317 srv.shutdown().await;
318 assert_eq!(cnt_sht.get(), 2);
319
320 let fac = factory(fn_service(|i: usize| async move { Ok::<_, ()>(i * 2) }))
321 .apply(Rc::new(Mw(Rc::new(Cell::new(0))).clone()))
322 .clone();
323
324 let srv = Pipeline::new((), fac.create(&()).await.unwrap().clone());
325 let res = srv.call(10).await;
326 assert!(res.is_ok());
327 assert_eq!(res.unwrap(), 20);
328 let _ = format!("{fac:?} {srv:?}");
329
330 assert_eq!(srv.ready().await, Ok(()));
331 }
332
333 #[ntex::test]
334 async fn middleware_apply() {
335 let cnt_sht = Rc::new(Cell::new(0));
336 let fac = Mw(cnt_sht.clone())
337 .apply_to(factory(async |i: usize| Ok::<_, ()>(i * 2)))
338 .boxed();
339
340 let srv = Pipeline::new((), fac.create(&()).await.unwrap());
341 let res = srv.call(10).await;
342 assert!(res.is_ok());
343 assert_eq!(res.unwrap(), 20);
344 let _ = format!("{fac:?} {srv:?}");
345
346 assert_eq!(srv.ready().await, Ok(()));
347 srv.shutdown().await;
348 assert_eq!(cnt_sht.get(), 2);
349 }
350
351 #[ntex::test]
352 async fn middleware_chain() {
353 let cnt_sht = Rc::new(Cell::new(0));
354 let fac = factory(fn_service(async move |i: usize| Ok::<_, ()>(i * 2)))
355 .apply(Mw(cnt_sht.clone()).clone());
356
357 let srv = Pipeline::new((), fac.create(&()).await.unwrap().clone());
358 let res = srv.call(10).await;
359 assert!(res.is_ok());
360 assert_eq!(res.unwrap(), 20);
361 let _ = format!("{fac:?} {srv:?}");
362
363 assert_eq!(srv.ready().await, Ok(()));
364 srv.shutdown().await;
365 assert_eq!(cnt_sht.get(), 2);
366 }
367
368 #[ntex::test]
369 async fn stack() {
370 let cnt_sht = Rc::new(Cell::new(0));
371 let mw = Stack::new(Identity, Mw(cnt_sht.clone()));
372 let _ = format!("{mw:?}");
373
374 let pl = Pipeline::new(
375 (),
376 Middleware::create(
377 &mw,
378 &(),
379 fn_service(|i: usize| async move { Ok::<_, ()>(i * 2) }),
380 ),
381 );
382 let res = pl.call(10).await;
383 assert!(res.is_ok());
384 assert_eq!(res.unwrap(), 20);
385 assert_eq!(pl.ready().await, Ok(()));
386 pl.shutdown().await;
387 assert_eq!(cnt_sht.get(), 2);
388 }
389
390 #[ntex::test]
391 async fn fn_middleware_service() {
392 let cnt_sht = Rc::new(Cell::new(0));
393 let cnt_sht2 = cnt_sht.clone();
394 let mw = fn_layer(async move |req: &'static str, svc| {
395 cnt_sht2.set(cnt_sht2.get() + 1);
396 let result = svc.call(1).await?;
397 Ok::<_, ()>((req, result))
398 })
399 .clone();
400 let _ = format!("{mw:?}");
401
402 let svc = Pipeline::new(
403 (),
404 mw.create(&(), fn_service(async move |i: usize| Ok::<_, ()>(i * 2))),
405 );
406
407 let res = svc.call("test").await;
408 assert!(res.is_ok());
409 assert_eq!(res.unwrap(), ("test", 2));
410 let _ = format!("{svc:?}");
411
412 assert_eq!(svc.ready().await, Ok(()));
413 svc.shutdown().await;
414 assert_eq!(cnt_sht.get(), 1);
415 }
416}