1use crate::error::{Error, Result};
64
65const VERSION: u8 = 3;
66const PARTITION: usize = 4096;
69const ESCAPE_K: u32 = 31;
72const ZERO_K: u32 = 30;
78const MAX_K: u32 = 29;
80const MAX_QUOTIENT: u32 = 48;
82const HEADER_LEN: usize = 16;
83const MAX_LPC_ORDER: usize = 12;
87const COEF_PRECISION: u32 = 15;
90
91#[derive(Debug, Clone, Copy, PartialEq, Eq)]
93pub enum SampleFormat {
94 SignedInt,
96 UnsignedByte,
98 Float32,
100}
101
102#[derive(Debug, Clone, Copy, PartialEq, Eq)]
104pub struct AudioFormat {
105 pub bits_per_sample: u16,
106 pub channels: u16,
107 pub data_start: u64,
109 pub block_align: u16,
111 pub sample_format: SampleFormat,
112}
113
114impl AudioFormat {
115 #[inline]
117 pub fn sample_bytes(&self) -> usize {
118 self.bits_per_sample as usize / 8
119 }
120
121 pub fn supported(&self) -> bool {
127 let width_ok = match self.sample_format {
128 SampleFormat::UnsignedByte => self.bits_per_sample == 8,
129 SampleFormat::SignedInt => matches!(self.bits_per_sample, 16 | 24 | 32),
130 SampleFormat::Float32 => self.bits_per_sample == 32,
131 };
132 width_ok
133 && (self.channels == 1 || self.channels == 2)
134 && self.block_align as usize == self.channels as usize * self.sample_bytes()
135 }
136}
137
138#[derive(Debug, Clone, Copy, PartialEq, Eq)]
140enum Decorrelation {
141 Independent = 0,
143 LeftSide = 1,
145 RightSide = 2,
147}
148
149impl Decorrelation {
150 fn from_u8(v: u8) -> Result<Self> {
151 Ok(match v {
152 0 => Decorrelation::Independent,
153 1 => Decorrelation::LeftSide,
154 2 => Decorrelation::RightSide,
155 other => return Err(Error::Compress(format!("unknown decorrelation {other}"))),
156 })
157 }
158}
159
160pub fn parse_wav_header(head: &[u8]) -> Option<AudioFormat> {
169 if head.len() < 44 || &head[0..4] != b"RIFF" || &head[8..12] != b"WAVE" {
170 return None;
171 }
172 let mut pos = 12usize;
173 let mut channels = 0u16;
174 let mut bits = 0u16;
175 let mut block_align = 0u16;
176 let mut sample_format = SampleFormat::SignedInt;
177 let mut seen_fmt = false;
178
179 while pos + 8 <= head.len() {
180 let id = &head[pos..pos + 4];
181 let size = u32::from_le_bytes(head[pos + 4..pos + 8].try_into().ok()?) as usize;
182 let body = pos + 8;
183
184 if id == b"fmt " {
185 if body + 16 > head.len() {
186 return None;
187 }
188 let audio_format = u16::from_le_bytes(head[body..body + 2].try_into().ok()?);
189 if !matches!(audio_format, 1 | 3 | 0xFFFE) {
193 return None;
194 }
195 channels = u16::from_le_bytes(head[body + 2..body + 4].try_into().ok()?);
196 block_align = u16::from_le_bytes(head[body + 12..body + 14].try_into().ok()?);
197 bits = u16::from_le_bytes(head[body + 14..body + 16].try_into().ok()?);
198 sample_format = match (audio_format, bits) {
199 (3, 32) => SampleFormat::Float32,
200 (_, 8) => SampleFormat::UnsignedByte,
201 _ => SampleFormat::SignedInt,
202 };
203 seen_fmt = true;
204 } else if id == b"data" {
205 if !seen_fmt {
206 return None;
207 }
208 let fmt = AudioFormat {
209 bits_per_sample: bits,
210 channels,
211 data_start: body as u64,
212 block_align,
213 sample_format,
214 };
215 return fmt.supported().then_some(fmt);
216 }
217
218 pos = body + size + (size & 1);
220 if size == 0 {
221 return None;
222 }
223 }
224 None
225}
226
227pub fn encode(
238 fmt: &AudioFormat,
239 chunk_offset: u64,
240 input: &[u8],
241 out: &mut Vec<u8>,
242) -> Option<usize> {
243 if !fmt.supported() {
244 return None;
245 }
246 let align = fmt.block_align as usize;
247 let channels = fmt.channels as usize;
248
249 let data_start = fmt.data_start;
252 let body_start = chunk_offset.max(data_start);
253 if body_start.saturating_sub(chunk_offset) as usize >= input.len() {
254 return None;
255 }
256 let mut prefix_len = (body_start - chunk_offset) as usize;
257 let rel = body_start - data_start;
258 let pad = (align - (rel % align as u64) as usize) % align;
259 prefix_len += pad;
260 if prefix_len >= input.len() {
261 return None;
262 }
263
264 let usable = input.len() - prefix_len;
265 let n_frames = usable / align;
266 if n_frames < 256 {
269 return None;
270 }
271 let suffix_len = usable - n_frames * align;
272 if prefix_len > u16::MAX as usize || suffix_len > u16::MAX as usize {
273 return None;
274 }
275
276 let scale = match fmt.sample_format {
279 SampleFormat::Float32 => float_scale(&input[prefix_len..prefix_len + n_frames * align])?,
280 _ => 0,
281 };
282
283 let width = fmt.sample_bytes();
285 let body = &input[prefix_len..prefix_len + n_frames * align];
286 let mut ch: Vec<Vec<i32>> = vec![Vec::with_capacity(n_frames); channels];
287 for f in 0..n_frames {
288 let base = f * align;
289 for (c, dst) in ch.iter_mut().enumerate() {
290 dst.push(read_sample(
291 &body[base + c * width..],
292 width,
293 fmt.sample_format,
294 scale,
295 ));
296 }
297 }
298
299 let max_abs = ch
304 .iter()
305 .flat_map(|c| c.iter())
306 .fold(0i64, |m, &v| m.max((v as i64).abs()));
307 let mut sample_bits = 1u32;
308 while sample_bits < 32 && max_abs >= (1i64 << (sample_bits - 1)) {
309 sample_bits += 1;
310 }
311 sample_bits = sample_bits.clamp(2, 32);
312
313 let mode = if channels == 2 && sample_bits < 32 {
318 choose_decorrelation(&ch[0], &ch[1])
319 } else {
320 Decorrelation::Independent
321 };
322 let coded: Vec<Vec<i32>> = match mode {
323 Decorrelation::Independent => ch,
324 Decorrelation::LeftSide => {
325 let side: Vec<i32> = ch[0].iter().zip(&ch[1]).map(|(l, r)| l - r).collect();
326 vec![std::mem::take(&mut ch[0]), side]
327 }
328 Decorrelation::RightSide => {
329 let side: Vec<i32> = ch[0].iter().zip(&ch[1]).map(|(l, r)| l - r).collect();
330 vec![std::mem::take(&mut ch[1]), side]
331 }
332 };
333
334 let start = out.len();
335 let format_code = match fmt.sample_format {
336 SampleFormat::SignedInt => 0u8,
337 SampleFormat::UnsignedByte => 1,
338 SampleFormat::Float32 => 2,
339 };
340 out.push(VERSION);
341 out.push(fmt.channels as u8);
342 out.push(fmt.bits_per_sample as u8);
343 out.push(mode as u8);
344 out.push(format_code);
345 out.push(scale as u8);
346 out.push(sample_bits as u8);
347 out.push(0); out.extend_from_slice(&(prefix_len as u16).to_le_bytes());
349 out.extend_from_slice(&(suffix_len as u16).to_le_bytes());
350 out.extend_from_slice(&(n_frames as u32).to_le_bytes());
351 out.extend_from_slice(&input[..prefix_len]);
352 out.extend_from_slice(&input[prefix_len + n_frames * align..]);
353
354 let verify_from = out.len();
355 let mut bits = BitWriter::new();
356 for (i, signal) in coded.iter().enumerate() {
359 let w = if i == 1 && mode != Decorrelation::Independent {
360 sample_bits + 1
361 } else {
362 sample_bits
363 };
364 encode_channel(signal, w, &mut bits);
365 }
366 bits.finish_into(out);
367 let _ = verify_from;
368
369 let mut check = Vec::with_capacity(input.len());
378 match decode(&out[start..], &mut check) {
379 Ok(()) if check == input => Some(out.len() - start),
380 _ => {
381 out.truncate(start);
382 debug_assert!(false, "pcm encoder produced output it could not decode");
383 tracing::warn!("pcm encode failed self-verification; falling back");
384 None
385 }
386 }
387}
388
389fn choose_decorrelation(left: &[i32], right: &[i32]) -> Decorrelation {
394 let cost = |signal: &[i32]| -> u64 {
395 signal
396 .windows(2)
397 .map(|w| (w[1] - w[0]).unsigned_abs() as u64)
398 .sum()
399 };
400 let side: Vec<i32> = left.iter().zip(right).map(|(l, r)| l - r).collect();
401 let (cl, cr, cs) = (cost(left), cost(right), cost(&side));
402 let independent = cl + cr;
403 let left_side = cl + cs;
404 let right_side = cr + cs;
405 if independent <= left_side && independent <= right_side {
406 Decorrelation::Independent
407 } else if left_side <= right_side {
408 Decorrelation::LeftSide
409 } else {
410 Decorrelation::RightSide
411 }
412}
413
414#[inline]
422fn residual(order: usize, s: &[i32], i: usize) -> i64 {
423 let x = |k: usize| s[i - k] as i64;
424 match order {
425 0 => x(0),
426 1 => x(0) - x(1),
427 2 => x(0) - 2 * x(1) + x(2),
428 3 => x(0) - 3 * x(1) + 3 * x(2) - x(3),
429 _ => x(0) - 4 * x(1) + 6 * x(2) - 4 * x(3) + x(4),
430 }
431}
432
433fn encode_channel(signal: &[i32], width: u32, bits: &mut BitWriter) {
434 const SCORE_SAMPLE: usize = 8192;
441 let scored = signal.len().min(SCORE_SAMPLE);
442 let mut fixed_order = 0usize;
443 let mut fixed_cost = u64::MAX;
444 for order in 0..=4usize {
445 if scored <= order {
446 break;
447 }
448 let cost: u64 = (order..scored)
449 .map(|i| residual(order, signal, i).unsigned_abs())
450 .sum();
451 if cost < fixed_cost {
452 fixed_cost = cost;
453 fixed_order = order;
454 }
455 }
456
457 let lpc_order = choose_lpc_order(&signal[..scored], fixed_cost);
463
464 match lpc_order {
465 None => {
466 bits.write(0, 1);
467 bits.write(fixed_order as u32, 3);
468 write_warmup(signal, fixed_order, width, bits);
469 encode_partitions(signal, fixed_order, None, bits);
470 }
471 Some(order) => {
472 bits.write(1, 1);
473 bits.write(order as u32 - 1, 5);
474 bits.write(COEF_PRECISION - 1, 4);
475 write_warmup(signal, order, width, bits);
476 encode_partitions(signal, order, Some(()), bits);
477 }
478 }
479}
480
481fn write_warmup(signal: &[i32], order: usize, width: u32, bits: &mut BitWriter) {
482 let mask = if width >= 32 {
483 u32::MAX
484 } else {
485 (1u32 << width) - 1
486 };
487 for &s in signal.iter().take(order) {
488 bits.write(s as u32 & mask, width);
489 }
490}
491
492thread_local! {
493 static WINDOW: std::cell::RefCell<Vec<f64>> = const { std::cell::RefCell::new(Vec::new()) };
499 static SCRATCH: std::cell::RefCell<Vec<f64>> = const { std::cell::RefCell::new(Vec::new()) };
501}
502
503fn with_window<R>(n: usize, f: impl FnOnce(&[f64]) -> R) -> R {
504 WINDOW.with(|w| {
505 let mut w = w.borrow_mut();
506 if w.len() != n {
507 w.clear();
508 w.reserve(n);
509 let scale = std::f64::consts::TAU / (n.max(2) - 1) as f64;
510 for i in 0..n {
511 w.push(0.5 - 0.5 * (i as f64 * scale).cos());
512 }
513 }
514 f(&w)
515 })
516}
517
518fn levinson(sample: &[i32], max_order: usize, coefs: &mut Vec<f64>, errors: &mut Vec<f64>) -> bool {
531 let n = sample.len();
532 coefs.clear();
533 errors.clear();
534 if n <= max_order + 1 || max_order == 0 || max_order > MAX_LPC_ORDER {
535 return false;
536 }
537
538 let mut autoc = [0.0f64; MAX_LPC_ORDER + 1];
539 with_window(n, |w| {
540 SCRATCH.with(|sc| {
541 let mut buf = sc.borrow_mut();
542 buf.clear();
543 buf.extend(sample.iter().zip(w).map(|(&v, &wi)| v as f64 * wi));
544 for (lag, slot) in autoc.iter_mut().enumerate().take(max_order + 1) {
545 *slot = buf[lag..].iter().zip(buf.iter()).map(|(a, b)| a * b).sum();
546 }
547 });
548 });
549 if autoc[0] <= 0.0 || !autoc[0].is_finite() {
550 return false;
551 }
552
553 let mut err = autoc[0];
554 coefs.resize(max_order, 0.0);
555 for i in 0..max_order {
556 let mut acc = autoc[i + 1];
557 for j in 0..i {
558 acc -= coefs[j] * autoc[i - j];
559 }
560 let k = acc / err;
561 if !k.is_finite() {
562 coefs.truncate(i);
563 return i > 0;
564 }
565 coefs[i] = k;
566 for j in 0..i / 2 {
567 let tmp = coefs[j];
568 coefs[j] = tmp - k * coefs[i - 1 - j];
569 coefs[i - 1 - j] -= k * tmp;
570 }
571 if i % 2 == 1 {
572 coefs[i / 2] -= k * coefs[i / 2];
573 }
574 err *= 1.0 - k * k;
575 errors.push(err.max(f64::MIN_POSITIVE));
576 if err <= 0.0 {
577 coefs.truncate(i + 1);
578 while errors.len() < max_order {
579 errors.push(f64::MIN_POSITIVE);
580 }
581 return true;
582 }
583 }
584 true
585}
586
587fn choose_lpc_order(sample: &[i32], fixed_cost: u64) -> Option<usize> {
592 if sample.len() < 4 * MAX_LPC_ORDER {
593 return None;
594 }
595 let mut coefs = Vec::new();
596 let mut errors = Vec::new();
597 if !levinson(sample, MAX_LPC_ORDER, &mut coefs, &mut errors) {
598 return None;
599 }
600
601 let n = sample.len() as f64;
605 let mut best: Option<(usize, f64)> = None;
606 for (idx, &err) in errors.iter().enumerate() {
607 let order = idx + 1;
608 if err <= 0.0 || !err.is_finite() {
609 continue;
610 }
611 let bits_per = 0.5 * (err / n).max(1e-9).log2();
612 let overhead = order as f64 * COEF_PRECISION as f64 / PARTITION as f64;
614 let total = bits_per + overhead;
615 if best.map_or(true, |(_, b)| total < b) {
616 best = Some((order, total));
617 }
618 }
619 let (order, est_bits) = best?;
620
621 let fixed_bits = if fixed_cost == 0 {
624 0.0
625 } else {
626 (fixed_cost as f64 / n).max(1.0).log2() + 1.0
627 };
628 (est_bits + 0.02 < fixed_bits).then_some(order)
631}
632
633#[derive(Clone)]
635struct Quantised {
636 coefs: Vec<i32>,
637 shift: u32,
638}
639
640fn quantise(coefs: &[f64]) -> Quantised {
644 let max = coefs.iter().fold(0.0f64, |m, c| m.max(c.abs()));
645 if max <= 0.0 || !max.is_finite() {
646 return Quantised {
647 coefs: vec![0; coefs.len()],
648 shift: 0,
649 };
650 }
651 let headroom = (COEF_PRECISION - 1) as i32;
652 let mut shift = headroom - (max.log2().floor() as i32) - 1;
653 shift = shift.clamp(0, 31);
654 let limit = 1i64 << (COEF_PRECISION - 1);
655
656 let mut error = 0.0f64;
660 let mut out = Vec::with_capacity(coefs.len());
661 for &c in coefs {
662 let scaled = c * (1u64 << shift) as f64 + error;
663 let q = scaled.round();
664 error = scaled - q;
665 out.push(q.clamp(-(limit as f64), (limit - 1) as f64) as i32);
666 }
667 Quantised {
668 coefs: out,
669 shift: shift as u32,
670 }
671}
672
673#[inline]
674fn lpc_residual(signal: &[i32], i: usize, q: &Quantised) -> i64 {
675 let mut acc: i64 = 0;
676 for (j, &c) in q.coefs.iter().enumerate() {
677 acc += c as i64 * signal[i - 1 - j] as i64;
678 }
679 signal[i] as i64 - (acc >> q.shift)
680}
681
682fn encode_partitions(signal: &[i32], order: usize, lpc: Option<()>, bits: &mut BitWriter) {
688 if signal.len() <= order {
689 return;
690 }
691 let mut start = order;
692 while start < signal.len() {
693 let end = (start + PARTITION).min(signal.len());
694 let residuals: Vec<i64> = match lpc {
695 None => (start..end).map(|i| residual(order, signal, i)).collect(),
696 Some(()) => {
697 let from = start - order;
699 let mut c = Vec::new();
700 let mut e = Vec::new();
701 let q = if levinson(&signal[from..end], order, &mut c, &mut e) && c.len() == order {
702 quantise(&c)
703 } else {
704 Quantised {
705 coefs: vec![0; order],
706 shift: 0,
707 }
708 };
709 bits.write(q.shift, 5);
710 for &c in &q.coefs {
711 bits.write(c as u32 & ((1u32 << COEF_PRECISION) - 1), COEF_PRECISION);
712 }
713 (start..end).map(|i| lpc_residual(signal, i, &q)).collect()
714 }
715 };
716 encode_residual_partition(&residuals, bits);
717 start = end;
718 }
719}
720
721fn encode_residual_partition(part: &[i64], bits: &mut BitWriter) {
722 let zig: Vec<u64> = part.iter().map(|&r| zigzag(r)).collect();
724 if zig.iter().all(|&z| z == 0) {
725 bits.write(ZERO_K, 5);
726 return;
727 }
728 let k = choose_rice_k(&zig);
729 if k == ESCAPE_K {
730 bits.write(ESCAPE_K, 5);
731 for &z in &zig {
732 bits.write64(z, 40);
733 }
734 return;
735 }
736 bits.write(k, 5);
737 for &z in &zig {
738 let q = (z >> k) as u32;
739 bits.write_unary(q);
740 if k > 0 {
741 bits.write64(z & ((1u64 << k) - 1), k);
742 }
743 }
744}
745
746fn choose_rice_k(zig: &[u64]) -> u32 {
752 if zig.is_empty() {
753 return 0;
754 }
755 let sum: u64 = zig.iter().fold(0u64, |a, &z| a.saturating_add(z));
756 let mean = sum / zig.len() as u64;
757 let guess = (64 - mean.leading_zeros()).saturating_sub(1).min(MAX_K);
758
759 let mut best_k = guess;
760 let mut best_bits = u64::MAX;
761 for k in guess.saturating_sub(2)..=(guess + 2).min(MAX_K) {
762 let mut total = 0u64;
763 let mut blown = false;
764 for &z in zig {
765 let q = z >> k;
766 if q > MAX_QUOTIENT as u64 {
767 blown = true;
768 break;
769 }
770 total += q + 1 + k as u64;
771 }
772 if !blown && total < best_bits {
773 best_bits = total;
774 best_k = k;
775 }
776 }
777 if best_bits == u64::MAX {
778 return ESCAPE_K;
780 }
781 best_k
782}
783
784#[inline]
791fn read_sample(b: &[u8], width: usize, format: SampleFormat, scale: u32) -> i32 {
792 match format {
793 SampleFormat::UnsignedByte => b[0] as i32 - 128,
794 SampleFormat::Float32 => {
795 let f = f32::from_le_bytes([b[0], b[1], b[2], b[3]]);
796 (f as f64 * (1u64 << scale) as f64) as i32
797 }
798 SampleFormat::SignedInt => match width {
799 2 => i16::from_le_bytes([b[0], b[1]]) as i32,
800 3 => i32::from_le_bytes([0, b[0], b[1], b[2]]) >> 8,
803 _ => i32::from_le_bytes([b[0], b[1], b[2], b[3]]),
804 },
805 }
806}
807
808#[inline]
810fn write_sample(v: i32, width: usize, format: SampleFormat, scale: u32, out: &mut Vec<u8>) {
811 match format {
812 SampleFormat::UnsignedByte => out.push((v + 128) as u8),
813 SampleFormat::Float32 => {
814 let f = (v as f64 / (1u64 << scale) as f64) as f32;
815 out.extend_from_slice(&f.to_le_bytes());
816 }
817 SampleFormat::SignedInt => {
818 let b = v.to_le_bytes();
819 match width {
820 2 => out.extend_from_slice(&b[..2]),
821 3 => out.extend_from_slice(&b[..3]),
822 _ => out.extend_from_slice(&b),
823 }
824 }
825 }
826}
827
828fn float_scale(body: &[u8]) -> Option<u32> {
837 for scale in [15u32, 23, 24, 31] {
838 let mul = (1u64 << scale) as f64;
839 let ok = body.chunks_exact(4).all(|b| {
840 let f = f32::from_le_bytes([b[0], b[1], b[2], b[3]]) as f64;
841 if !f.is_finite() {
842 return false;
843 }
844 let v = f * mul;
845 v.fract() == 0.0 && v.abs() <= i32::MAX as f64
846 });
847 if ok {
848 return Some(scale);
849 }
850 }
851 None
852}
853
854#[inline]
856fn sign_extend(v: u32, bits: u32) -> i32 {
857 if bits >= 32 {
858 return v as i32;
859 }
860 let shift = 32 - bits;
861 ((v << shift) as i32) >> shift
862}
863
864#[inline]
865fn zigzag(v: i64) -> u64 {
866 ((v << 1) ^ (v >> 63)) as u64
867}
868
869#[inline]
870fn unzigzag(z: u64) -> i64 {
871 ((z >> 1) as i64) ^ -((z & 1) as i64)
872}
873
874pub fn decode(input: &[u8], out: &mut Vec<u8>) -> Result<()> {
880 if input.len() < HEADER_LEN {
881 return Err(Error::Compress(
882 "pcm chunk is shorter than its header".into(),
883 ));
884 }
885 if input[0] != VERSION {
886 return Err(Error::Compress(format!(
887 "pcm version {} unsupported",
888 input[0]
889 )));
890 }
891 let channels = input[1] as usize;
892 let bits_per_sample = input[2] as u16;
893 let mode = Decorrelation::from_u8(input[3])?;
894 let sample_format = match input[4] {
895 0 => SampleFormat::SignedInt,
896 1 => SampleFormat::UnsignedByte,
897 2 => SampleFormat::Float32,
898 _ => {
899 return Err(Error::Compress(
900 "pcm chunk declares an unknown format".into(),
901 ))
902 }
903 };
904 let scale = input[5] as u32;
905 let sample_bits = input[6] as u32;
906 let prefix_len = u16::from_le_bytes([input[8], input[9]]) as usize;
907 let suffix_len = u16::from_le_bytes([input[10], input[11]]) as usize;
908 let n_frames = u32::from_le_bytes(input[12..16].try_into().unwrap()) as usize;
909 if !(2..=32).contains(&sample_bits) || scale > 40 {
910 return Err(Error::Compress(
911 "pcm chunk declares an impossible width".into(),
912 ));
913 }
914
915 let layout_ok = match sample_format {
916 SampleFormat::UnsignedByte => bits_per_sample == 8,
917 SampleFormat::SignedInt => matches!(bits_per_sample, 16 | 24 | 32),
918 SampleFormat::Float32 => bits_per_sample == 32,
919 };
920 if !layout_ok || !(1..=2).contains(&channels) {
921 return Err(Error::Compress(
922 "pcm chunk declares an unsupported layout".into(),
923 ));
924 }
925 let width = bits_per_sample as usize / 8;
926 let raw_end = HEADER_LEN
927 .checked_add(prefix_len)
928 .and_then(|v| v.checked_add(suffix_len))
929 .ok_or_else(|| Error::Compress("pcm chunk lengths overflow".into()))?;
930 if raw_end > input.len() {
931 return Err(Error::Compress("pcm chunk is truncated".into()));
932 }
933 if n_frames > 1 << 28 {
936 return Err(Error::Compress("pcm chunk declares too many frames".into()));
937 }
938
939 let prefix = &input[HEADER_LEN..HEADER_LEN + prefix_len];
940 let suffix = &input[HEADER_LEN + prefix_len..raw_end];
941 let mut bits = BitReader::new(&input[raw_end..]);
942
943 let mut coded: Vec<Vec<i32>> = Vec::with_capacity(channels);
944 for i in 0..channels {
945 let w = if i == 1 && mode != Decorrelation::Independent {
946 sample_bits + 1
947 } else {
948 sample_bits
949 };
950 coded.push(decode_channel(n_frames, w, &mut bits)?);
951 }
952
953 let channels_out: Vec<Vec<i32>> = match mode {
955 Decorrelation::Independent => coded,
956 Decorrelation::LeftSide => {
957 let left = &coded[0];
958 let side = &coded[1];
959 let right: Vec<i32> = left.iter().zip(side).map(|(l, s)| l - s).collect();
960 vec![coded[0].clone(), right]
961 }
962 Decorrelation::RightSide => {
963 let right = &coded[0];
964 let side = &coded[1];
965 let left: Vec<i32> = right.iter().zip(side).map(|(r, s)| r + s).collect();
966 vec![left, coded[0].clone()]
967 }
968 };
969
970 out.extend_from_slice(prefix);
971 for f in 0..n_frames {
972 for c in channels_out.iter() {
973 write_sample(c[f], width, sample_format, scale, out);
974 }
975 }
976 out.extend_from_slice(suffix);
977 Ok(())
978}
979
980fn decode_channel(n_frames: usize, width: u32, bits: &mut BitReader) -> Result<Vec<i32>> {
981 let is_lpc = bits.read(1)? == 1;
982 let (order, precision) = if is_lpc {
983 let order = bits.read(5)? as usize + 1;
984 let precision = bits.read(4)? + 1;
985 if precision > 32 {
986 return Err(Error::Compress("pcm coefficient precision too wide".into()));
987 }
988 (order, precision)
989 } else {
990 let order = bits.read(3)? as usize;
991 if order > 4 {
992 return Err(Error::Compress("pcm predictor order out of range".into()));
993 }
994 (order, 0)
995 };
996
997 let mut signal: Vec<i32> = Vec::with_capacity(n_frames);
998 for _ in 0..order.min(n_frames) {
999 signal.push(sign_extend(bits.read(width)?, width));
1000 }
1001 if n_frames <= order {
1002 return Ok(signal);
1003 }
1004
1005 let mut remaining = n_frames - order;
1006 while remaining > 0 {
1007 let count = remaining.min(PARTITION);
1008
1009 let quant = if is_lpc {
1012 let shift = bits.read(5)?;
1013 let mut coefs = Vec::with_capacity(order);
1014 for _ in 0..order {
1015 coefs.push(sign_extend(bits.read(precision)?, precision));
1016 }
1017 Some(Quantised { coefs, shift })
1018 } else {
1019 None
1020 };
1021
1022 let k = bits.read(5)?;
1023 for _ in 0..count {
1024 let z = if k == ZERO_K {
1025 0
1026 } else if k == ESCAPE_K {
1027 bits.read64(40)?
1028 } else {
1029 let q = bits.read_unary(MAX_QUOTIENT)? as u64;
1030 let low = if k > 0 { bits.read64(k)? } else { 0 };
1031 (q << k) | low
1032 };
1033 let r = unzigzag(z);
1034 let i = signal.len();
1035 let value = match &quant {
1036 Some(q) => {
1037 let mut acc: i64 = 0;
1038 for (j, &c) in q.coefs.iter().enumerate() {
1039 acc += c as i64 * signal[i - 1 - j] as i64;
1040 }
1041 r + (acc >> q.shift)
1042 }
1043 None => {
1044 let x = |back: usize| signal[i - back] as i64;
1046 match order {
1047 0 => r,
1048 1 => r + x(1),
1049 2 => r + 2 * x(1) - x(2),
1050 3 => r + 3 * x(1) - 3 * x(2) + x(3),
1051 _ => r + 4 * x(1) - 6 * x(2) + 4 * x(3) - x(4),
1052 }
1053 }
1054 };
1055 signal.push(value as i32);
1056 }
1057 remaining -= count;
1058 }
1059 Ok(signal)
1060}
1061
1062struct BitWriter {
1067 out: Vec<u8>,
1068 acc: u64,
1069 nbits: u32,
1070}
1071
1072impl BitWriter {
1073 fn new() -> Self {
1074 Self {
1075 out: Vec::new(),
1076 acc: 0,
1077 nbits: 0,
1078 }
1079 }
1080
1081 #[inline]
1082 fn write(&mut self, value: u32, bits: u32) {
1083 self.write64(value as u64, bits);
1084 }
1085
1086 #[inline]
1087 fn write64(&mut self, value: u64, bits: u32) {
1088 debug_assert!(bits <= 56);
1089 let masked = if bits >= 64 {
1090 value
1091 } else {
1092 value & ((1u64 << bits) - 1)
1093 };
1094 self.acc = (self.acc << bits) | masked;
1095 self.nbits += bits;
1096 while self.nbits >= 8 {
1097 self.nbits -= 8;
1098 self.out.push((self.acc >> self.nbits) as u8);
1099 }
1100 }
1101
1102 #[inline]
1104 fn write_unary(&mut self, q: u32) {
1105 let mut left = q;
1106 while left >= 32 {
1107 self.write64(0, 32);
1108 left -= 32;
1109 }
1110 if left > 0 {
1111 self.write64(0, left);
1112 }
1113 self.write64(1, 1);
1114 }
1115
1116 fn finish_into(mut self, out: &mut Vec<u8>) {
1117 if self.nbits > 0 {
1118 let pad = 8 - self.nbits;
1119 self.acc <<= pad;
1120 self.out.push(self.acc as u8);
1121 }
1122 out.extend_from_slice(&self.out);
1123 }
1124}
1125
1126struct BitReader<'a> {
1127 data: &'a [u8],
1128 pos: usize,
1129 acc: u64,
1130 nbits: u32,
1131}
1132
1133impl<'a> BitReader<'a> {
1134 fn new(data: &'a [u8]) -> Self {
1135 Self {
1136 data,
1137 pos: 0,
1138 acc: 0,
1139 nbits: 0,
1140 }
1141 }
1142
1143 #[inline]
1144 fn fill(&mut self) {
1145 while self.nbits <= 56 && self.pos < self.data.len() {
1146 self.acc = (self.acc << 8) | self.data[self.pos] as u64;
1147 self.pos += 1;
1148 self.nbits += 8;
1149 }
1150 }
1151
1152 #[inline]
1153 fn read(&mut self, bits: u32) -> Result<u32> {
1154 Ok(self.read64(bits)? as u32)
1155 }
1156
1157 #[inline]
1158 fn read64(&mut self, bits: u32) -> Result<u64> {
1159 if bits == 0 {
1160 return Ok(0);
1161 }
1162 self.fill();
1163 if self.nbits < bits {
1164 return Err(Error::Compress("pcm bitstream ended early".into()));
1165 }
1166 self.nbits -= bits;
1167 let value = (self.acc >> self.nbits) & ((1u64 << bits) - 1);
1168 Ok(value)
1169 }
1170
1171 #[inline]
1173 fn read_unary(&mut self, limit: u32) -> Result<u32> {
1174 let mut count = 0u32;
1175 loop {
1176 if self.read64(1)? == 1 {
1177 return Ok(count);
1178 }
1179 count += 1;
1180 if count > limit {
1181 return Err(Error::Compress("pcm unary run exceeds its limit".into()));
1182 }
1183 }
1184 }
1185}
1186
1187#[cfg(test)]
1188mod tests {
1189 use super::*;
1190
1191 fn wav_header(data_len: u32, channels: u16) -> Vec<u8> {
1193 let block_align = channels * 2;
1194 let byte_rate = 44_100 * block_align as u32;
1195 let mut h = Vec::new();
1196 h.extend_from_slice(b"RIFF");
1197 h.extend_from_slice(&(36 + data_len).to_le_bytes());
1198 h.extend_from_slice(b"WAVE");
1199 h.extend_from_slice(b"fmt ");
1200 h.extend_from_slice(&16u32.to_le_bytes());
1201 h.extend_from_slice(&1u16.to_le_bytes());
1202 h.extend_from_slice(&channels.to_le_bytes());
1203 h.extend_from_slice(&44_100u32.to_le_bytes());
1204 h.extend_from_slice(&byte_rate.to_le_bytes());
1205 h.extend_from_slice(&block_align.to_le_bytes());
1206 h.extend_from_slice(&16u16.to_le_bytes());
1207 h.extend_from_slice(b"data");
1208 h.extend_from_slice(&data_len.to_le_bytes());
1209 h
1210 }
1211
1212 fn tone(frames: usize, channels: usize, noise_shift: u32) -> Vec<u8> {
1213 let mut out = Vec::with_capacity(frames * channels * 2);
1214 let mut s = 0x1234_5678_9ABC_DEF0u64;
1215 for i in 0..frames {
1216 let t = i as f64 / 44_100.0;
1217 s ^= s << 13;
1218 s ^= s >> 7;
1219 s ^= s << 17;
1220 let dither = if noise_shift >= 63 {
1221 0
1222 } else {
1223 ((s >> noise_shift) as i16) / 4
1224 };
1225 for c in 0..channels {
1226 let f = if c == 0 { 440.0 } else { 659.25 };
1227 let v = ((t * f * std::f64::consts::TAU).sin() * 11_000.0) as i16;
1228 out.extend_from_slice(&v.wrapping_add(dither).to_le_bytes());
1229 }
1230 }
1231 out
1232 }
1233
1234 fn roundtrip(fmt: &AudioFormat, offset: u64, input: &[u8]) -> Option<usize> {
1235 let mut enc = Vec::new();
1236 let n = encode(fmt, offset, input, &mut enc)?;
1237 assert_eq!(n, enc.len());
1238 let mut dec = Vec::new();
1239 decode(&enc, &mut dec).expect("decode");
1240 assert_eq!(dec.len(), input.len(), "length changed");
1241 assert!(dec == input, "codec is not lossless");
1242 Some(enc.len())
1243 }
1244
1245 #[test]
1246 fn parses_a_canonical_wav_header() {
1247 let h = wav_header(1000, 2);
1248 let fmt = parse_wav_header(&h).expect("should parse");
1249 assert_eq!(fmt.channels, 2);
1250 assert_eq!(fmt.bits_per_sample, 16);
1251 assert_eq!(fmt.block_align, 4);
1252 assert_eq!(fmt.data_start, 44);
1253 assert!(fmt.supported());
1254 }
1255
1256 #[test]
1257 fn parses_a_header_with_extra_chunks_before_data() {
1258 let mut h = Vec::new();
1259 h.extend_from_slice(b"RIFF");
1260 h.extend_from_slice(&2000u32.to_le_bytes());
1261 h.extend_from_slice(b"WAVE");
1262 h.extend_from_slice(b"fmt ");
1263 h.extend_from_slice(&16u32.to_le_bytes());
1264 h.extend_from_slice(&1u16.to_le_bytes());
1265 h.extend_from_slice(&2u16.to_le_bytes());
1266 h.extend_from_slice(&44_100u32.to_le_bytes());
1267 h.extend_from_slice(&176_400u32.to_le_bytes());
1268 h.extend_from_slice(&4u16.to_le_bytes());
1269 h.extend_from_slice(&16u16.to_le_bytes());
1270 h.extend_from_slice(b"LIST");
1272 h.extend_from_slice(&5u32.to_le_bytes());
1273 h.extend_from_slice(b"INFOx");
1274 h.push(0);
1275 h.extend_from_slice(b"data");
1276 h.extend_from_slice(&1000u32.to_le_bytes());
1277 let fmt = parse_wav_header(&h).expect("should walk past LIST");
1278 assert_eq!(fmt.data_start as usize, h.len());
1279 }
1280
1281 #[test]
1282 fn rejects_non_wav_and_unsupported_layouts() {
1283 assert!(parse_wav_header(b"not a wav file at all, really truly not").is_none());
1284 assert!(parse_wav_header(&[]).is_none());
1285 let mut h = wav_header(1000, 2);
1287 h[34] = 24;
1288 h[32] = 6;
1289 let fmt = parse_wav_header(&h).expect("24-bit is supported");
1290 assert_eq!(fmt.bits_per_sample, 24);
1291 assert_eq!(fmt.sample_bytes(), 3);
1292
1293 let mut h8 = wav_header(1000, 2);
1296 h8[34] = 8;
1297 h8[32] = 2;
1298 let f8 = parse_wav_header(&h8).expect("8-bit is supported");
1299 assert_eq!(f8.sample_format, SampleFormat::UnsignedByte);
1300
1301 let mut h32 = wav_header(1000, 2);
1303 h32[34] = 32;
1304 h32[32] = 8;
1305 let f32i = parse_wav_header(&h32).expect("32-bit int is supported");
1306 assert_eq!(f32i.sample_format, SampleFormat::SignedInt);
1307
1308 let mut hf = wav_header(1000, 2);
1309 hf[34] = 32;
1310 hf[32] = 8;
1311 hf[20] = 3; let ff = parse_wav_header(&hf).expect("float is supported");
1313 assert_eq!(ff.sample_format, SampleFormat::Float32);
1314
1315 let mut h12 = wav_header(1000, 2);
1317 h12[34] = 12;
1318 h12[32] = 3;
1319 assert!(parse_wav_header(&h12).is_none());
1320 }
1321
1322 #[test]
1323 fn stereo_roundtrips_and_beats_zstd() {
1324 let fmt = AudioFormat {
1325 bits_per_sample: 16,
1326 channels: 2,
1327 data_start: 0,
1328 block_align: 4,
1329 sample_format: SampleFormat::SignedInt,
1330 };
1331 let pcm = tone(262_144, 2, 56);
1332 let size = roundtrip(&fmt, 0, &pcm).expect("should encode");
1333 let ratio = pcm.len() as f64 / size as f64;
1334 assert!(ratio > 1.3, "ratio only {ratio:.2}x");
1335 }
1336
1337 #[test]
1340 fn twenty_four_bit_roundtrips_and_compresses() {
1341 let fmt = AudioFormat {
1342 bits_per_sample: 24,
1343 channels: 2,
1344 data_start: 0,
1345 block_align: 6,
1346 sample_format: SampleFormat::SignedInt,
1347 };
1348 let frames = 150_000;
1350 let mut pcm = Vec::with_capacity(frames * 6);
1351 let mut s = 0x2545_F491_4F6C_DD1Du64;
1352 for i in 0..frames {
1353 let t = i as f64 / 44_100.0;
1354 s ^= s << 13;
1355 s ^= s >> 7;
1356 s ^= s << 17;
1357 let d = ((s >> 52) as i32 & 0x7FF) - 1024;
1358 for (f, amp) in [(440.0, 2_800_000.0), (659.25, 2_100_000.0)] {
1359 let v = ((t * f * std::f64::consts::TAU).sin() * amp) as i32 + d;
1360 pcm.extend_from_slice(&v.to_le_bytes()[..3]);
1361 }
1362 }
1363 let n = roundtrip(&fmt, 0, &pcm).expect("should encode 24-bit");
1364 let ratio = pcm.len() as f64 / n as f64;
1365 assert!(ratio > 1.3, "24-bit ratio only {ratio:.2}x");
1366 }
1367
1368 #[test]
1370 fn twenty_four_bit_extremes_survive() {
1371 let fmt = AudioFormat {
1372 bits_per_sample: 24,
1373 channels: 1,
1374 data_start: 0,
1375 block_align: 3,
1376 sample_format: SampleFormat::SignedInt,
1377 };
1378 let mut pcm = Vec::new();
1379 for i in 0..60_000i32 {
1380 let v = match i % 4 {
1382 0 => -8_388_608,
1383 1 => 8_388_607,
1384 2 => 0,
1385 _ => (i * 977) % 8_388_608 - 4_194_304,
1386 };
1387 pcm.extend_from_slice(&v.to_le_bytes()[..3]);
1388 }
1389 roundtrip(&fmt, 0, &pcm).expect("24-bit extremes must roundtrip");
1390 }
1391
1392 #[test]
1396 fn self_verification_catches_a_bad_encode() {
1397 let fmt = AudioFormat {
1398 bits_per_sample: 16,
1399 channels: 2,
1400 data_start: 0,
1401 block_align: 4,
1402 sample_format: SampleFormat::SignedInt,
1403 };
1404 let pcm = tone(80_000, 2, 56);
1405 let mut enc = Vec::new();
1406 encode(&fmt, 0, &pcm, &mut enc).expect("encode");
1407
1408 let mut broken = enc.clone();
1411 let at = HEADER_LEN + (broken.len() - HEADER_LEN) / 2;
1412 broken[at] ^= 0b0010_0000;
1413 let mut out = Vec::new();
1414 let matched = decode(&broken, &mut out).is_ok() && out == pcm;
1415 assert!(!matched, "a corrupted stream decoded as the original");
1416 }
1417
1418 #[test]
1422 fn every_accepted_encode_is_verified_lossless() {
1423 for (bits, channels, align) in [(16u16, 2u16, 4u16), (16, 1, 2), (24, 2, 6), (24, 1, 3)] {
1424 let fmt = AudioFormat {
1425 bits_per_sample: bits,
1426 channels,
1427 data_start: 0,
1428 block_align: align,
1429 sample_format: SampleFormat::SignedInt,
1430 };
1431 for noise in [63u32, 56, 48, 40] {
1432 let frames = 40_000;
1433 let mut pcm = Vec::new();
1434 let mut s = 0xDEAD_BEEF_CAFE_F00Du64 ^ noise as u64;
1435 for i in 0..frames {
1436 let t = i as f64 / 44_100.0;
1437 s ^= s << 13;
1438 s ^= s >> 7;
1439 s ^= s << 17;
1440 let amp = if bits == 16 { 11_000.0 } else { 2_800_000.0 };
1441 let d = if noise >= 63 {
1442 0
1443 } else {
1444 ((s >> noise) as i32) % 512
1445 };
1446 for c in 0..channels {
1447 let f = if c == 0 { 440.0 } else { 659.25 };
1448 let v = ((t * f * std::f64::consts::TAU).sin() * amp) as i32 + d;
1449 let b = v.to_le_bytes();
1450 pcm.extend_from_slice(&b[..(bits / 8) as usize]);
1451 }
1452 }
1453 roundtrip(&fmt, 0, &pcm)
1454 .unwrap_or_else(|| panic!("{bits}-bit {channels}ch noise={noise} refused"));
1455 }
1456 }
1457 }
1458
1459 #[test]
1462 fn eight_bit_unsigned_roundtrips() {
1463 let fmt = AudioFormat {
1464 bits_per_sample: 8,
1465 channels: 2,
1466 data_start: 0,
1467 block_align: 2,
1468 sample_format: SampleFormat::UnsignedByte,
1469 };
1470 let mut pcm = Vec::new();
1471 for i in 0..80_000 {
1472 let t = i as f64 / 8_000.0;
1473 for f in [440.0, 659.25] {
1474 let v = ((t * f * std::f64::consts::TAU).sin() * 100.0) as i32 + 128;
1475 pcm.push(v.clamp(0, 255) as u8);
1476 }
1477 }
1478 pcm.extend_from_slice(&[0, 255, 0, 255, 128, 128]);
1480 roundtrip(&fmt, 0, &pcm).expect("8-bit must roundtrip");
1481 }
1482
1483 #[test]
1484 fn thirty_two_bit_integer_roundtrips() {
1485 let fmt = AudioFormat {
1486 bits_per_sample: 32,
1487 channels: 2,
1488 data_start: 0,
1489 block_align: 8,
1490 sample_format: SampleFormat::SignedInt,
1491 };
1492 let mut pcm = Vec::new();
1493 for i in 0..60_000 {
1494 let t = i as f64 / 44_100.0;
1495 for f in [440.0, 659.25] {
1496 let v = ((t * f * std::f64::consts::TAU).sin() * 700_000_000.0) as i32;
1497 pcm.extend_from_slice(&v.to_le_bytes());
1498 }
1499 }
1500 for v in [i32::MIN, i32::MAX, 0, -1] {
1501 pcm.extend_from_slice(&v.to_le_bytes());
1502 pcm.extend_from_slice(&v.to_le_bytes());
1503 }
1504 roundtrip(&fmt, 0, &pcm).expect("32-bit int must roundtrip");
1505 }
1506
1507 #[test]
1510 fn integer_valued_float_roundtrips_bit_exactly() {
1511 let fmt = AudioFormat {
1512 bits_per_sample: 32,
1513 channels: 2,
1514 data_start: 0,
1515 block_align: 8,
1516 sample_format: SampleFormat::Float32,
1517 };
1518 let mut pcm = Vec::new();
1519 for i in 0..60_000 {
1520 let t = i as f64 / 44_100.0;
1521 for f in [440.0, 659.25] {
1522 let q = ((t * f * std::f64::consts::TAU).sin() * 6_000_000.0) as i32;
1524 let v = q as f32 / (1i32 << 23) as f32;
1525 pcm.extend_from_slice(&v.to_le_bytes());
1526 }
1527 }
1528 let n = roundtrip(&fmt, 0, &pcm).expect("integer-valued float must encode");
1529 assert!(pcm.len() > n, "float should have compressed");
1530 }
1531
1532 #[test]
1536 fn fractional_float_is_declined_not_rounded() {
1537 let fmt = AudioFormat {
1538 bits_per_sample: 32,
1539 channels: 2,
1540 data_start: 0,
1541 block_align: 8,
1542 sample_format: SampleFormat::Float32,
1543 };
1544 let mut pcm = Vec::new();
1545 let mut s = 0x1234_5678_9ABC_DEF0u64;
1546 for _ in 0..40_000 {
1547 for _ in 0..2 {
1548 s ^= s << 13;
1549 s ^= s >> 7;
1550 s ^= s << 17;
1551 let v = f32::from_bits(((s >> 32) as u32 & 0x7FFF_FFFF) | 0x3000_0000);
1553 pcm.extend_from_slice(&v.to_le_bytes());
1554 }
1555 }
1556 let mut out = Vec::new();
1557 let r = encode(&fmt, 0, &pcm, &mut out);
1558 if r.is_some() {
1559 let mut back = Vec::new();
1561 decode(&out, &mut back).expect("decode");
1562 assert_eq!(back, pcm, "float encode was not bit-exact");
1563 } else {
1564 assert!(out.is_empty(), "a refusal must not write anything");
1565 }
1566 }
1567
1568 #[test]
1569 fn mono_roundtrips() {
1570 let fmt = AudioFormat {
1571 bits_per_sample: 16,
1572 channels: 1,
1573 data_start: 0,
1574 block_align: 2,
1575 sample_format: SampleFormat::SignedInt,
1576 };
1577 let pcm = tone(100_000, 1, 56);
1578 let size = roundtrip(&fmt, 0, &pcm).expect("should encode");
1579 assert!(pcm.len() > size);
1580 }
1581
1582 #[test]
1583 fn survives_a_chunk_that_starts_mid_frame() {
1584 let fmt = AudioFormat {
1585 bits_per_sample: 16,
1586 channels: 2,
1587 data_start: 44,
1588 block_align: 4,
1589 sample_format: SampleFormat::SignedInt,
1590 };
1591 let pcm = tone(200_000, 2, 56);
1592 for offset in [0u64, 44, 45, 46, 47, 48, 1000, 1001, 1002, 1003] {
1595 let body: Vec<u8> = if offset < 44 {
1596 let mut v = wav_header(pcm.len() as u32, 2);
1597 v.extend_from_slice(&pcm);
1598 v[offset as usize..].to_vec()
1599 } else {
1600 let skip = (offset - 44) as usize;
1601 pcm[skip..].to_vec()
1602 };
1603 roundtrip(&fmt, offset, &body).unwrap_or_else(|| panic!("offset {offset}"));
1604 }
1605 }
1606
1607 #[test]
1608 fn handles_silence_and_full_scale() {
1609 let fmt = AudioFormat {
1610 bits_per_sample: 16,
1611 channels: 2,
1612 data_start: 0,
1613 block_align: 4,
1614 sample_format: SampleFormat::SignedInt,
1615 };
1616 let silence = vec![0u8; 4 * 100_000];
1618 let n = roundtrip(&fmt, 0, &silence).expect("encode silence");
1619 assert!(
1620 silence.len() as f64 / n as f64 > 50.0,
1621 "silence only reached {:.0}x",
1622 silence.len() as f64 / n as f64
1623 );
1624
1625 let mut extreme = Vec::new();
1628 for i in 0..100_000 {
1629 let (l, r) = if i % 2 == 0 {
1632 (i16::MIN, i16::MAX)
1633 } else {
1634 (i16::MAX, i16::MIN)
1635 };
1636 extreme.extend_from_slice(&l.to_le_bytes());
1637 extreme.extend_from_slice(&r.to_le_bytes());
1638 }
1639 roundtrip(&fmt, 0, &extreme).expect("encode extremes");
1640 }
1641
1642 #[test]
1643 fn random_bytes_still_roundtrip_exactly() {
1644 let fmt = AudioFormat {
1646 bits_per_sample: 16,
1647 channels: 2,
1648 data_start: 0,
1649 block_align: 4,
1650 sample_format: SampleFormat::SignedInt,
1651 };
1652 let mut s = 0x9E37_79B9_7F4A_7C15u64;
1653 let mut noise = Vec::with_capacity(400_000);
1654 while noise.len() < 400_000 {
1655 s ^= s << 13;
1656 s ^= s >> 7;
1657 s ^= s << 17;
1658 noise.extend_from_slice(&s.to_le_bytes());
1659 }
1660 roundtrip(&fmt, 0, &noise).expect("must still be lossless on noise");
1661 }
1662
1663 #[test]
1664 fn refuses_chunks_with_too_little_to_model() {
1665 let fmt = AudioFormat {
1666 bits_per_sample: 16,
1667 channels: 2,
1668 data_start: 0,
1669 block_align: 4,
1670 sample_format: SampleFormat::SignedInt,
1671 };
1672 let mut out = Vec::new();
1673 assert!(encode(&fmt, 0, &[1, 2, 3, 4], &mut out).is_none());
1674 assert!(encode(&fmt, 0, &[], &mut out).is_none());
1675 assert!(out.is_empty(), "a refusal must not write anything");
1676 }
1677
1678 #[test]
1679 fn decode_rejects_malformed_input() {
1680 let mut out = Vec::new();
1681 assert!(decode(&[], &mut out).is_err());
1682 assert!(decode(&[9, 2, 16, 0, 0, 0, 0, 0, 0, 0, 0, 0], &mut out).is_err());
1683 let mut bad = vec![VERSION, 2, 16, 0];
1685 bad.extend_from_slice(&u16::MAX.to_le_bytes());
1686 bad.extend_from_slice(&0u16.to_le_bytes());
1687 bad.extend_from_slice(&10u32.to_le_bytes());
1688 assert!(decode(&bad, &mut out).is_err());
1689 }
1690
1691 #[test]
1692 fn truncated_stream_is_an_error_not_a_panic() {
1693 let fmt = AudioFormat {
1694 bits_per_sample: 16,
1695 channels: 2,
1696 data_start: 0,
1697 block_align: 4,
1698 sample_format: SampleFormat::SignedInt,
1699 };
1700 let pcm = tone(50_000, 2, 56);
1701 let mut enc = Vec::new();
1702 encode(&fmt, 0, &pcm, &mut enc).unwrap();
1703 for cut in [HEADER_LEN + 1, enc.len() / 4, enc.len() / 2, enc.len() - 1] {
1704 let mut out = Vec::new();
1705 if decode(&enc[..cut], &mut out).is_ok() {
1708 assert_ne!(out, pcm, "truncated input decoded as complete");
1709 }
1710 }
1711 }
1712
1713 #[test]
1714 fn bit_io_roundtrips_arbitrary_widths() {
1715 let mut w = BitWriter::new();
1716 let values: Vec<(u64, u32)> = vec![
1717 (0, 1),
1718 (1, 1),
1719 (5, 3),
1720 (0xFFFF, 16),
1721 (0, 5),
1722 (12345, 20),
1723 (1, 40),
1724 ];
1725 for &(v, b) in &values {
1726 w.write64(v, b);
1727 }
1728 w.write_unary(0);
1729 w.write_unary(7);
1730 w.write_unary(40);
1731 let mut buf = Vec::new();
1732 w.finish_into(&mut buf);
1733
1734 let mut r = BitReader::new(&buf);
1735 for &(v, b) in &values {
1736 assert_eq!(r.read64(b).unwrap(), v, "width {b}");
1737 }
1738 assert_eq!(r.read_unary(64).unwrap(), 0);
1739 assert_eq!(r.read_unary(64).unwrap(), 7);
1740 assert_eq!(r.read_unary(64).unwrap(), 40);
1741 }
1742
1743 #[test]
1744 fn zigzag_is_a_bijection_over_the_range_we_use() {
1745 for v in [0i64, 1, -1, 2, -2, 32767, -32768, 1 << 40, -(1 << 40)] {
1746 assert_eq!(unzigzag(zigzag(v)), v, "zigzag failed for {v}");
1747 }
1748 }
1749}