use bytes::Bytes;
use http::{Request, Response, StatusCode};
use http_body_util::Full;
use http_cache::{CacheManager, HttpResponse, Result};
use http_cache_semantics::CachePolicy;
use http_cache_tower_server::{
CustomKeyer, DefaultKeyer, Keyer, QueryKeyer, ServerCacheLayer,
ServerCacheOptions,
};
use std::collections::HashMap;
use std::sync::{Arc, Mutex};
use tower::{Layer, Service, ServiceExt};
#[derive(Debug, Clone, PartialEq)]
struct PathParams {
id: String,
}
#[derive(Clone)]
struct MemoryCacheManager {
store: Arc<Mutex<HashMap<String, (HttpResponse, CachePolicy)>>>,
}
impl MemoryCacheManager {
fn new() -> Self {
Self { store: Arc::new(Mutex::new(HashMap::new())) }
}
}
impl CacheManager for MemoryCacheManager {
async fn get(
&self,
cache_key: &str,
) -> Result<Option<(HttpResponse, CachePolicy)>> {
Ok(self.store.lock().unwrap().get(cache_key).cloned())
}
async fn put(
&self,
cache_key: String,
res: HttpResponse,
policy: CachePolicy,
) -> Result<HttpResponse> {
self.store.lock().unwrap().insert(cache_key, (res.clone(), policy));
Ok(res)
}
async fn delete(&self, cache_key: &str) -> Result<()> {
self.store.lock().unwrap().remove(cache_key);
Ok(())
}
}
#[tokio::test]
async fn test_cache_hit_and_miss() {
let manager = MemoryCacheManager::new();
let layer = ServerCacheLayer::new(manager.clone());
let mut service =
layer.layer(tower::service_fn(|_req: Request<Full<Bytes>>| async {
Ok::<_, std::io::Error>(
Response::builder()
.status(StatusCode::OK)
.header("cache-control", "max-age=60")
.body(Full::new(Bytes::from("Hello, World!")))
.unwrap(),
)
}));
let req = Request::get("/test").body(Full::new(Bytes::new())).unwrap();
let res = service.ready().await.unwrap().call(req).await.unwrap();
assert_eq!(res.headers().get("x-cache").unwrap(), "MISS");
tokio::time::sleep(tokio::time::Duration::from_millis(100)).await;
let req = Request::get("/test").body(Full::new(Bytes::new())).unwrap();
let res = service.ready().await.unwrap().call(req).await.unwrap();
assert_eq!(res.headers().get("x-cache").unwrap(), "HIT");
}
#[tokio::test]
async fn test_no_store_directive() {
let manager = MemoryCacheManager::new();
let layer = ServerCacheLayer::new(manager.clone());
let mut service =
layer.layer(tower::service_fn(|_req: Request<Full<Bytes>>| async {
Ok::<_, std::io::Error>(
Response::builder()
.status(StatusCode::OK)
.header("cache-control", "no-store")
.body(Full::new(Bytes::from("Don't cache me")))
.unwrap(),
)
}));
let req = Request::get("/no-store").body(Full::new(Bytes::new())).unwrap();
let res = service.ready().await.unwrap().call(req).await.unwrap();
assert!(
res.headers().get("x-cache").is_none()
|| res.headers().get("x-cache").unwrap() != "MISS"
);
tokio::time::sleep(tokio::time::Duration::from_millis(100)).await;
let req = Request::get("/no-store").body(Full::new(Bytes::new())).unwrap();
let res = service.ready().await.unwrap().call(req).await.unwrap();
assert!(
res.headers().get("x-cache").is_none()
|| res.headers().get("x-cache").unwrap() != "HIT"
);
}
#[tokio::test]
async fn test_private_directive() {
let manager = MemoryCacheManager::new();
let layer = ServerCacheLayer::new(manager.clone());
let mut service =
layer.layer(tower::service_fn(|_req: Request<Full<Bytes>>| async {
Ok::<_, std::io::Error>(
Response::builder()
.status(StatusCode::OK)
.header("cache-control", "private, max-age=60")
.body(Full::new(Bytes::from("Private data")))
.unwrap(),
)
}));
let req = Request::get("/private").body(Full::new(Bytes::new())).unwrap();
let _res = service.ready().await.unwrap().call(req).await.unwrap();
tokio::time::sleep(tokio::time::Duration::from_millis(100)).await;
let req = Request::get("/private").body(Full::new(Bytes::new())).unwrap();
let res = service.ready().await.unwrap().call(req).await.unwrap();
assert!(
res.headers().get("x-cache").is_none()
|| res.headers().get("x-cache").unwrap() != "HIT"
);
}
#[tokio::test]
async fn test_s_maxage_override() {
let manager = MemoryCacheManager::new();
let layer = ServerCacheLayer::new(manager.clone());
let mut service =
layer.layer(tower::service_fn(|_req: Request<Full<Bytes>>| async {
Ok::<_, std::io::Error>(
Response::builder()
.status(StatusCode::OK)
.header("cache-control", "max-age=60, s-maxage=120")
.body(Full::new(Bytes::from("Shared cache data")))
.unwrap(),
)
}));
let req = Request::get("/s-maxage").body(Full::new(Bytes::new())).unwrap();
let res = service.ready().await.unwrap().call(req).await.unwrap();
assert_eq!(res.headers().get("x-cache").unwrap(), "MISS");
tokio::time::sleep(tokio::time::Duration::from_millis(100)).await;
let req = Request::get("/s-maxage").body(Full::new(Bytes::new())).unwrap();
let res = service.ready().await.unwrap().call(req).await.unwrap();
assert_eq!(res.headers().get("x-cache").unwrap(), "HIT");
}
#[tokio::test]
async fn test_only_cache_success_status() {
let manager = MemoryCacheManager::new();
let layer = ServerCacheLayer::new(manager.clone());
let mut service =
layer.layer(tower::service_fn(|_req: Request<Full<Bytes>>| async {
Ok::<_, std::io::Error>(
Response::builder()
.status(StatusCode::NOT_FOUND)
.header("cache-control", "max-age=60")
.body(Full::new(Bytes::from("Not found")))
.unwrap(),
)
}));
let req = Request::get("/not-found").body(Full::new(Bytes::new())).unwrap();
let res = service.ready().await.unwrap().call(req).await.unwrap();
assert_eq!(res.status(), StatusCode::NOT_FOUND);
tokio::time::sleep(tokio::time::Duration::from_millis(100)).await;
let req = Request::get("/not-found").body(Full::new(Bytes::new())).unwrap();
let res = service.ready().await.unwrap().call(req).await.unwrap();
assert!(
res.headers().get("x-cache").is_none()
|| res.headers().get("x-cache").unwrap() != "HIT"
);
}
#[tokio::test]
async fn test_default_keyer() {
let keyer = DefaultKeyer;
let req = Request::get("/users/123?page=1").body(()).unwrap();
let key = keyer.cache_key(&req);
assert_eq!(key, "GET /users/123");
}
#[tokio::test]
async fn test_query_keyer() {
let keyer = QueryKeyer;
let req = Request::get("/users/123?page=1").body(()).unwrap();
let key = keyer.cache_key(&req);
assert_eq!(key, "GET /users/123?page=1");
}
#[tokio::test]
async fn test_body_size_limit() {
let manager = MemoryCacheManager::new();
let options = ServerCacheOptions {
max_body_size: 10, ..Default::default()
};
let layer = ServerCacheLayer::new(manager.clone()).with_options(options);
let mut service =
layer.layer(tower::service_fn(|_req: Request<Full<Bytes>>| async {
Ok::<_, std::io::Error>(
Response::builder()
.status(StatusCode::OK)
.header("cache-control", "max-age=60")
.body(Full::new(Bytes::from(
"This is a long response body",
)))
.unwrap(),
)
}));
let req = Request::get("/large").body(Full::new(Bytes::new())).unwrap();
let res = service.ready().await.unwrap().call(req).await.unwrap();
assert_eq!(res.headers().get("x-cache").unwrap(), "MISS");
tokio::time::sleep(tokio::time::Duration::from_millis(100)).await;
let req = Request::get("/large").body(Full::new(Bytes::new())).unwrap();
let res = service.ready().await.unwrap().call(req).await.unwrap();
assert_eq!(res.headers().get("x-cache").unwrap(), "MISS");
}
#[tokio::test]
async fn test_ttl_constraints() {
let manager = MemoryCacheManager::new();
let options = ServerCacheOptions {
min_ttl: Some(std::time::Duration::from_secs(30)),
max_ttl: Some(std::time::Duration::from_secs(90)),
..Default::default()
};
let layer = ServerCacheLayer::new(manager.clone()).with_options(options);
let mut service =
layer.layer(tower::service_fn(|_req: Request<Full<Bytes>>| async {
Ok::<_, std::io::Error>(
Response::builder()
.status(StatusCode::OK)
.header("cache-control", "max-age=10") .body(Full::new(Bytes::from("Response")))
.unwrap(),
)
}));
let req = Request::get("/ttl").body(Full::new(Bytes::new())).unwrap();
let res = service.ready().await.unwrap().call(req).await.unwrap();
assert_eq!(res.headers().get("x-cache").unwrap(), "MISS");
tokio::time::sleep(tokio::time::Duration::from_millis(100)).await;
let req = Request::get("/ttl").body(Full::new(Bytes::new())).unwrap();
let res = service.ready().await.unwrap().call(req).await.unwrap();
assert_eq!(res.headers().get("x-cache").unwrap(), "HIT");
}
#[tokio::test]
async fn test_public_directive() {
let manager = MemoryCacheManager::new();
let layer = ServerCacheLayer::new(manager.clone());
let mut service =
layer.layer(tower::service_fn(|_req: Request<Full<Bytes>>| async {
Ok::<_, std::io::Error>(
Response::builder()
.status(StatusCode::OK)
.header("cache-control", "public")
.body(Full::new(Bytes::from("Public data")))
.unwrap(),
)
}));
let req = Request::get("/public").body(Full::new(Bytes::new())).unwrap();
let res = service.ready().await.unwrap().call(req).await.unwrap();
assert_eq!(res.headers().get("x-cache").unwrap(), "MISS");
tokio::time::sleep(tokio::time::Duration::from_millis(100)).await;
let req = Request::get("/public").body(Full::new(Bytes::new())).unwrap();
let res = service.ready().await.unwrap().call(req).await.unwrap();
assert_eq!(res.headers().get("x-cache").unwrap(), "HIT");
}
#[tokio::test]
async fn test_no_cache_directive() {
let manager = MemoryCacheManager::new();
let layer = ServerCacheLayer::new(manager.clone());
let mut service =
layer.layer(tower::service_fn(|_req: Request<Full<Bytes>>| async {
Ok::<_, std::io::Error>(
Response::builder()
.status(StatusCode::OK)
.header("cache-control", "no-cache")
.body(Full::new(Bytes::from("No cache")))
.unwrap(),
)
}));
let req = Request::get("/no-cache").body(Full::new(Bytes::new())).unwrap();
let _res = service.ready().await.unwrap().call(req).await.unwrap();
tokio::time::sleep(tokio::time::Duration::from_millis(100)).await;
let req = Request::get("/no-cache").body(Full::new(Bytes::new())).unwrap();
let res = service.ready().await.unwrap().call(req).await.unwrap();
assert!(
res.headers().get("x-cache").is_none()
|| res.headers().get("x-cache").unwrap() != "HIT"
);
}
#[tokio::test]
async fn test_expires_future_date() {
let manager = MemoryCacheManager::new();
let layer = ServerCacheLayer::new(manager.clone());
let future_time =
std::time::SystemTime::now() + std::time::Duration::from_secs(60);
let expires_date = httpdate::fmt_http_date(future_time);
let mut service =
layer.layer(tower::service_fn(move |_req: Request<Full<Bytes>>| {
let expires = expires_date.clone();
async move {
Ok::<_, std::io::Error>(
Response::builder()
.status(StatusCode::OK)
.header("expires", expires)
.body(Full::new(Bytes::from("Cacheable with Expires")))
.unwrap(),
)
}
}));
let req =
Request::get("/expires-future").body(Full::new(Bytes::new())).unwrap();
let res = service.ready().await.unwrap().call(req).await.unwrap();
assert_eq!(res.headers().get("x-cache").unwrap(), "MISS");
tokio::time::sleep(tokio::time::Duration::from_millis(100)).await;
let req =
Request::get("/expires-future").body(Full::new(Bytes::new())).unwrap();
let res = service.ready().await.unwrap().call(req).await.unwrap();
assert_eq!(res.headers().get("x-cache").unwrap(), "HIT");
}
#[tokio::test]
async fn test_expires_past_date() {
let manager = MemoryCacheManager::new();
let layer = ServerCacheLayer::new(manager.clone());
let past_time =
std::time::SystemTime::now() - std::time::Duration::from_secs(60);
let expires_date = httpdate::fmt_http_date(past_time);
let mut service =
layer.layer(tower::service_fn(move |_req: Request<Full<Bytes>>| {
let expires = expires_date.clone();
async move {
Ok::<_, std::io::Error>(
Response::builder()
.status(StatusCode::OK)
.header("expires", expires)
.body(Full::new(Bytes::from("Already expired")))
.unwrap(),
)
}
}));
let req =
Request::get("/expires-past").body(Full::new(Bytes::new())).unwrap();
let _res = service.ready().await.unwrap().call(req).await.unwrap();
tokio::time::sleep(tokio::time::Duration::from_millis(100)).await;
let req =
Request::get("/expires-past").body(Full::new(Bytes::new())).unwrap();
let res = service.ready().await.unwrap().call(req).await.unwrap();
assert!(
res.headers().get("x-cache").is_none()
|| res.headers().get("x-cache").unwrap() != "HIT"
);
}
#[tokio::test]
async fn test_expires_invalid_format() {
let manager = MemoryCacheManager::new();
let layer = ServerCacheLayer::new(manager.clone());
let mut service =
layer.layer(tower::service_fn(|_req: Request<Full<Bytes>>| async {
Ok::<_, std::io::Error>(
Response::builder()
.status(StatusCode::OK)
.header("expires", "not-a-valid-date")
.body(Full::new(Bytes::from("Invalid expires")))
.unwrap(),
)
}));
let req =
Request::get("/invalid-expires").body(Full::new(Bytes::new())).unwrap();
let res = service.ready().await.unwrap().call(req).await.unwrap();
assert_eq!(res.status(), StatusCode::OK);
tokio::time::sleep(tokio::time::Duration::from_millis(100)).await;
let req =
Request::get("/invalid-expires").body(Full::new(Bytes::new())).unwrap();
let res = service.ready().await.unwrap().call(req).await.unwrap();
assert!(
res.headers().get("x-cache").is_none()
|| res.headers().get("x-cache").unwrap() != "HIT"
);
}
#[tokio::test]
async fn test_cache_control_overrides_expires() {
let manager = MemoryCacheManager::new();
let layer = ServerCacheLayer::new(manager.clone());
let future_time =
std::time::SystemTime::now() + std::time::Duration::from_secs(10);
let expires_date = httpdate::fmt_http_date(future_time);
let mut service =
layer.layer(tower::service_fn(move |_req: Request<Full<Bytes>>| {
let expires = expires_date.clone();
async move {
Ok::<_, std::io::Error>(
Response::builder()
.status(StatusCode::OK)
.header("cache-control", "max-age=60")
.header("expires", expires)
.body(Full::new(Bytes::from("Both headers")))
.unwrap(),
)
}
}));
let req =
Request::get("/both-headers").body(Full::new(Bytes::new())).unwrap();
let res = service.ready().await.unwrap().call(req).await.unwrap();
assert_eq!(res.headers().get("x-cache").unwrap(), "MISS");
tokio::time::sleep(tokio::time::Duration::from_millis(100)).await;
let req =
Request::get("/both-headers").body(Full::new(Bytes::new())).unwrap();
let res = service.ready().await.unwrap().call(req).await.unwrap();
assert_eq!(res.headers().get("x-cache").unwrap(), "HIT");
}
#[tokio::test]
async fn test_expires_only_no_cache_control() {
let manager = MemoryCacheManager::new();
let layer = ServerCacheLayer::new(manager.clone());
let future_time =
std::time::SystemTime::now() + std::time::Duration::from_secs(60);
let expires_date = httpdate::fmt_http_date(future_time);
let mut service =
layer.layer(tower::service_fn(move |_req: Request<Full<Bytes>>| {
let expires = expires_date.clone();
async move {
Ok::<_, std::io::Error>(
Response::builder()
.status(StatusCode::OK)
.header("expires", expires)
.body(Full::new(Bytes::from("Expires only")))
.unwrap(),
)
}
}));
let req =
Request::get("/expires-only").body(Full::new(Bytes::new())).unwrap();
let res = service.ready().await.unwrap().call(req).await.unwrap();
assert_eq!(res.headers().get("x-cache").unwrap(), "MISS");
tokio::time::sleep(tokio::time::Duration::from_millis(100)).await;
let req =
Request::get("/expires-only").body(Full::new(Bytes::new())).unwrap();
let res = service.ready().await.unwrap().call(req).await.unwrap();
assert_eq!(res.headers().get("x-cache").unwrap(), "HIT");
}
#[tokio::test]
async fn test_expires_with_ttl_constraints() {
let manager = MemoryCacheManager::new();
let options = ServerCacheOptions {
max_ttl: Some(std::time::Duration::from_secs(30)),
..Default::default()
};
let layer = ServerCacheLayer::new(manager.clone()).with_options(options);
let future_time =
std::time::SystemTime::now() + std::time::Duration::from_secs(3600);
let expires_date = httpdate::fmt_http_date(future_time);
let mut service =
layer.layer(tower::service_fn(move |_req: Request<Full<Bytes>>| {
let expires = expires_date.clone();
async move {
Ok::<_, std::io::Error>(
Response::builder()
.status(StatusCode::OK)
.header("expires", expires)
.body(Full::new(Bytes::from(
"Long expires with max_ttl",
)))
.unwrap(),
)
}
}));
let req =
Request::get("/expires-capped").body(Full::new(Bytes::new())).unwrap();
let res = service.ready().await.unwrap().call(req).await.unwrap();
assert_eq!(res.headers().get("x-cache").unwrap(), "MISS");
tokio::time::sleep(tokio::time::Duration::from_millis(100)).await;
let req =
Request::get("/expires-capped").body(Full::new(Bytes::new())).unwrap();
let res = service.ready().await.unwrap().call(req).await.unwrap();
assert_eq!(res.headers().get("x-cache").unwrap(), "HIT");
}
#[tokio::test]
async fn test_concurrent_cache_requests() {
let manager = MemoryCacheManager::new();
let layer = ServerCacheLayer::new(manager.clone());
let request_count = Arc::new(Mutex::new(0));
let request_count_clone = request_count.clone();
let mut service =
layer.layer(tower::service_fn(move |_req: Request<Full<Bytes>>| {
let count = request_count_clone.clone();
async move {
*count.lock().unwrap() += 1;
tokio::time::sleep(tokio::time::Duration::from_millis(50))
.await;
Ok::<_, std::io::Error>(
Response::builder()
.status(StatusCode::OK)
.header("cache-control", "max-age=60")
.body(Full::new(Bytes::from("Concurrent response")))
.unwrap(),
)
}
}));
let mut handles = vec![];
for _ in 0..5 {
let req =
Request::get("/concurrent").body(Full::new(Bytes::new())).unwrap();
let mut svc = service.clone();
let handle = tokio::spawn(async move {
svc.ready().await.unwrap().call(req).await.unwrap()
});
handles.push(handle);
}
let mut responses = vec![];
for handle in handles {
responses.push(handle.await.unwrap());
}
assert_eq!(responses.len(), 5, "All concurrent requests should complete");
let miss_count = responses
.iter()
.filter(|r| {
r.headers().get("x-cache").map(|v| v == "MISS").unwrap_or(false)
})
.count();
assert!(miss_count >= 1, "At least one request should be a cache MISS");
tokio::time::sleep(tokio::time::Duration::from_millis(200)).await;
let req =
Request::get("/concurrent").body(Full::new(Bytes::new())).unwrap();
let res = service.ready().await.unwrap().call(req).await.unwrap();
assert_eq!(
res.headers().get("x-cache").unwrap(),
"HIT",
"Subsequent request should hit cache"
);
}
#[tokio::test]
async fn test_stale_cache_expiration() {
let manager = MemoryCacheManager::new();
let layer = ServerCacheLayer::new(manager.clone());
let mut service =
layer.layer(tower::service_fn(|_req: Request<Full<Bytes>>| async {
Ok::<_, std::io::Error>(
Response::builder()
.status(StatusCode::OK)
.header("cache-control", "max-age=1") .body(Full::new(Bytes::from("Expires soon")))
.unwrap(),
)
}));
let req = Request::get("/stale").body(Full::new(Bytes::new())).unwrap();
let res = service.ready().await.unwrap().call(req).await.unwrap();
assert_eq!(
res.headers().get("x-cache").unwrap(),
"MISS",
"First request should be a cache MISS"
);
tokio::time::sleep(tokio::time::Duration::from_millis(100)).await;
let req = Request::get("/stale").body(Full::new(Bytes::new())).unwrap();
let res = service.ready().await.unwrap().call(req).await.unwrap();
assert_eq!(
res.headers().get("x-cache").unwrap(),
"HIT",
"Request within TTL should be a cache HIT"
);
tokio::time::sleep(tokio::time::Duration::from_millis(1100)).await;
let req = Request::get("/stale").body(Full::new(Bytes::new())).unwrap();
let res = service.ready().await.unwrap().call(req).await.unwrap();
assert_eq!(
res.headers().get("x-cache").unwrap(),
"MISS",
"Request after expiration should be a cache MISS"
);
}
#[tokio::test]
async fn test_multiple_directives() {
let manager = MemoryCacheManager::new();
let layer = ServerCacheLayer::new(manager.clone());
let mut service =
layer.layer(tower::service_fn(|_req: Request<Full<Bytes>>| async {
Ok::<_, std::io::Error>(
Response::builder()
.status(StatusCode::OK)
.header(
"cache-control",
"max-age=60, public, must-revalidate",
)
.body(Full::new(Bytes::from("Multiple directives")))
.unwrap(),
)
}));
let req =
Request::get("/multi-directive").body(Full::new(Bytes::new())).unwrap();
let res = service.ready().await.unwrap().call(req).await.unwrap();
assert_eq!(
res.headers().get("x-cache").unwrap(),
"MISS",
"First request should be a cache MISS"
);
tokio::time::sleep(tokio::time::Duration::from_millis(100)).await;
let req =
Request::get("/multi-directive").body(Full::new(Bytes::new())).unwrap();
let res = service.ready().await.unwrap().call(req).await.unwrap();
assert_eq!(
res.headers().get("x-cache").unwrap(),
"HIT",
"Cache should handle multiple directives correctly"
);
let mut service2 =
layer.layer(tower::service_fn(|_req: Request<Full<Bytes>>| async {
Ok::<_, std::io::Error>(
Response::builder()
.status(StatusCode::OK)
.header("cache-control", "max-age=60,public,s-maxage=120")
.body(Full::new(Bytes::from("No spaces")))
.unwrap(),
)
}));
let req =
Request::get("/multi-no-space").body(Full::new(Bytes::new())).unwrap();
let res = service2.ready().await.unwrap().call(req).await.unwrap();
assert_eq!(
res.headers().get("x-cache").unwrap(),
"MISS",
"Should handle directives without spaces"
);
tokio::time::sleep(tokio::time::Duration::from_millis(100)).await;
let req =
Request::get("/multi-no-space").body(Full::new(Bytes::new())).unwrap();
let res = service2.ready().await.unwrap().call(req).await.unwrap();
assert_eq!(
res.headers().get("x-cache").unwrap(),
"HIT",
"Should cache with directives without spaces"
);
}
#[tokio::test]
async fn test_malformed_cache_control() {
let manager = MemoryCacheManager::new();
let layer = ServerCacheLayer::new(manager.clone());
let mut service1 = layer.clone().layer(tower::service_fn(
|_req: Request<Full<Bytes>>| async {
Ok::<_, std::io::Error>(
Response::builder()
.status(StatusCode::OK)
.header("cache-control", "invalid-directive")
.body(Full::new(Bytes::from("Invalid directive")))
.unwrap(),
)
},
));
let req = Request::get("/invalid-directive")
.body(Full::new(Bytes::new()))
.unwrap();
let res = service1.ready().await.unwrap().call(req).await.unwrap();
assert_eq!(
res.status(),
StatusCode::OK,
"Should handle invalid directive gracefully"
);
tokio::time::sleep(tokio::time::Duration::from_millis(100)).await;
let req = Request::get("/invalid-directive")
.body(Full::new(Bytes::new()))
.unwrap();
let res = service1.ready().await.unwrap().call(req).await.unwrap();
assert!(
res.headers().get("x-cache").is_none()
|| res.headers().get("x-cache").unwrap() != "HIT",
"Should not cache with invalid directive"
);
let mut service2 = layer.clone().layer(tower::service_fn(
|_req: Request<Full<Bytes>>| async {
Ok::<_, std::io::Error>(
Response::builder()
.status(StatusCode::OK)
.header("cache-control", "max-age=notanumber")
.body(Full::new(Bytes::from("Invalid max-age")))
.unwrap(),
)
},
));
let req =
Request::get("/bad-max-age").body(Full::new(Bytes::new())).unwrap();
let res = service2.ready().await.unwrap().call(req).await.unwrap();
assert_eq!(
res.status(),
StatusCode::OK,
"Should handle malformed max-age gracefully"
);
tokio::time::sleep(tokio::time::Duration::from_millis(100)).await;
let req =
Request::get("/bad-max-age").body(Full::new(Bytes::new())).unwrap();
let res = service2.ready().await.unwrap().call(req).await.unwrap();
assert!(
res.headers().get("x-cache").is_none()
|| res.headers().get("x-cache").unwrap() != "HIT",
"Should not cache with malformed max-age"
);
let mut service3 =
layer.layer(tower::service_fn(|_req: Request<Full<Bytes>>| async {
Ok::<_, std::io::Error>(
Response::builder()
.status(StatusCode::OK)
.header("cache-control", "")
.body(Full::new(Bytes::from("Empty cache-control")))
.unwrap(),
)
}));
let req = Request::get("/empty-cc").body(Full::new(Bytes::new())).unwrap();
let res = service3.ready().await.unwrap().call(req).await.unwrap();
assert_eq!(
res.status(),
StatusCode::OK,
"Should handle empty cache-control gracefully"
);
}
#[tokio::test]
async fn test_path_parameter_preservation() {
let manager = MemoryCacheManager::new();
let layer = ServerCacheLayer::new(manager.clone());
let call_count = Arc::new(Mutex::new(0));
let call_count_clone = call_count.clone();
let mut service =
layer.layer(tower::service_fn(move |req: Request<Full<Bytes>>| {
let count = call_count_clone.clone();
async move {
*count.lock().unwrap() += 1;
let path_params = req
.extensions()
.get::<PathParams>()
.expect("PathParams extension should be present");
let body = format!("User ID: {}", path_params.id);
Ok::<_, std::io::Error>(
Response::builder()
.status(StatusCode::OK)
.header("cache-control", "max-age=60")
.body(Full::new(Bytes::from(body)))
.unwrap(),
)
}
}));
let mut req1 =
Request::get("/users/123").body(Full::new(Bytes::new())).unwrap();
req1.extensions_mut().insert(PathParams { id: "123".to_string() });
let res1 = service.ready().await.unwrap().call(req1).await.unwrap();
assert_eq!(res1.status(), StatusCode::OK);
assert_eq!(res1.headers().get("x-cache").unwrap(), "MISS");
let body1 = http_body_util::BodyExt::collect(res1.into_body())
.await
.unwrap()
.to_bytes();
assert_eq!(
body1, "User ID: 123",
"Handler should receive path parameter on cache miss"
);
assert_eq!(
*call_count.lock().unwrap(),
1,
"Handler should be called on cache miss"
);
tokio::time::sleep(tokio::time::Duration::from_millis(100)).await;
let mut req2 =
Request::get("/users/123").body(Full::new(Bytes::new())).unwrap();
req2.extensions_mut().insert(PathParams { id: "123".to_string() });
let res2 = service.ready().await.unwrap().call(req2).await.unwrap();
assert_eq!(res2.status(), StatusCode::OK);
assert_eq!(res2.headers().get("x-cache").unwrap(), "HIT");
let body2 = http_body_util::BodyExt::collect(res2.into_body())
.await
.unwrap()
.to_bytes();
assert_eq!(
body2, "User ID: 123",
"Cached response should have correct content"
);
assert_eq!(
*call_count.lock().unwrap(),
1,
"Handler should not be called on cache hit"
);
let mut req3 =
Request::get("/users/456").body(Full::new(Bytes::new())).unwrap();
req3.extensions_mut().insert(PathParams { id: "456".to_string() });
let res3 = service.ready().await.unwrap().call(req3).await.unwrap();
assert_eq!(res3.status(), StatusCode::OK);
assert_eq!(res3.headers().get("x-cache").unwrap(), "MISS");
let body3 = http_body_util::BodyExt::collect(res3.into_body())
.await
.unwrap()
.to_bytes();
assert_eq!(
body3, "User ID: 456",
"Handler should receive different path parameter for different request"
);
assert_eq!(
*call_count.lock().unwrap(),
2,
"Handler should be called for new path"
);
}
#[tokio::test]
async fn test_request_extensions_not_stripped() {
let manager = MemoryCacheManager::new();
let layer = ServerCacheLayer::new(manager.clone());
#[derive(Debug, Clone, PartialEq)]
struct CustomExtension {
value: String,
}
let mut service = layer.layer(tower::service_fn(
|req: Request<Full<Bytes>>| async move {
let ext = req.extensions().get::<CustomExtension>();
assert!(
ext.is_some(),
"Extension should be preserved through cache layer"
);
assert_eq!(ext.unwrap().value, "test-value");
Ok::<_, std::io::Error>(
Response::builder()
.status(StatusCode::OK)
.header("cache-control", "max-age=60")
.body(Full::new(Bytes::from("OK")))
.unwrap(),
)
},
));
let mut req = Request::get("/test").body(Full::new(Bytes::new())).unwrap();
req.extensions_mut()
.insert(CustomExtension { value: "test-value".to_string() });
let res = service.ready().await.unwrap().call(req).await.unwrap();
assert_eq!(res.status(), StatusCode::OK);
assert_eq!(res.headers().get("x-cache").unwrap(), "MISS");
}
#[tokio::test]
async fn test_cache_by_default_option() {
let manager = MemoryCacheManager::new();
let options_disabled =
ServerCacheOptions { cache_by_default: false, ..Default::default() };
let layer =
ServerCacheLayer::new(manager.clone()).with_options(options_disabled);
let mut service =
layer.layer(tower::service_fn(|_req: Request<Full<Bytes>>| async {
Ok::<_, std::io::Error>(
Response::builder()
.status(StatusCode::OK)
.body(Full::new(Bytes::from("No directives")))
.unwrap(),
)
}));
let req =
Request::get("/no-directive").body(Full::new(Bytes::new())).unwrap();
let _res = service.ready().await.unwrap().call(req).await.unwrap();
tokio::time::sleep(tokio::time::Duration::from_millis(100)).await;
let req =
Request::get("/no-directive").body(Full::new(Bytes::new())).unwrap();
let res = service.ready().await.unwrap().call(req).await.unwrap();
assert!(
res.headers().get("x-cache").is_none()
|| res.headers().get("x-cache").unwrap() != "HIT",
"Should not cache without directives when cache_by_default is false"
);
let options_enabled =
ServerCacheOptions { cache_by_default: true, ..Default::default() };
let layer =
ServerCacheLayer::new(manager.clone()).with_options(options_enabled);
let mut service =
layer.layer(tower::service_fn(|_req: Request<Full<Bytes>>| async {
Ok::<_, std::io::Error>(
Response::builder()
.status(StatusCode::OK)
.body(Full::new(Bytes::from("No directives but cached")))
.unwrap(),
)
}));
let req = Request::get("/cache-by-default")
.body(Full::new(Bytes::new()))
.unwrap();
let res = service.ready().await.unwrap().call(req).await.unwrap();
assert_eq!(res.headers().get("x-cache").unwrap(), "MISS");
tokio::time::sleep(tokio::time::Duration::from_millis(100)).await;
let req = Request::get("/cache-by-default")
.body(Full::new(Bytes::new()))
.unwrap();
let res = service.ready().await.unwrap().call(req).await.unwrap();
assert_eq!(
res.headers().get("x-cache").unwrap(),
"HIT",
"Should cache without directives when cache_by_default is true"
);
}
#[tokio::test]
async fn test_custom_keyer() {
let manager = MemoryCacheManager::new();
let keyer = CustomKeyer::new(|req: &Request<()>| {
let lang = req
.headers()
.get("accept-language")
.and_then(|v| v.to_str().ok())
.unwrap_or("en");
format!("{} {} lang:{}", req.method(), req.uri().path(), lang)
});
let layer = ServerCacheLayer::with_keyer(manager.clone(), keyer);
let mut service = layer.layer(tower::service_fn(
|req: Request<Full<Bytes>>| async move {
let lang = req
.headers()
.get("accept-language")
.and_then(|v| v.to_str().ok())
.unwrap_or("en");
let body = format!("Response for {}", lang);
Ok::<_, std::io::Error>(
Response::builder()
.status(StatusCode::OK)
.header("cache-control", "max-age=60")
.body(Full::new(Bytes::from(body)))
.unwrap(),
)
},
));
let req = Request::get("/test")
.header("accept-language", "en")
.body(Full::new(Bytes::new()))
.unwrap();
let res = service.ready().await.unwrap().call(req).await.unwrap();
assert_eq!(res.headers().get("x-cache").unwrap(), "MISS");
tokio::time::sleep(tokio::time::Duration::from_millis(100)).await;
let req = Request::get("/test")
.header("accept-language", "en")
.body(Full::new(Bytes::new()))
.unwrap();
let res = service.ready().await.unwrap().call(req).await.unwrap();
assert_eq!(res.headers().get("x-cache").unwrap(), "HIT");
let req = Request::get("/test")
.header("accept-language", "fr")
.body(Full::new(Bytes::new()))
.unwrap();
let res = service.ready().await.unwrap().call(req).await.unwrap();
assert_eq!(
res.headers().get("x-cache").unwrap(),
"MISS",
"Different language should have different cache key"
);
}
#[tokio::test]
async fn test_directive_parsing_edge_cases() {
let manager = MemoryCacheManager::new();
let layer = ServerCacheLayer::new(manager.clone());
let mut service1 = layer.clone().layer(tower::service_fn(
|_req: Request<Full<Bytes>>| async {
Ok::<_, std::io::Error>(
Response::builder()
.status(StatusCode::OK)
.header("cache-control", "max-age=60, no-store-custom")
.body(Full::new(Bytes::from("Should be cached")))
.unwrap(),
)
},
));
let req =
Request::get("/no-store-custom").body(Full::new(Bytes::new())).unwrap();
let res = service1.ready().await.unwrap().call(req).await.unwrap();
assert_eq!(res.headers().get("x-cache").unwrap(), "MISS");
tokio::time::sleep(tokio::time::Duration::from_millis(100)).await;
let req =
Request::get("/no-store-custom").body(Full::new(Bytes::new())).unwrap();
let res = service1.ready().await.unwrap().call(req).await.unwrap();
assert_eq!(
res.headers().get("x-cache").unwrap(),
"HIT",
"no-store-custom should not prevent caching"
);
let mut service2 = layer.clone().layer(tower::service_fn(
|_req: Request<Full<Bytes>>| async {
Ok::<_, std::io::Error>(
Response::builder()
.status(StatusCode::OK)
.header("cache-control", "max-age=60, private-ext")
.body(Full::new(Bytes::from("Should be cached")))
.unwrap(),
)
},
));
let req =
Request::get("/private-ext").body(Full::new(Bytes::new())).unwrap();
let res = service2.ready().await.unwrap().call(req).await.unwrap();
assert_eq!(res.headers().get("x-cache").unwrap(), "MISS");
tokio::time::sleep(tokio::time::Duration::from_millis(100)).await;
let req =
Request::get("/private-ext").body(Full::new(Bytes::new())).unwrap();
let res = service2.ready().await.unwrap().call(req).await.unwrap();
assert_eq!(
res.headers().get("x-cache").unwrap(),
"HIT",
"private-ext should not prevent caching"
);
}
#[tokio::test]
async fn test_zero_max_age() {
let manager = MemoryCacheManager::new();
let layer = ServerCacheLayer::new(manager.clone());
let mut service =
layer.layer(tower::service_fn(|_req: Request<Full<Bytes>>| async {
Ok::<_, std::io::Error>(
Response::builder()
.status(StatusCode::OK)
.header("cache-control", "max-age=0")
.body(Full::new(Bytes::from("Zero TTL")))
.unwrap(),
)
}));
let req = Request::get("/zero-ttl").body(Full::new(Bytes::new())).unwrap();
let res = service.ready().await.unwrap().call(req).await.unwrap();
assert_eq!(res.headers().get("x-cache").unwrap(), "MISS");
tokio::time::sleep(tokio::time::Duration::from_millis(100)).await;
let req = Request::get("/zero-ttl").body(Full::new(Bytes::new())).unwrap();
let res = service.ready().await.unwrap().call(req).await.unwrap();
assert_eq!(
res.headers().get("x-cache").unwrap(),
"MISS",
"Zero max-age should result in immediately stale cache entry"
);
}
#[tokio::test]
async fn test_different_http_methods() {
let manager = MemoryCacheManager::new();
let layer = ServerCacheLayer::new(manager.clone());
let mut service = layer.layer(tower::service_fn(
|req: Request<Full<Bytes>>| async move {
let body = format!("Method: {}", req.method());
Ok::<_, std::io::Error>(
Response::builder()
.status(StatusCode::OK)
.header("cache-control", "max-age=60")
.body(Full::new(Bytes::from(body)))
.unwrap(),
)
},
));
let req =
Request::get("/method-test").body(Full::new(Bytes::new())).unwrap();
let res = service.ready().await.unwrap().call(req).await.unwrap();
assert_eq!(res.headers().get("x-cache").unwrap(), "MISS");
tokio::time::sleep(tokio::time::Duration::from_millis(100)).await;
let req =
Request::get("/method-test").body(Full::new(Bytes::new())).unwrap();
let res = service.ready().await.unwrap().call(req).await.unwrap();
assert_eq!(res.headers().get("x-cache").unwrap(), "HIT");
let req =
Request::post("/method-test").body(Full::new(Bytes::new())).unwrap();
let res = service.ready().await.unwrap().call(req).await.unwrap();
assert_eq!(
res.headers().get("x-cache").unwrap(),
"MISS",
"POST should have different cache key than GET"
);
}
#[tokio::test]
async fn test_cache_status_headers_disabled() {
let manager = MemoryCacheManager::new();
let options = ServerCacheOptions {
cache_status_headers: false,
..Default::default()
};
let layer = ServerCacheLayer::new(manager.clone()).with_options(options);
let mut service =
layer.layer(tower::service_fn(|_req: Request<Full<Bytes>>| async {
Ok::<_, std::io::Error>(
Response::builder()
.status(StatusCode::OK)
.header("cache-control", "max-age=60")
.body(Full::new(Bytes::from("No status headers")))
.unwrap(),
)
}));
let req = Request::get("/no-status").body(Full::new(Bytes::new())).unwrap();
let res = service.ready().await.unwrap().call(req).await.unwrap();
assert!(
res.headers().get("x-cache").is_none(),
"Should not have x-cache header when disabled"
);
tokio::time::sleep(tokio::time::Duration::from_millis(100)).await;
let req = Request::get("/no-status").body(Full::new(Bytes::new())).unwrap();
let res = service.ready().await.unwrap().call(req).await.unwrap();
assert!(
res.headers().get("x-cache").is_none(),
"Should not have x-cache header when disabled, even on HIT"
);
}