1use std::collections::HashMap;
4use std::sync::{Mutex, MutexGuard};
5
6use sva_formula::Hash;
7use sva_samples::Label;
8
9use super::{Entry, Expected, Payload};
10
11pub const DEFAULT_CACHE_BYTES: u64 = 2 << 30;
12
13#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
15pub enum CachePolicy {
16 #[default]
17 All,
18 Forks,
20 Target,
21 None,
22}
23
24impl CachePolicy {
25 pub const ALL: [CachePolicy; 4] = [
26 CachePolicy::All,
27 CachePolicy::Forks,
28 CachePolicy::Target,
29 CachePolicy::None,
30 ];
31
32 pub fn name(self) -> &'static str {
33 match self {
34 CachePolicy::All => "all",
35 CachePolicy::Forks => "forks",
36 CachePolicy::Target => "target",
37 CachePolicy::None => "none",
38 }
39 }
40
41 pub fn named(name: &str) -> Option<CachePolicy> {
42 CachePolicy::ALL.into_iter().find(|p| p.name() == name)
43 }
44
45 pub(crate) fn stores(self, fork: bool, target: bool) -> bool {
46 match self {
47 CachePolicy::All => true,
48 CachePolicy::Forks => fork || target,
49 CachePolicy::Target => target,
50 CachePolicy::None => false,
51 }
52 }
53}
54
55#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
56pub enum PrunePolicy {
57 #[default]
59 Oldest,
60 Forks,
62}
63
64impl PrunePolicy {
65 pub const ALL: [PrunePolicy; 2] = [PrunePolicy::Oldest, PrunePolicy::Forks];
66
67 pub fn name(self) -> &'static str {
68 match self {
69 PrunePolicy::Oldest => "oldest",
70 PrunePolicy::Forks => "forks",
71 }
72 }
73
74 pub fn named(name: &str) -> Option<PrunePolicy> {
75 PrunePolicy::ALL.into_iter().find(|p| p.name() == name)
76 }
77}
78
79#[derive(Clone, Copy, Debug)]
80pub(crate) struct Stamp {
81 pub tree: u64,
82 pub fork: bool,
83 pub slot: Option<Hash>,
85}
86
87#[derive(Clone, Copy, Debug, PartialEq, Eq)]
88pub(crate) enum Kept {
89 Held,
90 Replaced,
91 Refused,
92}
93
94struct Held {
95 payload: Payload,
96 label: Option<Label>,
97 read: u64,
98 tree: u64,
99 fork: bool,
100 slot: Option<Hash>,
101}
102
103impl Held {
104 fn bytes(&self) -> u64 {
105 self.payload.bytes() as u64
106 }
107}
108
109#[derive(Default)]
110struct State {
111 entries: HashMap<Hash, Held>,
112 slots: HashMap<Hash, Hash>,
113 bytes: u64,
114 max_bytes: u64,
115 policy: CachePolicy,
116 prune: PrunePolicy,
117 clock: u64,
118 tree: u64,
119 evictions: u64,
120}
121
122impl State {
123 fn tick(&mut self) -> u64 {
124 self.clock += 1;
125 self.clock
126 }
127
128 fn remove(&mut self, key: Hash) -> bool {
129 let Some(gone) = self.entries.remove(&key) else {
130 return false;
131 };
132 self.bytes -= gone.bytes();
133 if let Some(slot) = gone.slot
134 && self.slots.get(&slot) == Some(&key)
135 {
136 self.slots.remove(&slot);
137 }
138 true
139 }
140
141 fn evict(&mut self, key: Hash) {
142 if self.remove(key) {
143 self.evictions += 1;
144 }
145 }
146
147 fn prune(&mut self, policy: PrunePolicy, to: u64) {
150 let newest = self.tree;
151 let mut named: Vec<(u64, Hash)> = self
152 .entries
153 .iter()
154 .filter(|(_, held)| match policy {
155 PrunePolicy::Oldest => held.tree != newest,
156 PrunePolicy::Forks => !held.fork,
157 })
158 .map(|(key, held)| (held.read, *key))
159 .collect();
160 named.sort_unstable();
161 for (_, key) in named {
162 if self.bytes <= to {
163 break;
164 }
165 self.evict(key);
166 }
167 let mut trees: Vec<u64> = self.entries.values().map(|held| held.tree).collect();
168 trees.sort_unstable();
169 trees.dedup();
170 for tree in trees {
171 if self.bytes <= self.max_bytes {
172 break;
173 }
174 let whole: Vec<Hash> = self
175 .entries
176 .iter()
177 .filter(|(_, held)| held.tree == tree)
178 .map(|(key, _)| *key)
179 .collect();
180 for key in whole {
181 self.evict(key);
182 }
183 }
184 }
185
186 fn bounded(&mut self) {
187 if self.bytes > self.max_bytes {
188 self.prune(self.prune, self.max_bytes);
189 }
190 }
191}
192
193pub struct Cache {
196 state: Mutex<State>,
197}
198
199impl Default for Cache {
200 fn default() -> Cache {
201 Cache::new()
202 }
203}
204
205impl Cache {
206 pub fn new() -> Cache {
207 Cache::holding(DEFAULT_CACHE_BYTES)
208 }
209
210 pub fn holding(max_bytes: u64) -> Cache {
211 Cache {
212 state: Mutex::new(State {
213 max_bytes,
214 ..State::default()
215 }),
216 }
217 }
218
219 fn locked(&self) -> MutexGuard<'_, State> {
220 self.state.lock().unwrap_or_else(|poisoned| {
221 let mut state = poisoned.into_inner();
222 state.entries.clear();
223 state.slots.clear();
224 state.bytes = 0;
225 self.state.clear_poison();
226 state
227 })
228 }
229
230 pub fn max_bytes(&self) -> u64 {
231 self.locked().max_bytes
232 }
233
234 pub fn set_max_bytes(&self, max_bytes: u64) {
235 let mut state = self.locked();
236 state.max_bytes = max_bytes;
237 state.bounded();
238 }
239
240 pub fn policy(&self) -> CachePolicy {
242 self.locked().policy
243 }
244
245 pub fn set_policy(&self, policy: CachePolicy) {
246 self.locked().policy = policy;
247 }
248
249 pub fn prune_policy(&self) -> PrunePolicy {
250 self.locked().prune
251 }
252
253 pub fn set_prune_policy(&self, policy: PrunePolicy) {
254 self.locked().prune = policy;
255 }
256
257 pub fn prune(&self, policy: PrunePolicy) {
259 self.locked().prune(policy, 0);
260 }
261
262 pub fn clear(&self) {
263 let mut state = self.locked();
264 state.entries.clear();
265 state.slots.clear();
266 state.bytes = 0;
267 }
268
269 pub fn bytes(&self) -> u64 {
270 self.locked().bytes
271 }
272
273 pub fn entries(&self) -> usize {
274 self.locked().entries.len()
275 }
276
277 pub fn evictions(&self) -> u64 {
278 self.locked().evictions
279 }
280
281 pub fn holds(&self, key: Hash) -> bool {
282 self.locked().entries.contains_key(&key)
283 }
284
285 pub(crate) fn begin_tree(&self) -> u64 {
286 let mut state = self.locked();
287 state.tree += 1;
288 state.tree
289 }
290
291 pub(crate) fn load(&self, key: Hash, expected: Expected, stamp: Stamp) -> Option<Entry> {
292 let mut state = self.locked();
293 let tick = state.tick();
294 let held = state.entries.get_mut(&key)?;
295 if !held.payload.answers(expected) {
296 state.remove(key);
297 return None;
298 }
299 held.read = tick;
300 held.tree = stamp.tree;
301 held.fork = stamp.fork;
302 Some(Entry {
303 payload: held.payload.clone(),
304 label: held.label.clone(),
305 })
306 }
307
308 pub(crate) fn store(
309 &self,
310 key: Hash,
311 payload: &Payload,
312 label: Option<&Label>,
313 stamp: Stamp,
314 ) -> Kept {
315 let mut state = self.locked();
316 let bytes = payload.bytes() as u64;
317 if bytes > state.max_bytes {
318 return Kept::Refused;
319 }
320 let read = state.tick();
321 let replaced = match stamp.slot.and_then(|slot| state.slots.insert(slot, key)) {
322 Some(last) if last != key => state.remove(last),
323 _ => false,
324 };
325 let held = Held {
326 payload: payload.clone(),
327 label: label.cloned(),
328 read,
329 tree: stamp.tree,
330 fork: stamp.fork,
331 slot: stamp.slot,
332 };
333 if let Some(old) = state.entries.insert(key, held) {
334 state.bytes -= old.bytes();
335 }
336 state.bytes += bytes;
337 state.bounded();
338 match replaced {
339 true => Kept::Replaced,
340 false => Kept::Held,
341 }
342 }
343}
344
345#[cfg(test)]
346mod tests {
347 use super::*;
348 use sva_samples::Buffer;
349
350 #[test]
352 fn an_entry_that_does_not_answer_what_was_asked_is_a_miss_and_goes() {
353 let cache = Cache::new();
354 let key = Hash(7, 11);
355 let stamp = Stamp {
356 tree: cache.begin_tree(),
357 fork: false,
358 slot: None,
359 };
360 let four = Payload::Samples(Box::new(Buffer::mono(8_000, vec![0.25; 4])));
361 for (rate, samples) in [(8_000, 5), (48_000, 4)] {
362 cache.store(key, &four, None, stamp);
363 let asked = Expected::Samples {
364 rate,
365 width: 1,
366 samples,
367 };
368 assert!(cache.load(key, asked, stamp).is_none());
369 assert!(!cache.holds(key));
370 assert_eq!(cache.bytes(), 0);
371 }
372 }
373}