use smallvec::SmallVec;
use crate::DecodeError;
use mediaway_sw::h264::BitReader;
use super::hevc_vps_sps_pps::{HevcNalUnitType, Pps, Sps};
fn read_bit(r: &mut BitReader<'_>) -> Result<bool, DecodeError> {
Ok(r.read_bit().map_err(|_err| DecodeError::InvalidInput)? != 0)
}
fn read_bits(r: &mut BitReader<'_>, count: u32) -> Result<u32, DecodeError> {
r.read_bits(count).map_err(|_err| DecodeError::InvalidInput)
}
fn read_ue(r: &mut BitReader<'_>) -> Result<u32, DecodeError> {
r.read_ue().map_err(|_err| DecodeError::InvalidInput)
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub(super) enum SliceType {
P,
I,
}
impl SliceType {
const fn from_raw(value: u32) -> Option<Self> {
match value {
1 => Some(Self::P),
2 => Some(Self::I),
_ => None,
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub(super) struct ShortTermRefPicEntry {
pub(super) delta_poc: i32,
pub(super) used_by_curr_pic: bool,
}
#[derive(Debug, Clone, Default)]
pub(super) struct ShortTermRefPicSet {
pub(super) s0: SmallVec<[ShortTermRefPicEntry; 8]>,
pub(super) s1: SmallVec<[ShortTermRefPicEntry; 8]>,
}
impl ShortTermRefPicSet {
fn parse(r: &mut BitReader<'_>) -> Result<Self, DecodeError> {
let num_negative_pics = read_ue(r)?;
let num_positive_pics = read_ue(r)?;
if num_negative_pics > 8 || num_positive_pics > 8 {
return Err(DecodeError::Unsupported);
}
let mut s0 = SmallVec::new();
let mut delta_poc = 0i32;
for _ in 0..num_negative_pics {
let delta_poc_s0_minus1 = read_ue(r)?;
let used_by_curr_pic = read_bit(r)?;
let step = i32::try_from(delta_poc_s0_minus1)
.ok()
.and_then(|v| v.checked_add(1))
.ok_or(DecodeError::InvalidInput)?;
delta_poc = delta_poc
.checked_sub(step)
.ok_or(DecodeError::InvalidInput)?;
s0.push(ShortTermRefPicEntry {
delta_poc,
used_by_curr_pic,
});
}
let mut s1 = SmallVec::new();
let mut delta_poc = 0i32;
for _ in 0..num_positive_pics {
let delta_poc_s1_minus1 = read_ue(r)?;
let used_by_curr_pic = read_bit(r)?;
let step = i32::try_from(delta_poc_s1_minus1)
.ok()
.and_then(|v| v.checked_add(1))
.ok_or(DecodeError::InvalidInput)?;
delta_poc = delta_poc
.checked_add(step)
.ok_or(DecodeError::InvalidInput)?;
s1.push(ShortTermRefPicEntry {
delta_poc,
used_by_curr_pic,
});
}
Ok(Self { s0, s1 })
}
pub(super) fn num_curr_pics(&self) -> usize {
self.s0.iter().filter(|e| e.used_by_curr_pic).count()
+ self.s1.iter().filter(|e| e.used_by_curr_pic).count()
}
pub(super) fn curr_before_after_poc(
&self,
current_poc: i32,
) -> (SmallVec<[i32; 8]>, SmallVec<[i32; 8]>) {
let before = self
.s0
.iter()
.filter(|e| e.used_by_curr_pic)
.map(|e| current_poc + e.delta_poc)
.collect();
let after = self
.s1
.iter()
.filter(|e| e.used_by_curr_pic)
.map(|e| current_poc + e.delta_poc)
.collect();
(before, after)
}
pub(super) fn all_poc(&self, current_poc: i32) -> SmallVec<[i32; 16]> {
self.s0
.iter()
.chain(self.s1.iter())
.map(|e| current_poc + e.delta_poc)
.collect()
}
}
#[derive(Debug, Clone)]
pub(super) struct SliceHeader {
pub(super) slice_type: SliceType,
pub(super) pic_order_cnt_lsb: Option<u32>,
pub(super) short_term_rps: Option<ShortTermRefPicSet>,
pub(super) num_ref_idx_l0_active_minus1: u32,
pub(super) short_term_rps_bits: u32,
}
#[allow(
clippy::too_many_lines,
reason = "one linear slice_segment_header() parse sequence; splitting fragments the \
bit-position invariant, mirrors h264_slice.rs::parse_slice_header's identical shape"
)]
pub(super) fn parse_slice_header(
rbsp: &[u8],
nal_unit_type: HevcNalUnitType,
sps: &Sps,
pps: &Pps,
) -> Result<SliceHeader, DecodeError> {
let mut r = BitReader::new(rbsp);
let first_slice_segment_in_pic_flag = read_bit(&mut r)?;
if !first_slice_segment_in_pic_flag {
return Err(DecodeError::Unsupported);
}
if matches!(nal_unit_type, HevcNalUnitType::Idr | HevcNalUnitType::Cra) {
let _no_output_of_prior_pics_flag = read_bit(&mut r)?;
}
let _slice_pic_parameter_set_id = read_ue(&mut r)?;
for _ in 0..pps.num_extra_slice_header_bits {
let _slice_reserved_flag = read_bit(&mut r)?;
}
let slice_type_raw = read_ue(&mut r)?;
let slice_type = SliceType::from_raw(slice_type_raw).ok_or(DecodeError::Unsupported)?;
if pps.output_flag_present_flag {
let _pic_output_flag = read_bit(&mut r)?;
}
let is_idr = nal_unit_type.is_idr();
let (pic_order_cnt_lsb, short_term_rps, short_term_rps_bits) = if is_idr {
(None, None, 0)
} else {
let poc_lsb = read_bits(&mut r, sps.log2_max_pic_order_cnt_lsb)?;
let short_term_ref_pic_set_sps_flag = read_bit(&mut r)?;
if short_term_ref_pic_set_sps_flag {
return Err(DecodeError::Unsupported);
}
let rps_bits_start = r.bits_read();
let rps = ShortTermRefPicSet::parse(&mut r)?;
let rps_bits = u32::try_from(r.bits_read() - rps_bits_start).unwrap_or(u32::MAX);
if matches!(slice_type, SliceType::P) && rps.num_curr_pics() != 1 {
return Err(DecodeError::Unsupported);
}
if sps.sps_temporal_mvp_enabled_flag {
let _slice_temporal_mvp_enabled_flag = read_bit(&mut r)?;
}
(Some(poc_lsb), Some(rps), rps_bits)
};
if sps.sample_adaptive_offset_enabled_flag {
let _slice_sao_luma_flag = read_bit(&mut r)?;
let _slice_sao_chroma_flag = read_bit(&mut r)?;
}
let num_ref_idx_l0_active_minus1 = if matches!(slice_type, SliceType::P) {
let num_ref_idx_active_override_flag = read_bit(&mut r)?;
let value = if num_ref_idx_active_override_flag {
read_ue(&mut r)?
} else {
pps.num_ref_idx_l0_default_active_minus1
};
if value != 0 {
return Err(DecodeError::Unsupported);
}
value
} else {
0
};
Ok(SliceHeader {
slice_type,
pic_order_cnt_lsb,
short_term_rps,
num_ref_idx_l0_active_minus1,
short_term_rps_bits,
})
}
#[cfg(test)]
#[path = "hevc_slice_tests.rs"]
mod tests;