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"
);
}
}