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 async fn coalesce<F, Fut>(&self, key: String, load: F) -> Result<Arc<Document>>
55 where
56 F: FnOnce() -> Fut + Send + 'static,
57 Fut: Future<Output = Result<serde_json::Value>> + Send + 'static,
58 {
59 let started = Instant::now();
60 let entry = self.entry(key)?;
61 let guard = entry.loading.clone().lock_owned().await;
62 {
63 let state = entry.state.lock().unwrap();
64 if let Some((at, value)) = &state.value
65 && *at >= started
66 {
67 return Ok(value.clone());
68 }
69 if let Some((at, error)) = &state.failure
70 && *at >= started
71 {
72 return Err(error.clone());
73 }
74 }
75 tokio::spawn(async move {
76 let _guard = guard;
77 entry.store(load().await)
78 })
79 .await?
80 }
81
82 pub fn invalidate(&self, key: &str) {
83 self.0.lock().unwrap().remove(key);
84 }
85
86 pub fn invalidate_prefix(&self, prefix: &str) {
87 self.0
88 .lock()
89 .unwrap()
90 .retain(|key, _| !key.starts_with(prefix));
91 }
92
93 pub fn peek(&self, key: &str, ttl: Duration) -> Option<Arc<Document>> {
94 let entry = self.0.lock().unwrap().get(key).cloned()?;
95 let state = entry.state.lock().unwrap();
96 state
97 .value
98 .as_ref()
99 .filter(|(at, _)| at.elapsed() <= ttl + Duration::from_secs(60))
100 .map(|(_, value)| value.clone())
101 }
102
103 pub async fn get<F, Fut>(&self, key: String, ttl: Duration, load: F) -> Result<Arc<Document>>
104 where
105 F: FnOnce() -> Fut + Send + 'static,
106 Fut: Future<Output = Result<serde_json::Value>> + Send + 'static,
107 {
108 let entry = self.entry(key)?;
109 let stale = {
110 let state = entry.state.lock().unwrap();
111 let value = state
112 .value
113 .as_ref()
114 .filter(|(at, _)| at.elapsed() <= ttl + Duration::from_secs(60));
115 if let Some((at, value)) = value
116 && at.elapsed() < ttl
117 {
118 return Ok(value.clone());
119 }
120 if let Some((at, error)) = &state.failure
121 && at.elapsed() < ttl
122 {
123 return value
124 .map(|(_, value)| value.clone())
125 .ok_or_else(|| error.clone());
126 }
127 value.map(|(_, value)| value.clone())
128 };
129 if let Some(stale) = stale {
130 if let Ok(guard) = entry.loading.clone().try_lock_owned() {
131 tokio::spawn(async move {
132 let _guard = guard;
133 let _ = entry.store(load().await);
134 });
135 }
136 return Ok(stale);
137 }
138 let guard = entry.loading.clone().lock_owned().await;
139 {
140 let state = entry.state.lock().unwrap();
141 if let Some((at, value)) = &state.value
142 && at.elapsed() < ttl
143 {
144 return Ok(value.clone());
145 }
146 if let Some((at, error)) = &state.failure
147 && at.elapsed() < ttl
148 {
149 return Err(error.clone());
150 }
151 }
152 tokio::spawn(async move {
153 let _guard = guard;
154 entry.store(load().await)
155 })
156 .await?
157 }
158}
159
160impl Entry {
161 fn store(&self, result: Result<serde_json::Value>) -> Result<Arc<Document>> {
162 let mut state = self.state.lock().unwrap();
163 match result {
164 Ok(value) => {
165 let document = Arc::new(Document::new(value));
166 state.value = Some((Instant::now(), document.clone()));
167 state.failure = None;
168 Ok(document)
169 }
170 Err(error) => {
171 state.failure = Some((Instant::now(), error.clone()));
172 Err(error)
173 }
174 }
175 }
176}
177
178#[cfg(test)]
179mod tests {
180 use super::*;
181 use std::sync::atomic::{AtomicUsize, Ordering};
182 #[tokio::test]
183 async fn coalesced_identity_reads_do_not_reuse_completed_or_failed_results() {
184 let cache = Arc::new(Cache::default());
185 let calls = Arc::new(AtomicUsize::new(0));
186 let mut readers = Vec::new();
187 for _ in 0..100 {
188 let (cache, calls) = (cache.clone(), calls.clone());
189 readers.push(tokio::spawn(async move {
190 cache
191 .coalesce("identity:owner".into(), move || async move {
192 calls.fetch_add(1, Ordering::SeqCst);
193 tokio::time::sleep(Duration::from_millis(20)).await;
194 Ok(serde_json::json!({"enabled":true}))
195 })
196 .await
197 .unwrap()
198 }));
199 }
200 for reader in readers {
201 assert_eq!(reader.await.unwrap().value["enabled"], true);
202 }
203 assert_eq!(calls.load(Ordering::SeqCst), 1);
204 let disabled = cache
205 .coalesce("identity:owner".into(), || async {
206 Ok(serde_json::json!({"enabled":false}))
207 })
208 .await
209 .unwrap();
210 assert_eq!(disabled.value["enabled"], false);
211 assert!(
212 cache
213 .coalesce("identity:owner".into(), || async {
214 Err(Error::new(502, "unavailable"))
215 })
216 .await
217 .is_err()
218 );
219 let recovered = cache
220 .coalesce("identity:owner".into(), || async {
221 Ok(serde_json::json!({"enabled":true}))
222 })
223 .await
224 .unwrap();
225 assert_eq!(recovered.value["enabled"], true);
226 }
227 #[tokio::test]
228 async fn disconnecting_reader_does_not_cancel_shared_load() {
229 let cache = Arc::new(Cache::default());
230 let started = Arc::new(tokio::sync::Notify::new());
231 let release = Arc::new(tokio::sync::Notify::new());
232 let (state, ready, finish) = (cache.clone(), started.clone(), release.clone());
233 let first = tokio::spawn(async move {
234 state
235 .get(
236 "cancelled".into(),
237 Duration::from_secs(10),
238 move || async move {
239 ready.notify_one();
240 finish.notified().await;
241 Ok(serde_json::json!({"ready":true}))
242 },
243 )
244 .await
245 });
246 started.notified().await;
247 first.abort();
248 release.notify_one();
249 let value = cache
250 .get("cancelled".into(), Duration::from_secs(10), || async {
251 panic!("shared load was repeated");
252 })
253 .await
254 .unwrap();
255 assert_eq!(value.value["ready"], true);
256 }
257 #[tokio::test]
258 async fn coalesces_readers_and_backs_off_failures() {
259 let cache = Arc::new(Cache::default());
260 let calls = Arc::new(AtomicUsize::new(0));
261 let mut readers = Vec::new();
262 for _ in 0..100 {
263 let (cache, calls) = (cache.clone(), calls.clone());
264 readers.push(tokio::spawn(async move {
265 cache
266 .get("same".into(), Duration::from_secs(10), move || async move {
267 calls.fetch_add(1, Ordering::SeqCst);
268 tokio::time::sleep(Duration::from_millis(20)).await;
269 Err(Error::new(502, "unavailable"))
270 })
271 .await
272 }));
273 }
274 for reader in readers {
275 assert!(reader.await.unwrap().is_err());
276 }
277 assert_eq!(calls.load(Ordering::SeqCst), 1);
278 cache.invalidate("same");
279 assert!(
280 cache
281 .get("same".into(), Duration::from_secs(10), || async {
282 Ok(serde_json::json!({"ok":true}))
283 })
284 .await
285 .is_ok()
286 );
287 }
288}