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
6use async_trait::async_trait;
7use hyper::StatusCode;
8use reinhardt_http::{Handler, Middleware, Request, Response, Result};
9use serde::{Deserialize, Serialize};
10use sha2::{Digest, Sha256};
11use std::collections::HashMap;
12use std::sync::{Arc, RwLock};
13use std::time::{Duration, Instant};
14
15/// Cache Entry
16#[derive(Debug, Clone, Serialize, Deserialize)]
17pub struct CacheEntry {
18	/// Status code
19	status: u16,
20	/// Headers
21	headers: HashMap<String, String>,
22	/// Body
23	body: Vec<u8>,
24	/// Cached timestamp
25	#[serde(skip)]
26	cached_at: Option<Instant>,
27	/// TTL (seconds)
28	ttl_secs: u64,
29}
30
31impl CacheEntry {
32	/// Create a new entry
33	fn new(response: &Response, ttl: Duration) -> Self {
34		let mut headers = HashMap::new();
35		for (key, value) in response.headers.iter() {
36			if let Ok(value_str) = value.to_str() {
37				headers.insert(key.to_string(), value_str.to_string());
38			}
39		}
40
41		Self {
42			status: response.status.as_u16(),
43			headers,
44			body: response.body.to_vec(),
45			cached_at: Some(Instant::now()),
46			ttl_secs: ttl.as_secs(),
47		}
48	}
49
50	/// Check if expired
51	fn is_expired(&self) -> bool {
52		if let Some(cached_at) = self.cached_at {
53			cached_at.elapsed().as_secs() >= self.ttl_secs
54		} else {
55			true
56		}
57	}
58
59	/// Convert to response
60	fn to_response(&self) -> Response {
61		let status = StatusCode::from_u16(self.status).unwrap_or(StatusCode::OK);
62		let mut response = Response::new(status).with_body(self.body.clone());
63
64		for (key, value) in &self.headers {
65			if let (Ok(header_name), Ok(header_value)) =
66				(hyper::header::HeaderName::try_from(key), value.parse())
67			{
68				response.headers.insert(header_name, header_value);
69			}
70		}
71
72		// Add cache header
73		response.headers.insert(
74			hyper::header::HeaderName::from_static("x-cache"),
75			hyper::header::HeaderValue::from_static("HIT"),
76		);
77
78		response
79	}
80}
81
82/// Cache Storage
83#[derive(Debug, Default)]
84pub struct CacheStore {
85	/// Entries
86	entries: RwLock<HashMap<String, CacheEntry>>,
87}
88
89impl CacheStore {
90	/// Create a new store
91	pub fn new() -> Self {
92		Self::default()
93	}
94
95	/// Get an entry
96	pub fn get(&self, key: &str) -> Option<CacheEntry> {
97		let entries = self.entries.read().unwrap_or_else(|e| e.into_inner());
98		entries.get(key).cloned()
99	}
100
101	/// Set an entry
102	pub fn set(&self, key: String, entry: CacheEntry) {
103		let mut entries = self.entries.write().unwrap_or_else(|e| e.into_inner());
104		entries.insert(key, entry);
105	}
106
107	/// Delete an entry
108	pub fn delete(&self, key: &str) {
109		let mut entries = self.entries.write().unwrap_or_else(|e| e.into_inner());
110		entries.remove(key);
111	}
112
113	/// Clean up expired entries
114	pub fn cleanup(&self) {
115		let mut entries = self.entries.write().unwrap_or_else(|e| e.into_inner());
116		entries.retain(|_, entry| !entry.is_expired());
117	}
118
119	/// Clear the store
120	pub fn clear(&self) {
121		let mut entries = self.entries.write().unwrap_or_else(|e| e.into_inner());
122		entries.clear();
123	}
124
125	/// Get the number of entries
126	pub fn len(&self) -> usize {
127		let entries = self.entries.read().unwrap_or_else(|e| e.into_inner());
128		entries.len()
129	}
130
131	/// Check if the store is empty
132	pub fn is_empty(&self) -> bool {
133		let entries = self.entries.read().unwrap_or_else(|e| e.into_inner());
134		entries.is_empty()
135	}
136}
137
138/// Cache key generation strategy
139#[derive(Debug, Clone, Copy)]
140pub enum CacheKeyStrategy {
141	/// URL only
142	UrlOnly,
143	/// URL and method
144	UrlAndMethod,
145	/// URL and query parameters
146	UrlAndQuery,
147	/// URL and headers
148	UrlAndHeaders,
149}
150
151/// Cache configuration
152#[non_exhaustive]
153#[derive(Debug, Clone)]
154pub struct CacheConfig {
155	/// Default TTL
156	pub default_ttl: Duration,
157	/// Cache key generation strategy
158	pub key_strategy: CacheKeyStrategy,
159	/// Cacheable methods
160	pub cacheable_methods: Vec<String>,
161	/// Cacheable status codes
162	pub cacheable_status_codes: Vec<u16>,
163	/// Paths to exclude
164	pub exclude_paths: Vec<String>,
165	/// Maximum cache size
166	pub max_entries: Option<usize>,
167}
168
169impl CacheConfig {
170	/// Create a new configuration
171	///
172	/// # Examples
173	///
174	/// ```
175	/// use std::time::Duration;
176	/// use reinhardt_middleware::cache::{CacheConfig, CacheKeyStrategy};
177	///
178	/// let config = CacheConfig::new(Duration::from_secs(300), CacheKeyStrategy::UrlOnly);
179	/// assert_eq!(config.default_ttl, Duration::from_secs(300));
180	/// ```
181	pub fn new(default_ttl: Duration, key_strategy: CacheKeyStrategy) -> Self {
182		Self {
183			default_ttl,
184			key_strategy,
185			cacheable_methods: vec!["GET".to_string(), "HEAD".to_string()],
186			cacheable_status_codes: vec![200, 203, 204, 206, 300, 301, 404, 405, 410, 414, 501],
187			exclude_paths: Vec::new(),
188			max_entries: Some(1000),
189		}
190	}
191
192	/// Set cacheable methods
193	///
194	/// # Examples
195	///
196	/// ```
197	/// use std::time::Duration;
198	/// use reinhardt_middleware::cache::{CacheConfig, CacheKeyStrategy};
199	///
200	/// let config = CacheConfig::new(Duration::from_secs(300), CacheKeyStrategy::UrlOnly)
201	///     .with_cacheable_methods(vec!["GET".to_string()]);
202	/// ```
203	pub fn with_cacheable_methods(mut self, methods: Vec<String>) -> Self {
204		self.cacheable_methods = methods;
205		self
206	}
207
208	/// Add paths to exclude
209	///
210	/// # Examples
211	///
212	/// ```
213	/// use std::time::Duration;
214	/// use reinhardt_middleware::cache::{CacheConfig, CacheKeyStrategy};
215	///
216	/// let config = CacheConfig::new(Duration::from_secs(300), CacheKeyStrategy::UrlOnly)
217	///     .with_excluded_paths(vec!["/admin".to_string()]);
218	/// ```
219	pub fn with_excluded_paths(mut self, paths: Vec<String>) -> Self {
220		self.exclude_paths.extend(paths);
221		self
222	}
223
224	/// Set maximum number of entries
225	///
226	/// # Examples
227	///
228	/// ```
229	/// use std::time::Duration;
230	/// use reinhardt_middleware::cache::{CacheConfig, CacheKeyStrategy};
231	///
232	/// let config = CacheConfig::new(Duration::from_secs(300), CacheKeyStrategy::UrlOnly)
233	///     .with_max_entries(5000);
234	/// ```
235	pub fn with_max_entries(mut self, max_entries: usize) -> Self {
236		self.max_entries = Some(max_entries);
237		self
238	}
239}
240
241impl Default for CacheConfig {
242	fn default() -> Self {
243		Self::new(Duration::from_secs(300), CacheKeyStrategy::UrlOnly)
244	}
245}
246
247/// Cache Middleware
248///
249/// # Examples
250///
251/// ```
252/// use std::sync::Arc;
253/// use std::time::Duration;
254/// use reinhardt_middleware::cache::{CacheMiddleware, CacheConfig, CacheKeyStrategy};
255/// use reinhardt_http::{Handler, Middleware, Request, Response};
256/// use hyper::{StatusCode, Method, Version, HeaderMap};
257/// use bytes::Bytes;
258///
259/// struct TestHandler;
260///
261/// #[async_trait::async_trait]
262/// impl Handler for TestHandler {
263///     async fn handle(&self, _request: Request) -> reinhardt_core::exception::Result<Response> {
264///         Ok(Response::new(StatusCode::OK).with_body(Bytes::from("OK")))
265///     }
266/// }
267///
268/// # tokio_test::block_on(async {
269/// let config = CacheConfig::new(Duration::from_secs(60), CacheKeyStrategy::UrlOnly);
270/// let middleware = CacheMiddleware::new(config);
271/// let handler = Arc::new(TestHandler);
272///
273/// let request = Request::builder()
274///     .method(Method::GET)
275///     .uri("/api/data")
276///     .version(Version::HTTP_11)
277///     .headers(HeaderMap::new())
278///     .body(Bytes::new())
279///     .build()
280///     .unwrap();
281///
282/// let response = middleware.process(request, handler).await.unwrap();
283/// assert_eq!(response.status, StatusCode::OK);
284/// # });
285/// ```
286pub struct CacheMiddleware {
287	config: CacheConfig,
288	store: Arc<CacheStore>,
289}
290
291impl CacheMiddleware {
292	/// Create a new cache middleware
293	///
294	/// # Examples
295	///
296	/// ```
297	/// use std::time::Duration;
298	/// use reinhardt_middleware::cache::{CacheMiddleware, CacheConfig, CacheKeyStrategy};
299	///
300	/// let config = CacheConfig::new(Duration::from_secs(300), CacheKeyStrategy::UrlOnly);
301	/// let middleware = CacheMiddleware::new(config);
302	/// ```
303	pub fn new(config: CacheConfig) -> Self {
304		Self {
305			config,
306			store: Arc::new(CacheStore::new()),
307		}
308	}
309
310	/// Create with default configuration
311	pub fn with_defaults() -> Self {
312		Self::new(CacheConfig::default())
313	}
314
315	/// Create from an existing Arc-wrapped cache store
316	///
317	/// This is provided for cases where you already have an `Arc<CacheStore>`.
318	/// In most cases, you should use `new()` instead, which creates the store internally.
319	pub fn from_arc(config: CacheConfig, store: Arc<CacheStore>) -> Self {
320		Self { config, store }
321	}
322
323	/// Get a reference to the cache store
324	///
325	/// # Examples
326	///
327	/// ```
328	/// use std::time::Duration;
329	/// use reinhardt_middleware::cache::{CacheMiddleware, CacheConfig, CacheKeyStrategy};
330	///
331	/// let middleware = CacheMiddleware::new(
332	///     CacheConfig::new(Duration::from_secs(300), CacheKeyStrategy::UrlOnly)
333	/// );
334	///
335	/// // Access the store
336	/// let store = middleware.store();
337	/// assert_eq!(store.len(), 0);
338	/// ```
339	pub fn store(&self) -> &CacheStore {
340		&self.store
341	}
342
343	/// Get a cloned Arc of the store (for cases where you need ownership)
344	///
345	/// In most cases, you should use `store()` instead to get a reference.
346	pub fn store_arc(&self) -> Arc<CacheStore> {
347		Arc::clone(&self.store)
348	}
349
350	/// Check if path should be excluded
351	fn should_exclude(&self, path: &str) -> bool {
352		self.config
353			.exclude_paths
354			.iter()
355			.any(|p| path.starts_with(p))
356	}
357
358	/// Check if method is cacheable
359	fn is_cacheable_method(&self, method: &str) -> bool {
360		self.config.cacheable_methods.iter().any(|m| m == method)
361	}
362
363	/// Check if status code is cacheable
364	fn is_cacheable_status(&self, status: u16) -> bool {
365		self.config.cacheable_status_codes.contains(&status)
366	}
367
368	/// Generate cache key
369	fn generate_cache_key(&self, request: &Request) -> String {
370		let base = match self.config.key_strategy {
371			CacheKeyStrategy::UrlOnly => request.uri.path().to_string(),
372			CacheKeyStrategy::UrlAndMethod => {
373				format!("{}:{}", request.method.as_str(), request.uri.path())
374			}
375			CacheKeyStrategy::UrlAndQuery => {
376				let query = request.uri.query().unwrap_or("");
377				format!(
378					"{}:{}?{}",
379					request.method.as_str(),
380					request.uri.path(),
381					query
382				)
383			}
384			CacheKeyStrategy::UrlAndHeaders => {
385				let headers_str = request
386					.headers
387					.iter()
388					.map(|(k, v)| format!("{}={}", k, v.to_str().unwrap_or("")))
389					.collect::<Vec<_>>()
390					.join("&");
391				format!(
392					"{}:{}:{}",
393					request.method.as_str(),
394					request.uri.path(),
395					headers_str
396				)
397			}
398		};
399
400		// Hash with SHA256
401		let mut hasher = Sha256::new();
402		hasher.update(base.as_bytes());
403		let result = hasher.finalize();
404		hex::encode(result)
405	}
406}
407
408impl Default for CacheMiddleware {
409	fn default() -> Self {
410		Self::with_defaults()
411	}
412}
413
414#[async_trait]
415impl Middleware for CacheMiddleware {
416	async fn process(&self, request: Request, handler: Arc<dyn Handler>) -> Result<Response> {
417		let path = request.uri.path().to_string();
418		let method = request.method.as_str().to_string();
419
420		// Skip excluded paths
421		if self.should_exclude(&path) {
422			return handler.handle(request).await;
423		}
424
425		// Skip non-cacheable methods
426		if !self.is_cacheable_method(&method) {
427			return handler.handle(request).await;
428		}
429
430		// Generate cache key
431		let cache_key = self.generate_cache_key(&request);
432
433		// Check cache
434		if let Some(entry) = self.store.get(&cache_key) {
435			if !entry.is_expired() {
436				// Cache hit
437				return Ok(entry.to_response());
438			} else {
439				// Delete expired entry
440				self.store.delete(&cache_key);
441			}
442		}
443
444		// Convert errors to responses so post-processing always runs,
445		// even when invoked outside MiddlewareChain. (#3244)
446		let response = match handler.handle(request).await {
447			Ok(resp) => resp,
448			Err(e) => Response::from(e),
449		};
450
451		// Save to cache if status code is cacheable
452		if self.is_cacheable_status(response.status.as_u16()) {
453			let entry = CacheEntry::new(&response, self.config.default_ttl);
454			self.store.set(cache_key, entry);
455
456			// Clean up expired entries if max entries exceeded
457			if let Some(max_entries) = self.config.max_entries
458				&& self.store.len() > max_entries
459			{
460				self.store.cleanup();
461			}
462		}
463
464		// Add X-Cache header
465		let mut response = response;
466		response.headers.insert(
467			hyper::header::HeaderName::from_static("x-cache"),
468			hyper::header::HeaderValue::from_static("MISS"),
469		);
470
471		Ok(response)
472	}
473}
474
475#[cfg(test)]
476mod tests {
477	use super::*;
478	use bytes::Bytes;
479	use hyper::{HeaderMap, Method, StatusCode, Version};
480
481	struct TestHandler {
482		status: StatusCode,
483		call_count: Arc<RwLock<usize>>,
484	}
485
486	impl TestHandler {
487		fn new(status: StatusCode) -> Self {
488			Self {
489				status,
490				call_count: Arc::new(RwLock::new(0)),
491			}
492		}
493
494		fn get_call_count(&self) -> usize {
495			*self.call_count.read().unwrap()
496		}
497	}
498
499	#[async_trait]
500	impl Handler for TestHandler {
501		async fn handle(&self, _request: Request) -> Result<Response> {
502			*self.call_count.write().unwrap() += 1;
503			Ok(Response::new(self.status).with_body(Bytes::from("OK")))
504		}
505	}
506
507	#[tokio::test]
508	async fn test_cache_miss() {
509		let config = CacheConfig::new(Duration::from_secs(60), CacheKeyStrategy::UrlOnly);
510		let middleware = CacheMiddleware::new(config);
511		let handler = Arc::new(TestHandler::new(StatusCode::OK));
512
513		let request = Request::builder()
514			.method(Method::GET)
515			.uri("/test")
516			.version(Version::HTTP_11)
517			.headers(HeaderMap::new())
518			.body(Bytes::new())
519			.build()
520			.unwrap();
521
522		let response = middleware.process(request, handler).await.unwrap();
523
524		assert_eq!(response.status, StatusCode::OK);
525		assert_eq!(response.headers.get("x-cache").unwrap(), "MISS");
526	}
527
528	#[tokio::test]
529	async fn test_cache_hit() {
530		let config = CacheConfig::new(Duration::from_secs(60), CacheKeyStrategy::UrlOnly);
531		let middleware = Arc::new(CacheMiddleware::new(config));
532		let handler = Arc::new(TestHandler::new(StatusCode::OK));
533
534		// First request (cache miss)
535		let request1 = Request::builder()
536			.method(Method::GET)
537			.uri("/test")
538			.version(Version::HTTP_11)
539			.headers(HeaderMap::new())
540			.body(Bytes::new())
541			.build()
542			.unwrap();
543		let response1 = middleware.process(request1, handler.clone()).await.unwrap();
544		assert_eq!(response1.headers.get("x-cache").unwrap(), "MISS");
545		assert_eq!(handler.get_call_count(), 1);
546
547		// Second request (cache hit)
548		let request2 = Request::builder()
549			.method(Method::GET)
550			.uri("/test")
551			.version(Version::HTTP_11)
552			.headers(HeaderMap::new())
553			.body(Bytes::new())
554			.build()
555			.unwrap();
556		let response2 = middleware.process(request2, handler.clone()).await.unwrap();
557		assert_eq!(response2.headers.get("x-cache").unwrap(), "HIT");
558		assert_eq!(handler.get_call_count(), 1); // Handler is not called
559	}
560
561	#[tokio::test]
562	async fn test_cache_expiration() {
563		let config = CacheConfig::new(Duration::from_millis(100), CacheKeyStrategy::UrlOnly);
564		let middleware = Arc::new(CacheMiddleware::new(config));
565		let handler = Arc::new(TestHandler::new(StatusCode::OK));
566
567		// First request
568		let request1 = Request::builder()
569			.method(Method::GET)
570			.uri("/test")
571			.version(Version::HTTP_11)
572			.headers(HeaderMap::new())
573			.body(Bytes::new())
574			.build()
575			.unwrap();
576		let _response1 = middleware.process(request1, handler.clone()).await.unwrap();
577
578		// Wait for expiration
579		std::thread::sleep(Duration::from_millis(150));
580
581		// Request after expiration (cache miss)
582		let request2 = Request::builder()
583			.method(Method::GET)
584			.uri("/test")
585			.version(Version::HTTP_11)
586			.headers(HeaderMap::new())
587			.body(Bytes::new())
588			.build()
589			.unwrap();
590		let response2 = middleware.process(request2, handler.clone()).await.unwrap();
591		assert_eq!(response2.headers.get("x-cache").unwrap(), "MISS");
592		assert_eq!(handler.get_call_count(), 2);
593	}
594
595	#[tokio::test]
596	async fn test_non_cacheable_method() {
597		let config = CacheConfig::new(Duration::from_secs(60), CacheKeyStrategy::UrlOnly);
598		let middleware = CacheMiddleware::new(config);
599		let handler = Arc::new(TestHandler::new(StatusCode::OK));
600
601		let request = Request::builder()
602			.method(Method::POST)
603			.uri("/test")
604			.version(Version::HTTP_11)
605			.headers(HeaderMap::new())
606			.body(Bytes::new())
607			.build()
608			.unwrap();
609
610		let response = middleware.process(request, handler).await.unwrap();
611
612		assert_eq!(response.status, StatusCode::OK);
613		assert!(!response.headers.contains_key("x-cache"));
614	}
615
616	#[tokio::test]
617	async fn test_exclude_paths() {
618		let config = CacheConfig::new(Duration::from_secs(60), CacheKeyStrategy::UrlOnly)
619			.with_excluded_paths(vec!["/admin".to_string()]);
620		let middleware = CacheMiddleware::new(config);
621		let handler = Arc::new(TestHandler::new(StatusCode::OK));
622
623		let request = Request::builder()
624			.method(Method::GET)
625			.uri("/admin/users")
626			.version(Version::HTTP_11)
627			.headers(HeaderMap::new())
628			.body(Bytes::new())
629			.build()
630			.unwrap();
631
632		let response = middleware.process(request, handler).await.unwrap();
633
634		assert_eq!(response.status, StatusCode::OK);
635		assert!(!response.headers.contains_key("x-cache"));
636	}
637
638	#[tokio::test]
639	async fn test_different_urls() {
640		let config = CacheConfig::new(Duration::from_secs(60), CacheKeyStrategy::UrlOnly);
641		let middleware = Arc::new(CacheMiddleware::new(config));
642		let handler = Arc::new(TestHandler::new(StatusCode::OK));
643
644		// Request to /test1
645		let request1 = Request::builder()
646			.method(Method::GET)
647			.uri("/test1")
648			.version(Version::HTTP_11)
649			.headers(HeaderMap::new())
650			.body(Bytes::new())
651			.build()
652			.unwrap();
653		let _response1 = middleware.process(request1, handler.clone()).await.unwrap();
654
655		// Request to /test2 (different cache entry)
656		let request2 = Request::builder()
657			.method(Method::GET)
658			.uri("/test2")
659			.version(Version::HTTP_11)
660			.headers(HeaderMap::new())
661			.body(Bytes::new())
662			.build()
663			.unwrap();
664		let response2 = middleware.process(request2, handler.clone()).await.unwrap();
665
666		assert_eq!(response2.headers.get("x-cache").unwrap(), "MISS");
667		assert_eq!(handler.get_call_count(), 2);
668	}
669
670	#[tokio::test]
671	async fn test_cache_store() {
672		let store = CacheStore::new();
673
674		let response = Response::new(StatusCode::OK).with_body(Bytes::from("test"));
675		let entry = CacheEntry::new(&response, Duration::from_secs(60));
676
677		store.set("key1".to_string(), entry.clone());
678
679		assert_eq!(store.len(), 1);
680		assert!(!store.is_empty());
681
682		let retrieved = store.get("key1").unwrap();
683		assert_eq!(retrieved.status, 200);
684		assert_eq!(retrieved.body, b"test");
685	}
686
687	#[tokio::test]
688	async fn test_cache_cleanup() {
689		let store = CacheStore::new();
690
691		let response = Response::new(StatusCode::OK).with_body(Bytes::from("test"));
692		let mut entry = CacheEntry::new(&response, Duration::from_millis(10));
693		entry.cached_at = Some(Instant::now() - Duration::from_millis(20));
694
695		store.set("key1".to_string(), entry);
696
697		store.cleanup();
698
699		assert_eq!(store.len(), 0);
700		assert!(store.is_empty());
701	}
702
703	#[tokio::test]
704	async fn test_multiple_status_codes_cached() {
705		let config = CacheConfig::new(Duration::from_secs(60), CacheKeyStrategy::UrlOnly);
706		let middleware = Arc::new(CacheMiddleware::new(config));
707
708		// Test with 404 status (cached by default)
709		let handler_404 = Arc::new(TestHandler::new(StatusCode::NOT_FOUND));
710		let request1 = Request::builder()
711			.method(Method::GET)
712			.uri("/not-found")
713			.version(Version::HTTP_11)
714			.headers(HeaderMap::new())
715			.body(Bytes::new())
716			.build()
717			.unwrap();
718		let response1 = middleware
719			.process(request1, handler_404.clone())
720			.await
721			.unwrap();
722		assert_eq!(response1.status, StatusCode::NOT_FOUND);
723		assert_eq!(response1.headers.get("x-cache").unwrap(), "MISS");
724		assert_eq!(handler_404.get_call_count(), 1);
725
726		// Second request to same 404 URL (cache hit)
727		let request1b = Request::builder()
728			.method(Method::GET)
729			.uri("/not-found")
730			.version(Version::HTTP_11)
731			.headers(HeaderMap::new())
732			.body(Bytes::new())
733			.build()
734			.unwrap();
735		let response1b = middleware
736			.process(request1b, handler_404.clone())
737			.await
738			.unwrap();
739		assert_eq!(response1b.status, StatusCode::NOT_FOUND);
740		assert_eq!(response1b.headers.get("x-cache").unwrap(), "HIT");
741		assert_eq!(handler_404.get_call_count(), 1); // Not called again
742
743		// Test with 500 status (also cached by default)
744		let handler_500 = Arc::new(TestHandler::new(StatusCode::INTERNAL_SERVER_ERROR));
745		let request2 = Request::builder()
746			.method(Method::GET)
747			.uri("/error")
748			.version(Version::HTTP_11)
749			.headers(HeaderMap::new())
750			.body(Bytes::new())
751			.build()
752			.unwrap();
753		let response2 = middleware
754			.process(request2, handler_500.clone())
755			.await
756			.unwrap();
757		assert_eq!(response2.status, StatusCode::INTERNAL_SERVER_ERROR);
758		assert_eq!(response2.headers.get("x-cache").unwrap(), "MISS");
759	}
760
761	#[tokio::test]
762	async fn test_cache_key_strategy_url_and_method() {
763		let config = CacheConfig::new(Duration::from_secs(60), CacheKeyStrategy::UrlAndMethod);
764		let middleware = Arc::new(CacheMiddleware::new(config));
765		let handler = Arc::new(TestHandler::new(StatusCode::OK));
766
767		// GET request to /api
768		let request1 = Request::builder()
769			.method(Method::GET)
770			.uri("/api")
771			.version(Version::HTTP_11)
772			.headers(HeaderMap::new())
773			.body(Bytes::new())
774			.build()
775			.unwrap();
776		let response1 = middleware.process(request1, handler.clone()).await.unwrap();
777		assert_eq!(response1.headers.get("x-cache").unwrap(), "MISS");
778		assert_eq!(handler.get_call_count(), 1);
779
780		// HEAD request to same URL (different cache key due to method)
781		let handler2 = Arc::new(TestHandler::new(StatusCode::OK));
782		let request2 = Request::builder()
783			.method(Method::HEAD)
784			.uri("/api")
785			.version(Version::HTTP_11)
786			.headers(HeaderMap::new())
787			.body(Bytes::new())
788			.build()
789			.unwrap();
790		let response2 = middleware
791			.process(request2, handler2.clone())
792			.await
793			.unwrap();
794		// Different method should result in cache miss
795		assert_eq!(response2.headers.get("x-cache").unwrap(), "MISS");
796		assert_eq!(handler2.get_call_count(), 1);
797	}
798
799	#[rstest::rstest]
800	fn test_rwlock_poison_recovery_cache_store() {
801		// Arrange
802		let store = Arc::new(CacheStore::new());
803
804		// Act - poison the RwLock by panicking while holding a write guard
805		let store_clone = Arc::clone(&store);
806		let _ = std::thread::spawn(move || {
807			let _guard = store_clone.entries.write().unwrap();
808			panic!("intentional panic to poison lock");
809		})
810		.join();
811
812		// Assert - operations still work after poison recovery
813		let response = Response::new(StatusCode::OK).with_body(Bytes::from("test"));
814		let entry = CacheEntry::new(&response, Duration::from_secs(60));
815		store.set("key1".to_string(), entry);
816		assert_eq!(store.len(), 1);
817		assert!(!store.is_empty());
818		assert!(store.get("key1").is_some());
819		store.delete("key1");
820		assert_eq!(store.len(), 0);
821	}
822}