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