1use alloc::vec::Vec;
2use core::fmt;
3use core::num::NonZeroUsize;
4
5use crate::ValueId;
6
7const SHAPE_TAG: u64 = 0x5348_4150_4500_0001;
8const CTC_TAG: u64 = 0x4354_4300_0000_0002;
9
10#[derive(Clone, Copy, Debug, PartialEq, Eq)]
12pub enum SpecializationError {
13 ValuesOutOfOrder,
14 TooManyWords,
15 AllocationFailed,
16}
17
18impl fmt::Display for SpecializationError {
19 fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
20 write!(formatter, "{self:?}")
21 }
22}
23
24#[derive(Clone, Debug, PartialEq, Eq)]
28pub struct SpecializationKey {
29 words: Vec<u64>,
30 fingerprint: [u64; 2],
31}
32
33impl SpecializationKey {
34 pub fn words(&self) -> &[u64] {
35 &self.words
36 }
37
38 pub const fn fingerprint(&self) -> [u64; 2] {
39 self.fingerprint
40 }
41}
42
43#[derive(Debug)]
48pub struct SpecializationKeyBuilder {
49 words: Vec<u64>,
50 max_words: usize,
51 last_value: Option<ValueId>,
52}
53
54impl SpecializationKeyBuilder {
55 pub fn new(max_words: usize) -> Self {
56 Self {
57 words: Vec::new(),
58 max_words,
59 last_value: None,
60 }
61 }
62
63 pub fn push_shape(
64 &mut self,
65 value: ValueId,
66 dimensions: &[u64],
67 ) -> Result<(), SpecializationError> {
68 self.begin_value(value, SHAPE_TAG, dimensions.len(), dimensions.len())?;
69 self.words.extend_from_slice(dimensions);
70 Ok(())
71 }
72
73 pub fn push_ctc(&mut self, value: ValueId, bytes: &[u8]) -> Result<(), SpecializationError> {
74 let payload_words = bytes.len().div_ceil(8);
75 self.begin_value(value, CTC_TAG, bytes.len(), payload_words)?;
76 for chunk in bytes.chunks(8) {
77 let mut packed = [0_u8; 8];
78 packed[..chunk.len()].copy_from_slice(chunk);
79 self.words.push(u64::from_le_bytes(packed));
80 }
81 Ok(())
82 }
83
84 pub fn finish(self) -> SpecializationKey {
85 SpecializationKey {
86 fingerprint: fingerprint(&self.words),
87 words: self.words,
88 }
89 }
90
91 fn begin_value(
92 &mut self,
93 value: ValueId,
94 tag: u64,
95 payload_len: usize,
96 payload_words: usize,
97 ) -> Result<(), SpecializationError> {
98 if self.last_value.is_some_and(|prior| prior >= value) {
99 return Err(SpecializationError::ValuesOutOfOrder);
100 }
101 let additional = 3_usize
102 .checked_add(payload_words)
103 .ok_or(SpecializationError::TooManyWords)?;
104 self.reserve(additional)?;
107 self.words.push(tag);
108 self.words.push(u64::from(value.get()));
109 self.words
110 .push(u64::try_from(payload_len).map_err(|_| SpecializationError::TooManyWords)?);
111 self.last_value = Some(value);
112 Ok(())
113 }
114
115 fn reserve(&mut self, additional: usize) -> Result<(), SpecializationError> {
116 if self
117 .words
118 .len()
119 .checked_add(additional)
120 .is_none_or(|required| required > self.max_words)
121 {
122 return Err(SpecializationError::TooManyWords);
123 }
124 self.words
125 .try_reserve_exact(additional)
126 .map_err(|_| SpecializationError::AllocationFailed)
127 }
128}
129
130fn fingerprint(words: &[u64]) -> [u64; 2] {
131 let mut first = 0xcbf2_9ce4_8422_2325_u64;
132 let mut second = 0x9e37_79b9_7f4a_7c15_u64;
133 for &word in words {
134 first ^= word;
135 first = first.wrapping_mul(0x0000_0100_0000_01b3);
136 second ^= word.wrapping_add(first.rotate_left(17));
137 second = second.rotate_left(27).wrapping_mul(0x94d0_49bb_1331_11eb);
138 }
139 [first, second ^ words.len() as u64]
140}
141
142#[derive(Debug)]
143struct CacheEntry<V> {
144 key: SpecializationKey,
145 value: V,
146 last_used: u64,
147}
148
149#[derive(Debug)]
151pub struct CacheInsertError<V> {
152 pub key: SpecializationKey,
153 pub value: V,
154}
155
156#[derive(Debug)]
161pub struct SpecializationCache<V> {
162 capacity: NonZeroUsize,
163 clock: u64,
164 entries: Vec<CacheEntry<V>>,
165}
166
167impl<V> SpecializationCache<V> {
168 pub fn new(capacity: NonZeroUsize) -> Self {
169 Self {
170 capacity,
171 clock: 0,
172 entries: Vec::new(),
173 }
174 }
175
176 pub const fn capacity(&self) -> usize {
177 self.capacity.get()
178 }
179
180 pub fn len(&self) -> usize {
181 self.entries.len()
182 }
183
184 pub fn is_empty(&self) -> bool {
185 self.entries.is_empty()
186 }
187
188 pub fn get(&mut self, key: &SpecializationKey) -> Option<&V> {
189 let index = self.position(key)?;
190 let stamp = self.tick();
191 self.entries[index].last_used = stamp;
192 Some(&self.entries[index].value)
193 }
194
195 pub fn get_mut(&mut self, key: &SpecializationKey) -> Option<&mut V> {
196 let index = self.position(key)?;
197 let stamp = self.tick();
198 self.entries[index].last_used = stamp;
199 Some(&mut self.entries[index].value)
200 }
201
202 pub fn insert(
204 &mut self,
205 key: SpecializationKey,
206 value: V,
207 ) -> Result<Option<V>, CacheInsertError<V>> {
208 let stamp = self.tick();
209 if let Some(index) = self.position(&key) {
210 self.entries[index].last_used = stamp;
211 return Ok(Some(core::mem::replace(
212 &mut self.entries[index].value,
213 value,
214 )));
215 }
216 if self.entries.len() == self.capacity.get() {
217 let index = self
218 .entries
219 .iter()
220 .enumerate()
221 .min_by_key(|(_, entry)| entry.last_used)
222 .map(|(index, _)| index)
223 .expect("nonzero full cache");
224 let evicted = self.entries.swap_remove(index).value;
225 self.entries.push(CacheEntry {
226 key,
227 value,
228 last_used: stamp,
229 });
230 return Ok(Some(evicted));
231 }
232 if self.entries.try_reserve_exact(1).is_err() {
233 return Err(CacheInsertError { key, value });
234 }
235 self.entries.push(CacheEntry {
236 key,
237 value,
238 last_used: stamp,
239 });
240 Ok(None)
241 }
242
243 fn position(&self, key: &SpecializationKey) -> Option<usize> {
244 self.entries.iter().position(|entry| {
245 entry.key.fingerprint == key.fingerprint && entry.key.words == key.words
246 })
247 }
248
249 fn tick(&mut self) -> u64 {
250 if self.clock == u64::MAX {
251 self.entries.sort_unstable_by_key(|entry| entry.last_used);
252 for (index, entry) in self.entries.iter_mut().enumerate() {
253 entry.last_used = index as u64;
254 }
255 self.clock = self.entries.len() as u64;
256 }
257 let stamp = self.clock;
258 self.clock += 1;
259 stamp
260 }
261}
262
263#[cfg(test)]
264mod tests {
265 use super::*;
266
267 fn value(raw: u32) -> ValueId {
268 ValueId::from_raw(raw)
269 }
270
271 fn key(raw: u32) -> SpecializationKey {
272 let mut builder = SpecializationKeyBuilder::new(16);
273 builder
274 .push_shape(value(raw), &[1, u64::from(raw)])
275 .unwrap();
276 builder.finish()
277 }
278
279 #[test]
280 fn keys_are_canonical_bounded_and_collision_safe() {
281 let mut first = SpecializationKeyBuilder::new(16);
282 first.push_shape(value(1), &[1, 2, 3]).unwrap();
283 first
284 .push_ctc(value(2), &[1, 2, 3, 4, 5, 6, 7, 8, 9])
285 .unwrap();
286 let first = first.finish();
287
288 let mut second = SpecializationKeyBuilder::new(16);
289 second.push_shape(value(1), &[1, 2, 3]).unwrap();
290 second
291 .push_ctc(value(2), &[1, 2, 3, 4, 5, 6, 7, 8, 9])
292 .unwrap();
293 assert_eq!(first, second.finish());
294
295 let mut invalid = SpecializationKeyBuilder::new(8);
296 invalid.push_shape(value(2), &[1]).unwrap();
297 assert_eq!(
298 invalid.push_shape(value(1), &[1]),
299 Err(SpecializationError::ValuesOutOfOrder)
300 );
301
302 let mut retry = SpecializationKeyBuilder::new(4);
303 assert_eq!(
304 retry.push_shape(value(3), &[1, 2]),
305 Err(SpecializationError::TooManyWords)
306 );
307 retry.push_shape(value(3), &[1]).unwrap();
308 assert_eq!(retry.finish().words().len(), 4);
309 }
310
311 #[test]
312 fn cache_replaces_and_evicts_the_least_recently_used_value() {
313 let mut cache = SpecializationCache::new(NonZeroUsize::new(2).unwrap());
314 let one = key(1);
315 let two = key(2);
316 let three = key(3);
317 assert_eq!(cache.insert(one.clone(), 10).unwrap(), None);
318 assert_eq!(cache.insert(two.clone(), 20).unwrap(), None);
319 assert_eq!(cache.get(&one), Some(&10));
320 assert_eq!(cache.insert(three.clone(), 30).unwrap(), Some(20));
321 assert_eq!(cache.get(&two), None);
322 assert_eq!(cache.get(&three), Some(&30));
323 assert_eq!(cache.insert(one.clone(), 11).unwrap(), Some(10));
324 assert_eq!(cache.get(&one), Some(&11));
325 }
326}