1use crate::error::{Error, Result};
11
12pub const ALPHABET: usize = 256;
14pub const MODEL_VERSION_1: u8 = 1;
16pub const MODEL_VERSION_2: u8 = 2;
18pub const MIN_SCALE_BITS: u8 = 1;
20pub const MAX_SCALE_BITS: u8 = 15;
21
22const MODEL_FORM_SPARSE: u8 = 0;
24const MODEL_FORM_DENSE: u8 = 1;
26
27const MODEL_V1_HEADER_LEN: usize = 4;
29const MODEL_V1_ENCODED_LEN: usize = MODEL_V1_HEADER_LEN + ALPHABET * 2;
31const MODEL_V2_HEADER_LEN: usize = 3;
33const MODEL_V2_COUNT_LEN: usize = 2;
35const MODEL_V2_SPARSE_ENTRY_LEN: usize = 3;
37const MODEL_V2_DENSE_LEN: usize = MODEL_V2_HEADER_LEN + MODEL_V2_COUNT_LEN + ALPHABET * 2;
39
40#[derive(Debug, Clone, PartialEq, Eq)]
42pub struct EntropyModel {
43 pub scale_bits: u8,
45 pub frequencies: Vec<u32>,
47}
48
49impl EntropyModel {
50 pub fn total(&self) -> u32 {
52 1u32 << self.scale_bits
53 }
54
55 fn frequency_sum(&self) -> u64 {
58 self.frequencies.iter().map(|&f| u64::from(f)).sum()
59 }
60
61 pub fn from_counts(counts: &[u64; ALPHABET], scale_bits: u8) -> Result<EntropyModel> {
75 if !(MIN_SCALE_BITS..=MAX_SCALE_BITS).contains(&scale_bits) {
76 let (min, max) = (MIN_SCALE_BITS, MAX_SCALE_BITS);
77 return Err(Error::invalid_model(format!(
78 "scale_bits {scale_bits} outside {min}..={max}"
79 )));
80 }
81 let target: u32 = 1u32 << scale_bits;
82
83 let present: Vec<usize> = (0..ALPHABET).filter(|&i| counts[i] > 0).collect();
85 if present.is_empty() {
86 return Self::uniform(scale_bits);
87 }
88
89 if present.len() as u64 > u64::from(target) {
91 return Err(Error::invalid_model(format!(
92 "{} present symbols exceed target {target}",
93 present.len()
94 )));
95 }
96
97 let mut frequencies = vec![0u32; ALPHABET];
99 for &i in &present {
100 frequencies[i] = 1;
101 }
102 let remaining: u32 = target - present.len() as u32;
103
104 if remaining > 0 {
105 let total: u128 = present.iter().map(|&i| u128::from(counts[i])).sum();
106
107 let mut quota_sum: u64 = 0;
109 let mut rems: Vec<(usize, u128)> = Vec::with_capacity(present.len());
110 for &i in &present {
111 let product = u128::from(counts[i]) * u128::from(remaining);
112 let quota = product / total;
113 let rem = product % total;
114 frequencies[i] += u32::try_from(quota)
115 .map_err(|_| Error::internal_invariant("quota exceeded 32-bit range"))?;
116 quota_sum += u64::try_from(quota)
117 .map_err(|_| Error::internal_invariant("quota exceeded 64-bit range"))?;
118 rems.push((i, rem));
119 }
120 let leftover: u64 = u64::from(remaining) - quota_sum;
121
122 rems.sort_by(|a, b| b.1.cmp(&a.1).then_with(|| a.0.cmp(&b.0)));
126 let mut added: u64 = 0;
127 for &(i, _) in &rems {
128 if added >= leftover {
129 break;
130 }
131 frequencies[i] += 1;
132 added += 1;
133 }
134
135 while added < leftover {
140 let i = present
141 .iter()
142 .copied()
143 .max_by_key(|&i| (counts[i], core::cmp::Reverse(i)))
144 .expect("present is non-empty");
145 frequencies[i] += 1;
146 added += 1;
147 }
148 }
149
150 let model = EntropyModel {
151 scale_bits,
152 frequencies,
153 };
154 if model.frequency_sum() != u64::from(target) {
155 return Err(Error::internal_invariant(
156 "normalized frequencies do not sum to target",
157 ));
158 }
159 Ok(model)
160 }
161
162 pub fn uniform(scale_bits: u8) -> Result<EntropyModel> {
168 if !(MIN_SCALE_BITS..=MAX_SCALE_BITS).contains(&scale_bits) {
169 let (min, max) = (MIN_SCALE_BITS, MAX_SCALE_BITS);
170 return Err(Error::invalid_model(format!(
171 "scale_bits {scale_bits} outside {min}..={max}"
172 )));
173 }
174 let target: u32 = 1u32 << scale_bits;
175 if target < ALPHABET as u32 {
176 return Err(Error::invalid_model(format!(
177 "scale_bits {scale_bits} cannot hold {ALPHABET} equal frequencies"
178 )));
179 }
180 let per = target / ALPHABET as u32;
182 Ok(EntropyModel {
183 scale_bits,
184 frequencies: vec![per; ALPHABET],
185 })
186 }
187
188 fn validate(&self) -> Result<()> {
190 Self::validate_scale_bits(self.scale_bits)?;
191 if self.frequencies.len() != ALPHABET {
192 return Err(Error::invalid_model(format!(
193 "expected {ALPHABET} frequencies, got {}",
194 self.frequencies.len()
195 )));
196 }
197 if self.frequency_sum() != u64::from(self.total()) {
198 return Err(Error::invalid_model(
199 "frequencies do not sum to 1 << scale_bits",
200 ));
201 }
202 Ok(())
203 }
204
205 fn validate_scale_bits(scale_bits: u8) -> Result<()> {
207 if !(MIN_SCALE_BITS..=MAX_SCALE_BITS).contains(&scale_bits) {
208 let (min, max) = (MIN_SCALE_BITS, MAX_SCALE_BITS);
209 return Err(Error::invalid_model(format!(
210 "scale_bits {scale_bits} outside {min}..={max}"
211 )));
212 }
213 Ok(())
214 }
215
216 pub fn encode(&self) -> Result<Vec<u8>> {
227 self.validate()?;
228
229 let present: Vec<(u8, u16)> = self
231 .frequencies
232 .iter()
233 .enumerate()
234 .filter(|&(_, &f)| f > 0)
235 .map(|(i, &f)| {
236 let f = u16::try_from(f)
237 .map_err(|_| Error::invalid_model("frequency does not fit u16"))?;
238 Ok((i as u8, f))
239 })
240 .collect::<Result<Vec<_>>>()?;
241 let sparse_len = MODEL_V2_HEADER_LEN + MODEL_V2_COUNT_LEN + present.len() * 3;
242
243 if sparse_len < MODEL_V2_DENSE_LEN {
244 let mut out = Vec::with_capacity(sparse_len);
245 out.push(MODEL_VERSION_2);
246 out.push(MODEL_FORM_SPARSE);
247 out.push(self.scale_bits);
248 let count = u16::try_from(present.len())
249 .map_err(|_| Error::invalid_model("present count does not fit u16"))?;
250 out.extend_from_slice(&count.to_le_bytes());
251 for (symbol, freq) in present {
252 out.push(symbol);
253 out.extend_from_slice(&freq.to_le_bytes());
254 }
255 Ok(out)
256 } else {
257 let mut out = Vec::with_capacity(MODEL_V2_DENSE_LEN);
258 out.push(MODEL_VERSION_2);
259 out.push(MODEL_FORM_DENSE);
260 out.push(self.scale_bits);
261 out.extend_from_slice(&(ALPHABET as u16).to_le_bytes());
262 for &f in &self.frequencies {
263 let f = u16::try_from(f)
264 .map_err(|_| Error::invalid_model("frequency does not fit u16"))?;
265 out.extend_from_slice(&f.to_le_bytes());
266 }
267 Ok(out)
268 }
269 }
270
271 pub fn decode(bytes: &[u8]) -> Result<EntropyModel> {
278 let Some(&version) = bytes.first() else {
279 return Err(Error::invalid_model("empty model"));
280 };
281 match version {
282 MODEL_VERSION_1 => Self::decode_v1(bytes),
283 MODEL_VERSION_2 => Self::decode_v2(bytes),
284 other => Err(Error::unsupported_version(format!(
285 "model version {other}, expected {MODEL_VERSION_1} or {MODEL_VERSION_2}"
286 ))),
287 }
288 }
289
290 fn decode_v1(bytes: &[u8]) -> Result<EntropyModel> {
292 if bytes.len() < MODEL_V1_HEADER_LEN {
293 return Err(Error::invalid_model("model shorter than v1 header"));
294 }
295 let scale_bits = bytes[1];
296 Self::validate_scale_bits(scale_bits)?;
297 let count = u16::from_le_bytes([bytes[2], bytes[3]]);
298 if count as usize != ALPHABET {
299 return Err(Error::invalid_model(format!(
300 "v1 declared count {count}, expected {ALPHABET}"
301 )));
302 }
303 if bytes.len() != MODEL_V1_ENCODED_LEN {
304 return Err(Error::invalid_model(
305 "v1 model length does not match declared count (trailing or truncated)",
306 ));
307 }
308
309 let target = u64::from(1u32 << scale_bits);
310 let mut frequencies = Vec::with_capacity(ALPHABET);
311 let mut sum: u64 = 0;
312 let (chunks, _rest) = bytes[MODEL_V1_HEADER_LEN..].as_chunks::<2>();
313 for chunk in chunks {
314 let f = u32::from(u16::from_le_bytes(*chunk));
315 sum += u64::from(f);
316 frequencies.push(f);
317 }
318 if sum != target {
319 return Err(Error::invalid_model(
320 "v1 frequencies do not sum to 1 << scale_bits",
321 ));
322 }
323 Ok(EntropyModel {
324 scale_bits,
325 frequencies,
326 })
327 }
328
329 fn decode_v2(bytes: &[u8]) -> Result<EntropyModel> {
331 if bytes.len() < MODEL_V2_HEADER_LEN {
332 return Err(Error::invalid_model("model shorter than v2 header"));
333 }
334 let form = bytes[1];
335 let scale_bits = bytes[2];
336 Self::validate_scale_bits(scale_bits)?;
337 let target = u64::from(1u32 << scale_bits);
338 match form {
339 MODEL_FORM_SPARSE => Self::decode_v2_sparse(bytes, scale_bits, target),
340 MODEL_FORM_DENSE => Self::decode_v2_dense(bytes, scale_bits, target),
341 other => Err(Error::invalid_model(format!(
342 "model form {other} is not 0 (sparse) or 1 (dense)"
343 ))),
344 }
345 }
346
347 fn decode_v2_sparse(bytes: &[u8], scale_bits: u8, target: u64) -> Result<EntropyModel> {
349 let prefix = MODEL_V2_HEADER_LEN + MODEL_V2_COUNT_LEN;
350 if bytes.len() < prefix {
351 return Err(Error::invalid_model("sparse model shorter than its count"));
352 }
353 let present_count = u16::from_le_bytes([bytes[3], bytes[4]]) as usize;
354 if present_count > ALPHABET {
355 return Err(Error::invalid_model(format!(
356 "sparse present_count {present_count} exceeds alphabet {ALPHABET}"
357 )));
358 }
359 let expected = prefix + present_count * MODEL_V2_SPARSE_ENTRY_LEN;
360 if bytes.len() != expected {
361 return Err(Error::invalid_model(
362 "sparse entry count does not match payload (trailing or truncated)",
363 ));
364 }
365
366 let mut frequencies = vec![0u32; ALPHABET];
367 let mut sum: u64 = 0;
368 let mut previous: Option<u8> = None;
369 for entry in bytes[prefix..].as_chunks::<MODEL_V2_SPARSE_ENTRY_LEN>().0 {
370 let symbol = entry[0];
371 let freq = u32::from(u16::from_le_bytes([entry[1], entry[2]]));
372 if let Some(prev) = previous
373 && symbol <= prev
374 {
375 return Err(Error::invalid_model(
376 "sparse symbols must be strictly ascending and unique",
377 ));
378 }
379 if freq == 0 {
380 return Err(Error::invalid_model("sparse frequency must be >= 1"));
381 }
382 previous = Some(symbol);
383 sum += u64::from(freq);
384 frequencies[symbol as usize] = freq;
385 }
386
387 if sum != target {
388 return Err(Error::invalid_model(
389 "sparse frequencies do not sum to 1 << scale_bits",
390 ));
391 }
392 Ok(EntropyModel {
393 scale_bits,
394 frequencies,
395 })
396 }
397
398 fn decode_v2_dense(bytes: &[u8], scale_bits: u8, target: u64) -> Result<EntropyModel> {
400 let prefix = MODEL_V2_HEADER_LEN + MODEL_V2_COUNT_LEN;
401 if bytes.len() < prefix {
402 return Err(Error::invalid_model("dense model shorter than its count"));
403 }
404 let count = u16::from_le_bytes([bytes[3], bytes[4]]);
405 if count as usize != ALPHABET {
406 return Err(Error::invalid_model(format!(
407 "dense declared count {count}, expected {ALPHABET}"
408 )));
409 }
410 if bytes.len() != MODEL_V2_DENSE_LEN {
411 return Err(Error::invalid_model(
412 "dense model length does not match its count (trailing or truncated)",
413 ));
414 }
415
416 let mut frequencies = Vec::with_capacity(ALPHABET);
417 let mut sum: u64 = 0;
418 for chunk in bytes[prefix..].as_chunks::<2>().0 {
419 let f = u32::from(u16::from_le_bytes(*chunk));
420 sum += u64::from(f);
421 frequencies.push(f);
422 }
423 if sum != target {
424 return Err(Error::invalid_model(
425 "dense frequencies do not sum to 1 << scale_bits",
426 ));
427 }
428 Ok(EntropyModel {
429 scale_bits,
430 frequencies,
431 })
432 }
433}
434
435#[cfg(test)]
436mod tests {
437 use super::*;
438
439 fn counts_with(pairs: &[(usize, u64)]) -> [u64; ALPHABET] {
440 let mut counts = [0u64; ALPHABET];
441 for &(i, v) in pairs {
442 counts[i] = v;
443 }
444 counts
445 }
446
447 #[test]
448 fn uniform_sums_to_target() {
449 for bits in [8u8, 12, 15] {
450 let model = EntropyModel::uniform(bits).expect("uniform must succeed");
451 assert_eq!(model.frequencies.len(), ALPHABET);
452 assert_eq!(model.frequency_sum(), u64::from(1u32 << bits));
453 let per = model.frequencies[0];
454 assert!(model.frequencies.iter().all(|&f| f == per));
455 }
456 }
457
458 #[test]
459 fn from_counts_sums_exactly() {
460 let inputs: &[&[(usize, u64)]] = &[
461 &[(0, 1)],
462 &[(0, 1), (1, 1)],
463 &[(0, 300), (1, 1)],
464 &[(0, 1), (1, 2), (2, 3), (3, 4)],
465 &[(5, 10), (200, 1), (255, 7)],
466 &[(0, u64::MAX), (255, 1)],
467 ];
468 for bits in [8u8, 12] {
469 for input in inputs {
470 let counts = counts_with(input);
471 let model =
472 EntropyModel::from_counts(&counts, bits).expect("normalization must succeed");
473 assert_eq!(model.scale_bits, bits);
474 assert_eq!(model.frequency_sum(), u64::from(1u32 << bits));
475 for &(i, _) in *input {
476 assert!(model.frequencies[i] >= 1);
477 }
478 }
479 }
480 }
481
482 #[test]
483 fn determinism() {
484 let counts = counts_with(&[(0, 17), (3, 5), (7, 250), (255, 1)]);
485 let a = EntropyModel::from_counts(&counts, 12)
486 .expect("ok")
487 .encode()
488 .expect("ok");
489 let b = EntropyModel::from_counts(&counts, 12)
490 .expect("ok")
491 .encode()
492 .expect("ok");
493 assert_eq!(a, b);
494 }
495
496 #[test]
497 fn present_symbols_get_at_least_one() {
498 let counts = [1u64; ALPHABET];
501 let model = EntropyModel::from_counts(&counts, 12).expect("ok");
502 assert!(model.frequencies.iter().all(|&f| f >= 1));
503 assert_eq!(model.frequency_sum(), 4096);
504 }
505
506 #[test]
507 fn zero_symbols_get_zero() {
508 let counts = counts_with(&[(0, 5), (10, 9), (255, 3)]);
509 let model = EntropyModel::from_counts(&counts, 12).expect("ok");
510 for (i, &f) in model.frequencies.iter().enumerate() {
511 if counts[i] == 0 {
512 assert_eq!(f, 0, "symbol {i} had zero count but nonzero frequency");
513 }
514 }
515 }
516
517 #[test]
518 fn alphabet_too_big_errors() {
519 let counts = [1u64; ALPHABET];
520 for bits in 1u8..=7 {
521 assert!(EntropyModel::from_counts(&counts, bits).is_err());
522 }
523 }
524
525 #[test]
526 fn roundtrip_encode_decode() {
527 let counts = counts_with(&[(0, 1), (1, 2), (2, 3), (100, 400), (255, 7)]);
528 for bits in [8u8, 12, 15] {
529 let model = EntropyModel::from_counts(&counts, bits).expect("ok");
530 let bytes = model.encode().expect("ok");
531 assert_eq!(bytes[0], MODEL_VERSION_2);
533 assert_eq!(bytes[1], MODEL_FORM_SPARSE);
534 assert_eq!(bytes.len(), MODEL_V2_HEADER_LEN + 2 + 5 * 3);
535 let decoded = EntropyModel::decode(&bytes).expect("ok");
536 assert_eq!(decoded, model);
537 assert_eq!(decoded.encode().expect("ok"), bytes);
539 }
540 }
541
542 #[test]
543 fn sparse_form_chosen_for_low_alphabet() {
544 let counts = counts_with(&[(0, 5), (255, 3)]);
545 let model = EntropyModel::from_counts(&counts, 12).expect("ok");
546 let bytes = model.encode().expect("ok");
547 assert_eq!(bytes[0], MODEL_VERSION_2);
548 assert_eq!(bytes[1], MODEL_FORM_SPARSE);
549 assert_eq!(bytes.len(), MODEL_V2_HEADER_LEN + 2 + 2 * 3);
550 assert_eq!(EntropyModel::decode(&bytes).expect("ok"), model);
551 }
552
553 #[test]
554 fn dense_form_chosen_for_full_alphabet() {
555 let model = EntropyModel::uniform(8).expect("ok");
556 let bytes = model.encode().expect("ok");
557 assert_eq!(bytes[0], MODEL_VERSION_2);
558 assert_eq!(bytes[1], MODEL_FORM_DENSE);
559 assert_eq!(bytes.len(), MODEL_V2_DENSE_LEN);
560 assert_eq!(EntropyModel::decode(&bytes).expect("ok"), model);
561 }
562
563 #[test]
564 fn decode_accepts_legacy_v1_dense() {
565 let model = EntropyModel::uniform(8).expect("ok");
566 let mut bytes = Vec::with_capacity(MODEL_V1_ENCODED_LEN);
568 bytes.push(MODEL_VERSION_1);
569 bytes.push(model.scale_bits);
570 bytes.extend_from_slice(&(ALPHABET as u16).to_le_bytes());
571 for &f in &model.frequencies {
572 bytes.extend_from_slice(&(f as u16).to_le_bytes());
573 }
574 assert_eq!(bytes.len(), MODEL_V1_ENCODED_LEN);
575 assert_eq!(EntropyModel::decode(&bytes).expect("ok"), model);
576 }
577
578 #[test]
579 fn encode_is_deterministic() {
580 let counts = counts_with(&[(0, 17), (3, 5), (7, 250), (255, 1)]);
581 let model = EntropyModel::from_counts(&counts, 12).expect("ok");
582 let a = model.encode().expect("ok");
583 let b = model.encode().expect("ok");
584 assert_eq!(a, b);
585 assert_eq!(a[0], MODEL_VERSION_2);
586 }
587
588 #[test]
589 fn sparse_decode_rejects_unsorted_or_duplicate_symbols() {
590 let counts = counts_with(&[(1, 5), (2, 3)]);
591 let model = EntropyModel::from_counts(&counts, 12).expect("ok");
592 let bytes = model.encode().expect("ok");
593 assert_eq!(bytes[1], MODEL_FORM_SPARSE);
594 assert_eq!(bytes.len(), MODEL_V2_HEADER_LEN + 2 + 2 * 3);
595
596 let mut duplicate = bytes.clone();
598 duplicate[8] = duplicate[5];
599 assert!(EntropyModel::decode(&duplicate).is_err());
600
601 let mut swapped = bytes.clone();
603 swapped[5..8].copy_from_slice(&bytes[8..11]);
604 swapped[8..11].copy_from_slice(&bytes[5..8]);
605 assert!(EntropyModel::decode(&swapped).is_err());
606 }
607
608 #[test]
609 fn decode_rejects_bad_sum() {
610 let mut bytes = EntropyModel::uniform(8).expect("ok").encode().expect("ok");
611 let last = bytes.len();
613 let value = u16::from_le_bytes([bytes[last - 2], bytes[last - 1]]);
614 let bumped = value.checked_add(1).expect("fits");
615 bytes[last - 2..].copy_from_slice(&bumped.to_le_bytes());
616 assert!(EntropyModel::decode(&bytes).is_err());
617 }
618
619 #[test]
620 fn decode_rejects_bad_version() {
621 let mut bytes = EntropyModel::uniform(8).expect("ok").encode().expect("ok");
622 bytes[0] = 3;
623 let err = EntropyModel::decode(&bytes).expect_err("must reject");
624 assert_eq!(err.class(), crate::error::ErrorClass::UnsupportedVersion);
625 }
626
627 #[test]
628 fn decode_rejects_trailing() {
629 let mut bytes = EntropyModel::uniform(8).expect("ok").encode().expect("ok");
630 bytes.push(0);
631 assert!(EntropyModel::decode(&bytes).is_err());
632 }
633
634 #[test]
635 fn all_zero_is_uniform() {
636 let zero = [0u64; ALPHABET];
637 for bits in [8u8, 12, 15] {
638 let model = EntropyModel::from_counts(&zero, bits).expect("ok");
639 assert_eq!(model, EntropyModel::uniform(bits).expect("ok"));
640 }
641 }
642}