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