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 {
width: u32,
height: u32,
levels: Vec<Vec<Option<u32>>>,
num_levels: usize,
}
impl TagTree {
pub fn new(width: u32, height: u32) -> Self {
let mut levels = Vec::new();
let mut w = width;
let mut h = height;
loop {
let count = (w as usize) * (h as usize);
levels.push(vec![None; count]);
if w == 1 && h == 1 {
break;
}
w = w.div_ceil(2);
h = h.div_ceil(2);
}
let num_levels = levels.len();
Self {
width,
height,
levels,
num_levels,
}
}
pub fn decode_value(
&mut self,
cx: u32,
cy: u32,
threshold: u32,
bits: &mut BitReader<'_>,
) -> Result<bool> {
let mut path: Vec<(usize, usize)> = Vec::with_capacity(self.num_levels);
let mut lx = cx;
let mut ly = cy;
for lvl in 0..self.num_levels {
let level_idx = self.num_levels - 1 - lvl; let w = self.level_width(level_idx);
let idx = (ly as usize) * w + (lx as usize);
path.push((level_idx, idx));
lx /= 2;
ly /= 2;
}
let mut parent_lower = 0u32;
for &(level_idx, node_idx) in path.iter().rev() {
let current = self.levels[level_idx][node_idx].unwrap_or(parent_lower);
let mut lower = current.max(parent_lower);
loop {
if lower > threshold {
self.levels[level_idx][node_idx] = Some(lower);
return Ok(false);
}
let bit = bits.read_bit()?;
if bit == 1 {
self.levels[level_idx][node_idx] = Some(lower);
break; }
lower += 1;
}
parent_lower = lower;
}
Ok(true)
}
fn level_width(&self, level_idx: usize) -> usize {
let mut lw = self.width as usize;
let mut lh = self.height as usize;
let target = self.num_levels - 1 - level_idx; for _ in 0..target {
lw = lw.div_ceil(2);
lh = lh.div_ceil(2);
}
let _ = lh;
lw
}
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)
}
}
#[cfg(test)]
mod tests {
use super::*;
#[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);
}
}