1use crate::{Document, Error, Result};
2use std::{
3 collections::HashMap,
4 future::Future,
5 sync::{Arc, Mutex},
6 time::{Duration, Instant},
7};
8use tokio::sync::Mutex as AsyncMutex;
9
10#[derive(Default)]
11pub struct Cache(Mutex<HashMap<String, Arc<Entry>>>);
12
13#[derive(Default)]
14struct Entry {
15 state: Mutex<Cached>,
16 loading: Arc<AsyncMutex<()>>,
17}
18
19struct Cached {
20 value: Option<(Instant, Arc<Document>)>,
21 failure: Option<(Instant, Error)>,
22 used: Instant,
23}
24impl Default for Cached {
25 fn default() -> Self {
26 Self {
27 value: None,
28 failure: None,
29 used: Instant::now(),
30 }
31 }
32}
33
34impl Cache {
35 fn entry(&self, key: String) -> Result<Arc<Entry>> {
36 let mut entries = self.0.lock().unwrap();
37 if !entries.contains_key(&key) && entries.len() >= 256 {
38 let oldest = entries
39 .iter()
40 .filter(|(_, entry)| Arc::strong_count(entry) == 1)
41 .min_by_key(|(_, entry)| entry.state.lock().unwrap().used)
42 .map(|(key, _)| key.clone());
43 if let Some(oldest) = oldest {
44 entries.remove(&oldest);
45 } else {
46 return Err(Error::new(503, "The dashboard is busy. Retry in a moment."));
47 }
48 }
49 let entry = entries.entry(key).or_default().clone();
50 entry.state.lock().unwrap().used = Instant::now();
51 Ok(entry)
52 }
53
54 pub fn invalidate(&self, key: &str) {
55 self.0.lock().unwrap().remove(key);
56 }
57
58 pub fn invalidate_prefix(&self, prefix: &str) {
59 self.0
60 .lock()
61 .unwrap()
62 .retain(|key, _| !key.starts_with(prefix));
63 }
64
65 pub fn peek(&self, key: &str, ttl: Duration) -> Option<Arc<Document>> {
66 let entry = self.0.lock().unwrap().get(key).cloned()?;
67 let state = entry.state.lock().unwrap();
68 state
69 .value
70 .as_ref()
71 .filter(|(at, _)| at.elapsed() <= ttl + Duration::from_secs(60))
72 .map(|(_, value)| value.clone())
73 }
74
75 pub async fn get<F, Fut>(&self, key: String, ttl: Duration, load: F) -> Result<Arc<Document>>
76 where
77 F: FnOnce() -> Fut + Send + 'static,
78 Fut: Future<Output = Result<serde_json::Value>> + Send + 'static,
79 {
80 let entry = self.entry(key)?;
81 let stale = {
82 let state = entry.state.lock().unwrap();
83 let value = state
84 .value
85 .as_ref()
86 .filter(|(at, _)| at.elapsed() <= ttl + Duration::from_secs(60));
87 if let Some((at, value)) = value
88 && at.elapsed() < ttl
89 {
90 return Ok(value.clone());
91 }
92 if let Some((at, error)) = &state.failure
93 && at.elapsed() < ttl
94 {
95 return value
96 .map(|(_, value)| value.clone())
97 .ok_or_else(|| error.clone());
98 }
99 value.map(|(_, value)| value.clone())
100 };
101 if let Some(stale) = stale {
102 if let Ok(guard) = entry.loading.clone().try_lock_owned() {
103 tokio::spawn(async move {
104 let _guard = guard;
105 let _ = entry.store(load().await);
106 });
107 }
108 return Ok(stale);
109 }
110 let guard = entry.loading.clone().lock_owned().await;
111 {
112 let state = entry.state.lock().unwrap();
113 if let Some((at, value)) = &state.value
114 && at.elapsed() < ttl
115 {
116 return Ok(value.clone());
117 }
118 if let Some((at, error)) = &state.failure
119 && at.elapsed() < ttl
120 {
121 return Err(error.clone());
122 }
123 }
124 tokio::spawn(async move {
125 let _guard = guard;
126 entry.store(load().await)
127 })
128 .await?
129 }
130}
131
132impl Entry {
133 fn store(&self, result: Result<serde_json::Value>) -> Result<Arc<Document>> {
134 let mut state = self.state.lock().unwrap();
135 match result {
136 Ok(value) => {
137 let document = Arc::new(Document::new(value));
138 state.value = Some((Instant::now(), document.clone()));
139 state.failure = None;
140 Ok(document)
141 }
142 Err(error) => {
143 state.failure = Some((Instant::now(), error.clone()));
144 Err(error)
145 }
146 }
147 }
148}
149
150#[cfg(test)]
151mod tests {
152 use super::*;
153 use std::sync::atomic::{AtomicUsize, Ordering};
154 #[tokio::test]
155 async fn disconnecting_reader_does_not_cancel_shared_load() {
156 let cache = Arc::new(Cache::default());
157 let started = Arc::new(tokio::sync::Notify::new());
158 let release = Arc::new(tokio::sync::Notify::new());
159 let (state, ready, finish) = (cache.clone(), started.clone(), release.clone());
160 let first = tokio::spawn(async move {
161 state
162 .get(
163 "cancelled".into(),
164 Duration::from_secs(10),
165 move || async move {
166 ready.notify_one();
167 finish.notified().await;
168 Ok(serde_json::json!({"ready":true}))
169 },
170 )
171 .await
172 });
173 started.notified().await;
174 first.abort();
175 release.notify_one();
176 let value = cache
177 .get("cancelled".into(), Duration::from_secs(10), || async {
178 panic!("shared load was repeated");
179 })
180 .await
181 .unwrap();
182 assert_eq!(value.value["ready"], true);
183 }
184 #[tokio::test]
185 async fn coalesces_readers_and_backs_off_failures() {
186 let cache = Arc::new(Cache::default());
187 let calls = Arc::new(AtomicUsize::new(0));
188 let mut readers = Vec::new();
189 for _ in 0..100 {
190 let (cache, calls) = (cache.clone(), calls.clone());
191 readers.push(tokio::spawn(async move {
192 cache
193 .get("same".into(), Duration::from_secs(10), move || async move {
194 calls.fetch_add(1, Ordering::SeqCst);
195 tokio::time::sleep(Duration::from_millis(20)).await;
196 Err(Error::new(502, "unavailable"))
197 })
198 .await
199 }));
200 }
201 for reader in readers {
202 assert!(reader.await.unwrap().is_err());
203 }
204 assert_eq!(calls.load(Ordering::SeqCst), 1);
205 cache.invalidate("same");
206 assert!(
207 cache
208 .get("same".into(), Duration::from_secs(10), || async {
209 Ok(serde_json::json!({"ok":true}))
210 })
211 .await
212 .is_ok()
213 );
214 }
215}