ntex 4.0.0-beta.17

Framework for composable network services
//! Middleware stack types used by applications, scopes, and resources.
use std::marker::PhantomData;

use crate::error::Failure;
use crate::service::{Ctx, Middleware, Service, ServiceFactory};
use crate::web::{State, WebError, WebRequest, WebResponse, WebResponseError};

/// Stack of middlewares.
#[derive(Debug, Clone)]
pub struct WebStack<St, Inner, Outer> {
    inner: Inner,
    outer: Outer,
    err: PhantomData<St>,
}

impl<St, Inner, Outer> WebStack<St, Inner, Outer> {
    /// Create a stack that applies `inner` first and then wraps it with `outer`.
    pub fn new(inner: Inner, outer: Outer) -> Self {
        WebStack {
            inner,
            outer,
            err: PhantomData,
        }
    }
}

impl<S, St, Inner, Outer> Middleware<S, St> for WebStack<St, Inner, Outer>
where
    St: State,
    Inner: Middleware<S, St>,
    Outer: Middleware<Inner::Service, St>,
{
    type Service = WebMiddleware<Outer::Service, St>;

    fn create(&self, st: &St, service: S) -> Self::Service {
        WebMiddleware {
            svc: self.outer.create(st, self.inner.create(st, service)),
            err: PhantomData,
        }
    }
}

/// Service produced by [`WebStack`] layers.
///
/// Wraps a middleware service and converts its errors into [`WebError`].
#[derive(Debug)]
pub struct WebMiddleware<S, St> {
    svc: S,
    err: PhantomData<St>,
}

impl<S, St> Clone for WebMiddleware<S, St>
where
    S: Clone,
{
    fn clone(&self) -> Self {
        Self {
            svc: self.svc.clone(),
            err: PhantomData,
        }
    }
}

impl<S, St, In> Service<St, WebRequest<In>> for WebMiddleware<S, St>
where
    S: Service<St, WebRequest<In>, Res = WebResponse>,
    S::Error: WebResponseError<St, St::Error>,
    St: State,
{
    type Res = WebResponse;
    type Error = WebError<St, St::Error>;

    #[inline]
    async fn call(
        &self,
        req: WebRequest<In>,
        ctx: Ctx<'_, Self, St>,
    ) -> Result<Self::Res, Self::Error> {
        ctx.call(&self.svc, req).await.map_err(WebError::from_err)
    }

    crate::forward_ready!(St, svc, WebError::from_err);
    crate::forward_shutdown!(St, svc);
}

/// Identity request filter.
///
/// The default filter of applications, scopes, and resources. It passes
/// requests through unchanged.
#[derive(derive_more::Debug)]
#[debug("Filter")]
pub struct Filter<St, In>(PhantomData<(St, In)>);

impl<St, In> Filter<St, In> {
    pub(super) fn new() -> Self {
        Filter(PhantomData)
    }
}

impl<St: State, In> ServiceFactory<St, WebRequest<In>> for Filter<St, In> {
    type Res = WebRequest<In>;
    type Error = WebError<St, St::Error>;

    type Service = Filter<St, In>;
    type InitError = Failure;

    async fn create(&self, _: &St) -> Result<Self::Service, Self::InitError> {
        Ok(Filter(PhantomData))
    }
}

impl<St: State, In> Service<St, WebRequest<In>> for Filter<St, In> {
    type Res = WebRequest<In>;
    type Error = WebError<St, St::Error>;

    async fn call(
        &self,
        req: WebRequest<In>,
        _: Ctx<'_, Self, St>,
    ) -> Result<Self::Res, Self::Error> {
        Ok(req)
    }
}

#[cfg(test)]
mod tests {
    use std::io;

    use super::*;
    use crate::http::StatusCode;
    use crate::service::{Identity, Pipeline, fn_service};
    use crate::web::{HttpResponse, test::TestRequest};

    #[crate::rt_test]
    async fn test_web_middleware() {
        let svc = fn_service(async |req: WebRequest<()>| {
            if req.path() == "/err" {
                Err(io::Error::new(io::ErrorKind::NotFound, "not found"))
            } else {
                Ok(req.into_response(HttpResponse::Ok().build()))
            }
        });
        let mw = WebStack::<(), _, _>::new(Identity, Identity).create(&(), svc);
        let srv = Pipeline::new((), mw.clone());

        let res = srv
            .call(TestRequest::default().to_srv_request())
            .await
            .unwrap();
        assert_eq!(res.status(), StatusCode::OK);

        let err = srv
            .call(TestRequest::with_uri("/err").to_srv_request())
            .await
            .unwrap_err();
        assert_eq!(err.to_string(), "not found");
        let res = WebResponseError::error_response(&err, &());
        assert_eq!(res.status(), StatusCode::NOT_FOUND);
    }
}