Skip to main content

reinhardt_middleware/
cache.rs

1//! Cache Middleware
2//!
3//! Provides caching for HTTP responses.
4//! Supports various cache backends (memory, Redis, file).
5//! Requests carrying credentials or authenticated state bypass the cache. Responses marked
6//! private, non-reusable, cookie-setting, or unsupported variant-dependent are never shared.
7
8use async_trait::async_trait;
9use hyper::StatusCode;
10use hyper::header::{AUTHORIZATION, CACHE_CONTROL, COOKIE, SET_COOKIE, VARY};
11use reinhardt_http::{AuthState, Handler, IsAuthenticated, Middleware, Request, Response, Result};
12use serde::{Deserialize, Serialize};
13use sha2::{Digest, Sha256};
14use std::collections::HashMap;
15use std::sync::{Arc, RwLock};
16use std::time::{Duration, Instant};
17
18/// Cache Entry
19#[derive(Debug, Clone, Serialize, Deserialize)]
20pub struct CacheEntry {
21	/// Status code
22	status: u16,
23	/// Headers
24	headers: HashMap<String, String>,
25	/// Body
26	body: Vec<u8>,
27	/// Cached timestamp
28	#[serde(skip)]
29	cached_at: Option<Instant>,
30	/// TTL (seconds)
31	ttl_secs: u64,
32}
33
34impl CacheEntry {
35	/// Create a new entry
36	fn new(response: &Response, ttl: Duration) -> Self {
37		let mut headers = HashMap::new();
38		for (key, value) in response.headers.iter() {
39			if let Ok(value_str) = value.to_str() {
40				headers.insert(key.to_string(), value_str.to_string());
41			}
42		}
43
44		Self {
45			status: response.status.as_u16(),
46			headers,
47			body: response.body.to_vec(),
48			cached_at: Some(Instant::now()),
49			ttl_secs: ttl.as_secs(),
50		}
51	}
52
53	/// Check if expired
54	fn is_expired(&self) -> bool {
55		if let Some(cached_at) = self.cached_at {
56			cached_at.elapsed().as_secs() >= self.ttl_secs
57		} else {
58			true
59		}
60	}
61
62	/// Check if an entry is safe to serve from a shared cache.
63	fn is_shareable(&self, key_strategy: CacheKeyStrategy) -> bool {
64		!self.headers.contains_key(SET_COOKIE.as_str())
65			&& !self
66				.headers
67				.get(CACHE_CONTROL.as_str())
68				.is_some_and(|value| cache_control_forbids_shared_storage(value))
69			&& self.headers.get(VARY.as_str()).is_none_or(|value| {
70				matches!(key_strategy, CacheKeyStrategy::UrlAndHeaders)
71					&& !value.split(',').any(|field| field.trim() == "*")
72			})
73	}
74
75	/// Convert to response
76	fn to_response(&self) -> Response {
77		let status = StatusCode::from_u16(self.status).unwrap_or(StatusCode::OK);
78		let mut response = Response::new(status).with_body(self.body.clone());
79
80		for (key, value) in &self.headers {
81			if let (Ok(header_name), Ok(header_value)) =
82				(hyper::header::HeaderName::try_from(key), value.parse())
83			{
84				response.headers.insert(header_name, header_value);
85			}
86		}
87
88		// Add cache header
89		response.headers.insert(
90			hyper::header::HeaderName::from_static("x-cache"),
91			hyper::header::HeaderValue::from_static("HIT"),
92		);
93
94		response
95	}
96}
97
98fn cache_control_forbids_shared_storage(value: &str) -> bool {
99	value.split(',').any(|directive| {
100		let directive = directive.trim();
101		let name = directive
102			.split_once('=')
103			.map_or(directive, |(name, _)| name)
104			.trim();
105		name.eq_ignore_ascii_case("private")
106			|| name.eq_ignore_ascii_case("no-store")
107			|| name.eq_ignore_ascii_case("no-cache")
108	})
109}
110
111/// Cache Storage
112#[derive(Debug, Default)]
113pub struct CacheStore {
114	/// Entries
115	entries: RwLock<HashMap<String, CacheEntry>>,
116}
117
118impl CacheStore {
119	/// Create a new store
120	pub fn new() -> Self {
121		Self::default()
122	}
123
124	/// Get an entry
125	pub fn get(&self, key: &str) -> Option<CacheEntry> {
126		let entries = self.entries.read().unwrap_or_else(|e| e.into_inner());
127		entries.get(key).cloned()
128	}
129
130	/// Set an entry
131	pub fn set(&self, key: String, entry: CacheEntry) {
132		let mut entries = self.entries.write().unwrap_or_else(|e| e.into_inner());
133		entries.insert(key, entry);
134	}
135
136	/// Delete an entry
137	pub fn delete(&self, key: &str) {
138		let mut entries = self.entries.write().unwrap_or_else(|e| e.into_inner());
139		entries.remove(key);
140	}
141
142	/// Clean up expired entries
143	pub fn cleanup(&self) {
144		let mut entries = self.entries.write().unwrap_or_else(|e| e.into_inner());
145		entries.retain(|_, entry| !entry.is_expired());
146	}
147
148	/// Clear the store
149	pub fn clear(&self) {
150		let mut entries = self.entries.write().unwrap_or_else(|e| e.into_inner());
151		entries.clear();
152	}
153
154	/// Get the number of entries
155	pub fn len(&self) -> usize {
156		let entries = self.entries.read().unwrap_or_else(|e| e.into_inner());
157		entries.len()
158	}
159
160	/// Check if the store is empty
161	pub fn is_empty(&self) -> bool {
162		let entries = self.entries.read().unwrap_or_else(|e| e.into_inner());
163		entries.is_empty()
164	}
165}
166
167/// Cache key generation strategy
168#[derive(Debug, Clone, Copy)]
169pub enum CacheKeyStrategy {
170	/// URL only
171	UrlOnly,
172	/// URL and method
173	UrlAndMethod,
174	/// URL and query parameters
175	UrlAndQuery,
176	/// URL and headers
177	UrlAndHeaders,
178}
179
180/// Cache configuration
181#[non_exhaustive]
182#[derive(Debug, Clone)]
183pub struct CacheConfig {
184	/// Default TTL
185	pub default_ttl: Duration,
186	/// Cache key generation strategy
187	pub key_strategy: CacheKeyStrategy,
188	/// Cacheable methods
189	pub cacheable_methods: Vec<String>,
190	/// Cacheable status codes
191	pub cacheable_status_codes: Vec<u16>,
192	/// Paths to exclude
193	pub exclude_paths: Vec<String>,
194	/// Maximum cache size
195	pub max_entries: Option<usize>,
196}
197
198impl CacheConfig {
199	/// Create a new configuration
200	///
201	/// # Examples
202	///
203	/// ```
204	/// use std::time::Duration;
205	/// use reinhardt_middleware::cache::{CacheConfig, CacheKeyStrategy};
206	///
207	/// let config = CacheConfig::new(Duration::from_secs(300), CacheKeyStrategy::UrlOnly);
208	/// assert_eq!(config.default_ttl, Duration::from_secs(300));
209	/// ```
210	pub fn new(default_ttl: Duration, key_strategy: CacheKeyStrategy) -> Self {
211		Self {
212			default_ttl,
213			key_strategy,
214			cacheable_methods: vec!["GET".to_string(), "HEAD".to_string()],
215			cacheable_status_codes: vec![200, 203, 204, 206, 300, 301, 404, 405, 410, 414, 501],
216			exclude_paths: Vec::new(),
217			max_entries: Some(1000),
218		}
219	}
220
221	/// Set cacheable methods
222	///
223	/// # Examples
224	///
225	/// ```
226	/// use std::time::Duration;
227	/// use reinhardt_middleware::cache::{CacheConfig, CacheKeyStrategy};
228	///
229	/// let config = CacheConfig::new(Duration::from_secs(300), CacheKeyStrategy::UrlOnly)
230	///     .with_cacheable_methods(vec!["GET".to_string()]);
231	/// ```
232	pub fn with_cacheable_methods(mut self, methods: Vec<String>) -> Self {
233		self.cacheable_methods = methods;
234		self
235	}
236
237	/// Add paths to exclude
238	///
239	/// # Examples
240	///
241	/// ```
242	/// use std::time::Duration;
243	/// use reinhardt_middleware::cache::{CacheConfig, CacheKeyStrategy};
244	///
245	/// let config = CacheConfig::new(Duration::from_secs(300), CacheKeyStrategy::UrlOnly)
246	///     .with_excluded_paths(vec!["/admin".to_string()]);
247	/// ```
248	pub fn with_excluded_paths(mut self, paths: Vec<String>) -> Self {
249		self.exclude_paths.extend(paths);
250		self
251	}
252
253	/// Set maximum number of entries
254	///
255	/// # Examples
256	///
257	/// ```
258	/// use std::time::Duration;
259	/// use reinhardt_middleware::cache::{CacheConfig, CacheKeyStrategy};
260	///
261	/// let config = CacheConfig::new(Duration::from_secs(300), CacheKeyStrategy::UrlOnly)
262	///     .with_max_entries(5000);
263	/// ```
264	pub fn with_max_entries(mut self, max_entries: usize) -> Self {
265		self.max_entries = Some(max_entries);
266		self
267	}
268}
269
270impl Default for CacheConfig {
271	fn default() -> Self {
272		Self::new(Duration::from_secs(300), CacheKeyStrategy::UrlOnly)
273	}
274}
275
276/// Cache Middleware
277///
278/// # Examples
279///
280/// ```
281/// use std::sync::Arc;
282/// use std::time::Duration;
283/// use reinhardt_middleware::cache::{CacheMiddleware, CacheConfig, CacheKeyStrategy};
284/// use reinhardt_http::{Handler, Middleware, Request, Response};
285/// use hyper::{StatusCode, Method, Version, HeaderMap};
286/// use bytes::Bytes;
287///
288/// struct TestHandler;
289///
290/// #[async_trait::async_trait]
291/// impl Handler for TestHandler {
292///     async fn handle(&self, _request: Request) -> reinhardt_core::exception::Result<Response> {
293///         Ok(Response::new(StatusCode::OK).with_body(Bytes::from("OK")))
294///     }
295/// }
296///
297/// # tokio_test::block_on(async {
298/// let config = CacheConfig::new(Duration::from_secs(60), CacheKeyStrategy::UrlOnly);
299/// let middleware = CacheMiddleware::new(config);
300/// let handler = Arc::new(TestHandler);
301///
302/// let request = Request::builder()
303///     .method(Method::GET)
304///     .uri("/api/data")
305///     .version(Version::HTTP_11)
306///     .headers(HeaderMap::new())
307///     .body(Bytes::new())
308///     .build()
309///     .unwrap();
310///
311/// let response = middleware.process(request, handler).await.unwrap();
312/// assert_eq!(response.status, StatusCode::OK);
313/// # });
314/// ```
315pub struct CacheMiddleware {
316	config: CacheConfig,
317	store: Arc<CacheStore>,
318}
319
320impl CacheMiddleware {
321	/// Create a new cache middleware
322	///
323	/// # Examples
324	///
325	/// ```
326	/// use std::time::Duration;
327	/// use reinhardt_middleware::cache::{CacheMiddleware, CacheConfig, CacheKeyStrategy};
328	///
329	/// let config = CacheConfig::new(Duration::from_secs(300), CacheKeyStrategy::UrlOnly);
330	/// let middleware = CacheMiddleware::new(config);
331	/// ```
332	pub fn new(config: CacheConfig) -> Self {
333		Self {
334			config,
335			store: Arc::new(CacheStore::new()),
336		}
337	}
338
339	/// Create with default configuration
340	pub fn with_defaults() -> Self {
341		Self::new(CacheConfig::default())
342	}
343
344	/// Create from an existing Arc-wrapped cache store
345	///
346	/// This is provided for cases where you already have an `Arc<CacheStore>`.
347	/// In most cases, you should use `new()` instead, which creates the store internally.
348	pub fn from_arc(config: CacheConfig, store: Arc<CacheStore>) -> Self {
349		Self { config, store }
350	}
351
352	/// Get a reference to the cache store
353	///
354	/// # Examples
355	///
356	/// ```
357	/// use std::time::Duration;
358	/// use reinhardt_middleware::cache::{CacheMiddleware, CacheConfig, CacheKeyStrategy};
359	///
360	/// let middleware = CacheMiddleware::new(
361	///     CacheConfig::new(Duration::from_secs(300), CacheKeyStrategy::UrlOnly)
362	/// );
363	///
364	/// // Access the store
365	/// let store = middleware.store();
366	/// assert_eq!(store.len(), 0);
367	/// ```
368	pub fn store(&self) -> &CacheStore {
369		&self.store
370	}
371
372	/// Get a cloned Arc of the store (for cases where you need ownership)
373	///
374	/// In most cases, you should use `store()` instead to get a reference.
375	pub fn store_arc(&self) -> Arc<CacheStore> {
376		Arc::clone(&self.store)
377	}
378
379	/// Check if path should be excluded
380	fn should_exclude(&self, path: &str) -> bool {
381		self.config
382			.exclude_paths
383			.iter()
384			.any(|p| path.starts_with(p))
385	}
386
387	/// Check if method is cacheable
388	fn is_cacheable_method(&self, method: &str) -> bool {
389		self.config.cacheable_methods.iter().any(|m| m == method)
390	}
391
392	/// Check if status code is cacheable
393	fn is_cacheable_status(&self, status: u16) -> bool {
394		self.config.cacheable_status_codes.contains(&status)
395	}
396
397	/// Check if a request may carry user-specific state.
398	fn is_private_request(request: &Request) -> bool {
399		request.headers.contains_key(AUTHORIZATION)
400			|| request.headers.contains_key(COOKIE)
401			|| request.headers.contains_key("remote_user")
402			|| AuthState::from_extensions(&request.extensions)
403				.is_some_and(|state| state.is_authenticated())
404			|| request
405				.extensions
406				.get::<IsAuthenticated>()
407				.is_some_and(|state| state.0)
408	}
409
410	/// Check if a response is safe to store in a shared cache.
411	fn is_shareable_response(&self, response: &Response) -> bool {
412		!response.headers.contains_key(SET_COOKIE)
413			&& response.headers.get_all(CACHE_CONTROL).iter().all(|value| {
414				value
415					.to_str()
416					.is_ok_and(|value| !cache_control_forbids_shared_storage(value))
417			}) && response.headers.get_all(VARY).iter().all(|value| {
418			matches!(self.config.key_strategy, CacheKeyStrategy::UrlAndHeaders)
419				&& value
420					.to_str()
421					.is_ok_and(|value| !value.split(',').any(|field| field.trim() == "*"))
422		})
423	}
424
425	/// Generate cache key
426	fn generate_cache_key(&self, request: &Request) -> String {
427		let base = match self.config.key_strategy {
428			CacheKeyStrategy::UrlOnly => request.uri.path().to_string(),
429			CacheKeyStrategy::UrlAndMethod => {
430				format!("{}:{}", request.method.as_str(), request.uri.path())
431			}
432			CacheKeyStrategy::UrlAndQuery => {
433				let query = request.uri.query().unwrap_or("");
434				format!(
435					"{}:{}?{}",
436					request.method.as_str(),
437					request.uri.path(),
438					query
439				)
440			}
441			CacheKeyStrategy::UrlAndHeaders => {
442				let headers_str = request
443					.headers
444					.iter()
445					.map(|(k, v)| format!("{}={}", k, v.to_str().unwrap_or("")))
446					.collect::<Vec<_>>()
447					.join("&");
448				format!(
449					"{}:{}:{}",
450					request.method.as_str(),
451					request.uri.path(),
452					headers_str
453				)
454			}
455		};
456
457		// Hash with SHA256
458		let mut hasher = Sha256::new();
459		hasher.update(base.as_bytes());
460		let result = hasher.finalize();
461		hex::encode(result)
462	}
463}
464
465impl Default for CacheMiddleware {
466	fn default() -> Self {
467		Self::with_defaults()
468	}
469}
470
471#[async_trait]
472impl Middleware for CacheMiddleware {
473	async fn process(&self, request: Request, handler: Arc<dyn Handler>) -> Result<Response> {
474		let path = request.uri.path().to_string();
475		let method = request.method.as_str().to_string();
476
477		// Skip excluded paths
478		if self.should_exclude(&path) {
479			return handler.handle(request).await;
480		}
481
482		// Skip non-cacheable methods
483		if !self.is_cacheable_method(&method) {
484			return handler.handle(request).await;
485		}
486
487		// Credential-bearing and authenticated requests must never use a shared cache entry.
488		let cache_key =
489			(!Self::is_private_request(&request)).then(|| self.generate_cache_key(&request));
490
491		// Check cache
492		if let Some((cache_key, entry)) = cache_key
493			.as_deref()
494			.and_then(|key| self.store.get(key).map(|entry| (key, entry)))
495		{
496			if !entry.is_shareable(self.config.key_strategy) {
497				self.store.delete(cache_key);
498			} else if !entry.is_expired() {
499				// Cache hit
500				return Ok(entry.to_response());
501			} else {
502				// Delete expired entry
503				self.store.delete(cache_key);
504			}
505		}
506
507		// Convert errors to responses so post-processing always runs,
508		// even when invoked outside MiddlewareChain. (#3244)
509		let response = match handler.handle(request).await {
510			Ok(resp) => resp,
511			Err(e) => Response::from(e),
512		};
513
514		// Save to cache if status code is cacheable
515		if let Some(cache_key) = cache_key
516			&& self.is_cacheable_status(response.status.as_u16())
517			&& self.is_shareable_response(&response)
518		{
519			let entry = CacheEntry::new(&response, self.config.default_ttl);
520			self.store.set(cache_key, entry);
521
522			// Clean up expired entries if max entries exceeded
523			if let Some(max_entries) = self.config.max_entries
524				&& self.store.len() > max_entries
525			{
526				self.store.cleanup();
527			}
528		}
529
530		// Add X-Cache header
531		let mut response = response;
532		response.headers.insert(
533			hyper::header::HeaderName::from_static("x-cache"),
534			hyper::header::HeaderValue::from_static("MISS"),
535		);
536
537		Ok(response)
538	}
539}
540
541#[cfg(test)]
542mod tests {
543	use super::*;
544	use bytes::Bytes;
545	use hyper::{HeaderMap, Method, StatusCode, Version};
546
547	struct TestHandler {
548		status: StatusCode,
549		call_count: Arc<RwLock<usize>>,
550	}
551
552	impl TestHandler {
553		fn new(status: StatusCode) -> Self {
554			Self {
555				status,
556				call_count: Arc::new(RwLock::new(0)),
557			}
558		}
559
560		fn get_call_count(&self) -> usize {
561			*self.call_count.read().unwrap()
562		}
563	}
564
565	#[async_trait]
566	impl Handler for TestHandler {
567		async fn handle(&self, _request: Request) -> Result<Response> {
568			*self.call_count.write().unwrap() += 1;
569			Ok(Response::new(self.status).with_body(Bytes::from("OK")))
570		}
571	}
572
573	struct IdentityHandler;
574
575	#[async_trait]
576	impl Handler for IdentityHandler {
577		async fn handle(&self, request: Request) -> Result<Response> {
578			let identity = AuthState::from_extensions(&request.extensions)
579				.filter(|state| state.is_authenticated())
580				.map(|state| state.user_id().to_string())
581				.or_else(|| {
582					request
583						.headers
584						.get(AUTHORIZATION)
585						.and_then(|value| value.to_str().ok())
586						.map(str::to_string)
587				})
588				.or_else(|| {
589					request
590						.headers
591						.get(COOKIE)
592						.and_then(|value| value.to_str().ok())
593						.map(str::to_string)
594				})
595				.or_else(|| {
596					request
597						.headers
598						.get("remote_user")
599						.and_then(|value| value.to_str().ok())
600						.map(str::to_string)
601				})
602				.unwrap_or_else(|| "public".to_string());
603			Ok(Response::new(StatusCode::OK).with_body(identity))
604		}
605	}
606
607	#[tokio::test]
608	async fn authenticated_responses_are_not_shared_by_url_only_cache() {
609		let middleware = CacheMiddleware::with_defaults();
610		let handler = Arc::new(IdentityHandler);
611
612		for identity in ["Bearer alice", "Bearer bob"] {
613			let mut headers = HeaderMap::new();
614			headers.insert(AUTHORIZATION, identity.parse().unwrap());
615			let request = Request::builder()
616				.method(Method::GET)
617				.uri("/account")
618				.version(Version::HTTP_11)
619				.headers(headers)
620				.body(Bytes::new())
621				.build()
622				.unwrap();
623
624			let response = middleware.process(request, handler.clone()).await.unwrap();
625			assert_eq!(response.body, identity);
626			assert_eq!(response.headers.get("x-cache").unwrap(), "MISS");
627		}
628
629		for identity in ["session=alice", "session=bob"] {
630			let mut headers = HeaderMap::new();
631			headers.insert(COOKIE, identity.parse().unwrap());
632			let request = Request::builder()
633				.method(Method::GET)
634				.uri("/account")
635				.version(Version::HTTP_11)
636				.headers(headers)
637				.body(Bytes::new())
638				.build()
639				.unwrap();
640
641			let response = middleware.process(request, handler.clone()).await.unwrap();
642			assert_eq!(response.body, identity);
643			assert_eq!(response.headers.get("x-cache").unwrap(), "MISS");
644		}
645
646		for identity in ["alice", "bob"] {
647			let mut headers = HeaderMap::new();
648			headers.insert("remote_user", identity.parse().unwrap());
649			let request = Request::builder()
650				.method(Method::GET)
651				.uri("/account")
652				.version(Version::HTTP_11)
653				.headers(headers)
654				.body(Bytes::new())
655				.build()
656				.unwrap();
657
658			let response = middleware.process(request, handler.clone()).await.unwrap();
659			assert_eq!(response.body, identity);
660			assert_eq!(response.headers.get("x-cache").unwrap(), "MISS");
661		}
662
663		let request = Request::builder()
664			.method(Method::GET)
665			.uri("/account")
666			.version(Version::HTTP_11)
667			.headers(HeaderMap::new())
668			.body(Bytes::new())
669			.build()
670			.unwrap();
671		request
672			.extensions
673			.insert(AuthState::authenticated("extension-user", false, true));
674		let response = middleware.process(request, handler.clone()).await.unwrap();
675		assert_eq!(response.body, "extension-user");
676		assert_eq!(response.headers.get("x-cache").unwrap(), "MISS");
677
678		for expected_cache in ["MISS", "HIT"] {
679			let request = Request::builder()
680				.method(Method::GET)
681				.uri("/public")
682				.version(Version::HTTP_11)
683				.headers(HeaderMap::new())
684				.body(Bytes::new())
685				.build()
686				.unwrap();
687			let response = middleware.process(request, handler.clone()).await.unwrap();
688			assert_eq!(response.body, "public");
689			assert_eq!(response.headers.get("x-cache").unwrap(), expected_cache);
690		}
691	}
692
693	#[tokio::test]
694	async fn authenticated_request_bypasses_an_existing_public_cache_entry() {
695		let middleware = CacheMiddleware::with_defaults();
696		let handler = Arc::new(IdentityHandler);
697
698		let public_request = Request::builder()
699			.method(Method::GET)
700			.uri("/account")
701			.version(Version::HTTP_11)
702			.headers(HeaderMap::new())
703			.body(Bytes::new())
704			.build()
705			.unwrap();
706		let public_response = middleware
707			.process(public_request, handler.clone())
708			.await
709			.unwrap();
710		assert_eq!(public_response.body, "public");
711		assert_eq!(public_response.headers.get("x-cache").unwrap(), "MISS");
712
713		let mut headers = HeaderMap::new();
714		headers.insert(AUTHORIZATION, "Bearer alice".parse().unwrap());
715		let authenticated_request = Request::builder()
716			.method(Method::GET)
717			.uri("/account")
718			.version(Version::HTTP_11)
719			.headers(headers)
720			.body(Bytes::new())
721			.build()
722			.unwrap();
723		let authenticated_response = middleware
724			.process(authenticated_request, handler)
725			.await
726			.unwrap();
727		assert_eq!(authenticated_response.body, "Bearer alice");
728		assert_eq!(
729			authenticated_response.headers.get("x-cache").unwrap(),
730			"MISS"
731		);
732	}
733
734	#[rstest::rstest]
735	#[case("Cache-Control", b"private")]
736	#[case("Cache-Control", b"PUBLIC, NO-STORE=\"field\"")]
737	#[case("Cache-Control", b"no-cache")]
738	#[case("Cache-Control", b"private=\"field-\x80\"")]
739	#[case("Set-Cookie", b"session=alice")]
740	#[case("Vary", b"Authorization")]
741	#[tokio::test]
742	async fn private_response_headers_prevent_shared_storage(
743		#[case] header_name: &str,
744		#[case] header_value: &[u8],
745	) {
746		struct PrivateResponseHandler {
747			header_name: hyper::header::HeaderName,
748			header_value: hyper::header::HeaderValue,
749			call_count: RwLock<usize>,
750		}
751
752		#[async_trait]
753		impl Handler for PrivateResponseHandler {
754			async fn handle(&self, _request: Request) -> Result<Response> {
755				let mut count = self.call_count.write().unwrap();
756				*count += 1;
757				let mut response = Response::new(StatusCode::OK).with_body(count.to_string());
758				response
759					.headers
760					.insert(self.header_name.clone(), self.header_value.clone());
761				Ok(response)
762			}
763		}
764
765		let middleware = CacheMiddleware::with_defaults();
766		let handler = Arc::new(PrivateResponseHandler {
767			header_name: header_name.parse().unwrap(),
768			header_value: hyper::header::HeaderValue::from_bytes(header_value).unwrap(),
769			call_count: RwLock::new(0),
770		});
771
772		for expected_body in ["1", "2"] {
773			let request = Request::builder()
774				.method(Method::GET)
775				.uri("/account")
776				.version(Version::HTTP_11)
777				.headers(HeaderMap::new())
778				.body(Bytes::new())
779				.build()
780				.unwrap();
781			let response = middleware.process(request, handler.clone()).await.unwrap();
782			assert_eq!(response.body, expected_body);
783			assert_eq!(response.headers.get("x-cache").unwrap(), "MISS");
784		}
785	}
786
787	#[tokio::test]
788	async fn url_and_headers_cache_preserves_supported_vary_responses() {
789		struct VaryHandler;
790
791		#[async_trait]
792		impl Handler for VaryHandler {
793			async fn handle(&self, _request: Request) -> Result<Response> {
794				Ok(Response::new(StatusCode::OK)
795					.with_body("public")
796					.with_header("Vary", "Accept-Encoding"))
797			}
798		}
799
800		let middleware = CacheMiddleware::new(CacheConfig::new(
801			Duration::from_secs(60),
802			CacheKeyStrategy::UrlAndHeaders,
803		));
804		let handler = Arc::new(VaryHandler);
805
806		for expected_cache in ["MISS", "HIT"] {
807			let request = Request::builder()
808				.method(Method::GET)
809				.uri("/public")
810				.version(Version::HTTP_11)
811				.headers(HeaderMap::new())
812				.body(Bytes::new())
813				.build()
814				.unwrap();
815			let response = middleware.process(request, handler.clone()).await.unwrap();
816			assert_eq!(response.headers.get("x-cache").unwrap(), expected_cache);
817		}
818	}
819
820	#[tokio::test]
821	async fn test_cache_miss() {
822		let config = CacheConfig::new(Duration::from_secs(60), CacheKeyStrategy::UrlOnly);
823		let middleware = CacheMiddleware::new(config);
824		let handler = Arc::new(TestHandler::new(StatusCode::OK));
825
826		let request = Request::builder()
827			.method(Method::GET)
828			.uri("/test")
829			.version(Version::HTTP_11)
830			.headers(HeaderMap::new())
831			.body(Bytes::new())
832			.build()
833			.unwrap();
834
835		let response = middleware.process(request, handler).await.unwrap();
836
837		assert_eq!(response.status, StatusCode::OK);
838		assert_eq!(response.headers.get("x-cache").unwrap(), "MISS");
839	}
840
841	#[tokio::test]
842	async fn test_cache_hit() {
843		let config = CacheConfig::new(Duration::from_secs(60), CacheKeyStrategy::UrlOnly);
844		let middleware = Arc::new(CacheMiddleware::new(config));
845		let handler = Arc::new(TestHandler::new(StatusCode::OK));
846
847		// First request (cache miss)
848		let request1 = Request::builder()
849			.method(Method::GET)
850			.uri("/test")
851			.version(Version::HTTP_11)
852			.headers(HeaderMap::new())
853			.body(Bytes::new())
854			.build()
855			.unwrap();
856		let response1 = middleware.process(request1, handler.clone()).await.unwrap();
857		assert_eq!(response1.headers.get("x-cache").unwrap(), "MISS");
858		assert_eq!(handler.get_call_count(), 1);
859
860		// Second request (cache hit)
861		let request2 = Request::builder()
862			.method(Method::GET)
863			.uri("/test")
864			.version(Version::HTTP_11)
865			.headers(HeaderMap::new())
866			.body(Bytes::new())
867			.build()
868			.unwrap();
869		let response2 = middleware.process(request2, handler.clone()).await.unwrap();
870		assert_eq!(response2.headers.get("x-cache").unwrap(), "HIT");
871		assert_eq!(handler.get_call_count(), 1); // Handler is not called
872	}
873
874	#[tokio::test]
875	async fn test_cache_expiration() {
876		let config = CacheConfig::new(Duration::from_millis(100), CacheKeyStrategy::UrlOnly);
877		let middleware = Arc::new(CacheMiddleware::new(config));
878		let handler = Arc::new(TestHandler::new(StatusCode::OK));
879
880		// First request
881		let request1 = Request::builder()
882			.method(Method::GET)
883			.uri("/test")
884			.version(Version::HTTP_11)
885			.headers(HeaderMap::new())
886			.body(Bytes::new())
887			.build()
888			.unwrap();
889		let _response1 = middleware.process(request1, handler.clone()).await.unwrap();
890
891		// Wait for expiration
892		std::thread::sleep(Duration::from_millis(150));
893
894		// Request after expiration (cache miss)
895		let request2 = Request::builder()
896			.method(Method::GET)
897			.uri("/test")
898			.version(Version::HTTP_11)
899			.headers(HeaderMap::new())
900			.body(Bytes::new())
901			.build()
902			.unwrap();
903		let response2 = middleware.process(request2, handler.clone()).await.unwrap();
904		assert_eq!(response2.headers.get("x-cache").unwrap(), "MISS");
905		assert_eq!(handler.get_call_count(), 2);
906	}
907
908	#[tokio::test]
909	async fn test_non_cacheable_method() {
910		let config = CacheConfig::new(Duration::from_secs(60), CacheKeyStrategy::UrlOnly);
911		let middleware = CacheMiddleware::new(config);
912		let handler = Arc::new(TestHandler::new(StatusCode::OK));
913
914		let request = Request::builder()
915			.method(Method::POST)
916			.uri("/test")
917			.version(Version::HTTP_11)
918			.headers(HeaderMap::new())
919			.body(Bytes::new())
920			.build()
921			.unwrap();
922
923		let response = middleware.process(request, handler).await.unwrap();
924
925		assert_eq!(response.status, StatusCode::OK);
926		assert!(!response.headers.contains_key("x-cache"));
927	}
928
929	#[tokio::test]
930	async fn test_exclude_paths() {
931		let config = CacheConfig::new(Duration::from_secs(60), CacheKeyStrategy::UrlOnly)
932			.with_excluded_paths(vec!["/admin".to_string()]);
933		let middleware = CacheMiddleware::new(config);
934		let handler = Arc::new(TestHandler::new(StatusCode::OK));
935
936		let request = Request::builder()
937			.method(Method::GET)
938			.uri("/admin/users")
939			.version(Version::HTTP_11)
940			.headers(HeaderMap::new())
941			.body(Bytes::new())
942			.build()
943			.unwrap();
944
945		let response = middleware.process(request, handler).await.unwrap();
946
947		assert_eq!(response.status, StatusCode::OK);
948		assert!(!response.headers.contains_key("x-cache"));
949	}
950
951	#[tokio::test]
952	async fn test_different_urls() {
953		let config = CacheConfig::new(Duration::from_secs(60), CacheKeyStrategy::UrlOnly);
954		let middleware = Arc::new(CacheMiddleware::new(config));
955		let handler = Arc::new(TestHandler::new(StatusCode::OK));
956
957		// Request to /test1
958		let request1 = Request::builder()
959			.method(Method::GET)
960			.uri("/test1")
961			.version(Version::HTTP_11)
962			.headers(HeaderMap::new())
963			.body(Bytes::new())
964			.build()
965			.unwrap();
966		let _response1 = middleware.process(request1, handler.clone()).await.unwrap();
967
968		// Request to /test2 (different cache entry)
969		let request2 = Request::builder()
970			.method(Method::GET)
971			.uri("/test2")
972			.version(Version::HTTP_11)
973			.headers(HeaderMap::new())
974			.body(Bytes::new())
975			.build()
976			.unwrap();
977		let response2 = middleware.process(request2, handler.clone()).await.unwrap();
978
979		assert_eq!(response2.headers.get("x-cache").unwrap(), "MISS");
980		assert_eq!(handler.get_call_count(), 2);
981	}
982
983	#[tokio::test]
984	async fn test_cache_store() {
985		let store = CacheStore::new();
986
987		let response = Response::new(StatusCode::OK).with_body(Bytes::from("test"));
988		let entry = CacheEntry::new(&response, Duration::from_secs(60));
989
990		store.set("key1".to_string(), entry.clone());
991
992		assert_eq!(store.len(), 1);
993		assert!(!store.is_empty());
994
995		let retrieved = store.get("key1").unwrap();
996		assert_eq!(retrieved.status, 200);
997		assert_eq!(retrieved.body, b"test");
998	}
999
1000	#[tokio::test]
1001	async fn test_cache_cleanup() {
1002		let store = CacheStore::new();
1003
1004		let response = Response::new(StatusCode::OK).with_body(Bytes::from("test"));
1005		let mut entry = CacheEntry::new(&response, Duration::from_millis(10));
1006		entry.cached_at = Some(Instant::now() - Duration::from_millis(20));
1007
1008		store.set("key1".to_string(), entry);
1009
1010		store.cleanup();
1011
1012		assert_eq!(store.len(), 0);
1013		assert!(store.is_empty());
1014	}
1015
1016	#[tokio::test]
1017	async fn test_multiple_status_codes_cached() {
1018		let config = CacheConfig::new(Duration::from_secs(60), CacheKeyStrategy::UrlOnly);
1019		let middleware = Arc::new(CacheMiddleware::new(config));
1020
1021		// Test with 404 status (cached by default)
1022		let handler_404 = Arc::new(TestHandler::new(StatusCode::NOT_FOUND));
1023		let request1 = Request::builder()
1024			.method(Method::GET)
1025			.uri("/not-found")
1026			.version(Version::HTTP_11)
1027			.headers(HeaderMap::new())
1028			.body(Bytes::new())
1029			.build()
1030			.unwrap();
1031		let response1 = middleware
1032			.process(request1, handler_404.clone())
1033			.await
1034			.unwrap();
1035		assert_eq!(response1.status, StatusCode::NOT_FOUND);
1036		assert_eq!(response1.headers.get("x-cache").unwrap(), "MISS");
1037		assert_eq!(handler_404.get_call_count(), 1);
1038
1039		// Second request to same 404 URL (cache hit)
1040		let request1b = Request::builder()
1041			.method(Method::GET)
1042			.uri("/not-found")
1043			.version(Version::HTTP_11)
1044			.headers(HeaderMap::new())
1045			.body(Bytes::new())
1046			.build()
1047			.unwrap();
1048		let response1b = middleware
1049			.process(request1b, handler_404.clone())
1050			.await
1051			.unwrap();
1052		assert_eq!(response1b.status, StatusCode::NOT_FOUND);
1053		assert_eq!(response1b.headers.get("x-cache").unwrap(), "HIT");
1054		assert_eq!(handler_404.get_call_count(), 1); // Not called again
1055
1056		// Test with 500 status (also cached by default)
1057		let handler_500 = Arc::new(TestHandler::new(StatusCode::INTERNAL_SERVER_ERROR));
1058		let request2 = Request::builder()
1059			.method(Method::GET)
1060			.uri("/error")
1061			.version(Version::HTTP_11)
1062			.headers(HeaderMap::new())
1063			.body(Bytes::new())
1064			.build()
1065			.unwrap();
1066		let response2 = middleware
1067			.process(request2, handler_500.clone())
1068			.await
1069			.unwrap();
1070		assert_eq!(response2.status, StatusCode::INTERNAL_SERVER_ERROR);
1071		assert_eq!(response2.headers.get("x-cache").unwrap(), "MISS");
1072	}
1073
1074	#[tokio::test]
1075	async fn test_cache_key_strategy_url_and_method() {
1076		let config = CacheConfig::new(Duration::from_secs(60), CacheKeyStrategy::UrlAndMethod);
1077		let middleware = Arc::new(CacheMiddleware::new(config));
1078		let handler = Arc::new(TestHandler::new(StatusCode::OK));
1079
1080		// GET request to /api
1081		let request1 = Request::builder()
1082			.method(Method::GET)
1083			.uri("/api")
1084			.version(Version::HTTP_11)
1085			.headers(HeaderMap::new())
1086			.body(Bytes::new())
1087			.build()
1088			.unwrap();
1089		let response1 = middleware.process(request1, handler.clone()).await.unwrap();
1090		assert_eq!(response1.headers.get("x-cache").unwrap(), "MISS");
1091		assert_eq!(handler.get_call_count(), 1);
1092
1093		// HEAD request to same URL (different cache key due to method)
1094		let handler2 = Arc::new(TestHandler::new(StatusCode::OK));
1095		let request2 = Request::builder()
1096			.method(Method::HEAD)
1097			.uri("/api")
1098			.version(Version::HTTP_11)
1099			.headers(HeaderMap::new())
1100			.body(Bytes::new())
1101			.build()
1102			.unwrap();
1103		let response2 = middleware
1104			.process(request2, handler2.clone())
1105			.await
1106			.unwrap();
1107		// Different method should result in cache miss
1108		assert_eq!(response2.headers.get("x-cache").unwrap(), "MISS");
1109		assert_eq!(handler2.get_call_count(), 1);
1110	}
1111
1112	#[rstest::rstest]
1113	fn test_rwlock_poison_recovery_cache_store() {
1114		// Arrange
1115		let store = Arc::new(CacheStore::new());
1116
1117		// Act - poison the RwLock by panicking while holding a write guard
1118		let store_clone = Arc::clone(&store);
1119		let _ = std::thread::spawn(move || {
1120			let _guard = store_clone.entries.write().unwrap();
1121			panic!("intentional panic to poison lock");
1122		})
1123		.join();
1124
1125		// Assert - operations still work after poison recovery
1126		let response = Response::new(StatusCode::OK).with_body(Bytes::from("test"));
1127		let entry = CacheEntry::new(&response, Duration::from_secs(60));
1128		store.set("key1".to_string(), entry);
1129		assert_eq!(store.len(), 1);
1130		assert!(!store.is_empty());
1131		assert!(store.get("key1").is_some());
1132		store.delete("key1");
1133		assert_eq!(store.len(), 0);
1134	}
1135}