Skip to main content

runsync_transfer/codec/
pcm.rs

1//! Lossless codec for uncompressed PCM audio.
2//!
3//! # Why this exists
4//!
5//! zstd manages about 1.05x on 16-bit PCM — below the gate that decides a chunk
6//! is worth compressing at all, so the engine ships raw `.wav` untouched. That
7//! is the correct call for a general-purpose compressor, and it is the wrong
8//! outcome for a transfer whose payload *is* audio.
9//!
10//! PCM is not high-entropy, it is *correlated*. Consecutive samples are close
11//! together and the two channels of a stereo pair are nearly the same signal.
12//! A general compressor looks for repeated byte strings and finds none, because
13//! a waveform never repeats exactly. Modelling the correlation instead is what
14//! FLAC does, and most of FLAC's ratio comes from two cheap stages:
15//!
16//! 1. **Decorrelate the channels.** Send left and (left − right) rather than
17//!    left and right; the difference of two similar signals is small.
18//! 2. **Predict from previous samples** and send only the error. A fixed
19//!    polynomial predictor of order 0–4 costs a few adds per sample and turns a
20//!    smooth waveform into residuals clustered near zero.
21//!
22//! Small numbers clustered near zero is exactly what Rice coding encodes well,
23//! so that is the third stage. The predictor is solved per partition, by
24//! autocorrelation and Levinson–Durbin, so it follows the music through a
25//! track rather than assuming the signal is locally polynomial; fixed
26//! polynomial predictors are kept as the cheaper option and used whenever they
27//! win. On a real library this reaches **1.58x**, against reference FLAC's
28//! ~1.53x on the same files.
29//!
30//! # Formats
31//!
32//! Every sample layout a WAV can hold: 8-bit unsigned, 16/24/32-bit signed,
33//! and 32-bit IEEE float. Each reaches the same signed-integer predictor by a
34//! different exactly-reversible step — 8-bit is biased by 128, 24-bit has no
35//! native type, float is scaled by a power of two chosen so the whole chunk
36//! lands on integers. Float that is genuinely fractional has no such scale, and
37//! is declined rather than rounded; nothing here ever approximates.
38//!
39//! # Why you can trust it with your only copy
40//!
41//! Every encode is decoded again and compared against its input before it is
42//! accepted. If they differ for any reason at all, the encode is discarded and
43//! the caller falls back to zstd or to raw bytes. A bug in this file can
44//! therefore cost throughput, never a file.
45//!
46//! # Guarantees
47//!
48//! Exactly lossless: `decode(encode(x)) == x` for every input, including
49//! inputs that are not really audio — and checked, per chunk, at encode time
50//! rather than merely intended. Anything the codec cannot model — an
51//! unsupported sample width, a chunk with no whole frames in it — is refused
52//! rather than approximated, and the caller falls back to zstd or to raw.
53//!
54//! # Chunk alignment
55//!
56//! The engine splits files into fixed-size chunks that know nothing about frame
57//! boundaries, and a `.wav` file starts with a header, so a chunk generally
58//! begins mid-frame. Each encoded chunk therefore carries a raw prefix and
59//! suffix around the region it could actually model, and is entirely
60//! self-describing: the decoder needs no manifest, no format side-channel, and
61//! no neighbouring chunk.
62
63use crate::error::{Error, Result};
64
65const VERSION: u8 = 3;
66/// Residuals are Rice-coded in runs of this many, each with its own parameter,
67/// so a quiet passage and a loud one in the same chunk do not have to share.
68const PARTITION: usize = 4096;
69/// Rice parameter reserved to mean "this partition is pathological, the
70/// residuals follow verbatim".
71const ESCAPE_K: u32 = 31;
72/// Reserved to mean "every residual in this partition is zero".
73///
74/// Without it, digital silence still costs one bit per sample — a lead-in, a
75/// fade-out, or a padded track would be coded 4096 times over to say nothing.
76/// With it the whole partition is five bits.
77const ZERO_K: u32 = 30;
78/// Largest real Rice parameter, the two above being reserved.
79const MAX_K: u32 = 29;
80/// Refuse to unary-code a quotient longer than this; escape the partition instead.
81const MAX_QUOTIENT: u32 = 48;
82const HEADER_LEN: usize = 16;
83/// Highest LPC order considered. Beyond about this the coefficients cost more
84/// than the prediction saves, and the autocorrelation gets proportionally
85/// dearer; FLAC's own high presets stop at 12 for the same reason.
86const MAX_LPC_ORDER: usize = 12;
87/// Bits per quantised coefficient. 15 keeps the accumulator comfortably inside
88/// i64 for 24-bit samples at order 12, with room to spare.
89const COEF_PRECISION: u32 = 15;
90
91/// How samples are laid out within a frame.
92#[derive(Debug, Clone, Copy, PartialEq, Eq)]
93pub enum SampleFormat {
94    /// Signed little-endian two's complement, 16/24/32-bit.
95    SignedInt,
96    /// Unsigned, biased by 128. WAV stores 8-bit this way and no other way.
97    UnsignedByte,
98    /// IEEE 754 binary32.
99    Float32,
100}
101
102/// PCM layout, as read from a container header.
103#[derive(Debug, Clone, Copy, PartialEq, Eq)]
104pub struct AudioFormat {
105    pub bits_per_sample: u16,
106    pub channels: u16,
107    /// Byte offset in the file where sample data begins.
108    pub data_start: u64,
109    /// Bytes per frame: `channels * bits_per_sample / 8`.
110    pub block_align: u16,
111    pub sample_format: SampleFormat,
112}
113
114impl AudioFormat {
115    /// Bytes per sample.
116    #[inline]
117    pub fn sample_bytes(&self) -> usize {
118        self.bits_per_sample as usize / 8
119    }
120
121    /// Is this something the codec models, rather than merely tolerates?
122    ///
123    /// Every width WAV can hold: 8-bit unsigned, 16/24/32-bit signed, and
124    /// 32-bit float. What each needs to reach the same signed-integer predictor
125    /// differs, and is handled at the point the samples are read.
126    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/// How a stereo pair was rewritten before prediction.
139#[derive(Debug, Clone, Copy, PartialEq, Eq)]
140enum Decorrelation {
141    /// Channels coded as they arrived.
142    Independent = 0,
143    /// left, and (left − right).
144    LeftSide = 1,
145    /// right, and (left − right).
146    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
160// ---------------------------------------------------------------------------
161// Container parsing
162// ---------------------------------------------------------------------------
163
164/// Read a RIFF/WAVE header, returning the sample layout and where data starts.
165///
166/// Walks the chunk list rather than assuming the canonical 44-byte layout,
167/// because real files carry `LIST`, `fact` and other chunks before `data`.
168pub 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            // 1 = integer PCM, 3 = IEEE float, 0xFFFE = extensible (still one of
190            // the two, described by a sub-format GUID we do not need to read
191            // because the bit width already tells us which).
192            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        // Chunks are word-aligned, and a zero size would not advance.
219        pos = body + size + (size & 1);
220        if size == 0 {
221            return None;
222        }
223    }
224    None
225}
226
227// ---------------------------------------------------------------------------
228// Encode
229// ---------------------------------------------------------------------------
230
231/// Encode one chunk of a PCM file.
232///
233/// `chunk_offset` is where this chunk starts in the file, which is what lets
234/// the codec find the frame boundaries inside it. Returns `None` when there is
235/// nothing worth modelling — too few whole frames — so the caller can fall
236/// back rather than pay for a pointless pass.
237pub 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    // Where the modellable region starts: past any container header this chunk
250    // happens to contain, then rounded up to the next frame boundary.
251    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    // Below this the header and per-partition overhead dominate and the whole
267    // exercise is a loss.
268    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    // Float needs a scale that makes the whole chunk exact before anything
277    // else can happen; if there is none, decline and let zstd have it.
278    let scale = match fmt.sample_format {
279        SampleFormat::Float32 => float_scale(&input[prefix_len..prefix_len + n_frames * align])?,
280        _ => 0,
281    };
282
283    // Deinterleave into signed channels.
284    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    // How wide the samples actually are, once in the integer domain. Measured
300    // rather than assumed: a float scaled by 2^23 needs 24-ish bits, not 32,
301    // and coding the warm-up samples at the container's width would waste the
302    // difference on every partition.
303    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    // Choose a channel representation by trying each and keeping the cheapest.
314    //
315    // Only when the difference is guaranteed to fit: at a full 32 bits, `l - r`
316    // would overflow, and no ratio is worth a wrong sample.
317    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); // reserved
348    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    // The side channel of a stereo difference needs one extra bit of range.
357    // `sample_bits` was measured above.
358    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    // Decode what was just produced and compare it against the input.
370    //
371    // The engine already hashes every chunk end to end, so a bad encode could
372    // never be delivered as good — but it would fail the transfer. Checking
373    // here converts an encoder bug from a failed transfer into a silent
374    // fallback to zstd or raw: the file always arrives, byte for byte,
375    // whatever this codec does. That is worth a decode pass, and the decoder
376    // is the faster half.
377    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
389/// Pick the stereo representation whose residuals will cost least.
390///
391/// Estimated by summed absolute first difference, which tracks the eventual
392/// Rice cost closely enough and costs one pass instead of three full encodes.
393fn 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/// Fixed polynomial predictors, orders 0 through 4.
415///
416/// Order *p* predicts the next sample from the previous *p* by repeated
417/// differencing, which is exact integer arithmetic — no coefficients to
418/// transmit and nothing to round. Higher orders model smoother signals; the
419/// best order is whichever leaves the smallest residuals, so all five are
420/// scored and the cheapest wins.
421#[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    // Score each fixed order by the magnitude of the residuals it leaves.
435    //
436    // Scored on a sample rather than the whole channel: this is choosing
437    // between a handful of options, not measuring anything, and the best order
438    // for the first few thousand samples is almost always the best order for
439    // the rest.
440    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    // Then see whether a solved predictor does better on the same sample.
458    // Fixed predictors assume the signal is locally polynomial; LPC fits the
459    // actual spectrum, which is worth several percent on real music and much
460    // more on tonal material. It is only worth its cost when it wins, so both
461    // are costed and the cheaper is used.
462    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    /// The Hann window, kept between partitions.
494    ///
495    /// Every partition is the same length, so recomputing it meant a `cos()`
496    /// per sample per partition — half a million transcendental calls per MiB,
497    /// which cost more than the autocorrelation it was preparing for.
498    static WINDOW: std::cell::RefCell<Vec<f64>> = const { std::cell::RefCell::new(Vec::new()) };
499    /// Windowed samples, reused so the autocorrelation allocates nothing.
500    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
518/// Autocorrelation plus Levinson–Durbin.
519///
520/// Writes the coefficients for `max_order` into `coefs` and the prediction
521/// error at each order into `errors`, which is what the order search wants.
522///
523/// A Hann window is applied first: the autocorrelation of an abruptly truncated
524/// block describes the truncation as much as the signal, and tapering the edges
525/// keeps the solved predictor about the audio.
526///
527/// Everything goes into caller-owned buffers. Returning a `Vec` per order meant
528/// a dozen allocations for every partition of every channel, which cost more
529/// than the arithmetic they carried.
530fn 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
587/// Is a solved predictor worth it, and at what order?
588///
589/// Returns `None` when no order beats the fixed predictor that would otherwise
590/// be used, so the cheaper path stays the default.
591fn 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    // Levinson already yields each order's prediction error, and expected bits
602    // per residual go as log2 of it. Using that costs nothing, where trial
603    // encoding every order cost a full pass each.
604    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        // Charge the coefficients, in every partition they will appear in.
613        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    // Compare against the fixed predictor on the same footing: its cost was
622    // summed absolute residual, so convert to bits per sample the same way.
623    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    // Require a real margin: an LPC frame carries coefficients in every
629    // partition, so a marginal win is a loss once that is paid.
630    (est_bits + 0.02 < fixed_bits).then_some(order)
631}
632
633/// Quantised predictor: integer coefficients and the shift they were scaled by.
634#[derive(Clone)]
635struct Quantised {
636    coefs: Vec<i32>,
637    shift: u32,
638}
639
640/// Scale real coefficients into integers, because the predictor has to be
641/// evaluated identically on both sides and floating point is not identical
642/// across machines.
643fn 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    // Carry the rounding error forward: quantising each coefficient in
657    // isolation biases the predictor, and feeding the error into its neighbour
658    // recovers most of what that costs.
659    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
682/// Emit the residuals partition by partition.
683///
684/// With LPC, each partition carries its own coefficients. They cost about
685/// 180 bits against 4096 samples, which is nothing, and letting the predictor
686/// follow the music through a track is worth considerably more than that.
687fn 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                // Solve on this partition plus the history the predictor needs.
698                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    // Zigzag so small negatives are small unsigned values.
723    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
746/// Rice parameter for one partition.
747///
748/// The optimum is close to `log2(mean)`, so that is the starting point and the
749/// neighbours are checked exactly — a full search over 32 values would cost
750/// more than it saves.
751fn 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        // No parameter keeps the unary parts bounded: store the partition raw.
779        return ESCAPE_K;
780    }
781    best_k
782}
783
784/// Bring one stored sample into the signed integer domain the predictor works in.
785///
786/// Each format needs a different step, and each is exactly reversible:
787/// 8-bit is unsigned and biased, 24-bit has no native type, and float is
788/// scaled by a power of two chosen so every sample in the chunk lands on an
789/// integer (see [`float_scale`]).
790#[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            // 24-bit has no native type: place it in the high three bytes and
801            // arithmetic-shift down, which sign-extends in one step.
802            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/// The exact inverse of [`read_sample`].
809#[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
828/// Find a power-of-two scale that turns every float in the chunk into an exact
829/// integer, or `None` if no single scale does.
830///
831/// Float audio that came from an integer source — which is most of what a DAW
832/// exports — sits on a regular grid, so one scale makes the whole chunk exact
833/// and the full predictor applies. Genuinely fractional float (heavily
834/// processed material) does not, and is declined rather than approximated:
835/// this codec does not round anything, ever.
836fn 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/// Sign-extend a `bits`-wide two's complement value read from the bitstream.
855#[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
874// ---------------------------------------------------------------------------
875// Decode
876// ---------------------------------------------------------------------------
877
878/// Decode a chunk produced by [`encode`], appending the original bytes to `out`.
879pub 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    // A frame count that cannot fit any plausible chunk is a malformed or
934    // hostile header; refuse before allocating from it.
935    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    // Undo the channel decorrelation.
954    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        // Mirror the encoder exactly: coefficients first when this is an LPC
1010        // frame, then the Rice parameter, then the residuals.
1011        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                    // Reverse the fixed predictor from samples already rebuilt.
1045                    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
1062// ---------------------------------------------------------------------------
1063// Bit I/O, MSB first
1064// ---------------------------------------------------------------------------
1065
1066struct 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    /// `q` zero bits then a one.
1103    #[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    /// Count zeros up to and including the terminating one.
1172    #[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    /// A 44-byte canonical WAV header for 16-bit stereo at 44.1 kHz.
1192    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        // A LIST chunk of odd length, so the pad byte matters.
1271        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        // 24-bit is modelled, so it must now be accepted.
1286        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        // 8-bit is unsigned in WAV; recognised as such rather than mistaken
1294        // for signed, which would invert every sample.
1295        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        // 32-bit signed, and 32-bit IEEE float (audio_format 3).
1302        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; // audio_format = IEEE float
1312        let ff = parse_wav_header(&hf).expect("float is supported");
1313        assert_eq!(ff.sample_format, SampleFormat::Float32);
1314
1315        // A width that does not exist is still refused.
1316        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    /// 24-bit integer PCM is what sample packs and hi-res masters carry, and
1338    /// it is most of a production library by volume.
1339    #[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        // Build 24-bit tones with dither in the low bits.
1349        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    /// Every 24-bit value, including both extremes of the range, must survive.
1369    #[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            // Sweep the whole signed 24-bit range, extremes included.
1381            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    /// The encoder must never emit something it cannot itself decode. Simulate
1393    /// a broken encode by corrupting the bitstream and confirm the check that
1394    /// guards the real path would catch it.
1395    #[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        // Flip a bit deep in the coded residuals. Decoding must either fail or
1409        // produce something that differs — never silently pass as correct.
1410        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    /// Every encode that succeeds has already been decoded and compared, so a
1419    /// successful return is itself the guarantee. Sweep a spread of signals to
1420    /// exercise that path rather than trusting one.
1421    #[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    /// 8-bit WAV is unsigned and biased by 128. Getting that wrong would not
1460    /// merely compress badly, it would invert the waveform.
1461    #[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        // Include both rails, which are the values a biased format gets wrong.
1479        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    /// Float exported from an integer source lands on a regular grid, so one
1508    /// scale makes the whole chunk exact and the full predictor applies.
1509    #[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                // A 24-bit integer scaled to unity, which is what a DAW writes.
1523                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    /// Genuinely fractional float cannot be represented exactly on any single
1533    /// integer grid, and is declined rather than rounded. Nothing is ever
1534    /// approximated.
1535    #[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                // Arbitrary mantissas at wildly differing exponents.
1552                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            // If it did accept, it must still be exact.
1560            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        // Every possible phase relative to a 4-byte frame, and a chunk that
1593        // still contains part of the header.
1594        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        // Digital silence should compress enormously.
1617        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        // Alternating extremes: the worst case for a predictor, and the case
1626        // that would overflow a careless one.
1627        let mut extreme = Vec::new();
1628        for i in 0..100_000 {
1629            // Opposite rails on the two channels. Written out rather than
1630            // negated: `-i16::MIN` does not exist.
1631            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        // Not audio at all. It must not compress, and must not corrupt.
1645        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        // Header claims a prefix far larger than the buffer.
1684        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            // Either a clean error, or a short result — never a panic, and
1706            // never silently the wrong bytes presented as right.
1707            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}