Skip to main content

proof_engine/audio/
wav.rs

1//! Reading and writing WAV files, through the [`hound`] crate.
2//!
3//! Samples are always `f32` in `[-1, 1]`, interleaved by channel, which is
4//! what the synthesiser produces and what [`crate::asset::SoundAsset`]
5//! holds. Writing produces 16-bit PCM, which every player and editor reads;
6//! reading accepts 8, 16, 24 and 32-bit PCM and 32-bit float.
7//!
8//! ```rust
9//! use proof_engine::audio::wav::{read_wav, write_wav};
10//! # let path = std::env::temp_dir().join("proof_engine_wav_doc.wav");
11//! // A tenth of a second of a 440 Hz sine, mono.
12//! let sine: Vec<f32> = (0..4410)
13//!     .map(|i| (i as f32 / 44100.0 * 440.0 * std::f32::consts::TAU).sin() * 0.5)
14//!     .collect();
15//! write_wav(&path, 44100, 1, &sine).unwrap();
16//! let clip = read_wav(&std::fs::read(&path).unwrap()).unwrap();
17//! assert_eq!((clip.sample_rate, clip.channels, clip.samples.len()), (44100, 1, 4410));
18//! ```
19
20use std::io::{self, Cursor};
21use std::path::Path;
22
23/// A decoded WAV file.
24#[derive(Debug, Clone, PartialEq)]
25pub struct WavClip {
26    /// Frames per second.
27    pub sample_rate: u32,
28    /// Interleaved channels per frame.
29    pub channels: u16,
30    /// Interleaved samples in `[-1, 1]`.
31    pub samples: Vec<f32>,
32}
33
34impl WavClip {
35    /// Length in seconds.
36    pub fn duration_secs(&self) -> f32 {
37        if self.sample_rate == 0 || self.channels == 0 {
38            return 0.0;
39        }
40        self.samples.len() as f32 / (self.sample_rate as f32 * self.channels as f32)
41    }
42}
43
44fn to_io(e: hound::Error) -> io::Error {
45    match e {
46        hound::Error::IoError(e) => e,
47        other => io::Error::new(io::ErrorKind::InvalidData, other.to_string()),
48    }
49}
50
51/// Write interleaved `f32` samples to `path` as 16-bit PCM WAV.
52///
53/// Samples outside `[-1, 1]` are clipped. `samples.len()` must be a whole
54/// number of frames.
55pub fn write_wav(path: impl AsRef<Path>, sample_rate: u32, channels: u16, samples: &[f32]) -> io::Result<()> {
56    if channels == 0 || samples.len() % channels as usize != 0 {
57        return Err(io::Error::new(
58            io::ErrorKind::InvalidInput,
59            format!("{} samples is not a whole number of {channels} channel frames", samples.len()),
60        ));
61    }
62    let spec = hound::WavSpec {
63        channels,
64        sample_rate,
65        bits_per_sample: 16,
66        sample_format: hound::SampleFormat::Int,
67    };
68    let mut w = hound::WavWriter::create(path, spec).map_err(to_io)?;
69    {
70        let mut w16 = w.get_i16_writer(samples.len() as u32);
71        for &s in samples {
72            // Scale by 32768 so reading back (which divides by 32768) is exact
73            // to the step; only +1.0 itself saturates.
74            w16.write_sample((s.clamp(-1.0, 1.0) * 32768.0).round().min(32767.0) as i16);
75        }
76        w16.flush().map_err(to_io)?;
77    }
78    w.finalize().map_err(to_io)
79}
80
81/// Decode a WAV file held in memory.
82pub fn read_wav(bytes: &[u8]) -> io::Result<WavClip> {
83    let mut r = hound::WavReader::new(Cursor::new(bytes)).map_err(to_io)?;
84    let spec = r.spec();
85    let samples: Vec<f32> = match spec.sample_format {
86        hound::SampleFormat::Float => r.samples::<f32>().collect::<Result<_, _>>().map_err(to_io)?,
87        hound::SampleFormat::Int => {
88            let scale = 1.0 / (1u64 << (spec.bits_per_sample - 1)) as f32;
89            r.samples::<i32>()
90                .map(|s| s.map(|v| v as f32 * scale))
91                .collect::<Result<_, _>>()
92                .map_err(to_io)?
93        }
94    };
95    Ok(WavClip { sample_rate: spec.sample_rate, channels: spec.channels, samples })
96}
97
98/// Read a WAV file from disk.
99pub fn load_wav(path: impl AsRef<Path>) -> io::Result<WavClip> {
100    read_wav(&std::fs::read(path)?)
101}
102
103#[cfg(test)]
104mod tests {
105    use super::*;
106
107    fn tmp(name: &str) -> std::path::PathBuf {
108        let dir = std::env::temp_dir().join("proof_engine_wav_tests");
109        std::fs::create_dir_all(&dir).unwrap();
110        dir.join(name)
111    }
112
113    #[test]
114    fn sixteen_bit_round_trip_is_within_one_step() {
115        let path = tmp("rt.wav");
116        let input: Vec<f32> = (0..2000).map(|i| ((i as f32) * 0.013).sin() * 0.9).collect();
117        write_wav(&path, 32000, 2, &input).unwrap();
118        let clip = load_wav(&path).unwrap();
119        assert_eq!((clip.sample_rate, clip.channels), (32000, 2));
120        assert_eq!(clip.samples.len(), input.len());
121        assert!((clip.duration_secs() - 1000.0 / 32000.0).abs() < 1e-6);
122        for (a, b) in clip.samples.iter().zip(&input) {
123            assert!((a - b).abs() <= 0.5 / 32768.0 + 1e-7, "{a} vs {b}");
124        }
125    }
126
127    #[test]
128    fn out_of_range_samples_are_clipped_not_wrapped() {
129        let path = tmp("clip.wav");
130        write_wav(&path, 8000, 1, &[2.0, -3.0]).unwrap();
131        let clip = load_wav(&path).unwrap();
132        assert!(clip.samples[0] > 0.99 && clip.samples[1] < -0.99);
133    }
134
135    #[test]
136    fn reads_float_and_24_bit_files() {
137        for (bits, format) in [(32u16, hound::SampleFormat::Float), (24, hound::SampleFormat::Int)] {
138            let path = tmp(&format!("in_{bits}.wav"));
139            let spec = hound::WavSpec { channels: 1, sample_rate: 48000, bits_per_sample: bits, sample_format: format };
140            let mut w = hound::WavWriter::create(&path, spec).unwrap();
141            if format == hound::SampleFormat::Float {
142                w.write_sample(0.25f32).unwrap();
143            } else {
144                w.write_sample((0.25 * (1 << 23) as f32) as i32).unwrap();
145            }
146            w.finalize().unwrap();
147            let clip = load_wav(&path).unwrap();
148            assert!((clip.samples[0] - 0.25).abs() < 1e-5, "{bits} bit: {}", clip.samples[0]);
149        }
150    }
151
152    #[test]
153    fn bad_input_is_an_error() {
154        assert!(read_wav(b"RIFF....not really").is_err());
155        assert!(write_wav(tmp("odd.wav"), 8000, 2, &[0.0; 3]).is_err());
156    }
157
158    #[test]
159    fn a_synth_bounce_survives_the_trip_to_disk() {
160        use crate::audio::{math_source::MathAudioSource, AudioEvent, OfflineRenderer};
161        use glam::Vec3;
162        let mut synth = OfflineRenderer::new(44100);
163        synth.emit(AudioEvent::SpawnSource { source: MathAudioSource::death_knell(Vec3::ZERO), position: Vec3::ZERO });
164        let stereo = synth.render(0.4);
165        let path = tmp("bounce.wav");
166        write_wav(&path, synth.sample_rate(), 2, &stereo).unwrap();
167        let clip = load_wav(&path).unwrap();
168        assert_eq!(clip.samples.len(), stereo.len());
169        let peak = clip.samples.iter().fold(0.0f32, |m, s| m.max(s.abs()));
170        assert!(peak > 0.05, "the knell is silent on disk: {peak}");
171    }
172}