1mod 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
71pub 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 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 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 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 #[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 #[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 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 #[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 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 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); }
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 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); }
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 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 service
451 .ready()
452 .await
453 .unwrap()
454 .call("key3".to_string())
455 .await
456 .unwrap();
457
458 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 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 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 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 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 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 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}