1use 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#[derive(Debug, Clone, PartialEq, Eq)]
25pub struct Capsule {
26 pub initial_state: u32,
28 pub payload: Vec<u8>,
33 pub symbol_count: u64,
35 pub decoded_length: u64,
37}
38
39fn 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
82pub 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 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 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 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
147pub 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 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 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 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 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 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 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}