use oxideav_core::{Error, Result};
use crate::bitstream::{BitReader, BitWriter};
pub fn read_combo(
br: &mut BitReader<'_>,
last_rice_q: u32,
k_rice: u32,
k_exp: u32,
) -> Result<u32> {
const MAX_PREFIX: u32 = 64;
let mut q = 0u32;
while br.read_bit()? == 0 {
q += 1;
if q > MAX_PREFIX {
return Err(Error::invalid("prores entropy: unary prefix too long"));
}
}
if q <= last_rice_q {
let tail = if k_rice == 0 {
0
} else {
br.read_bits(k_rice)?
};
let value = (((q as u64) << k_rice) + tail as u64)
.try_into()
.map_err(|_| Error::invalid("prores entropy: combo rice value out of range"))?;
Ok(value)
} else {
let q_exp = q - (last_rice_q + 1);
let suffix_bits = q_exp + k_exp;
if suffix_bits > 32 {
return Err(Error::invalid(
"prores entropy: combo exp-golomb suffix too wide",
));
}
let suffix = if suffix_bits == 0 {
0
} else {
br.read_bits(suffix_bits)?
};
let exp_val = ((1u64 << q_exp) << k_exp) - (1u64 << k_exp) + suffix as u64;
let value = (((last_rice_q as u64 + 1) << k_rice) + exp_val)
.try_into()
.map_err(|_| Error::invalid("prores entropy: combo exp value out of range"))?;
Ok(value)
}
}
pub fn write_combo(bw: &mut BitWriter, n: u32, last_rice_q: u32, k_rice: u32, k_exp: u32) {
let switch_value = (last_rice_q + 1) << k_rice;
if n < switch_value {
let q = n >> k_rice;
let r = n & ((1u32 << k_rice) - 1);
for _ in 0..q {
bw.write_bit(0);
}
bw.write_bit(1);
if k_rice > 0 {
bw.write_bits(r, k_rice);
}
} else {
let n2 = n - switch_value;
let x = (n2 as u64) + (1u64 << k_exp);
let q_exp = (63 - x.leading_zeros()) - k_exp;
for _ in 0..(last_rice_q + 1 + q_exp) {
bw.write_bit(0);
}
bw.write_bit(1);
let suffix_bits = q_exp + k_exp;
if suffix_bits > 0 {
let suffix_val = (x - (1u64 << (q_exp + k_exp))) as u32;
bw.write_bits(suffix_val, suffix_bits);
}
}
}
pub fn read_exp_golomb(br: &mut BitReader<'_>, k: u32) -> Result<u32> {
const MAX_PREFIX: u32 = 32;
let mut q = 0u32;
while br.read_bit()? == 0 {
q += 1;
if q > MAX_PREFIX {
return Err(Error::invalid("prores entropy: exp-golomb prefix too long"));
}
}
let bits = q + k;
if bits > 32 {
return Err(Error::invalid("prores entropy: exp-golomb suffix too wide"));
}
let suffix = if bits == 0 { 0 } else { br.read_bits(bits)? };
let val = ((1u64 << q) << k) - (1u64 << k) + suffix as u64;
val.try_into()
.map_err(|_| Error::invalid("prores entropy: exp-golomb value out of range"))
}
pub fn write_exp_golomb(bw: &mut BitWriter, n: u32, k: u32) {
let x = (n as u64) + (1u64 << k);
let q = (63 - x.leading_zeros()) - k;
for _ in 0..q {
bw.write_bit(0);
}
bw.write_bit(1);
let bits = q + k;
if bits > 0 {
let suffix = (x - (1u64 << (q + k))) as u32;
bw.write_bits(suffix, bits);
}
}
#[derive(Copy, Clone, Debug)]
pub enum Codebook {
ExpGolomb(u32),
RiceExp {
last_rice_q: u32,
k_rice: u32,
k_exp: u32,
},
}
impl Codebook {
pub fn read(self, br: &mut BitReader<'_>) -> Result<u32> {
match self {
Codebook::ExpGolomb(k) => read_exp_golomb(br, k),
Codebook::RiceExp {
last_rice_q,
k_rice,
k_exp,
} => read_combo(br, last_rice_q, k_rice, k_exp),
}
}
pub fn write(self, bw: &mut BitWriter, n: u32) {
match self {
Codebook::ExpGolomb(k) => write_exp_golomb(bw, n, k),
Codebook::RiceExp {
last_rice_q,
k_rice,
k_exp,
} => write_combo(bw, n, last_rice_q, k_rice, k_exp),
}
}
}
pub fn dc_diff_codebook(prev_abs: u32) -> Codebook {
match prev_abs {
0 => Codebook::ExpGolomb(0),
1 => Codebook::ExpGolomb(1),
2 => Codebook::RiceExp {
last_rice_q: 1,
k_rice: 2,
k_exp: 3,
},
_ => Codebook::ExpGolomb(3),
}
}
pub fn run_codebook(prev_run: u32) -> Codebook {
match prev_run {
0 | 1 => Codebook::RiceExp {
last_rice_q: 2,
k_rice: 0,
k_exp: 1,
},
2 | 3 => Codebook::RiceExp {
last_rice_q: 1,
k_rice: 0,
k_exp: 1,
},
4 => Codebook::ExpGolomb(0),
5..=8 => Codebook::RiceExp {
last_rice_q: 1,
k_rice: 1,
k_exp: 2,
},
9..=14 => Codebook::ExpGolomb(1),
_ => Codebook::ExpGolomb(2), }
}
pub fn level_codebook(prev_level_symbol: u32) -> Codebook {
match prev_level_symbol {
0 => Codebook::RiceExp {
last_rice_q: 2,
k_rice: 0,
k_exp: 2,
},
1 => Codebook::RiceExp {
last_rice_q: 1,
k_rice: 0,
k_exp: 1,
},
2 => Codebook::RiceExp {
last_rice_q: 2,
k_rice: 0,
k_exp: 1,
},
3 => Codebook::ExpGolomb(0),
4..=7 => Codebook::ExpGolomb(1),
_ => Codebook::ExpGolomb(2), }
}
pub fn decode_scanned_coefficients(data: &[u8], num_blocks: usize) -> Result<Vec<i32>> {
let total = num_blocks
.checked_mul(64)
.ok_or_else(|| Error::invalid("prores entropy: num_blocks overflow"))?;
let mut coeffs = vec![0i32; total];
if num_blocks == 0 {
return Ok(coeffs);
}
let mut br = BitReader::new(data);
let s = read_exp_golomb(&mut br, 5)?;
let first_dc = inv_signed_mapping(s);
coeffs[0] = first_dc;
let mut previous_dc_coeff = first_dc;
let mut previous_dc_diff: i32 = 3; for n in 1..num_blocks {
let cb = dc_diff_codebook(previous_dc_diff.unsigned_abs());
let s = cb.read(&mut br)?;
let mut diff = inv_signed_mapping(s);
if previous_dc_diff < 0 {
diff = diff.wrapping_neg();
}
let dc = previous_dc_coeff.wrapping_add(diff);
coeffs[n] = dc;
previous_dc_coeff = dc;
previous_dc_diff = diff;
}
let mut n = num_blocks; let mut previous_run: u32 = 4; let mut previous_level_symbol: u32 = 1;
while n < total && !br.end_of_data() {
let run_cb = run_codebook(previous_run);
let run = run_cb.read(&mut br)?;
previous_run = run;
let advance = run as usize;
if n + advance >= total {
return Err(Error::invalid("prores entropy: AC run overruns array"));
}
n += advance;
let lvl_cb = level_codebook(previous_level_symbol);
let abs_minus_1 = lvl_cb.read(&mut br)?;
previous_level_symbol = abs_minus_1;
let abs_level = (abs_minus_1 as i32).wrapping_add(1);
let sign = br.read_bit()?;
let level = abs_level.wrapping_mul(1 - 2 * sign as i32);
coeffs[n] = level;
n += 1;
}
Ok(coeffs)
}
pub fn encode_scanned_coefficients(coeffs: &[i32], num_blocks: usize) -> Result<Vec<u8>> {
let total = num_blocks * 64;
if coeffs.len() != total {
return Err(Error::invalid(
"prores entropy: coefficient buffer size mismatch",
));
}
let mut bw = BitWriter::new();
if num_blocks == 0 {
return Ok(bw.finish());
}
let first_dc = coeffs[0];
write_exp_golomb(&mut bw, signed_mapping(first_dc), 5);
let mut previous_dc_coeff = first_dc;
let mut previous_dc_diff: i32 = 3;
for n in 1..num_blocks {
let dc = coeffs[n];
let diff = dc - previous_dc_coeff;
let cb = dc_diff_codebook(previous_dc_diff.unsigned_abs());
let stored = if previous_dc_diff < 0 { -diff } else { diff };
cb.write(&mut bw, signed_mapping(stored));
previous_dc_coeff = dc;
previous_dc_diff = diff;
}
let mut previous_run: u32 = 4;
let mut previous_level_symbol: u32 = 1;
let mut run_acc: u32 = 0;
for n in num_blocks..total {
let v = coeffs[n];
if v == 0 {
run_acc += 1;
continue;
}
let run_cb = run_codebook(previous_run);
run_cb.write(&mut bw, run_acc);
previous_run = run_acc;
let abs_level = v.unsigned_abs();
let abs_minus_1 = abs_level - 1;
let lvl_cb = level_codebook(previous_level_symbol);
lvl_cb.write(&mut bw, abs_minus_1);
previous_level_symbol = abs_minus_1;
bw.write_bit(if v < 0 { 1 } else { 0 });
run_acc = 0;
}
bw.align_byte();
Ok(bw.finish())
}
pub fn signed_mapping(n: i32) -> u32 {
if n >= 0 {
(2 * n) as u32
} else {
(2 * (-n) - 1) as u32
}
}
pub fn inv_signed_mapping(s: u32) -> i32 {
if (s & 1) == 0 {
(s >> 1) as i32
} else {
-((s.wrapping_add(1) >> 1) as i32)
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn zero_block_slice_decodes_to_empty_without_reading_bits() {
for payload in [&[][..], &[0x88][..], &[0xff, 0x00, 0x13][..]] {
let out = decode_scanned_coefficients(payload, 0).expect("zero-block decode");
assert!(out.is_empty(), "zero-block scan array must be empty");
}
}
#[test]
fn zero_block_slice_encodes_to_empty_payload() {
let out = encode_scanned_coefficients(&[], 0).expect("zero-block encode");
assert!(out.is_empty(), "zero-block payload must be empty");
let back = decode_scanned_coefficients(&out, 0).expect("zero-block decode");
assert!(back.is_empty());
}
#[test]
fn signed_mapping_roundtrip() {
for n in -1000..=1000 {
let s = signed_mapping(n);
assert_eq!(inv_signed_mapping(s), n, "roundtrip n={n}");
}
}
#[test]
fn signed_mapping_table8() {
assert_eq!(signed_mapping(0), 0);
assert_eq!(signed_mapping(-1), 1);
assert_eq!(signed_mapping(1), 2);
assert_eq!(signed_mapping(-2), 3);
assert_eq!(signed_mapping(2), 4);
assert_eq!(signed_mapping(-3), 5);
assert_eq!(signed_mapping(3), 6);
}
#[test]
fn exp_golomb_roundtrip() {
for k in [0u32, 1, 2, 3, 5] {
for v in [0u32, 1, 2, 3, 7, 15, 16, 100, 1000, 65_535] {
let mut bw = BitWriter::new();
write_exp_golomb(&mut bw, v, k);
let buf = bw.finish();
let mut br = BitReader::new(&buf);
let got = read_exp_golomb(&mut br, k).expect("decode");
assert_eq!(got, v, "k={k} v={v}");
}
}
}
#[test]
fn combo_roundtrip() {
let params = [
(1u32, 2u32, 3u32),
(2, 0, 1),
(1, 0, 1),
(1, 1, 2),
(2, 0, 2),
];
for (lq, kr, ke) in params {
for v in 0u32..200 {
let mut bw = BitWriter::new();
write_combo(&mut bw, v, lq, kr, ke);
let buf = bw.finish();
let mut br = BitReader::new(&buf);
let got = read_combo(&mut br, lq, kr, ke).expect("decode");
assert_eq!(got, v, "(lq={lq},kr={kr},ke={ke}) v={v}");
}
}
}
#[test]
fn combo_eq_exp_golomb_when_lq_zero() {
for k in [0u32, 1, 2, 3, 5] {
for v in [0u32, 1, 5, 100, 1000] {
let mut a = BitWriter::new();
write_exp_golomb(&mut a, v, k);
let mut b = BitWriter::new();
write_combo(&mut b, v, 0, k, k + 1);
assert_eq!(a.finish(), b.finish(), "k={k} v={v}");
}
}
}
#[test]
fn rdd36_scanned_coeffs_roundtrip() {
let num_blocks = 4;
let mut coeffs = vec![0i32; num_blocks * 64];
coeffs[0] = 100;
coeffs[1] = 102;
coeffs[2] = 99;
coeffs[3] = 101;
coeffs[num_blocks] = 5; coeffs[num_blocks * 3 + 1] = -3; coeffs[num_blocks * 10 + 2] = 1;
coeffs[num_blocks * 20 + 3] = -1;
let buf = encode_scanned_coefficients(&coeffs, num_blocks).unwrap();
let decoded = decode_scanned_coefficients(&buf, num_blocks).unwrap();
assert_eq!(decoded.len(), coeffs.len());
for (i, (a, b)) in coeffs.iter().zip(decoded.iter()).enumerate() {
assert_eq!(a, b, "coeff {i}");
}
}
#[test]
fn rdd36_dc_only_roundtrip() {
let num_blocks = 8;
let mut coeffs = vec![0i32; num_blocks * 64];
let dcs = [10, -10, 20, 0, -50, 50, 100, -100];
for (i, dc) in dcs.iter().enumerate() {
coeffs[i] = *dc;
}
let buf = encode_scanned_coefficients(&coeffs, num_blocks).unwrap();
let decoded = decode_scanned_coefficients(&buf, num_blocks).unwrap();
for i in 0..num_blocks {
assert_eq!(decoded[i], coeffs[i], "DC {i}");
}
for i in num_blocks..coeffs.len() {
assert_eq!(decoded[i], 0, "AC {i} should stay zero");
}
}
#[test]
fn rdd36_random_pattern_roundtrip() {
let num_blocks = 16;
let mut coeffs = vec![0i32; num_blocks * 64];
let mut dc = 200i32;
for n in 0..num_blocks {
dc += if n % 2 == 0 { -5 } else { 7 };
coeffs[n] = dc;
}
for f in 1..16 {
for b in 0..num_blocks {
let idx = f * num_blocks + b;
let v = ((b as i32 + f as i32 * 3) % 7) - 3;
if v != 0 && (b + f) % 3 == 0 {
coeffs[idx] = v;
}
}
}
let buf = encode_scanned_coefficients(&coeffs, num_blocks).unwrap();
let decoded = decode_scanned_coefficients(&buf, num_blocks).unwrap();
assert_eq!(decoded, coeffs);
}
}