1use std::collections::HashMap;
4use std::sync::{Arc, 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
13pub const DEFAULT_MARK_EVERY: usize = 16_384;
15
16#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
18pub enum CachePolicy {
19 #[default]
20 All,
21 Forks,
23 Target,
24 None,
25}
26
27impl CachePolicy {
28 pub const ALL: [CachePolicy; 4] = [
29 CachePolicy::All,
30 CachePolicy::Forks,
31 CachePolicy::Target,
32 CachePolicy::None,
33 ];
34
35 pub fn name(self) -> &'static str {
36 match self {
37 CachePolicy::All => "all",
38 CachePolicy::Forks => "forks",
39 CachePolicy::Target => "target",
40 CachePolicy::None => "none",
41 }
42 }
43
44 pub fn named(name: &str) -> Option<CachePolicy> {
45 CachePolicy::ALL.into_iter().find(|p| p.name() == name)
46 }
47
48 pub(crate) fn stores(self, fork: bool, target: bool) -> bool {
49 match self {
50 CachePolicy::All => true,
51 CachePolicy::Forks => fork || target,
52 CachePolicy::Target => target,
53 CachePolicy::None => false,
54 }
55 }
56}
57
58#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
59pub enum PrunePolicy {
60 #[default]
62 Oldest,
63 Forks,
65}
66
67impl PrunePolicy {
68 pub const ALL: [PrunePolicy; 2] = [PrunePolicy::Oldest, PrunePolicy::Forks];
69
70 pub fn name(self) -> &'static str {
71 match self {
72 PrunePolicy::Oldest => "oldest",
73 PrunePolicy::Forks => "forks",
74 }
75 }
76
77 pub fn named(name: &str) -> Option<PrunePolicy> {
78 PrunePolicy::ALL.into_iter().find(|p| p.name() == name)
79 }
80}
81
82#[derive(Clone, Copy, Debug)]
83pub(crate) struct Stamp {
84 pub tree: u64,
85 pub fork: bool,
86 pub slot: Option<Hash>,
88}
89
90#[derive(Clone, Copy, Debug, PartialEq, Eq)]
91pub(crate) enum Kept {
92 Held,
93 Replaced,
94 Refused,
95}
96
97struct Held {
98 payload: Payload,
99 label: Option<Label>,
100 read: u64,
101 tree: u64,
102 fork: bool,
103 slot: Option<Hash>,
104}
105
106impl Held {
107 fn bytes(&self) -> u64 {
108 self.payload.bytes() as u64
109 }
110}
111
112#[derive(Default)]
113struct State {
114 entries: HashMap<Hash, Held>,
115 slots: HashMap<Hash, Hash>,
116 bytes: u64,
117 max_bytes: u64,
118 policy: CachePolicy,
119 prune: PrunePolicy,
120 clock: u64,
121 tree: u64,
122 evictions: u64,
123 mark_every: usize,
124}
125
126impl State {
127 fn tick(&mut self) -> u64 {
128 self.clock += 1;
129 self.clock
130 }
131
132 fn remove(&mut self, key: Hash) -> bool {
133 let Some(gone) = self.entries.remove(&key) else {
134 return false;
135 };
136 self.bytes -= gone.bytes();
137 if let Some(slot) = gone.slot
138 && self.slots.get(&slot) == Some(&key)
139 {
140 self.slots.remove(&slot);
141 }
142 true
143 }
144
145 fn evict(&mut self, key: Hash) {
146 if self.remove(key) {
147 self.evictions += 1;
148 }
149 }
150
151 fn prune(&mut self, policy: PrunePolicy, to: u64) {
154 let newest = self.tree;
155 let mut named: Vec<(u64, Hash)> = self
156 .entries
157 .iter()
158 .filter(|(_, held)| match policy {
159 PrunePolicy::Oldest => held.tree != newest,
160 PrunePolicy::Forks => !held.fork,
161 })
162 .map(|(key, held)| (held.read, *key))
163 .collect();
164 named.sort_unstable();
165 for (_, key) in named {
166 if self.bytes <= to {
167 break;
168 }
169 self.evict(key);
170 }
171 let mut trees: Vec<u64> = self.entries.values().map(|held| held.tree).collect();
172 trees.sort_unstable();
173 trees.dedup();
174 for tree in trees {
175 if self.bytes <= self.max_bytes {
176 break;
177 }
178 let whole: Vec<Hash> = self
179 .entries
180 .iter()
181 .filter(|(_, held)| held.tree == tree)
182 .map(|(key, _)| *key)
183 .collect();
184 for key in whole {
185 self.evict(key);
186 }
187 }
188 }
189
190 fn bounded(&mut self) {
191 if self.bytes > self.max_bytes {
192 self.prune(self.prune, self.max_bytes);
193 }
194 }
195}
196
197#[derive(Clone)]
201pub struct Cache {
202 state: Arc<Mutex<State>>,
203}
204
205impl Default for Cache {
206 fn default() -> Cache {
207 Cache::new()
208 }
209}
210
211impl Cache {
212 pub fn new() -> Cache {
213 Cache::holding(DEFAULT_CACHE_BYTES)
214 }
215
216 pub fn holding(max_bytes: u64) -> Cache {
217 Cache {
218 state: Arc::new(Mutex::new(State {
219 max_bytes,
220 mark_every: DEFAULT_MARK_EVERY,
221 ..State::default()
222 })),
223 }
224 }
225
226 fn locked(&self) -> MutexGuard<'_, State> {
227 self.state.lock().unwrap_or_else(|poisoned| {
228 let mut state = poisoned.into_inner();
229 state.entries.clear();
230 state.slots.clear();
231 state.bytes = 0;
232 self.state.clear_poison();
233 state
234 })
235 }
236
237 pub fn max_bytes(&self) -> u64 {
238 self.locked().max_bytes
239 }
240
241 pub fn set_max_bytes(&self, max_bytes: u64) {
242 let mut state = self.locked();
243 state.max_bytes = max_bytes;
244 state.bounded();
245 }
246
247 pub fn policy(&self) -> CachePolicy {
249 self.locked().policy
250 }
251
252 pub fn set_policy(&self, policy: CachePolicy) {
253 self.locked().policy = policy;
254 }
255
256 pub fn prune_policy(&self) -> PrunePolicy {
257 self.locked().prune
258 }
259
260 pub fn set_prune_policy(&self, policy: PrunePolicy) {
261 self.locked().prune = policy;
262 }
263
264 pub fn prune(&self, policy: PrunePolicy) {
266 self.locked().prune(policy, 0);
267 }
268
269 pub fn clear(&self) {
270 let mut state = self.locked();
271 state.entries.clear();
272 state.slots.clear();
273 state.bytes = 0;
274 }
275
276 pub fn bytes(&self) -> u64 {
277 self.locked().bytes
278 }
279
280 pub fn entries(&self) -> usize {
281 self.locked().entries.len()
282 }
283
284 pub fn evictions(&self) -> u64 {
285 self.locked().evictions
286 }
287
288 pub fn holds(&self, key: Hash) -> bool {
289 self.locked().entries.contains_key(&key)
290 }
291
292 pub fn mark_every(&self) -> usize {
294 self.locked().mark_every
295 }
296
297 pub fn set_mark_every(&self, samples: usize) {
298 self.locked().mark_every = samples.max(1);
299 }
300
301 pub(crate) fn begin_tree(&self) -> u64 {
302 let mut state = self.locked();
303 state.tree += 1;
304 state.tree
305 }
306
307 pub(crate) fn load(&self, key: Hash, expected: Expected, stamp: Stamp) -> Option<Entry> {
308 let mut state = self.locked();
309 let tick = state.tick();
310 let held = state.entries.get_mut(&key)?;
311 if !held.payload.answers(expected) {
312 state.remove(key);
313 return None;
314 }
315 held.read = tick;
316 held.tree = stamp.tree;
317 held.fork = stamp.fork;
318 Some(Entry {
319 payload: held.payload.clone(),
320 label: held.label.clone(),
321 })
322 }
323
324 pub(crate) fn merge(
327 &self,
328 key: Hash,
329 payload: Payload,
330 label: Option<&Label>,
331 stamp: Stamp,
332 ) -> Kept {
333 let mut state = self.locked();
334 let tick = state.tick();
335 let joined = match (state.entries.get_mut(&key), payload) {
336 (Some(held), payload) if held.slot == stamp.slot => {
337 let before = held.bytes();
338 let payload = match (&mut held.payload, payload) {
339 (Payload::Segments(parts), Payload::Segments(more)) => {
340 joined(parts, more);
341 None
342 }
343 (Payload::Run(run), Payload::Run(more)) if overlaps(run, &more) => {
344 let from = (run.end() - more.samples.start).max(0) as usize;
345 for (held, more) in run.samples.planes.iter_mut().zip(&more.samples.planes)
346 {
347 held.extend_from_slice(&more[from.min(more.len())..]);
348 }
349 run.marks.extend(more.marks);
350 None
351 }
352 (_, payload) => Some(payload),
353 };
354 match payload {
355 None => {
356 held.read = tick;
357 held.tree = stamp.tree;
358 held.fork = stamp.fork;
359 let after = held.bytes();
360 Ok((before, after))
361 }
362 Some(payload) => Err(payload),
363 }
364 }
365 (_, payload) => Err(payload),
366 };
367 match joined {
368 Ok((before, after)) => {
369 state.bytes = state.bytes - before + after;
370 state.bounded();
371 Kept::Held
372 }
373 Err(payload) => {
374 drop(state);
375 self.store(key, &payload, label, stamp)
376 }
377 }
378 }
379
380 pub(crate) fn store(
381 &self,
382 key: Hash,
383 payload: &Payload,
384 label: Option<&Label>,
385 stamp: Stamp,
386 ) -> Kept {
387 let mut state = self.locked();
388 let bytes = payload.bytes() as u64;
389 if bytes > state.max_bytes {
390 return Kept::Refused;
391 }
392 let read = state.tick();
393 let replaced = match stamp.slot.and_then(|slot| state.slots.insert(slot, key)) {
394 Some(last) if last != key => state.remove(last),
395 _ => false,
396 };
397 let held = Held {
398 payload: payload.clone(),
399 label: label.cloned(),
400 read,
401 tree: stamp.tree,
402 fork: stamp.fork,
403 slot: stamp.slot,
404 };
405 if let Some(old) = state.entries.insert(key, held) {
406 state.bytes -= old.bytes();
407 }
408 state.bytes += bytes;
409 state.bounded();
410 match replaced {
411 true => Kept::Replaced,
412 false => Kept::Held,
413 }
414 }
415}
416
417fn overlaps(held: &super::Run, more: &super::Run) -> bool {
420 let (a, b) = (held.samples.start, held.end());
421 a <= more.samples.start && more.samples.start <= b
422}
423
424pub(crate) fn joined(parts: &mut Vec<sva_samples::Buffer>, more: Vec<sva_samples::Buffer>) {
426 for part in more {
427 parts.push(part);
428 }
429 parts.sort_by_key(|b| b.start);
430 let mut out: Vec<sva_samples::Buffer> = Vec::with_capacity(parts.len());
431 for part in parts.drain(..) {
432 match out.last_mut() {
433 Some(last) if last.extent().end >= part.start => {
434 let from = (last.extent().end - part.start) as usize;
435 for (held, more) in last.planes.iter_mut().zip(&part.planes) {
436 held.extend_from_slice(&more[from.min(more.len())..]);
437 }
438 }
439 _ => out.push(part),
440 }
441 }
442 *parts = out;
443}
444
445#[cfg(test)]
446mod tests {
447 use super::*;
448 use sva_samples::Buffer;
449
450 #[test]
452 fn an_entry_that_does_not_answer_what_was_asked_is_a_miss_and_goes() {
453 let cache = Cache::new();
454 let key = Hash(7, 11);
455 let stamp = Stamp {
456 tree: cache.begin_tree(),
457 fork: false,
458 slot: None,
459 };
460 let four = Payload::Segments(vec![Buffer::mono(8_000, vec![0.25; 4])]);
461 for (rate, width) in [(48_000, 1), (8_000, 2)] {
462 cache.store(key, &four, None, stamp);
463 let asked = Expected::Segments { rate, width };
464 assert!(cache.load(key, asked, stamp).is_none());
465 assert!(!cache.holds(key));
466 assert_eq!(cache.bytes(), 0);
467 }
468 }
469}