use crate::celt_band_layout::CeltFrameSize;
use crate::celt_tf_adjust::{celt_tf_adjustment, celt_tf_select_can_affect, TfAdjustment};
use crate::range_decoder::RangeDecoder;
const LOGP_3_1_4: u32 = 2;
const LOGP_15_1_16: u32 = 4;
const LOGP_31_1_32: u32 = 5;
const LOGP_1_1_2: u32 = 1;
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct TfDecode {
pub tf_change: Vec<bool>,
pub tf_select: bool,
pub adjustments: Vec<TfAdjustment>,
}
impl TfDecode {
#[inline]
pub fn len(&self) -> usize {
self.tf_change.len()
}
#[inline]
pub fn is_empty(&self) -> bool {
self.tf_change.is_empty()
}
}
pub fn decode_tf(
rd: &mut RangeDecoder<'_>,
frame_size: CeltFrameSize,
transient: bool,
start_band: usize,
end_band: usize,
) -> TfDecode {
let n = end_band.saturating_sub(start_band);
let mut tf_change: Vec<bool> = Vec::with_capacity(n);
let mut prev = false;
for i in 0..n {
let choice = if i == 0 {
let logp = if transient { LOGP_3_1_4 } else { LOGP_15_1_16 };
rd.dec_bit_logp(logp) == 1
} else {
let logp = if transient {
LOGP_15_1_16
} else {
LOGP_31_1_32
};
let diff = rd.dec_bit_logp(logp) == 1;
prev ^ diff
};
tf_change.push(choice);
prev = choice;
}
let tf_select = if celt_tf_select_can_affect(frame_size, transient, &tf_change) {
rd.dec_bit_logp(LOGP_1_1_2) == 1
} else {
false
};
let adjustments = tf_change
.iter()
.map(|&c| celt_tf_adjustment(frame_size, transient, tf_select, c))
.collect();
TfDecode {
tf_change,
tf_select,
adjustments,
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::celt_tf_adjust::{
TF_ADJ_NONTRANSIENT_SELECT1, TF_ADJ_TRANSIENT_SELECT0, TF_ADJ_TRANSIENT_SELECT1,
};
struct RangeEncoder {
val: u32,
rng: u32,
out: Vec<u8>,
rem: i32,
ext: u32,
}
impl RangeEncoder {
fn new() -> Self {
Self {
val: 0,
rng: 0x8000_0000,
out: Vec::new(),
rem: -1,
ext: 0,
}
}
fn carry_out(&mut self, c: u32) {
if c == 0xFF {
self.ext += 1;
return;
}
let b = (c >> 8) & 1; if self.rem != -1 {
self.out.push((self.rem as u32 + b) as u8);
}
if self.ext > 0 {
let fill: u8 = if b == 1 { 0x00 } else { 0xFF };
for _ in 0..self.ext {
self.out.push(fill);
}
self.ext = 0;
}
self.rem = (c & 0xFF) as i32;
}
fn renorm(&mut self) {
while self.rng <= (1 << 23) {
self.carry_out(self.val >> 23);
self.val = (self.val << 8) & 0x7FFF_FFFF;
self.rng <<= 8;
}
}
fn bit_logp(&mut self, bit: u32, logp: u32) {
let r = self.rng;
let s = r >> logp; if bit == 0 {
self.rng = r - s;
} else {
self.val += r - s;
self.rng = s;
}
self.renorm();
}
fn finish(mut self) -> Vec<u8> {
let val = self.val;
let rng = self.rng;
let mut end = val;
for b in (0..=31u32).rev() {
let step = 1u32 << b;
let aligned = val.wrapping_add(step - 1) & !(step - 1);
let hi = u64::from(val) + u64::from(rng);
if u64::from(aligned) >= u64::from(val)
&& u64::from(aligned) + u64::from(step) - 1 < hi
{
end = aligned;
break;
}
}
while end != 0 {
self.carry_out(end >> 23);
end = (end << 8) & 0x7FFF_FFFF;
}
if self.rem != -1 || self.ext > 0 {
self.carry_out(0);
}
if self.rem != -1 {
self.out.push(self.rem as u8);
}
if self.ext > 0 {
for _ in 0..self.ext {
self.out.push(0xFF);
}
}
self.out
}
}
fn encode_bits(symbols: &[(u32, u32)]) -> Vec<u8> {
let mut enc = RangeEncoder::new();
for &(bit, logp) in symbols {
enc.bit_logp(bit, logp);
}
enc.finish()
}
#[test]
fn fixture_encoder_round_trips_through_real_decoder() {
let symbols = [
(1u32, 2u32),
(0, 4),
(1, 5),
(0, 1),
(1, 4),
(1, 1),
(0, 5),
(0, 2),
];
let buf = encode_bits(&symbols);
let mut rd = RangeDecoder::new(&buf);
for &(bit, logp) in &symbols {
assert_eq!(rd.dec_bit_logp(logp), bit, "logp={logp}");
}
assert!(!rd.has_error());
}
#[test]
fn empty_band_range_reads_nothing() {
let buf = encode_bits(&[(1, 2), (1, 2)]);
let mut rd = RangeDecoder::new(&buf);
let before = rd.tell_frac();
let res = decode_tf(&mut rd, CeltFrameSize::Ms20, true, 21, 21);
assert!(res.is_empty());
assert!(!res.tf_select);
assert!(res.adjustments.is_empty());
assert_eq!(rd.tell_frac(), before);
}
#[test]
fn ms2_5_nontransient_no_tf_select_all_choice0() {
let buf = encode_bits(&[(0, 4), (0, 5), (0, 5), (0, 5)]);
let mut rd = RangeDecoder::new(&buf);
let res = decode_tf(&mut rd, CeltFrameSize::Ms2_5, false, 17, 21);
assert_eq!(res.tf_change, vec![false, false, false, false]);
assert!(!res.tf_select);
assert_eq!(res.adjustments, vec![0, 0, 0, 0]);
}
#[test]
fn first_band_transient_choice1() {
let buf = encode_bits(&[(1, 2), (0, 1)]);
let mut rd = RangeDecoder::new(&buf);
let res = decode_tf(&mut rd, CeltFrameSize::Ms20, true, 20, 21);
assert_eq!(res.tf_change, vec![true]);
assert!(!res.tf_select);
assert_eq!(res.adjustments, vec![TF_ADJ_TRANSIENT_SELECT0[3][1]]);
}
#[test]
fn subsequent_bands_are_relative_toggles() {
let buf = encode_bits(&[(0, 4), (1, 5), (0, 5), (1, 5), (1, 1)]);
let mut rd = RangeDecoder::new(&buf);
let res = decode_tf(&mut rd, CeltFrameSize::Ms10, false, 17, 21);
assert_eq!(res.tf_change, vec![false, true, true, false]);
assert!(
res.tf_select,
"tf_select must be read (Tables 60/61 differ)"
);
assert_eq!(
res.adjustments,
vec![
TF_ADJ_NONTRANSIENT_SELECT1[2][0], TF_ADJ_NONTRANSIENT_SELECT1[2][1], TF_ADJ_NONTRANSIENT_SELECT1[2][1], TF_ADJ_NONTRANSIENT_SELECT1[2][0], ]
);
}
#[test]
fn tf_select_skipped_when_all_choice0_10ms() {
let buf = encode_bits(&[(0, 4), (0, 5), (0, 5), (0, 5)]);
let mut rd = RangeDecoder::new(&buf);
let res = decode_tf(&mut rd, CeltFrameSize::Ms10, false, 17, 21);
assert_eq!(res.tf_change, vec![false, false, false, false]);
assert!(!res.tf_select);
assert_eq!(res.adjustments, vec![0, 0, 0, 0]);
}
#[test]
fn ms20_transient_tf_select_always_read_nonempty() {
let buf = encode_bits(&[(0, 2), (1, 1)]);
let mut rd = RangeDecoder::new(&buf);
let res = decode_tf(&mut rd, CeltFrameSize::Ms20, true, 20, 21);
assert_eq!(res.tf_change, vec![false]);
assert!(res.tf_select);
assert_eq!(res.adjustments, vec![TF_ADJ_TRANSIENT_SELECT1[3][0]]);
}
#[test]
fn full_celt_only_5ms_transient_all_choice0() {
let mut symbols = vec![(0u32, 2u32)];
for _ in 0..20 {
symbols.push((0, 4));
}
let buf = encode_bits(&symbols);
let mut rd = RangeDecoder::new(&buf);
let res = decode_tf(&mut rd, CeltFrameSize::Ms5, true, 0, 21);
assert_eq!(res.len(), 21);
assert!(res.tf_change.iter().all(|&c| !c));
assert!(
!res.tf_select,
"5ms transient all-choice0 agrees → no tf_select"
);
assert!(res
.adjustments
.iter()
.all(|&a| a == TF_ADJ_TRANSIENT_SELECT0[1][0]));
}
#[test]
fn tf_select_one_routes_transient_to_table_63() {
let buf = encode_bits(&[(1, 2), (0, 4), (1, 1)]);
let mut rd = RangeDecoder::new(&buf);
let res = decode_tf(&mut rd, CeltFrameSize::Ms10, true, 19, 21);
assert_eq!(res.tf_change, vec![true, true]);
assert!(res.tf_select);
assert_eq!(
res.adjustments,
vec![
TF_ADJ_TRANSIENT_SELECT1[2][1],
TF_ADJ_TRANSIENT_SELECT1[2][1]
]
);
}
#[test]
fn len_and_is_empty_track_band_count() {
let buf = encode_bits(&[(0, 4), (0, 5), (0, 5), (0, 5)]);
let mut rd = RangeDecoder::new(&buf);
let res = decode_tf(&mut rd, CeltFrameSize::Ms2_5, false, 17, 21);
assert_eq!(res.len(), 4);
assert!(!res.is_empty());
let buf2 = encode_bits(&[(0, 1)]);
let mut rd2 = RangeDecoder::new(&buf2);
let empty = decode_tf(&mut rd2, CeltFrameSize::Ms20, false, 21, 21);
assert_eq!(empty.len(), 0);
assert!(empty.is_empty());
}
#[test]
fn adjustments_match_direct_table_lookup() {
let buf = encode_bits(&[(1, 2), (1, 4), (0, 4), (1, 4), (1, 1)]);
let mut rd = RangeDecoder::new(&buf);
let res = decode_tf(&mut rd, CeltFrameSize::Ms20, true, 17, 21);
for (i, &c) in res.tf_change.iter().enumerate() {
let expected = celt_tf_adjustment(CeltFrameSize::Ms20, true, res.tf_select, c);
assert_eq!(res.adjustments[i], expected, "band {i}");
}
}
}