1use crate::metrics::*;
2use dashmap::DashMap;
3use kumo_prometheus::prometheus::{IntCounter, IntGauge};
4use kumo_server_memory::subscribe_to_memory_status_changes_async;
5pub use linkme;
6use parking_lot::Mutex;
7pub use pastey as paste;
8use scopeguard::defer;
9use std::borrow::Borrow;
10use std::collections::HashMap;
11use std::fmt::Debug;
12use std::future::Future;
13use std::hash::Hash;
14use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering};
15use std::sync::{Arc, LazyLock, Weak};
16use tokio::sync::{OwnedSemaphorePermit, Semaphore};
17use tokio::time::{timeout, timeout_at, Duration, Instant};
18
19mod metrics;
20
21static CACHES: LazyLock<Mutex<Vec<Weak<dyn CachePurger + Send + Sync>>>> =
22 LazyLock::new(Mutex::default);
23
24struct Inner<K: Clone + Hash + Eq + Debug, V: Clone + Send + Sync + Debug> {
25 name: String,
26 tick: AtomicUsize,
27 capacity: AtomicUsize,
28 allow_stale_reads: AtomicBool,
29 cache: DashMap<K, Item<V>>,
30 lru_samples: AtomicUsize,
31 sema_timeout_milliseconds: AtomicUsize,
32 lookup_counter: IntCounter,
33 evict_counter: IntCounter,
34 expire_counter: IntCounter,
35 hit_counter: IntCounter,
36 miss_counter: IntCounter,
37 populate_counter: IntCounter,
38 insert_counter: IntCounter,
39 stale_counter: IntCounter,
40 error_counter: IntCounter,
41 wait_gauge: IntGauge,
42 size_gauge: IntGauge,
43}
44
45impl<
46 K: Clone + Debug + Send + Sync + Hash + Eq + 'static,
47 V: Clone + Debug + Send + Sync + 'static,
48 > Inner<K, V>
49{
50 pub fn clear(&self) -> usize {
51 let num_entries = self.cache.len();
52
53 self.cache.retain(|_k, item| {
58 if let ItemState::Pending(sema) = &item.item {
59 sema.close();
61 }
62 false
63 });
64
65 self.size_gauge.set(self.cache.len() as i64);
66 num_entries
67 }
68
69 pub fn evict_some(&self, target: usize) -> usize {
88 let now = Instant::now();
89
90 let cache_size = self.cache.len();
93 let num_samples = self.lru_samples.load(Ordering::Relaxed).min(cache_size);
95
96 let mut expired_keys = vec![];
98 let mut samples = vec![];
100
101 {
116 let mut rng = rand::thread_rng();
117 let mut indices =
118 rand::seq::index::sample(&mut rng, cache_size, num_samples).into_vec();
119 indices.sort();
120 let mut iter = self.cache.iter();
121 let mut current_idx = 0;
122
123 fn advance_by(iter: &mut impl Iterator, skip_amount: usize) {
128 for _ in 0..skip_amount {
129 if iter.next().is_none() {
130 return;
131 }
132 }
133 }
134
135 for idx in indices {
136 let skip_amount = idx - current_idx;
140 advance_by(&mut iter, skip_amount);
141
142 match iter.next() {
143 Some(map_entry) => {
144 current_idx = idx + 1;
145 let item = map_entry.value();
146 match &item.item {
147 ItemState::Pending(_) | ItemState::Refreshing { .. } => {
148 }
150 ItemState::Present(_) | ItemState::Failed(_) => {
151 if now >= item.expiration {
152 expired_keys.push(map_entry.key().clone());
153 } else {
154 let last_tick = item.last_tick.load(Ordering::Relaxed);
155 samples.push((map_entry.key().clone(), last_tick));
156 }
157 }
158 }
159 }
160 None => {
161 break;
162 }
163 }
164 }
165 }
166
167 let mut num_removed = 0;
168 for key in expired_keys {
169 let removed = self
173 .cache
174 .remove_if(&key, |_k, entry| now >= entry.expiration)
175 .is_some();
176 if removed {
177 tracing::trace!("{} expired {key:?}", self.name);
178 num_removed += 1;
179 self.expire_counter.inc();
180 }
181 }
182
183 let target = target.min(samples.len() / 2).max(1);
190
191 if num_removed >= target {
193 self.size_gauge.set(self.cache.len() as i64);
194 tracing::trace!("{} expired {num_removed} of target {target}", self.name);
195 return num_removed;
196 }
197
198 samples.sort_by(|(_ka, tick_a), (_kb, tick_b)| tick_a.cmp(tick_b));
201
202 for (key, tick) in samples {
203 if self
207 .cache
208 .remove_if(&key, |_k, item| {
209 item.last_tick.load(Ordering::Relaxed) == tick
210 })
211 .is_some()
212 {
213 tracing::debug!("{} evicted {key:?}", self.name);
214 num_removed += 1;
215 self.evict_counter.inc();
216 self.size_gauge.set(self.cache.len() as i64);
217 if num_removed >= target {
218 return num_removed;
219 }
220 }
221 }
222
223 if num_removed == 0 {
224 tracing::debug!(
225 "{} did not find anything to evict, target was {target}",
226 self.name
227 );
228 }
229
230 tracing::trace!("{} removed {num_removed} of target {target}", self.name);
231
232 num_removed
233 }
234
235 pub fn maybe_evict(&self) -> usize {
238 let cache_size = self.cache.len();
239 let capacity = self.capacity.load(Ordering::Relaxed);
240 if cache_size > capacity {
241 self.evict_some(cache_size - capacity)
242 } else {
243 0
244 }
245 }
246}
247
248trait CachePurger {
249 fn name(&self) -> &str;
250 fn purge(&self) -> usize;
251 fn process_expirations(&self) -> usize;
252 fn update_capacity(&self, capacity: usize);
253}
254
255impl<
256 K: Clone + Debug + Send + Sync + Hash + Eq + 'static,
257 V: Clone + Debug + Send + Sync + 'static,
258 > CachePurger for Inner<K, V>
259{
260 fn name(&self) -> &str {
261 &self.name
262 }
263 fn purge(&self) -> usize {
264 self.clear()
265 }
266 fn process_expirations(&self) -> usize {
267 let now = Instant::now();
268 let mut expired_keys = vec![];
269 for map_entry in self.cache.iter() {
270 let item = map_entry.value();
271 match &item.item {
272 ItemState::Pending(_) | ItemState::Refreshing { .. } => {
273 }
275 ItemState::Present(_) | ItemState::Failed(_) => {
276 if now >= item.expiration {
277 expired_keys.push(map_entry.key().clone());
278 }
279 }
280 }
281 }
282
283 let mut num_removed = 0;
284 for key in expired_keys {
285 let removed = self
289 .cache
290 .remove_if(&key, |_k, entry| now >= entry.expiration)
291 .is_some();
292 if removed {
293 num_removed += 1;
294 self.expire_counter.inc();
295 self.size_gauge.set(self.cache.len() as i64);
296 }
297 }
298
299 num_removed + self.maybe_evict()
300 }
301
302 fn update_capacity(&self, capacity: usize) {
303 self.capacity.store(capacity, Ordering::Relaxed);
304 self.process_expirations();
309 }
310}
311
312fn all_caches() -> Vec<Arc<dyn CachePurger + Send + Sync>> {
313 let mut result = vec![];
314 let mut caches = CACHES.lock();
315 caches.retain(|entry| match entry.upgrade() {
316 Some(purger) => {
317 result.push(purger);
318 true
319 }
320 None => false,
321 });
322 result
323}
324
325pub fn purge_all_caches() {
326 let purgers = all_caches();
327
328 tracing::error!("purging {} caches", purgers.len());
329 for purger in purgers {
330 let name = purger.name();
331 let num_entries = purger.purge();
332 tracing::error!("cleared {num_entries} entries from cache {name}");
333 }
334}
335
336async fn prune_expired_caches() {
337 loop {
338 tokio::time::sleep(tokio::time::Duration::from_secs(30)).await;
339 let purgers = all_caches();
340
341 for p in purgers {
342 let n = p.process_expirations();
343 if n > 0 {
344 tracing::debug!("expired {n} entries from cache {}", p.name());
345 }
346 }
347 }
348}
349
350#[linkme::distributed_slice]
351pub static LRUTTL_VIVIFY: [fn() -> CacheDefinition];
352
353#[macro_export]
354macro_rules! optional_doc {
355 ($doc:expr) => {
356 Some($doc.trim())
357 };
358 ($($doc:expr)+) => {
359 Some(concat!($($doc,)+).trim())
360 };
361 () => {
362 None
363 };
364}
365
366#[macro_export]
373macro_rules! declare_cache {
374 (
375 $(#[doc = $doc:expr])*
376 $vis:vis
377 static $sym:ident:
378 LruCacheWithTtl<$key:ty, $value:ty>::new($name:expr, $capacity:expr);
379 ) => {
380 $(#[doc = $doc])*
381 $vis static $sym: ::std::sync::LazyLock<$crate::LruCacheWithTtl<$key, $value>> =
382 ::std::sync::LazyLock::new(
383 || $crate::LruCacheWithTtl::new($name, $capacity));
384
385 $crate::paste::paste! {
387 #[linkme::distributed_slice($crate::LRUTTL_VIVIFY)]
388 static [<VIVIFY_ $sym>]: fn() -> $crate::CacheDefinition = || {
389 ::std::sync::LazyLock::force(&$sym);
390 $crate::CacheDefinition {
391 name: $name,
392 capacity: $capacity,
393 doc: $crate::optional_doc!($($doc)*),
394 }
395 };
396 }
397 };
398}
399
400fn vivify() {
403 LazyLock::force(&PREDEFINED_CACHES);
404}
405
406fn vivify_impl() -> HashMap<&'static str, CacheDefinition> {
407 let mut map = HashMap::new();
408
409 for vivify_func in LRUTTL_VIVIFY {
410 let definition = vivify_func();
411 assert!(
412 !map.contains_key(definition.name),
413 "duplicate cache name {}",
414 definition.name
415 );
416 map.insert(definition.name, definition);
417 }
418
419 map
420}
421
422#[derive(serde::Serialize)]
423pub struct CacheDefinition {
424 pub name: &'static str,
425 pub capacity: usize,
426 pub doc: Option<&'static str>,
427}
428
429static PREDEFINED_CACHES: LazyLock<HashMap<&'static str, CacheDefinition>> =
430 LazyLock::new(vivify_impl);
431
432pub fn get_definitions() -> Vec<&'static CacheDefinition> {
433 let mut defs = PREDEFINED_CACHES.values().collect::<Vec<_>>();
434 defs.sort_by(|a, b| a.name.cmp(&b.name));
435 defs
436}
437
438pub fn is_name_available(name: &str) -> bool {
439 !PREDEFINED_CACHES.contains_key(name)
440}
441
442pub fn set_cache_capacity(name: &str, capacity: usize) -> bool {
444 if !PREDEFINED_CACHES.contains_key(name) {
445 return false;
446 }
447 let caches = all_caches();
448 match caches.iter().find(|p| p.name() == name) {
449 Some(p) => {
450 p.update_capacity(capacity);
451 true
452 }
453 None => false,
454 }
455}
456
457pub fn spawn_memory_monitor() {
458 vivify();
459 tokio::spawn(purge_caches_on_memory_shortage());
460 tokio::spawn(prune_expired_caches());
461}
462
463async fn purge_caches_on_memory_shortage() {
464 tracing::debug!("starting memory monitor");
465 let mut memory_status = subscribe_to_memory_status_changes_async().await;
466 while let Ok(()) = memory_status.changed().await {
467 if kumo_server_memory::get_headroom() == 0 {
468 purge_all_caches();
469
470 tokio::time::sleep(tokio::time::Duration::from_secs(30)).await;
474 }
475 }
476}
477
478#[derive(Debug, Clone)]
479enum ItemState<V>
480where
481 V: Send,
482 V: Sync,
483{
484 Present(V),
485 Pending(Arc<Semaphore>),
486 Failed(Arc<anyhow::Error>),
487 Refreshing {
488 stale_value: V,
489 pending: Arc<Semaphore>,
490 },
491}
492
493#[derive(Debug)]
494struct Item<V>
495where
496 V: Send,
497 V: Sync,
498{
499 item: ItemState<V>,
500 expiration: Instant,
501 last_tick: AtomicUsize,
502}
503
504impl<V: Clone + Send + Sync> Clone for Item<V> {
505 fn clone(&self) -> Self {
506 Self {
507 item: self.item.clone(),
508 expiration: self.expiration,
509 last_tick: self.last_tick.load(Ordering::Relaxed).into(),
510 }
511 }
512}
513
514#[derive(Debug)]
515pub struct ItemLookup<V: Debug> {
516 pub item: V,
518 pub is_fresh: bool,
521 pub expiration: Instant,
523}
524
525pub struct LruCacheWithTtl<K: Clone + Debug + Hash + Eq, V: Clone + Debug + Send + Sync> {
526 inner: Arc<Inner<K, V>>,
527}
528
529impl<K: Clone + Debug + Hash + Eq, V: Clone + Debug + Send + Sync> Clone for LruCacheWithTtl<K, V> {
530 fn clone(&self) -> Self {
531 Self {
532 inner: Arc::clone(&self.inner),
533 }
534 }
535}
536
537enum Acquisition<V: Debug> {
539 Resolved(Result<ItemLookup<V>, Arc<anyhow::Error>>),
542 Owner(OwnedSemaphorePermit),
546}
547
548impl<
549 K: Clone + Debug + Hash + Eq + Send + Sync + std::fmt::Debug + 'static,
550 V: Clone + Debug + Send + Sync + 'static,
551 > LruCacheWithTtl<K, V>
552{
553 pub fn new<S: Into<String>>(name: S, capacity: usize) -> Self {
554 let name = name.into();
555 let cache = DashMap::new();
556
557 let lookup_counter = CACHE_LOOKUP
558 .get_metric_with_label_values(&[&name])
559 .expect("failed to get counter");
560 let hit_counter = CACHE_HIT
561 .get_metric_with_label_values(&[&name])
562 .expect("failed to get counter");
563 let stale_counter = CACHE_STALE
564 .get_metric_with_label_values(&[&name])
565 .expect("failed to get counter");
566 let evict_counter = CACHE_EVICT
567 .get_metric_with_label_values(&[&name])
568 .expect("failed to get counter");
569 let expire_counter = CACHE_EXPIRE
570 .get_metric_with_label_values(&[&name])
571 .expect("failed to get counter");
572 let miss_counter = CACHE_MISS
573 .get_metric_with_label_values(&[&name])
574 .expect("failed to get counter");
575 let populate_counter = CACHE_POPULATED
576 .get_metric_with_label_values(&[&name])
577 .expect("failed to get counter");
578 let insert_counter = CACHE_INSERT
579 .get_metric_with_label_values(&[&name])
580 .expect("failed to get counter");
581 let error_counter = CACHE_ERROR
582 .get_metric_with_label_values(&[&name])
583 .expect("failed to get counter");
584 let wait_gauge = CACHE_WAIT
585 .get_metric_with_label_values(&[&name])
586 .expect("failed to get counter");
587 let size_gauge = CACHE_SIZE
588 .get_metric_with_label_values(&[&name])
589 .expect("failed to get counter");
590
591 let inner = Arc::new(Inner {
592 name,
593 cache,
594 tick: AtomicUsize::new(0),
595 allow_stale_reads: AtomicBool::new(false),
596 capacity: AtomicUsize::new(capacity),
597 lru_samples: AtomicUsize::new(10),
598 sema_timeout_milliseconds: AtomicUsize::new(120_000),
599 lookup_counter,
600 evict_counter,
601 expire_counter,
602 hit_counter,
603 miss_counter,
604 populate_counter,
605 error_counter,
606 wait_gauge,
607 insert_counter,
608 stale_counter,
609 size_gauge,
610 });
611
612 {
616 let generic: Arc<dyn CachePurger + Send + Sync> = inner.clone();
617 CACHES.lock().push(Arc::downgrade(&generic));
618 tracing::debug!(
619 "registered cache {} with capacity {capacity}",
620 generic.name()
621 );
622 }
623
624 Self { inner }
625 }
626
627 fn allow_stale_reads(&self) -> bool {
628 self.inner.allow_stale_reads.load(Ordering::Relaxed)
629 }
630
631 pub fn set_allow_stale_reads(&self, value: bool) {
632 self.inner.allow_stale_reads.store(value, Ordering::Relaxed);
633 }
634
635 pub fn set_sema_timeout(&self, duration: Duration) {
636 self.inner
637 .sema_timeout_milliseconds
638 .store(duration.as_millis() as usize, Ordering::Relaxed);
639 }
640
641 pub fn get_sema_timeout(&self) -> Duration {
642 Duration::from_millis(self.inner.sema_timeout_milliseconds.load(Ordering::Relaxed) as u64)
643 }
644
645 pub fn clear(&self) -> usize {
646 self.inner.clear()
647 }
648
649 fn inc_tick(&self) -> usize {
650 self.inner.tick.fetch_add(1, Ordering::Relaxed) + 1
651 }
652
653 fn update_tick(&self, item: &Item<V>) {
654 let v = self.inc_tick();
655 item.last_tick.store(v, Ordering::Relaxed);
656 }
657
658 pub fn lookup<Q: ?Sized>(&self, name: &Q) -> Option<ItemLookup<V>>
659 where
660 K: Borrow<Q>,
661 Q: Hash + Eq,
662 {
663 self.inner.lookup_counter.inc();
664 match self.inner.cache.get_mut(name) {
665 None => {
666 self.inner.miss_counter.inc();
667 None
668 }
669 Some(entry) => {
670 match &entry.item {
671 ItemState::Present(item) => {
672 let now = Instant::now();
673 if now >= entry.expiration {
674 if self.allow_stale_reads() {
676 self.inner.miss_counter.inc();
683 return None;
684 }
685
686 drop(entry);
690 if self
691 .inner
692 .cache
693 .remove_if(name, |_k, entry| now >= entry.expiration)
694 .is_some()
695 {
696 self.inner.expire_counter.inc();
697 self.inner.size_gauge.set(self.inner.cache.len() as i64);
698 }
699 self.inner.miss_counter.inc();
700 return None;
701 }
702 self.inner.hit_counter.inc();
703 self.update_tick(&entry);
704 Some(ItemLookup {
705 item: item.clone(),
706 expiration: entry.expiration,
707 is_fresh: false,
708 })
709 }
710 ItemState::Refreshing { .. } | ItemState::Pending(_) | ItemState::Failed(_) => {
711 self.inner.miss_counter.inc();
712 None
713 }
714 }
715 }
716 }
717 }
718
719 pub fn get<Q: ?Sized>(&self, name: &Q) -> Option<V>
720 where
721 K: Borrow<Q>,
722 Q: Hash + Eq,
723 {
724 self.lookup(name).map(|lookup| lookup.item)
725 }
726
727 pub async fn insert(&self, name: K, item: V, expiration: Instant) -> V {
728 self.inner.cache.insert(
729 name,
730 Item {
731 item: ItemState::Present(item.clone()),
732 expiration,
733 last_tick: self.inc_tick().into(),
734 },
735 );
736
737 self.inner.insert_counter.inc();
738 self.inner.size_gauge.set(self.inner.cache.len() as i64);
739 self.inner.maybe_evict();
740
741 item
742 }
743
744 fn clone_item_state(
745 &self,
746 name: &K,
747 deadline: Instant,
748 timeout_duration: Duration,
749 ) -> (ItemState<V>, Instant) {
750 let mut is_new = false;
751 let mut entry = self.inner.cache.entry(name.clone()).or_insert_with(|| {
752 is_new = true;
753 Item {
754 item: ItemState::Pending(Arc::new(Semaphore::new(1))),
755 expiration: deadline,
756 last_tick: self.inc_tick().into(),
757 }
758 });
759
760 match &entry.value().item {
761 ItemState::Pending(sema) => {
762 if sema.is_closed() {
763 entry.value_mut().item = ItemState::Pending(Arc::new(Semaphore::new(1)));
764 } else {
765 let now = Instant::now();
766 if now >= entry.expiration {
767 tracing::warn!(
772 "{} semaphore for {name:?} remains open, \
773 but the lookup state has not been satisfied within the \
774 populate deadline. Assuming that something is stuck and \
775 making the now-active caller responsible for populating \
776 this entry",
777 self.inner.name
778 );
779 entry.value_mut().item = ItemState::Pending(Arc::new(Semaphore::new(1)));
780 entry.value_mut().expiration = now + timeout_duration;
781 }
782 }
783 }
784 ItemState::Refreshing {
785 stale_value,
786 pending,
787 } => {
788 if pending.is_closed() {
789 entry.value_mut().item = ItemState::Refreshing {
790 stale_value: stale_value.clone(),
791 pending: Arc::new(Semaphore::new(1)),
792 };
793 }
794 }
795 ItemState::Present(item) => {
796 let now = Instant::now();
797 if now >= entry.expiration {
798 let pending = Arc::new(Semaphore::new(1));
800 if self.allow_stale_reads() {
801 entry.value_mut().item = ItemState::Refreshing {
802 stale_value: item.clone(),
803 pending,
804 };
805 } else {
806 entry.value_mut().item = ItemState::Pending(pending);
807 }
808 }
809 }
810 ItemState::Failed(_) => {
811 let now = Instant::now();
812 if now >= entry.expiration {
813 entry.value_mut().item = ItemState::Pending(Arc::new(Semaphore::new(1)));
815 entry.value_mut().expiration = now + timeout_duration;
816 }
817 }
818 }
819
820 self.update_tick(&entry);
821 let item = entry.value();
822 let result = (item.item.clone(), entry.expiration);
823 drop(entry);
824
825 if is_new {
826 self.inner.size_gauge.set(self.inner.cache.len() as i64);
827 self.inner.maybe_evict();
828 }
829
830 result
831 }
832
833 async fn acquire(
837 &self,
838 name: &K,
839 deadline: Instant,
840 timeout_duration: Duration,
841 ) -> Acquisition<V> {
842 loop {
845 let (stale_value, sema) = match self.clone_item_state(name, deadline, timeout_duration)
846 {
847 (ItemState::Present(item), expiration) => {
848 return Acquisition::Resolved(Ok(ItemLookup {
849 item,
850 expiration,
851 is_fresh: false,
852 }));
853 }
854 (ItemState::Failed(error), _) => {
855 return Acquisition::Resolved(Err(error));
856 }
857 (
858 ItemState::Refreshing {
859 stale_value,
860 pending,
861 },
862 expiration,
863 ) => (Some((stale_value, expiration)), pending),
864 (ItemState::Pending(sema), _) => (None, sema),
865 };
866
867 let wait_result = {
868 self.inner.wait_gauge.inc();
869 defer! {
870 self.inner.wait_gauge.dec();
871 }
872
873 match timeout_at(deadline, sema.acquire_owned()).await {
877 Err(_) => {
878 if let Some((item, expiration)) = stale_value {
879 tracing::debug!(
880 "{} semaphore acquire for {name:?} timed out after \
881 {timeout_duration:?}, allowing stale value to satisfy the lookup",
882 self.inner.name
883 );
884 self.inner.stale_counter.inc();
885 return Acquisition::Resolved(Ok(ItemLookup {
886 item,
887 expiration,
888 is_fresh: false,
889 }));
890 }
891 tracing::debug!(
892 "{} semaphore acquire for {name:?} timed out after \
893 {timeout_duration:?}, returning error",
894 self.inner.name
895 );
896
897 self.inner.error_counter.inc();
898 return Acquisition::Resolved(Err(Arc::new(anyhow::anyhow!(
899 "{} lookup for {name:?} \
900 timed out after {timeout_duration:?} \
901 on semaphore acquire while waiting for cache to populate",
902 self.inner.name
903 ))));
904 }
905 Ok(r) => r,
906 }
907 };
908
909 let current_sema = match self.clone_item_state(name, deadline, timeout_duration) {
912 (ItemState::Present(item), expiration) => {
913 return Acquisition::Resolved(Ok(ItemLookup {
914 item,
915 expiration,
916 is_fresh: false,
917 }));
918 }
919 (ItemState::Failed(error), _) => {
920 self.inner.hit_counter.inc();
921 return Acquisition::Resolved(Err(error));
922 }
923 (
924 ItemState::Refreshing {
925 stale_value: _,
926 pending,
927 },
928 _,
929 ) => pending,
930 (ItemState::Pending(current_sema), _) => current_sema,
931 };
932
933 match wait_result {
935 Ok(permit) => {
936 if !Arc::ptr_eq(¤t_sema, permit.semaphore()) {
937 self.inner.error_counter.inc();
938 tracing::warn!(
939 "{} mismatched semaphores for {name:?}, \
940 will restart cache resolve.",
941 self.inner.name
942 );
943 permit.semaphore().close();
947 continue;
949 }
950
951 return Acquisition::Owner(permit);
954 }
955 Err(_) => {
956 self.inner.error_counter.inc();
957
958 tracing::debug!(
961 "{} lookup for {name:?} woke up semaphores \
962 but is still marked pending, \
963 will restart cache lookup",
964 self.inner.name
965 );
966 continue;
967 }
968 }
969 }
970 }
971
972 pub async fn get_or_try_insert<E: Into<anyhow::Error>, TTL: FnOnce(&V) -> Duration>(
980 &self,
981 name: &K,
982 ttl_func: TTL,
983 fut: impl Future<Output = Result<V, E>>,
984 ) -> Result<ItemLookup<V>, Arc<anyhow::Error>> {
985 if let Some(entry) = self.lookup(name) {
987 return Ok(entry);
988 }
989
990 let timeout_duration = Duration::from_millis(
991 self.inner.sema_timeout_milliseconds.load(Ordering::Relaxed) as u64,
992 );
993 let start = Instant::now();
994 let deadline = start + timeout_duration;
995
996 let permit = match self.acquire(name, deadline, timeout_duration).await {
997 Acquisition::Resolved(result) => return result,
998 Acquisition::Owner(permit) => permit,
999 };
1000
1001 defer! {
1004 permit.semaphore().close();
1005 }
1006
1007 self.inner.populate_counter.inc();
1008 let mut ttl = Duration::from_secs(60);
1009 let future_result = fut.await;
1010 let now = Instant::now();
1011
1012 let (item_result, return_value) = match future_result {
1013 Ok(item) => {
1014 ttl = ttl_func(&item);
1015 (
1016 ItemState::Present(item.clone()),
1017 Ok(ItemLookup {
1018 item,
1019 expiration: now + ttl,
1020 is_fresh: true,
1021 }),
1022 )
1023 }
1024 Err(err) => {
1025 self.inner.error_counter.inc();
1026 let err = Arc::new(err.into());
1027 (ItemState::Failed(err.clone()), Err(err))
1028 }
1029 };
1030
1031 self.inner.cache.insert(
1033 name.clone(),
1034 Item {
1035 item: item_result,
1036 expiration: Instant::now() + ttl,
1037 last_tick: self.inc_tick().into(),
1038 },
1039 );
1040 self.inner.maybe_evict();
1041 return_value
1042 }
1043
1044 pub async fn get_or_try_insert_detached<E, TTL, MK, F>(
1053 &self,
1054 name: &K,
1055 ttl_func: TTL,
1056 make_fut: MK,
1057 populate_timeout: Duration,
1058 ) -> Result<ItemLookup<V>, Arc<anyhow::Error>>
1059 where
1060 E: Into<anyhow::Error> + Send + 'static,
1061 TTL: Clone + Send + 'static + FnOnce(&V) -> Duration,
1062 MK: Fn() -> F,
1063 F: Future<Output = Result<V, E>> + Send + 'static,
1064 {
1065 if let Some(entry) = self.lookup(name) {
1066 return Ok(entry);
1067 }
1068
1069 let timeout_duration = Duration::from_millis(
1070 self.inner.sema_timeout_milliseconds.load(Ordering::Relaxed) as u64,
1071 );
1072 let deadline = Instant::now() + timeout_duration;
1073
1074 loop {
1075 let permit = match self.acquire(name, deadline, timeout_duration).await {
1076 Acquisition::Resolved(result) => return result,
1077 Acquisition::Owner(permit) => permit,
1078 };
1079
1080 let this = self.clone();
1081 let key = name.clone();
1082 let ttl_func = ttl_func.clone();
1083 let fut = make_fut();
1084 let sema = permit.semaphore().clone();
1087
1088 self.set_pending_expiration(name, Instant::now() + populate_timeout);
1092
1093 kumo_server_runtime::spawn("lruttl-populate", async move {
1094 defer! {
1097 permit.semaphore().close();
1098 }
1099
1100 let our_sema = permit.semaphore();
1104
1105 this.inner.populate_counter.inc();
1106 let (item, ttl) = match timeout(populate_timeout, fut).await {
1107 Ok(Ok(value)) => {
1108 let ttl = ttl_func(&value);
1109 (ItemState::Present(value), ttl)
1110 }
1111 Ok(Err(err)) => {
1112 this.inner.error_counter.inc();
1113 (
1117 ItemState::Failed(Arc::new(err.into())),
1118 Duration::from_secs(60),
1119 )
1120 }
1121 Err(_) => {
1122 this.inner.error_counter.inc();
1123 (
1124 ItemState::Failed(Arc::new(anyhow::anyhow!(
1125 "{} populate for {key:?} timed out after {populate_timeout:?}",
1126 this.inner.name
1127 ))),
1128 Duration::from_secs(60),
1129 )
1130 }
1131 };
1132
1133 if let Some(mut entry) = this.inner.cache.get_mut(&key) {
1137 let is_current = match &entry.item {
1138 ItemState::Pending(current)
1139 | ItemState::Refreshing {
1140 pending: current, ..
1141 } => Arc::ptr_eq(current, our_sema),
1142 ItemState::Present(_) | ItemState::Failed(_) => false,
1143 };
1144 if is_current {
1145 entry.item = item;
1146 entry.expiration = Instant::now() + ttl;
1147 entry.last_tick = this.inc_tick().into();
1148 }
1149 }
1150 this.inner.maybe_evict();
1151 })
1152 .map_err(|err| {
1153 sema.close();
1157 Arc::new(anyhow::anyhow!(
1158 "{} failed to spawn detached populate for {name:?}: {err:#}",
1159 self.inner.name
1160 ))
1161 })?;
1162
1163 }
1168 }
1169
1170 fn set_pending_expiration(&self, name: &K, expiration: Instant) {
1174 if let Some(mut entry) = self.inner.cache.get_mut(name) {
1175 if matches!(
1176 entry.item,
1177 ItemState::Pending(_) | ItemState::Refreshing { .. }
1178 ) {
1179 entry.expiration = expiration;
1180 }
1181 }
1182 }
1183}
1184
1185#[cfg(test)]
1186mod test {
1187 use super::*;
1188 use test_log::test; #[test(tokio::test)]
1191 async fn test_capacity() {
1192 let cache = LruCacheWithTtl::new("test_capacity", 40);
1193
1194 let expiration = Instant::now() + Duration::from_secs(60);
1195 for i in 0..100 {
1196 cache.insert(i, i, expiration).await;
1197 }
1198
1199 assert_eq!(cache.inner.cache.len(), 40, "capacity is respected");
1200 }
1201
1202 #[test(tokio::test)]
1203 async fn test_expiration() {
1204 let cache = LruCacheWithTtl::new("test_expiration", 1);
1205
1206 tokio::time::pause();
1207 let expiration = Instant::now() + Duration::from_secs(1);
1208 cache.insert(0, 0, expiration).await;
1209
1210 cache.get(&0).expect("still in cache");
1211 tokio::time::advance(Duration::from_secs(2)).await;
1212 assert!(cache.get(&0).is_none(), "evicted due to ttl");
1213 }
1214
1215 #[test(tokio::test)]
1216 async fn test_zero_ttl_requeries() {
1217 let cache = LruCacheWithTtl::<u32, u32>::new("test_zero_ttl_requeries", 1);
1219
1220 let first = cache
1221 .get_or_try_insert(&0, |_| Duration::ZERO, async { Ok::<_, anyhow::Error>(1) })
1222 .await
1223 .unwrap();
1224 assert_eq!(first.item, 1);
1225 assert!(first.is_fresh);
1226
1227 let second = cache
1228 .get_or_try_insert(&0, |_| Duration::ZERO, async { Ok::<_, anyhow::Error>(2) })
1229 .await
1230 .unwrap();
1231 assert_eq!(second.item, 2, "the zero-ttl entry was not reused");
1232 assert!(second.is_fresh);
1233 }
1234
1235 #[test(tokio::test)]
1236 async fn test_over_capacity_slow_resolve() {
1237 let cache = Arc::new(LruCacheWithTtl::<String, u64>::new(
1238 "test_over_capacity_slow_resolve",
1239 1,
1240 ));
1241
1242 let mut foos = vec![];
1243 for idx in 0..2 {
1244 let cache = cache.clone();
1245 foos.push(tokio::spawn(async move {
1246 eprintln!("spawned task {idx} is running");
1247 cache
1248 .get_or_try_insert(&"foo".to_string(), |_| Duration::from_secs(86400), async {
1249 if idx == 0 {
1250 eprintln!("foo {idx} getter sleeping");
1251 tokio::time::sleep(Duration::from_secs(300)).await;
1252 }
1253 eprintln!("foo {idx} getter done");
1254 Ok::<_, anyhow::Error>(idx)
1255 })
1256 .await
1257 }));
1258 }
1259
1260 tokio::task::yield_now().await;
1261
1262 eprintln!("calling again with immediate getter");
1263 let result = cache
1264 .get_or_try_insert(&"bar".to_string(), |_| Duration::from_secs(60), async {
1265 eprintln!("bar immediate getter running");
1266 Ok::<_, anyhow::Error>(42)
1267 })
1268 .await
1269 .unwrap();
1270
1271 assert_eq!(result.item, 42);
1272 assert_eq!(cache.inner.cache.len(), 1);
1273
1274 eprintln!("aborting first one");
1275 foos.remove(0).abort();
1276
1277 eprintln!("try new key");
1278 let result = cache
1279 .get_or_try_insert(&"baz".to_string(), |_| Duration::from_secs(60), async {
1280 eprintln!("baz immediate getter running");
1281 Ok::<_, anyhow::Error>(32)
1282 })
1283 .await
1284 .unwrap();
1285 assert_eq!(result.item, 32);
1286 assert_eq!(cache.inner.cache.len(), 1);
1287
1288 eprintln!("waiting second one");
1289 assert_eq!(1, foos.pop().unwrap().await.unwrap().unwrap().item);
1290
1291 assert_eq!(cache.inner.cache.len(), 1);
1292 }
1293
1294 #[test(tokio::test)]
1295 async fn detached_populate_survives_caller_cancellation() {
1296 use tokio::sync::Notify;
1297
1298 let cache = Arc::new(LruCacheWithTtl::<String, u64>::new(
1299 "detached_survives_cancel",
1300 16,
1301 ));
1302 let key = "k".to_string();
1303 let started = Arc::new(Notify::new());
1304 let gate = Arc::new(Notify::new());
1305 let runs = Arc::new(AtomicUsize::new(0));
1306
1307 let caller = {
1308 let cache = cache.clone();
1309 let key = key.clone();
1310 let started = started.clone();
1311 let gate = gate.clone();
1312 let runs = runs.clone();
1313 tokio::spawn(async move {
1314 cache
1315 .get_or_try_insert_detached(
1316 &key,
1317 |_| Duration::from_secs(86400),
1318 move || {
1319 let started = started.clone();
1320 let gate = gate.clone();
1321 let runs = runs.clone();
1322 async move {
1323 runs.fetch_add(1, Ordering::SeqCst);
1324 started.notify_one();
1325 gate.notified().await;
1326 Ok::<_, anyhow::Error>(42u64)
1327 }
1328 },
1329 Duration::from_secs(300),
1330 )
1331 .await
1332 })
1333 };
1334
1335 started.notified().await;
1338 caller.abort();
1339 gate.notify_one();
1341
1342 loop {
1343 if let Some(value) = cache.get(&key) {
1344 assert_eq!(value, 42);
1345 break;
1346 }
1347 tokio::task::yield_now().await;
1348 }
1349 assert_eq!(runs.load(Ordering::SeqCst), 1);
1350 }
1351
1352 #[test(tokio::test)]
1353 async fn detached_second_caller_shares_populate() {
1354 use tokio::sync::Notify;
1355
1356 let cache = Arc::new(LruCacheWithTtl::<String, u64>::new(
1357 "detached_shared_populate",
1358 16,
1359 ));
1360 let key = "k".to_string();
1361 let gate = Arc::new(Notify::new());
1362 let runs = Arc::new(AtomicUsize::new(0));
1363
1364 let make_caller = || {
1365 let cache = cache.clone();
1366 let key = key.clone();
1367 let gate = gate.clone();
1368 let runs = runs.clone();
1369 tokio::spawn(async move {
1370 cache
1371 .get_or_try_insert_detached(
1372 &key,
1373 |_| Duration::from_secs(86400),
1374 move || {
1375 let gate = gate.clone();
1376 let runs = runs.clone();
1377 async move {
1378 runs.fetch_add(1, Ordering::SeqCst);
1379 gate.notified().await;
1380 Ok::<_, anyhow::Error>(7u64)
1381 }
1382 },
1383 Duration::from_secs(300),
1384 )
1385 .await
1386 })
1387 };
1388
1389 let first = make_caller();
1390 let second = make_caller();
1391
1392 while runs.load(Ordering::SeqCst) == 0 {
1394 tokio::task::yield_now().await;
1395 }
1396 gate.notify_one();
1397
1398 let first = first.await.unwrap().unwrap();
1399 let second = second.await.unwrap().unwrap();
1400 assert_eq!(first.item, 7);
1401 assert_eq!(second.item, 7);
1402 assert_eq!(
1403 runs.load(Ordering::SeqCst),
1404 1,
1405 "the two callers must share one populate"
1406 );
1407 }
1408
1409 #[test(tokio::test)]
1410 async fn detached_populate_timeout_yields_one_error() {
1411 let cache = Arc::new(LruCacheWithTtl::<String, u64>::new(
1412 "detached_populate_timeout",
1413 16,
1414 ));
1415 let key = "k".to_string();
1416 let runs = Arc::new(AtomicUsize::new(0));
1417
1418 let runs_in_fut = runs.clone();
1419 let result = cache
1420 .get_or_try_insert_detached(
1421 &key,
1422 |_| Duration::from_secs(86400),
1423 move || {
1424 let runs = runs_in_fut.clone();
1425 async move {
1426 runs.fetch_add(1, Ordering::SeqCst);
1427 std::future::pending::<()>().await;
1430 Ok::<_, anyhow::Error>(0u64)
1431 }
1432 },
1433 Duration::from_millis(50),
1434 )
1435 .await;
1436
1437 let err = result.unwrap_err();
1438 assert!(
1439 err.to_string().contains("timed out after"),
1440 "unexpected error: {err}"
1441 );
1442 assert_eq!(
1443 runs.load(Ordering::SeqCst),
1444 1,
1445 "a timed-out populate must leave a Failed entry, not spin up repeated populates"
1446 );
1447 }
1448}