use crate::error::{Error, Result};
use super::{bzip2, ByteReader, PackMeta, MAX_CODEC_LEN};
mod flag {
pub const ORDER: u8 = 1;
pub const EXT: 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 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(),
"an arith stream sets NoSize but nothing outside it knows the size",
)
})?
};
if flags & flag::STRIPE != 0 {
return decode_stripe(reader, len, depth);
}
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));
}
if len > MAX_CODEC_LEN {
return Err(Error::corrupt(
reader.path(),
reader.offset(),
format!("an arith 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::EXT != 0 {
bzip2::decode(
reader.take(reader.remaining())?,
len,
reader.path(),
reader.offset(),
)?
} else {
let rle = flags & flag::RLE != 0;
let order_1 = flags & flag::ORDER != 0;
match (rle, order_1) {
(false, false) => decode_order_0(reader, len)?,
(false, true) => decode_order_1(reader, len)?,
(true, false) => decode_rle_0(reader, len)?,
(true, true) => decode_rle_1(reader, len)?,
}
};
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(),
"arith stripes nested past this reader's limit",
));
}
let n = reader.u8()? as usize;
if n == 0 {
return Err(Error::corrupt(
reader.path(),
reader.offset(),
"an arith 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!(
"an arith 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)
}
pub(super) struct RangeCoder {
range: u32,
code: u32,
}
impl RangeCoder {
pub(super) fn new(reader: &mut ByteReader<'_>) -> Result<Self> {
let mut code: u32 = 0;
for _ in 0..5 {
code = (code << 8) | u32::from(reader.u8()?);
}
Ok(Self {
range: u32::MAX,
code,
})
}
pub(super) fn frequency(&mut self, total: u32) -> u32 {
self.range /= total;
self.code / self.range
}
pub(super) fn decode(
&mut self,
low: u32,
freq: u32,
reader: &mut ByteReader<'_>,
) -> Result<()> {
self.code = self.code.wrapping_sub(low.wrapping_mul(self.range));
self.range = self.range.wrapping_mul(freq);
while self.range < (1 << 24) {
self.range <<= 8;
self.code = (self.code << 8) | u32::from(reader.u8()?);
}
Ok(())
}
}
pub(super) struct Model {
symbols: Vec<u8>,
freq: Vec<u32>,
total: u32,
}
const MAX_TOTAL: u32 = (1 << 16) - 17;
const STEP: u32 = 16;
impl Model {
pub(super) fn new(n_symbols: usize) -> Self {
Self {
symbols: (0..n_symbols).map(|s| s as u8).collect(),
freq: vec![1; n_symbols],
total: n_symbols as u32,
}
}
fn renormalise(&mut self) {
self.total = 0;
for f in self.freq.iter_mut() {
*f -= *f / 2;
self.total += *f;
}
}
pub(super) fn decode(
&mut self,
rc: &mut RangeCoder,
reader: &mut ByteReader<'_>,
) -> Result<u8> {
let target = rc.frequency(self.total);
let mut acc = 0u32;
let mut x = 0usize;
while x < self.freq.len() && acc + self.freq[x] <= target {
acc += self.freq[x];
x += 1;
}
if x >= self.freq.len() {
return Err(Error::corrupt(
reader.path(),
reader.offset(),
format!(
"an arith model selecting frequency {target} of {}, which no symbol covers",
self.total
),
));
}
rc.decode(acc, self.freq[x], reader)?;
let symbol = self.symbols[x];
self.freq[x] += STEP;
self.total += STEP;
if self.total > MAX_TOTAL {
self.renormalise();
}
if x > 0 && self.freq[x] > self.freq[x - 1] {
self.freq.swap(x, x - 1);
self.symbols.swap(x, x - 1);
}
Ok(symbol)
}
}
fn read_max_sym(reader: &mut ByteReader<'_>) -> Result<usize> {
let max_sym = reader.u8()?;
Ok(if max_sym == 0 { 256 } else { max_sym as usize })
}
fn decode_order_0(reader: &mut ByteReader<'_>, len: usize) -> Result<Vec<u8>> {
let n_symbols = read_max_sym(reader)?;
let mut model = Model::new(n_symbols);
let mut rc = RangeCoder::new(reader)?;
let mut out = Vec::with_capacity(len.min(1 << 20));
for _ in 0..len {
out.push(model.decode(&mut rc, reader)?);
}
Ok(out)
}
fn decode_order_1(reader: &mut ByteReader<'_>, len: usize) -> Result<Vec<u8>> {
let n_symbols = read_max_sym(reader)?;
let mut models: Vec<Model> = (0..n_symbols).map(|_| Model::new(n_symbols)).collect();
let mut rc = RangeCoder::new(reader)?;
let mut out = Vec::with_capacity(len.min(1 << 20));
let mut last = 0usize;
for _ in 0..len {
let model = models.get_mut(last).ok_or_else(|| context(reader, last))?;
let symbol = model.decode(&mut rc, reader)?;
out.push(symbol);
last = symbol as usize;
}
Ok(out)
}
fn context(reader: &ByteReader<'_>, symbol: usize) -> Error {
Error::corrupt(
reader.path(),
reader.offset(),
format!("an arith order-1 stream reaching context {symbol}, past the alphabet it declared"),
)
}
fn run_models() -> Vec<Model> {
(0..258).map(|_| Model::new(4)).collect()
}
fn decode_run(
runs: &mut [Model],
first_context: usize,
rc: &mut RangeCoder,
reader: &mut ByteReader<'_>,
) -> Result<usize> {
let mut part = runs[first_context].decode(rc, reader)? as usize;
let mut run = part;
let mut context = 256;
while part == 3 {
part = runs[context].decode(rc, reader)? as usize;
context = 257;
run += part;
if run > MAX_CODEC_LEN {
return Err(Error::corrupt(
reader.path(),
reader.offset(),
"an arith run longer than this reader's ceiling",
));
}
}
Ok(run)
}
fn decode_rle_0(reader: &mut ByteReader<'_>, len: usize) -> Result<Vec<u8>> {
let n_symbols = read_max_sym(reader)?;
let mut literals = Model::new(n_symbols);
let mut runs = run_models();
let mut rc = RangeCoder::new(reader)?;
let mut out = Vec::with_capacity(len.min(1 << 20));
while out.len() < len {
let symbol = literals.decode(&mut rc, reader)?;
let run = decode_run(&mut runs, symbol as usize, &mut rc, reader)?;
let wanted = (run + 1).min(len - out.len());
out.resize(out.len() + wanted, symbol);
}
Ok(out)
}
fn decode_rle_1(reader: &mut ByteReader<'_>, len: usize) -> Result<Vec<u8>> {
let n_symbols = read_max_sym(reader)?;
let mut literals: Vec<Model> = (0..n_symbols).map(|_| Model::new(n_symbols)).collect();
let mut runs = run_models();
let mut rc = RangeCoder::new(reader)?;
let mut out = Vec::with_capacity(len.min(1 << 20));
let mut last = 0usize;
while out.len() < len {
let model = literals
.get_mut(last)
.ok_or_else(|| context(reader, last))?;
let symbol = model.decode(&mut rc, reader)?;
last = symbol as usize;
let run = decode_run(&mut runs, symbol as usize, &mut rc, reader)?;
let wanted = (run + 1).min(len - out.len());
out.resize(out.len() + wanted, symbol);
}
Ok(out)
}
#[cfg(test)]
pub(crate) mod testing {
use super::{MAX_TOTAL, STEP};
pub(crate) struct RangeEncoder {
low: u64,
range: u32,
cache: u8,
ff_num: u64,
out: Vec<u8>,
}
impl RangeEncoder {
pub(crate) fn new() -> Self {
Self {
low: 0,
range: u32::MAX,
cache: 0,
ff_num: 0,
out: Vec::new(),
}
}
pub(crate) fn shift_low(&mut self) {
if self.low < 0xff00_0000 || self.low > 0xffff_ffff {
let carry = (self.low >> 32) as u8;
self.out.push(self.cache.wrapping_add(carry));
while self.ff_num > 0 {
self.out.push(0xffu8.wrapping_add(carry));
self.ff_num -= 1;
}
self.cache = (self.low >> 24) as u8;
} else {
self.ff_num += 1;
}
self.low = (self.low << 8) & 0xffff_ffff;
}
pub(crate) fn encode(&mut self, low: u32, freq: u32, total: u32) {
self.range /= total;
self.low += u64::from(low) * u64::from(self.range);
self.range *= freq;
while self.range < (1 << 24) {
self.range <<= 8;
self.shift_low();
}
}
pub(crate) fn finish(mut self) -> Vec<u8> {
for _ in 0..5 {
self.shift_low();
}
self.out
}
}
pub(crate) struct ModelEncoder {
symbols: Vec<u8>,
freq: Vec<u32>,
total: u32,
}
impl ModelEncoder {
pub(crate) fn new(n_symbols: usize) -> Self {
Self {
symbols: (0..n_symbols).map(|s| s as u8).collect(),
freq: vec![1; n_symbols],
total: n_symbols as u32,
}
}
pub(crate) fn encode(&mut self, rc: &mut RangeEncoder, symbol: u8) {
let mut acc = 0u32;
let mut x = 0usize;
while self.symbols[x] != symbol {
acc += self.freq[x];
x += 1;
}
rc.encode(acc, self.freq[x], self.total);
self.freq[x] += STEP;
self.total += STEP;
if self.total > MAX_TOTAL {
self.total = 0;
for f in self.freq.iter_mut() {
*f -= *f / 2;
self.total += *f;
}
}
if x > 0 && self.freq[x] > self.freq[x - 1] {
self.freq.swap(x, x - 1);
self.symbols.swap(x, x - 1);
}
}
}
}
#[cfg(test)]
mod tests {
use super::testing::{ModelEncoder, RangeEncoder};
use super::*;
fn alphabet(data: &[u8]) -> usize {
data.iter().map(|b| *b as usize + 1).max().unwrap_or(1)
}
fn encode_order_0(data: &[u8]) -> Vec<u8> {
let n = alphabet(data);
let mut model = ModelEncoder::new(n);
let mut rc = RangeEncoder::new();
for byte in data {
model.encode(&mut rc, *byte);
}
let mut out = vec![(n & 0xff) as u8];
out.extend_from_slice(&rc.finish());
out
}
fn encode_order_1(data: &[u8]) -> Vec<u8> {
let n = alphabet(data);
let mut models: Vec<ModelEncoder> = (0..n).map(|_| ModelEncoder::new(n)).collect();
let mut rc = RangeEncoder::new();
let mut last = 0usize;
for byte in data {
models[last].encode(&mut rc, *byte);
last = *byte as usize;
}
let mut out = vec![(n & 0xff) as u8];
out.extend_from_slice(&rc.finish());
out
}
fn runs_of(data: &[u8]) -> Vec<(u8, usize)> {
let mut out = Vec::new();
let mut i = 0;
while i < data.len() {
let symbol = data[i];
let mut run = 0;
while i + run + 1 < data.len() && data[i + run + 1] == symbol {
run += 1;
}
out.push((symbol, run));
i += run + 1;
}
out
}
fn encode_run(rc: &mut RangeEncoder, runs: &mut [ModelEncoder], first: usize, mut run: usize) {
let mut context = first;
loop {
let part = run.min(3);
runs[context].encode(rc, part as u8);
run -= part;
if part < 3 {
return;
}
context = if context == first { 256 } else { 257 };
}
}
fn encode_rle(data: &[u8], order_1: bool) -> Vec<u8> {
let n = alphabet(data);
let mut literals: Vec<ModelEncoder> = if order_1 {
(0..n).map(|_| ModelEncoder::new(n)).collect()
} else {
vec![ModelEncoder::new(n)]
};
let mut runs: Vec<ModelEncoder> = (0..258).map(|_| ModelEncoder::new(4)).collect();
let mut rc = RangeEncoder::new();
let mut last = 0usize;
for (symbol, run) in runs_of(data) {
let index = if order_1 { last } else { 0 };
literals[index].encode(&mut rc, symbol);
last = symbol as usize;
encode_run(&mut rc, &mut runs, symbol as usize, run);
}
let mut out = vec![(n & 0xff) as u8];
out.extend_from_slice(&rc.finish());
out
}
fn wrap(flags: u8, len: usize, body: &[u8]) -> Vec<u8> {
let mut out = vec![flags];
let mut value = len as u32;
let mut seven = Vec::new();
loop {
seven.push((value & 0x7f) as u8);
value >>= 7;
if value == 0 {
break;
}
}
for (i, byte) in seven.iter().enumerate().rev() {
out.push(if i == 0 { *byte } else { byte | 0x80 });
}
out.extend_from_slice(body);
out
}
fn roundtrip(data: &[u8], flags: u8) {
let body = match (flags & flag::RLE != 0, flags & flag::ORDER != 0) {
(false, false) => encode_order_0(data),
(false, true) => encode_order_1(data),
(true, order_1) => encode_rle(data, order_1),
};
let stream = wrap(flags, data.len(), &body);
let decoded = decode(&stream, "test", 0)
.unwrap_or_else(|e| panic!("flags {flags}, {} bytes: {e}", data.len()));
assert_eq!(decoded, data, "flags {flags}, {} bytes", data.len());
}
#[test]
fn order_0_round_trips_text() {
roundtrip(&b"abracadabra".repeat(30), 0);
}
#[test]
fn order_1_round_trips_text() {
roundtrip(&b"abracadabra".repeat(30), flag::ORDER);
}
#[test]
fn run_length_round_trips_at_both_orders() {
let data = b"ABBCCCCDDDDD".repeat(10);
roundtrip(&data, flag::RLE);
roundtrip(&data, flag::RLE | flag::ORDER);
}
#[test]
fn the_specifications_run_length_example_decodes_to_twelve_bytes() {
let data = b"ABBCCCCDDDDD";
assert_eq!(
runs_of(data),
vec![(b'A', 0), (b'B', 1), (b'C', 3), (b'D', 4)]
);
roundtrip(data, flag::RLE);
}
#[test]
fn a_run_longer_than_three_is_split_into_continuations() {
let data = vec![b'Z'; 300];
roundtrip(&data, flag::RLE);
roundtrip(&data, flag::RLE | flag::ORDER);
}
#[test]
fn the_whole_byte_alphabet_round_trips() {
let data: Vec<u8> = (0..=255u8).chain((0..=255u8).rev()).collect();
roundtrip(&data, 0);
roundtrip(&data, flag::ORDER);
}
#[test]
fn a_stream_long_enough_to_renormalise_its_model_round_trips() {
let data: Vec<u8> = (0..20_000).map(|i| b"acgt"[i % 4]).collect();
roundtrip(&data, 0);
roundtrip(&data, flag::ORDER);
}
#[test]
fn a_single_byte_and_an_empty_stream_round_trip() {
roundtrip(b"", 0);
roundtrip(b"Q", 0);
roundtrip(b"Q", flag::ORDER);
roundtrip(b"Q", flag::RLE);
}
#[test]
fn the_cat_flag_gives_the_bytes_back_verbatim() {
let stream = wrap(flag::CAT, 5, b"hello");
assert_eq!(decode(&stream, "test", 0).expect("cat"), b"hello");
}
#[test]
fn a_stripe_interleaves_its_sub_streams() {
let subs: Vec<Vec<u8>> = (0..4)
.map(|j| {
let body = encode_order_0(&[j * 10, j * 10 + 1]);
let mut sub = vec![flag::NO_SIZE];
sub.extend_from_slice(&body);
sub
})
.collect();
let mut stream = wrap(flag::STRIPE, 8, &[]);
stream.push(4);
for sub in &subs {
stream.push(sub.len() as u8);
}
for sub in &subs {
stream.extend_from_slice(sub);
}
assert_eq!(
decode(&stream, "test", 0).expect("stripe"),
vec![0, 10, 20, 30, 1, 11, 21, 31]
);
}
#[test]
fn a_model_selecting_a_frequency_no_symbol_covers_is_refused() {
let mut stream = wrap(flag::RLE, 64, &[4u8]);
stream.extend_from_slice(&[0xff; 32]);
let error = decode(&stream, "test", 0).expect_err("no symbol covers it");
assert!(
error.to_string().contains("which no symbol covers"),
"{error}"
);
}
#[test]
fn every_prefix_of_a_real_stream_fails_without_panicking() {
let stream = wrap(
flag::ORDER,
330,
&encode_order_1(&b"abracadabra".repeat(30)),
);
for cut in 0..stream.len() {
let _ = decode(&stream[..cut], "test", 0);
}
for byte in 0..=255u8 {
let mut damaged = stream.clone();
damaged[0] = byte;
let _ = decode(&damaged, "test", 0);
}
}
}