use crate::error::{Error, Result};
const VERSION: u8 = 3;
const PARTITION: usize = 4096;
const ESCAPE_K: u32 = 31;
const ZERO_K: u32 = 30;
const MAX_K: u32 = 29;
const MAX_QUOTIENT: u32 = 48;
const HEADER_LEN: usize = 16;
const MAX_LPC_ORDER: usize = 12;
const COEF_PRECISION: u32 = 15;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum SampleFormat {
SignedInt,
UnsignedByte,
Float32,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct AudioFormat {
pub bits_per_sample: u16,
pub channels: u16,
pub data_start: u64,
pub block_align: u16,
pub sample_format: SampleFormat,
}
impl AudioFormat {
#[inline]
pub fn sample_bytes(&self) -> usize {
self.bits_per_sample as usize / 8
}
pub fn supported(&self) -> bool {
let width_ok = match self.sample_format {
SampleFormat::UnsignedByte => self.bits_per_sample == 8,
SampleFormat::SignedInt => matches!(self.bits_per_sample, 16 | 24 | 32),
SampleFormat::Float32 => self.bits_per_sample == 32,
};
width_ok
&& (self.channels == 1 || self.channels == 2)
&& self.block_align as usize == self.channels as usize * self.sample_bytes()
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum Decorrelation {
Independent = 0,
LeftSide = 1,
RightSide = 2,
}
impl Decorrelation {
fn from_u8(v: u8) -> Result<Self> {
Ok(match v {
0 => Decorrelation::Independent,
1 => Decorrelation::LeftSide,
2 => Decorrelation::RightSide,
other => return Err(Error::Compress(format!("unknown decorrelation {other}"))),
})
}
}
pub fn parse_wav_header(head: &[u8]) -> Option<AudioFormat> {
if head.len() < 44 || &head[0..4] != b"RIFF" || &head[8..12] != b"WAVE" {
return None;
}
let mut pos = 12usize;
let mut channels = 0u16;
let mut bits = 0u16;
let mut block_align = 0u16;
let mut sample_format = SampleFormat::SignedInt;
let mut seen_fmt = false;
while pos + 8 <= head.len() {
let id = &head[pos..pos + 4];
let size = u32::from_le_bytes(head[pos + 4..pos + 8].try_into().ok()?) as usize;
let body = pos + 8;
if id == b"fmt " {
if body + 16 > head.len() {
return None;
}
let audio_format = u16::from_le_bytes(head[body..body + 2].try_into().ok()?);
if !matches!(audio_format, 1 | 3 | 0xFFFE) {
return None;
}
channels = u16::from_le_bytes(head[body + 2..body + 4].try_into().ok()?);
block_align = u16::from_le_bytes(head[body + 12..body + 14].try_into().ok()?);
bits = u16::from_le_bytes(head[body + 14..body + 16].try_into().ok()?);
sample_format = match (audio_format, bits) {
(3, 32) => SampleFormat::Float32,
(_, 8) => SampleFormat::UnsignedByte,
_ => SampleFormat::SignedInt,
};
seen_fmt = true;
} else if id == b"data" {
if !seen_fmt {
return None;
}
let fmt = AudioFormat {
bits_per_sample: bits,
channels,
data_start: body as u64,
block_align,
sample_format,
};
return fmt.supported().then_some(fmt);
}
pos = body + size + (size & 1);
if size == 0 {
return None;
}
}
None
}
pub fn encode(
fmt: &AudioFormat,
chunk_offset: u64,
input: &[u8],
out: &mut Vec<u8>,
) -> Option<usize> {
if !fmt.supported() {
return None;
}
let align = fmt.block_align as usize;
let channels = fmt.channels as usize;
let data_start = fmt.data_start;
let body_start = chunk_offset.max(data_start);
if body_start.saturating_sub(chunk_offset) as usize >= input.len() {
return None;
}
let mut prefix_len = (body_start - chunk_offset) as usize;
let rel = body_start - data_start;
let pad = (align - (rel % align as u64) as usize) % align;
prefix_len += pad;
if prefix_len >= input.len() {
return None;
}
let usable = input.len() - prefix_len;
let n_frames = usable / align;
if n_frames < 256 {
return None;
}
let suffix_len = usable - n_frames * align;
if prefix_len > u16::MAX as usize || suffix_len > u16::MAX as usize {
return None;
}
let scale = match fmt.sample_format {
SampleFormat::Float32 => float_scale(&input[prefix_len..prefix_len + n_frames * align])?,
_ => 0,
};
let width = fmt.sample_bytes();
let body = &input[prefix_len..prefix_len + n_frames * align];
let mut ch: Vec<Vec<i32>> = vec![Vec::with_capacity(n_frames); channels];
for f in 0..n_frames {
let base = f * align;
for (c, dst) in ch.iter_mut().enumerate() {
dst.push(read_sample(
&body[base + c * width..],
width,
fmt.sample_format,
scale,
));
}
}
let max_abs = ch
.iter()
.flat_map(|c| c.iter())
.fold(0i64, |m, &v| m.max((v as i64).abs()));
let mut sample_bits = 1u32;
while sample_bits < 32 && max_abs >= (1i64 << (sample_bits - 1)) {
sample_bits += 1;
}
sample_bits = sample_bits.clamp(2, 32);
let mode = if channels == 2 && sample_bits < 32 {
choose_decorrelation(&ch[0], &ch[1])
} else {
Decorrelation::Independent
};
let coded: Vec<Vec<i32>> = match mode {
Decorrelation::Independent => ch,
Decorrelation::LeftSide => {
let side: Vec<i32> = ch[0].iter().zip(&ch[1]).map(|(l, r)| l - r).collect();
vec![std::mem::take(&mut ch[0]), side]
}
Decorrelation::RightSide => {
let side: Vec<i32> = ch[0].iter().zip(&ch[1]).map(|(l, r)| l - r).collect();
vec![std::mem::take(&mut ch[1]), side]
}
};
let start = out.len();
let format_code = match fmt.sample_format {
SampleFormat::SignedInt => 0u8,
SampleFormat::UnsignedByte => 1,
SampleFormat::Float32 => 2,
};
out.push(VERSION);
out.push(fmt.channels as u8);
out.push(fmt.bits_per_sample as u8);
out.push(mode as u8);
out.push(format_code);
out.push(scale as u8);
out.push(sample_bits as u8);
out.push(0); out.extend_from_slice(&(prefix_len as u16).to_le_bytes());
out.extend_from_slice(&(suffix_len as u16).to_le_bytes());
out.extend_from_slice(&(n_frames as u32).to_le_bytes());
out.extend_from_slice(&input[..prefix_len]);
out.extend_from_slice(&input[prefix_len + n_frames * align..]);
let verify_from = out.len();
let mut bits = BitWriter::new();
for (i, signal) in coded.iter().enumerate() {
let w = if i == 1 && mode != Decorrelation::Independent {
sample_bits + 1
} else {
sample_bits
};
encode_channel(signal, w, &mut bits);
}
bits.finish_into(out);
let _ = verify_from;
let mut check = Vec::with_capacity(input.len());
match decode(&out[start..], &mut check) {
Ok(()) if check == input => Some(out.len() - start),
_ => {
out.truncate(start);
debug_assert!(false, "pcm encoder produced output it could not decode");
tracing::warn!("pcm encode failed self-verification; falling back");
None
}
}
}
fn choose_decorrelation(left: &[i32], right: &[i32]) -> Decorrelation {
let cost = |signal: &[i32]| -> u64 {
signal
.windows(2)
.map(|w| (w[1] - w[0]).unsigned_abs() as u64)
.sum()
};
let side: Vec<i32> = left.iter().zip(right).map(|(l, r)| l - r).collect();
let (cl, cr, cs) = (cost(left), cost(right), cost(&side));
let independent = cl + cr;
let left_side = cl + cs;
let right_side = cr + cs;
if independent <= left_side && independent <= right_side {
Decorrelation::Independent
} else if left_side <= right_side {
Decorrelation::LeftSide
} else {
Decorrelation::RightSide
}
}
#[inline]
fn residual(order: usize, s: &[i32], i: usize) -> i64 {
let x = |k: usize| s[i - k] as i64;
match order {
0 => x(0),
1 => x(0) - x(1),
2 => x(0) - 2 * x(1) + x(2),
3 => x(0) - 3 * x(1) + 3 * x(2) - x(3),
_ => x(0) - 4 * x(1) + 6 * x(2) - 4 * x(3) + x(4),
}
}
fn encode_channel(signal: &[i32], width: u32, bits: &mut BitWriter) {
const SCORE_SAMPLE: usize = 8192;
let scored = signal.len().min(SCORE_SAMPLE);
let mut fixed_order = 0usize;
let mut fixed_cost = u64::MAX;
for order in 0..=4usize {
if scored <= order {
break;
}
let cost: u64 = (order..scored)
.map(|i| residual(order, signal, i).unsigned_abs())
.sum();
if cost < fixed_cost {
fixed_cost = cost;
fixed_order = order;
}
}
let lpc_order = choose_lpc_order(&signal[..scored], fixed_cost);
match lpc_order {
None => {
bits.write(0, 1);
bits.write(fixed_order as u32, 3);
write_warmup(signal, fixed_order, width, bits);
encode_partitions(signal, fixed_order, None, bits);
}
Some(order) => {
bits.write(1, 1);
bits.write(order as u32 - 1, 5);
bits.write(COEF_PRECISION - 1, 4);
write_warmup(signal, order, width, bits);
encode_partitions(signal, order, Some(()), bits);
}
}
}
fn write_warmup(signal: &[i32], order: usize, width: u32, bits: &mut BitWriter) {
let mask = if width >= 32 {
u32::MAX
} else {
(1u32 << width) - 1
};
for &s in signal.iter().take(order) {
bits.write(s as u32 & mask, width);
}
}
thread_local! {
static WINDOW: std::cell::RefCell<Vec<f64>> = const { std::cell::RefCell::new(Vec::new()) };
static SCRATCH: std::cell::RefCell<Vec<f64>> = const { std::cell::RefCell::new(Vec::new()) };
}
fn with_window<R>(n: usize, f: impl FnOnce(&[f64]) -> R) -> R {
WINDOW.with(|w| {
let mut w = w.borrow_mut();
if w.len() != n {
w.clear();
w.reserve(n);
let scale = std::f64::consts::TAU / (n.max(2) - 1) as f64;
for i in 0..n {
w.push(0.5 - 0.5 * (i as f64 * scale).cos());
}
}
f(&w)
})
}
fn levinson(sample: &[i32], max_order: usize, coefs: &mut Vec<f64>, errors: &mut Vec<f64>) -> bool {
let n = sample.len();
coefs.clear();
errors.clear();
if n <= max_order + 1 || max_order == 0 || max_order > MAX_LPC_ORDER {
return false;
}
let mut autoc = [0.0f64; MAX_LPC_ORDER + 1];
with_window(n, |w| {
SCRATCH.with(|sc| {
let mut buf = sc.borrow_mut();
buf.clear();
buf.extend(sample.iter().zip(w).map(|(&v, &wi)| v as f64 * wi));
for (lag, slot) in autoc.iter_mut().enumerate().take(max_order + 1) {
*slot = buf[lag..].iter().zip(buf.iter()).map(|(a, b)| a * b).sum();
}
});
});
if autoc[0] <= 0.0 || !autoc[0].is_finite() {
return false;
}
let mut err = autoc[0];
coefs.resize(max_order, 0.0);
for i in 0..max_order {
let mut acc = autoc[i + 1];
for j in 0..i {
acc -= coefs[j] * autoc[i - j];
}
let k = acc / err;
if !k.is_finite() {
coefs.truncate(i);
return i > 0;
}
coefs[i] = k;
for j in 0..i / 2 {
let tmp = coefs[j];
coefs[j] = tmp - k * coefs[i - 1 - j];
coefs[i - 1 - j] -= k * tmp;
}
if i % 2 == 1 {
coefs[i / 2] -= k * coefs[i / 2];
}
err *= 1.0 - k * k;
errors.push(err.max(f64::MIN_POSITIVE));
if err <= 0.0 {
coefs.truncate(i + 1);
while errors.len() < max_order {
errors.push(f64::MIN_POSITIVE);
}
return true;
}
}
true
}
fn choose_lpc_order(sample: &[i32], fixed_cost: u64) -> Option<usize> {
if sample.len() < 4 * MAX_LPC_ORDER {
return None;
}
let mut coefs = Vec::new();
let mut errors = Vec::new();
if !levinson(sample, MAX_LPC_ORDER, &mut coefs, &mut errors) {
return None;
}
let n = sample.len() as f64;
let mut best: Option<(usize, f64)> = None;
for (idx, &err) in errors.iter().enumerate() {
let order = idx + 1;
if err <= 0.0 || !err.is_finite() {
continue;
}
let bits_per = 0.5 * (err / n).max(1e-9).log2();
let overhead = order as f64 * COEF_PRECISION as f64 / PARTITION as f64;
let total = bits_per + overhead;
if best.map_or(true, |(_, b)| total < b) {
best = Some((order, total));
}
}
let (order, est_bits) = best?;
let fixed_bits = if fixed_cost == 0 {
0.0
} else {
(fixed_cost as f64 / n).max(1.0).log2() + 1.0
};
(est_bits + 0.02 < fixed_bits).then_some(order)
}
#[derive(Clone)]
struct Quantised {
coefs: Vec<i32>,
shift: u32,
}
fn quantise(coefs: &[f64]) -> Quantised {
let max = coefs.iter().fold(0.0f64, |m, c| m.max(c.abs()));
if max <= 0.0 || !max.is_finite() {
return Quantised {
coefs: vec![0; coefs.len()],
shift: 0,
};
}
let headroom = (COEF_PRECISION - 1) as i32;
let mut shift = headroom - (max.log2().floor() as i32) - 1;
shift = shift.clamp(0, 31);
let limit = 1i64 << (COEF_PRECISION - 1);
let mut error = 0.0f64;
let mut out = Vec::with_capacity(coefs.len());
for &c in coefs {
let scaled = c * (1u64 << shift) as f64 + error;
let q = scaled.round();
error = scaled - q;
out.push(q.clamp(-(limit as f64), (limit - 1) as f64) as i32);
}
Quantised {
coefs: out,
shift: shift as u32,
}
}
#[inline]
fn lpc_residual(signal: &[i32], i: usize, q: &Quantised) -> i64 {
let mut acc: i64 = 0;
for (j, &c) in q.coefs.iter().enumerate() {
acc += c as i64 * signal[i - 1 - j] as i64;
}
signal[i] as i64 - (acc >> q.shift)
}
fn encode_partitions(signal: &[i32], order: usize, lpc: Option<()>, bits: &mut BitWriter) {
if signal.len() <= order {
return;
}
let mut start = order;
while start < signal.len() {
let end = (start + PARTITION).min(signal.len());
let residuals: Vec<i64> = match lpc {
None => (start..end).map(|i| residual(order, signal, i)).collect(),
Some(()) => {
let from = start - order;
let mut c = Vec::new();
let mut e = Vec::new();
let q = if levinson(&signal[from..end], order, &mut c, &mut e) && c.len() == order {
quantise(&c)
} else {
Quantised {
coefs: vec![0; order],
shift: 0,
}
};
bits.write(q.shift, 5);
for &c in &q.coefs {
bits.write(c as u32 & ((1u32 << COEF_PRECISION) - 1), COEF_PRECISION);
}
(start..end).map(|i| lpc_residual(signal, i, &q)).collect()
}
};
encode_residual_partition(&residuals, bits);
start = end;
}
}
fn encode_residual_partition(part: &[i64], bits: &mut BitWriter) {
let zig: Vec<u64> = part.iter().map(|&r| zigzag(r)).collect();
if zig.iter().all(|&z| z == 0) {
bits.write(ZERO_K, 5);
return;
}
let k = choose_rice_k(&zig);
if k == ESCAPE_K {
bits.write(ESCAPE_K, 5);
for &z in &zig {
bits.write64(z, 40);
}
return;
}
bits.write(k, 5);
for &z in &zig {
let q = (z >> k) as u32;
bits.write_unary(q);
if k > 0 {
bits.write64(z & ((1u64 << k) - 1), k);
}
}
}
fn choose_rice_k(zig: &[u64]) -> u32 {
if zig.is_empty() {
return 0;
}
let sum: u64 = zig.iter().fold(0u64, |a, &z| a.saturating_add(z));
let mean = sum / zig.len() as u64;
let guess = (64 - mean.leading_zeros()).saturating_sub(1).min(MAX_K);
let mut best_k = guess;
let mut best_bits = u64::MAX;
for k in guess.saturating_sub(2)..=(guess + 2).min(MAX_K) {
let mut total = 0u64;
let mut blown = false;
for &z in zig {
let q = z >> k;
if q > MAX_QUOTIENT as u64 {
blown = true;
break;
}
total += q + 1 + k as u64;
}
if !blown && total < best_bits {
best_bits = total;
best_k = k;
}
}
if best_bits == u64::MAX {
return ESCAPE_K;
}
best_k
}
#[inline]
fn read_sample(b: &[u8], width: usize, format: SampleFormat, scale: u32) -> i32 {
match format {
SampleFormat::UnsignedByte => b[0] as i32 - 128,
SampleFormat::Float32 => {
let f = f32::from_le_bytes([b[0], b[1], b[2], b[3]]);
(f as f64 * (1u64 << scale) as f64) as i32
}
SampleFormat::SignedInt => match width {
2 => i16::from_le_bytes([b[0], b[1]]) as i32,
3 => i32::from_le_bytes([0, b[0], b[1], b[2]]) >> 8,
_ => i32::from_le_bytes([b[0], b[1], b[2], b[3]]),
},
}
}
#[inline]
fn write_sample(v: i32, width: usize, format: SampleFormat, scale: u32, out: &mut Vec<u8>) {
match format {
SampleFormat::UnsignedByte => out.push((v + 128) as u8),
SampleFormat::Float32 => {
let f = (v as f64 / (1u64 << scale) as f64) as f32;
out.extend_from_slice(&f.to_le_bytes());
}
SampleFormat::SignedInt => {
let b = v.to_le_bytes();
match width {
2 => out.extend_from_slice(&b[..2]),
3 => out.extend_from_slice(&b[..3]),
_ => out.extend_from_slice(&b),
}
}
}
}
fn float_scale(body: &[u8]) -> Option<u32> {
for scale in [15u32, 23, 24, 31] {
let mul = (1u64 << scale) as f64;
let ok = body.chunks_exact(4).all(|b| {
let f = f32::from_le_bytes([b[0], b[1], b[2], b[3]]) as f64;
if !f.is_finite() {
return false;
}
let v = f * mul;
v.fract() == 0.0 && v.abs() <= i32::MAX as f64
});
if ok {
return Some(scale);
}
}
None
}
#[inline]
fn sign_extend(v: u32, bits: u32) -> i32 {
if bits >= 32 {
return v as i32;
}
let shift = 32 - bits;
((v << shift) as i32) >> shift
}
#[inline]
fn zigzag(v: i64) -> u64 {
((v << 1) ^ (v >> 63)) as u64
}
#[inline]
fn unzigzag(z: u64) -> i64 {
((z >> 1) as i64) ^ -((z & 1) as i64)
}
pub fn decode(input: &[u8], out: &mut Vec<u8>) -> Result<()> {
if input.len() < HEADER_LEN {
return Err(Error::Compress(
"pcm chunk is shorter than its header".into(),
));
}
if input[0] != VERSION {
return Err(Error::Compress(format!(
"pcm version {} unsupported",
input[0]
)));
}
let channels = input[1] as usize;
let bits_per_sample = input[2] as u16;
let mode = Decorrelation::from_u8(input[3])?;
let sample_format = match input[4] {
0 => SampleFormat::SignedInt,
1 => SampleFormat::UnsignedByte,
2 => SampleFormat::Float32,
_ => {
return Err(Error::Compress(
"pcm chunk declares an unknown format".into(),
))
}
};
let scale = input[5] as u32;
let sample_bits = input[6] as u32;
let prefix_len = u16::from_le_bytes([input[8], input[9]]) as usize;
let suffix_len = u16::from_le_bytes([input[10], input[11]]) as usize;
let n_frames = u32::from_le_bytes(input[12..16].try_into().unwrap()) as usize;
if !(2..=32).contains(&sample_bits) || scale > 40 {
return Err(Error::Compress(
"pcm chunk declares an impossible width".into(),
));
}
let layout_ok = match sample_format {
SampleFormat::UnsignedByte => bits_per_sample == 8,
SampleFormat::SignedInt => matches!(bits_per_sample, 16 | 24 | 32),
SampleFormat::Float32 => bits_per_sample == 32,
};
if !layout_ok || !(1..=2).contains(&channels) {
return Err(Error::Compress(
"pcm chunk declares an unsupported layout".into(),
));
}
let width = bits_per_sample as usize / 8;
let raw_end = HEADER_LEN
.checked_add(prefix_len)
.and_then(|v| v.checked_add(suffix_len))
.ok_or_else(|| Error::Compress("pcm chunk lengths overflow".into()))?;
if raw_end > input.len() {
return Err(Error::Compress("pcm chunk is truncated".into()));
}
if n_frames > 1 << 28 {
return Err(Error::Compress("pcm chunk declares too many frames".into()));
}
let prefix = &input[HEADER_LEN..HEADER_LEN + prefix_len];
let suffix = &input[HEADER_LEN + prefix_len..raw_end];
let mut bits = BitReader::new(&input[raw_end..]);
let mut coded: Vec<Vec<i32>> = Vec::with_capacity(channels);
for i in 0..channels {
let w = if i == 1 && mode != Decorrelation::Independent {
sample_bits + 1
} else {
sample_bits
};
coded.push(decode_channel(n_frames, w, &mut bits)?);
}
let channels_out: Vec<Vec<i32>> = match mode {
Decorrelation::Independent => coded,
Decorrelation::LeftSide => {
let left = &coded[0];
let side = &coded[1];
let right: Vec<i32> = left.iter().zip(side).map(|(l, s)| l - s).collect();
vec![coded[0].clone(), right]
}
Decorrelation::RightSide => {
let right = &coded[0];
let side = &coded[1];
let left: Vec<i32> = right.iter().zip(side).map(|(r, s)| r + s).collect();
vec![left, coded[0].clone()]
}
};
out.extend_from_slice(prefix);
for f in 0..n_frames {
for c in channels_out.iter() {
write_sample(c[f], width, sample_format, scale, out);
}
}
out.extend_from_slice(suffix);
Ok(())
}
fn decode_channel(n_frames: usize, width: u32, bits: &mut BitReader) -> Result<Vec<i32>> {
let is_lpc = bits.read(1)? == 1;
let (order, precision) = if is_lpc {
let order = bits.read(5)? as usize + 1;
let precision = bits.read(4)? + 1;
if precision > 32 {
return Err(Error::Compress("pcm coefficient precision too wide".into()));
}
(order, precision)
} else {
let order = bits.read(3)? as usize;
if order > 4 {
return Err(Error::Compress("pcm predictor order out of range".into()));
}
(order, 0)
};
let mut signal: Vec<i32> = Vec::with_capacity(n_frames);
for _ in 0..order.min(n_frames) {
signal.push(sign_extend(bits.read(width)?, width));
}
if n_frames <= order {
return Ok(signal);
}
let mut remaining = n_frames - order;
while remaining > 0 {
let count = remaining.min(PARTITION);
let quant = if is_lpc {
let shift = bits.read(5)?;
let mut coefs = Vec::with_capacity(order);
for _ in 0..order {
coefs.push(sign_extend(bits.read(precision)?, precision));
}
Some(Quantised { coefs, shift })
} else {
None
};
let k = bits.read(5)?;
for _ in 0..count {
let z = if k == ZERO_K {
0
} else if k == ESCAPE_K {
bits.read64(40)?
} else {
let q = bits.read_unary(MAX_QUOTIENT)? as u64;
let low = if k > 0 { bits.read64(k)? } else { 0 };
(q << k) | low
};
let r = unzigzag(z);
let i = signal.len();
let value = match &quant {
Some(q) => {
let mut acc: i64 = 0;
for (j, &c) in q.coefs.iter().enumerate() {
acc += c as i64 * signal[i - 1 - j] as i64;
}
r + (acc >> q.shift)
}
None => {
let x = |back: usize| signal[i - back] as i64;
match order {
0 => r,
1 => r + x(1),
2 => r + 2 * x(1) - x(2),
3 => r + 3 * x(1) - 3 * x(2) + x(3),
_ => r + 4 * x(1) - 6 * x(2) + 4 * x(3) - x(4),
}
}
};
signal.push(value as i32);
}
remaining -= count;
}
Ok(signal)
}
struct BitWriter {
out: Vec<u8>,
acc: u64,
nbits: u32,
}
impl BitWriter {
fn new() -> Self {
Self {
out: Vec::new(),
acc: 0,
nbits: 0,
}
}
#[inline]
fn write(&mut self, value: u32, bits: u32) {
self.write64(value as u64, bits);
}
#[inline]
fn write64(&mut self, value: u64, bits: u32) {
debug_assert!(bits <= 56);
let masked = if bits >= 64 {
value
} else {
value & ((1u64 << bits) - 1)
};
self.acc = (self.acc << bits) | masked;
self.nbits += bits;
while self.nbits >= 8 {
self.nbits -= 8;
self.out.push((self.acc >> self.nbits) as u8);
}
}
#[inline]
fn write_unary(&mut self, q: u32) {
let mut left = q;
while left >= 32 {
self.write64(0, 32);
left -= 32;
}
if left > 0 {
self.write64(0, left);
}
self.write64(1, 1);
}
fn finish_into(mut self, out: &mut Vec<u8>) {
if self.nbits > 0 {
let pad = 8 - self.nbits;
self.acc <<= pad;
self.out.push(self.acc as u8);
}
out.extend_from_slice(&self.out);
}
}
struct BitReader<'a> {
data: &'a [u8],
pos: usize,
acc: u64,
nbits: u32,
}
impl<'a> BitReader<'a> {
fn new(data: &'a [u8]) -> Self {
Self {
data,
pos: 0,
acc: 0,
nbits: 0,
}
}
#[inline]
fn fill(&mut self) {
while self.nbits <= 56 && self.pos < self.data.len() {
self.acc = (self.acc << 8) | self.data[self.pos] as u64;
self.pos += 1;
self.nbits += 8;
}
}
#[inline]
fn read(&mut self, bits: u32) -> Result<u32> {
Ok(self.read64(bits)? as u32)
}
#[inline]
fn read64(&mut self, bits: u32) -> Result<u64> {
if bits == 0 {
return Ok(0);
}
self.fill();
if self.nbits < bits {
return Err(Error::Compress("pcm bitstream ended early".into()));
}
self.nbits -= bits;
let value = (self.acc >> self.nbits) & ((1u64 << bits) - 1);
Ok(value)
}
#[inline]
fn read_unary(&mut self, limit: u32) -> Result<u32> {
let mut count = 0u32;
loop {
if self.read64(1)? == 1 {
return Ok(count);
}
count += 1;
if count > limit {
return Err(Error::Compress("pcm unary run exceeds its limit".into()));
}
}
}
}
#[cfg(test)]
mod tests {
use super::*;
fn wav_header(data_len: u32, channels: u16) -> Vec<u8> {
let block_align = channels * 2;
let byte_rate = 44_100 * block_align as u32;
let mut h = Vec::new();
h.extend_from_slice(b"RIFF");
h.extend_from_slice(&(36 + data_len).to_le_bytes());
h.extend_from_slice(b"WAVE");
h.extend_from_slice(b"fmt ");
h.extend_from_slice(&16u32.to_le_bytes());
h.extend_from_slice(&1u16.to_le_bytes());
h.extend_from_slice(&channels.to_le_bytes());
h.extend_from_slice(&44_100u32.to_le_bytes());
h.extend_from_slice(&byte_rate.to_le_bytes());
h.extend_from_slice(&block_align.to_le_bytes());
h.extend_from_slice(&16u16.to_le_bytes());
h.extend_from_slice(b"data");
h.extend_from_slice(&data_len.to_le_bytes());
h
}
fn tone(frames: usize, channels: usize, noise_shift: u32) -> Vec<u8> {
let mut out = Vec::with_capacity(frames * channels * 2);
let mut s = 0x1234_5678_9ABC_DEF0u64;
for i in 0..frames {
let t = i as f64 / 44_100.0;
s ^= s << 13;
s ^= s >> 7;
s ^= s << 17;
let dither = if noise_shift >= 63 {
0
} else {
((s >> noise_shift) as i16) / 4
};
for c in 0..channels {
let f = if c == 0 { 440.0 } else { 659.25 };
let v = ((t * f * std::f64::consts::TAU).sin() * 11_000.0) as i16;
out.extend_from_slice(&v.wrapping_add(dither).to_le_bytes());
}
}
out
}
fn roundtrip(fmt: &AudioFormat, offset: u64, input: &[u8]) -> Option<usize> {
let mut enc = Vec::new();
let n = encode(fmt, offset, input, &mut enc)?;
assert_eq!(n, enc.len());
let mut dec = Vec::new();
decode(&enc, &mut dec).expect("decode");
assert_eq!(dec.len(), input.len(), "length changed");
assert!(dec == input, "codec is not lossless");
Some(enc.len())
}
#[test]
fn parses_a_canonical_wav_header() {
let h = wav_header(1000, 2);
let fmt = parse_wav_header(&h).expect("should parse");
assert_eq!(fmt.channels, 2);
assert_eq!(fmt.bits_per_sample, 16);
assert_eq!(fmt.block_align, 4);
assert_eq!(fmt.data_start, 44);
assert!(fmt.supported());
}
#[test]
fn parses_a_header_with_extra_chunks_before_data() {
let mut h = Vec::new();
h.extend_from_slice(b"RIFF");
h.extend_from_slice(&2000u32.to_le_bytes());
h.extend_from_slice(b"WAVE");
h.extend_from_slice(b"fmt ");
h.extend_from_slice(&16u32.to_le_bytes());
h.extend_from_slice(&1u16.to_le_bytes());
h.extend_from_slice(&2u16.to_le_bytes());
h.extend_from_slice(&44_100u32.to_le_bytes());
h.extend_from_slice(&176_400u32.to_le_bytes());
h.extend_from_slice(&4u16.to_le_bytes());
h.extend_from_slice(&16u16.to_le_bytes());
h.extend_from_slice(b"LIST");
h.extend_from_slice(&5u32.to_le_bytes());
h.extend_from_slice(b"INFOx");
h.push(0);
h.extend_from_slice(b"data");
h.extend_from_slice(&1000u32.to_le_bytes());
let fmt = parse_wav_header(&h).expect("should walk past LIST");
assert_eq!(fmt.data_start as usize, h.len());
}
#[test]
fn rejects_non_wav_and_unsupported_layouts() {
assert!(parse_wav_header(b"not a wav file at all, really truly not").is_none());
assert!(parse_wav_header(&[]).is_none());
let mut h = wav_header(1000, 2);
h[34] = 24;
h[32] = 6;
let fmt = parse_wav_header(&h).expect("24-bit is supported");
assert_eq!(fmt.bits_per_sample, 24);
assert_eq!(fmt.sample_bytes(), 3);
let mut h8 = wav_header(1000, 2);
h8[34] = 8;
h8[32] = 2;
let f8 = parse_wav_header(&h8).expect("8-bit is supported");
assert_eq!(f8.sample_format, SampleFormat::UnsignedByte);
let mut h32 = wav_header(1000, 2);
h32[34] = 32;
h32[32] = 8;
let f32i = parse_wav_header(&h32).expect("32-bit int is supported");
assert_eq!(f32i.sample_format, SampleFormat::SignedInt);
let mut hf = wav_header(1000, 2);
hf[34] = 32;
hf[32] = 8;
hf[20] = 3; let ff = parse_wav_header(&hf).expect("float is supported");
assert_eq!(ff.sample_format, SampleFormat::Float32);
let mut h12 = wav_header(1000, 2);
h12[34] = 12;
h12[32] = 3;
assert!(parse_wav_header(&h12).is_none());
}
#[test]
fn stereo_roundtrips_and_beats_zstd() {
let fmt = AudioFormat {
bits_per_sample: 16,
channels: 2,
data_start: 0,
block_align: 4,
sample_format: SampleFormat::SignedInt,
};
let pcm = tone(262_144, 2, 56);
let size = roundtrip(&fmt, 0, &pcm).expect("should encode");
let ratio = pcm.len() as f64 / size as f64;
assert!(ratio > 1.3, "ratio only {ratio:.2}x");
}
#[test]
fn twenty_four_bit_roundtrips_and_compresses() {
let fmt = AudioFormat {
bits_per_sample: 24,
channels: 2,
data_start: 0,
block_align: 6,
sample_format: SampleFormat::SignedInt,
};
let frames = 150_000;
let mut pcm = Vec::with_capacity(frames * 6);
let mut s = 0x2545_F491_4F6C_DD1Du64;
for i in 0..frames {
let t = i as f64 / 44_100.0;
s ^= s << 13;
s ^= s >> 7;
s ^= s << 17;
let d = ((s >> 52) as i32 & 0x7FF) - 1024;
for (f, amp) in [(440.0, 2_800_000.0), (659.25, 2_100_000.0)] {
let v = ((t * f * std::f64::consts::TAU).sin() * amp) as i32 + d;
pcm.extend_from_slice(&v.to_le_bytes()[..3]);
}
}
let n = roundtrip(&fmt, 0, &pcm).expect("should encode 24-bit");
let ratio = pcm.len() as f64 / n as f64;
assert!(ratio > 1.3, "24-bit ratio only {ratio:.2}x");
}
#[test]
fn twenty_four_bit_extremes_survive() {
let fmt = AudioFormat {
bits_per_sample: 24,
channels: 1,
data_start: 0,
block_align: 3,
sample_format: SampleFormat::SignedInt,
};
let mut pcm = Vec::new();
for i in 0..60_000i32 {
let v = match i % 4 {
0 => -8_388_608,
1 => 8_388_607,
2 => 0,
_ => (i * 977) % 8_388_608 - 4_194_304,
};
pcm.extend_from_slice(&v.to_le_bytes()[..3]);
}
roundtrip(&fmt, 0, &pcm).expect("24-bit extremes must roundtrip");
}
#[test]
fn self_verification_catches_a_bad_encode() {
let fmt = AudioFormat {
bits_per_sample: 16,
channels: 2,
data_start: 0,
block_align: 4,
sample_format: SampleFormat::SignedInt,
};
let pcm = tone(80_000, 2, 56);
let mut enc = Vec::new();
encode(&fmt, 0, &pcm, &mut enc).expect("encode");
let mut broken = enc.clone();
let at = HEADER_LEN + (broken.len() - HEADER_LEN) / 2;
broken[at] ^= 0b0010_0000;
let mut out = Vec::new();
let matched = decode(&broken, &mut out).is_ok() && out == pcm;
assert!(!matched, "a corrupted stream decoded as the original");
}
#[test]
fn every_accepted_encode_is_verified_lossless() {
for (bits, channels, align) in [(16u16, 2u16, 4u16), (16, 1, 2), (24, 2, 6), (24, 1, 3)] {
let fmt = AudioFormat {
bits_per_sample: bits,
channels,
data_start: 0,
block_align: align,
sample_format: SampleFormat::SignedInt,
};
for noise in [63u32, 56, 48, 40] {
let frames = 40_000;
let mut pcm = Vec::new();
let mut s = 0xDEAD_BEEF_CAFE_F00Du64 ^ noise as u64;
for i in 0..frames {
let t = i as f64 / 44_100.0;
s ^= s << 13;
s ^= s >> 7;
s ^= s << 17;
let amp = if bits == 16 { 11_000.0 } else { 2_800_000.0 };
let d = if noise >= 63 {
0
} else {
((s >> noise) as i32) % 512
};
for c in 0..channels {
let f = if c == 0 { 440.0 } else { 659.25 };
let v = ((t * f * std::f64::consts::TAU).sin() * amp) as i32 + d;
let b = v.to_le_bytes();
pcm.extend_from_slice(&b[..(bits / 8) as usize]);
}
}
roundtrip(&fmt, 0, &pcm)
.unwrap_or_else(|| panic!("{bits}-bit {channels}ch noise={noise} refused"));
}
}
}
#[test]
fn eight_bit_unsigned_roundtrips() {
let fmt = AudioFormat {
bits_per_sample: 8,
channels: 2,
data_start: 0,
block_align: 2,
sample_format: SampleFormat::UnsignedByte,
};
let mut pcm = Vec::new();
for i in 0..80_000 {
let t = i as f64 / 8_000.0;
for f in [440.0, 659.25] {
let v = ((t * f * std::f64::consts::TAU).sin() * 100.0) as i32 + 128;
pcm.push(v.clamp(0, 255) as u8);
}
}
pcm.extend_from_slice(&[0, 255, 0, 255, 128, 128]);
roundtrip(&fmt, 0, &pcm).expect("8-bit must roundtrip");
}
#[test]
fn thirty_two_bit_integer_roundtrips() {
let fmt = AudioFormat {
bits_per_sample: 32,
channels: 2,
data_start: 0,
block_align: 8,
sample_format: SampleFormat::SignedInt,
};
let mut pcm = Vec::new();
for i in 0..60_000 {
let t = i as f64 / 44_100.0;
for f in [440.0, 659.25] {
let v = ((t * f * std::f64::consts::TAU).sin() * 700_000_000.0) as i32;
pcm.extend_from_slice(&v.to_le_bytes());
}
}
for v in [i32::MIN, i32::MAX, 0, -1] {
pcm.extend_from_slice(&v.to_le_bytes());
pcm.extend_from_slice(&v.to_le_bytes());
}
roundtrip(&fmt, 0, &pcm).expect("32-bit int must roundtrip");
}
#[test]
fn integer_valued_float_roundtrips_bit_exactly() {
let fmt = AudioFormat {
bits_per_sample: 32,
channels: 2,
data_start: 0,
block_align: 8,
sample_format: SampleFormat::Float32,
};
let mut pcm = Vec::new();
for i in 0..60_000 {
let t = i as f64 / 44_100.0;
for f in [440.0, 659.25] {
let q = ((t * f * std::f64::consts::TAU).sin() * 6_000_000.0) as i32;
let v = q as f32 / (1i32 << 23) as f32;
pcm.extend_from_slice(&v.to_le_bytes());
}
}
let n = roundtrip(&fmt, 0, &pcm).expect("integer-valued float must encode");
assert!(pcm.len() > n, "float should have compressed");
}
#[test]
fn fractional_float_is_declined_not_rounded() {
let fmt = AudioFormat {
bits_per_sample: 32,
channels: 2,
data_start: 0,
block_align: 8,
sample_format: SampleFormat::Float32,
};
let mut pcm = Vec::new();
let mut s = 0x1234_5678_9ABC_DEF0u64;
for _ in 0..40_000 {
for _ in 0..2 {
s ^= s << 13;
s ^= s >> 7;
s ^= s << 17;
let v = f32::from_bits(((s >> 32) as u32 & 0x7FFF_FFFF) | 0x3000_0000);
pcm.extend_from_slice(&v.to_le_bytes());
}
}
let mut out = Vec::new();
let r = encode(&fmt, 0, &pcm, &mut out);
if r.is_some() {
let mut back = Vec::new();
decode(&out, &mut back).expect("decode");
assert_eq!(back, pcm, "float encode was not bit-exact");
} else {
assert!(out.is_empty(), "a refusal must not write anything");
}
}
#[test]
fn mono_roundtrips() {
let fmt = AudioFormat {
bits_per_sample: 16,
channels: 1,
data_start: 0,
block_align: 2,
sample_format: SampleFormat::SignedInt,
};
let pcm = tone(100_000, 1, 56);
let size = roundtrip(&fmt, 0, &pcm).expect("should encode");
assert!(pcm.len() > size);
}
#[test]
fn survives_a_chunk_that_starts_mid_frame() {
let fmt = AudioFormat {
bits_per_sample: 16,
channels: 2,
data_start: 44,
block_align: 4,
sample_format: SampleFormat::SignedInt,
};
let pcm = tone(200_000, 2, 56);
for offset in [0u64, 44, 45, 46, 47, 48, 1000, 1001, 1002, 1003] {
let body: Vec<u8> = if offset < 44 {
let mut v = wav_header(pcm.len() as u32, 2);
v.extend_from_slice(&pcm);
v[offset as usize..].to_vec()
} else {
let skip = (offset - 44) as usize;
pcm[skip..].to_vec()
};
roundtrip(&fmt, offset, &body).unwrap_or_else(|| panic!("offset {offset}"));
}
}
#[test]
fn handles_silence_and_full_scale() {
let fmt = AudioFormat {
bits_per_sample: 16,
channels: 2,
data_start: 0,
block_align: 4,
sample_format: SampleFormat::SignedInt,
};
let silence = vec![0u8; 4 * 100_000];
let n = roundtrip(&fmt, 0, &silence).expect("encode silence");
assert!(
silence.len() as f64 / n as f64 > 50.0,
"silence only reached {:.0}x",
silence.len() as f64 / n as f64
);
let mut extreme = Vec::new();
for i in 0..100_000 {
let (l, r) = if i % 2 == 0 {
(i16::MIN, i16::MAX)
} else {
(i16::MAX, i16::MIN)
};
extreme.extend_from_slice(&l.to_le_bytes());
extreme.extend_from_slice(&r.to_le_bytes());
}
roundtrip(&fmt, 0, &extreme).expect("encode extremes");
}
#[test]
fn random_bytes_still_roundtrip_exactly() {
let fmt = AudioFormat {
bits_per_sample: 16,
channels: 2,
data_start: 0,
block_align: 4,
sample_format: SampleFormat::SignedInt,
};
let mut s = 0x9E37_79B9_7F4A_7C15u64;
let mut noise = Vec::with_capacity(400_000);
while noise.len() < 400_000 {
s ^= s << 13;
s ^= s >> 7;
s ^= s << 17;
noise.extend_from_slice(&s.to_le_bytes());
}
roundtrip(&fmt, 0, &noise).expect("must still be lossless on noise");
}
#[test]
fn refuses_chunks_with_too_little_to_model() {
let fmt = AudioFormat {
bits_per_sample: 16,
channels: 2,
data_start: 0,
block_align: 4,
sample_format: SampleFormat::SignedInt,
};
let mut out = Vec::new();
assert!(encode(&fmt, 0, &[1, 2, 3, 4], &mut out).is_none());
assert!(encode(&fmt, 0, &[], &mut out).is_none());
assert!(out.is_empty(), "a refusal must not write anything");
}
#[test]
fn decode_rejects_malformed_input() {
let mut out = Vec::new();
assert!(decode(&[], &mut out).is_err());
assert!(decode(&[9, 2, 16, 0, 0, 0, 0, 0, 0, 0, 0, 0], &mut out).is_err());
let mut bad = vec![VERSION, 2, 16, 0];
bad.extend_from_slice(&u16::MAX.to_le_bytes());
bad.extend_from_slice(&0u16.to_le_bytes());
bad.extend_from_slice(&10u32.to_le_bytes());
assert!(decode(&bad, &mut out).is_err());
}
#[test]
fn truncated_stream_is_an_error_not_a_panic() {
let fmt = AudioFormat {
bits_per_sample: 16,
channels: 2,
data_start: 0,
block_align: 4,
sample_format: SampleFormat::SignedInt,
};
let pcm = tone(50_000, 2, 56);
let mut enc = Vec::new();
encode(&fmt, 0, &pcm, &mut enc).unwrap();
for cut in [HEADER_LEN + 1, enc.len() / 4, enc.len() / 2, enc.len() - 1] {
let mut out = Vec::new();
if decode(&enc[..cut], &mut out).is_ok() {
assert_ne!(out, pcm, "truncated input decoded as complete");
}
}
}
#[test]
fn bit_io_roundtrips_arbitrary_widths() {
let mut w = BitWriter::new();
let values: Vec<(u64, u32)> = vec![
(0, 1),
(1, 1),
(5, 3),
(0xFFFF, 16),
(0, 5),
(12345, 20),
(1, 40),
];
for &(v, b) in &values {
w.write64(v, b);
}
w.write_unary(0);
w.write_unary(7);
w.write_unary(40);
let mut buf = Vec::new();
w.finish_into(&mut buf);
let mut r = BitReader::new(&buf);
for &(v, b) in &values {
assert_eq!(r.read64(b).unwrap(), v, "width {b}");
}
assert_eq!(r.read_unary(64).unwrap(), 0);
assert_eq!(r.read_unary(64).unwrap(), 7);
assert_eq!(r.read_unary(64).unwrap(), 40);
}
#[test]
fn zigzag_is_a_bijection_over_the_range_we_use() {
for v in [0i64, 1, -1, 2, -2, 32767, -32768, 1 << 40, -(1 << 40)] {
assert_eq!(unzigzag(zigzag(v)), v, "zigzag failed for {v}");
}
}
}