use oxideav_core::bits::BitReader;
use crate::cce::CouplingChannelElement;
use crate::ics_body::IcsBody;
use crate::raw_data_block::{Element, IdSynEle, Walker};
use crate::spectral_data::SpectralData;
use crate::{Error, Result};
#[doc(hidden)]
pub const ADTS_CRC_POLY: u32 = 0x8005;
#[doc(hidden)]
pub const ADTS_CRC_INIT: u32 = 0xFFFF;
#[doc(hidden)]
pub const SBR_CRC_POLY: u32 = 0x0233;
#[doc(hidden)]
#[derive(Debug, Clone, Copy)]
pub struct CrcRegister {
reg: u32,
poly: u32,
mask: u32,
top: u32,
}
impl CrcRegister {
pub fn adts() -> Self {
CrcRegister {
reg: ADTS_CRC_INIT,
poly: ADTS_CRC_POLY,
mask: 0xFFFF,
top: 0x8000,
}
}
pub fn sbr() -> Self {
CrcRegister {
reg: 0,
poly: SBR_CRC_POLY,
mask: 0x03FF,
top: 0x0200,
}
}
#[inline]
pub fn feed_bit(&mut self, bit: bool) {
let feedback = ((self.reg & self.top) != 0) ^ bit;
self.reg = (self.reg << 1) & self.mask;
if feedback {
self.reg ^= self.poly;
}
}
pub fn feed_zeros(&mut self, n: u64) {
for _ in 0..n {
self.feed_bit(false);
}
}
pub fn feed_bit_range(&mut self, data: &[u8], start_bit: u64, end_bit: u64) {
for pos in start_bit..end_bit {
let byte = (pos / 8) as usize;
let bit = data.get(byte).is_some_and(|b| b & (0x80 >> (pos % 8)) != 0);
self.feed_bit(bit);
}
}
pub fn value(&self) -> u16 {
self.reg as u16
}
}
#[doc(hidden)]
pub fn sbr_crc(data: &[u8], start_bit: u64, end_bit: u64) -> u16 {
let mut reg = CrcRegister::sbr();
reg.feed_bit_range(data, start_bit, end_bit);
reg.value()
}
#[doc(hidden)]
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct ProtectedRegion {
pub start_bit: u64,
pub end_bit: u64,
pub pad_to: Option<u32>,
}
impl ProtectedRegion {
fn feed(&self, reg: &mut CrcRegister, payload: &[u8]) {
let len = self.end_bit.saturating_sub(self.start_bit);
match self.pad_to {
Some(pad) => {
let take = len.min(u64::from(pad));
reg.feed_bit_range(payload, self.start_bit, self.start_bit + take);
reg.feed_zeros(u64::from(pad) - take);
}
None => reg.feed_bit_range(payload, self.start_bit, self.end_bit),
}
}
}
#[doc(hidden)]
pub fn adts_single_crc(header: &[u8], payload: &[u8], regions: &[ProtectedRegion]) -> u16 {
let mut reg = CrcRegister::adts();
reg.feed_bit_range(header, 0, 56);
for r in regions {
r.feed(&mut reg, payload);
}
reg.value()
}
#[doc(hidden)]
pub fn adts_header_crc(header: &[u8], positions: &[u16]) -> u16 {
let mut reg = CrcRegister::adts();
reg.feed_bit_range(header, 0, 56);
for &p in positions {
for i in (0..16).rev() {
reg.feed_bit((p >> i) & 1 != 0);
}
}
reg.value()
}
#[doc(hidden)]
pub fn adts_rdb_crc(payload: &[u8], regions: &[ProtectedRegion]) -> u16 {
let mut reg = CrcRegister::adts();
for r in regions {
r.feed(&mut reg, payload);
}
reg.value()
}
#[doc(hidden)]
pub fn collect_block_regions(
reader: &mut BitReader<'_>,
aot: u8,
fs: u8,
) -> Result<Vec<ProtectedRegion>> {
let mut regions = Vec::new();
loop {
let elem_start = reader.bit_position();
let Some(elem) = Walker::new(reader).next_element()? else {
return Ok(regions);
};
match elem {
Element::ChannelElement {
kind: IdSynEle::Sce | IdSynEle::Lfe,
..
} => {
let body = IcsBody::parse(reader, aot, fs, false)?;
let ics = body.ics_info.clone().ok_or(Error::ElementDecodeInvalid)?;
SpectralData::parse(reader, &ics, &body.section_data, fs)?;
regions.push(ProtectedRegion {
start_bit: elem_start + 3,
end_bit: reader.bit_position(),
pad_to: Some(192),
});
}
Element::ChannelElement {
kind: IdSynEle::Cpe,
..
} => {
let parsed = crate::decode::parse_cpe(reader, aot, fs)?;
let end = reader.bit_position();
regions.push(ProtectedRegion {
start_bit: elem_start + 3,
end_bit: end,
pad_to: Some(192),
});
regions.push(ProtectedRegion {
start_bit: parsed.second_ics_start_bit,
end_bit: end,
pad_to: Some(128),
});
}
Element::ChannelElement {
kind: IdSynEle::Cce,
element_instance_tag,
} => {
CouplingChannelElement::parse_after_tag(reader, element_instance_tag, aot, fs)?;
regions.push(ProtectedRegion {
start_bit: elem_start + 3,
end_bit: reader.bit_position(),
pad_to: Some(192),
});
}
Element::ChannelElement { .. } => return Err(Error::ElementDecodeInvalid),
Element::Data { .. } | Element::ProgramConfig(_) => {
regions.push(ProtectedRegion {
start_bit: elem_start + 3,
end_bit: reader.bit_position(),
pad_to: None,
});
}
Element::Fill { .. } => {}
Element::End => return Ok(regions),
}
}
}
pub fn protect_adts_frame(frame: &[u8]) -> Result<Vec<u8>> {
let (header, payload_offset) = crate::adts::AdtsHeader::parse(frame)?;
let frame_len = header.aac_frame_length as usize;
if frame_len < payload_offset || frame.len() < frame_len {
return Err(Error::UnexpectedEnd);
}
let frame = &frame[..frame_len];
if !header.protection_absent {
return Ok(frame.to_vec());
}
if header.number_of_raw_data_blocks_in_frame != 1 {
return Err(Error::NotImplemented);
}
let new_len = header.aac_frame_length + 2;
if new_len >= (1 << 13) {
return Err(Error::AdtsEncodeInvalid);
}
let mut h = [0u8; 7];
h.copy_from_slice(&frame[..7]);
h[1] &= 0xFE;
h[3] = (h[3] & 0xFC) | ((new_len >> 11) as u8 & 0x03);
h[4] = (new_len >> 3) as u8;
h[5] = (h[5] & 0x1F) | (((new_len & 0x07) as u8) << 5);
let payload = &frame[payload_offset..];
let mut reader = BitReader::new(payload);
let regions = collect_block_regions(
&mut reader,
header.audio_object_type(),
header.sampling_frequency_index,
)?;
let crc = adts_single_crc(&h, payload, ®ions);
let mut out = Vec::with_capacity(frame.len() + 2);
out.extend_from_slice(&h);
out.extend_from_slice(&crc.to_be_bytes());
out.extend_from_slice(payload);
Ok(out)
}
pub fn protect_adts_stream(data: &[u8]) -> Result<Vec<u8>> {
let mut out = Vec::with_capacity(data.len());
let mut pos = 0usize;
while pos + crate::adts::ADTS_HEADER_BYTES_NO_CRC <= data.len() {
let (header, _) = crate::adts::AdtsHeader::parse(&data[pos..])?;
let frame_len = header.aac_frame_length as usize;
if pos + frame_len > data.len() {
return Err(Error::UnexpectedEnd);
}
out.extend_from_slice(&protect_adts_frame(&data[pos..pos + frame_len])?);
pos += frame_len;
}
Ok(out)
}
#[cfg(test)]
mod tests {
use super::*;
fn reference(poly_low: u32, k: u32, init: u32, bits: &[bool]) -> u32 {
let full = u64::from(poly_low) | (1u64 << k);
let mut dividend: Vec<bool> = bits.to_vec();
dividend.extend(std::iter::repeat(false).take(k as usize));
for (i, d) in dividend.iter_mut().enumerate().take(k as usize) {
*d ^= (init >> (k as usize - 1 - i)) & 1 != 0;
}
let mut reg: u64 = 0;
let topbit = 1u64 << k;
for &b in ÷nd {
reg = (reg << 1) | u64::from(b);
if reg & topbit != 0 {
reg ^= full;
}
}
(reg & ((1u64 << k) - 1)) as u32
}
fn to_bits(bytes: &[u8]) -> Vec<bool> {
bytes
.iter()
.flat_map(|&b| (0..8).rev().map(move |i| (b >> i) & 1 != 0))
.collect()
}
#[test]
fn adts_register_matches_long_division_reference() {
for msg in [
&[][..],
&[0x00][..],
&[0xFF, 0xF1][..],
&[0x12, 0x34, 0x56, 0x78, 0x9A][..],
&[0xDE, 0xAD, 0xBE, 0xEF, 0x01, 0x02, 0x03][..],
] {
let bits = to_bits(msg);
let mut reg = CrcRegister::adts();
for &b in &bits {
reg.feed_bit(b);
}
assert_eq!(
u32::from(reg.value()),
reference(ADTS_CRC_POLY, 16, ADTS_CRC_INIT, &bits),
"message {msg:x?}"
);
}
}
#[test]
fn sbr_register_is_plain_remainder() {
for msg in [&[0x5Au8, 0x33][..], &[0xFF, 0x00, 0xAB, 0xCD][..]] {
let bits = to_bits(msg);
let mut reg = CrcRegister::sbr();
for &b in &bits {
reg.feed_bit(b);
}
assert_eq!(
u32::from(reg.value()),
reference(SBR_CRC_POLY, 10, 0, &bits),
"message {msg:x?}"
);
}
}
#[test]
fn sbr_poly_matches_crate_crc10_generator() {
assert_eq!(
u64::from(SBR_CRC_POLY),
crate::crc::CrcPoly::Crc10.generator()
);
assert_eq!(
u64::from(ADTS_CRC_POLY),
crate::crc::CrcPoly::Crc16.generator()
);
}
#[test]
fn empty_message_yields_init_for_adts() {
let reg = CrcRegister::adts();
assert_eq!(u32::from(reg.value()), ADTS_CRC_INIT);
assert_eq!(CrcRegister::sbr().value(), 0);
}
#[test]
fn appending_checksum_cancels_the_register() {
for msg in [&[0x53u8, 0x91, 0x2C][..], &[0xFF, 0xF9, 0x5C, 0x80][..]] {
let bits = to_bits(msg);
let mut reg = CrcRegister::adts();
for &b in &bits {
reg.feed_bit(b);
}
let crc = reg.value();
for i in (0..16).rev() {
reg.feed_bit((crc >> i) & 1 != 0);
}
assert_eq!(reg.value(), 0, "message {msg:x?}");
}
}
#[test]
fn region_pads_short_elements_with_zeros() {
let payload = [0xA5u8; 8];
let region = ProtectedRegion {
start_bit: 3,
end_bit: 43,
pad_to: Some(192),
};
let mut a = CrcRegister::adts();
region.feed(&mut a, &payload);
let mut b = CrcRegister::adts();
b.feed_bit_range(&payload, 3, 43);
b.feed_zeros(152);
assert_eq!(a.value(), b.value());
}
#[test]
fn region_caps_long_elements_at_pad_to() {
let payload = [0x3Cu8; 64];
let region = ProtectedRegion {
start_bit: 5,
end_bit: 305,
pad_to: Some(192),
};
let mut a = CrcRegister::adts();
region.feed(&mut a, &payload);
let mut b = CrcRegister::adts();
b.feed_bit_range(&payload, 5, 5 + 192);
assert_eq!(a.value(), b.value());
}
#[test]
fn header_crc_covers_positions() {
let header = [0xFFu8, 0xF1, 0x50, 0x80, 0x2F, 0xFF, 0xFC];
let a = adts_header_crc(&header, &[]);
let b = adts_header_crc(&header, &[0x1234]);
assert_ne!(a, b, "positions must alter the header CRC");
}
}