use crate::base::{Error, Result};
pub struct BitReader<'a> {
data: &'a [u8],
pos: usize,
bit_buf: u32,
bit_count: u32,
}
impl<'a> BitReader<'a> {
pub fn new(data: &'a [u8]) -> BitReader<'a> {
BitReader {
data,
pos: 0,
bit_buf: 0,
bit_count: 0,
}
}
pub fn byte_pos(&self) -> usize {
self.pos
}
fn load_byte(&mut self) -> Result<()> {
match self.data.get(self.pos) {
None => Err(Error::Parse("jpeg: truncated entropy data".into())),
Some(&0xFF) => match self.data.get(self.pos + 1) {
Some(&0x00) => {
self.pos += 2;
self.bit_buf = (self.bit_buf << 8) | 0xFF;
self.bit_count += 8;
Ok(())
}
_ => Err(Error::Parse(
"jpeg: entropy data ended at a marker mid-block".into(),
)),
},
Some(&b) => {
self.pos += 1;
self.bit_buf = (self.bit_buf << 8) | b as u32;
self.bit_count += 8;
Ok(())
}
}
}
#[inline]
pub fn next_bit(&mut self) -> Result<u32> {
if self.bit_count == 0 {
self.load_byte()?;
}
self.bit_count -= 1;
Ok((self.bit_buf >> self.bit_count) & 1)
}
pub fn receive(&mut self, n: u32) -> Result<u32> {
debug_assert!(n <= 16);
let mut v = 0u32;
for _ in 0..n {
v = (v << 1) | self.next_bit()?;
}
Ok(v)
}
pub fn expect_restart(&mut self, n: u8) -> Result<()> {
self.bit_buf = 0;
self.bit_count = 0;
let want = 0xD0 + (n & 7);
match (self.data.get(self.pos), self.data.get(self.pos + 1)) {
(Some(&0xFF), Some(&m)) if m == want => {
self.pos += 2;
Ok(())
}
(Some(&0xFF), Some(&m)) => Err(Error::Parse(format!(
"jpeg: expected restart marker RST{} but found 0xFF{m:02X}",
n & 7
))),
_ => Err(Error::Parse("jpeg: missing restart marker".into())),
}
}
}
pub struct HuffTable {
min_code: [i32; 17],
max_code: [i32; 17],
val_ptr: [usize; 17],
values: Vec<u8>,
}
impl HuffTable {
pub fn build(counts: &[u8; 16], values: &[u8]) -> Result<HuffTable> {
let total: usize = counts.iter().map(|&c| c as usize).sum();
if total != values.len() || total > 256 {
return Err(Error::Parse(format!(
"jpeg: DHT declares {total} symbols, carries {}",
values.len()
)));
}
let mut min_code = [0i32; 17];
let mut max_code = [-1i32; 17];
let mut val_ptr = [0usize; 17];
let mut code = 0i32;
let mut k = 0usize;
for len in 1..=16usize {
let n = counts[len - 1] as i32;
if n > 0 {
val_ptr[len] = k;
min_code[len] = code;
code += n;
max_code[len] = code - 1;
k += n as usize;
if max_code[len] >= (1 << len) {
return Err(Error::Parse(
"jpeg: DHT codes overflow their bit length".into(),
));
}
}
code <<= 1;
}
Ok(HuffTable {
min_code,
max_code,
val_ptr,
values: values.to_vec(),
})
}
pub fn decode(&self, r: &mut BitReader<'_>) -> Result<u8> {
let mut code = 0i32;
for len in 1..=16usize {
code = (code << 1) | r.next_bit()? as i32;
if self.max_code[len] >= 0 && code <= self.max_code[len] {
let idx = self.val_ptr[len] + (code - self.min_code[len]) as usize;
return self
.values
.get(idx)
.copied()
.ok_or_else(|| Error::Parse("jpeg: Huffman value index out of range".into()));
}
}
Err(Error::Parse("jpeg: invalid Huffman code (>16 bits)".into()))
}
}
#[inline]
pub fn extend(v: u32, size: u32) -> i32 {
if size == 0 {
return 0;
}
if v < (1 << (size - 1)) {
v as i32 - (1 << size) + 1
} else {
v as i32
}
}
#[inline]
fn coef(v: i32) -> i16 {
v.clamp(i16::MIN as i32, i16::MAX as i32) as i16
}
pub fn decode_block(
r: &mut BitReader<'_>,
dc: &HuffTable,
ac: &HuffTable,
dc_pred: &mut i32,
out: &mut [i16],
) -> Result<()> {
debug_assert_eq!(out.len(), 64);
let s = dc.decode(r)? as u32;
if s > 11 {
return Err(Error::Parse(format!("jpeg: DC size class {s} > 11")));
}
let diff = extend(r.receive(s)?, s);
*dc_pred += diff;
out[0] = coef(*dc_pred);
let mut k = 1usize;
while k < 64 {
let rs = ac.decode(r)? as u32;
let run = rs >> 4;
let size = rs & 0x0F;
if size == 0 {
if run == 15 {
k += 16; continue;
}
break; }
k += run as usize;
if k > 63 {
return Err(Error::Parse("jpeg: AC run past end of block".into()));
}
if size > 10 {
return Err(Error::Parse(format!("jpeg: AC size class {size} > 10")));
}
out[k] = coef(extend(r.receive(size)?, size));
k += 1;
}
Ok(())
}
pub fn decode_dc_first(
r: &mut BitReader<'_>,
dc: &HuffTable,
dc_pred: &mut i32,
al: u32,
out: &mut [i16],
) -> Result<()> {
debug_assert_eq!(out.len(), 64);
let s = dc.decode(r)? as u32;
if s > 11 {
return Err(Error::Parse(format!("jpeg: DC size class {s} > 11")));
}
let diff = extend(r.receive(s)?, s);
*dc_pred += diff;
out[0] = coef((((*dc_pred) as i64) << al).clamp(i32::MIN as i64, i32::MAX as i64) as i32);
Ok(())
}
pub fn decode_dc_refine(r: &mut BitReader<'_>, al: u32, out: &mut [i16]) -> Result<()> {
debug_assert_eq!(out.len(), 64);
if r.next_bit()? != 0 {
out[0] |= 1i16 << al;
}
Ok(())
}
pub fn decode_ac_first(
r: &mut BitReader<'_>,
ac: &HuffTable,
ss: usize,
se: usize,
al: u32,
eobrun: &mut u32,
out: &mut [i16],
) -> Result<()> {
debug_assert_eq!(out.len(), 64);
if *eobrun > 0 {
*eobrun -= 1;
return Ok(());
}
let mut k = ss;
while k <= se {
let rs = ac.decode(r)? as u32;
let run = (rs >> 4) as usize;
let size = rs & 0x0F;
if size == 0 {
if run != 15 {
*eobrun = (1u32 << run) - 1;
if run > 0 {
*eobrun += r.receive(run as u32)?;
}
break;
}
k += 16; continue;
}
k += run;
if k > se {
return Err(Error::Parse(
"jpeg: AC run past the end of the spectral band".into(),
));
}
if size > 10 {
return Err(Error::Parse(format!("jpeg: AC size class {size} > 10")));
}
out[k] = coef(extend(r.receive(size)?, size) << al);
k += 1;
}
Ok(())
}
pub fn decode_ac_refine(
r: &mut BitReader<'_>,
ac: &HuffTable,
ss: usize,
se: usize,
al: u32,
eobrun: &mut u32,
out: &mut [i16],
) -> Result<()> {
debug_assert_eq!(out.len(), 64);
let p1 = 1i16 << al; let m1 = -1i16 << al; let mut k = ss;
if *eobrun == 0 {
while k <= se {
let rs = ac.decode(r)? as u32;
let mut run = (rs >> 4) as i32;
let size = rs & 0x0F;
let mut newval = 0i16;
if size != 0 {
if size != 1 {
return Err(Error::Parse(
"jpeg: AC refinement size class must be 1".into(),
));
}
newval = if r.next_bit()? != 0 { p1 } else { m1 };
} else if run != 15 {
*eobrun = 1u32 << run;
if run > 0 {
*eobrun += r.receive(run as u32)?;
}
break;
}
loop {
if out[k] != 0 {
if r.next_bit()? != 0 && (out[k] & p1) == 0 {
out[k] = out[k].saturating_add(if out[k] >= 0 { p1 } else { m1 });
}
} else {
run -= 1;
if run < 0 {
break;
}
}
k += 1;
if k > se {
break;
}
}
if newval != 0 && k <= se {
out[k] = newval;
}
k += 1;
}
}
if *eobrun > 0 {
while k <= se {
if out[k] != 0 && r.next_bit()? != 0 && (out[k] & p1) == 0 {
out[k] = out[k].saturating_add(if out[k] >= 0 { p1 } else { m1 });
}
k += 1;
}
*eobrun -= 1;
}
Ok(())
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn bit_reader_stuffing_and_markers() {
let data = [0b1010_1010, 0xFF, 0x00, 0xFF, 0xD9];
let mut r = BitReader::new(&data);
assert_eq!(r.receive(8).unwrap(), 0b1010_1010);
assert_eq!(r.receive(8).unwrap(), 0xFF);
assert!(r.receive(1).is_err(), "marker ends entropy data");
}
#[test]
fn restart_consumption() {
let data = [0xAB, 0xFF, 0xD3, 0xCD];
let mut r = BitReader::new(&data);
assert_eq!(r.receive(4).unwrap(), 0xA);
r.expect_restart(3).unwrap();
assert_eq!(r.receive(8).unwrap(), 0xCD);
let data = [0xFF, 0xD4];
let mut r = BitReader::new(&data);
assert!(r
.expect_restart(3)
.unwrap_err()
.to_string()
.contains("RST3"));
}
#[test]
fn extend_matches_spec_table() {
assert_eq!(extend(0b00, 2), -3);
assert_eq!(extend(0b01, 2), -2);
assert_eq!(extend(0b10, 2), 2);
assert_eq!(extend(0b11, 2), 3);
assert_eq!(extend(0, 0), 0);
}
#[test]
fn huffman_canonical_decode() {
let mut counts = [0u8; 16];
counts[0] = 1;
counts[1] = 1;
let t = HuffTable::build(&counts, &[5, 9]).unwrap();
#[allow(clippy::unusual_byte_groupings)]
let data = [0b0_10_0_10_0 << 1];
let mut r = BitReader::new(&data);
assert_eq!(t.decode(&mut r).unwrap(), 5);
assert_eq!(t.decode(&mut r).unwrap(), 9);
assert_eq!(t.decode(&mut r).unwrap(), 5);
}
#[test]
fn huffman_build_rejects_lies() {
let mut counts = [0u8; 16];
counts[0] = 2; assert!(
HuffTable::build(&counts, &[1]).is_err(),
"count/value mismatch"
);
counts[0] = 3;
assert!(
HuffTable::build(&counts, &[1, 2, 3]).is_err(),
"codes overflow length"
);
}
}