Skip to main content

sva_engine/cache/
store.rs

1// Concern: the one bounded store, what it evicts and the cap it never passes | Non-concern: what a render asks of it (stats.rs) | IO: (Hash) -> a payload + label
2
3use 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
13/// Samples between two states a run keeps, so a reader resumes from one at most this far back.
14pub const DEFAULT_MARK_EVERY: usize = 16_384;
15
16/// Which values a render stores; whatever is not stored is computed again when asked.
17#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
18pub enum CachePolicy {
19    #[default]
20    All,
21    /// The values two or more nodes read, and the render target.
22    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    /// Every entry the newest render neither stored nor read.
61    #[default]
62    Oldest,
63    /// Every entry whose node fewer than two nodes read.
64    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    /// A volatile node's value replaces the last one stored under the same slot.
87    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    /// Every entry `policy` names, oldest-read first, until `bytes` is at most `to`; then whole
152    /// trees oldest-first, until the cap holds.
153    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/// Every value a render computed and chose to keep, under its content hash. A hit is the value
198/// the cold run wrote, so what is kept or evicted decides only what is computed again. A clone
199/// is a handle on the same store, which a stream holds for as long as it plays.
200#[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    /// What a render stores where it names no policy of its own.
248    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    /// Evicts every entry `policy` names, and more where the cap still needs it.
265    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    /// How many samples apart a run keeps the states it passes.
293    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    /// A value's segments join those held under `key`, and a run continuing the one held
325    /// there extends it, each in place; anything else replaces what `key` held.
326    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
417/// A run that starts inside or at the end of the one held continues it: what it holds past
418/// that one's end is laid on, the samples both hold being the same.
419fn 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
424/// `more` laid among `parts`, each touching pair joined into one.
425pub(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    /// Only a colliding key reaches this, so no render can: the entry is a miss, and goes.
451    #[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}