Skip to main content

tower_resilience_cache/
lib.rs

1//! Response caching middleware for Tower services.
2//!
3//! This crate provides a Tower middleware for caching service responses,
4//! reducing load on downstream services by storing and reusing responses
5//! for identical requests.
6//!
7//! # Features
8//!
9//! - **Multiple Eviction Policies**: LRU, LFU, and FIFO eviction strategies
10//! - **TTL Support**: Optional time-to-live for cache entries
11//! - **Event System**: Observability through cache events (Hit, Miss, Eviction)
12//! - **Flexible Key Extraction**: User-defined key extraction from requests
13//!
14//! # Examples
15//!
16//! ```
17//! use tower_resilience_cache::{CacheLayer, EvictionPolicy};
18//! use tower::ServiceBuilder;
19//! use std::time::Duration;
20//!
21//! # async fn example() -> Result<(), Box<dyn std::error::Error>> {
22//! // Create a cache layer with LFU eviction
23//! let cache_layer = CacheLayer::builder()
24//!     .max_size(100)
25//!     .ttl(Duration::from_secs(60))
26//!     .eviction_policy(EvictionPolicy::Lfu)  // or Lru (default), Fifo
27//!     .key_extractor(|req: &String| req.clone())
28//!     .on_hit(|| println!("Cache hit!"))
29//!     .on_miss(|| println!("Cache miss!"))
30//!     .build()?;
31//!
32//! // Apply to a service
33//! let service = ServiceBuilder::new()
34//!     .layer(cache_layer)
35//!     .service(tower::service_fn(|req: String| async move {
36//!         Ok::<_, std::io::Error>(format!("Response: {}", req))
37//!     }));
38//! # Ok(())
39//! # }
40//! ```
41
42mod config;
43mod error;
44mod events;
45mod eviction;
46mod layer;
47mod shared_layer;
48mod store;
49
50pub use config::{CacheConfig, CacheConfigBuilder, KeyExtractor};
51pub use error::{CacheBuildError, CacheError};
52pub use events::CacheEvent;
53pub use eviction::EvictionPolicy;
54pub use layer::CacheLayer;
55pub use shared_layer::{SharedCacheConfigBuilder, SharedCacheLayer};
56
57use futures::future::BoxFuture;
58use std::hash::Hash;
59use std::sync::{Arc, Mutex};
60use std::task::{Context, Poll};
61use std::time::Instant;
62use store::CacheStore;
63use tower::Service;
64
65#[cfg(feature = "metrics")]
66use metrics::{counter, describe_counter, describe_gauge, gauge};
67
68#[cfg(feature = "tracing")]
69use tracing::{debug, info};
70
71/// A Tower [`Service`] that caches responses.
72///
73/// This service wraps an inner service and caches successful responses.
74/// When a request comes in, the cache checks if a valid cached response
75/// exists. If so, it returns the cached value immediately without calling
76/// the inner service.
77///
78/// Responses must implement `Clone` to be cacheable.
79pub struct Cache<S, Req, K, Resp> {
80    inner: S,
81    config: Arc<CacheConfig<Req, K>>,
82    store: Arc<Mutex<CacheStore<K, Resp>>>,
83}
84
85impl<S, Req, K, Resp> Cache<S, Req, K, Resp>
86where
87    K: Hash + Eq + Clone + Send + 'static,
88    Resp: Clone + Send + 'static,
89{
90    /// Creates a new `Cache` wrapping the given service.
91    pub fn new(inner: S, config: Arc<CacheConfig<Req, K>>) -> Self {
92        #[cfg(feature = "metrics")]
93        {
94            describe_counter!(
95                "cache_requests_total",
96                "Total number of cache requests (hits and misses)"
97            );
98            describe_counter!("cache_evictions_total", "Total number of cache evictions");
99            describe_gauge!("cache_size", "Current number of entries in the cache");
100        }
101
102        let store = Arc::new(Mutex::new(CacheStore::new(
103            config.max_size,
104            config.ttl,
105            config.eviction_policy,
106        )));
107        Self {
108            inner,
109            config,
110            store,
111        }
112    }
113
114    /// Creates a new `Cache` wrapping the given service with a pre-existing store.
115    ///
116    /// This is used by [`SharedCacheLayer`] to share the same cache store across
117    /// multiple services.
118    pub(crate) fn with_store(
119        inner: S,
120        config: Arc<CacheConfig<Req, K>>,
121        store: Arc<Mutex<CacheStore<K, Resp>>>,
122    ) -> Self {
123        #[cfg(feature = "metrics")]
124        {
125            describe_counter!(
126                "cache_requests_total",
127                "Total number of cache requests (hits and misses)"
128            );
129            describe_counter!("cache_evictions_total", "Total number of cache evictions");
130            describe_gauge!("cache_size", "Current number of entries in the cache");
131        }
132
133        Self {
134            inner,
135            config,
136            store,
137        }
138    }
139}
140
141impl<S, Req, K, Resp> Clone for Cache<S, Req, K, Resp>
142where
143    S: Clone,
144{
145    fn clone(&self) -> Self {
146        Self {
147            inner: self.inner.clone(),
148            config: Arc::clone(&self.config),
149            store: Arc::clone(&self.store),
150        }
151    }
152}
153
154impl<S, Req, K> Service<Req> for Cache<S, Req, K, S::Response>
155where
156    S: Service<Req>,
157    S::Response: Clone + Send + 'static,
158    K: Hash + Eq + Clone + Send + 'static,
159    Req: Send + 'static,
160    S::Future: Send + 'static,
161{
162    type Response = S::Response;
163    type Error = CacheError<S::Error>;
164    type Future = BoxFuture<'static, Result<Self::Response, Self::Error>>;
165
166    fn poll_ready(&mut self, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
167        self.inner.poll_ready(cx).map_err(CacheError::Inner)
168    }
169
170    fn call(&mut self, req: Req) -> Self::Future {
171        let key = (self.config.key_extractor)(&req);
172        let cache_name = self.config.name.clone();
173
174        // Check cache first
175        let cached = {
176            let mut store = self.store.lock().unwrap_or_else(|e| e.into_inner());
177            store.get(&key)
178        };
179
180        if let Some(response) = cached {
181            // Cache hit
182            #[cfg(feature = "metrics")]
183            {
184                counter!("cache_requests_total", "cache" => cache_name.clone(), "result" => "hit")
185                    .increment(1);
186            }
187
188            #[cfg(feature = "tracing")]
189            debug!(cache = %cache_name, "Cache hit");
190
191            let event = CacheEvent::Hit {
192                pattern_name: cache_name,
193                timestamp: Instant::now(),
194            };
195            self.config.event_listeners.emit(&event);
196            return Box::pin(async move { Ok(response) });
197        }
198
199        // Cache miss
200        #[cfg(feature = "metrics")]
201        {
202            counter!("cache_requests_total", "cache" => cache_name.clone(), "result" => "miss")
203                .increment(1);
204        }
205
206        #[cfg(feature = "tracing")]
207        debug!(cache = %cache_name, "Cache miss");
208
209        let miss_event = CacheEvent::Miss {
210            pattern_name: cache_name.clone(),
211            timestamp: Instant::now(),
212        };
213        self.config.event_listeners.emit(&miss_event);
214
215        let future = self.inner.call(req);
216        let store = Arc::clone(&self.store);
217        let config = Arc::clone(&self.config);
218
219        Box::pin(async move {
220            let response = future.await.map_err(CacheError::Inner)?;
221
222            // Store successful response in cache
223            let was_evicted = {
224                let mut store = store.lock().unwrap_or_else(|e| e.into_inner());
225                let was_full = store.len() >= config.max_size;
226                store.insert(key, response.clone());
227
228                // Update cache size gauge
229                #[cfg(feature = "metrics")]
230                {
231                    let new_size = store.len();
232                    gauge!("cache_size", "cache" => config.name.clone()).set(new_size as f64);
233                }
234
235                was_full
236            };
237
238            if was_evicted {
239                #[cfg(feature = "metrics")]
240                {
241                    counter!("cache_evictions_total", "cache" => config.name.clone()).increment(1);
242                }
243
244                #[cfg(feature = "tracing")]
245                info!(cache = %config.name, "Cache eviction occurred");
246
247                let event = CacheEvent::Eviction {
248                    pattern_name: config.name.clone(),
249                    timestamp: Instant::now(),
250                };
251                config.event_listeners.emit(&event);
252            }
253
254            Ok(response)
255        })
256    }
257}
258
259#[cfg(test)]
260mod tests {
261    use super::*;
262    use std::sync::atomic::{AtomicUsize, Ordering};
263    use std::time::Duration;
264    use tower::service_fn;
265    use tower::Layer;
266    use tower::ServiceExt;
267
268    #[tokio::test]
269    async fn cache_hit_returns_cached_response() {
270        let call_count = Arc::new(AtomicUsize::new(0));
271        let cc = Arc::clone(&call_count);
272
273        let service = service_fn(move |req: String| {
274            let cc = Arc::clone(&cc);
275            async move {
276                cc.fetch_add(1, Ordering::SeqCst);
277                Ok::<_, std::io::Error>(format!("Response: {}", req))
278            }
279        });
280
281        let layer = CacheLayer::builder()
282            .max_size(10)
283            .key_extractor(|req: &String| req.clone())
284            .build()
285            .unwrap();
286
287        let mut service = layer.layer(service);
288
289        // First call - cache miss
290        let response1 = service
291            .ready()
292            .await
293            .unwrap()
294            .call("test".to_string())
295            .await
296            .unwrap();
297        assert_eq!(response1, "Response: test");
298        assert_eq!(call_count.load(Ordering::SeqCst), 1);
299
300        // Second call - cache hit
301        let response2 = service
302            .ready()
303            .await
304            .unwrap()
305            .call("test".to_string())
306            .await
307            .unwrap();
308        assert_eq!(response2, "Response: test");
309        assert_eq!(call_count.load(Ordering::SeqCst), 1); // Not called again
310    }
311
312    #[tokio::test]
313    async fn cache_miss_calls_inner_service() {
314        let service = service_fn(|req: String| async move {
315            Ok::<_, std::io::Error>(format!("Response: {}", req))
316        });
317
318        let layer = CacheLayer::builder()
319            .max_size(10)
320            .key_extractor(|req: &String| req.clone())
321            .build()
322            .unwrap();
323
324        let mut service = layer.layer(service);
325
326        let response = service
327            .ready()
328            .await
329            .unwrap()
330            .call("test".to_string())
331            .await
332            .unwrap();
333        assert_eq!(response, "Response: test");
334    }
335
336    #[tokio::test]
337    async fn different_keys_not_cached_together() {
338        let call_count = Arc::new(AtomicUsize::new(0));
339        let cc = Arc::clone(&call_count);
340
341        let service = service_fn(move |req: String| {
342            let cc = Arc::clone(&cc);
343            async move {
344                cc.fetch_add(1, Ordering::SeqCst);
345                Ok::<_, std::io::Error>(format!("Response: {}", req))
346            }
347        });
348
349        let layer = CacheLayer::builder()
350            .max_size(10)
351            .key_extractor(|req: &String| req.clone())
352            .build()
353            .unwrap();
354
355        let mut service = layer.layer(service);
356
357        service
358            .ready()
359            .await
360            .unwrap()
361            .call("test1".to_string())
362            .await
363            .unwrap();
364        service
365            .ready()
366            .await
367            .unwrap()
368            .call("test2".to_string())
369            .await
370            .unwrap();
371
372        assert_eq!(call_count.load(Ordering::SeqCst), 2);
373    }
374
375    #[tokio::test]
376    async fn ttl_expiration_causes_cache_miss() {
377        let call_count = Arc::new(AtomicUsize::new(0));
378        let cc = Arc::clone(&call_count);
379
380        let service = service_fn(move |req: String| {
381            let cc = Arc::clone(&cc);
382            async move {
383                cc.fetch_add(1, Ordering::SeqCst);
384                Ok::<_, std::io::Error>(format!("Response: {}", req))
385            }
386        });
387
388        let layer = CacheLayer::builder()
389            .max_size(10)
390            .ttl(Duration::from_millis(50))
391            .key_extractor(|req: &String| req.clone())
392            .build()
393            .unwrap();
394
395        let mut service = layer.layer(service);
396
397        service
398            .ready()
399            .await
400            .unwrap()
401            .call("test".to_string())
402            .await
403            .unwrap();
404        assert_eq!(call_count.load(Ordering::SeqCst), 1);
405
406        // Wait for TTL to expire
407        tokio::time::sleep(Duration::from_millis(100)).await;
408
409        service
410            .ready()
411            .await
412            .unwrap()
413            .call("test".to_string())
414            .await
415            .unwrap();
416        assert_eq!(call_count.load(Ordering::SeqCst), 2); // Called again
417    }
418
419    #[tokio::test]
420    async fn lru_eviction_removes_least_recently_used() {
421        let service = service_fn(|req: String| async move {
422            Ok::<_, std::io::Error>(format!("Response: {}", req))
423        });
424
425        let layer = CacheLayer::builder()
426            .max_size(2)
427            .key_extractor(|req: &String| req.clone())
428            .build()
429            .unwrap();
430
431        let mut service = layer.layer(service);
432
433        // Fill cache with 2 items
434        service
435            .ready()
436            .await
437            .unwrap()
438            .call("key1".to_string())
439            .await
440            .unwrap();
441        service
442            .ready()
443            .await
444            .unwrap()
445            .call("key2".to_string())
446            .await
447            .unwrap();
448
449        // Add third item, should evict key1
450        service
451            .ready()
452            .await
453            .unwrap()
454            .call("key3".to_string())
455            .await
456            .unwrap();
457
458        // Verify cache state by checking call counts
459        let call_count = Arc::new(AtomicUsize::new(0));
460        let cc = Arc::clone(&call_count);
461
462        let service2 = service_fn(move |req: String| {
463            let cc = Arc::clone(&cc);
464            async move {
465                cc.fetch_add(1, Ordering::SeqCst);
466                Ok::<_, std::io::Error>(format!("Response: {}", req))
467            }
468        });
469
470        let mut service2 = layer.layer(service2);
471
472        // key1 should be evicted (cache miss)
473        service2
474            .ready()
475            .await
476            .unwrap()
477            .call("key1".to_string())
478            .await
479            .unwrap();
480        assert_eq!(call_count.load(Ordering::SeqCst), 1);
481    }
482
483    #[tokio::test]
484    async fn event_listeners_called() {
485        let hit_count = Arc::new(AtomicUsize::new(0));
486        let miss_count = Arc::new(AtomicUsize::new(0));
487        let eviction_count = Arc::new(AtomicUsize::new(0));
488
489        let hc = Arc::clone(&hit_count);
490        let mc = Arc::clone(&miss_count);
491        let ec = Arc::clone(&eviction_count);
492
493        let service = service_fn(|req: String| async move {
494            Ok::<_, std::io::Error>(format!("Response: {}", req))
495        });
496
497        let layer = CacheLayer::builder()
498            .max_size(1)
499            .key_extractor(|req: &String| req.clone())
500            .on_hit(move || {
501                hc.fetch_add(1, Ordering::SeqCst);
502            })
503            .on_miss(move || {
504                mc.fetch_add(1, Ordering::SeqCst);
505            })
506            .on_eviction(move || {
507                ec.fetch_add(1, Ordering::SeqCst);
508            })
509            .build()
510            .unwrap();
511
512        let mut service = layer.layer(service);
513
514        // First call - miss
515        service
516            .ready()
517            .await
518            .unwrap()
519            .call("test".to_string())
520            .await
521            .unwrap();
522        assert_eq!(miss_count.load(Ordering::SeqCst), 1);
523        assert_eq!(hit_count.load(Ordering::SeqCst), 0);
524
525        // Second call - hit
526        service
527            .ready()
528            .await
529            .unwrap()
530            .call("test".to_string())
531            .await
532            .unwrap();
533        assert_eq!(hit_count.load(Ordering::SeqCst), 1);
534        assert_eq!(miss_count.load(Ordering::SeqCst), 1);
535
536        // Third call with different key - eviction
537        service
538            .ready()
539            .await
540            .unwrap()
541            .call("other".to_string())
542            .await
543            .unwrap();
544        assert_eq!(eviction_count.load(Ordering::SeqCst), 1);
545    }
546
547    #[tokio::test]
548    async fn errors_not_cached() {
549        let call_count = Arc::new(AtomicUsize::new(0));
550        let cc = Arc::clone(&call_count);
551
552        let service = service_fn(move |_req: String| {
553            let cc = Arc::clone(&cc);
554            async move {
555                cc.fetch_add(1, Ordering::SeqCst);
556                Err::<String, _>(std::io::Error::other("error"))
557            }
558        });
559
560        let layer = CacheLayer::builder()
561            .max_size(10)
562            .key_extractor(|req: &String| req.clone())
563            .build()
564            .unwrap();
565
566        let mut service = layer.layer(service);
567
568        // First call - error
569        let _ = service
570            .ready()
571            .await
572            .unwrap()
573            .call("test".to_string())
574            .await;
575        assert_eq!(call_count.load(Ordering::SeqCst), 1);
576
577        // Second call - should call inner again (error not cached)
578        let _ = service
579            .ready()
580            .await
581            .unwrap()
582            .call("test".to_string())
583            .await;
584        assert_eq!(call_count.load(Ordering::SeqCst), 2);
585    }
586}