#![forbid(unsafe_code)]
use crate::DecodeError;
use mediaway_sw::h264::{BitReader, H264Error};
use super::av1_sequence_header::SequenceHeader;
fn map_bit_err<T>(r: Result<T, H264Error>) -> Result<T, DecodeError> {
r.map_err(|_err| DecodeError::InvalidInput)
}
fn read_bit(r: &mut BitReader<'_>) -> Result<bool, DecodeError> {
Ok(map_bit_err(r.read_bit())? != 0)
}
fn read_bits(r: &mut BitReader<'_>, count: u32) -> Result<u32, DecodeError> {
map_bit_err(r.read_bits(count))
}
fn read_su(r: &mut BitReader<'_>, n: u32) -> Result<i32, DecodeError> {
let value = i64::from(read_bits(r, n)?);
let sign_mask = 1i64 << (n - 1);
let signed = if value & sign_mask != 0 {
value - (sign_mask << 1)
} else {
value
};
i32::try_from(signed).map_err(|_err| DecodeError::InvalidInput)
}
const fn tile_log2(blk_size: u32, target: u32) -> u32 {
let mut k = 0;
while (blk_size << k) < target {
k += 1;
}
k
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub(super) struct TileInfo {
pub(super) tile_width_sb: u32,
pub(super) tile_height_sb: u32,
}
fn parse_tile_info(
r: &mut BitReader<'_>,
use_128x128_superblock: bool,
mi_cols: u32,
mi_rows: u32,
) -> Result<TileInfo, DecodeError> {
const MAX_TILE_WIDTH: u32 = 4096;
const MAX_TILE_AREA: u32 = 4096 * 2304;
const MAX_TILE_COLS: u32 = 64;
const MAX_TILE_ROWS: u32 = 64;
let sb_shift = if use_128x128_superblock { 5 } else { 4 };
let sb_size = sb_shift + 2;
let sb_cols = if use_128x128_superblock {
(mi_cols + 31) >> 5
} else {
(mi_cols + 15) >> 4
};
let sb_rows = if use_128x128_superblock {
(mi_rows + 31) >> 5
} else {
(mi_rows + 15) >> 4
};
let max_tile_width_sb = MAX_TILE_WIDTH >> sb_size;
let max_tile_area_sb = MAX_TILE_AREA >> (2 * sb_size);
let min_log2_tile_cols = tile_log2(max_tile_width_sb, sb_cols);
let max_log2_tile_cols = tile_log2(1, sb_cols.min(MAX_TILE_COLS));
let max_log2_tile_rows = tile_log2(1, sb_rows.min(MAX_TILE_ROWS));
let min_log2_tiles =
min_log2_tile_cols.max(tile_log2(max_tile_area_sb, sb_rows.saturating_mul(sb_cols)));
let uniform_tile_spacing_flag = read_bit(r)?;
if !uniform_tile_spacing_flag {
return Err(DecodeError::Unsupported);
}
let mut tile_cols_log2 = min_log2_tile_cols;
while tile_cols_log2 < max_log2_tile_cols {
if read_bit(r)? {
tile_cols_log2 += 1;
} else {
break;
}
}
let min_log2_tile_rows = min_log2_tiles.saturating_sub(tile_cols_log2);
let mut tile_rows_log2 = min_log2_tile_rows;
while tile_rows_log2 < max_log2_tile_rows {
if read_bit(r)? {
tile_rows_log2 += 1;
} else {
break;
}
}
if tile_cols_log2 > 0 || tile_rows_log2 > 0 {
return Err(DecodeError::Unsupported);
}
Ok(TileInfo {
tile_width_sb: sb_cols,
tile_height_sb: sb_rows,
})
}
fn read_delta_q(r: &mut BitReader<'_>) -> Result<i32, DecodeError> {
if read_bit(r)? { read_su(r, 7) } else { Ok(0) }
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
pub(super) struct Quantization {
pub(super) base_q_idx: u32,
pub(super) delta_q_y_dc: i32,
pub(super) delta_q_u_dc: i32,
pub(super) delta_q_u_ac: i32,
pub(super) delta_q_v_dc: i32,
pub(super) delta_q_v_ac: i32,
}
#[allow(
clippy::similar_names,
reason = "delta_q_y_dc/delta_q_u_dc/delta_q_v_dc are the real AV1 spec DeltaQYDc/\
DeltaQUDc/DeltaQVDc variable names (§5.9.12 quantization_params()) — renaming to look \
less similar would obscure the 1:1 spec mapping, mirrors hevc_vps_sps_pps.rs's \
log2_min_cb_size/log2_min_tb_size identical allow"
)]
fn parse_quantization_params(
r: &mut BitReader<'_>,
separate_uv_delta_q: bool,
) -> Result<Quantization, DecodeError> {
let base_q_idx = read_bits(r, 8)?;
let delta_q_y_dc = read_delta_q(r)?;
let diff_uv_delta = if separate_uv_delta_q {
read_bit(r)?
} else {
false
};
let delta_q_u_dc = read_delta_q(r)?;
let delta_q_u_ac = read_delta_q(r)?;
let (delta_q_v_dc, delta_q_v_ac) = if diff_uv_delta {
(read_delta_q(r)?, read_delta_q(r)?)
} else {
(delta_q_u_dc, delta_q_u_ac)
};
let using_qmatrix = read_bit(r)?;
if using_qmatrix {
return Err(DecodeError::Unsupported);
}
Ok(Quantization {
base_q_idx,
delta_q_y_dc,
delta_q_u_dc,
delta_q_u_ac,
delta_q_v_dc,
delta_q_v_ac,
})
}
fn parse_segmentation_params(r: &mut BitReader<'_>) -> Result<(), DecodeError> {
if read_bit(r)? {
return Err(DecodeError::Unsupported);
}
Ok(())
}
fn parse_delta_q_params(
r: &mut BitReader<'_>,
base_q_idx: u32,
) -> Result<(bool, u32), DecodeError> {
let delta_q_present = if base_q_idx > 0 { read_bit(r)? } else { false };
let delta_q_res = if delta_q_present { read_bits(r, 2)? } else { 0 };
Ok((delta_q_present, delta_q_res))
}
fn parse_delta_lf_params(
r: &mut BitReader<'_>,
delta_q_present: bool,
) -> Result<(bool, u32, bool), DecodeError> {
if !delta_q_present {
return Ok((false, 0, false));
}
let delta_lf_present = read_bit(r)?;
if !delta_lf_present {
return Ok((false, 0, false));
}
let delta_lf_res = read_bits(r, 2)?;
let delta_lf_multi = read_bit(r)?;
Ok((delta_lf_present, delta_lf_res, delta_lf_multi))
}
const DEFAULT_REF_DELTAS: [i32; 8] = [1, 0, 0, 0, -1, 0, 0, -1];
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub(super) struct LoopFilter {
pub(super) level: [u32; 2],
pub(super) level_u: u32,
pub(super) level_v: u32,
pub(super) sharpness: u32,
pub(super) delta_enabled: bool,
pub(super) delta_update: bool,
pub(super) ref_deltas: [i32; 8],
pub(super) mode_deltas: [i32; 2],
}
fn parse_loop_filter_params(
r: &mut BitReader<'_>,
coded_lossless: bool,
) -> Result<LoopFilter, DecodeError> {
if coded_lossless {
return Ok(LoopFilter {
level: [0, 0],
level_u: 0,
level_v: 0,
sharpness: 0,
delta_enabled: false,
delta_update: false,
ref_deltas: DEFAULT_REF_DELTAS,
mode_deltas: [0, 0],
});
}
let level0 = read_bits(r, 6)?;
let level1 = read_bits(r, 6)?;
let (level_u, level_v) = if level0 != 0 || level1 != 0 {
(read_bits(r, 6)?, read_bits(r, 6)?)
} else {
(0, 0)
};
let sharpness = read_bits(r, 3)?;
let delta_enabled = read_bit(r)?;
let mut ref_deltas = DEFAULT_REF_DELTAS;
let mut mode_deltas = [0i32, 0];
let mut delta_update = false;
if delta_enabled {
delta_update = read_bit(r)?;
if delta_update {
for delta in &mut ref_deltas {
if read_bit(r)? {
*delta = read_su(r, 7)?;
}
}
for delta in &mut mode_deltas {
if read_bit(r)? {
*delta = read_su(r, 7)?;
}
}
}
}
Ok(LoopFilter {
level: [level0, level1],
level_u,
level_v,
sharpness,
delta_enabled,
delta_update,
ref_deltas,
mode_deltas,
})
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
#[allow(
clippy::struct_excessive_bools,
reason = "each bool is a real, independent AV1 frame-header flag that must be echoed \
into DXVA_PicParams_AV1 exactly as signaled — same reasoning \
hevc_vps_sps_pps.rs's Pps gives for its own identical allow"
)]
pub(super) struct FrameHeader {
pub(super) width: u32,
pub(super) height: u32,
pub(super) disable_cdf_update: bool,
pub(super) disable_frame_end_update_cdf: bool,
pub(super) order_hint: u32,
pub(super) quantization: Quantization,
pub(super) delta_q_present: bool,
pub(super) delta_q_res: u32,
pub(super) delta_lf_present: bool,
pub(super) delta_lf_res: u32,
pub(super) delta_lf_multi: bool,
pub(super) loop_filter: LoopFilter,
pub(super) tx_mode: u8,
pub(super) reduced_tx_set: bool,
pub(super) tile: TileInfo,
}
#[allow(
clippy::too_many_lines,
reason = "one linear AV1 spec §5.9.2 syntax-element sequence through the fields this \
module needs; mirrors hevc_slice.rs::parse_slice_header's identical shape"
)]
pub(super) fn parse_frame_header(
payload: &[u8],
seq: &SequenceHeader,
) -> Result<(FrameHeader, usize), DecodeError> {
const KEY_FRAME: u32 = 0;
let mut r = BitReader::new(payload);
let show_existing_frame = read_bit(&mut r)?;
if show_existing_frame {
return Err(DecodeError::Unsupported);
}
let frame_type = read_bits(&mut r, 2)?;
if frame_type != KEY_FRAME {
return Err(DecodeError::Unsupported);
}
let show_frame = read_bit(&mut r)?;
if !show_frame {
return Err(DecodeError::Unsupported);
}
let disable_cdf_update = read_bit(&mut r)?;
let frame_size_override_flag = read_bit(&mut r)?;
let order_hint = if seq.order_hint_bits > 0 {
read_bits(&mut r, seq.order_hint_bits)?
} else {
0
};
let (width, height) = if frame_size_override_flag {
let w = read_bits(&mut r, seq.frame_width_bits)?
.checked_add(1)
.ok_or(DecodeError::InvalidInput)?;
let h = read_bits(&mut r, seq.frame_height_bits)?
.checked_add(1)
.ok_or(DecodeError::InvalidInput)?;
(w, h)
} else {
(seq.max_frame_width, seq.max_frame_height)
};
let render_and_frame_size_different = read_bit(&mut r)?;
if render_and_frame_size_different {
let _render_width_minus_1 = read_bits(&mut r, 16)?;
let _render_height_minus_1 = read_bits(&mut r, 16)?;
}
let disable_frame_end_update_cdf = if disable_cdf_update {
true
} else {
read_bit(&mut r)?
};
let mi_cols = 2 * ((width + 7) >> 3);
let mi_rows = 2 * ((height + 7) >> 3);
let tile = parse_tile_info(&mut r, seq.use_128x128_superblock, mi_cols, mi_rows)?;
let quantization = parse_quantization_params(&mut r, seq.separate_uv_delta_q)?;
parse_segmentation_params(&mut r)?;
let (delta_q_present, delta_q_res) = parse_delta_q_params(&mut r, quantization.base_q_idx)?;
let (delta_lf_present, delta_lf_res, delta_lf_multi) =
parse_delta_lf_params(&mut r, delta_q_present)?;
let coded_lossless = quantization.base_q_idx == 0
&& quantization.delta_q_y_dc == 0
&& quantization.delta_q_u_dc == 0
&& quantization.delta_q_u_ac == 0
&& quantization.delta_q_v_dc == 0
&& quantization.delta_q_v_ac == 0;
let loop_filter = parse_loop_filter_params(&mut r, coded_lossless)?;
let tx_mode = if coded_lossless {
0u8 } else if read_bit(&mut r)? {
2u8 } else {
1u8 };
let reduced_tx_set = read_bit(&mut r)?;
Ok((
FrameHeader {
width,
height,
disable_cdf_update,
disable_frame_end_update_cdf,
order_hint,
quantization,
delta_q_present,
delta_q_res,
delta_lf_present,
delta_lf_res,
delta_lf_multi,
loop_filter,
tx_mode,
reduced_tx_set,
tile,
},
r.bits_read(),
))
}
#[cfg(test)]
#[path = "av1_frame_header_tests.rs"]
mod tests;