1use std::collections::{BTreeMap, HashMap};
2
3use crate::{CacheBlobStore, CacheDedupeStats, ExactStatePayload};
4
5#[derive(Debug)]
6pub struct ExactStateCache<E> {
7 max_entries: usize,
8 max_bytes: u64,
9 clock: u64,
10 logical_bytes: u64,
11 blobs: CacheBlobStore,
12 entries: HashMap<String, ExactStateEntry<E>>,
13 token_count_refs: BTreeMap<u64, usize>,
14}
15
16#[derive(Debug, Clone)]
17struct ExactStateEntry<E> {
18 token_count: u64,
19 logical_bytes: u64,
20 last_used: u64,
21 payload: ExactStatePayload,
22 extra: E,
23}
24
25#[derive(Debug, Clone)]
26pub struct ExactStateLookup<E> {
27 pub page_id: String,
28 pub token_count: u64,
29 pub logical_bytes: u64,
30 pub entries: usize,
31 pub payload: ExactStatePayload,
32 pub extra: E,
33}
34
35#[derive(Debug, Clone)]
36pub struct ExactStateRecordOutcome {
37 pub stored: bool,
38 pub page_id: String,
39 pub token_count: u64,
40 pub logical_bytes: u64,
41 pub physical_bytes: u64,
42 pub entries: usize,
43 pub evicted_entries: usize,
44 pub evicted_logical_bytes: u64,
45 pub dedupe: CacheDedupeStats,
46}
47
48#[derive(Debug, Clone, Copy, Default)]
49pub struct ExactStateCacheStats {
50 pub entries: usize,
51 pub logical_bytes: u64,
52 pub physical_bytes: u64,
53 pub block_count: usize,
54 pub max_entries: usize,
55 pub max_bytes: u64,
56}
57
58impl<E: Clone> ExactStateCache<E> {
59 pub fn new(max_entries: usize, max_bytes: u64) -> Self {
60 Self {
61 max_entries: max_entries.max(1),
62 max_bytes,
63 clock: 0,
64 logical_bytes: 0,
65 blobs: CacheBlobStore::default(),
66 entries: HashMap::new(),
67 token_count_refs: BTreeMap::new(),
68 }
69 }
70
71 pub fn lookup(&mut self, page_id: &str) -> Option<ExactStateLookup<E>> {
72 self.clock = self.clock.saturating_add(1);
73 let entries = self.entries.len();
74 let entry = self.entries.get_mut(page_id)?;
75 entry.last_used = self.clock;
76 Some(ExactStateLookup {
77 page_id: page_id.to_string(),
78 token_count: entry.token_count,
79 logical_bytes: entry.logical_bytes,
80 entries,
81 payload: entry.payload.clone(),
82 extra: entry.extra.clone(),
83 })
84 }
85
86 pub fn token_counts_at_most(&self, max_token_count: u64) -> Vec<u64> {
87 self.token_count_refs
88 .range(..=max_token_count)
89 .rev()
90 .map(|(token_count, _)| *token_count)
91 .collect()
92 }
93
94 pub fn record(
95 &mut self,
96 page_id: String,
97 token_count: u64,
98 payload: ExactStatePayload,
99 extra: E,
100 ) -> ExactStateRecordOutcome {
101 self.clock = self.clock.saturating_add(1);
102 if let Some(previous) = self.entries.remove(&page_id) {
103 self.remove_entry(previous);
104 }
105
106 let logical_bytes = payload.byte_len();
107 let (payload, dedupe) = payload.dedupe_into(&mut self.blobs);
108 self.logical_bytes = self.logical_bytes.saturating_add(logical_bytes);
109 self.add_token_count(token_count);
110 self.entries.insert(
111 page_id.clone(),
112 ExactStateEntry {
113 token_count,
114 logical_bytes,
115 last_used: self.clock,
116 payload,
117 extra,
118 },
119 );
120
121 let (mut evicted_entries, mut evicted_logical_bytes) = self.evict_until_within_limits();
122 if self.max_bytes > 0
123 && self.blobs.physical_bytes() > self.max_bytes
124 && let Some(entry) = self.entries.remove(&page_id)
125 {
126 evicted_entries = evicted_entries.saturating_add(1);
127 evicted_logical_bytes = evicted_logical_bytes.saturating_add(entry.logical_bytes);
128 self.remove_entry(entry);
129 }
130 let stored = self.entries.contains_key(&page_id);
131 let stats = self.stats();
132 ExactStateRecordOutcome {
133 stored,
134 page_id,
135 token_count,
136 logical_bytes,
137 physical_bytes: stats.physical_bytes,
138 entries: stats.entries,
139 evicted_entries,
140 evicted_logical_bytes,
141 dedupe,
142 }
143 }
144
145 pub fn stats(&self) -> ExactStateCacheStats {
146 ExactStateCacheStats {
147 entries: self.entries.len(),
148 logical_bytes: self.logical_bytes,
149 physical_bytes: self.blobs.physical_bytes(),
150 block_count: self.blobs.block_count(),
151 max_entries: self.max_entries,
152 max_bytes: self.max_bytes,
153 }
154 }
155
156 fn evict_until_within_limits(&mut self) -> (usize, u64) {
157 let mut evicted_entries = 0usize;
158 let mut evicted_logical_bytes = 0u64;
159 loop {
160 let over_entries = self.entries.len() > self.max_entries;
161 let over_bytes = self.max_bytes > 0 && self.blobs.physical_bytes() > self.max_bytes;
162 if !over_entries && !over_bytes {
163 break;
164 }
165 let Some(victim) = self
166 .entries
167 .iter()
168 .min_by_key(|(_, entry)| entry.last_used)
169 .map(|(page_id, _)| page_id.clone())
170 else {
171 break;
172 };
173 if let Some(entry) = self.entries.remove(&victim) {
174 evicted_entries = evicted_entries.saturating_add(1);
175 evicted_logical_bytes = evicted_logical_bytes.saturating_add(entry.logical_bytes);
176 self.remove_entry(entry);
177 }
178 }
179 (evicted_entries, evicted_logical_bytes)
180 }
181
182 fn remove_entry(&mut self, entry: ExactStateEntry<E>) {
183 self.logical_bytes = self.logical_bytes.saturating_sub(entry.logical_bytes);
184 self.remove_token_count(entry.token_count);
185 entry.payload.release_from(&mut self.blobs);
186 }
187
188 fn add_token_count(&mut self, token_count: u64) {
189 *self.token_count_refs.entry(token_count).or_default() += 1;
190 }
191
192 fn remove_token_count(&mut self, token_count: u64) {
193 let Some(count) = self.token_count_refs.get_mut(&token_count) else {
194 debug_assert!(false, "exact-state token-count index drifted");
195 return;
196 };
197 *count = count.saturating_sub(1);
198 if *count == 0 {
199 self.token_count_refs.remove(&token_count);
200 }
201 }
202}
203
204#[cfg(test)]
205mod tests {
206 use crate::{ExactStatePayload, exact_state::ExactStateCache};
207
208 #[test]
209 fn exact_state_cache_evicts_lru_by_entry_cap() {
210 let mut cache = ExactStateCache::new(1, 0);
211 cache.record(
212 "first".to_string(),
213 2,
214 ExactStatePayload::full_state(vec![1, 2]),
215 (),
216 );
217 cache.record(
218 "second".to_string(),
219 2,
220 ExactStatePayload::full_state(vec![3, 4]),
221 (),
222 );
223
224 assert!(cache.lookup("first").is_none());
225 assert!(cache.lookup("second").is_some());
226 assert_eq!(cache.stats().entries, 1);
227 }
228
229 #[test]
230 fn cached_token_counts_are_bounded_sorted_and_deduplicated() {
231 let mut cache = ExactStateCache::new(4, 0);
232 for (page_id, token_count) in [("a", 96), ("b", 160), ("c", 96), ("d", 224)] {
233 cache.record(
234 page_id.to_string(),
235 token_count,
236 ExactStatePayload::full_state(vec![1]),
237 (),
238 );
239 }
240
241 assert_eq!(cache.token_counts_at_most(200), vec![160, 96]);
242 }
243
244 #[test]
245 fn cached_token_counts_track_replacement_and_eviction() {
246 let mut cache = ExactStateCache::new(1, 0);
247 cache.record(
248 "first".to_string(),
249 96,
250 ExactStatePayload::full_state(vec![1]),
251 (),
252 );
253 cache.record(
254 "first".to_string(),
255 160,
256 ExactStatePayload::full_state(vec![2]),
257 (),
258 );
259 assert_eq!(cache.token_counts_at_most(160), vec![160]);
260
261 cache.record(
262 "second".to_string(),
263 224,
264 ExactStatePayload::full_state(vec![3]),
265 (),
266 );
267 assert_eq!(cache.token_counts_at_most(224), vec![224]);
268 }
269
270 #[test]
271 fn exact_state_cache_releases_deduped_blocks_on_eviction() {
272 let mut cache = ExactStateCache::new(1, 0);
273 cache.record(
274 "first".to_string(),
275 8,
276 ExactStatePayload::full_state(vec![7; 1024 * 1024]),
277 (),
278 );
279 cache.record(
280 "second".to_string(),
281 8,
282 ExactStatePayload::full_state(vec![7; 1024 * 1024]),
283 (),
284 );
285
286 assert_eq!(cache.stats().entries, 1);
287 assert_eq!(cache.stats().physical_bytes, 1024 * 1024);
288 assert_eq!(cache.stats().block_count, 1);
289 }
290}