use std::{fmt, marker::PhantomData, rc::Rc};
use crate::dev::{Apply, ApplyCtx};
use crate::{IntoServiceFactory, Service, ServiceChainFactory, ServiceFactory};
pub fn apply<Sf, St, Req, M>(
mw: M,
factory: impl IntoServiceFactory<Sf, St, Req>,
) -> ServiceChainFactory<ApplyMiddleware<M, Sf>, St, Req>
where
Sf: ServiceFactory<St, Req>,
M: Middleware<Sf::Service, St>,
{
ServiceChainFactory {
factory: ApplyMiddleware::new(mw, factory.into_factory()),
_t: PhantomData,
}
}
pub trait Middleware<S, St> {
type Service;
fn create(&self, st: &St, service: S) -> Self::Service;
fn apply_to<Sf, Req>(
self,
factory: Sf,
) -> ServiceChainFactory<ApplyMiddleware<Self, Sf>, St, Req>
where
Sf: ServiceFactory<St, Req, Service = S>,
Self: Sized,
Self::Service: Service<St, Req>,
{
crate::factory(ApplyMiddleware::new(self, factory))
}
}
impl<M, S, St> Middleware<S, St> for Rc<M>
where
M: Middleware<S, St>,
{
type Service = M::Service;
fn create(&self, st: &St, service: S) -> M::Service {
self.as_ref().create(st, service)
}
}
pub struct ApplyMiddleware<M, Sf>(Rc<(M, Sf)>);
impl<M, Sf> ApplyMiddleware<M, Sf> {
pub(crate) fn new(mw: M, sf: Sf) -> Self {
Self(Rc::new((mw, sf)))
}
}
impl<M, Sf> Clone for ApplyMiddleware<M, Sf> {
fn clone(&self) -> Self {
Self(self.0.clone())
}
}
impl<M, Sf> fmt::Debug for ApplyMiddleware<M, Sf>
where
M: fmt::Debug,
Sf: fmt::Debug,
{
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("ApplyMiddleware")
.field("factory", &self.0.1)
.field("middleware", &self.0.0)
.finish()
}
}
impl<M, Sf, St, Req> ServiceFactory<St, Req> for ApplyMiddleware<M, Sf>
where
Sf: ServiceFactory<St, Req>,
M: Middleware<Sf::Service, St>,
M::Service: Service<St, Req>,
{
type Res = <M::Service as Service<St, Req>>::Res;
type Error = <M::Service as Service<St, Req>>::Error;
type Service = M::Service;
type InitError = Sf::InitError;
#[inline]
async fn create(&self, st: &St) -> Result<Self::Service, Self::InitError> {
Ok(self.0.0.create(st, self.0.1.create(st).await?))
}
}
#[derive(Debug, Clone, Copy)]
pub struct Identity;
impl<S, St> Middleware<S, St> for Identity {
type Service = S;
#[inline]
fn create(&self, _: &St, service: S) -> Self::Service {
service
}
}
#[derive(Debug, Clone)]
pub struct Stack<Inner, Outer> {
inner: Inner,
outer: Outer,
}
impl<Inner, Outer> Stack<Inner, Outer> {
pub fn new(inner: Inner, outer: Outer) -> Self {
Stack { inner, outer }
}
}
impl<S, St, Inner, Outer> Middleware<S, St> for Stack<Inner, Outer>
where
Inner: Middleware<S, St>,
Outer: Middleware<Inner::Service, St>,
{
type Service = Outer::Service;
fn create(&self, st: &St, service: S) -> Self::Service {
self.outer.create(st, self.inner.create(st, service))
}
}
#[doc(hidden)]
pub fn fn_layer<F, S, St, Req, In, Out, Err>(f: F) -> FnMiddleware<F, S, St, Req, In, Out, Err>
where
F: AsyncFn(In, &ApplyCtx<'_, S, St, Req>) -> Result<Out, Err> + Clone,
S: Service<St, Req>,
{
FnMiddleware { f, r: PhantomData }
}
pub struct FnMiddleware<F, S, St, Req, In, Out, Err> {
f: F,
r: PhantomData<fn(S, St, Req) -> (In, Out, Err)>,
}
impl<F, S, St, Req, In, Out, Err> Clone for FnMiddleware<F, S, St, Req, In, Out, Err>
where
F: Clone,
{
fn clone(&self) -> Self {
FnMiddleware {
f: self.f.clone(),
r: PhantomData,
}
}
}
impl<F, S, St, Req, In, Out, Err> fmt::Debug for FnMiddleware<F, S, St, Req, In, Out, Err> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("FnMiddleware")
.field("layer", &std::any::type_name::<F>())
.finish()
}
}
impl<F, S, St, Req, In, Out, Err> Middleware<S, St> for FnMiddleware<F, S, St, Req, In, Out, Err>
where
S: Service<St, Req>,
F: AsyncFn(In, &ApplyCtx<'_, S, St, Req>) -> Result<Out, Err> + Clone,
Err: From<S::Error>,
{
type Service = Apply<S, St, Req, F, In, Out, Err>;
fn create(&self, _: &St, service: S) -> Self::Service {
Apply::new(service, self.f.clone())
}
}
#[cfg(test)]
#[allow(clippy::redundant_clone)]
mod tests {
use std::{cell::Cell, rc::Rc};
use super::*;
use crate::{Ctx, Pipeline, factory, fn_service};
#[derive(Debug, Clone)]
struct Mw(Rc<Cell<usize>>);
impl<S, St> Middleware<S, St> for Mw {
type Service = Srv<S>;
fn create(&self, _: &St, service: S) -> Self::Service {
self.0.set(self.0.get() + 1);
Srv(service, self.0.clone())
}
}
#[derive(Debug, Clone)]
struct Srv<S>(S, Rc<Cell<usize>>);
impl<S: Service<(), R>, R> Service<(), R> for Srv<S> {
type Res = S::Res;
type Error = S::Error;
async fn ready(&self, ctx: Ctx<'_, Self>) -> Result<(), Self::Error> {
ctx.ready(&self.0).await
}
async fn call(&self, req: R, ctx: Ctx<'_, Self>) -> Result<S::Res, S::Error> {
ctx.call(&self.0, req).await
}
async fn shutdown(&self, _: Ctx<'_, Self, ()>) {
self.1.set(self.1.get() + 1);
}
}
#[ntex::test]
async fn middleware() {
let cnt_sht = Rc::new(Cell::new(0));
let fac = apply(
Rc::new(Mw(cnt_sht.clone()).clone()),
fn_service(|i: usize| async move { Ok::<_, ()>(i * 2) }),
)
.clone();
let srv = Pipeline::new((), fac.create(&()).await.unwrap().clone());
let res = srv.call(10).await;
assert!(res.is_ok());
assert_eq!(res.unwrap(), 20);
let _ = format!("{fac:?} {srv:?}");
assert_eq!(srv.ready().await, Ok(()));
srv.shutdown().await;
assert_eq!(cnt_sht.get(), 2);
let fac = factory(fn_service(|i: usize| async move { Ok::<_, ()>(i * 2) }))
.apply(Rc::new(Mw(Rc::new(Cell::new(0))).clone()))
.clone();
let srv = Pipeline::new((), fac.create(&()).await.unwrap().clone());
let res = srv.call(10).await;
assert!(res.is_ok());
assert_eq!(res.unwrap(), 20);
let _ = format!("{fac:?} {srv:?}");
assert_eq!(srv.ready().await, Ok(()));
}
#[ntex::test]
async fn middleware_apply() {
let cnt_sht = Rc::new(Cell::new(0));
let fac = Mw(cnt_sht.clone())
.apply_to(factory(async |i: usize| Ok::<_, ()>(i * 2)))
.boxed();
let srv = Pipeline::new((), fac.create(&()).await.unwrap());
let res = srv.call(10).await;
assert!(res.is_ok());
assert_eq!(res.unwrap(), 20);
let _ = format!("{fac:?} {srv:?}");
assert_eq!(srv.ready().await, Ok(()));
srv.shutdown().await;
assert_eq!(cnt_sht.get(), 2);
}
#[ntex::test]
async fn middleware_chain() {
let cnt_sht = Rc::new(Cell::new(0));
let fac = factory(fn_service(async move |i: usize| Ok::<_, ()>(i * 2)))
.apply(Mw(cnt_sht.clone()).clone());
let srv = Pipeline::new((), fac.create(&()).await.unwrap().clone());
let res = srv.call(10).await;
assert!(res.is_ok());
assert_eq!(res.unwrap(), 20);
let _ = format!("{fac:?} {srv:?}");
assert_eq!(srv.ready().await, Ok(()));
srv.shutdown().await;
assert_eq!(cnt_sht.get(), 2);
}
#[ntex::test]
async fn stack() {
let cnt_sht = Rc::new(Cell::new(0));
let mw = Stack::new(Identity, Mw(cnt_sht.clone()));
let _ = format!("{mw:?}");
let pl = Pipeline::new(
(),
Middleware::create(
&mw,
&(),
fn_service(|i: usize| async move { Ok::<_, ()>(i * 2) }),
),
);
let res = pl.call(10).await;
assert!(res.is_ok());
assert_eq!(res.unwrap(), 20);
assert_eq!(pl.ready().await, Ok(()));
pl.shutdown().await;
assert_eq!(cnt_sht.get(), 2);
}
#[ntex::test]
async fn fn_middleware_service() {
let cnt_sht = Rc::new(Cell::new(0));
let cnt_sht2 = cnt_sht.clone();
let mw = fn_layer(async move |req: &'static str, svc| {
cnt_sht2.set(cnt_sht2.get() + 1);
let result = svc.call(1).await?;
Ok::<_, ()>((req, result))
})
.clone();
let _ = format!("{mw:?}");
let svc = Pipeline::new(
(),
mw.create(&(), fn_service(async move |i: usize| Ok::<_, ()>(i * 2))),
);
let res = svc.call("test").await;
assert!(res.is_ok());
assert_eq!(res.unwrap(), ("test", 2));
let _ = format!("{svc:?}");
assert_eq!(svc.ready().await, Ok(()));
svc.shutdown().await;
assert_eq!(cnt_sht.get(), 1);
}
}