use rusty_esp_core::error::{Error, Result};
pub const STEP_TABLE: [i32; 89] = [
7, 8, 9, 10, 11, 12, 13, 14, 16, 17, 19, 21, 23, 25, 28, 31, 34, 37, 41, 45, 50, 55, 60, 66,
73, 80, 88, 97, 107, 118, 130, 143, 157, 173, 190, 209, 230, 253, 279, 307, 337, 371, 408, 449,
494, 544, 598, 658, 724, 796, 876, 963, 1060, 1166, 1282, 1411, 1552, 1707, 1878, 2066, 2272,
2499, 2749, 3024, 3327, 3660, 4026, 4428, 4871, 5358, 5894, 6484, 7132, 7845, 8630, 9493,
10442, 11487, 12635, 13899, 15289, 16818, 18500, 20350, 22385, 24623, 27086, 29794, 32767,
];
pub const INDEX_TABLE: [i8; 16] = [-1, -1, -1, -1, 2, 4, 6, 8, -1, -1, -1, -1, 2, 4, 6, 8];
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
pub enum PredictionRule {
#[default]
DecoderExact,
FfmpegEncoder,
}
#[derive(Debug, Clone, Copy, Default, PartialEq, Eq)]
pub struct ChannelState {
pub predictor: i32,
pub step_index: u8,
}
impl ChannelState {
#[inline]
#[must_use]
pub const fn diff_reference(step: i32, nibble: u8) -> i32 {
let mut d = step >> 3;
if nibble & 4 != 0 {
d += step;
}
if nibble & 2 != 0 {
d += step >> 1;
}
if nibble & 1 != 0 {
d += step >> 2;
}
d
}
#[inline]
#[must_use]
pub const fn diff_ffmpeg_encoder(step: i32, nibble: u8) -> i32 {
((2 * (nibble & 7) as i32 + 1) * step) >> 3
}
#[inline]
fn advance(&mut self, nibble: u8) {
let idx = i32::from(self.step_index) + i32::from(INDEX_TABLE[usize::from(nibble & 15)]);
self.step_index = idx.clamp(0, 88) as u8;
}
#[inline]
fn apply(&mut self, d: i32, negative: bool) {
self.predictor += if negative { -d } else { d };
self.predictor = self
.predictor
.clamp(i32::from(i16::MIN), i32::from(i16::MAX));
}
#[inline]
pub fn encode(&mut self, sample: i16, rule: PredictionRule) -> u8 {
let step = STEP_TABLE[usize::from(self.step_index)];
let delta = i32::from(sample) - self.predictor;
let mag = (delta.unsigned_abs() * 4 / step as u32).min(7) as u8;
let nibble = mag | if delta < 0 { 8 } else { 0 };
let d = match rule {
PredictionRule::DecoderExact => Self::diff_reference(step, nibble),
PredictionRule::FfmpegEncoder => Self::diff_ffmpeg_encoder(step, nibble),
};
self.apply(d, delta < 0);
self.advance(nibble);
nibble
}
#[inline]
pub fn decode(&mut self, nibble: u8) -> i16 {
let step = STEP_TABLE[usize::from(self.step_index)];
self.advance(nibble);
self.apply(Self::diff_reference(step, nibble), nibble & 8 != 0);
self.predictor as i16
}
}
pub const MAX_CHANNELS: usize = 2;
pub const fn block_align(channels: u8, frames_per_block: usize) -> usize {
4 * channels as usize + (frames_per_block - 1) * channels as usize / 2
}
pub const fn frames_per_block(channels: u8, block_align: usize) -> usize {
(block_align - 4 * channels as usize) * 2 / channels as usize + 1
}
pub const fn ffmpeg_block_align() -> usize {
1024
}
fn check(channels: u8, frames_per_block: usize) -> Result<()> {
if channels == 0 || channels as usize > MAX_CHANNELS {
return Err(Error::Unsupported);
}
if frames_per_block < 9 || (frames_per_block - 1) % 8 != 0 {
return Err(Error::InvalidFormat);
}
Ok(())
}
#[derive(Debug, Clone)]
pub struct Encoder {
channels: u8,
frames_per_block: usize,
rule: PredictionRule,
state: [ChannelState; MAX_CHANNELS],
pub blocks: u64,
}
impl Encoder {
pub fn new(channels: u8, frames_per_block: usize) -> Result<Self> {
Self::with_rule(channels, frames_per_block, PredictionRule::DecoderExact)
}
pub fn ffmpeg_compatible(channels: u8, frames_per_block: usize) -> Result<Self> {
Self::with_rule(channels, frames_per_block, PredictionRule::FfmpegEncoder)
}
pub fn with_rule(channels: u8, frames_per_block: usize, rule: PredictionRule) -> Result<Self> {
check(channels, frames_per_block)?;
Ok(Encoder {
channels,
frames_per_block,
rule,
state: [ChannelState::default(); MAX_CHANNELS],
blocks: 0,
})
}
#[must_use]
pub const fn frames_per_block(&self) -> usize {
self.frames_per_block
}
#[must_use]
pub const fn block_align(&self) -> usize {
block_align(self.channels, self.frames_per_block)
}
#[must_use]
pub const fn rule(&self) -> PredictionRule {
self.rule
}
pub fn encode_block(&mut self, pcm: &[u8], out: &mut [u8]) -> Result<usize> {
let ch = self.channels as usize;
if pcm.len() != self.frames_per_block * ch * 2 {
return Err(Error::InvalidGeometry);
}
let need = self.block_align();
if out.len() < need {
return Err(Error::BufferTooSmall { needed: need });
}
let sample = |frame: usize, c: usize| -> i16 {
let i = (frame * ch + c) * 2;
i16::from_le_bytes([pcm[i], pcm[i + 1]])
};
let mut w = 0usize;
for c in 0..ch {
let first = sample(0, c);
self.state[c].predictor = i32::from(first);
out[w..w + 2].copy_from_slice(&first.to_le_bytes());
out[w + 2] = self.state[c].step_index;
out[w + 3] = 0;
w += 4;
}
let groups = (self.frames_per_block - 1) / 8;
for g in 0..groups {
for c in 0..ch {
let base = 1 + g * 8;
for j in 0..4 {
let lo = self.state[c].encode(sample(base + j * 2, c), self.rule);
let hi = self.state[c].encode(sample(base + j * 2 + 1, c), self.rule);
out[w] = lo | (hi << 4);
w += 1;
}
}
}
self.blocks += 1;
Ok(w)
}
}
#[derive(Debug, Clone)]
pub struct Decoder {
channels: u8,
frames_per_block: usize,
pub blocks: u64,
}
impl Decoder {
pub fn new(channels: u8, frames_per_block: usize) -> Result<Self> {
check(channels, frames_per_block)?;
Ok(Decoder {
channels,
frames_per_block,
blocks: 0,
})
}
#[must_use]
pub const fn block_align(&self) -> usize {
block_align(self.channels, self.frames_per_block)
}
#[must_use]
pub const fn pcm_bytes_per_block(&self) -> usize {
self.frames_per_block * self.channels as usize * 2
}
pub fn decode_block(&mut self, block: &[u8], out: &mut [u8]) -> Result<usize> {
let ch = self.channels as usize;
if block.len() != self.block_align() {
return Err(Error::InvalidGeometry);
}
let need = self.pcm_bytes_per_block();
if out.len() < need {
return Err(Error::BufferTooSmall { needed: need });
}
let mut state = [ChannelState::default(); MAX_CHANNELS];
let mut r = 0usize;
let put = |out: &mut [u8], frame: usize, c: usize, v: i16| {
let i = (frame * ch + c) * 2;
out[i..i + 2].copy_from_slice(&v.to_le_bytes());
};
for (c, st) in state.iter_mut().enumerate().take(ch) {
let first = i16::from_le_bytes([block[r], block[r + 1]]);
let idx = block[r + 2];
if idx > 88 {
return Err(Error::Corrupt);
}
*st = ChannelState {
predictor: i32::from(first),
step_index: idx,
};
put(out, 0, c, first);
r += 4;
}
let groups = (self.frames_per_block - 1) / 8;
for g in 0..groups {
for (c, st) in state.iter_mut().enumerate().take(ch) {
let base = 1 + g * 8;
for j in 0..4 {
let byte = block[r];
r += 1;
let a = st.decode(byte & 0x0F);
let b = st.decode(byte >> 4);
put(out, base + j * 2, c, a);
put(out, base + j * 2 + 1, c, b);
}
}
}
self.blocks += 1;
Ok(need)
}
}
#[cfg(test)]
mod tests {
use super::*;
fn pcm_of(samples: &[i16]) -> Vec<u8> {
samples.iter().flat_map(|s| s.to_le_bytes()).collect()
}
fn i16s(bytes: &[u8]) -> Vec<i16> {
bytes
.chunks_exact(2)
.map(|b| i16::from_le_bytes([b[0], b[1]]))
.collect()
}
#[test]
fn sizes_match_the_wav_convention() {
assert_eq!(block_align(1, 2041), 1024);
assert_eq!(block_align(2, 2041), 2048);
assert_eq!(frames_per_block(1, 1024), 2041);
assert_eq!(frames_per_block(2, 2048), 2041);
assert_eq!(frames_per_block(2, ffmpeg_block_align()), 1017);
assert_eq!(block_align(2, 1017), 1024);
assert_eq!(Encoder::new(3, 2041).err(), Some(Error::Unsupported));
assert_eq!(Encoder::new(1, 2040).err(), Some(Error::InvalidFormat));
assert_eq!(Decoder::new(1, 8).err(), Some(Error::InvalidFormat));
}
#[test]
fn the_two_expansions_differ_by_truncation_only() {
assert_eq!(ChannelState::diff_reference(7, 1), 1);
assert_eq!(ChannelState::diff_ffmpeg_encoder(7, 1), 2);
for n in 0..16u8 {
assert_eq!(
ChannelState::diff_reference(1024, n),
ChannelState::diff_ffmpeg_encoder(1024, n)
);
}
for step in STEP_TABLE {
for n in 0..16u8 {
let d = (ChannelState::diff_reference(step, n)
- ChannelState::diff_ffmpeg_encoder(step, n))
.abs();
assert!(d <= 2, "step {step} nibble {n}: {d}");
}
}
}
#[test]
fn round_trip_tracks_a_ramp_closely() {
let mut enc = Encoder::new(1, 17).unwrap();
let mut dec = Decoder::new(1, 17).unwrap();
let samples: Vec<i16> = (0..17).map(|i| (i * 100 - 800) as i16).collect();
let pcm = pcm_of(&samples);
let mut blk = [0u8; 12];
assert_eq!(enc.encode_block(&pcm, &mut blk).unwrap(), 12);
assert_eq!(&blk[..2], &(-800i16).to_le_bytes());
assert_eq!(blk[2], 0);
let mut out = [0u8; 34];
assert_eq!(dec.decode_block(&blk, &mut out).unwrap(), 34);
let got = i16s(&out);
assert_eq!(got[0], -800);
for (i, (g, s)) in got.iter().zip(&samples).enumerate().skip(8) {
assert!(
(i32::from(*g) - i32::from(*s)).abs() <= 60,
"sample {i}: {g} vs {s}"
);
}
assert_eq!(enc.blocks, 1);
assert_eq!(dec.blocks, 1);
}
#[test]
fn closed_loop_encoder_predictor_equals_decoder_output() {
let mut enc = Encoder::new(1, 41).unwrap();
let mut dec = Decoder::new(1, 41).unwrap();
let samples: Vec<i16> = (0..41)
.map(|i| ((i * 977) % 20_000 - 10_000) as i16)
.collect();
let pcm = pcm_of(&samples);
let mut blk = [0u8; 24];
enc.encode_block(&pcm, &mut blk).unwrap();
let mut out = [0u8; 82];
dec.decode_block(&blk, &mut out).unwrap();
assert_eq!(enc.state[0].predictor, i32::from(i16s(&out)[40]));
}
#[test]
fn stereo_interleaves_by_word_and_step_index_carries() {
let mut enc = Encoder::new(2, 9).unwrap();
let mut dec = Decoder::new(2, 9).unwrap();
let mut frames = Vec::new();
for i in 0..9 {
frames.push(if i % 2 == 0 { 20_000 } else { -20_000 });
frames.push(0);
}
let pcm = pcm_of(&frames);
let mut blk = [0u8; 16];
assert_eq!(enc.encode_block(&pcm, &mut blk).unwrap(), 16);
assert_eq!(&blk[..4], &[0x20, 0x4E, 0, 0]);
assert_eq!(&blk[4..8], &[0, 0, 0, 0]);
assert_eq!(&blk[12..16], &[0, 0, 0, 0]);
let mut out = [0u8; 36];
dec.decode_block(&blk, &mut out).unwrap();
let got = i16s(&out);
assert!(got.iter().skip(1).step_by(2).all(|&r| r == 0));
let mut blk2 = [0u8; 16];
enc.encode_block(&pcm, &mut blk2).unwrap();
assert!(blk2[2] > 0, "left step index should have climbed");
assert_eq!(blk2[6], 0, "right stayed at 0 (nibble 0 is -1 clamped)");
}
#[test]
fn decoder_rejects_bad_headers_and_sizes() {
let mut dec = Decoder::new(1, 9).unwrap();
let mut out = [0u8; 18];
let mut blk = [0u8; 8];
blk[2] = 89;
assert_eq!(dec.decode_block(&blk, &mut out).err(), Some(Error::Corrupt));
assert_eq!(
dec.decode_block(&blk[..7], &mut out).err(),
Some(Error::InvalidGeometry)
);
blk[2] = 0;
assert_eq!(
dec.decode_block(&blk, &mut out[..10]).err(),
Some(Error::BufferTooSmall { needed: 18 })
);
}
}