use oxideav_core::{Error, Result};
use crate::entropy::{decode_scanned_coefficients, encode_scanned_coefficients};
use crate::frame::ChromaFormat;
use crate::quant::{BLOCK_SCAN_INTERLACED, BLOCK_SCAN_PROGRESSIVE};
pub const MAX_MBS_PER_SLICE: usize = 8;
pub const LUMA_BLOCKS_PER_MB: usize = 4;
pub fn chroma_blocks_per_mb(chroma: ChromaFormat) -> usize {
match chroma {
ChromaFormat::Y422 => 2,
ChromaFormat::Y444 => 4,
}
}
pub fn blocks_per_mb(chroma: ChromaFormat) -> usize {
LUMA_BLOCKS_PER_MB + 2 * chroma_blocks_per_mb(chroma)
}
pub struct DecodedSlice {
pub mb_count: u8,
pub quant_index: u8,
pub blocks: Vec<[i32; 64]>,
}
pub fn encode_slice_components(
mb_count: usize,
chroma: ChromaFormat,
interlaced: bool,
blocks: &[[i32; 64]],
) -> Result<(Vec<u8>, Vec<u8>, Vec<u8>)> {
if mb_count == 0 || mb_count > MAX_MBS_PER_SLICE {
return Err(Error::invalid("prores: slice mb_count out of range"));
}
let cb_per_mb = chroma_blocks_per_mb(chroma);
let per_mb = LUMA_BLOCKS_PER_MB + 2 * cb_per_mb;
if blocks.len() != mb_count * per_mb {
return Err(Error::invalid("prores: slice block-count mismatch"));
}
let scan = if interlaced {
&BLOCK_SCAN_INTERLACED
} else {
&BLOCK_SCAN_PROGRESSIVE
};
let y_coeffs = build_slice_scan(blocks, mb_count, per_mb, 0, LUMA_BLOCKS_PER_MB, scan);
let cb_coeffs = build_slice_scan(
blocks,
mb_count,
per_mb,
LUMA_BLOCKS_PER_MB,
cb_per_mb,
scan,
);
let cr_coeffs = build_slice_scan(
blocks,
mb_count,
per_mb,
LUMA_BLOCKS_PER_MB + cb_per_mb,
cb_per_mb,
scan,
);
let y_bits = encode_scanned_coefficients(&y_coeffs, mb_count * LUMA_BLOCKS_PER_MB)?;
let cb_bits = encode_scanned_coefficients(&cb_coeffs, mb_count * cb_per_mb)?;
let cr_bits = encode_scanned_coefficients(&cr_coeffs, mb_count * cb_per_mb)?;
Ok((y_bits, cb_bits, cr_bits))
}
pub fn decode_slice_components(
y_data: &[u8],
cb_data: &[u8],
cr_data: &[u8],
mb_count: usize,
chroma: ChromaFormat,
interlaced: bool,
) -> Result<Vec<[i32; 64]>> {
let cb_per_mb = chroma_blocks_per_mb(chroma);
let per_mb = LUMA_BLOCKS_PER_MB + 2 * cb_per_mb;
let scan = if interlaced {
&BLOCK_SCAN_INTERLACED
} else {
&BLOCK_SCAN_PROGRESSIVE
};
let y_blocks = mb_count * LUMA_BLOCKS_PER_MB;
let c_blocks = mb_count * cb_per_mb;
let y_coeffs = decode_scanned_coefficients(y_data, y_blocks)?;
let cb_coeffs = decode_scanned_coefficients(cb_data, c_blocks)?;
let cr_coeffs = decode_scanned_coefficients(cr_data, c_blocks)?;
let mut out = vec![[0i32; 64]; mb_count * per_mb];
inverse_slice_scan(
&y_coeffs,
&mut out,
mb_count,
per_mb,
0,
LUMA_BLOCKS_PER_MB,
scan,
);
inverse_slice_scan(
&cb_coeffs,
&mut out,
mb_count,
per_mb,
LUMA_BLOCKS_PER_MB,
cb_per_mb,
scan,
);
inverse_slice_scan(
&cr_coeffs,
&mut out,
mb_count,
per_mb,
LUMA_BLOCKS_PER_MB + cb_per_mb,
cb_per_mb,
scan,
);
Ok(out)
}
fn build_slice_scan(
blocks: &[[i32; 64]],
mb_count: usize,
per_mb: usize,
component_offset: usize,
blocks_per_mb_component: usize,
block_scan: &[u8; 64],
) -> Vec<i32> {
let total_blocks = mb_count * blocks_per_mb_component;
let mut out = vec![0i32; total_blocks * 64];
let mut inv = [0u8; 64];
for (nat_idx, &k) in block_scan.iter().enumerate() {
inv[k as usize] = nat_idx as u8;
}
for n in 0..64 {
let nat_pos = inv[n] as usize;
for m in 0..mb_count {
for b in 0..blocks_per_mb_component {
let block_idx = m * per_mb + component_offset + b;
let blk = &blocks[block_idx];
let dst_idx = n * total_blocks + m * blocks_per_mb_component + b;
out[dst_idx] = blk[nat_pos];
}
}
}
out
}
#[allow(clippy::too_many_arguments)]
fn inverse_slice_scan(
coeffs: &[i32],
blocks_out: &mut [[i32; 64]],
mb_count: usize,
per_mb: usize,
component_offset: usize,
blocks_per_mb_component: usize,
block_scan: &[u8; 64],
) {
let total_blocks = mb_count * blocks_per_mb_component;
let mut inv = [0u8; 64];
for (nat_idx, &k) in block_scan.iter().enumerate() {
inv[k as usize] = nat_idx as u8;
}
for n in 0..64 {
let nat_pos = inv[n] as usize;
for m in 0..mb_count {
for b in 0..blocks_per_mb_component {
let block_idx = m * per_mb + component_offset + b;
let src_idx = n * total_blocks + m * blocks_per_mb_component + b;
blocks_out[block_idx][nat_pos] = coeffs[src_idx];
}
}
}
}
#[cfg(test)]
mod tests {
use super::*;
fn synth_blocks(mb_count: usize, chroma: ChromaFormat) -> Vec<[i32; 64]> {
let per_mb = blocks_per_mb(chroma);
let mut blocks = Vec::with_capacity(mb_count * per_mb);
for m in 0..mb_count {
for b in 0..per_mb {
let mut blk = [0i32; 64];
blk[0] = ((m * 13 + b * 7) as i32 % 41) - 20; for k in 1..8 {
blk[k] = (((m + b + k) as i32) % 5) - 2;
}
blocks.push(blk);
}
}
blocks
}
#[test]
fn slice_components_roundtrip_422() {
let mb_count = 4;
let chroma = ChromaFormat::Y422;
let blocks = synth_blocks(mb_count, chroma);
let (y, cb, cr) = encode_slice_components(mb_count, chroma, false, &blocks).unwrap();
let decoded = decode_slice_components(&y, &cb, &cr, mb_count, chroma, false).unwrap();
assert_eq!(decoded.len(), blocks.len());
for (i, (a, b)) in blocks.iter().zip(decoded.iter()).enumerate() {
assert_eq!(a, b, "block {i} differs");
}
}
#[test]
fn slice_components_roundtrip_444() {
let mb_count = 2;
let chroma = ChromaFormat::Y444;
let blocks = synth_blocks(mb_count, chroma);
let (y, cb, cr) = encode_slice_components(mb_count, chroma, false, &blocks).unwrap();
let decoded = decode_slice_components(&y, &cb, &cr, mb_count, chroma, false).unwrap();
for (i, (a, b)) in blocks.iter().zip(decoded.iter()).enumerate() {
assert_eq!(a, b, "block {i} differs");
}
}
#[test]
fn slice_components_roundtrip_8mb() {
let mb_count = 8;
let chroma = ChromaFormat::Y422;
let blocks = synth_blocks(mb_count, chroma);
let (y, cb, cr) = encode_slice_components(mb_count, chroma, false, &blocks).unwrap();
let decoded = decode_slice_components(&y, &cb, &cr, mb_count, chroma, false).unwrap();
assert_eq!(decoded, blocks);
}
}