use oxideav_core::{Error, Result};
use crate::bitstream::{BitReader, BitWriter};
use crate::frame::ChromaFormat;
use crate::quant::ZIGZAG;
pub const BLOCKS_PER_MB_422: usize = 8;
pub const BLOCKS_PER_MB_444: usize = 12;
pub const BLOCKS_PER_MB: usize = BLOCKS_PER_MB_422;
pub const MAX_MBS_PER_SLICE: usize = 8;
pub fn blocks_per_mb(chroma: ChromaFormat) -> usize {
match chroma {
ChromaFormat::Y422 => BLOCKS_PER_MB_422,
ChromaFormat::Y444 => BLOCKS_PER_MB_444,
}
}
pub struct DecodedSlice {
pub mb_count: u8,
pub quant_index: u8,
pub blocks: Vec<[i32; 64]>,
}
pub fn encode_slice(
mb_count: u8,
quant_index: u8,
chroma: ChromaFormat,
blocks: &[[i32; 64]],
) -> Result<Vec<u8>> {
if mb_count == 0 || mb_count as usize > MAX_MBS_PER_SLICE {
return Err(Error::invalid("prores: slice mb_count out of range"));
}
let per_mb = blocks_per_mb(chroma);
if blocks.len() != mb_count as usize * per_mb {
return Err(Error::invalid("prores: slice block-count mismatch"));
}
let mut out = Vec::with_capacity(2 + 64 * blocks.len());
out.push(mb_count);
out.push(quant_index);
let mut bw = BitWriter::new();
for blk in blocks {
bw.write_se(blk[0]);
for k in 1..64 {
bw.write_se(blk[ZIGZAG[k] as usize]);
}
}
out.extend(bw.finish());
Ok(out)
}
pub fn decode_slice(data: &[u8], chroma: ChromaFormat) -> Result<DecodedSlice> {
if data.len() < 2 {
return Err(Error::invalid("prores: slice header truncated"));
}
let mb_count = data[0];
let quant_index = data[1];
if mb_count == 0 || mb_count as usize > MAX_MBS_PER_SLICE {
return Err(Error::invalid("prores: slice mb_count out of range"));
}
let per_mb = blocks_per_mb(chroma);
let expected_blocks = mb_count as usize * per_mb;
let mut blocks = Vec::with_capacity(expected_blocks);
let mut br = BitReader::new(&data[2..]);
for _ in 0..expected_blocks {
let mut blk = [0i32; 64];
blk[0] = br.read_se()?;
for k in 1..64 {
blk[ZIGZAG[k] as usize] = br.read_se()?;
}
blocks.push(blk);
}
Ok(DecodedSlice {
mb_count,
quant_index,
blocks,
})
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn slice_roundtrip() {
let mut blocks = Vec::new();
for mb in 0..3 {
for blk_in_mb in 0..BLOCKS_PER_MB_422 {
let mut blk = [0i32; 64];
for k in 0..64 {
blk[k] = (((mb * 13 + blk_in_mb * 7 + k) as i32) % 31) - 15;
}
blocks.push(blk);
}
}
let encoded = encode_slice(3, 4, ChromaFormat::Y422, &blocks).unwrap();
let decoded = decode_slice(&encoded, ChromaFormat::Y422).unwrap();
assert_eq!(decoded.mb_count, 3);
assert_eq!(decoded.quant_index, 4);
assert_eq!(decoded.blocks.len(), blocks.len());
for (i, b) in decoded.blocks.iter().enumerate() {
assert_eq!(*b, blocks[i], "block {i} differs");
}
}
#[test]
fn slice_roundtrip_444() {
let mut blocks = Vec::new();
for mb in 0..2 {
for blk_in_mb in 0..BLOCKS_PER_MB_444 {
let mut blk = [0i32; 64];
for k in 0..64 {
blk[k] = (((mb * 17 + blk_in_mb * 5 + k) as i32) % 29) - 14;
}
blocks.push(blk);
}
}
let encoded = encode_slice(2, 4, ChromaFormat::Y444, &blocks).unwrap();
let decoded = decode_slice(&encoded, ChromaFormat::Y444).unwrap();
assert_eq!(decoded.mb_count, 2);
assert_eq!(decoded.blocks.len(), blocks.len());
for (i, b) in decoded.blocks.iter().enumerate() {
assert_eq!(*b, blocks[i], "444 block {i} differs");
}
}
}