use rusty_h264_common::cabac_tables::{CTX_INIT, RANGE_LPS, STATE_TRANS};
const fn build_lps_range() -> [u8; 4 * 128] {
let mut t = [0u8; 4 * 128];
let mut q = 0;
while q < 4 {
let mut s = 0;
while s < 128 {
t[q * 128 + s] = RANGE_LPS[s >> 1][q];
s += 1;
}
q += 1;
}
t
}
const fn build_trans() -> [u8; 256] {
let mut t = [0u8; 256];
let mut s = 0;
while s < 128 {
let mps = s as u8 & 1;
t[s] = (STATE_TRANS[s >> 1][1] << 1) | mps;
let new_mps = if s >> 1 == 0 { 1 - mps } else { mps };
t[128 + s] = (STATE_TRANS[s >> 1][0] << 1) | new_mps;
s += 1;
}
t
}
static LPS_RANGE: [u8; 4 * 128] = build_lps_range();
static TRANS: [u8; 256] = build_trans();
pub struct Cabac<'a> {
data: &'a [u8],
byte_pos: usize,
window: u64,
wbits: u32,
range: u32,
offset: u32,
ctx: [u8; 460],
trace: bool,
sym: u64,
}
impl Cabac<'_> {
#[inline]
fn tr(&mut self, kind: &str) {
if self.trace {
eprintln!("{} {} r={} o={}", self.sym, kind, self.range, self.offset);
self.sym += 1;
}
}
}
impl<'a> Cabac<'a> {
pub fn new(data: &'a [u8], start_byte: usize, qp: i32, init_idc: u32, is_i: bool) -> Self {
let model = if is_i { 0 } else { ((init_idc + 1) as usize).min(3) };
let q = qp.clamp(0, 51);
let mut ctx = [0u8; 460];
for (i, c) in ctx.iter_mut().enumerate() {
let (m, n) = CTX_INIT[i][model];
let pre = (((m as i32 * q) >> 4) + n as i32).clamp(1, 126);
*c = if pre <= 63 {
((63 - pre) as u8) << 1
} else {
(((pre - 64) as u8) << 1) | 1
};
}
let trace = std::env::var_os("RH_CABAC_TRACE").is_some();
let mut e = Cabac { data, byte_pos: start_byte, window: 0, wbits: 0, range: 510, offset: 0, ctx, trace, sym: 0 };
e.offset = e.take(9);
e
}
pub fn dbg_state(&self) -> (u32, u32) {
(self.range, self.offset)
}
#[inline]
fn refill(&mut self) {
if let Some(chunk) = self.data.get(self.byte_pos..self.byte_pos + 8) {
let take_bytes = ((64 - self.wbits) / 8) as usize; let keep = (take_bytes * 8) as u32;
let v = u64::from_be_bytes(chunk.try_into().unwrap());
let v = if keep == 64 { v } else { v & (!0u64 << (64 - keep)) };
self.window |= v >> self.wbits;
self.byte_pos += take_bytes;
self.wbits += keep;
return;
}
while self.wbits <= 56 {
let b = self.data.get(self.byte_pos).copied().unwrap_or(0);
self.window |= (b as u64) << (56 - self.wbits);
self.byte_pos += 1;
self.wbits += 8;
}
}
#[inline(always)]
fn take(&mut self, n: u32) -> u32 {
if self.wbits < n {
self.refill();
}
let v = ((self.window >> (63 - n)) >> 1) as u32;
self.window <<= n;
self.wbits -= n;
v
}
#[inline(always)]
fn renorm(&mut self) {
let n = self.range.leading_zeros() - 23;
self.range <<= n;
self.offset = (self.offset << n) | self.take(n);
}
pub fn decode_decision(&mut self, ctx_idx: usize) -> u32 {
self.tr("D");
let s = (self.ctx[ctx_idx] & 127) as usize;
let q = ((self.range >> 6) & 3) as usize;
let lps = LPS_RANGE[q * 128 + s] as u32;
debug_assert!(self.range >= 256, "renorm invariant broken: range={}", self.range);
self.range -= lps;
let mask = ((self.range as i32 - self.offset as i32 - 1) >> 31) as u32;
self.offset -= self.range & mask;
self.range = self.range.wrapping_add(lps.wrapping_sub(self.range) & mask);
self.ctx[ctx_idx] = TRANS[s | (mask as usize & 128)];
let bin = (s as u32 ^ mask) & 1;
self.renorm();
bin
}
#[inline(always)]
pub fn decode_bypass(&mut self) -> u32 {
self.tr("B");
self.offset = (self.offset << 1) | self.take(1);
if self.offset >= self.range {
self.offset -= self.range;
1
} else {
0
}
}
#[allow(dead_code)] pub fn decode_bypass_bits(&mut self, n: u32) -> u32 {
let mut v = 0;
for _ in 0..n {
v = (v << 1) | self.decode_bypass();
}
v
}
pub fn decode_terminate(&mut self) -> bool {
self.tr("T");
self.range -= 2;
if self.offset >= self.range {
true
} else {
self.renorm();
false
}
}
}
#[cfg(test)]
mod tests {
use super::*;
struct Enc {
low: u32,
range: u32,
outstanding: u32,
first: bool,
bits: Vec<u8>,
ctx: Vec<(u8, u8)>, }
fn init_ctx(qp: i32, init_idc: u32, is_i: bool) -> Vec<(u8, u8)> {
let model = if is_i { 0 } else { ((init_idc + 1) as usize).min(3) };
let q = qp.clamp(0, 51);
(0..460)
.map(|i| {
let (m, n) = CTX_INIT[i][model];
let pre = (((m as i32 * q) >> 4) + n as i32).clamp(1, 126);
if pre <= 63 {
((63 - pre) as u8, 0)
} else {
((pre - 64) as u8, 1)
}
})
.collect()
}
impl Enc {
fn new(qp: i32, init_idc: u32, is_i: bool) -> Self {
Enc {
low: 0,
range: 510,
outstanding: 0,
first: true,
bits: Vec::new(),
ctx: init_ctx(qp, init_idc, is_i),
}
}
fn put_bit(&mut self, b: u32) {
if self.first {
self.first = false;
} else {
self.bits.push(b as u8);
}
while self.outstanding > 0 {
self.bits.push((1 - b) as u8);
self.outstanding -= 1;
}
}
fn renorm(&mut self) {
while self.range < 256 {
if self.low < 256 {
self.put_bit(0);
} else if self.low >= 512 {
self.low -= 512;
self.put_bit(1);
} else {
self.low -= 256;
self.outstanding += 1;
}
self.range <<= 1;
self.low <<= 1;
}
}
fn encode(&mut self, ctx_idx: usize, bin: u32) {
let (state, mps) = self.ctx[ctx_idx];
let q = ((self.range >> 6) & 3) as usize;
let lps = RANGE_LPS[state as usize][q] as u32;
self.range -= lps;
if bin != mps as u32 {
self.low += self.range;
self.range = lps;
let nm = if state == 0 { 1 - mps } else { mps };
self.ctx[ctx_idx] = (STATE_TRANS[state as usize][0], nm);
} else {
self.ctx[ctx_idx].0 = STATE_TRANS[state as usize][1];
}
self.renorm();
}
fn encode_bypass(&mut self, bin: u32) {
self.low <<= 1;
if bin != 0 {
self.low += self.range;
}
if self.low >= 1024 {
self.put_bit(1);
self.low -= 1024;
} else if self.low < 512 {
self.put_bit(0);
} else {
self.low -= 512;
self.outstanding += 1;
}
}
fn finish(&mut self) -> Vec<u8> {
self.range -= 2;
self.low += self.range;
self.range = 2;
self.renorm();
self.put_bit((self.low >> 9) & 1);
let v = ((self.low >> 7) & 3) | 1;
self.bits.push(((v >> 1) & 1) as u8);
self.bits.push((v & 1) as u8);
let mut out = vec![0u8; self.bits.len().div_ceil(8)];
for (i, &b) in self.bits.iter().enumerate() {
out[i / 8] |= b << (7 - (i % 8));
}
out
}
}
struct Rng(u32);
impl Rng {
fn next(&mut self) -> u32 {
self.0 ^= self.0 << 13;
self.0 ^= self.0 >> 17;
self.0 ^= self.0 << 5;
self.0
}
}
fn roundtrip(qp: i32, init_idc: u32, is_i: bool, seed: u32, n: usize) {
let mut rng = Rng(seed);
let mut script: Vec<(u8, usize, u32)> = Vec::with_capacity(n);
let mut enc = Enc::new(qp, init_idc, is_i);
for _ in 0..n {
let r = rng.next();
let kind = (r & 1) as u8;
let ctx = (r >> 1) as usize % 460;
let bin = (r >> 12) & 1;
script.push((kind, ctx, bin));
if kind == 0 {
enc.encode(ctx, bin);
} else {
enc.encode_bypass(bin);
}
}
let bytes = enc.finish();
let mut dec = Cabac::new(&bytes, 0, qp, init_idc, is_i);
for (i, &(kind, ctx, bin)) in script.iter().enumerate() {
let got = if kind == 0 {
dec.decode_decision(ctx)
} else {
dec.decode_bypass()
};
assert_eq!(got, bin, "bin {i} (kind {kind}, ctx {ctx}) mismatched");
}
assert!(dec.decode_terminate(), "terminate should signal end-of-stream");
}
#[test]
fn engine_roundtrip_many() {
for &qp in &[0, 12, 26, 37, 51] {
for &(idc, is_i) in &[(0u32, true), (0, false), (1, false), (2, false)] {
for seed in 1..=40u32 {
roundtrip(qp, idc, is_i, seed.wrapping_mul(2654435761), seed as usize * 53);
}
}
}
}
#[test]
fn engine_init_matches_spec() {
let dec = Cabac::new(&[0xFF, 0xFF, 0xFF], 0, 26, 0, true);
assert_eq!(dec.ctx[0] >> 1, 46, "state");
assert_eq!(dec.ctx[0] & 1, 0, "mps");
assert_eq!(dec.range, 510);
assert_eq!(dec.offset, 0x1FF);
}
#[test]
fn packed_state_tables_match_spec_form() {
for s in 0usize..128 {
let (state, mps) = ((s >> 1) as u8, (s & 1) as u8);
for q in 0usize..4 {
assert_eq!(LPS_RANGE[q * 128 + s], RANGE_LPS[state as usize][q], "lps s={s} q={q}");
}
let mps_t = TRANS[s];
assert_eq!(mps_t >> 1, STATE_TRANS[state as usize][1], "mps-trans state s={s}");
assert_eq!(mps_t & 1, mps, "mps-trans mps s={s}");
let lps_t = TRANS[128 + s];
let want_mps = if state == 0 { 1 - mps } else { mps };
assert_eq!(lps_t >> 1, STATE_TRANS[state as usize][0], "lps-trans state s={s}");
assert_eq!(lps_t & 1, want_mps, "lps-trans mps s={s}");
}
}
#[test]
fn branchless_mask_matches_conditional_form() {
for s in 0usize..128 {
for range in [256u32, 257, 300, 383, 384, 400, 448, 509, 510] {
for offset in [0u32, 1, 127, 128, 255, 256, 300, 383, 384, 509] {
if offset >= range {
continue;
}
let q = ((range >> 6) & 3) as usize;
let lps = LPS_RANGE[q * 128 + s] as u32;
let r1 = range.wrapping_sub(lps);
let (mut lr, mut lo, lbin, lctx) = if offset >= r1 {
(lps, offset - r1, (s as u32 & 1) ^ 1, TRANS[128 + s])
} else {
(r1, offset, s as u32 & 1, TRANS[s])
};
let mask = ((r1 as i32 - offset as i32 - 1) >> 31) as u32;
let bo = offset - (r1 & mask);
let br = r1.wrapping_add(lps.wrapping_sub(r1) & mask);
let bctx = TRANS[s | (mask as usize & 128)];
let bbin = (s as u32 ^ mask) & 1;
lr += 0;
lo += 0;
assert_eq!((lr, lo, lbin, lctx), (br, bo, bbin, bctx), "s={s} range={range} offset={offset}");
}
}
}
}
#[test]
fn tables_match_spec_boundaries() {
assert_eq!(RANGE_LPS[0], [128, 176, 208, 240]);
assert_eq!(RANGE_LPS[63], [2, 2, 2, 2]);
assert_eq!(STATE_TRANS[0], [0, 1]);
assert_eq!(STATE_TRANS[63], [63, 63]);
assert_eq!(CTX_INIT[0][0], (20, -15));
}
}