1use 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#[derive(Debug, Clone, Serialize, Deserialize)]
20pub struct CacheEntry {
21 status: u16,
23 headers: HashMap<String, String>,
25 body: Vec<u8>,
27 #[serde(skip)]
29 cached_at: Option<Instant>,
30 ttl_secs: u64,
32}
33
34impl CacheEntry {
35 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 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 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 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 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#[derive(Debug, Default)]
113pub struct CacheStore {
114 entries: RwLock<HashMap<String, CacheEntry>>,
116}
117
118impl CacheStore {
119 pub fn new() -> Self {
121 Self::default()
122 }
123
124 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 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 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 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 pub fn clear(&self) {
150 let mut entries = self.entries.write().unwrap_or_else(|e| e.into_inner());
151 entries.clear();
152 }
153
154 pub fn len(&self) -> usize {
156 let entries = self.entries.read().unwrap_or_else(|e| e.into_inner());
157 entries.len()
158 }
159
160 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#[derive(Debug, Clone, Copy)]
169pub enum CacheKeyStrategy {
170 UrlOnly,
172 UrlAndMethod,
174 UrlAndQuery,
176 UrlAndHeaders,
178}
179
180#[non_exhaustive]
182#[derive(Debug, Clone)]
183pub struct CacheConfig {
184 pub default_ttl: Duration,
186 pub key_strategy: CacheKeyStrategy,
188 pub cacheable_methods: Vec<String>,
190 pub cacheable_status_codes: Vec<u16>,
192 pub exclude_paths: Vec<String>,
194 pub max_entries: Option<usize>,
196}
197
198impl CacheConfig {
199 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 pub fn with_cacheable_methods(mut self, methods: Vec<String>) -> Self {
233 self.cacheable_methods = methods;
234 self
235 }
236
237 pub fn with_excluded_paths(mut self, paths: Vec<String>) -> Self {
249 self.exclude_paths.extend(paths);
250 self
251 }
252
253 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
276pub struct CacheMiddleware {
316 config: CacheConfig,
317 store: Arc<CacheStore>,
318}
319
320impl CacheMiddleware {
321 pub fn new(config: CacheConfig) -> Self {
333 Self {
334 config,
335 store: Arc::new(CacheStore::new()),
336 }
337 }
338
339 pub fn with_defaults() -> Self {
341 Self::new(CacheConfig::default())
342 }
343
344 pub fn from_arc(config: CacheConfig, store: Arc<CacheStore>) -> Self {
349 Self { config, store }
350 }
351
352 pub fn store(&self) -> &CacheStore {
369 &self.store
370 }
371
372 pub fn store_arc(&self) -> Arc<CacheStore> {
376 Arc::clone(&self.store)
377 }
378
379 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 fn is_cacheable_method(&self, method: &str) -> bool {
389 self.config.cacheable_methods.iter().any(|m| m == method)
390 }
391
392 fn is_cacheable_status(&self, status: u16) -> bool {
394 self.config.cacheable_status_codes.contains(&status)
395 }
396
397 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 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 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 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 if self.should_exclude(&path) {
479 return handler.handle(request).await;
480 }
481
482 if !self.is_cacheable_method(&method) {
484 return handler.handle(request).await;
485 }
486
487 let cache_key =
489 (!Self::is_private_request(&request)).then(|| self.generate_cache_key(&request));
490
491 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 return Ok(entry.to_response());
501 } else {
502 self.store.delete(cache_key);
504 }
505 }
506
507 let response = match handler.handle(request).await {
510 Ok(resp) => resp,
511 Err(e) => Response::from(e),
512 };
513
514 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 if let Some(max_entries) = self.config.max_entries
524 && self.store.len() > max_entries
525 {
526 self.store.cleanup();
527 }
528 }
529
530 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 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 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); }
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 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 std::thread::sleep(Duration::from_millis(150));
893
894 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 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 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 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 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); 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 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 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 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 let store = Arc::new(CacheStore::new());
1116
1117 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 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}