use core::fmt;
use std::io::Read;
use std::path::Path;
#[cfg(feature = "tnc")]
use crate::tnc::{DefaultTncReceiver, OwnedFrame, TncConfig, TncStats};
use crate::types::{SAMPLE_RATE_MAX, SAMPLE_RATE_MIN, SampleRate};
#[derive(Debug)]
pub enum WavError {
UnsupportedFormat {
channels: u16,
bits_per_sample: u16,
float: bool,
},
UnsupportedRate {
hz: u32,
},
Wav(hound::Error),
Config(crate::ConfigError),
RateRequired,
RateContradiction {
header_hz: u32,
given_hz: u32,
},
}
impl fmt::Display for WavError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match *self {
WavError::UnsupportedFormat {
channels,
bits_per_sample,
float,
} => write!(
f,
"got {channels} channel(s), {bits_per_sample} bits, {} samples; \
16-bit mono integer PCM is required",
if float { "float" } else { "integer" }
),
WavError::UnsupportedRate { hz } => write!(
f,
"got {hz} Hz, supported: {SAMPLE_RATE_MIN}..={SAMPLE_RATE_MAX} Hz"
),
WavError::Wav(ref e) => write!(f, "WAV codec: {e}"),
WavError::Config(ref e) => write!(f, "configuration: {e}"),
WavError::RateRequired => write!(
f,
"raw PCM carries no sample-rate header; a sample rate is required"
),
WavError::RateContradiction {
header_hz,
given_hz,
} => write!(
f,
"the given sample rate ({given_hz} Hz) contradicts the WAV header \
({header_hz} Hz)"
),
}
}
}
impl std::error::Error for WavError {}
impl From<hound::Error> for WavError {
fn from(e: hound::Error) -> Self {
WavError::Wav(e)
}
}
impl From<crate::ConfigError> for WavError {
fn from(e: crate::ConfigError) -> Self {
WavError::Config(e)
}
}
pub fn check_spec(spec: &hound::WavSpec) -> Result<SampleRate, WavError> {
if spec.channels != 1
|| spec.bits_per_sample != 16
|| spec.sample_format != hound::SampleFormat::Int
{
return Err(WavError::UnsupportedFormat {
channels: spec.channels,
bits_per_sample: spec.bits_per_sample,
float: spec.sample_format == hound::SampleFormat::Float,
});
}
SampleRate::new(spec.sample_rate).map_err(|_| WavError::UnsupportedRate {
hz: spec.sample_rate,
})
}
#[cfg(feature = "tnc")]
pub fn decode_frames<P, F>(path: P, sink: F) -> Result<TncStats, WavError>
where
P: AsRef<Path>,
F: FnMut(OwnedFrame) -> bool,
{
let mut reader = hound::WavReader::open(path)?;
let rate = check_spec(&reader.spec())?;
run_receiver(
rate,
reader.samples::<i16>().map(|s| s.map_err(WavError::from)),
sink,
)
}
pub type Replayed<R> = std::io::Chain<std::io::Cursor<Vec<u8>>, R>;
pub enum SniffedPcm<R: Read> {
Wav {
rate: SampleRate,
reader: hound::WavReader<Replayed<R>>,
},
Raw {
rate: SampleRate,
reader: Replayed<R>,
},
}
impl<R: Read> SniffedPcm<R> {
#[must_use]
pub fn rate(&self) -> SampleRate {
match *self {
SniffedPcm::Wav { rate, .. } | SniffedPcm::Raw { rate, .. } => rate,
}
}
}
pub fn sniff_pcm<R: Read>(
mut reader: R,
rate: Option<SampleRate>,
) -> Result<SniffedPcm<R>, WavError> {
let mut head = [0u8; 4];
let mut got = 0usize;
while got < head.len() {
match reader.read(&mut head[got..]) {
Ok(0) => break,
Ok(n) => got += n,
Err(e) if e.kind() == std::io::ErrorKind::Interrupted => {}
Err(e) => return Err(WavError::Wav(e.into())),
}
}
let replay = std::io::Cursor::new(head[..got].to_vec()).chain(reader);
if got == head.len() && head == *b"RIFF" {
let wav = hound::WavReader::new(replay)?;
let header = check_spec(&wav.spec())?;
if let Some(given) = rate
&& given.hz() != header.hz()
{
return Err(WavError::RateContradiction {
header_hz: header.hz(),
given_hz: given.hz(),
});
}
return Ok(SniffedPcm::Wav {
rate: header,
reader: wav,
});
}
let rate = rate.ok_or(WavError::RateRequired)?;
Ok(SniffedPcm::Raw {
rate,
reader: replay,
})
}
#[cfg(feature = "tnc")]
pub fn decode_sniffed<R, F>(input: SniffedPcm<R>, sink: F) -> Result<TncStats, WavError>
where
R: Read,
F: FnMut(OwnedFrame) -> bool,
{
match input {
SniffedPcm::Wav { rate, mut reader } => run_receiver(
rate,
reader.samples::<i16>().map(|s| s.map_err(WavError::from)),
sink,
),
SniffedPcm::Raw { rate, mut reader } => {
run_receiver(rate, raw_s16le_samples(&mut reader), sink)
}
}
}
#[cfg(feature = "tnc")]
fn run_receiver<F>(
rate: SampleRate,
samples: impl Iterator<Item = Result<i16, WavError>>,
mut sink: F,
) -> Result<TncStats, WavError>
where
F: FnMut(OwnedFrame) -> bool,
{
let config = TncConfig::bell_202(rate)?;
let mut rx = DefaultTncReceiver::new(config)?;
for sample in samples {
if let Some(frame) = rx.push_i16(sample?) {
let Ok(owned) = OwnedFrame::new(&frame) else {
continue;
};
if !sink(owned) {
break;
}
}
}
Ok(rx.stats())
}
#[cfg(feature = "tnc")]
fn raw_s16le_samples<R: Read>(reader: &mut R) -> impl Iterator<Item = Result<i16, WavError>> + '_ {
let mut done = false;
std::iter::from_fn(move || {
if done {
return None;
}
let mut bytes = [0u8; 2];
let mut filled = 0usize;
while filled < bytes.len() {
match reader.read(&mut bytes[filled..]) {
Ok(0) if filled == 0 => {
done = true;
return None;
}
Ok(0) => {
done = true;
let e = std::io::Error::new(
std::io::ErrorKind::UnexpectedEof,
"truncated sample (odd byte count) at EOF",
);
return Some(Err(WavError::Wav(e.into())));
}
Ok(n) => filled += n,
Err(e) if e.kind() == std::io::ErrorKind::Interrupted => {}
Err(e) => {
done = true;
return Some(Err(WavError::Wav(e.into())));
}
}
}
Some(Ok(i16::from_le_bytes(bytes)))
})
}
#[cfg(test)]
mod tests {
use super::*;
fn wav_bytes(hz: u32, samples: &[i16]) -> Vec<u8> {
let spec = hound::WavSpec {
channels: 1,
sample_rate: hz,
bits_per_sample: 16,
sample_format: hound::SampleFormat::Int,
};
let mut cursor = std::io::Cursor::new(Vec::new());
let mut writer = hound::WavWriter::new(&mut cursor, spec).unwrap();
for &s in samples {
writer.write_sample(s).unwrap();
}
writer.finalize().unwrap();
cursor.into_inner()
}
#[test]
fn sniff_honors_wav_header() {
let bytes = wav_bytes(44_100, &[1, -2, 3]);
let sniffed = sniff_pcm(std::io::Cursor::new(bytes), None).unwrap();
assert_eq!(sniffed.rate().hz(), 44_100);
let SniffedPcm::Wav { mut reader, .. } = sniffed else {
panic!("WAV bytes classified as raw");
};
let samples: Vec<i16> = reader.samples::<i16>().map(|s| s.unwrap()).collect();
assert_eq!(samples, [1, -2, 3], "sniff must not eat header bytes");
}
#[test]
fn sniff_matching_rate_hint_is_accepted() {
let bytes = wav_bytes(44_100, &[0; 4]);
let hint = SampleRate::new(44_100).unwrap();
let sniffed = sniff_pcm(std::io::Cursor::new(bytes), Some(hint)).unwrap();
assert!(matches!(sniffed, SniffedPcm::Wav { .. }));
}
#[test]
fn sniff_rejects_contradicting_rate() {
let bytes = wav_bytes(44_100, &[0; 4]);
let hint = SampleRate::new(48_000).unwrap();
let err = match sniff_pcm(std::io::Cursor::new(bytes), Some(hint)) {
Err(e) => e,
Ok(_) => panic!("contradicting rate hint accepted"),
};
assert!(matches!(
err,
WavError::RateContradiction {
header_hz: 44_100,
given_hz: 48_000,
}
));
}
#[test]
fn sniff_raw_needs_a_rate() {
let err = match sniff_pcm(std::io::Cursor::new(vec![0u8; 64]), None) {
Err(e) => e,
Ok(_) => panic!("raw stream without a rate accepted"),
};
assert!(matches!(err, WavError::RateRequired));
}
#[test]
fn sniff_raw_with_rate_replays_the_prefix() {
let rate = SampleRate::new(48_000).unwrap();
let bytes = vec![1, 0, 2, 0, 3, 0];
let sniffed = sniff_pcm(std::io::Cursor::new(bytes.clone()), Some(rate)).unwrap();
let SniffedPcm::Raw { mut reader, rate } = sniffed else {
panic!("raw bytes classified as WAV");
};
assert_eq!(rate.hz(), 48_000);
let mut back = Vec::new();
reader.read_to_end(&mut back).unwrap();
assert_eq!(back, bytes, "the sniffed prefix must be replayed");
}
#[test]
fn sniff_short_stream_is_raw() {
let rate = SampleRate::new(48_000).unwrap();
let sniffed = sniff_pcm(std::io::Cursor::new(vec![7u8, 0]), Some(rate)).unwrap();
let SniffedPcm::Raw { mut reader, .. } = sniffed else {
panic!("short stream classified as WAV");
};
let mut back = Vec::new();
reader.read_to_end(&mut back).unwrap();
assert_eq!(back, [7, 0]);
}
}