1use crate::error::{Error, Result};
11
12pub const ALPHABET: usize = 256;
14pub const MODEL_VERSION_1: u8 = 1;
16pub const MIN_SCALE_BITS: u8 = 1;
18pub const MAX_SCALE_BITS: u8 = 15;
19
20const MODEL_HEADER_LEN: usize = 4;
22const MODEL_ENCODED_LEN: usize = MODEL_HEADER_LEN + ALPHABET * 2;
24
25#[derive(Debug, Clone, PartialEq, Eq)]
27pub struct EntropyModel {
28 pub scale_bits: u8,
30 pub frequencies: Vec<u32>,
32}
33
34impl EntropyModel {
35 pub fn total(&self) -> u32 {
37 1u32 << self.scale_bits
38 }
39
40 fn frequency_sum(&self) -> u64 {
43 self.frequencies.iter().map(|&f| u64::from(f)).sum()
44 }
45
46 pub fn from_counts(counts: &[u64; ALPHABET], scale_bits: u8) -> Result<EntropyModel> {
60 if !(MIN_SCALE_BITS..=MAX_SCALE_BITS).contains(&scale_bits) {
61 let (min, max) = (MIN_SCALE_BITS, MAX_SCALE_BITS);
62 return Err(Error::invalid_model(format!(
63 "scale_bits {scale_bits} outside {min}..={max}"
64 )));
65 }
66 let target: u32 = 1u32 << scale_bits;
67
68 let present: Vec<usize> = (0..ALPHABET).filter(|&i| counts[i] > 0).collect();
70 if present.is_empty() {
71 return Self::uniform(scale_bits);
72 }
73
74 if present.len() as u64 > u64::from(target) {
76 return Err(Error::invalid_model(format!(
77 "{} present symbols exceed target {target}",
78 present.len()
79 )));
80 }
81
82 let mut frequencies = vec![0u32; ALPHABET];
84 for &i in &present {
85 frequencies[i] = 1;
86 }
87 let remaining: u32 = target - present.len() as u32;
88
89 if remaining > 0 {
90 let total: u128 = present.iter().map(|&i| u128::from(counts[i])).sum();
91
92 let mut quota_sum: u64 = 0;
94 let mut rems: Vec<(usize, u128)> = Vec::with_capacity(present.len());
95 for &i in &present {
96 let product = u128::from(counts[i]) * u128::from(remaining);
97 let quota = product / total;
98 let rem = product % total;
99 frequencies[i] += u32::try_from(quota)
100 .map_err(|_| Error::internal_invariant("quota exceeded 32-bit range"))?;
101 quota_sum += u64::try_from(quota)
102 .map_err(|_| Error::internal_invariant("quota exceeded 64-bit range"))?;
103 rems.push((i, rem));
104 }
105 let leftover: u64 = u64::from(remaining) - quota_sum;
106
107 rems.sort_by(|a, b| b.1.cmp(&a.1).then_with(|| a.0.cmp(&b.0)));
111 let mut added: u64 = 0;
112 for &(i, _) in &rems {
113 if added >= leftover {
114 break;
115 }
116 frequencies[i] += 1;
117 added += 1;
118 }
119
120 while added < leftover {
125 let i = present
126 .iter()
127 .copied()
128 .max_by_key(|&i| (counts[i], core::cmp::Reverse(i)))
129 .expect("present is non-empty");
130 frequencies[i] += 1;
131 added += 1;
132 }
133 }
134
135 let model = EntropyModel {
136 scale_bits,
137 frequencies,
138 };
139 if model.frequency_sum() != u64::from(target) {
140 return Err(Error::internal_invariant(
141 "normalized frequencies do not sum to target",
142 ));
143 }
144 Ok(model)
145 }
146
147 pub fn uniform(scale_bits: u8) -> Result<EntropyModel> {
153 if !(MIN_SCALE_BITS..=MAX_SCALE_BITS).contains(&scale_bits) {
154 let (min, max) = (MIN_SCALE_BITS, MAX_SCALE_BITS);
155 return Err(Error::invalid_model(format!(
156 "scale_bits {scale_bits} outside {min}..={max}"
157 )));
158 }
159 let target: u32 = 1u32 << scale_bits;
160 if target < ALPHABET as u32 {
161 return Err(Error::invalid_model(format!(
162 "scale_bits {scale_bits} cannot hold {ALPHABET} equal frequencies"
163 )));
164 }
165 let per = target / ALPHABET as u32;
167 Ok(EntropyModel {
168 scale_bits,
169 frequencies: vec![per; ALPHABET],
170 })
171 }
172
173 pub fn encode(&self) -> Result<Vec<u8>> {
179 if !(MIN_SCALE_BITS..=MAX_SCALE_BITS).contains(&self.scale_bits) {
180 let (min, max) = (MIN_SCALE_BITS, MAX_SCALE_BITS);
181 return Err(Error::invalid_model(format!(
182 "scale_bits {} outside {min}..={max}",
183 self.scale_bits
184 )));
185 }
186 if self.frequencies.len() != ALPHABET {
187 return Err(Error::invalid_model(format!(
188 "expected {ALPHABET} frequencies, got {}",
189 self.frequencies.len()
190 )));
191 }
192 let target = u64::from(self.total());
193 if self.frequency_sum() != target {
194 return Err(Error::invalid_model(
195 "frequencies do not sum to 1 << scale_bits",
196 ));
197 }
198
199 let mut out = Vec::with_capacity(MODEL_ENCODED_LEN);
200 out.push(MODEL_VERSION_1);
201 out.push(self.scale_bits);
202 out.extend_from_slice(&(ALPHABET as u16).to_le_bytes());
203 for &f in &self.frequencies {
204 let f =
205 u16::try_from(f).map_err(|_| Error::invalid_model("frequency does not fit u16"))?;
206 out.extend_from_slice(&f.to_le_bytes());
207 }
208 Ok(out)
209 }
210
211 pub fn decode(bytes: &[u8]) -> Result<EntropyModel> {
213 if bytes.len() < MODEL_HEADER_LEN {
214 return Err(Error::invalid_model("model shorter than header"));
215 }
216 let version = bytes[0];
217 if version != MODEL_VERSION_1 {
218 return Err(Error::unsupported_version(format!(
219 "model version {version}, expected {MODEL_VERSION_1}"
220 )));
221 }
222 let scale_bits = bytes[1];
223 if !(MIN_SCALE_BITS..=MAX_SCALE_BITS).contains(&scale_bits) {
224 let (min, max) = (MIN_SCALE_BITS, MAX_SCALE_BITS);
225 return Err(Error::invalid_model(format!(
226 "scale_bits {scale_bits} outside {min}..={max}"
227 )));
228 }
229 let count = u16::from_le_bytes([bytes[2], bytes[3]]);
230 if count as usize != ALPHABET {
231 return Err(Error::invalid_model(format!(
232 "declared count {count}, expected {ALPHABET}"
233 )));
234 }
235 if bytes.len() != MODEL_ENCODED_LEN {
236 return Err(Error::invalid_model(
237 "model length does not match declared count (trailing or truncated)",
238 ));
239 }
240
241 let target = u64::from(1u32 << scale_bits);
242 let mut frequencies = Vec::with_capacity(ALPHABET);
243 let mut sum: u64 = 0;
244 let (chunks, _rest) = bytes[MODEL_HEADER_LEN..].as_chunks::<2>();
245 for chunk in chunks {
246 let f = u32::from(u16::from_le_bytes(*chunk));
247 sum += u64::from(f);
248 frequencies.push(f);
249 }
250 if sum != target {
251 return Err(Error::invalid_model(
252 "frequencies do not sum to 1 << scale_bits",
253 ));
254 }
255 Ok(EntropyModel {
256 scale_bits,
257 frequencies,
258 })
259 }
260}
261
262#[cfg(test)]
263mod tests {
264 use super::*;
265
266 fn counts_with(pairs: &[(usize, u64)]) -> [u64; ALPHABET] {
267 let mut counts = [0u64; ALPHABET];
268 for &(i, v) in pairs {
269 counts[i] = v;
270 }
271 counts
272 }
273
274 #[test]
275 fn uniform_sums_to_target() {
276 for bits in [8u8, 12, 15] {
277 let model = EntropyModel::uniform(bits).expect("uniform must succeed");
278 assert_eq!(model.frequencies.len(), ALPHABET);
279 assert_eq!(model.frequency_sum(), u64::from(1u32 << bits));
280 let per = model.frequencies[0];
281 assert!(model.frequencies.iter().all(|&f| f == per));
282 }
283 }
284
285 #[test]
286 fn from_counts_sums_exactly() {
287 let inputs: &[&[(usize, u64)]] = &[
288 &[(0, 1)],
289 &[(0, 1), (1, 1)],
290 &[(0, 300), (1, 1)],
291 &[(0, 1), (1, 2), (2, 3), (3, 4)],
292 &[(5, 10), (200, 1), (255, 7)],
293 &[(0, u64::MAX), (255, 1)],
294 ];
295 for bits in [8u8, 12] {
296 for input in inputs {
297 let counts = counts_with(input);
298 let model =
299 EntropyModel::from_counts(&counts, bits).expect("normalization must succeed");
300 assert_eq!(model.scale_bits, bits);
301 assert_eq!(model.frequency_sum(), u64::from(1u32 << bits));
302 for &(i, _) in *input {
303 assert!(model.frequencies[i] >= 1);
304 }
305 }
306 }
307 }
308
309 #[test]
310 fn determinism() {
311 let counts = counts_with(&[(0, 17), (3, 5), (7, 250), (255, 1)]);
312 let a = EntropyModel::from_counts(&counts, 12)
313 .expect("ok")
314 .encode()
315 .expect("ok");
316 let b = EntropyModel::from_counts(&counts, 12)
317 .expect("ok")
318 .encode()
319 .expect("ok");
320 assert_eq!(a, b);
321 }
322
323 #[test]
324 fn present_symbols_get_at_least_one() {
325 let counts = [1u64; ALPHABET];
328 let model = EntropyModel::from_counts(&counts, 12).expect("ok");
329 assert!(model.frequencies.iter().all(|&f| f >= 1));
330 assert_eq!(model.frequency_sum(), 4096);
331 }
332
333 #[test]
334 fn zero_symbols_get_zero() {
335 let counts = counts_with(&[(0, 5), (10, 9), (255, 3)]);
336 let model = EntropyModel::from_counts(&counts, 12).expect("ok");
337 for (i, &f) in model.frequencies.iter().enumerate() {
338 if counts[i] == 0 {
339 assert_eq!(f, 0, "symbol {i} had zero count but nonzero frequency");
340 }
341 }
342 }
343
344 #[test]
345 fn alphabet_too_big_errors() {
346 let counts = [1u64; ALPHABET];
347 for bits in 1u8..=7 {
348 assert!(EntropyModel::from_counts(&counts, bits).is_err());
349 }
350 }
351
352 #[test]
353 fn roundtrip_encode_decode() {
354 let counts = counts_with(&[(0, 1), (1, 2), (2, 3), (100, 400), (255, 7)]);
355 for bits in [8u8, 12, 15] {
356 let model = EntropyModel::from_counts(&counts, bits).expect("ok");
357 let bytes = model.encode().expect("ok");
358 assert_eq!(bytes.len(), MODEL_ENCODED_LEN);
359 let decoded = EntropyModel::decode(&bytes).expect("ok");
360 assert_eq!(decoded, model);
361 }
362 }
363
364 #[test]
365 fn decode_rejects_bad_sum() {
366 let mut bytes = EntropyModel::uniform(8).expect("ok").encode().expect("ok");
367 let last = bytes.len();
369 let value = u16::from_le_bytes([bytes[last - 2], bytes[last - 1]]);
370 let bumped = value.checked_add(1).expect("fits");
371 bytes[last - 2..].copy_from_slice(&bumped.to_le_bytes());
372 assert!(EntropyModel::decode(&bytes).is_err());
373 }
374
375 #[test]
376 fn decode_rejects_bad_version() {
377 let mut bytes = EntropyModel::uniform(8).expect("ok").encode().expect("ok");
378 bytes[0] = 2;
379 let err = EntropyModel::decode(&bytes).expect_err("must reject");
380 assert_eq!(err.class(), crate::error::ErrorClass::UnsupportedVersion);
381 }
382
383 #[test]
384 fn decode_rejects_trailing() {
385 let mut bytes = EntropyModel::uniform(8).expect("ok").encode().expect("ok");
386 bytes.push(0);
387 assert!(EntropyModel::decode(&bytes).is_err());
388 }
389
390 #[test]
391 fn all_zero_is_uniform() {
392 let zero = [0u64; ALPHABET];
393 for bits in [8u8, 12, 15] {
394 let model = EntropyModel::from_counts(&zero, bits).expect("ok");
395 assert_eq!(model, EntropyModel::uniform(bits).expect("ok"));
396 }
397 }
398}