1use std::io::{self, Seek, SeekFrom, Write};
23
24pub const SAMPLE_RATE_HZ: u32 = 24_000;
26
27pub const SAMPLES_PER_FRAME: usize = 1_920;
29
30pub const CHANNELS: u16 = 1;
32
33pub const BITS_PER_SAMPLE: u16 = 16;
35
36pub const WAV_HEADER_BYTES: usize = 44;
38
39#[must_use]
41pub const fn samples_for_frames(frames: usize) -> usize {
42 frames * SAMPLES_PER_FRAME
43}
44
45#[must_use]
56pub fn sample_to_i16(sample: f32) -> i16 {
57 if !sample.is_finite() {
58 return 0;
59 }
60 let clamped = sample.clamp(-1.0, 1.0);
61 (clamped * 32_767.0).round() as i16
64}
65
66#[must_use]
68pub fn pcm_f32_to_i16(pcm: &[f32]) -> Vec<i16> {
69 pcm.iter().copied().map(sample_to_i16).collect()
70}
71
72#[must_use]
78pub fn mean_square_energy(pcm: &[f32]) -> f64 {
79 if pcm.is_empty() {
80 return 0.0;
81 }
82 let total: f64 = pcm.iter().map(|s| f64::from(*s) * f64::from(*s)).sum();
83 total / pcm.len() as f64
84}
85
86#[must_use]
91pub fn wav_header(sample_rate: u32, sample_count: usize) -> [u8; WAV_HEADER_BYTES] {
92 let data_bytes = (sample_count * usize::from(BITS_PER_SAMPLE / 8)) as u32;
93 let byte_rate = sample_rate * u32::from(CHANNELS) * u32::from(BITS_PER_SAMPLE / 8);
94 let block_align = CHANNELS * (BITS_PER_SAMPLE / 8);
95
96 let mut header = [0u8; WAV_HEADER_BYTES];
97 header[0..4].copy_from_slice(b"RIFF");
98 header[4..8].copy_from_slice(&(36u32.saturating_add(data_bytes)).to_le_bytes());
100 header[8..12].copy_from_slice(b"WAVE");
101 header[12..16].copy_from_slice(b"fmt ");
102 header[16..20].copy_from_slice(&16u32.to_le_bytes()); header[20..22].copy_from_slice(&1u16.to_le_bytes()); header[22..24].copy_from_slice(&CHANNELS.to_le_bytes());
105 header[24..28].copy_from_slice(&sample_rate.to_le_bytes());
106 header[28..32].copy_from_slice(&byte_rate.to_le_bytes());
107 header[32..34].copy_from_slice(&block_align.to_le_bytes());
108 header[34..36].copy_from_slice(&BITS_PER_SAMPLE.to_le_bytes());
109 header[36..40].copy_from_slice(b"data");
110 header[40..44].copy_from_slice(&data_bytes.to_le_bytes());
111 header
112}
113
114#[must_use]
119pub fn encode_wav(pcm: &[f32], sample_rate: u32) -> Vec<u8> {
120 let mut bytes = Vec::with_capacity(WAV_HEADER_BYTES + pcm.len() * 2);
121 bytes.extend_from_slice(&wav_header(sample_rate, pcm.len()));
122 for sample in pcm {
123 bytes.extend_from_slice(&sample_to_i16(*sample).to_le_bytes());
124 }
125 bytes
126}
127
128pub struct WavWriter<W: Write + Seek> {
138 sink: Option<W>,
144 sample_rate: u32,
145 samples_written: usize,
146}
147
148impl<W: Write + Seek> WavWriter<W> {
149 pub fn new(mut sink: W, sample_rate: u32) -> io::Result<Self> {
155 sink.write_all(&wav_header(sample_rate, 0))?;
156 Ok(Self {
157 sink: Some(sink),
158 sample_rate,
159 samples_written: 0,
160 })
161 }
162
163 pub fn write_samples(&mut self, pcm: &[f32]) -> io::Result<()> {
170 let mut bytes = Vec::with_capacity(pcm.len() * 2);
173 for sample in pcm {
174 bytes.extend_from_slice(&sample_to_i16(*sample).to_le_bytes());
175 }
176 let sink = self
177 .sink
178 .as_mut()
179 .ok_or_else(|| io::Error::other("WavWriter already finished"))?;
180 sink.write_all(&bytes)?;
181 self.samples_written += pcm.len();
182 Ok(())
183 }
184
185 #[must_use]
187 pub const fn samples_written(&self) -> usize {
188 self.samples_written
189 }
190
191 #[must_use]
193 pub const fn duration_millis(&self) -> u64 {
194 if self.sample_rate == 0 {
195 return 0;
196 }
197 (self.samples_written as u64) * 1000 / (self.sample_rate as u64)
198 }
199
200 pub fn finish(mut self) -> io::Result<W> {
206 self.finalize_header()?;
207 self.sink
208 .take()
209 .ok_or_else(|| io::Error::other("WavWriter already finished"))
210 }
211
212 fn finalize_header(&mut self) -> io::Result<()> {
213 let header = wav_header(self.sample_rate, self.samples_written);
214 let Some(sink) = self.sink.as_mut() else {
215 return Ok(());
216 };
217 sink.seek(SeekFrom::Start(0))?;
218 sink.write_all(&header)?;
219 sink.seek(SeekFrom::End(0))?;
220 sink.flush()
221 }
222}
223
224impl<W: Write + Seek> Drop for WavWriter<W> {
225 fn drop(&mut self) {
232 if self.sink.is_some() {
234 let _ = self.finalize_header();
235 }
236 }
237}
238
239#[cfg(test)]
240mod tests {
241 use super::*;
242 use std::io::Cursor;
243
244 #[test]
245 fn frame_count_maps_to_the_codec_sample_rate() {
246 assert_eq!(SAMPLES_PER_FRAME * 25, SAMPLE_RATE_HZ as usize * 2);
249 assert_eq!(samples_for_frames(1), 1_920);
250 assert_eq!(samples_for_frames(125), SAMPLE_RATE_HZ as usize * 10);
251 }
252
253 #[test]
254 fn conversion_clamps_before_scaling() {
255 assert_eq!(sample_to_i16(1.2), i16::MAX);
258 assert_eq!(sample_to_i16(-1.2), -i16::MAX);
259 assert_eq!(sample_to_i16(1.0), i16::MAX);
260 assert_eq!(sample_to_i16(-1.0), -i16::MAX);
261 assert_eq!(sample_to_i16(0.0), 0);
262 }
263
264 #[test]
265 fn conversion_rounds_rather_than_truncates() {
266 let half_step = 0.5 / 32_767.0;
268 assert_eq!(sample_to_i16(half_step), 1);
269 assert_eq!(sample_to_i16(-half_step), -1);
270 }
271
272 #[test]
273 fn non_finite_becomes_silence_not_noise() {
274 assert_eq!(sample_to_i16(f32::NAN), 0);
275 assert_eq!(sample_to_i16(f32::INFINITY), 0);
276 assert_eq!(sample_to_i16(f32::NEG_INFINITY), 0);
277 }
278
279 #[test]
280 fn energy_separates_silence_from_audio() {
281 assert_eq!(mean_square_energy(&[]), 0.0);
282 assert_eq!(mean_square_energy(&[0.0; 64]), 0.0);
283 assert!(mean_square_energy(&[0.5; 64]) > 0.2);
284 }
285
286 #[test]
287 fn the_header_describes_exactly_the_payload() {
288 let pcm = vec![0.25f32; 1_920];
289 let wav = encode_wav(&pcm, SAMPLE_RATE_HZ);
290 assert_eq!(wav.len(), WAV_HEADER_BYTES + pcm.len() * 2);
291 assert_eq!(&wav[0..4], b"RIFF");
292 assert_eq!(&wav[8..12], b"WAVE");
293 assert_eq!(&wav[36..40], b"data");
294
295 let declared_data = u32::from_le_bytes(wav[40..44].try_into().expect("data size"));
296 let actual_data = (wav.len() - WAV_HEADER_BYTES) as u32;
297 assert_eq!(
298 declared_data, actual_data,
299 "data size must match the payload"
300 );
301
302 let declared_riff = u32::from_le_bytes(wav[4..8].try_into().expect("riff size"));
303 assert_eq!(declared_riff, 36 + actual_data, "RIFF size must agree");
304
305 let rate = u32::from_le_bytes(wav[24..28].try_into().expect("rate"));
306 assert_eq!(rate, SAMPLE_RATE_HZ);
307 let channels = u16::from_le_bytes(wav[22..24].try_into().expect("channels"));
308 assert_eq!(channels, 1, "the model is mono");
309 let bits = u16::from_le_bytes(wav[34..36].try_into().expect("bits"));
310 assert_eq!(bits, 16);
311 }
312
313 #[test]
314 fn a_streamed_file_is_byte_identical_to_the_offline_encoding() {
315 let pcm: Vec<f32> = (0..1_920)
318 .map(|i| (i as f32 / 1_920.0 * std::f32::consts::TAU).sin() * 0.5)
319 .collect();
320
321 let mut writer = WavWriter::new(Cursor::new(Vec::new()), SAMPLE_RATE_HZ).expect("header");
322 for packet in pcm.chunks(480) {
323 writer.write_samples(packet).expect("packet");
324 }
325 assert_eq!(writer.samples_written(), pcm.len());
326 assert_eq!(writer.duration_millis(), 80);
327 let streamed = writer.finish().expect("finish").into_inner();
328
329 assert_eq!(streamed, encode_wav(&pcm, SAMPLE_RATE_HZ));
330 }
331
332 #[test]
333 fn a_truncated_run_still_finalises_a_valid_header() {
334 let mut writer = WavWriter::new(Cursor::new(Vec::new()), SAMPLE_RATE_HZ).expect("header");
337 writer.write_samples(&[0.5f32; 960]).expect("packet");
338 let file = writer.finish().expect("finish").into_inner();
339
340 let declared = u32::from_le_bytes(file[40..44].try_into().expect("data size"));
341 assert_eq!(declared, 960 * 2);
342 assert_eq!(file.len(), WAV_HEADER_BYTES + 960 * 2);
343 }
344
345 #[test]
346 fn dropping_without_finish_still_patches_the_length() {
347 let mut sink = Cursor::new(Vec::new());
349 {
350 let mut writer = WavWriter::new(&mut sink, SAMPLE_RATE_HZ).expect("header");
351 writer.write_samples(&[0.25f32; 128]).expect("packet");
352 }
354 let file = sink.into_inner();
355 let declared = u32::from_le_bytes(file[40..44].try_into().expect("data size"));
356 assert_eq!(declared, 128 * 2, "Drop must finalise the length");
357 }
358
359 #[test]
360 fn an_empty_run_is_a_valid_zero_length_wav() {
361 let wav = encode_wav(&[], SAMPLE_RATE_HZ);
362 assert_eq!(wav.len(), WAV_HEADER_BYTES);
363 assert_eq!(u32::from_le_bytes(wav[40..44].try_into().expect("size")), 0);
364 assert_eq!(u32::from_le_bytes(wav[4..8].try_into().expect("riff")), 36);
365 }
366}