use oxideav_core::bits::{BitReader, BitWriter};
use crate::ics_info::WindowSequence;
use crate::swb_offset::FrameFamily;
use crate::{Error, Result};
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct TnsFilter {
pub length: u8,
pub order: u8,
pub direction: bool,
pub coef_compress: bool,
pub coef: Vec<u8>,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct TnsWindow {
pub coef_res: bool,
pub filters: Vec<TnsFilter>,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct TnsData {
pub windows: Vec<TnsWindow>,
}
pub const N_FILT_BITS_SHORT: u32 = 1;
pub const N_FILT_BITS_LONG: u32 = 2;
pub const LENGTH_BITS_SHORT: u32 = 4;
pub const LENGTH_BITS_LONG: u32 = 6;
pub const ORDER_BITS_SHORT: u32 = 3;
pub const ORDER_BITS_LONG: u32 = 5;
pub const COEF_RES_BITS: u32 = 1;
pub const DIRECTION_BITS: u32 = 1;
pub const COEF_COMPRESS_BITS: u32 = 1;
pub fn field_widths(seq: WindowSequence) -> (u32, u32, u32) {
if seq.is_eight_short() {
(N_FILT_BITS_SHORT, LENGTH_BITS_SHORT, ORDER_BITS_SHORT)
} else {
(N_FILT_BITS_LONG, LENGTH_BITS_LONG, ORDER_BITS_LONG)
}
}
pub fn field_widths_family(family: FrameFamily, seq: WindowSequence) -> (u32, u32, u32) {
if family.is_ld() {
(N_FILT_BITS_SHORT, LENGTH_BITS_SHORT, ORDER_BITS_SHORT)
} else {
field_widths(seq)
}
}
pub fn num_windows(seq: WindowSequence) -> usize {
if seq.is_eight_short() {
8
} else {
1
}
}
pub fn coef_bits(coef_res: bool, coef_compress: bool) -> u32 {
let coef_res_bits = 3 + u32::from(coef_res);
coef_res_bits - u32::from(coef_compress)
}
impl TnsData {
pub fn parse(reader: &mut BitReader<'_>, window_sequence: WindowSequence) -> Result<Self> {
Self::parse_family(reader, FrameFamily::Lc1024, window_sequence)
}
pub fn parse_family(
reader: &mut BitReader<'_>,
family: FrameFamily,
window_sequence: WindowSequence,
) -> Result<Self> {
Self::parse_widths(
reader,
field_widths_family(family, window_sequence),
window_sequence,
)
}
pub fn parse_widths(
reader: &mut BitReader<'_>,
widths: (u32, u32, u32),
window_sequence: WindowSequence,
) -> Result<Self> {
let (n_filt_bits, length_bits, order_bits) = widths;
if !widths_valid(widths) {
return Err(Error::TnsDataEncodeInvalid);
}
let nw = num_windows(window_sequence);
let mut windows = Vec::with_capacity(nw);
for _ in 0..nw {
let n_filt = read_u8(reader, n_filt_bits)?;
let coef_res = if n_filt > 0 {
reader.read_bit().map_err(|_| Error::UnexpectedEnd)?
} else {
false
};
let mut filters = Vec::with_capacity(n_filt as usize);
for _ in 0..n_filt {
let length = read_u8(reader, length_bits)?;
let order = read_u8(reader, order_bits)?;
let (direction, coef_compress, coef) = if order > 0 {
let direction = reader.read_bit().map_err(|_| Error::UnexpectedEnd)?;
let coef_compress = reader.read_bit().map_err(|_| Error::UnexpectedEnd)?;
let bits = coef_bits(coef_res, coef_compress);
let mut coef = Vec::with_capacity(order as usize);
for _ in 0..order {
coef.push(read_u8(reader, bits)?);
}
(direction, coef_compress, coef)
} else {
(false, false, Vec::new())
};
filters.push(TnsFilter {
length,
order,
direction,
coef_compress,
coef,
});
}
windows.push(TnsWindow { coef_res, filters });
}
Ok(TnsData { windows })
}
pub fn write(&self, writer: &mut BitWriter, window_sequence: WindowSequence) -> Result<()> {
self.write_family(writer, FrameFamily::Lc1024, window_sequence)
}
pub fn write_family(
&self,
writer: &mut BitWriter,
family: FrameFamily,
window_sequence: WindowSequence,
) -> Result<()> {
self.write_widths(
writer,
field_widths_family(family, window_sequence),
window_sequence,
)
}
pub fn write_widths(
&self,
writer: &mut BitWriter,
widths: (u32, u32, u32),
window_sequence: WindowSequence,
) -> Result<()> {
let (n_filt_bits, length_bits, order_bits) = widths;
if !widths_valid(widths) {
return Err(Error::TnsDataEncodeInvalid);
}
let nw = num_windows(window_sequence);
if self.windows.len() != nw {
return Err(Error::TnsDataEncodeInvalid);
}
let n_filt_max = (1u32 << n_filt_bits) - 1;
let length_max = (1u32 << length_bits) - 1;
let order_max = (1u32 << order_bits) - 1;
for w in &self.windows {
if (w.filters.len() as u32) > n_filt_max {
return Err(Error::TnsDataEncodeInvalid);
}
for f in &w.filters {
if u32::from(f.length) > length_max || u32::from(f.order) > order_max {
return Err(Error::TnsDataEncodeInvalid);
}
if f.order as usize != f.coef.len() {
return Err(Error::TnsDataEncodeInvalid);
}
if f.order == 0 && (f.direction || f.coef_compress) {
return Err(Error::TnsDataEncodeInvalid);
}
if f.order > 0 {
let bits = coef_bits(w.coef_res, f.coef_compress);
let coef_max = (1u32 << bits) - 1;
for c in &f.coef {
if u32::from(*c) > coef_max {
return Err(Error::TnsDataEncodeInvalid);
}
}
}
}
}
for w in &self.windows {
writer.write_u32(w.filters.len() as u32, n_filt_bits);
if !w.filters.is_empty() {
writer.write_bit(w.coef_res);
}
for f in &w.filters {
writer.write_u32(u32::from(f.length), length_bits);
writer.write_u32(u32::from(f.order), order_bits);
if f.order > 0 {
writer.write_bit(f.direction);
writer.write_bit(f.coef_compress);
let bits = coef_bits(w.coef_res, f.coef_compress);
for c in &f.coef {
writer.write_u32(u32::from(*c), bits);
}
}
}
}
Ok(())
}
}
fn widths_valid((n_filt_bits, length_bits, order_bits): (u32, u32, u32)) -> bool {
(1..=8).contains(&n_filt_bits)
&& (1..=8).contains(&length_bits)
&& (1..=8).contains(&order_bits)
}
fn read_u8(reader: &mut BitReader<'_>, n: u32) -> Result<u8> {
debug_assert!(n <= 8);
Ok(reader.read_u32(n).map_err(|_| Error::UnexpectedEnd)? as u8)
}