proof_engine/audio/
wav.rs1use std::io::{self, Cursor};
21use std::path::Path;
22
23#[derive(Debug, Clone, PartialEq)]
25pub struct WavClip {
26 pub sample_rate: u32,
28 pub channels: u16,
30 pub samples: Vec<f32>,
32}
33
34impl WavClip {
35 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
51pub 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 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
81pub 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
98pub 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}