use crate::error::{Jpeg2000Error, Result};
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct CodeBlockInclusion {
pub included: bool,
pub new_passes: u8,
pub data_length: u32,
}
#[derive(Debug, Clone)]
pub struct PacketHeader {
pub is_empty: bool,
pub inclusions: Vec<CodeBlockInclusion>,
}
#[derive(Debug, Clone)]
pub struct Packet {
pub header: PacketHeader,
pub code_block_data: Vec<Vec<u8>>,
}
pub struct BitReader<'a> {
data: &'a [u8],
byte_pos: usize,
current_byte: u8,
bits_left: u8,
prev_was_ff: bool,
}
impl<'a> BitReader<'a> {
pub fn new(data: &'a [u8]) -> Self {
Self {
data,
byte_pos: 0,
current_byte: 0,
bits_left: 0,
prev_was_ff: false,
}
}
pub fn read_bit(&mut self) -> Result<u8> {
if self.bits_left == 0 {
self.refill()?;
}
self.bits_left -= 1;
let bit = (self.current_byte >> self.bits_left) & 1;
Ok(bit)
}
pub fn read_bits(&mut self, n: u8) -> Result<u32> {
let mut value = 0u32;
for _ in 0..n {
value = (value << 1) | u32::from(self.read_bit()?);
}
Ok(value)
}
pub fn bytes_consumed(&self) -> usize {
self.byte_pos
}
fn refill(&mut self) -> Result<()> {
if self.byte_pos >= self.data.len() {
return Err(Jpeg2000Error::InsufficientData {
expected: 1,
actual: 0,
});
}
let byte = self.data[self.byte_pos];
let skip_msb = self.prev_was_ff;
self.prev_was_ff = byte == 0xFF;
self.byte_pos += 1;
self.current_byte = byte;
self.bits_left = if skip_msb { 7 } else { 8 };
Ok(())
}
}
#[derive(Debug, Clone)]
pub struct TagTree {
level_dims: Vec<(u32, u32)>,
low: Vec<Vec<u32>>,
value: Vec<Vec<u32>>,
num_levels: usize,
}
impl TagTree {
pub fn new(width: u32, height: u32) -> Self {
let mut level_dims = Vec::new();
let mut low = Vec::new();
let mut value = Vec::new();
let mut w = width.max(1);
let mut h = height.max(1);
loop {
let count = (w as usize) * (h as usize);
level_dims.push((w, h));
low.push(vec![0u32; count]);
value.push(vec![u32::MAX; count]);
if w == 1 && h == 1 {
break;
}
w = w.div_ceil(2);
h = h.div_ceil(2);
}
let num_levels = level_dims.len();
Self {
level_dims,
low,
value,
num_levels,
}
}
fn decode_lt(
&mut self,
cx: u32,
cy: u32,
threshold: u32,
bits: &mut BitReader<'_>,
) -> Result<bool> {
let mut low = 0u32;
for k in (0..self.num_levels).rev() {
let (wk, _hk) = self.level_dims[k];
let nx = cx >> k;
let ny = cy >> k;
let idx = (ny as usize) * (wk as usize) + (nx as usize);
if low > self.low[k][idx] {
self.low[k][idx] = low;
} else {
low = self.low[k][idx];
}
while low < threshold && low < self.value[k][idx] {
if bits.read_bit()? == 1 {
self.value[k][idx] = low; } else {
low += 1;
}
}
self.low[k][idx] = low;
}
Ok(low < threshold)
}
pub fn decode_value(
&mut self,
cx: u32,
cy: u32,
threshold: u32,
bits: &mut BitReader<'_>,
) -> Result<bool> {
self.decode_lt(cx, cy, threshold.saturating_add(1), bits)
}
pub fn decode_full(&mut self, cx: u32, cy: u32, bits: &mut BitReader<'_>) -> Result<u32> {
let mut t = 0u32;
loop {
if self.decode_lt(cx, cy, t + 1, bits)? {
return Ok(t);
}
t += 1;
if t > (1 << 20) {
return Err(Jpeg2000Error::Tier2Error(
"tag-tree value did not converge".to_string(),
));
}
}
}
pub fn num_levels(&self) -> usize {
self.num_levels
}
}
pub struct PacketDecoder;
impl PacketDecoder {
pub fn decode(
data: &[u8],
num_code_blocks_x: u32,
num_code_blocks_y: u32,
previously_included: &mut Vec<bool>,
) -> Result<(Packet, usize)> {
let num_blocks = (num_code_blocks_x * num_code_blocks_y) as usize;
if previously_included.len() < num_blocks {
previously_included.resize(num_blocks, false);
}
let mut bits = BitReader::new(data);
let ppkt = bits.read_bit()?;
if ppkt == 0 {
let consumed = bits.bytes_consumed();
return Ok((
Packet {
header: PacketHeader {
is_empty: true,
inclusions: vec![
CodeBlockInclusion {
included: false,
new_passes: 0,
data_length: 0,
};
num_blocks
],
},
code_block_data: Vec::new(),
},
consumed,
));
}
let mut inclusion_tree = TagTree::new(num_code_blocks_x, num_code_blocks_y);
let mut zbp_tree = TagTree::new(num_code_blocks_x, num_code_blocks_y);
let mut inclusions = Vec::with_capacity(num_blocks);
let mut included_indices = Vec::new();
for by in 0..num_code_blocks_y {
for bx in 0..num_code_blocks_x {
let idx = (by * num_code_blocks_x + bx) as usize;
let already = previously_included[idx];
let included = if already {
bits.read_bit()? == 1
} else {
inclusion_tree.decode_value(bx, by, 0, &mut bits)?
};
if included {
previously_included[idx] = true;
let new_passes = Self::decode_num_passes(&mut bits)?;
let data_length = Self::decode_block_length(&mut bits, new_passes)?;
let _zbp = if !already {
zbp_tree
.decode_value(bx, by, 255, &mut bits)
.unwrap_or(false)
} else {
false
};
inclusions.push(CodeBlockInclusion {
included: true,
new_passes,
data_length,
});
included_indices.push(idx);
} else {
inclusions.push(CodeBlockInclusion {
included: false,
new_passes: 0,
data_length: 0,
});
}
}
}
let header_bytes = bits.bytes_consumed();
let mut pos = header_bytes;
let mut code_block_data = Vec::new();
for incl in &inclusions {
if incl.included {
let len = incl.data_length as usize;
if pos + len > data.len() {
return Err(Jpeg2000Error::InsufficientData {
expected: pos + len,
actual: data.len(),
});
}
code_block_data.push(data[pos..pos + len].to_vec());
pos += len;
}
}
Ok((
Packet {
header: PacketHeader {
is_empty: false,
inclusions,
},
code_block_data,
},
pos,
))
}
fn decode_num_passes(bits: &mut BitReader<'_>) -> Result<u8> {
if bits.read_bit()? == 1 {
return Ok(1);
}
if bits.read_bit()? == 1 {
return Ok(2);
}
if bits.read_bit()? == 1 {
let extra = bits.read_bits(2)? as u8;
return Ok(3 + extra);
}
if bits.read_bit()? == 1 {
let extra = bits.read_bits(4)? as u8;
return Ok(7 + extra);
}
let extra = bits.read_bits(6)? as u8;
Ok(23 + extra)
}
fn decode_block_length(bits: &mut BitReader<'_>, _new_passes: u8) -> Result<u32> {
let mut extra_bits = 0u32;
loop {
let b = bits.read_bit()?;
if b == 0 {
break;
}
extra_bits += 1;
}
let nbits = 3 + extra_bits;
if nbits > 31 {
return Err(Jpeg2000Error::Tier2Error(
"Block length exceeds 31 bits".to_string(),
));
}
let length = bits.read_bits(nbits as u8)?;
Ok(length)
}
}
#[derive(Debug, Clone, Default, PartialEq, Eq)]
pub struct CblkContribution {
pub included: bool,
pub num_passes: u32,
pub zbp: u32,
pub data_offset: usize,
pub data_len: usize,
}
#[inline]
pub fn floor_log2(x: u32) -> u32 {
if x == 0 { 0 } else { 31 - x.leading_zeros() }
}
pub fn read_num_coding_passes(bits: &mut BitReader<'_>) -> Result<u32> {
if bits.read_bit()? == 0 {
return Ok(1);
}
if bits.read_bit()? == 0 {
return Ok(2);
}
let n = bits.read_bits(2)?;
if n != 3 {
return Ok(3 + n);
}
let n = bits.read_bits(5)?;
if n != 31 {
return Ok(6 + n);
}
let n = bits.read_bits(7)?;
Ok(37 + n)
}
pub fn read_length_increment(bits: &mut BitReader<'_>) -> Result<u32> {
let mut increment = 0u32;
while bits.read_bit()? == 1 {
increment += 1;
if increment > 32 {
return Err(Jpeg2000Error::Tier2Error(
"Lblock length increment out of range".to_string(),
));
}
}
Ok(increment)
}
fn maybe_skip_eph(data: &[u8], pos: usize, has_eph: bool) -> usize {
if has_eph && pos + 2 <= data.len() && data[pos] == 0xFF && data[pos + 1] == 0x92 {
pos + 2
} else {
pos
}
}
pub fn parse_precinct_packet(
data: &[u8],
subband_grids: &[(u32, u32)],
layer: u16,
has_eph: bool,
) -> Result<(usize, Vec<Vec<CblkContribution>>)> {
let mut bits = BitReader::new(data);
let present = bits.read_bit()? == 1;
if !present {
let consumed = maybe_skip_eph(data, bits.bytes_consumed(), has_eph);
let empty = subband_grids
.iter()
.map(|&(nx, ny)| vec![CblkContribution::default(); (nx * ny) as usize])
.collect();
return Ok((consumed, empty));
}
let layer_threshold = u32::from(layer);
let mut per_subband: Vec<Vec<CblkContribution>> = Vec::with_capacity(subband_grids.len());
for &(nx, ny) in subband_grids {
let count = (nx * ny) as usize;
let mut inclusion_tree = TagTree::new(nx, ny);
let mut zbp_tree = TagTree::new(nx, ny);
let mut contributions = Vec::with_capacity(count);
for by in 0..ny {
for bx in 0..nx {
let included = inclusion_tree.decode_value(bx, by, layer_threshold, &mut bits)?;
if !included {
contributions.push(CblkContribution::default());
continue;
}
let zbp = zbp_tree.decode_full(bx, by, &mut bits)?;
let num_passes = read_num_coding_passes(&mut bits)?;
let lblock = 3 + read_length_increment(&mut bits)?;
let nbits = lblock + floor_log2(num_passes);
if nbits > 31 {
return Err(Jpeg2000Error::Tier2Error(
"code-block length field exceeds 31 bits".to_string(),
));
}
let length = bits.read_bits(nbits as u8)?;
contributions.push(CblkContribution {
included: true,
num_passes,
zbp,
data_offset: 0,
data_len: length as usize,
});
}
}
per_subband.push(contributions);
}
let mut pos = maybe_skip_eph(data, bits.bytes_consumed(), has_eph);
for subband in &mut per_subband {
for contribution in subband.iter_mut() {
if !contribution.included {
continue;
}
let end = pos.checked_add(contribution.data_len).ok_or_else(|| {
Jpeg2000Error::Tier2Error("code-block length overflow".to_string())
})?;
if end > data.len() {
return Err(Jpeg2000Error::InsufficientData {
expected: end,
actual: data.len(),
});
}
contribution.data_offset = pos;
pos = end;
}
}
Ok((pos, per_subband))
}
#[cfg(test)]
mod tests {
#![allow(clippy::unwrap_used, clippy::expect_used, clippy::panic)]
use super::*;
#[derive(Default)]
struct BitWriter {
bytes: Vec<u8>,
cur: u8,
nbits: u8,
}
impl BitWriter {
fn write_bit(&mut self, bit: u8) {
self.cur = (self.cur << 1) | (bit & 1);
self.nbits += 1;
if self.nbits == 8 {
self.bytes.push(self.cur);
self.cur = 0;
self.nbits = 0;
}
}
fn write_bits(&mut self, value: u32, n: u8) {
for i in (0..n).rev() {
self.write_bit(((value >> i) & 1) as u8);
}
}
fn align(&mut self) {
while self.nbits != 0 {
self.write_bit(0);
}
}
fn into_bytes(mut self) -> Vec<u8> {
self.align();
self.bytes
}
}
fn write_num_passes(w: &mut BitWriter, passes: u32) {
match passes {
1 => w.write_bit(0),
2 => {
w.write_bit(1);
w.write_bit(0);
}
3..=5 => {
w.write_bit(1);
w.write_bit(1);
w.write_bits(passes - 3, 2);
}
6..=36 => {
w.write_bit(1);
w.write_bit(1);
w.write_bits(3, 2);
w.write_bits(passes - 6, 5);
}
_ => panic!("test helper only supports <=36 passes"),
}
}
fn write_single_cblk(w: &mut BitWriter, zbp: u32, passes: u32, length: u32) {
w.write_bit(1);
for _ in 0..zbp {
w.write_bit(0);
}
w.write_bit(1);
write_num_passes(w, passes);
let need = if length == 0 {
1
} else {
32 - length.leading_zeros()
};
let base = 3 + floor_log2(passes);
let inc = need.saturating_sub(base);
for _ in 0..inc {
w.write_bit(1);
}
w.write_bit(0); let nbits = (base + inc) as u8;
w.write_bits(length, nbits);
}
#[test]
fn test_floor_log2() {
assert_eq!(floor_log2(0), 0);
assert_eq!(floor_log2(1), 0);
assert_eq!(floor_log2(2), 1);
assert_eq!(floor_log2(3), 1);
assert_eq!(floor_log2(4), 2);
assert_eq!(floor_log2(255), 7);
}
#[test]
fn test_num_passes_round_trip() {
for passes in [1u32, 2, 3, 4, 5, 6, 10, 20, 36] {
let mut w = BitWriter::default();
write_num_passes(&mut w, passes);
let bytes = w.into_bytes();
let mut br = BitReader::new(&bytes);
assert_eq!(
read_num_coding_passes(&mut br).expect("decode passes"),
passes
);
}
}
#[test]
fn test_tag_tree_full_value_1x1() {
let data = [0b0001_0000u8];
let mut br = BitReader::new(&data);
let mut tt = TagTree::new(1, 1);
assert_eq!(tt.decode_full(0, 0, &mut br).expect("full"), 3);
}
#[test]
fn test_tag_tree_multi_leaf_inclusion() {
let mut w = BitWriter::default();
w.write_bit(1); w.write_bit(1); w.write_bit(1); let bytes = w.into_bytes();
let mut br = BitReader::new(&bytes);
let mut tt = TagTree::new(2, 1);
assert!(tt.decode_value(0, 0, 0, &mut br).expect("leaf0"));
assert!(tt.decode_value(1, 0, 0, &mut br).expect("leaf1"));
}
#[test]
fn test_parse_empty_packet() {
let data = [0x00u8, 0xAA];
let (consumed, subs) =
parse_precinct_packet(&data, &[(1, 1)], 0, false).expect("empty packet");
assert_eq!(consumed, 1);
assert_eq!(subs.len(), 1);
assert!(!subs[0][0].included);
}
#[test]
fn test_parse_single_block_packet() {
let mut w = BitWriter::default();
w.write_bit(1); write_single_cblk(&mut w, 2, 1, 5);
let mut bytes = w.into_bytes();
let header_len = bytes.len();
let body = [0x11u8, 0x22, 0x33, 0x44, 0x55];
bytes.extend_from_slice(&body);
let (consumed, subs) =
parse_precinct_packet(&bytes, &[(1, 1)], 0, false).expect("single block");
assert_eq!(subs.len(), 1);
let c = &subs[0][0];
assert!(c.included);
assert_eq!(c.zbp, 2);
assert_eq!(c.num_passes, 1);
assert_eq!(c.data_len, 5);
assert_eq!(c.data_offset, header_len);
assert_eq!(&bytes[c.data_offset..c.data_offset + c.data_len], &body);
assert_eq!(consumed, header_len + 5);
}
#[test]
fn test_parse_multi_subband_packet() {
let mut w = BitWriter::default();
w.write_bit(1); write_single_cblk(&mut w, 0, 1, 2); write_single_cblk(&mut w, 1, 1, 3); write_single_cblk(&mut w, 0, 2, 4); let mut bytes = w.into_bytes();
let header_len = bytes.len();
let body: Vec<u8> = (0..(2 + 3 + 4)).map(|i| i as u8 + 1).collect();
bytes.extend_from_slice(&body);
let grids = [(1u32, 1u32), (1, 1), (1, 1)];
let (consumed, subs) =
parse_precinct_packet(&bytes, &grids, 0, false).expect("multi subband");
assert_eq!(subs.len(), 3);
assert_eq!(subs[0][0].data_len, 2);
assert_eq!(subs[1][0].data_len, 3);
assert_eq!(subs[2][0].data_len, 4);
assert_eq!(subs[0][0].data_offset, header_len);
assert_eq!(subs[1][0].data_offset, header_len + 2);
assert_eq!(subs[2][0].data_offset, header_len + 5);
assert_eq!(consumed, header_len + 9);
}
#[test]
fn test_parse_length_overrun_errors() {
let mut w = BitWriter::default();
w.write_bit(1);
write_single_cblk(&mut w, 0, 1, 200);
let bytes = w.into_bytes();
let err = parse_precinct_packet(&bytes, &[(1, 1)], 0, false);
assert!(matches!(err, Err(Jpeg2000Error::InsufficientData { .. })));
}
#[test]
fn test_bit_reader_basic() {
let data = [0b1011_0100u8];
let mut br = BitReader::new(&data);
assert_eq!(br.read_bit().expect("read bit 0"), 1);
assert_eq!(br.read_bit().expect("read bit 1"), 0);
assert_eq!(br.read_bit().expect("read bit 2"), 1);
assert_eq!(br.read_bit().expect("read bit 3"), 1);
assert_eq!(br.read_bit().expect("read bit 4"), 0);
assert_eq!(br.read_bit().expect("read bit 5"), 1);
assert_eq!(br.read_bit().expect("read bit 6"), 0);
assert_eq!(br.read_bit().expect("read bit 7"), 0);
}
#[test]
fn test_bit_reader_read_bits() {
let data = [0b1010_1010u8, 0b1111_0000u8];
let mut br = BitReader::new(&data);
assert_eq!(br.read_bits(4).expect("read nibble 0"), 0b1010);
assert_eq!(br.read_bits(4).expect("read nibble 1"), 0b1010);
assert_eq!(br.read_bits(4).expect("read nibble 2"), 0b1111);
assert_eq!(br.read_bits(4).expect("read nibble 3"), 0b0000);
}
#[test]
fn test_bit_reader_exhaustion() {
let data = [0x00u8];
let mut br = BitReader::new(&data);
for _ in 0..8 {
br.read_bit().expect("read bit before exhaustion");
}
assert!(br.read_bit().is_err());
}
#[test]
fn test_tag_tree_new() {
let tt = TagTree::new(4, 3);
assert_eq!(tt.num_levels(), 3);
}
#[test]
fn test_tag_tree_1x1() {
let tt = TagTree::new(1, 1);
assert_eq!(tt.num_levels(), 1);
}
#[test]
fn test_empty_packet_decode() {
let data = [0b0000_0000u8];
let mut prev = vec![];
let (pkt, consumed) =
PacketDecoder::decode(&data, 2, 2, &mut prev).expect("decode empty packet");
assert!(pkt.header.is_empty);
assert!(pkt.code_block_data.is_empty());
assert_eq!(consumed, 1);
}
#[test]
fn test_packet_header_is_empty_flag() {
let data = [0x00u8];
let mut prev = vec![];
let (pkt, _) = PacketDecoder::decode(&data, 1, 1, &mut prev).expect("decode packet header");
assert!(pkt.header.is_empty);
assert_eq!(pkt.header.inclusions.len(), 1);
assert!(!pkt.header.inclusions[0].included);
}
#[test]
fn test_code_block_inclusion_default() {
let incl = CodeBlockInclusion {
included: true,
new_passes: 3,
data_length: 128,
};
assert!(incl.included);
assert_eq!(incl.new_passes, 3);
assert_eq!(incl.data_length, 128);
}
#[test]
fn test_packet_struct_fields() {
let pkt = Packet {
header: PacketHeader {
is_empty: false,
inclusions: vec![CodeBlockInclusion {
included: true,
new_passes: 1,
data_length: 4,
}],
},
code_block_data: vec![vec![1, 2, 3, 4]],
};
assert!(!pkt.header.is_empty);
assert_eq!(pkt.code_block_data.len(), 1);
assert_eq!(pkt.code_block_data[0].len(), 4);
}
}