use alloc::vec::Vec;
use pith_digest::{Error, Result};
use crate::bits::{Bits, extend};
use crate::huffman::{Table, ZIGZAG};
use crate::parser::{Component, Frame, Parser, Tables};
pub(crate) fn decode_scan(
p: &mut Parser<'_>,
sos: &[u8],
frame: &mut Frame,
tables: &Tables,
) -> Result<()> {
let (order, ss, se, ah, al) = parse_sos(sos, frame)?;
let mut bits = Bits::new(p.raw(), p.offset());
let ri = tables.restart_interval;
if ss == 0 {
if ah == 0 {
dc_first(&mut bits, &order, frame, tables, al, ri)?;
} else {
dc_refine(&mut bits, &order, frame, al, ri)?;
}
} else {
if order.len() != 1 {
return Err(Error::BadValue("progressive AC scan with Ns > 1"));
}
if se < ss || se > 63 {
return Err(Error::BadValue("AC scan spectral band invalid"));
}
let ci = order[0];
if ah == 0 {
ac_first(&mut bits, ci, frame, tables, Band { ss, se, al }, ri)?;
} else {
if ah.saturating_sub(1) != al {
}
ac_refine(&mut bits, ci, frame, tables, Band { ss, se, al }, ri)?;
}
}
p.seek(bits.pos);
Ok(())
}
fn parse_sos(sos: &[u8], frame: &mut Frame) -> Result<(Vec<usize>, usize, usize, usize, usize)> {
if sos.is_empty() {
return Err(Error::truncated("SOS header", 1, 0));
}
let ns = sos[0] as usize;
if !(1..=4).contains(&ns) {
return Err(Error::BadValue("scan component count outside 1..4"));
}
let need = 1 + 2 * ns + 3;
if sos.len() < need {
return Err(Error::truncated("SOS header", need, sos.len()));
}
let mut order = Vec::with_capacity(ns);
for i in 0..ns {
let id = sos[1 + 2 * i];
let t = sos[2 + 2 * i];
let ci = frame
.comps
.iter()
.position(|c| c.id == id)
.ok_or(Error::BadValue("SOS references unknown component"))?;
frame.comps[ci].td = (t >> 4) as usize;
frame.comps[ci].ta = (t & 15) as usize;
if frame.comps[ci].td > 3 || frame.comps[ci].ta > 3 {
return Err(Error::BadValue("Huffman table selector > 3"));
}
order.push(ci);
}
let ss = sos[1 + 2 * ns] as usize;
let se = sos[2 + 2 * ns] as usize;
let ahal = sos[3 + 2 * ns];
let (ah, al) = ((ahal >> 4) as usize, (ahal & 15) as usize);
if ss > 63 || se > 63 || ah > 13 || al > 13 {
return Err(Error::BadValue("Ss/Se/Ah/Al out of range"));
}
if ss == 0 && se != 0 {
return Err(Error::BadValue("DC scan with Se != 0"));
}
Ok((order, ss, se, ah, al))
}
#[derive(Copy, Clone)]
struct Band {
ss: usize,
se: usize,
al: usize,
}
struct Restart {
interval: usize,
togo: usize,
n: u8,
}
impl Restart {
fn new(interval: usize) -> Self {
Restart {
interval,
togo: interval,
n: 0,
}
}
fn check(&mut self, bits: &mut Bits<'_>) -> Result<bool> {
if self.interval == 0 || self.togo != 0 {
return Ok(false);
}
if !bits.consume_rst(self.n) {
return Err(Error::BadValue("missing restart marker"));
}
self.n = (self.n + 1) & 7;
self.togo = self.interval;
Ok(true)
}
fn tick(&mut self) {
self.togo = self.togo.saturating_sub(1);
}
}
fn dc_table<'t>(comp: &Component, tables: &'t Tables) -> Result<&'t Table> {
tables.huff_dc[comp.td]
.as_ref()
.ok_or(Error::BadValue("scan uses missing DC Huffman table"))
}
fn ac_table<'t>(comp: &Component, tables: &'t Tables) -> Result<&'t Table> {
tables.huff_ac[comp.ta]
.as_ref()
.ok_or(Error::BadValue("scan uses missing AC Huffman table"))
}
fn true_grid(comp: &Component) -> (usize, usize) {
(comp.down_w.div_ceil(8), comp.down_h.div_ceil(8))
}
fn dc_first(
bits: &mut Bits<'_>,
order: &[usize],
frame: &mut Frame,
tables: &Tables,
al: usize,
ri: usize,
) -> Result<()> {
let mut preds = alloc::vec![0i32; order.len()];
let mut tabs: Vec<&Table> = Vec::new();
for &ci in order {
tabs.push(dc_table(&frame.comps[ci], tables)?);
}
let mut rst = Restart::new(ri);
for my in 0..frame.mcus_y {
for mx in 0..frame.mcus_x {
if rst.check(bits)? {
preds.iter_mut().for_each(|p| *p = 0);
}
for (si, &ci) in order.iter().enumerate() {
let comp = &frame.comps[ci];
let (bw, h, v) = (comp.blocks_w, comp.h, comp.v);
for by in 0..v {
for bx in 0..h {
let bi = (my * v + by) * bw + mx * h + bx;
let s = tabs[si].decode(bits)?;
if s > 15 {
return Err(Error::BadValue("DC coefficient size > 15"));
}
let diff = bits.extend(u32::from(s))?;
let v0 = preds[si].wrapping_add(diff);
preds[si] = v0;
frame.comps[ci].coefs[bi * 64] = v0 << al;
}
}
}
rst.tick();
}
}
Ok(())
}
fn dc_refine(
bits: &mut Bits<'_>,
order: &[usize],
frame: &mut Frame,
al: usize,
ri: usize,
) -> Result<()> {
let p1: i32 = 1 << al;
let mut rst = Restart::new(ri);
for my in 0..frame.mcus_y {
for mx in 0..frame.mcus_x {
rst.check(bits)?;
for &ci in order {
let comp = &frame.comps[ci];
let (bw, h, v) = (comp.blocks_w, comp.h, comp.v);
for by in 0..v {
for bx in 0..h {
let bi = (my * v + by) * bw + mx * h + bx;
let b = bits.bits(1)? as i32;
frame.comps[ci].coefs[bi * 64] |= b * p1;
}
}
}
rst.tick();
}
}
Ok(())
}
struct Blocks {
bw: usize,
bh: usize,
bw_store: usize,
bx: usize,
by: usize,
rst: Restart,
eobrun: u32,
done: bool,
}
impl Blocks {
fn new(comp: &Component, ri: usize) -> Self {
let (bw, bh) = true_grid(comp);
Blocks {
bw,
bh,
bw_store: comp.blocks_w,
bx: 0,
by: 0,
rst: Restart::new(ri),
eobrun: 0,
done: false,
}
}
fn next(&mut self, bits: &mut Bits<'_>) -> Result<Option<usize>> {
if self.done || self.by >= self.bh {
self.done = true;
return Ok(None);
}
if self.rst.check(bits)? {
self.eobrun = 0;
}
let bi = self.by * self.bw_store + self.bx;
self.bx += 1;
if self.bx >= self.bw {
self.bx = 0;
self.by += 1;
}
self.rst.tick();
Ok(Some(bi * 64))
}
}
#[inline]
fn refine_bit(block: &mut [i32], idx: usize, p1: i32, m1: i32, set: bool) {
if !set {
return;
}
let cf = block[idx];
if cf != 0 && (cf & p1) == 0 {
block[idx] = cf + if cf >= 0 { p1 } else { m1 };
}
}
fn ac_first(
bits: &mut Bits<'_>,
ci: usize,
frame: &mut Frame,
tables: &Tables,
band: Band,
ri: usize,
) -> Result<()> {
let (ss, se, al) = (band.ss, band.se, band.al);
let tab = ac_table(&frame.comps[ci], tables)?.clone();
let mut walk = Blocks::new(&frame.comps[ci], ri);
let comp = &mut frame.comps[ci];
while let Some(bi) = walk.next(bits)? {
if walk.eobrun > 0 {
walk.eobrun -= 1;
continue;
}
let block = &mut comp.coefs[bi..bi + 64];
let mut k = ss;
while k <= se {
let rs = tab.decode(bits)?;
let r = (rs >> 4) as usize;
let s = (rs & 15) as usize;
if s != 0 {
k += r;
if k > 63 {
break; }
let v = extend(bits.receive(s as u32)?, s as u32);
block[ZIGZAG[k] as usize] = v << al;
k += 1;
} else if r != 15 {
let mut eob = 1u32 << r;
if r > 0 {
eob += bits.receive(r as u32)?;
}
walk.eobrun = eob.saturating_sub(1);
break;
} else {
k += 16; }
}
}
Ok(())
}
fn ac_refine(
bits: &mut Bits<'_>,
ci: usize,
frame: &mut Frame,
tables: &Tables,
band: Band,
ri: usize,
) -> Result<()> {
let (ss, se, al) = (band.ss, band.se, band.al);
let tab = ac_table(&frame.comps[ci], tables)?.clone();
let p1: i32 = 1 << al;
let m1: i32 = -1 << al;
let mut walk = Blocks::new(&frame.comps[ci], ri);
let comp = &mut frame.comps[ci];
while let Some(bi) = walk.next(bits)? {
let block = &mut comp.coefs[bi..bi + 64];
let mut k = ss;
if walk.eobrun == 0 {
while k <= se {
let rs = tab.decode(bits)?;
let mut r = (rs >> 4) as i32;
let mut s = (rs & 15) as i32;
if s != 0 {
s = if bits.receive(1)? != 0 { p1 } else { m1 };
} else if r != 15 {
let mut eob = 1u32 << r;
if r > 0 {
eob += bits.receive(r as u32)?;
}
walk.eobrun = eob;
break;
}
loop {
if k > se {
break;
}
let idx = ZIGZAG[k] as usize;
if block[idx] != 0 {
let set = bits.receive(1)? != 0;
refine_bit(block, idx, p1, m1, set);
} else {
r -= 1;
if r < 0 {
break;
}
}
k += 1;
}
if s != 0 && k <= se {
block[ZIGZAG[k] as usize] = s;
}
k += 1;
}
}
if walk.eobrun > 0 {
while k <= se {
let idx = ZIGZAG[k] as usize;
if block[idx] != 0 {
let set = bits.receive(1)? != 0;
refine_bit(block, idx, p1, m1, set);
}
k += 1;
}
walk.eobrun = walk.eobrun.saturating_sub(1);
}
}
Ok(())
}
#[cfg(test)]
mod tests {
use super::parse_sos;
use crate::parser::parse_sof;
#[test]
fn progressive_sos_empty_is_truncated() {
let body = [8u8, 0, 8, 0, 8, 1, 1, 0x11, 0];
let mut f = parse_sof(0xc2, &body).expect("gray progressive frame");
let err = parse_sos(&[], &mut f).expect_err("empty progressive SOS");
assert!(err.to_string().contains("SOS header"));
}
}