Skip to content
use actix_web::Error;
use actix_web::dev::{Service, ServiceRequest, ServiceResponse, Transform, forward_ready};
use actix_web::http::header::{CACHE_CONTROL, HeaderValue};
use std::future::{Future, Ready, ready};
use std::pin::Pin;

#[derive(Clone, Debug, Default)]
pub struct CacheControl {
    header_value: Option<HeaderValue>,
}

impl CacheControl {
    pub fn new(header_value: Option<HeaderValue>) -> Self {
        Self { header_value }
    }

    #[cfg(test)]
    pub fn no_cache() -> Self {
        Self {
            header_value: Some(HeaderValue::from_static("no-cache")),
        }
    }
}

impl<S, B> Transform<S, ServiceRequest> for CacheControl
where
    S: Service<ServiceRequest, Response = ServiceResponse<B>, Error = Error> + 'static,
    S::Future: 'static,
    B: 'static,
{
    type Response = ServiceResponse<B>;
    type Error = Error;
    type InitError = ();
    type Transform = CacheControlMiddleware<S>;
    type Future = Ready<Result<Self::Transform, Self::InitError>>;

    fn new_transform(&self, service: S) -> Self::Future {
        ready(Ok(CacheControlMiddleware {
            service,
            header_value: self.header_value.clone(),
        }))
    }
}

pub struct CacheControlMiddleware<S> {
    service: S,
    header_value: Option<HeaderValue>,
}

impl<S, B> Service<ServiceRequest> for CacheControlMiddleware<S>
where
    S: Service<ServiceRequest, Response = ServiceResponse<B>, Error = Error> + 'static,
    S::Future: 'static,
    B: 'static,
{
    type Response = ServiceResponse<B>;
    type Error = Error;
    type Future = Pin<Box<dyn Future<Output = Result<Self::Response, Self::Error>>>>;

    forward_ready!(service);

    fn call(&self, req: ServiceRequest) -> Self::Future {
        let header_val = self.header_value.clone();
        let fut = self.service.call(req);
        Box::pin(async move {
            let mut res = fut.await?;
            if let Some(val) = header_val {
                res.headers_mut().insert(CACHE_CONTROL, val);
            }
            Ok(res)
        })
    }
}

#[cfg(test)]
mod tests {
    use super::*;
    use actix_web::{App, HttpResponse, test as aw_test, web};

    #[actix_web::test]
    async fn test_middleware_skips_cache_when_set() {
        let app = aw_test::init_service(
            App::new()
                .wrap(CacheControl::no_cache())
                .route(
                    "/",
                    web::get().to(|| async { HttpResponse::Ok().body("hello") }),
                )
                .route(
                    "/cached",
                    web::get().to(|| async {
                        HttpResponse::Ok()
                            .insert_header((CACHE_CONTROL, "public, max-age=3600"))
                            .body("cached")
                    }),
                ),
        )
        .await;

        let req = aw_test::TestRequest::get().uri("/").to_request();
        let resp = aw_test::call_service(&app, req).await;
        assert_eq!(resp.status(), actix_web::http::StatusCode::OK);
        assert_eq!(
            resp.headers()
                .get("cache-control")
                .unwrap()
                .to_str()
                .unwrap(),
            "no-cache"
        );

        let req = aw_test::TestRequest::get().uri("/cached").to_request();
        let resp = aw_test::call_service(&app, req).await;
        assert_eq!(resp.status(), actix_web::http::StatusCode::OK);
        assert_eq!(
            resp.headers()
                .get("cache-control")
                .unwrap()
                .to_str()
                .unwrap(),
            "no-cache"
        );
    }

    #[actix_web::test]
    async fn test_middleware_does_not_modify_when_unset() {
        let app = aw_test::init_service(
            App::new()
                .wrap(CacheControl::new(None))
                .route(
                    "/",
                    web::get().to(|| async { HttpResponse::Ok().body("hello") }),
                )
                .route(
                    "/cached",
                    web::get().to(|| async {
                        HttpResponse::Ok()
                            .insert_header((CACHE_CONTROL, "public, max-age=3600"))
                            .body("cached")
                    }),
                ),
        )
        .await;

        let req = aw_test::TestRequest::get().uri("/").to_request();
        let resp = aw_test::call_service(&app, req).await;
        assert_eq!(resp.status(), actix_web::http::StatusCode::OK);
        assert!(resp.headers().get("cache-control").is_none());

        let req = aw_test::TestRequest::get().uri("/cached").to_request();
        let resp = aw_test::call_service(&app, req).await;
        assert_eq!(resp.status(), actix_web::http::StatusCode::OK);
        assert_eq!(
            resp.headers()
                .get("cache-control")
                .unwrap()
                .to_str()
                .unwrap(),
            "public, max-age=3600"
        );
    }
}