#![allow(dead_code)]
use anyhow::Result;
pub const TARGET_SAMPLE_RATE: u32 = 16_000;
#[derive(Debug, Clone)]
pub struct PcmAudio {
pub samples: Vec<f32>,
pub sample_rate: u32,
}
impl PcmAudio {
pub fn new(samples: Vec<f32>, sample_rate: u32) -> Self {
Self {
samples,
sample_rate,
}
}
pub fn duration_secs(&self) -> f32 {
if self.sample_rate == 0 {
return 0.0;
}
self.samples.len() as f32 / self.sample_rate as f32
}
pub fn peak(&self) -> f32 {
self.samples.iter().fold(0.0f32, |acc, s| acc.max(s.abs()))
}
pub fn rms(&self) -> f32 {
if self.samples.is_empty() {
return 0.0;
}
let sum_sq: f32 = self.samples.iter().map(|s| s * s).sum();
(sum_sq / self.samples.len() as f32).sqrt()
}
}
pub fn decode_to_pcm(bytes: &[u8]) -> Result<PcmAudio> {
if bytes.len() < 44 {
anyhow::bail!("decode_to_pcm: bytes too short to be a WAV header");
}
if &bytes[0..4] != b"RIFF" || &bytes[8..12] != b"WAVE" {
anyhow::bail!(
"decode_to_pcm: not a RIFF/WAVE container; container demux \
(MP3/OGG/FLAC) lands when symphonia is wired"
);
}
let audio_format = u16::from_le_bytes([bytes[20], bytes[21]]);
if audio_format != 1 {
anyhow::bail!(
"decode_to_pcm: WAV audio_format = {audio_format} (only PCM=1 supported here)"
);
}
let num_channels = u16::from_le_bytes([bytes[22], bytes[23]]) as usize;
let sample_rate = u32::from_le_bytes([bytes[24], bytes[25], bytes[26], bytes[27]]);
let bits_per_sample = u16::from_le_bytes([bytes[34], bytes[35]]) as usize;
if num_channels == 0 || sample_rate == 0 {
anyhow::bail!(
"decode_to_pcm: invalid header (channels={num_channels}, rate={sample_rate})"
);
}
let mut cursor = 12usize;
let mut data_offset = None;
let mut data_len = 0usize;
while cursor + 8 <= bytes.len() {
let id = &bytes[cursor..cursor + 4];
let len = u32::from_le_bytes([
bytes[cursor + 4],
bytes[cursor + 5],
bytes[cursor + 6],
bytes[cursor + 7],
]) as usize;
if id == b"data" {
data_offset = Some(cursor + 8);
data_len = len;
break;
}
cursor += 8 + len;
}
let data_start =
data_offset.ok_or_else(|| anyhow::anyhow!("decode_to_pcm: no `data` chunk"))?;
let data = &bytes[data_start..(data_start + data_len).min(bytes.len())];
let bytes_per_sample = bits_per_sample / 8;
if !(bytes_per_sample == 1 || bytes_per_sample == 2) {
anyhow::bail!(
"decode_to_pcm: bits_per_sample={bits_per_sample} not supported (8/16-bit PCM only)"
);
}
let frame_size = bytes_per_sample * num_channels;
if frame_size == 0 || data.len() % frame_size != 0 {
anyhow::bail!(
"decode_to_pcm: misaligned data block ({} bytes, frame_size={frame_size})",
data.len()
);
}
let num_frames = data.len() / frame_size;
let mut mono: Vec<f32> = Vec::with_capacity(num_frames);
for frame in data.chunks_exact(frame_size) {
let mut sum = 0.0f32;
for ch_offset in 0..num_channels {
let s = if bytes_per_sample == 1 {
let raw = frame[ch_offset] as i16;
(raw - 128) as f32 / 128.0
} else {
let i = ch_offset * 2;
let raw = i16::from_le_bytes([frame[i], frame[i + 1]]);
raw as f32 / i16::MAX as f32
};
sum += s;
}
mono.push(sum / num_channels as f32);
}
let mut audio = PcmAudio::new(mono, sample_rate);
if audio.sample_rate != TARGET_SAMPLE_RATE {
audio = resample_linear(audio, TARGET_SAMPLE_RATE);
}
Ok(audio)
}
pub fn resample_linear(input: PcmAudio, target_rate: u32) -> PcmAudio {
if input.sample_rate == target_rate || input.samples.is_empty() {
return PcmAudio::new(input.samples, target_rate);
}
let ratio = input.sample_rate as f64 / target_rate as f64;
let out_len = ((input.samples.len() as f64) / ratio).floor() as usize;
let mut out = Vec::with_capacity(out_len);
for i in 0..out_len {
let src_pos = i as f64 * ratio;
let src_idx = src_pos.floor() as usize;
let frac = (src_pos - src_idx as f64) as f32;
let s0 = input.samples[src_idx];
let s1 = input.samples.get(src_idx + 1).copied().unwrap_or(s0);
out.push(s0 * (1.0 - frac) + s1 * frac);
}
PcmAudio::new(out, target_rate)
}
pub fn spectral_subtraction_denoise(audio: &PcmAudio, head_ms: u32) -> PcmAudio {
let head_samples = ((head_ms as u64) * (audio.sample_rate as u64) / 1000) as usize;
let head_samples = head_samples.min(audio.samples.len());
if head_samples == 0 {
return audio.clone();
}
let head = &audio.samples[..head_samples];
let noise_rms = {
let sum_sq: f32 = head.iter().map(|s| s * s).sum();
(sum_sq / head.len() as f32).sqrt()
};
let denoised: Vec<f32> = audio
.samples
.iter()
.map(|s| {
let abs_s = s.abs();
if abs_s <= noise_rms {
s * 0.1
} else {
let scale = (abs_s - noise_rms) / abs_s.max(f32::EPSILON);
s * scale
}
})
.collect();
PcmAudio::new(denoised, audio.sample_rate)
}
pub fn bandpass_filter(audio: &PcmAudio, low_hz: f32, high_hz: f32) -> PcmAudio {
if audio.samples.is_empty() {
return audio.clone();
}
let sr = audio.sample_rate as f32;
let hp = highpass_1pole(&audio.samples, sr, low_hz);
let lp = lowpass_1pole(&hp, sr, high_hz);
PcmAudio::new(lp, audio.sample_rate)
}
fn highpass_1pole(input: &[f32], sample_rate: f32, cutoff_hz: f32) -> Vec<f32> {
let rc = 1.0 / (2.0 * std::f32::consts::PI * cutoff_hz);
let dt = 1.0 / sample_rate;
let alpha = rc / (rc + dt);
let mut out = Vec::with_capacity(input.len());
let mut prev_in = 0.0f32;
let mut prev_out = 0.0f32;
for &x in input {
let y = alpha * (prev_out + x - prev_in);
out.push(y);
prev_in = x;
prev_out = y;
}
out
}
fn lowpass_1pole(input: &[f32], sample_rate: f32, cutoff_hz: f32) -> Vec<f32> {
let rc = 1.0 / (2.0 * std::f32::consts::PI * cutoff_hz);
let dt = 1.0 / sample_rate;
let alpha = dt / (rc + dt);
let mut out = Vec::with_capacity(input.len());
let mut prev_out = 0.0f32;
for &x in input {
let y = prev_out + alpha * (x - prev_out);
out.push(y);
prev_out = y;
}
out
}
pub fn peak_normalise(audio: &PcmAudio, target_dbfs: f32) -> PcmAudio {
let peak = audio.peak();
if peak == 0.0 {
return audio.clone();
}
let target_linear = 10f32.powf(target_dbfs / 20.0);
let gain = target_linear / peak;
let scaled: Vec<f32> = audio.samples.iter().map(|s| s * gain).collect();
PcmAudio::new(scaled, audio.sample_rate)
}
pub fn time_stretch(audio: &PcmAudio, factor: f32) -> PcmAudio {
if (factor - 1.0).abs() < 1e-3 || audio.samples.is_empty() {
return audio.clone();
}
let frame_size = (0.040 * audio.sample_rate as f32) as usize; let hop_in = (frame_size / 2) as f32;
let hop_out = hop_in / factor;
let mut out: Vec<f32> = Vec::with_capacity((audio.samples.len() as f32 / factor) as usize);
let mut read_pos = 0.0f32;
while (read_pos as usize) + frame_size < audio.samples.len() {
let start = read_pos as usize;
let end = start + frame_size;
out.extend_from_slice(&audio.samples[start..end]);
read_pos += hop_out;
}
PcmAudio::new(out, audio.sample_rate)
}
pub fn preprocess_for_stt(bytes: &[u8]) -> Result<PcmAudio> {
let raw = decode_to_pcm(bytes)?;
let denoised = spectral_subtraction_denoise(&raw, 200);
let bp = bandpass_filter(&denoised, 80.0, 3500.0);
let normalised = peak_normalise(&bp, -1.0);
Ok(normalised)
}
pub fn encode_wav_pcm16(audio: &PcmAudio) -> Vec<u8> {
let num_samples = audio.samples.len();
let bytes_per_sample = 2;
let num_channels: u16 = 1;
let byte_rate = audio.sample_rate * (bytes_per_sample as u32) * (num_channels as u32);
let block_align = bytes_per_sample as u16 * num_channels;
let data_size = num_samples * bytes_per_sample;
let riff_size = 36 + data_size;
let mut out = Vec::with_capacity(44 + data_size);
out.extend_from_slice(b"RIFF");
out.extend_from_slice(&(riff_size as u32).to_le_bytes());
out.extend_from_slice(b"WAVE");
out.extend_from_slice(b"fmt ");
out.extend_from_slice(&16u32.to_le_bytes()); out.extend_from_slice(&1u16.to_le_bytes()); out.extend_from_slice(&num_channels.to_le_bytes());
out.extend_from_slice(&audio.sample_rate.to_le_bytes());
out.extend_from_slice(&byte_rate.to_le_bytes());
out.extend_from_slice(&block_align.to_le_bytes());
out.extend_from_slice(&((bytes_per_sample * 8) as u16).to_le_bytes());
out.extend_from_slice(b"data");
out.extend_from_slice(&(data_size as u32).to_le_bytes());
for s in &audio.samples {
let clamped = s.clamp(-1.0, 1.0);
let i = (clamped * i16::MAX as f32) as i16;
out.extend_from_slice(&i.to_le_bytes());
}
out
}
#[cfg(test)]
mod tests {
use super::*;
fn synth_wav(samples: &[i16], sample_rate: u32) -> Vec<u8> {
let num_samples = samples.len();
let bytes_per_sample = 2;
let data_size = num_samples * bytes_per_sample;
let riff_size = 36 + data_size;
let mut out = Vec::with_capacity(44 + data_size);
out.extend_from_slice(b"RIFF");
out.extend_from_slice(&(riff_size as u32).to_le_bytes());
out.extend_from_slice(b"WAVE");
out.extend_from_slice(b"fmt ");
out.extend_from_slice(&16u32.to_le_bytes());
out.extend_from_slice(&1u16.to_le_bytes());
out.extend_from_slice(&1u16.to_le_bytes());
out.extend_from_slice(&sample_rate.to_le_bytes());
out.extend_from_slice(&(sample_rate * 2).to_le_bytes());
out.extend_from_slice(&2u16.to_le_bytes());
out.extend_from_slice(&16u16.to_le_bytes());
out.extend_from_slice(b"data");
out.extend_from_slice(&(data_size as u32).to_le_bytes());
for s in samples {
out.extend_from_slice(&s.to_le_bytes());
}
out
}
#[test]
fn decode_to_pcm_rejects_too_short_input() {
let err = decode_to_pcm(b"short").unwrap_err().to_string();
assert!(err.contains("too short"));
}
#[test]
fn decode_to_pcm_rejects_non_riff_container() {
let mut bytes = vec![0u8; 100];
bytes[..4].copy_from_slice(b"NOPE");
let err = decode_to_pcm(&bytes).unwrap_err().to_string();
assert!(err.contains("RIFF"));
}
#[test]
fn decode_to_pcm_round_trips_synthetic_wav() {
let samples: Vec<i16> = (0..1000)
.map(|i| ((i as f32 * 0.05).sin() * 16000.0) as i16)
.collect();
let wav = synth_wav(&samples, TARGET_SAMPLE_RATE);
let pcm = decode_to_pcm(&wav).expect("decode synth WAV");
assert_eq!(pcm.sample_rate, TARGET_SAMPLE_RATE);
assert_eq!(pcm.samples.len(), samples.len());
for (i, (got, want)) in pcm.samples.iter().zip(&samples).enumerate() {
let want_f = *want as f32 / i16::MAX as f32;
assert!(
(got - want_f).abs() < 1e-4,
"mismatch at {i}: {got} vs {want_f}"
);
}
}
#[test]
fn decode_to_pcm_resamples_non_target_rate_to_target() {
let samples: Vec<i16> = (0..400)
.map(|i| ((i as f32 * 0.1).sin() * 8000.0) as i16)
.collect();
let wav = synth_wav(&samples, 8_000);
let pcm = decode_to_pcm(&wav).expect("decode 8kHz");
assert_eq!(pcm.sample_rate, TARGET_SAMPLE_RATE);
assert!(pcm.samples.len() >= samples.len() * 19 / 10);
assert!(pcm.samples.len() <= samples.len() * 21 / 10);
}
#[test]
fn pcm_duration_secs_handles_zero_sample_rate() {
let p = PcmAudio {
samples: vec![0.0; 100],
sample_rate: 0,
};
assert_eq!(p.duration_secs(), 0.0);
}
#[test]
fn pcm_peak_returns_max_abs_sample() {
let p = PcmAudio::new(vec![0.0, -0.5, 0.3, -0.7, 0.2], 16000);
assert!((p.peak() - 0.7).abs() < 1e-6);
}
#[test]
fn pcm_rms_returns_root_mean_square() {
let p = PcmAudio::new(vec![1.0; 4], 16000);
assert!((p.rms() - 1.0).abs() < 1e-6);
let p = PcmAudio::new(vec![0.0; 4], 16000);
assert_eq!(p.rms(), 0.0);
let p = PcmAudio::new(vec![], 16000);
assert_eq!(p.rms(), 0.0);
}
#[test]
fn peak_normalise_brings_peak_to_target_dbfs() {
let p = PcmAudio::new(vec![0.1, -0.2, 0.05], 16000);
let normalised = peak_normalise(&p, -1.0);
let expected_peak = 10f32.powf(-1.0 / 20.0); assert!((normalised.peak() - expected_peak).abs() < 1e-4);
}
#[test]
fn peak_normalise_handles_silent_input_without_div_by_zero() {
let silent = PcmAudio::new(vec![0.0; 100], 16000);
let out = peak_normalise(&silent, -1.0);
assert_eq!(out.peak(), 0.0);
assert_eq!(out.samples.len(), 100);
}
#[test]
fn bandpass_filter_attenuates_dc_offset() {
let p = PcmAudio::new(vec![0.5; 16_000], 16_000);
let filtered = bandpass_filter(&p, 80.0, 3500.0);
let tail_rms = {
let tail = &filtered.samples[8000..];
let sum_sq: f32 = tail.iter().map(|s| s * s).sum();
(sum_sq / tail.len() as f32).sqrt()
};
assert!(
tail_rms < 0.05,
"DC offset survived bandpass: tail RMS = {tail_rms}"
);
}
#[test]
fn spectral_subtraction_attenuates_constant_noise_below_signal() {
let mut samples = vec![0.05f32; 16_000];
samples[8000] = 0.5;
let p = PcmAudio::new(samples, 16_000);
let denoised = spectral_subtraction_denoise(&p, 200);
assert!(
denoised.samples[8000].abs() > 0.4,
"spike was attenuated: {}",
denoised.samples[8000]
);
let tail_pre = &denoised.samples[1000..7000];
let pre_rms = {
let sum_sq: f32 = tail_pre.iter().map(|s| s * s).sum();
(sum_sq / tail_pre.len() as f32).sqrt()
};
assert!(
pre_rms < 0.05,
"noise floor not attenuated: pre-spike RMS = {pre_rms}"
);
}
#[test]
fn time_stretch_factor_one_returns_unchanged() {
let p = PcmAudio::new(vec![0.1, 0.2, 0.3], 16000);
let out = time_stretch(&p, 1.0);
assert_eq!(out.samples, p.samples);
}
#[test]
fn time_stretch_speeds_up_when_factor_below_one() {
let p = PcmAudio::new(vec![0.1; 16_000], 16_000);
let out = time_stretch(&p, 0.5);
assert!(
out.samples.len() < p.samples.len(),
"factor=0.5 should shorten: {} vs {}",
out.samples.len(),
p.samples.len()
);
}
#[test]
fn encode_wav_pcm16_produces_valid_header() {
let p = PcmAudio::new(vec![0.0, 0.5, -0.5, 0.25], 16000);
let wav = encode_wav_pcm16(&p);
assert_eq!(&wav[..4], b"RIFF");
assert_eq!(&wav[8..12], b"WAVE");
assert_eq!(&wav[12..16], b"fmt ");
assert_eq!(&wav[36..40], b"data");
assert_eq!(wav.len(), 52);
}
#[test]
fn encode_then_decode_round_trips_pcm() {
let original = PcmAudio::new(
(0..1000).map(|i| (i as f32 * 0.05).sin() * 0.5).collect(),
16_000,
);
let wav = encode_wav_pcm16(&original);
let decoded = decode_to_pcm(&wav).expect("round-trip decode");
assert_eq!(decoded.sample_rate, original.sample_rate);
assert_eq!(decoded.samples.len(), original.samples.len());
for (a, b) in decoded.samples.iter().zip(original.samples.iter()) {
assert!((a - b).abs() < 1e-3, "round-trip diff {a} vs {b}");
}
}
#[test]
fn encode_wav_clamps_overdriven_samples_instead_of_overflowing() {
let p = PcmAudio::new(vec![2.0, -2.0, 1.0, -1.0], 16000);
let wav = encode_wav_pcm16(&p);
let decoded = decode_to_pcm(&wav).unwrap();
assert!((decoded.samples[0] - 1.0).abs() < 1e-3);
assert!((decoded.samples[1] - -1.0).abs() < 1e-3);
}
#[test]
fn preprocess_for_stt_runs_full_pipeline_without_panic() {
let raw_samples: Vec<i16> = (0..16_000 * 2)
.map(|i| {
let signal = (i as f32 * 0.05).sin() * 8000.0;
let noise = ((i * 7919) % 100 - 50) as f32 * 80.0;
(signal + noise) as i16
})
.collect();
let wav = synth_wav(&raw_samples, 16_000);
let processed = preprocess_for_stt(&wav).expect("pipeline must not panic");
assert_eq!(processed.sample_rate, TARGET_SAMPLE_RATE);
assert!(!processed.samples.is_empty());
let expected_peak = 10f32.powf(-1.0 / 20.0);
assert!(
(processed.peak() - expected_peak).abs() < 0.05,
"peak normalisation off: {}",
processed.peak()
);
}
#[test]
fn preprocess_for_stt_rejects_invalid_input_with_actionable_error() {
let err = preprocess_for_stt(b"not a wav").unwrap_err().to_string();
assert!(err.contains("too short") || err.contains("RIFF"));
}
}