use std::io::{self, Seek, SeekFrom, Write};
pub const SAMPLE_RATE_HZ: u32 = 24_000;
pub const SAMPLES_PER_FRAME: usize = 1_920;
pub const CHANNELS: u16 = 1;
pub const BITS_PER_SAMPLE: u16 = 16;
pub const WAV_HEADER_BYTES: usize = 44;
#[must_use]
pub const fn samples_for_frames(frames: usize) -> usize {
frames * SAMPLES_PER_FRAME
}
#[must_use]
pub fn sample_to_i16(sample: f32) -> i16 {
if !sample.is_finite() {
return 0;
}
let clamped = sample.clamp(-1.0, 1.0);
(clamped * 32_767.0).round() as i16
}
#[must_use]
pub fn pcm_f32_to_i16(pcm: &[f32]) -> Vec<i16> {
pcm.iter().copied().map(sample_to_i16).collect()
}
#[must_use]
pub fn mean_square_energy(pcm: &[f32]) -> f64 {
if pcm.is_empty() {
return 0.0;
}
let total: f64 = pcm.iter().map(|s| f64::from(*s) * f64::from(*s)).sum();
total / pcm.len() as f64
}
#[must_use]
pub fn wav_header(sample_rate: u32, sample_count: usize) -> [u8; WAV_HEADER_BYTES] {
let data_bytes = (sample_count * usize::from(BITS_PER_SAMPLE / 8)) as u32;
let byte_rate = sample_rate * u32::from(CHANNELS) * u32::from(BITS_PER_SAMPLE / 8);
let block_align = CHANNELS * (BITS_PER_SAMPLE / 8);
let mut header = [0u8; WAV_HEADER_BYTES];
header[0..4].copy_from_slice(b"RIFF");
header[4..8].copy_from_slice(&(36u32.saturating_add(data_bytes)).to_le_bytes());
header[8..12].copy_from_slice(b"WAVE");
header[12..16].copy_from_slice(b"fmt ");
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());
header[24..28].copy_from_slice(&sample_rate.to_le_bytes());
header[28..32].copy_from_slice(&byte_rate.to_le_bytes());
header[32..34].copy_from_slice(&block_align.to_le_bytes());
header[34..36].copy_from_slice(&BITS_PER_SAMPLE.to_le_bytes());
header[36..40].copy_from_slice(b"data");
header[40..44].copy_from_slice(&data_bytes.to_le_bytes());
header
}
#[must_use]
pub fn encode_wav(pcm: &[f32], sample_rate: u32) -> Vec<u8> {
let mut bytes = Vec::with_capacity(WAV_HEADER_BYTES + pcm.len() * 2);
bytes.extend_from_slice(&wav_header(sample_rate, pcm.len()));
for sample in pcm {
bytes.extend_from_slice(&sample_to_i16(*sample).to_le_bytes());
}
bytes
}
pub struct WavWriter<W: Write + Seek> {
sink: Option<W>,
sample_rate: u32,
samples_written: usize,
}
impl<W: Write + Seek> WavWriter<W> {
pub fn new(mut sink: W, sample_rate: u32) -> io::Result<Self> {
sink.write_all(&wav_header(sample_rate, 0))?;
Ok(Self {
sink: Some(sink),
sample_rate,
samples_written: 0,
})
}
pub fn write_samples(&mut self, pcm: &[f32]) -> io::Result<()> {
let mut bytes = Vec::with_capacity(pcm.len() * 2);
for sample in pcm {
bytes.extend_from_slice(&sample_to_i16(*sample).to_le_bytes());
}
let sink = self
.sink
.as_mut()
.ok_or_else(|| io::Error::other("WavWriter already finished"))?;
sink.write_all(&bytes)?;
self.samples_written += pcm.len();
Ok(())
}
#[must_use]
pub const fn samples_written(&self) -> usize {
self.samples_written
}
#[must_use]
pub const fn duration_millis(&self) -> u64 {
if self.sample_rate == 0 {
return 0;
}
(self.samples_written as u64) * 1000 / (self.sample_rate as u64)
}
pub fn finish(mut self) -> io::Result<W> {
self.finalize_header()?;
self.sink
.take()
.ok_or_else(|| io::Error::other("WavWriter already finished"))
}
fn finalize_header(&mut self) -> io::Result<()> {
let header = wav_header(self.sample_rate, self.samples_written);
let Some(sink) = self.sink.as_mut() else {
return Ok(());
};
sink.seek(SeekFrom::Start(0))?;
sink.write_all(&header)?;
sink.seek(SeekFrom::End(0))?;
sink.flush()
}
}
impl<W: Write + Seek> Drop for WavWriter<W> {
fn drop(&mut self) {
if self.sink.is_some() {
let _ = self.finalize_header();
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::io::Cursor;
#[test]
fn frame_count_maps_to_the_codec_sample_rate() {
assert_eq!(SAMPLES_PER_FRAME * 25, SAMPLE_RATE_HZ as usize * 2);
assert_eq!(samples_for_frames(1), 1_920);
assert_eq!(samples_for_frames(125), SAMPLE_RATE_HZ as usize * 10);
}
#[test]
fn conversion_clamps_before_scaling() {
assert_eq!(sample_to_i16(1.2), i16::MAX);
assert_eq!(sample_to_i16(-1.2), -i16::MAX);
assert_eq!(sample_to_i16(1.0), i16::MAX);
assert_eq!(sample_to_i16(-1.0), -i16::MAX);
assert_eq!(sample_to_i16(0.0), 0);
}
#[test]
fn conversion_rounds_rather_than_truncates() {
let half_step = 0.5 / 32_767.0;
assert_eq!(sample_to_i16(half_step), 1);
assert_eq!(sample_to_i16(-half_step), -1);
}
#[test]
fn non_finite_becomes_silence_not_noise() {
assert_eq!(sample_to_i16(f32::NAN), 0);
assert_eq!(sample_to_i16(f32::INFINITY), 0);
assert_eq!(sample_to_i16(f32::NEG_INFINITY), 0);
}
#[test]
fn energy_separates_silence_from_audio() {
assert_eq!(mean_square_energy(&[]), 0.0);
assert_eq!(mean_square_energy(&[0.0; 64]), 0.0);
assert!(mean_square_energy(&[0.5; 64]) > 0.2);
}
#[test]
fn the_header_describes_exactly_the_payload() {
let pcm = vec![0.25f32; 1_920];
let wav = encode_wav(&pcm, SAMPLE_RATE_HZ);
assert_eq!(wav.len(), WAV_HEADER_BYTES + pcm.len() * 2);
assert_eq!(&wav[0..4], b"RIFF");
assert_eq!(&wav[8..12], b"WAVE");
assert_eq!(&wav[36..40], b"data");
let declared_data = u32::from_le_bytes(wav[40..44].try_into().expect("data size"));
let actual_data = (wav.len() - WAV_HEADER_BYTES) as u32;
assert_eq!(
declared_data, actual_data,
"data size must match the payload"
);
let declared_riff = u32::from_le_bytes(wav[4..8].try_into().expect("riff size"));
assert_eq!(declared_riff, 36 + actual_data, "RIFF size must agree");
let rate = u32::from_le_bytes(wav[24..28].try_into().expect("rate"));
assert_eq!(rate, SAMPLE_RATE_HZ);
let channels = u16::from_le_bytes(wav[22..24].try_into().expect("channels"));
assert_eq!(channels, 1, "the model is mono");
let bits = u16::from_le_bytes(wav[34..36].try_into().expect("bits"));
assert_eq!(bits, 16);
}
#[test]
fn a_streamed_file_is_byte_identical_to_the_offline_encoding() {
let pcm: Vec<f32> = (0..1_920)
.map(|i| (i as f32 / 1_920.0 * std::f32::consts::TAU).sin() * 0.5)
.collect();
let mut writer = WavWriter::new(Cursor::new(Vec::new()), SAMPLE_RATE_HZ).expect("header");
for packet in pcm.chunks(480) {
writer.write_samples(packet).expect("packet");
}
assert_eq!(writer.samples_written(), pcm.len());
assert_eq!(writer.duration_millis(), 80);
let streamed = writer.finish().expect("finish").into_inner();
assert_eq!(streamed, encode_wav(&pcm, SAMPLE_RATE_HZ));
}
#[test]
fn a_truncated_run_still_finalises_a_valid_header() {
let mut writer = WavWriter::new(Cursor::new(Vec::new()), SAMPLE_RATE_HZ).expect("header");
writer.write_samples(&[0.5f32; 960]).expect("packet");
let file = writer.finish().expect("finish").into_inner();
let declared = u32::from_le_bytes(file[40..44].try_into().expect("data size"));
assert_eq!(declared, 960 * 2);
assert_eq!(file.len(), WAV_HEADER_BYTES + 960 * 2);
}
#[test]
fn dropping_without_finish_still_patches_the_length() {
let mut sink = Cursor::new(Vec::new());
{
let mut writer = WavWriter::new(&mut sink, SAMPLE_RATE_HZ).expect("header");
writer.write_samples(&[0.25f32; 128]).expect("packet");
}
let file = sink.into_inner();
let declared = u32::from_le_bytes(file[40..44].try_into().expect("data size"));
assert_eq!(declared, 128 * 2, "Drop must finalise the length");
}
#[test]
fn an_empty_run_is_a_valid_zero_length_wav() {
let wav = encode_wav(&[], SAMPLE_RATE_HZ);
assert_eq!(wav.len(), WAV_HEADER_BYTES);
assert_eq!(u32::from_le_bytes(wav[40..44].try_into().expect("size")), 0);
assert_eq!(u32::from_le_bytes(wav[4..8].try_into().expect("riff")), 36);
}
}