Skip to main content

vole_document/entropy/
model.rs

1//! Canonical entropy models: a deterministic, integer-only (no floating point)
2//! frequency-table normalizer and its wire serialization.
3//!
4//! A model is a frequency table over a fixed byte alphabet ([`ALPHABET`] = 256)
5//! whose entries sum to exactly `1 << scale_bits`. Normalization from raw symbol
6//! counts is a *pure function of the counts and `scale_bits`*: identical inputs
7//! always produce identical output bytes. This is the substrate rANS will later
8//! consume; no entropy coder lives here yet.
9
10use crate::error::{Error, Result};
11
12/// Supported alphabet size (bytes).
13pub const ALPHABET: usize = 256;
14/// Legacy model wire version (dense-only, `[1][scale_bits][count=256][u16 x 256]`).
15pub const MODEL_VERSION_1: u8 = 1;
16/// Compact model wire version (sparse/dense form selection).
17pub const MODEL_VERSION_2: u8 = 2;
18/// Minimum / maximum scale bits. Frequencies must fit in u16, so <= 15.
19pub const MIN_SCALE_BITS: u8 = 1;
20pub const MAX_SCALE_BITS: u8 = 15;
21
22/// v2 form selector: sparse `[symbol u8][freq u16]` entries for present symbols.
23const MODEL_FORM_SPARSE: u8 = 0;
24/// v2 form selector: full dense 256-entry `u16` frequency table.
25const MODEL_FORM_DENSE: u8 = 1;
26
27/// Legacy v1 header: `version`, `scale_bits`, `count`.
28const MODEL_V1_HEADER_LEN: usize = 4;
29/// Legacy v1 total length: header plus `ALPHABET` little-endian u16s.
30const MODEL_V1_ENCODED_LEN: usize = MODEL_V1_HEADER_LEN + ALPHABET * 2;
31/// v2 header: `version`, `form`, `scale_bits`.
32const MODEL_V2_HEADER_LEN: usize = 3;
33/// v2 payload prefix: `count` little-endian u16.
34const MODEL_V2_COUNT_LEN: usize = 2;
35/// v2 sparse entry: `symbol` u8 plus `freq` little-endian u16.
36const MODEL_V2_SPARSE_ENTRY_LEN: usize = 3;
37/// v2 dense total length: header plus `count` plus `ALPHABET` little-endian u16s.
38const MODEL_V2_DENSE_LEN: usize = MODEL_V2_HEADER_LEN + MODEL_V2_COUNT_LEN + ALPHABET * 2;
39
40/// A normalized frequency table over the byte alphabet.
41#[derive(Debug, Clone, PartialEq, Eq)]
42pub struct EntropyModel {
43    /// Number of scale bits; the table sums to `1 << scale_bits`.
44    pub scale_bits: u8,
45    /// Exactly `ALPHABET` entries; sum == `1 << scale_bits`.
46    pub frequencies: Vec<u32>,
47}
48
49impl EntropyModel {
50    /// Total frequency (== `1 << scale_bits`).
51    pub fn total(&self) -> u32 {
52        1u32 << self.scale_bits
53    }
54
55    /// Sum of all frequencies as a widened integer (avoids debug overflow on
56    /// hand-built, out-of-contract tables).
57    fn frequency_sum(&self) -> u64 {
58        self.frequencies.iter().map(|&f| u64::from(f)).sum()
59    }
60
61    /// Canonical normalization of observed symbol counts into a frequency table.
62    ///
63    /// Requirements (all hold):
64    /// - integer-only; no floating point anywhere;
65    /// - deterministic: same input yields identical output bytes;
66    /// - every symbol with `count > 0` gets frequency `>= 1`;
67    /// - symbols with `count == 0` get frequency `0`;
68    /// - `sum(frequencies) == 1 << scale_bits` exactly;
69    /// - if all counts are `0`, return the canonical uniform model.
70    ///
71    /// Algorithm: after granting each *present* symbol a guaranteed minimum of
72    /// 1, the remaining budget is apportioned by the largest-remainder (Hare)
73    /// method using `u128` products, with ties broken by lower symbol index.
74    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        // 2. Present symbols.
84        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        // 3. Alphabet (as a set of present symbols) must fit in the target.
90        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        // 4. Guaranteed minimum of 1 for each present symbol.
98        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            // 5. Integer quota + remainder per present symbol.
108            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            // 6. Largest remainder, ties by lower index. There are always at
123            // least `leftover` symbols with a nonzero remainder, so this pass
124            // consumes the whole leftover in practice.
125            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            // Deterministic fallback, kept for total-sum safety. It is
136            // unreachable for well-formed inputs (the remainder pass above
137            // always suffices), and adds to the present symbol with the largest
138            // count, breaking ties by lower index.
139            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    /// Uniform model with every symbol getting an equal frequency.
163    ///
164    /// Requires `scale_bits >= 8`, i.e. `1 << scale_bits >= ALPHABET`; below
165    /// that no integer table of 256 equal entries can sum to the target, so the
166    /// request is rejected as [`ErrorClass::InvalidModel`](crate::error::ErrorClass::InvalidModel).
167    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        // `target` is a power of two >= 256, hence exactly divisible by 256.
181        let per = target / ALPHABET as u32;
182        Ok(EntropyModel {
183            scale_bits,
184            frequencies: vec![per; ALPHABET],
185        })
186    }
187
188    /// Structural and arithmetic validation shared by every serializer.
189    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    /// Reject `scale_bits` outside [`MIN_SCALE_BITS`]..=[`MAX_SCALE_BITS`].
206    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    /// Canonical wire encoding, **version 2** (little-endian):
217    /// `[version u8 = 2][form u8][scale_bits u8]` followed by the form payload.
218    ///
219    /// * form `0` (SPARSE): `[present_count u16][symbol u8][freq u16] * n`, with
220    ///   `symbol` strictly ascending and `freq >= 1`.
221    /// * form `1` (DENSE): `[count u16 = 256][freq u16] * 256`.
222    ///
223    /// The strictly smaller serialization is emitted; on a tie DENSE is chosen
224    /// so the mapping from model to bytes stays deterministic. The model is
225    /// validated before serialization so `encode` cannot emit a bad table.
226    pub fn encode(&self) -> Result<Vec<u8>> {
227        self.validate()?;
228
229        // Sparse cost is header + count + one 3-byte entry per present symbol.
230        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    /// Parse and validate a canonical model; `bytes` must be exactly consumed.
272    ///
273    /// Both legacy version 1 (dense) and version 2 (sparse or dense) are
274    /// accepted. All failures are typed as
275    /// [`ErrorClass::InvalidModel`](crate::error::ErrorClass::InvalidModel) or
276    /// [`ErrorClass::UnsupportedVersion`](crate::error::ErrorClass::UnsupportedVersion).
277    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    /// Legacy version 1: `[1][scale_bits][count u16 = 256][freq u16 x 256]`.
291    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    /// Version 2: `[2][form][scale_bits]` followed by the form payload.
330    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    /// v2 sparse payload: `[present_count u16]` then `[symbol u8][freq u16]`.
348    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    /// v2 dense payload: `[count u16 = 256]` then 256 little-endian u16s.
399    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        // Every symbol present; at scale_bits 12 the target (4096) comfortably
499        // exceeds the alphabet size (256).
500        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            // Five present symbols select the sparse form.
532            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            // The winner's own form re-encodes identically.
538            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        // Hand-build the legacy v1 wire form.
567        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        // Duplicate: second entry's symbol forced equal to the first's.
597        let mut duplicate = bytes.clone();
598        duplicate[8] = duplicate[5];
599        assert!(EntropyModel::decode(&duplicate).is_err());
600
601        // Unsorted: swap the two entries so symbols descend.
602        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        // Bump the last frequency so the total no longer equals 1 << 8.
612        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}