use crate::error::{Error, Result};
use super::{ByteReader, MAX_CODEC_LEN};
const RANS_L: u32 = 1 << 23;
const TOTAL_BITS: u32 = 12;
const TOTAL: usize = 1 << TOTAL_BITS;
pub fn decode(data: &[u8], path: &str, offset: u64) -> Result<Vec<u8>> {
let mut reader = ByteReader::new(data, path, offset);
let order = reader.u8()?;
let compressed = reader.u32()? as usize;
let raw = reader.u32()? as usize;
if raw > MAX_CODEC_LEN {
return Err(Error::corrupt(
path,
offset,
format!("a rans4x8 block naming {raw} raw bytes, past this reader's ceiling"),
));
}
if compressed > reader.remaining() {
return Err(Error::corrupt(
path,
offset,
format!(
"a rans4x8 block naming {compressed} compressed bytes with {} left",
reader.remaining()
),
));
}
if raw == 0 {
return Ok(Vec::new());
}
match order {
0 => decode_order_0(&mut reader, raw),
1 => decode_order_1(&mut reader, raw),
other => Err(Error::corrupt(
path,
offset,
format!("a rans4x8 block of order {other}, which is neither 0 nor 1"),
)),
}
}
struct SymbolTable {
freq: [u32; 256],
cumulative: [u32; 256],
lookup: Vec<u8>,
covered: u32,
}
impl SymbolTable {
fn build(freq: [u32; 256], path: &str, offset: u64) -> Result<Self> {
let total: u64 = freq.iter().map(|f| u64::from(*f)).sum();
if total == 0 {
return Err(Error::corrupt(
path,
offset,
"a rans4x8 frequency table whose frequencies are all zero",
));
}
if total > TOTAL as u64 {
return Err(Error::corrupt(
path,
offset,
format!("rans4x8 frequencies summing to {total}, past the {TOTAL} they must fit"),
));
}
let mut cumulative = [0u32; 256];
let mut lookup = vec![0u8; TOTAL];
let mut running = 0u32;
for symbol in 0..256 {
cumulative[symbol] = running;
let f = freq[symbol] as usize;
lookup[running as usize..running as usize + f].fill(symbol as u8);
running += freq[symbol];
}
Ok(Self {
freq,
cumulative,
lookup,
covered: running,
})
}
fn symbol(&self, c: u32, reader: &ByteReader<'_>) -> Result<u8> {
if c >= self.covered {
return Err(Error::corrupt(
reader.path(),
reader.offset(),
format!(
"a rans4x8 state selecting frequency {c}, past the {} its table covers",
self.covered
),
));
}
Ok(self.lookup[c as usize])
}
fn advance(&self, state: u32, symbol: u8, c: u32) -> u32 {
self.freq[symbol as usize] * (state >> TOTAL_BITS) + c - self.cumulative[symbol as usize]
}
}
fn read_symbol_list(
reader: &mut ByteReader<'_>,
mut body: impl FnMut(&mut ByteReader<'_>, u8) -> Result<()>,
) -> Result<()> {
let mut symbol = i32::from(reader.u8()?);
let mut last = symbol;
let mut run = 0u32;
let mut seen = 0usize;
loop {
if symbol > 255 {
return Err(Error::corrupt(
reader.path(),
reader.offset(),
"a rans4x8 run of symbols walking past 255",
));
}
body(reader, symbol as u8)?;
seen += 1;
if seen > 256 {
return Err(Error::corrupt(
reader.path(),
reader.offset(),
"a rans4x8 symbol list longer than the 256 symbols there are",
));
}
if run > 0 {
run -= 1;
symbol += 1;
} else {
symbol = i32::from(reader.u8()?);
if symbol == last + 1 {
run = u32::from(reader.u8()?);
}
}
last = symbol;
if symbol == 0 {
return Ok(());
}
}
}
fn read_frequencies_0(reader: &mut ByteReader<'_>) -> Result<SymbolTable> {
let offset = reader.offset();
let path = reader.path();
let mut freq = [0u32; 256];
read_symbol_list(reader, |reader, symbol| {
freq[symbol as usize] = reader.itf8()?;
Ok(())
})?;
SymbolTable::build(freq, path, offset)
}
fn read_frequencies_1(reader: &mut ByteReader<'_>) -> Result<Vec<Option<SymbolTable>>> {
let mut tables: Vec<Option<SymbolTable>> = (0..256).map(|_| None).collect();
read_symbol_list(reader, |reader, context| {
tables[context as usize] = Some(read_frequencies_0(reader)?);
Ok(())
})?;
Ok(tables)
}
fn renorm(mut state: u32, reader: &mut ByteReader<'_>) -> Result<u32> {
while state < RANS_L {
state = (state << 8) | u32::from(reader.u8()?);
}
Ok(state)
}
fn decode_order_0(reader: &mut ByteReader<'_>, len: usize) -> Result<Vec<u8>> {
let table = read_frequencies_0(reader)?;
let mut states = [0u32; 4];
for state in states.iter_mut() {
*state = reader.u32()?;
}
let mask = (1u32 << TOTAL_BITS) - 1;
let mut out = Vec::with_capacity(len.min(1 << 20));
for i in 0..len {
let j = i & 3;
let c = states[j] & mask;
let symbol = table.symbol(c, reader)?;
out.push(symbol);
states[j] = renorm(table.advance(states[j], symbol, c), reader)?;
}
Ok(out)
}
fn decode_order_1(reader: &mut ByteReader<'_>, len: usize) -> Result<Vec<u8>> {
let tables = read_frequencies_1(reader)?;
let mut states = [0u32; 4];
for state in states.iter_mut() {
*state = reader.u32()?;
}
let mut contexts = [0u8; 4];
let mask = (1u32 << TOTAL_BITS) - 1;
let missing = |context: u8, reader: &ByteReader<'_>| {
Error::corrupt(
reader.path(),
reader.offset(),
format!("a rans4x8 order-1 stream reaching context {context}, which its table omits"),
)
};
let mut out = vec![0u8; len];
let quarter = len / 4;
for i in 0..quarter {
for j in 0..4 {
let table = tables[contexts[j] as usize]
.as_ref()
.ok_or_else(|| missing(contexts[j], reader))?;
let c = states[j] & mask;
let symbol = table.symbol(c, reader)?;
out[i + j * quarter] = symbol;
states[j] = renorm(table.advance(states[j], symbol, c), reader)?;
contexts[j] = symbol;
}
}
for slot in out.iter_mut().take(len).skip(quarter * 4) {
let table = tables[contexts[3] as usize]
.as_ref()
.ok_or_else(|| missing(contexts[3], reader))?;
let c = states[3] & mask;
let symbol = table.symbol(c, reader)?;
*slot = symbol;
states[3] = renorm(table.advance(states[3], symbol, c), reader)?;
contexts[3] = symbol;
}
Ok(out)
}
#[cfg(test)]
mod tests {
use super::*;
fn encode(data: &[u8], order: u8) -> Vec<u8> {
let blob = if order == 0 {
encode_order_0(data)
} else {
encode_order_1(data)
};
let mut out = vec![order];
out.extend_from_slice(&(blob.len() as u32).to_le_bytes());
out.extend_from_slice(&(data.len() as u32).to_le_bytes());
out.extend_from_slice(&blob);
out
}
fn encode_symbol(state: u32, freq: u32, cumulative: u32, bytes: &mut Vec<u8>) -> u32 {
let mut x = state;
let ceiling = ((RANS_L >> TOTAL_BITS) << 8) * freq;
while x >= ceiling {
bytes.push((x & 0xff) as u8);
x >>= 8;
}
((x / freq) << TOTAL_BITS) + cumulative + (x % freq)
}
fn cumulative_of(freq: &[u32; 256]) -> [u32; 256] {
let mut cumulative = [0u32; 256];
let mut running = 0;
for symbol in 0..256 {
cumulative[symbol] = running;
running += freq[symbol];
}
cumulative
}
fn encode_order_0(data: &[u8]) -> Vec<u8> {
let mut freq = [0u32; 256];
for byte in data {
freq[*byte as usize] += 1;
}
normalise(&mut freq);
let cumulative = cumulative_of(&freq);
let mut states = [RANS_L; 4];
let mut bytes = Vec::new();
for i in (0..data.len()).rev() {
let symbol = data[i] as usize;
states[i & 3] =
encode_symbol(states[i & 3], freq[symbol], cumulative[symbol], &mut bytes);
}
bytes.reverse();
let mut out = Vec::new();
write_frequencies_0(&mut out, &freq);
for state in states {
out.extend_from_slice(&state.to_le_bytes());
}
out.extend_from_slice(&bytes);
out
}
fn order_1_plan(len: usize) -> Vec<(usize, usize)> {
let quarter = len / 4;
let mut plan = Vec::with_capacity(len);
for i in 0..quarter {
for j in 0..4 {
plan.push((i + j * quarter, j));
}
}
for i in quarter * 4..len {
plan.push((i, 3));
}
plan
}
fn order_1_context(data: &[u8], index: usize, stream: usize, len: usize) -> u8 {
if index == stream * (len / 4) {
0
} else {
data[index - 1]
}
}
fn encode_order_1(data: &[u8]) -> Vec<u8> {
let len = data.len();
let plan = order_1_plan(len);
let mut freq = vec![[0u32; 256]; 256];
for &(index, stream) in &plan {
let context = order_1_context(data, index, stream, len);
freq[context as usize][data[index] as usize] += 1;
}
for table in freq.iter_mut() {
normalise(table);
}
let cumulative: Vec<[u32; 256]> = freq.iter().map(cumulative_of).collect();
let mut states = [RANS_L; 4];
let mut bytes = Vec::new();
for &(index, stream) in plan.iter().rev() {
let context = order_1_context(data, index, stream, len) as usize;
let symbol = data[index] as usize;
states[stream] = encode_symbol(
states[stream],
freq[context][symbol],
cumulative[context][symbol],
&mut bytes,
);
}
bytes.reverse();
let mut out = Vec::new();
write_frequencies_1(&mut out, &freq);
for state in states {
out.extend_from_slice(&state.to_le_bytes());
}
out.extend_from_slice(&bytes);
out
}
fn normalise(freq: &mut [u32; 256]) {
let total: u64 = freq.iter().map(|f| u64::from(*f)).sum();
if total == 0 {
return;
}
for f in freq.iter_mut() {
if *f > 0 {
*f = ((u64::from(*f) * TOTAL as u64 / total).max(1)) as u32;
}
}
let sum: i64 = freq.iter().map(|f| i64::from(*f)).sum();
let widest = freq
.iter()
.enumerate()
.max_by_key(|(_, f)| **f)
.map(|(s, _)| s)
.expect("256 entries");
freq[widest] = (i64::from(freq[widest]) + TOTAL as i64 - sum) as u32;
}
fn write_itf8(out: &mut Vec<u8>, value: u32) {
if value < 0x80 {
out.push(value as u8);
} else if value < 0x4000 {
out.push(0x80 | (value >> 8) as u8);
out.push(value as u8);
} else {
panic!("this test's frequencies never exceed {TOTAL}");
}
}
fn write_symbol_list(
out: &mut Vec<u8>,
symbols: &[i32],
mut body: impl FnMut(&mut Vec<u8>, i32),
) {
if symbols.is_empty() {
out.push(0);
return;
}
let mut index = 0usize;
let mut symbol = symbols[0];
out.push(symbol as u8);
let mut last = symbol;
let mut run = 0u32;
loop {
body(out, symbol);
if run > 0 {
run -= 1;
symbol += 1;
} else {
index += 1;
let next = symbols.get(index).copied().unwrap_or(0);
out.push(next as u8);
if next == last + 1 {
let mut length = 0u32;
while symbols.get(index + 1 + length as usize).copied()
== Some(next + 1 + length as i32)
{
length += 1;
}
out.push(length as u8);
run = length;
index += length as usize;
}
symbol = next;
last = next;
}
if symbol == 0 {
return;
}
}
}
fn present(freq: &[u32; 256]) -> Vec<i32> {
(0..256).filter(|s| freq[*s as usize] > 0).collect()
}
fn write_frequencies_0(out: &mut Vec<u8>, freq: &[u32; 256]) {
write_symbol_list(out, &present(freq), |out, symbol| {
write_itf8(out, freq[symbol as usize]);
});
}
fn write_frequencies_1(out: &mut Vec<u8>, freq: &[[u32; 256]]) {
let contexts: Vec<i32> = (0..256)
.filter(|c| freq[*c as usize].iter().any(|f| *f > 0))
.collect();
write_symbol_list(out, &contexts, |out, context| {
write_frequencies_0(out, &freq[context as usize]);
});
}
fn roundtrip(data: &[u8], order: u8) {
let encoded = encode(data, order);
let decoded = decode(&encoded, "test", 0)
.unwrap_or_else(|e| panic!("order {order}, {} bytes: {e}", data.len()));
assert_eq!(decoded, data, "order {order}, {} bytes", data.len());
}
#[test]
fn order_0_round_trips_text() {
let data = b"abracadabra".repeat(20);
roundtrip(&data, 0);
}
#[test]
fn order_1_round_trips_text() {
let data = b"abracadabra".repeat(20);
roundtrip(&data, 1);
}
#[test]
fn lengths_that_are_not_a_multiple_of_four_round_trip() {
for extra in 0..4 {
let data: Vec<u8> = (0..40 + extra).map(|i| b"acgtn"[i % 5]).collect();
roundtrip(&data, 0);
roundtrip(&data, 1);
}
}
#[test]
fn a_stream_shorter_than_its_four_states_round_trips() {
for len in 1..=4usize {
let data: Vec<u8> = (0..len).map(|i| b"acgt"[i]).collect();
roundtrip(&data, 0);
roundtrip(&data, 1);
}
}
#[test]
fn a_single_symbol_stream_round_trips() {
roundtrip(&[7u8; 33], 0);
roundtrip(&[7u8; 33], 1);
}
#[test]
fn an_alphabet_that_reaches_255_does_not_eat_the_byte_after_its_terminator() {
let data: Vec<u8> = (0..=255u8).chain(250..=255u8).collect();
roundtrip(&data, 0);
roundtrip(&data, 1);
}
#[test]
fn an_alphabet_starting_at_zero_round_trips() {
let data: Vec<u8> = (0..64u8).map(|i| i % 3).collect();
roundtrip(&data, 0);
roundtrip(&data, 1);
}
#[test]
fn the_specifications_abracadabra_frequency_table_reads_as_documented() {
let bytes = [
0x61, 0x87, 0x47, 0x62, 0x02, 0x82, 0xe8, 0x81, 0x74, 0x81, 0x74, 0x72, 0x82, 0xe8, 0x00,
];
let mut reader = ByteReader::new(&bytes, "spec", 0);
let table = read_frequencies_0(&mut reader).expect("the spec's own table");
assert_eq!(table.freq[b'a' as usize], 1863);
assert_eq!(table.freq[b'b' as usize], 744);
assert_eq!(table.freq[b'c' as usize], 372);
assert_eq!(table.freq[b'd' as usize], 372);
assert_eq!(table.freq[b'r' as usize], 744);
assert_eq!(table.covered, 1863 + 744 + 372 + 372 + 744);
assert_eq!(table.freq.iter().filter(|f| **f > 0).count(), 5);
assert!(reader.is_empty());
}
#[test]
fn the_specifications_order_1_frequency_table_reads_as_documented() {
let bytes = [
0x00, 0x61, 0x8f, 0xff, 0x00, 0x61, 0x61, 0x82, 0x66, 0x62, 0x02, 0x86, 0x67, 0x83, 0x33, 0x83, 0xff, 0x00, 0x62, 0x02, 0x72, 0x8f, 0xff, 0x00, 0x61, 0x8f, 0xff, 0x00, 0x61, 0x8f, 0xff, 0x00, 0x72, 0x61, 0x8f, 0xff, 0x00, 0x00, ];
let mut reader = ByteReader::new(&bytes, "spec", 0);
let tables = read_frequencies_1(&mut reader).expect("the spec's own table");
let table = |context: u8| tables[context as usize].as_ref().expect("context");
assert_eq!(table(0).freq[b'a' as usize], 4095);
assert_eq!(table(b'a').freq[b'a' as usize], 614);
assert_eq!(table(b'a').freq[b'b' as usize], 1639);
assert_eq!(table(b'a').freq[b'c' as usize], 819);
assert_eq!(table(b'a').freq[b'd' as usize], 1023);
assert_eq!(table(b'b').freq[b'r' as usize], 4095);
assert_eq!(table(b'c').freq[b'a' as usize], 4095);
assert_eq!(table(b'd').freq[b'a' as usize], 4095);
assert_eq!(table(b'r').freq[b'a' as usize], 4095);
assert_eq!(tables.iter().filter(|t| t.is_some()).count(), 6);
assert!(reader.is_empty());
}
#[test]
fn a_state_landing_past_the_frequencies_is_refused() {
let mut bytes = vec![0u8; 0];
bytes.push(0); let mut blob = Vec::new();
let mut freq = [0u32; 256];
freq[b'a' as usize] = 4095;
write_frequencies_0(&mut blob, &freq);
blob.extend_from_slice(&0x0000_0fffu32.to_le_bytes()); for _ in 0..3 {
blob.extend_from_slice(&RANS_L.to_le_bytes());
}
blob.extend_from_slice(&[0u8; 64]);
bytes.extend_from_slice(&(blob.len() as u32).to_le_bytes());
bytes.extend_from_slice(&4u32.to_le_bytes());
bytes.extend_from_slice(&blob);
let error = decode(&bytes, "test", 0).expect_err("the slot belongs to no symbol");
assert!(
error.to_string().contains("past the 4095 its table covers"),
"{error}"
);
}
#[test]
fn a_block_of_an_unknown_order_is_refused() {
let mut bytes = vec![2u8];
bytes.extend_from_slice(&0u32.to_le_bytes());
bytes.extend_from_slice(&8u32.to_le_bytes());
let error = decode(&bytes, "test", 0).expect_err("order 2 does not exist");
assert!(error.to_string().contains("neither 0 nor 1"), "{error}");
}
#[test]
fn a_truncated_block_is_refused_rather_than_decoded() {
let encoded = encode(&b"abracadabra".repeat(20), 0);
for cut in [9, 12, 20, encoded.len() - 1] {
let error = decode(&encoded[..cut], "test", 0);
assert!(error.is_err(), "a block cut to {cut} bytes decoded");
}
}
#[test]
fn every_prefix_of_a_real_stream_fails_without_panicking() {
let encoded = encode(&b"abracadabra".repeat(20), 1);
for cut in 0..encoded.len() {
let _ = decode(&encoded[..cut], "test", 0);
}
for byte in 0..=255u8 {
let mut damaged = encoded.clone();
damaged[9] = byte;
let _ = decode(&damaged, "test", 0);
}
}
}