1use std::{marker::PhantomData, num::NonZeroUsize, sync::Arc, time::Duration};
4
5use saluki_common::{hash::FastBuildHasher, task::spawn_traced};
6use saluki_error::GenericError;
7use saluki_metrics::{static_metrics, Counter, Gauge, Histogram};
8use tokio::time::sleep;
9use tokio_util::sync::{CancellationToken, DropGuard};
10use tracing::debug;
11
12mod expiry;
13use self::expiry::{Expiration, ExpirationBuilder, ExpiryCapableLifecycle};
14
15pub mod weight;
16use self::weight::{ItemCountWeighter, Weighter, WrappedWeighter};
17
18type RawCache<K, V, W, H> = quick_cache::sync::Cache<K, V, WrappedWeighter<W>, H, ExpiryCapableLifecycle<K>>;
19
20#[static_metrics(prefix = cache, labels(cache_id))]
21#[derive(Clone)]
22struct Telemetry {
23 current_items: Gauge,
24 current_weight: Gauge,
25 weight_limit: Gauge,
26 hits_total: Counter,
27 misses_total: Counter,
28 items_evicted_total: Counter,
29 #[metric(level = debug)]
30 items_inserted_total: Counter,
31 #[metric(level = debug)]
32 items_removed_total: Counter,
33 #[metric(level = debug)]
34 items_expired_total: Counter,
35 #[metric(level = trace)]
36 items_expired_batch_size: Histogram,
37}
38
39struct InnerCache<K, V, W, H> {
40 cache: Arc<RawCache<K, V, W, H>>,
41 _task_shutdown_guard: DropGuard,
42}
43
44pub struct CacheBuilder<K, V, W = ItemCountWeighter, H = FastBuildHasher> {
46 identifier: String,
47 capacity: NonZeroUsize,
48 weighter: W,
49 idle_period: Option<Duration>,
50 expiration_interval: Option<Duration>,
51 telemetry_enabled: bool,
52 _key: PhantomData<K>,
53 _value: PhantomData<V>,
54 _hasher: PhantomData<H>,
55}
56
57impl<K, V> CacheBuilder<K, V> {
58 pub fn from_identifier<N: Into<String>>(identifier: N) -> Result<CacheBuilder<K, V>, GenericError> {
68 let identifier = identifier.into();
69 if identifier.is_empty() {
70 return Err(GenericError::msg("cache identifier must not be empty"));
71 }
72
73 Ok(CacheBuilder {
74 identifier,
75 capacity: NonZeroUsize::MAX,
76 weighter: ItemCountWeighter,
77 idle_period: None,
78 expiration_interval: None,
79 telemetry_enabled: true,
80 _key: PhantomData,
81 _value: PhantomData,
82 _hasher: PhantomData,
83 })
84 }
85
86 pub fn for_tests() -> CacheBuilder<K, V> {
97 CacheBuilder::from_identifier("noop")
98 .expect("cache identifier not empty")
99 .with_telemetry(false)
100 }
101}
102
103impl<K, V, W, H> CacheBuilder<K, V, W, H> {
104 pub fn with_capacity(mut self, capacity: NonZeroUsize) -> Self {
113 self.capacity = capacity;
114 self
115 }
116
117 pub fn with_time_to_idle(mut self, idle_period: Option<Duration>) -> Self {
127 self.idle_period = idle_period;
128
129 if self.idle_period.is_some() {
131 self.expiration_interval = self.expiration_interval.or(Some(Duration::from_secs(1)));
132 }
133
134 self
135 }
136
137 pub fn with_expiration_interval(mut self, expiration_interval: Duration) -> Self {
150 self.expiration_interval = Some(expiration_interval);
151 self
152 }
153
154 pub fn with_item_weighter<W2>(self, weighter: W2) -> CacheBuilder<K, V, W2, H> {
167 CacheBuilder {
168 identifier: self.identifier,
169 capacity: self.capacity,
170 weighter,
171 idle_period: self.idle_period,
172 expiration_interval: self.expiration_interval,
173 telemetry_enabled: self.telemetry_enabled,
174 _key: PhantomData,
175 _value: PhantomData,
176 _hasher: PhantomData,
177 }
178 }
179
180 pub fn with_hasher<H2>(self) -> CacheBuilder<K, V, W, H2> {
188 CacheBuilder {
189 identifier: self.identifier,
190 capacity: self.capacity,
191 weighter: self.weighter,
192 idle_period: self.idle_period,
193 expiration_interval: self.expiration_interval,
194 telemetry_enabled: self.telemetry_enabled,
195 _key: PhantomData,
196 _value: PhantomData,
197 _hasher: PhantomData,
198 }
199 }
200
201 pub fn with_telemetry(mut self, enabled: bool) -> Self {
210 self.telemetry_enabled = enabled;
211 self
212 }
213}
214
215impl<K, V, W, H> CacheBuilder<K, V, W, H>
216where
217 K: Eq + std::hash::Hash + Clone + Send + Sync + 'static,
218 V: Clone + Send + Sync + 'static,
219 W: Weighter<K, V> + Clone + Send + Sync + 'static,
220 H: std::hash::BuildHasher + Clone + Default + Send + Sync + 'static,
221{
222 pub fn build(self) -> Cache<K, V, W, H> {
224 let capacity = self.capacity.get();
225
226 let telemetry = Telemetry::new(self.identifier);
227 telemetry.weight_limit().set(capacity as f64);
228
229 let eviction_counter = if self.telemetry_enabled {
231 telemetry.items_evicted_total().clone()
232 } else {
233 Counter::noop()
234 };
235 let mut expiration_builder = ExpirationBuilder::new(eviction_counter);
236 if let Some(time_to_idle) = self.idle_period {
237 expiration_builder = expiration_builder.with_time_to_idle(time_to_idle);
238 }
239 let (expiration, expiry_lifecycle) = expiration_builder.build();
240
241 let shutdown_token = CancellationToken::new();
243 let raw_cache = Arc::new(RawCache::with(
244 capacity,
245 capacity as u64,
246 WrappedWeighter::from(self.weighter),
247 H::default(),
248 expiry_lifecycle,
249 ));
250
251 let cache = Cache {
252 inner: Arc::new(InnerCache {
253 cache: Arc::clone(&raw_cache),
254 _task_shutdown_guard: shutdown_token.clone().drop_guard(),
255 }),
256 expiration: expiration.clone(),
257 telemetry: telemetry.clone(),
258 };
259
260 if let Some(expiration_interval) = self.expiration_interval {
262 let expiration = expiration.clone();
263
264 spawn_traced(drive_expiration(
265 Arc::clone(&raw_cache),
266 telemetry.clone(),
267 expiration,
268 expiration_interval,
269 shutdown_token.clone(),
270 ));
271 }
272
273 if self.telemetry_enabled {
275 spawn_traced(drive_telemetry(Arc::clone(&raw_cache), telemetry, shutdown_token));
276 }
277
278 cache
279 }
280}
281
282#[derive(Clone)]
284pub struct Cache<K, V, W = ItemCountWeighter, H = FastBuildHasher> {
285 inner: Arc<InnerCache<K, V, W, H>>,
286 expiration: Expiration<K>,
287 telemetry: Telemetry,
288}
289
290impl<K, V, W, H> Cache<K, V, W, H>
291where
292 K: Eq + std::hash::Hash + Clone,
293 V: Clone,
294 W: Weighter<K, V> + Clone,
295 H: std::hash::BuildHasher + Clone,
296{
297 pub fn is_empty(&self) -> bool {
299 self.inner.cache.is_empty()
300 }
301
302 pub fn len(&self) -> usize {
304 self.inner.cache.len()
305 }
306
307 pub fn weight(&self) -> u64 {
309 self.inner.cache.weight()
310 }
311
312 pub fn insert(&self, key: K, value: V) {
318 self.inner.cache.insert(key.clone(), value);
319 self.expiration.mark_entry_accessed(key);
320 self.telemetry.items_inserted_total().increment(1);
321 }
322
323 pub fn get(&self, key: &K) -> Option<V> {
327 let value = self.inner.cache.get(key);
328 if value.is_some() {
329 self.expiration.mark_entry_accessed(key.clone());
330 self.telemetry.hits_total().increment(1);
331 } else {
332 self.telemetry.misses_total().increment(1);
333 }
334 value
335 }
336
337 pub fn remove(&self, key: &K) {
339 self.inner.cache.remove(key);
340 self.expiration.mark_entry_removed(key.clone());
341 self.telemetry.items_removed_total().increment(1);
342 }
343}
344
345async fn drive_expiration<K, V, W, H>(
346 cache: Arc<RawCache<K, V, W, H>>, telemetry: Telemetry, expiration: Expiration<K>, expiration_interval: Duration,
347 shutdown: CancellationToken,
348) where
349 K: Eq + std::hash::Hash + Clone,
350 V: Clone,
351 W: Weighter<K, V> + Clone,
352 H: std::hash::BuildHasher + Clone,
353{
354 let mut expired_item_keys = Vec::new();
355
356 loop {
357 tokio::select! {
358 _ = shutdown.cancelled() => break,
359 _ = sleep(expiration_interval) => {}
360 }
361
362 expiration.drain_expired_items(&mut expired_item_keys);
364
365 let num_expired_items = expired_item_keys.len();
366 if num_expired_items != 0 {
367 telemetry.items_expired_total().increment(num_expired_items as u64);
368 telemetry.items_expired_batch_size().record(num_expired_items as f64);
369 }
370
371 debug!(num_expired_items, "Found expired items.");
372
373 for item_key in expired_item_keys.drain(..) {
374 cache.remove(&item_key);
375 telemetry.items_removed_total().increment(1);
376 expiration.mark_entry_removed(item_key);
377 }
378
379 debug!(num_expired_items, "Removed expired items.");
380 }
381}
382
383async fn drive_telemetry<K, V, W, H>(
384 cache: Arc<RawCache<K, V, W, H>>, telemetry: Telemetry, shutdown: CancellationToken,
385) where
386 K: Eq + std::hash::Hash + Clone,
387 V: Clone,
388 W: Weighter<K, V> + Clone,
389 H: std::hash::BuildHasher + Clone,
390{
391 loop {
392 tokio::select! {
393 _ = shutdown.cancelled() => break,
394 _ = sleep(Duration::from_secs(1)) => {}
395 }
396
397 telemetry.current_items().set(cache.len() as f64);
398 telemetry.current_weight().set(cache.weight() as f64);
399 }
400}
401
402#[cfg(test)]
403mod tests {
404 use super::*;
405
406 #[derive(Clone)]
407 pub struct ItemValueWeighter;
408
409 impl<K> Weighter<K, usize> for ItemValueWeighter {
410 fn item_weight(&self, _key: &K, value: &usize) -> u64 {
411 *value as u64
412 }
413 }
414
415 #[test]
416 fn empty_cache_identifier() {
417 let result = CacheBuilder::<u64, u64>::from_identifier("");
418 assert!(result.is_err(), "expected error for empty cache identifier");
419 }
420
421 #[test]
422 fn basic() {
423 const CACHE_KEY: usize = 42;
424 const CACHE_VALUE: &str = "value1";
425
426 let cache = CacheBuilder::for_tests().build();
427
428 assert_eq!(cache.len(), 0);
429 assert_eq!(cache.weight(), 0);
430
431 cache.insert(CACHE_KEY, CACHE_VALUE);
432 assert_eq!(cache.len(), 1);
433 assert_eq!(cache.weight(), 1);
434
435 assert_eq!(cache.get(&CACHE_KEY), Some(CACHE_VALUE));
436
437 cache.remove(&CACHE_KEY);
438 assert_eq!(cache.len(), 0);
439 assert_eq!(cache.weight(), 0);
440 }
441
442 #[test]
443 fn evict_at_capacity() {
444 const CAPACITY: usize = 3;
445
446 let cache = CacheBuilder::for_tests()
447 .with_capacity(NonZeroUsize::new(CAPACITY).unwrap())
448 .build();
449
450 for i in 0..CAPACITY {
452 cache.insert(i, "value");
453 }
454
455 assert_eq!(cache.len(), CAPACITY);
456 assert_eq!(cache.weight(), CAPACITY as u64);
457
458 cache.insert(CAPACITY, "new_value");
461 assert_eq!(cache.len(), CAPACITY);
462 assert_eq!(cache.weight(), CAPACITY as u64);
463
464 let mut evicted = false;
465 for i in 0..CAPACITY {
466 if cache.get(&i).is_none() {
467 evicted = true;
468 break;
469 }
470 }
471 assert!(evicted, "expected at least one original item to be evicted");
472 }
473
474 #[test]
475 fn overweight_item() {
476 const CAPACITY: usize = 10;
477
478 let cache = CacheBuilder::for_tests()
480 .with_capacity(NonZeroUsize::new(CAPACITY).unwrap())
481 .with_item_weighter(ItemValueWeighter)
482 .build();
483
484 assert_eq!(cache.len(), 0);
486 assert_eq!(cache.weight(), 0);
487
488 cache.insert(1, CAPACITY + 1);
489 assert_eq!(cache.len(), 0);
490 assert_eq!(cache.weight(), 0);
491 assert_eq!(cache.get(&1), None);
492 }
493
494 #[test]
495 fn evict_on_insert_by_weight() {
496 const CAPACITY: usize = 10;
497
498 let cache = CacheBuilder::for_tests()
500 .with_capacity(NonZeroUsize::new(CAPACITY).unwrap())
501 .with_item_weighter(ItemValueWeighter)
502 .build();
503
504 assert_eq!(cache.len(), 0);
506 assert_eq!(cache.weight(), 0);
507
508 cache.insert(1, 3);
509 cache.insert(2, 4);
510 cache.insert(3, 3);
511 assert_eq!(cache.len(), 3);
512 assert_eq!(cache.weight(), CAPACITY as u64);
513
514 cache.insert(4, CAPACITY - 1);
517 assert_eq!(cache.len(), 1);
518 assert_eq!(cache.weight(), (CAPACITY - 1) as u64);
519
520 assert_eq!(cache.get(&1), None);
521 assert_eq!(cache.get(&2), None);
522 assert_eq!(cache.get(&3), None);
523 assert_eq!(cache.get(&4), Some(CAPACITY - 1));
524 }
525
526 #[tokio::test]
527 async fn tasks_stop_when_cache_dropped() {
528 let cache = CacheBuilder::<u64, u64>::from_identifier("test-drop")
529 .expect("valid identifier")
530 .with_time_to_idle(Some(Duration::from_secs(60)))
531 .with_expiration_interval(Duration::from_millis(50))
532 .build();
533
534 let weak_cache = Arc::downgrade(&cache.inner.cache);
536
537 drop(cache);
538
539 sleep(Duration::from_millis(100)).await;
546
547 assert!(
548 weak_cache.upgrade().is_none(),
549 "raw cache should be released after background tasks exit"
550 );
551 }
552}