use crate::error::{Error, Result};
use super::{ByteReader, PackMeta, MAX_CODEC_LEN};
mod flag {
pub const ORDER: u8 = 1;
pub const N32: u8 = 4;
pub const STRIPE: u8 = 8;
pub const NO_SIZE: u8 = 16;
pub const CAT: u8 = 32;
pub const RLE: u8 = 64;
pub const PACK: u8 = 128;
}
const RANS_L: u32 = 1 << 15;
const MAX_STRIPE_DEPTH: u32 = 4;
pub fn decode(data: &[u8], path: &str, offset: u64) -> Result<Vec<u8>> {
let mut reader = ByteReader::new(data, path, offset);
decode_stream(&mut reader, None, 0)
}
fn decode_stream(
reader: &mut ByteReader<'_>,
known_len: Option<usize>,
depth: u32,
) -> Result<Vec<u8>> {
let flags = reader.u8()?;
let mut len = if flags & flag::NO_SIZE == 0 {
reader.length()?
} else {
known_len.ok_or_else(|| {
Error::corrupt(
reader.path(),
reader.offset(),
"a rans4x16 stream sets NoSize but nothing outside it knows the size",
)
})?
};
if flags & flag::STRIPE != 0 {
return decode_stripe(reader, len, depth);
}
let n_states = if flags & flag::N32 != 0 { 32 } else { 4 };
let mut pack = None;
if flags & flag::PACK != 0 {
let unpacked_len = len;
let meta = PackMeta::read(reader)?;
len = meta.packed_len;
pack = Some((meta, unpacked_len));
}
let mut rle = None;
if flags & flag::RLE != 0 {
let expanded_len = len;
let meta = RleMeta::read(reader, n_states)?;
len = meta.pre_expansion_len;
rle = Some((meta, expanded_len));
}
if len > MAX_CODEC_LEN {
return Err(Error::corrupt(
reader.path(),
reader.offset(),
format!("a rans4x16 stream declares {len} bytes, past this reader's ceiling"),
));
}
let mut data = if flags & flag::CAT != 0 {
reader.take(len)?.to_vec()
} else if flags & flag::ORDER != 0 {
decode_order_1(reader, len, n_states)?
} else {
decode_order_0(reader, len, n_states)?
};
if let Some((meta, expanded_len)) = rle {
data = meta.expand(&data, expanded_len, reader.path(), reader.offset())?;
}
if let Some((meta, unpacked_len)) = pack {
data = meta.unpack(&data, unpacked_len, reader.path(), reader.offset())?;
}
Ok(data)
}
fn decode_stripe(reader: &mut ByteReader<'_>, len: usize, depth: u32) -> Result<Vec<u8>> {
if depth >= MAX_STRIPE_DEPTH {
return Err(Error::corrupt(
reader.path(),
reader.offset(),
"rans4x16 stripes nested past this reader's limit",
));
}
let n = reader.u8()? as usize;
if n == 0 {
return Err(Error::corrupt(
reader.path(),
reader.offset(),
"a rans4x16 stripe of zero sub-streams",
));
}
let mut lengths = Vec::with_capacity(n);
for _ in 0..n {
lengths.push(reader.length()?);
}
let mut out = vec![0u8; len];
for (j, compressed_len) in lengths.into_iter().enumerate() {
let sub_len = len / n + usize::from(len % n > j);
let bytes = reader.take(compressed_len)?;
let mut sub = ByteReader::new(bytes, reader.path(), reader.offset());
let decoded = decode_stream(&mut sub, Some(sub_len), depth + 1)?;
if decoded.len() != sub_len {
return Err(Error::corrupt(
reader.path(),
reader.offset(),
format!(
"a rans4x16 stripe sub-stream gave {} bytes where {sub_len} were due",
decoded.len()
),
));
}
for (i, byte) in decoded.into_iter().enumerate() {
out[i * n + j] = byte;
}
}
Ok(out)
}
struct SymbolTable {
freq: [u32; 256],
cumulative: [u32; 256],
lookup: Vec<u8>,
}
impl SymbolTable {
fn build(mut freq: [u32; 256], bits: u32, path: &str, offset: u64) -> Result<Self> {
let total: u64 = freq.iter().map(|f| u64::from(*f)).sum();
let target = 1u64 << bits;
if total == 0 {
return Err(Error::corrupt(
path,
offset,
"a rans4x16 frequency table whose frequencies are all zero",
));
}
if total > target {
return Err(Error::corrupt(
path,
offset,
format!("rans4x16 frequencies summing to {total}, past the {target} they must fit"),
));
}
let mut scaled = total;
let mut shift = 0;
while scaled < target {
scaled *= 2;
shift += 1;
}
if scaled != target {
return Err(Error::corrupt(
path,
offset,
format!("rans4x16 frequencies summing to {total}, which is not a power of two"),
));
}
let target = target as u32;
for f in freq.iter_mut() {
*f <<= shift;
}
let mut cumulative = [0u32; 256];
let mut lookup = vec![0u8; target as usize];
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,
})
}
}
fn read_alphabet(reader: &mut ByteReader<'_>) -> Result<Vec<u8>> {
let mut alphabet = Vec::new();
let mut symbol = i32::from(reader.u8()?);
let mut last = symbol;
let mut run = 0u32;
loop {
if symbol > 255 {
break;
}
alphabet.push(symbol as u8);
if alphabet.len() > 256 {
return Err(Error::corrupt(
reader.path(),
reader.offset(),
"a rans4x16 alphabet of more than 256 symbols",
));
}
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 {
break;
}
}
Ok(alphabet)
}
fn read_frequencies_0(reader: &mut ByteReader<'_>) -> Result<SymbolTable> {
let alphabet = read_alphabet(reader)?;
let mut freq = [0u32; 256];
for symbol in alphabet {
freq[symbol as usize] = reader.uint7()?;
}
SymbolTable::build(freq, 12, reader.path(), reader.offset())
}
fn read_frequencies_1(reader: &mut ByteReader<'_>) -> Result<(Vec<Option<SymbolTable>>, u32)> {
let comp = reader.u8()?;
let bits = u32::from(comp >> 4);
if bits != 10 && bits != 12 {
return Err(Error::corrupt(
reader.path(),
reader.offset(),
format!("a rans4x16 order-1 frequency table scaled to {bits} bits, not 10 or 12"),
));
}
let tables = if comp & 1 != 0 {
let raw_len = reader.length()?;
let compressed_len = reader.length()?;
let bytes = reader.take(compressed_len)?;
let mut inner = ByteReader::new(bytes, reader.path(), reader.offset());
let table_bytes = decode_order_0(&mut inner, raw_len, 4)?;
let mut table_reader = ByteReader::new(&table_bytes, reader.path(), reader.offset());
read_frequencies_1_tables(&mut table_reader, bits)?
} else {
read_frequencies_1_tables(reader, bits)?
};
Ok((tables, bits))
}
fn read_frequencies_1_tables(
reader: &mut ByteReader<'_>,
bits: u32,
) -> Result<Vec<Option<SymbolTable>>> {
let alphabet = read_alphabet(reader)?;
let mut tables: Vec<Option<SymbolTable>> = (0..256).map(|_| None).collect();
for &context in &alphabet {
let mut freq = [0u32; 256];
let mut run = 0u32;
for &symbol in &alphabet {
if run > 0 {
run -= 1;
continue;
}
let f = reader.uint7()?;
freq[symbol as usize] = f;
if f == 0 {
run = u32::from(reader.u8()?);
}
}
if freq.iter().any(|f| *f > 0) {
tables[context as usize] = Some(SymbolTable::build(
freq,
bits,
reader.path(),
reader.offset(),
)?);
}
}
Ok(tables)
}
#[inline]
fn renorm(state: u32, reader: &mut ByteReader<'_>) -> Result<u32> {
if state < RANS_L {
return Ok((state << 16) + u32::from(reader.u16()?));
}
Ok(state)
}
fn decode_order_0(reader: &mut ByteReader<'_>, len: usize, n: usize) -> Result<Vec<u8>> {
let table = read_frequencies_0(reader)?;
let mut states = Vec::with_capacity(n);
for _ in 0..n {
states.push(reader.u32()?);
}
let mask = (1u32 << 12) - 1;
let mut out = Vec::with_capacity(len.min(1 << 20));
for i in 0..len {
let j = i % n;
let state = states[j];
let symbol = table.lookup[(state & mask) as usize];
let f = table.freq[symbol as usize];
let c = table.cumulative[symbol as usize];
out.push(symbol);
states[j] = renorm(f * (state >> 12) + (state & mask) - c, reader)?;
}
Ok(out)
}
fn decode_order_1(reader: &mut ByteReader<'_>, len: usize, n: usize) -> Result<Vec<u8>> {
let (tables, bits) = read_frequencies_1(reader)?;
let mut states = Vec::with_capacity(n);
for _ in 0..n {
states.push(reader.u32()?);
}
let mut contexts = vec![0u8; n];
let mask = (1u32 << bits) - 1;
let missing = |context: u8, reader: &ByteReader<'_>| {
Error::corrupt(
reader.path(),
reader.offset(),
format!("a rans4x16 order-1 stream reaching context {context}, which its table omits"),
)
};
let mut out = vec![0u8; len];
let stride = len / n;
for i in 0..stride {
for j in 0..n {
let table = tables[contexts[j] as usize]
.as_ref()
.ok_or_else(|| missing(contexts[j], reader))?;
let state = states[j];
let symbol = table.lookup[(state & mask) as usize];
let f = table.freq[symbol as usize];
let c = table.cumulative[symbol as usize];
out[i + j * stride] = symbol;
states[j] = renorm(f * (state >> bits) + (state & mask) - c, reader)?;
contexts[j] = symbol;
}
}
let last = n - 1;
for slot in out.iter_mut().take(len).skip(stride * n) {
let table = tables[contexts[last] as usize]
.as_ref()
.ok_or_else(|| missing(contexts[last], reader))?;
let state = states[last];
let symbol = table.lookup[(state & mask) as usize];
let f = table.freq[symbol as usize];
let c = table.cumulative[symbol as usize];
*slot = symbol;
states[last] = renorm(f * (state >> bits) + (state & mask) - c, reader)?;
contexts[last] = symbol;
}
Ok(out)
}
struct RleMeta {
has_run: [bool; 256],
runs: Vec<u8>,
pre_expansion_len: usize,
}
impl RleMeta {
fn read(reader: &mut ByteReader<'_>, n: usize) -> Result<Self> {
let meta_len = reader.length()?;
let pre_expansion_len = reader.length()?;
let raw_meta_len = meta_len / 2;
let meta = if meta_len & 1 != 0 {
reader.take(raw_meta_len)?.to_vec()
} else {
let compressed_len = reader.length()?;
let bytes = reader.take(compressed_len)?;
let mut inner = ByteReader::new(bytes, reader.path(), reader.offset());
decode_order_0(&mut inner, raw_meta_len, n)?
};
let mut meta = ByteReader::new(&meta, reader.path(), reader.offset());
let count = meta.u8()?;
let count = if count == 0 { 256 } else { count as usize };
let mut has_run = [false; 256];
for _ in 0..count {
has_run[meta.u8()? as usize] = true;
}
Ok(Self {
has_run,
runs: meta.take(meta.remaining())?.to_vec(),
pre_expansion_len,
})
}
fn expand(&self, data: &[u8], len: usize, path: &str, offset: u64) -> Result<Vec<u8>> {
let mut runs = ByteReader::new(&self.runs, path, offset);
let mut out = Vec::with_capacity(len.min(1 << 20));
for &symbol in data {
if out.len() >= len {
break;
}
if self.has_run[symbol as usize] {
let run = runs.uint7()? as usize;
let take = (run + 1).min(len - out.len());
out.resize(out.len() + take, symbol);
} else {
out.push(symbol);
}
}
if out.len() != len {
return Err(Error::corrupt(
path,
offset,
format!(
"a rans4x16 run-length stream expanded to {} bytes where {len} were due",
out.len()
),
));
}
Ok(out)
}
}
#[cfg(test)]
mod tests {
use super::*;
fn encode_order_0(data: &[u8], n: usize) -> Vec<u8> {
let mut freq = [0u32; 256];
for &b in data {
freq[b as usize] += 1;
}
normalise(&mut freq, 12);
let mut cumulative = [0u32; 256];
let mut running = 0;
for s in 0..256 {
cumulative[s] = running;
running += freq[s];
}
let mut words: Vec<u16> = Vec::new();
let mut states = vec![RANS_L; n];
for i in (0..data.len()).rev() {
let j = i % n;
let s = data[i] as usize;
let (f, c) = (freq[s], cumulative[s]);
let max = ((RANS_L >> 12) << 16) * f;
if states[j] >= max {
words.push((states[j] & 0xffff) as u16);
states[j] >>= 16;
}
states[j] = ((states[j] / f) << 12) + (states[j] % f) + c;
}
let mut out = Vec::new();
out.push(if n == 32 { flag::N32 } else { 0 });
write_uint7(&mut out, data.len() as u32);
write_frequencies(&mut out, &freq);
for state in &states {
out.extend_from_slice(&state.to_le_bytes());
}
for word in words.iter().rev() {
out.extend_from_slice(&word.to_le_bytes());
}
out
}
fn encode_order_1(data: &[u8], n: usize) -> Vec<u8> {
let len = data.len();
let stride = len / n;
assert!(stride > 0, "order-1 wants at least one byte per state");
let context_of = |i: usize| -> u8 {
if i >= stride * n {
if i == stride * n {
if stride == 0 {
0
} else {
data[(n - 1) * stride + stride - 1]
}
} else {
data[i - 1]
}
} else if i % stride == 0 {
0
} else {
data[i - 1]
}
};
let mut freq = [[0u32; 256]; 256];
for i in 0..len {
freq[context_of(i) as usize][data[i] as usize] += 1;
}
let used: Vec<u8> = freq
.iter()
.enumerate()
.filter(|(symbol, row)| row.iter().any(|f| *f > 0) || data.contains(&(*symbol as u8)))
.map(|(symbol, _)| symbol as u8)
.collect();
for row in freq.iter_mut() {
if row.iter().any(|f| *f > 0) {
normalise(row, 12);
}
}
let mut cumulative = [[0u32; 256]; 256];
for s in 0..256 {
let mut running = 0;
for t in 0..256 {
cumulative[s][t] = running;
running += freq[s][t];
}
}
let mut words: Vec<u16> = Vec::new();
let mut states = vec![RANS_L; n];
let push = |i: usize, states: &mut Vec<u32>, words: &mut Vec<u16>, j: usize| {
let ctx = context_of(i) as usize;
let s = data[i] as usize;
let (f, c) = (freq[ctx][s], cumulative[ctx][s]);
let max = ((RANS_L >> 12) << 16) * f;
if states[j] >= max {
words.push((states[j] & 0xffff) as u16);
states[j] >>= 16;
}
states[j] = ((states[j] / f) << 12) + (states[j] % f) + c;
};
for i in (stride * n..len).rev() {
push(i, &mut states, &mut words, n - 1);
}
for i in (0..stride).rev() {
for j in (0..n).rev() {
push(i + j * stride, &mut states, &mut words, j);
}
}
let mut out = Vec::new();
out.push(flag::ORDER | if n == 32 { flag::N32 } else { 0 });
write_uint7(&mut out, len as u32);
out.push(12 << 4); write_alphabet(&mut out, &used);
for &ctx in &used {
let mut run: usize = 0;
let mut pending: Vec<u8> = Vec::new();
for (k, &sym) in used.iter().enumerate() {
if run > 0 {
run -= 1;
continue;
}
let f = freq[ctx as usize][sym as usize];
write_uint7(&mut pending, f);
if f == 0 {
let mut zeros = 0u8;
for &next in &used[k + 1..] {
if freq[ctx as usize][next as usize] == 0 && zeros < 255 {
zeros += 1;
} else {
break;
}
}
pending.push(zeros);
run = zeros as usize;
}
}
out.extend_from_slice(&pending);
}
for state in &states {
out.extend_from_slice(&state.to_le_bytes());
}
for word in words.iter().rev() {
out.extend_from_slice(&word.to_le_bytes());
}
out
}
fn normalise(freq: &mut [u32; 256], bits: u32) {
let target = 1u32 << bits;
let total: u32 = freq.iter().sum();
if total == 0 {
return;
}
let mut running = 0u32;
let mut last = 0usize;
for (symbol, count) in freq.iter_mut().enumerate() {
if *count == 0 {
continue;
}
let scaled = ((*count as u64 * target as u64) / total as u64).max(1) as u32;
*count = scaled;
running += scaled;
last = symbol;
}
while running > target {
let take = (running - target).min(freq[last] - 1);
freq[last] -= take;
running -= take;
if freq[last] == 1 {
last = (0..256).rev().find(|s| freq[*s] > 1).unwrap_or(last);
}
}
if running < target {
freq[last] += target - running;
}
}
fn write_uint7(out: &mut Vec<u8>, value: u32) {
let mut groups = Vec::new();
let mut value = value;
loop {
groups.push((value & 0x7f) as u8);
value >>= 7;
if value == 0 {
break;
}
}
for (i, group) in groups.iter().enumerate().rev() {
out.push(if i == 0 { *group } else { group | 0x80 });
}
}
fn write_alphabet(out: &mut Vec<u8>, symbols: &[u8]) {
let mut i = 0;
while i < symbols.len() {
out.push(symbols[i]);
if i + 1 < symbols.len() && symbols[i + 1] == symbols[i].wrapping_add(1) {
let mut run = 0u8;
while run < 255
&& i + 1 + (run as usize) < symbols.len()
&& symbols[i + 1 + run as usize] == symbols[i].wrapping_add(1).wrapping_add(run)
{
run += 1;
}
out.push(symbols[i + 1]);
out.push(run - 1);
i += 1 + run as usize;
} else {
i += 1;
}
}
out.push(0);
}
fn write_frequencies(out: &mut Vec<u8>, freq: &[u32; 256]) {
let symbols: Vec<u8> = (0..256).filter(|s| freq[*s] > 0).map(|s| s as u8).collect();
write_alphabet(out, &symbols);
for &s in &symbols {
write_uint7(out, freq[s as usize]);
}
}
fn roundtrip_0(data: &[u8], n: usize) {
let encoded = encode_order_0(data, n);
let decoded = decode(&encoded, "test", 0).expect("decodes");
assert_eq!(decoded, data, "order-0, N={n}");
}
fn roundtrip_1(data: &[u8], n: usize) {
let encoded = encode_order_1(data, n);
let decoded = decode(&encoded, "test", 0).expect("decodes");
assert_eq!(decoded, data, "order-1, N={n}");
}
#[test]
fn order_0_round_trips_at_both_interleavings() {
let data: Vec<u8> = (0..5000u32).map(|i| (i * 7 % 41) as u8).collect();
roundtrip_0(&data, 4);
roundtrip_0(&data, 32);
}
#[test]
fn order_0_round_trips_text() {
let data =
b"the quick brown fox jumps over the lazy dog, repeatedly and at length. ".repeat(40);
roundtrip_0(&data, 4);
roundtrip_0(&data, 32);
}
#[test]
fn order_1_round_trips_at_both_interleavings() {
let data = b"ACGTACGTTTTTACGNNNNNNACGTACGTACGGGGTTTACGT".repeat(120);
roundtrip_1(&data, 4);
roundtrip_1(&data, 32);
}
#[test]
fn order_1_round_trips_a_length_that_is_not_a_multiple_of_n() {
let base = b"ACGTNACGTNACGTTTTACG".repeat(30);
for extra in 0..8 {
let data = &base[..base.len() - extra];
roundtrip_1(data, 4);
}
}
#[test]
fn an_alphabet_that_reaches_255_does_not_eat_the_byte_after_its_terminator() {
let mut symbols = vec![0u8];
symbols.extend(226..=255);
let mut out = Vec::new();
write_alphabet(&mut out, &symbols);
let mut reader = ByteReader::new(&out, "test", 0);
assert_eq!(read_alphabet(&mut reader).expect("an alphabet"), symbols);
assert_eq!(reader.remaining(), 0);
let data: Vec<u8> = (0..4000).map(|i| symbols[i % symbols.len()]).collect();
roundtrip_0(&data, 4);
roundtrip_1(&data, 4);
}
#[test]
fn an_alphabet_starting_at_zero_round_trips() {
for symbols in [vec![0u8], vec![0, 1], vec![0, 1, 2], vec![0, 5], vec![255]] {
let mut out = Vec::new();
write_alphabet(&mut out, &symbols);
let mut reader = ByteReader::new(&out, "test", 0);
assert_eq!(
read_alphabet(&mut reader).expect("an alphabet"),
symbols,
"{symbols:?}"
);
}
}
#[test]
fn a_single_symbol_stream_round_trips() {
let data = vec![b'Q'; 1000];
roundtrip_0(&data, 4);
roundtrip_1(&data, 4);
}
#[test]
fn the_cat_flag_gives_the_bytes_back_verbatim() {
let mut stream = vec![flag::CAT];
write_uint7(&mut stream, 5);
stream.extend_from_slice(b"hello");
assert_eq!(decode(&stream, "test", 0).expect("decodes"), b"hello");
}
#[test]
fn a_pack_of_one_symbol_expands_without_reading_any_data() {
let mut stream = vec![flag::PACK | flag::CAT];
write_uint7(&mut stream, 6); stream.push(1); stream.push(b'N');
write_uint7(&mut stream, 0); assert_eq!(decode(&stream, "test", 0).expect("decodes"), b"NNNNNN");
}
#[test]
fn a_two_symbol_pack_unpacks_eight_values_to_the_byte() {
let mut stream = vec![flag::PACK | flag::CAT];
write_uint7(&mut stream, 8);
stream.push(2);
stream.extend_from_slice(b"AB");
write_uint7(&mut stream, 1);
stream.push(0b0000_1101);
assert_eq!(decode(&stream, "test", 0).expect("decodes"), b"BABBAAAA");
}
#[test]
fn a_run_length_stream_expands_its_runs() {
let mut meta = vec![1u8, b'N'];
write_uint7(&mut meta, 3); let mut stream = vec![flag::RLE | flag::CAT];
write_uint7(&mut stream, 6); write_uint7(&mut stream, (meta.len() * 2 + 1) as u32); write_uint7(&mut stream, 3); stream.extend_from_slice(&meta);
stream.extend_from_slice(b"ANA");
assert_eq!(decode(&stream, "test", 0).expect("decodes"), b"ANNNNA");
}
#[test]
fn a_stripe_transposes_its_sub_streams_back() {
let sub0 = {
let mut s = vec![flag::CAT | flag::NO_SIZE];
s.extend_from_slice(b"abc");
s
};
let sub1 = {
let mut s = vec![flag::CAT | flag::NO_SIZE];
s.extend_from_slice(b"AB");
s
};
let mut stream = vec![flag::STRIPE];
write_uint7(&mut stream, 5);
stream.push(2);
write_uint7(&mut stream, sub0.len() as u32);
write_uint7(&mut stream, sub1.len() as u32);
stream.extend_from_slice(&sub0);
stream.extend_from_slice(&sub1);
assert_eq!(decode(&stream, "test", 0).expect("decodes"), b"aAbBc");
}
#[test]
fn a_stripe_nested_past_the_limit_is_refused_rather_than_recursed() {
let mut stream = Vec::new();
for _ in 0..MAX_STRIPE_DEPTH + 2 {
let mut outer = vec![flag::STRIPE];
write_uint7(&mut outer, 4);
outer.push(1);
write_uint7(&mut outer, stream.len() as u32);
outer.extend_from_slice(&stream);
stream = outer;
}
assert!(decode(&stream, "test", 0).is_err());
}
#[test]
fn truncated_streams_are_errors_rather_than_short_output() {
let data: Vec<u8> = (0..2000u32).map(|i| (i % 37) as u8).collect();
let encoded = encode_order_0(&data, 4);
for cut in [1, 5, 20, encoded.len() / 2, encoded.len() - 1] {
assert!(
decode(&encoded[..cut], "test", 0).is_err(),
"a stream cut to {cut} bytes decoded"
);
}
}
}