Skip to main content

vole_document/entropy/
rans.rs

1//! Single-lane order-0 byte rANS over the safe manual ryg-rans-rs API.
2//!
3//! This module deliberately avoids `ryg_rans_rs::alloc_utils`, whose `decode`
4//! panics on truncated input. Everything here goes through the checked manual
5//! API (`rans_byte_enc_put_symbol`, `rans_byte_enc_flush`,
6//! `rans_byte_dec_get`, `rans_byte_dec_advance_symbol`). Symbol lookup is a
7//! scalar, deterministic cumulative table; every value derived from untrusted
8//! input is range-checked or uses checked arithmetic, and the decoder never
9//! unwinds on malformed input.
10
11use ryg_rans_rs::byte::{
12    BackwardByteWriter, ByteReader, RANS_BYTE_L, RansByteDecSymbol, RansByteEncSymbol,
13    RansByteState, rans_byte_dec_advance_symbol, rans_byte_dec_get, rans_byte_enc_flush,
14    rans_byte_enc_put_symbol,
15};
16
17use crate::entropy::model::{ALPHABET, EntropyModel, MAX_SCALE_BITS, MIN_SCALE_BITS};
18use crate::error::{Error, Result};
19use crate::limits::Limits;
20
21/// Complete decoder-entry capsule for one entropy channel.
22/// The persisted physical form (states + renormalization payload + counts) —
23/// never a bare scalar "seed".
24#[derive(Debug, Clone, PartialEq, Eq)]
25pub struct Capsule {
26    /// One decoder initial state (scalar single lane). Encoder final state.
27    pub initial_state: u32,
28    /// Renormalization payload bytes, in the exact orientation the decoder
29    /// consumes them: forward from index 0, the reverse of encoder emission
30    /// order. This excludes the 4-byte flush state, which lives in
31    /// [`Capsule::initial_state`].
32    pub payload: Vec<u8>,
33    /// Number of symbols encoded (== decoded_length for the byte alphabet).
34    pub symbol_count: u64,
35    /// Exact decoded length in bytes.
36    pub decoded_length: u64,
37}
38
39/// Validate a model for use with the byte-rANS codec.
40///
41/// Returns the model total (`1 << scale_bits`). `decode` selects the error
42/// class: decode paths report [`Error::entropy_decode`], encode paths report
43/// [`Error::invalid_model`].
44fn validate_model(model: &EntropyModel, decode: bool) -> Result<u32> {
45    let fail = |msg: String| {
46        if decode {
47            Error::entropy_decode(msg)
48        } else {
49            Error::invalid_model(msg)
50        }
51    };
52
53    if !(MIN_SCALE_BITS..=MAX_SCALE_BITS).contains(&model.scale_bits) {
54        return Err(fail(format!(
55            "scale_bits {} outside {}..={}",
56            model.scale_bits, MIN_SCALE_BITS, MAX_SCALE_BITS
57        )));
58    }
59    if model.frequencies.len() != ALPHABET {
60        return Err(fail(format!(
61            "expected {ALPHABET} frequencies, got {}",
62            model.frequencies.len()
63        )));
64    }
65
66    let target = 1u32 << model.scale_bits;
67    let mut sum: u64 = 0;
68    for &f in &model.frequencies {
69        if u64::from(f) > u64::from(target) {
70            return Err(fail("frequency exceeds 1 << scale_bits".to_string()));
71        }
72        sum += u64::from(f);
73    }
74    if sum != u64::from(target) {
75        return Err(fail(
76            "frequencies do not sum to 1 << scale_bits".to_string(),
77        ));
78    }
79    Ok(target)
80}
81
82/// Encode `data` with `model`. Deterministic.
83///
84/// Symbols are consumed last-to-first (rANS stack discipline) and
85/// renormalization bytes are emitted backward; the resulting payload is
86/// returned in forward decoder-consumption order.
87pub fn encode_channel(model: &EntropyModel, data: &[u8]) -> Result<Capsule> {
88    let scale_bits = u32::from(model.scale_bits);
89    let target = validate_model(model, false)?;
90
91    // Build the per-symbol encoder table. Symbols with zero frequency cannot
92    // be encoded; they are left absent and rejected if they occur in `data`.
93    let mut enc_syms: [Option<RansByteEncSymbol>; ALPHABET] = [None; ALPHABET];
94    let mut start: u32 = 0;
95    for (symbol, &freq) in model.frequencies.iter().enumerate() {
96        if freq > 0 {
97            let sym = RansByteEncSymbol::new(start, freq, scale_bits)
98                .map_err(|e| Error::invalid_model(format!("encoder symbol {symbol}: {e}")))?;
99            enc_syms[symbol] = Some(sym);
100        }
101        start = start
102            .checked_add(freq)
103            .ok_or_else(|| Error::invalid_model("cumulative frequency overflow"))?;
104    }
105    if start != target {
106        return Err(Error::invalid_model("cumulative table mismatch"));
107    }
108
109    // Worst-case output bound matching the upstream convenience API:
110    // at most 4 bytes per symbol plus flush headroom.
111    let max_size = data
112        .len()
113        .checked_mul(4)
114        .and_then(|n| n.checked_add(24))
115        .ok_or_else(|| Error::resource_limit("encoded size estimate overflow"))?;
116    let mut buf = vec![0u8; max_size];
117    let mut writer = BackwardByteWriter::new(&mut buf);
118
119    let mut state = RansByteState::new();
120    for &byte in data.iter().rev() {
121        let sym = enc_syms[byte as usize]
122            .as_ref()
123            .ok_or_else(|| Error::invalid_model(format!("data byte {byte} has zero frequency")))?;
124        rans_byte_enc_put_symbol(&mut state, &mut writer, sym)
125            .map_err(|_| Error::internal_invariant("rANS encoder buffer exhausted"))?;
126    }
127    rans_byte_enc_flush(&state, &mut writer)
128        .map_err(|_| Error::internal_invariant("rANS flush buffer exhausted"))?;
129
130    // `encoded` begins with the 4-byte little-endian flush state (the last
131    // write lands at the lowest address); the remainder is the renormalization
132    // payload already in forward decoder-consumption order.
133    let encoded = writer.encoded();
134    let payload = encoded
135        .get(4..)
136        .ok_or_else(|| Error::internal_invariant("flush state missing from encoder output"))?
137        .to_vec();
138
139    Ok(Capsule {
140        initial_state: state.get(),
141        payload,
142        symbol_count: data.len() as u64,
143        decoded_length: data.len() as u64,
144    })
145}
146
147/// Decode a channel, hostile-safe and bounded.
148///
149/// Never panics: truncated state/payload, illegal `symbol_count` /
150/// `decoded_length`, an unsupported model, a length mismatch, or a failed
151/// stream-integrity check all return [`Error::entropy_decode`].
152pub fn decode_channel(model: &EntropyModel, capsule: &Capsule, limits: Limits) -> Result<Vec<u8>> {
153    let scale_bits = u32::from(model.scale_bits);
154    let target = validate_model(model, true)?;
155
156    // Resource bounds checked before any large allocation or work.
157    if capsule.symbol_count > limits.max_channel_symbols {
158        return Err(Error::entropy_decode(format!(
159            "symbol_count {} exceeds limit {}",
160            capsule.symbol_count, limits.max_channel_symbols
161        )));
162    }
163    if capsule.decoded_length > limits.max_output_bytes {
164        return Err(Error::entropy_decode(format!(
165            "decoded_length {} exceeds limit {}",
166            capsule.decoded_length, limits.max_output_bytes
167        )));
168    }
169    if capsule.payload.len() as u64 > u64::from(limits.max_record_len) {
170        return Err(Error::entropy_decode(format!(
171            "payload length {} exceeds limit {}",
172            capsule.payload.len(),
173            limits.max_record_len
174        )));
175    }
176    if capsule.symbol_count != capsule.decoded_length {
177        return Err(Error::entropy_decode(
178            "symbol_count does not equal decoded_length".to_string(),
179        ));
180    }
181
182    let symbol_count = usize::try_from(capsule.symbol_count)
183        .map_err(|_| Error::entropy_decode("symbol_count does not fit platform usize"))?;
184
185    // Stable cumulative decode table: slot in [0, 1<<scale_bits) maps to the
186    // unique symbol `s` with cum[s] <= slot < cum[s] + freq[s].
187    let mut dec_syms: [Option<RansByteDecSymbol>; ALPHABET] = [None; ALPHABET];
188    let mut cum2sym = vec![0u8; target as usize];
189    let mut start: u32 = 0;
190    for (symbol, &freq) in model.frequencies.iter().enumerate() {
191        if freq > 0 {
192            let dsym = RansByteDecSymbol::new(start, freq)
193                .map_err(|e| Error::entropy_decode(format!("decoder symbol {symbol}: {e}")))?;
194            dec_syms[symbol] = Some(dsym);
195            let end = start
196                .checked_add(freq)
197                .ok_or_else(|| Error::entropy_decode("cumulative frequency overflow"))?;
198            for slot in &mut cum2sym[start as usize..end as usize] {
199                *slot = symbol as u8;
200            }
201            start = end;
202        }
203    }
204    if start != target {
205        return Err(Error::entropy_decode(
206            "cumulative table mismatch".to_string(),
207        ));
208    }
209
210    let mut output: Vec<u8> = Vec::new();
211    output
212        .try_reserve(symbol_count)
213        .map_err(|_| Error::entropy_decode("cannot allocate decode buffer"))?;
214
215    let mut reader = ByteReader::new(&capsule.payload);
216    // The public field is safe to construct directly; the value is untrusted
217    // and every downstream arithmetic step is overflow-free for any u32
218    // because model frequencies never exceed `1 << scale_bits`.
219    let mut state = RansByteState(capsule.initial_state);
220
221    for _ in 0..symbol_count {
222        let slot = rans_byte_dec_get(&state, scale_bits);
223        let symbol = cum2sym[slot as usize];
224        output.push(symbol);
225        let dsym = dec_syms[symbol as usize]
226            .as_ref()
227            .ok_or_else(|| Error::entropy_decode("slot mapped to zero-frequency symbol"))?;
228        rans_byte_dec_advance_symbol(&mut state, &mut reader, dsym, scale_bits)
229            .map_err(|_| Error::entropy_decode("truncated renormalization payload"))?;
230    }
231
232    if output.len() as u64 != capsule.decoded_length {
233        return Err(Error::entropy_decode(format!(
234            "decoded {} bytes, expected {}",
235            output.len(),
236            capsule.decoded_length
237        )));
238    }
239    // A well-formed stream consumes the payload exactly and returns the state
240    // to the encoder's initial lower bound. This rejects truncation and most
241    // corrupted state/payload combinations outright.
242    if state.get() != RANS_BYTE_L || reader.remaining() != 0 {
243        return Err(Error::entropy_decode(
244            "entropy stream integrity check failed".to_string(),
245        ));
246    }
247
248    Ok(output)
249}
250
251#[cfg(test)]
252mod tests {
253    use super::*;
254    use crate::error::ErrorClass;
255
256    /// Deterministic xorshift64 PRNG for reproducible test corpora.
257    struct XorShift(u64);
258
259    impl XorShift {
260        fn new(seed: u64) -> Self {
261            Self(seed | 1)
262        }
263
264        fn next_u32(&mut self) -> u32 {
265            let mut x = self.0;
266            x ^= x << 13;
267            x ^= x >> 7;
268            x ^= x << 17;
269            self.0 = x;
270            (x >> 32) as u32
271        }
272
273        fn next_byte(&mut self) -> u8 {
274            (self.next_u32() & 0xff) as u8
275        }
276    }
277
278    fn model_from_data(data: &[u8], scale_bits: u8) -> EntropyModel {
279        let mut counts = [0u64; ALPHABET];
280        for &b in data {
281            counts[b as usize] += 1;
282        }
283        EntropyModel::from_counts(&counts, scale_bits).expect("model normalizes")
284    }
285
286    fn roundtrip(data: &[u8], scale_bits: u8) {
287        let model = model_from_data(data, scale_bits);
288        let capsule = encode_channel(&model, data).expect("encode");
289        assert_eq!(capsule.symbol_count, data.len() as u64);
290        assert_eq!(capsule.decoded_length, data.len() as u64);
291
292        let decoded = decode_channel(&model, &capsule, Limits::DEFAULT).expect("decode");
293        assert_eq!(decoded, data);
294
295        let again = encode_channel(&model, data).expect("encode again");
296        assert_eq!(capsule, again, "encode must be deterministic");
297    }
298
299    #[test]
300    fn roundtrip_empty() {
301        for bits in [8u8, 12] {
302            roundtrip(&[], bits);
303        }
304    }
305
306    #[test]
307    fn roundtrip_single_repeated_byte() {
308        for bits in [8u8, 12] {
309            roundtrip(&[0x41u8; 1000], bits);
310        }
311    }
312
313    #[test]
314    fn roundtrip_all_256_values() {
315        let mut data = Vec::new();
316        for _ in 0..8 {
317            data.extend(0u8..=255);
318        }
319        for bits in [8u8, 12] {
320            roundtrip(&data, bits);
321        }
322    }
323
324    #[test]
325    fn roundtrip_uniform_random() {
326        let mut rng = XorShift::new(0x1234_5678_9abc_def0);
327        let data: Vec<u8> = (0..4096).map(|_| rng.next_byte()).collect();
328        for bits in [8u8, 12] {
329            roundtrip(&data, bits);
330        }
331    }
332
333    #[test]
334    fn roundtrip_heavily_skewed() {
335        let mut rng = XorShift::new(0xdead_beef_cafe_f00d);
336        let data: Vec<u8> = (0..4096)
337            .map(|i| if i % 100 == 0 { rng.next_byte() } else { 0x00 })
338            .collect();
339        for bits in [8u8, 12] {
340            roundtrip(&data, bits);
341        }
342    }
343
344    #[test]
345    fn truncated_payload_is_rejected() {
346        let mut rng = XorShift::new(0x0f0f_0f0f_1234_5678);
347        let data: Vec<u8> = (0..4096).map(|_| rng.next_byte()).collect();
348        let model = model_from_data(&data, 12);
349        let capsule = encode_channel(&model, &data).expect("encode");
350        assert!(!capsule.payload.is_empty());
351
352        let mut truncated = capsule.clone();
353        truncated.payload.pop();
354        assert!(decode_channel(&model, &truncated, Limits::DEFAULT).is_err());
355    }
356
357    #[test]
358    fn missing_payload_is_rejected() {
359        let mut rng = XorShift::new(0x9988_7766_5544_3322);
360        let data: Vec<u8> = (0..1024).map(|_| rng.next_byte()).collect();
361        let model = model_from_data(&data, 12);
362        let capsule = encode_channel(&model, &data).expect("encode");
363
364        let mut missing = capsule.clone();
365        missing.payload.clear();
366        assert!(decode_channel(&model, &missing, Limits::DEFAULT).is_err());
367    }
368
369    #[test]
370    fn corrupted_state_and_payload_never_panic() {
371        let mut rng = XorShift::new(0xabcd_ef01_2345_6789);
372        let data: Vec<u8> = (0..2048).map(|_| rng.next_byte()).collect();
373        let model = model_from_data(&data, 12);
374        let capsule = encode_channel(&model, &data).expect("encode");
375
376        let mut states = vec![0u32, u32::MAX, RANS_BYTE_L, capsule.initial_state];
377        for delta in [1u32, 0x8000, 0xffff_ffff] {
378            states.push(capsule.initial_state.wrapping_add(delta));
379        }
380        for value in states {
381            let mut c = capsule.clone();
382            c.initial_state = value;
383            match decode_channel(&model, &c, Limits::DEFAULT) {
384                Ok(out) => assert_eq!(out.len() as u64, c.decoded_length),
385                Err(e) => assert_eq!(e.class(), ErrorClass::EntropyDecode),
386            }
387        }
388
389        for (i, _) in capsule.payload.iter().enumerate().take(64) {
390            let mut c = capsule.clone();
391            c.payload[i] ^= 0xff;
392            match decode_channel(&model, &c, Limits::DEFAULT) {
393                Ok(out) => assert_eq!(out.len() as u64, c.decoded_length),
394                Err(e) => assert_eq!(e.class(), ErrorClass::EntropyDecode),
395            }
396        }
397    }
398
399    #[test]
400    fn limits_reject_oversized_fields() {
401        let mut rng = XorShift::new(0x1357_9bdf_2468_ace0);
402        let data: Vec<u8> = (0..2048).map(|_| rng.next_byte()).collect();
403        let model = model_from_data(&data, 12);
404        let capsule = encode_channel(&model, &data).expect("encode");
405        assert!(capsule.symbol_count > 0);
406        assert!(capsule.decoded_length > 0);
407        assert!(!capsule.payload.is_empty());
408
409        let by_symbols = Limits {
410            max_channel_symbols: capsule.symbol_count - 1,
411            ..Limits::DEFAULT
412        };
413        assert!(decode_channel(&model, &capsule, by_symbols).is_err());
414
415        let by_output = Limits {
416            max_output_bytes: capsule.decoded_length - 1,
417            ..Limits::DEFAULT
418        };
419        assert!(decode_channel(&model, &capsule, by_output).is_err());
420
421        let by_record = Limits {
422            max_record_len: capsule.payload.len() as u32 - 1,
423            ..Limits::DEFAULT
424        };
425        assert!(decode_channel(&model, &capsule, by_record).is_err());
426
427        // The strict profile still accepts this small channel.
428        assert!(decode_channel(&model, &capsule, Limits::STRICT).is_ok());
429    }
430
431    #[test]
432    fn inconsistent_lengths_are_rejected() {
433        let data = b"length check".to_vec();
434        let model = model_from_data(&data, 12);
435        let capsule = encode_channel(&model, &data).expect("encode");
436
437        let mut mismatched = capsule.clone();
438        mismatched.symbol_count = capsule.symbol_count + 1;
439        assert!(decode_channel(&model, &mismatched, Limits::DEFAULT).is_err());
440
441        let mut bad_len = capsule.clone();
442        bad_len.decoded_length = capsule.decoded_length + 1;
443        assert!(decode_channel(&model, &bad_len, Limits::DEFAULT).is_err());
444    }
445}