use crate::audio::format::{AudioFormat, SampleFormat};
use crate::audio::samples::AudioBuffer;
pub fn detect_audio_format(data: &[u8]) -> AudioFormat {
if data.len() >= 12 && &data[0..4] == b"RIFF" && &data[8..12] == b"WAVE" {
return AudioFormat::Wav;
}
if data.len() >= 4 && &data[0..4] == b"fLaC" {
return AudioFormat::Flac;
}
if data.len() >= 4 && &data[0..4] == b"OggS" {
return AudioFormat::Ogg;
}
if data.len() >= 3 && &data[0..3] == b"ID3" {
return AudioFormat::Mp3;
}
if data.len() >= 2 && data[0] == 0xFF {
if (data[1] & 0xF6) == 0xF0 {
return AudioFormat::Aac;
}
if (data[1] & 0xE0) == 0xE0 && (data[1] & 0x06) != 0x00 && (data[1] & 0x18) != 0x08 {
return AudioFormat::Mp3;
}
}
AudioFormat::Unknown
}
pub fn decode(data: &[u8]) -> Result<AudioBuffer, String> {
let format = detect_audio_format(data);
match format {
AudioFormat::Wav => decode_wav(data),
AudioFormat::Pcm => {
if !data.len().is_multiple_of(4) {
return Err(format!(
"raw PCM data is {} bytes, which is not a whole number of 32-bit \
little-endian f32 samples; pad it to a multiple of 4 bytes",
data.len()
));
}
let samples: Vec<f32> = data
.as_chunks::<4>()
.0
.iter()
.map(|c| f32::from_le_bytes([c[0], c[1], c[2], c[3]]))
.collect();
Ok(AudioBuffer::new(44100, samples, 1))
}
AudioFormat::Mp3 => decode_mp3(data),
AudioFormat::Flac => decode_flac(data),
AudioFormat::Ogg => decode_ogg_vorbis(data),
AudioFormat::Aac => decode_aac(data),
AudioFormat::Opus => decode_opus(data),
AudioFormat::Unknown => Err("Unknown audio format — cannot decode".into()),
}
}
fn decode_wav(data: &[u8]) -> Result<AudioBuffer, String> {
if data.len() < 44 || &data[0..4] != b"RIFF" || &data[8..12] != b"WAVE" {
return Err("Invalid WAV header".into());
}
let mut pos = 12;
let mut sample_rate = 0u32;
let mut channels = 0u16;
let mut bits_per_sample = 0u16;
let mut format_tag = 0u16;
let mut data_chunk: Option<&[u8]> = None;
while pos + 8 <= data.len() {
let chunk_id = &data[pos..pos + 4];
let chunk_size =
u32::from_le_bytes([data[pos + 4], data[pos + 5], data[pos + 6], data[pos + 7]])
as usize;
let chunk_start = pos + 8;
let chunk_end = chunk_start.checked_add(chunk_size).ok_or("WAV chunk size overflow")?;
if chunk_end > data.len() {
return Err(format!(
"WAV chunk {:?} truncated: need {chunk_size} bytes, got {}",
chunk_id,
data.len().saturating_sub(chunk_start)
));
}
let chunk_data = &data[chunk_start..chunk_end];
if chunk_id == b"fmt " && chunk_data.len() >= 16 {
format_tag = u16::from_le_bytes([chunk_data[0], chunk_data[1]]);
channels = u16::from_le_bytes([chunk_data[2], chunk_data[3]]);
sample_rate =
u32::from_le_bytes([chunk_data[4], chunk_data[5], chunk_data[6], chunk_data[7]]);
bits_per_sample = u16::from_le_bytes([chunk_data[14], chunk_data[15]]);
if format_tag == 0xFFFE {
if chunk_data.len() < 40 {
return Err("WAV fmt chunk uses WAVE_FORMAT_EXTENSIBLE but is too short \
(needs 40 bytes to carry the SubFormat GUID)"
.into());
}
format_tag = u16::from_le_bytes([chunk_data[24], chunk_data[25]]);
}
} else if chunk_id == b"data" {
data_chunk = Some(chunk_data);
}
let step = 8usize
.checked_add(chunk_size)
.and_then(|size| size.checked_add(chunk_size % 2))
.ok_or("WAV chunk offset overflow")?;
pos = pos.checked_add(step).ok_or("WAV chunk offset overflow")?;
}
if sample_rate == 0 || channels == 0 {
return Err("WAV fmt chunk is missing or invalid".into());
}
let channels_u8 = u8::try_from(channels)
.map_err(|_| format!("unsupported WAV channel count: {channels} (must be 1..=255)"))?;
let raw_samples = data_chunk.ok_or("No data chunk in WAV")?;
let fmt = sample_format_for_tag(format_tag, bits_per_sample)?;
let bytes_per_sample = fmt.bytes_per_sample();
if raw_samples.len() % bytes_per_sample != 0 {
return Err(format!(
"WAV data chunk is {} bytes, which is not a whole number of {bits_per_sample}-bit \
samples ({} bytes each); the data chunk is truncated",
raw_samples.len(),
bytes_per_sample
));
}
let bytes_per_frame = bytes_per_sample * channels as usize;
if raw_samples.len() % bytes_per_frame != 0 {
return Err(format!(
"WAV data chunk is {} bytes, which is not a whole number of {channels}-channel \
frames ({bytes_per_frame} bytes each); the data chunk holds an incomplete final frame",
raw_samples.len()
));
}
let samples = fmt.to_f32(raw_samples);
let mut buf = AudioBuffer::new(sample_rate.max(1), samples, channels_u8);
buf.original_format = fmt;
Ok(buf)
}
fn sample_format_for_tag(format_tag: u16, bits_per_sample: u16) -> Result<SampleFormat, String> {
match format_tag {
0x0001 => match bits_per_sample {
8 => Ok(SampleFormat::U8),
16 => Ok(SampleFormat::I16),
24 => Ok(SampleFormat::I24),
32 => Ok(SampleFormat::I32),
_ => Err(format!("unsupported PCM bits per sample: {bits_per_sample}")),
},
0x0003 => {
if bits_per_sample == 32 {
Ok(SampleFormat::F32)
} else {
Err(format!(
"unsupported IEEE float bits per sample: {bits_per_sample} \
(only 32-bit float is supported)"
))
}
}
other => Err(format!("unsupported WAV format tag: {other:#06x}")),
}
}
fn decode_mp3(data: &[u8]) -> Result<AudioBuffer, String> {
use minimp3_fixed::Decoder as Mp3Decoder;
let mp3_data = if data.len() > 10 && &data[0..3] == b"ID3" {
let tag_size = ((data[6] as usize) << 21)
| ((data[7] as usize) << 14)
| ((data[8] as usize) << 7)
| (data[9] as usize);
let header_size = 10 + tag_size;
if header_size < data.len() {
&data[header_size..]
} else {
data
}
} else {
data
};
let mut decoder = Mp3Decoder::new(mp3_data);
let mut all_samples: Vec<f32> = Vec::new();
let mut sample_rate = 44100u32;
let mut channels = 2u8;
loop {
match decoder.next_frame() {
Ok(frame) => {
sample_rate = frame.sample_rate as u32;
channels = frame.channels as u8;
for &sample in &frame.data {
all_samples.push(sample as f32 / 32768.0);
}
}
Err(minimp3_fixed::Error::Eof) => break,
Err(e) => {
return Err(format!(
"MP3 frame after {} decoded sample(s) could not be decoded: {e:?}",
all_samples.len()
))
}
}
}
if all_samples.is_empty() {
return Err(format!(
"MP3 decoded zero audio frames from {} bytes: the stream has no usable MPEG frame \
(check that the data is really MP3 and not a raw PCM or container payload)",
data.len()
));
}
let mut buf = AudioBuffer::new(sample_rate, all_samples, channels);
buf.original_format = SampleFormat::I16;
Ok(buf)
}
#[cfg(feature = "symphonia-codecs")]
fn decode_with_symphonia(data: &[u8], format: AudioFormat) -> Result<AudioBuffer, String> {
use symphonia::core::audio::SampleBuffer;
use symphonia::core::codecs::{DecoderOptions, CODEC_TYPE_NULL};
use symphonia::core::formats::FormatOptions;
use symphonia::core::io::MediaSourceStream;
use symphonia::core::meta::MetadataOptions;
use symphonia::core::probe::Hint;
let owned_data = data.to_vec();
let cursor = std::io::Cursor::new(owned_data);
let mss = MediaSourceStream::new(Box::new(cursor), Default::default());
let format_opts = FormatOptions::default();
let metadata_opts = MetadataOptions::default();
let mut hint = Hint::new();
if format == AudioFormat::Flac {
hint.with_extension("flac");
} else if format == AudioFormat::Ogg {
hint.with_extension("ogg");
} else if format == AudioFormat::Aac {
hint.with_extension("aac");
} else if format == AudioFormat::Opus {
hint.with_extension("opus");
}
let probed = symphonia::default::get_probe()
.format(&hint, mss, &format_opts, &metadata_opts)
.map_err(|e| {
format!(
"symphonia could not probe the stream container (unsupported or corrupt format): {e:?}"
)
})?;
let mut format_reader = probed.format;
let track = format_reader
.tracks()
.iter()
.find(|t| t.codec_params.codec != CODEC_TYPE_NULL)
.ok_or_else(|| {
format!(
"the stream has no decodable audio track: symphonia found {} track(s), none with \
a real codec — the container may hold only video or metadata",
format_reader.tracks().len()
)
})?;
let codec_params = track.codec_params.clone();
let track_id = track.id;
let sample_rate = codec_params.sample_rate.unwrap_or(44100);
let channels = codec_params.channels.map(|c| c.count() as u8).unwrap_or(2);
let decode_opts = DecoderOptions::default();
let mut decoder =
symphonia::default::get_codecs().make(&codec_params, &decode_opts).map_err(|e| {
format!(
"symphonia could not build a decoder for codec {codec:?}: {e:?} (the codec may \
require a different `symphonia-*` feature)",
codec = codec_params.codec
)
})?;
let mut all_samples: Vec<f32> = Vec::new();
loop {
let packet = match format_reader.next_packet() {
Ok(packet) => packet,
Err(error) => match classify_symphonia_error(&error) {
RecoverableError::EndOfStream => break,
RecoverableError::Skip => continue,
RecoverableError::Unrecoverable(reason) => {
return Err(format!(
"audio stream could not be read after {} decoded sample(s): {reason}",
all_samples.len()
));
}
},
};
if packet.track_id() != track_id {
continue;
}
let decoded = match decoder.decode(&packet) {
Ok(decoded) => decoded,
Err(error) => match classify_symphonia_error(&error) {
RecoverableError::EndOfStream => break,
RecoverableError::Skip => continue,
RecoverableError::Unrecoverable(reason) => {
return Err(format!(
"audio packet at {} decoded sample(s) could not be decoded: {reason}",
all_samples.len()
));
}
},
};
let spec = *decoded.spec();
let num_frames = decoded.frames() as usize;
if num_frames == 0 {
continue;
}
let mut sample_buf = SampleBuffer::<f32>::new(num_frames as u64, spec);
sample_buf.copy_interleaved_ref(decoded);
all_samples.extend_from_slice(sample_buf.samples());
}
if all_samples.is_empty() {
return Err(format!(
"symphonia decoded zero audio samples from {} bytes: every packet was empty or \
skipped, so there is no PCM to return",
data.len()
));
}
let mut buf = AudioBuffer::new(sample_rate, all_samples, channels);
buf.original_format = SampleFormat::F32;
Ok(buf)
}
#[cfg(feature = "symphonia-codecs")]
#[derive(Debug, Clone, PartialEq, Eq)]
enum RecoverableError {
EndOfStream,
Skip,
Unrecoverable(String),
}
#[cfg(feature = "symphonia-codecs")]
fn classify_symphonia_error(error: &symphonia::core::errors::Error) -> RecoverableError {
use symphonia::core::errors::Error;
match error {
Error::IoError(e) if e.kind() == std::io::ErrorKind::Interrupted => RecoverableError::Skip,
Error::IoError(e) if e.kind() == std::io::ErrorKind::UnexpectedEof => {
RecoverableError::EndOfStream
}
Error::ResetRequired => RecoverableError::Skip,
other => RecoverableError::Unrecoverable(format!("{other:?}")),
}
}
fn decode_flac(data: &[u8]) -> Result<AudioBuffer, String> {
#[cfg(feature = "symphonia-codecs")]
{
decode_with_symphonia(data, AudioFormat::Flac)
}
#[cfg(not(feature = "symphonia-codecs"))]
{
if data.len() < 4 || &data[0..4] != b"fLaC" {
return Err(format!(
"FLAC must start with the magic \"fLaC\", but this {}-byte input starts with \
{:02X?}",
data.len(),
&data[..data.len().min(4)]
));
}
Err("decoding FLAC requires the `symphonia-codecs` feature".to_string())
}
}
fn decode_ogg_vorbis(data: &[u8]) -> Result<AudioBuffer, String> {
#[cfg(feature = "symphonia-codecs")]
{
decode_with_symphonia(data, AudioFormat::Ogg)
}
#[cfg(not(feature = "symphonia-codecs"))]
{
if data.len() < 28 || &data[0..4] != b"OggS" {
return Err(format!(
"OGG must start with the capture pattern \"OggS\", but this {}-byte input \
starts with {:02X?}",
data.len(),
&data[..data.len().min(4)]
));
}
Err("decoding OGG Vorbis requires the `symphonia-codecs` feature".to_string())
}
}
fn decode_aac(data: &[u8]) -> Result<AudioBuffer, String> {
#[cfg(feature = "symphonia-codecs")]
{
decode_with_symphonia(data, AudioFormat::Aac)
}
#[cfg(not(feature = "symphonia-codecs"))]
{
let mut pos = 0usize;
while pos + 7 <= data.len() {
if data[pos] == 0xFF && (data[pos + 1] & 0xF6) == 0xF0 {
let sample_rate_index = ((data[pos + 2] >> 2) & 0x0F) as usize;
let frame_length = (((data[pos + 3] as u16 & 0x03) << 11) as usize)
| ((data[pos + 4] as usize) << 3)
| ((data[pos + 5] >> 5) as usize);
let plausible_frame = sample_rate_index <= 12
&& frame_length >= 7
&& pos + frame_length <= data.len();
if plausible_frame {
return Err("decoding AAC requires the `symphonia-codecs` feature".to_string());
}
}
pos += 1;
}
Err(format!(
"AAC data ({} bytes) contains no valid ADTS frame: every 0xFFF sync candidate \
failed the header checks",
data.len()
))
}
}
fn decode_opus(data: &[u8]) -> Result<AudioBuffer, String> {
#[cfg(feature = "symphonia-codecs")]
{
decode_with_symphonia(data, AudioFormat::Opus)
}
#[cfg(not(feature = "symphonia-codecs"))]
{
if data.len() < 28 || &data[0..4] != b"OggS" {
return Err(format!(
"Opus must be wrapped in an Ogg container starting with \"OggS\", but this \
{}-byte input starts with {:02X?}",
data.len(),
&data[..data.len().min(4)]
));
}
if !data.windows(8).any(|w| w == b"OpusHead") {
return Err(format!(
"Opus stream has no OpusHead identification header in its {} bytes: the Ogg \
container holds no Opus logical bitstream",
data.len()
));
}
Err("decoding Opus requires the `symphonia-codecs` feature".to_string())
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_detect_wav() {
let mut wav = b"RIFF".to_vec();
wav.extend_from_slice(&[0u8; 4]);
wav.extend_from_slice(b"WAVE");
assert_eq!(detect_audio_format(&wav), AudioFormat::Wav);
}
#[test]
fn test_detect_mp3_id3() {
assert_eq!(detect_audio_format(b"ID3xxxx"), AudioFormat::Mp3);
}
#[test]
fn test_detect_flac() {
assert_eq!(detect_audio_format(b"fLaCxxxx"), AudioFormat::Flac);
}
#[test]
fn test_detect_ogg() {
assert_eq!(detect_audio_format(b"OggSxxxx"), AudioFormat::Ogg);
}
#[test]
fn test_detect_unknown() {
assert_eq!(detect_audio_format(b"NotAudio"), AudioFormat::Unknown);
}
#[test]
fn test_detect_mp3_mpeg1_layer3_without_id3() {
assert_eq!(detect_audio_format(&[0xFF, 0xFB, 0x90, 0x64]), AudioFormat::Mp3);
}
#[test]
fn test_detect_mp3_mpeg1_all_layers() {
assert_eq!(detect_audio_format(&[0xFF, 0xFB, 0x00, 0x00]), AudioFormat::Mp3);
assert_eq!(detect_audio_format(&[0xFF, 0xFD, 0x00, 0x00]), AudioFormat::Mp3);
assert_eq!(detect_audio_format(&[0xFF, 0xFF, 0x00, 0x00]), AudioFormat::Mp3);
}
#[test]
fn test_detect_mp3_mpeg2_and_mpeg25() {
assert_eq!(detect_audio_format(&[0xFF, 0xF3, 0x00, 0x00]), AudioFormat::Mp3);
assert_eq!(detect_audio_format(&[0xFF, 0xE3, 0x00, 0x00]), AudioFormat::Mp3);
}
#[test]
fn test_detect_aac_adts() {
assert_eq!(
detect_audio_format(&[0xFF, 0xF1, 0x50, 0x80, 0x00, 0xFF, 0xFC]),
AudioFormat::Aac
);
assert_eq!(detect_audio_format(&[0xFF, 0xF0, 0x00, 0x00]), AudioFormat::Aac);
assert_eq!(detect_audio_format(&[0xFF, 0xF8, 0x00, 0x00]), AudioFormat::Aac);
}
#[test]
fn test_detect_short_or_reserved_headers_are_not_guessed() {
assert_eq!(detect_audio_format(&[0xFF]), AudioFormat::Unknown);
assert_eq!(detect_audio_format(&[]), AudioFormat::Unknown);
assert_eq!(detect_audio_format(&[0xFF, 0xE1, 0x00, 0x00]), AudioFormat::Unknown);
assert_eq!(detect_audio_format(&[0xFF, 0xEF, 0x00, 0x00]), AudioFormat::Unknown);
}
#[test]
fn test_decode_wav_valid() {
let data_size = 100;
let file_size = 36 + data_size;
let mut wav = Vec::new();
wav.extend_from_slice(b"RIFF");
wav.extend_from_slice(&(file_size as u32).to_le_bytes());
wav.extend_from_slice(b"WAVE");
wav.extend_from_slice(b"fmt ");
wav.extend_from_slice(&16u32.to_le_bytes());
wav.extend_from_slice(&1u16.to_le_bytes());
wav.extend_from_slice(&1u16.to_le_bytes());
wav.extend_from_slice(&44100u32.to_le_bytes());
wav.extend_from_slice(&(44100u32 * 2).to_le_bytes());
wav.extend_from_slice(&2u16.to_le_bytes());
wav.extend_from_slice(&16u16.to_le_bytes());
wav.extend_from_slice(b"data");
wav.extend_from_slice(&(data_size as u32).to_le_bytes());
wav.extend_from_slice(&[0u8; 100]);
let result = decode_wav(&wav);
assert!(result.is_ok());
let buf = result.unwrap();
assert_eq!(buf.sample_rate, 44100);
assert_eq!(buf.channels(), 1);
}
#[test]
fn test_decode_wav_rejects_truncated_chunk() {
let mut wav = b"RIFF".to_vec();
wav.extend_from_slice(&42u32.to_le_bytes());
wav.extend_from_slice(b"WAVEfmt ");
wav.extend_from_slice(&16u32.to_le_bytes());
wav.extend_from_slice(&1u16.to_le_bytes());
wav.extend_from_slice(&1u16.to_le_bytes());
wav.extend_from_slice(&44100u32.to_le_bytes());
wav.extend_from_slice(&88200u32.to_le_bytes());
wav.extend_from_slice(&2u16.to_le_bytes());
wav.extend_from_slice(&16u16.to_le_bytes());
wav.extend_from_slice(b"data");
wav.extend_from_slice(&100u32.to_le_bytes());
wav.extend_from_slice(&[0u8; 2]);
let error = decode_wav(&wav).unwrap_err();
assert!(error.contains("truncated"), "unexpected error: {error}");
}
#[test]
fn test_decode_wav_invalid() {
assert!(decode_wav(b"not a wav").is_err());
}
fn build_wav(
format_tag: u16,
channels: u16,
sample_rate: u32,
bits_per_sample: u16,
data: &[u8],
) -> Vec<u8> {
let bytes_per_sample = bits_per_sample / 8;
let block_align = channels * bytes_per_sample;
let byte_rate = sample_rate * block_align as u32;
let fmt_chunk_len: u32 = if format_tag == 0xFFFE { 40 } else { 16 };
let mut wav = Vec::new();
wav.extend_from_slice(b"RIFF");
let riff_size = 4 + (8 + fmt_chunk_len) + (8 + data.len() as u32);
wav.extend_from_slice(&riff_size.to_le_bytes());
wav.extend_from_slice(b"WAVE");
wav.extend_from_slice(b"fmt ");
wav.extend_from_slice(&fmt_chunk_len.to_le_bytes());
wav.extend_from_slice(&format_tag.to_le_bytes());
wav.extend_from_slice(&channels.to_le_bytes());
wav.extend_from_slice(&sample_rate.to_le_bytes());
wav.extend_from_slice(&byte_rate.to_le_bytes());
wav.extend_from_slice(&block_align.to_le_bytes());
wav.extend_from_slice(&bits_per_sample.to_le_bytes());
if format_tag == 0xFFFE {
wav.extend_from_slice(&22u16.to_le_bytes()); wav.extend_from_slice(&bits_per_sample.to_le_bytes()); wav.extend_from_slice(&0u32.to_le_bytes()); wav.extend_from_slice(&[
0x01, 0x00, 0x00, 0x00, 0x00, 0x00, 0x10, 0x00, 0x80, 0x00, 0x00, 0xaa, 0x00, 0x38,
0x9b, 0x71,
]);
}
wav.extend_from_slice(b"data");
wav.extend_from_slice(&(data.len() as u32).to_le_bytes());
wav.extend_from_slice(data);
wav
}
#[test]
fn test_decode_wav_ieee_float() {
let mut data = Vec::new();
for s in [1.0f32, -1.0f32, 0.5f32, 0.0f32] {
data.extend_from_slice(&s.to_le_bytes());
}
let wav = build_wav(0x0003, 1, 44100, 32, &data);
let buf = decode_wav(&wav).unwrap();
assert_eq!(buf.sample_rate, 44100);
assert_eq!(buf.channels(), 1);
assert_eq!(buf.original_format, SampleFormat::F32);
assert!((buf.samples[0] - 1.0).abs() < 1e-6, "got {}", buf.samples[0]);
assert!((buf.samples[1] + 1.0).abs() < 1e-6, "got {}", buf.samples[1]);
assert!((buf.samples[2] - 0.5).abs() < 1e-6, "got {}", buf.samples[2]);
assert_eq!(buf.samples[3], 0.0);
}
#[test]
fn test_decode_wav_rejects_unsupported_format_tag() {
let wav = build_wav(0x0006, 1, 44100, 8, &[0x80]);
let err = decode_wav(&wav).unwrap_err();
assert!(err.contains("format tag"), "got: {err}");
}
#[test]
fn test_decode_wav_rejects_unsupported_float_bit_depth() {
let wav = build_wav(0x0003, 1, 44100, 64, &[0u8; 8]);
let err = decode_wav(&wav).unwrap_err();
assert!(err.contains("IEEE float"), "got: {err}");
}
#[test]
fn test_decode_wav_extensible_pcm() {
let mut data = Vec::new();
for s in [0i16, 32767, -32768] {
data.extend_from_slice(&s.to_le_bytes());
}
let wav = build_wav(0xFFFE, 1, 44100, 16, &data);
let buf = decode_wav(&wav).unwrap();
assert_eq!(buf.original_format, SampleFormat::I16);
assert!((buf.samples[0] - 0.0).abs() < 0.01, "got {}", buf.samples[0]);
assert!((buf.samples[1] - 32767.0 / 32768.0).abs() < 0.01, "got {}", buf.samples[1]);
assert!((buf.samples[2] + 1.0).abs() < 0.01, "got {}", buf.samples[2]);
}
#[test]
fn test_decode_wav_rejects_unsupported_channel_counts_without_narrowing() {
let wav = build_wav(0x0001, 257, 44100, 16, &[0u8; 4]);
let err = decode_wav(&wav).unwrap_err();
assert!(err.contains("channel count"), "got: {err}");
let wav = build_wav(0x0001, 256, 44100, 16, &[0u8; 4]);
let err = decode_wav(&wav).unwrap_err();
assert!(err.contains("channel count"), "got: {err}");
}
#[test]
fn test_decode_wav_rejects_incomplete_final_frame() {
let wav = build_wav(0x0001, 2, 44100, 16, &[0u8, 0]);
let err = decode_wav(&wav).unwrap_err();
assert!(err.contains("incomplete final frame"), "got: {err}");
let wav = build_wav(0x0001, 2, 44100, 16, &[0u8; 4]);
let buf = decode_wav(&wav).unwrap();
assert_eq!(buf.channels(), 2);
assert_eq!(buf.samples.len(), 2);
}
#[test]
fn test_decode_mp3_empty_data_returns_error() {
assert!(decode_mp3(b"").is_err());
}
#[test]
fn test_decode_flac_empty_data_returns_error() {
assert!(decode_flac(b"").is_err());
}
#[test]
fn test_decode_ogg_empty_data_returns_error() {
assert!(decode_ogg_vorbis(b"").is_err());
}
#[test]
fn test_decode_aac_empty_data_returns_error() {
assert!(decode_aac(b"").is_err());
}
#[test]
fn test_decode_opus_empty_data_returns_error() {
assert!(decode_opus(b"").is_err());
}
#[test]
fn test_decode_unknown_format() {
assert!(decode(b"unknown format data").is_err());
}
#[test]
#[cfg(not(feature = "symphonia-codecs"))]
fn test_decode_compressed_formats_report_missing_feature() {
let mut ogg = b"OggS".to_vec();
ogg.resize(48, 0u8);
let mut opus = b"OggS".to_vec();
opus.extend_from_slice(b"OpusHead");
opus.resize(48, 0u8);
let opus_ok = opus.clone();
let cases = vec![
b"fLaC".to_vec(),
ogg,
vec![0xFF, 0xF1, 0x50, 0x80, 0x00, 0xFF, 0xFC],
opus,
];
for data in cases {
let err = decode(&data).unwrap_err();
assert!(
err.contains("symphonia-codecs"),
"expected a `symphonia-codecs` feature error, got: {err}"
);
}
let err = decode_opus(&opus_ok).unwrap_err();
assert!(err.contains("symphonia-codecs"), "got: {err}");
}
#[test]
fn test_decode_mp3_id3_only_no_frames() {
let mut id3 = b"ID3".to_vec();
id3.extend_from_slice(&[0x04, 0x00]); id3.push(0x00); id3.extend_from_slice(&[0, 0, 0, 0]); assert!(decode_mp3(&id3).is_err());
}
#[test]
#[cfg(feature = "symphonia-codecs")]
fn test_decode_with_symphonia_invalid_data_returns_error() {
let result = decode_with_symphonia(b"not valid audio data", AudioFormat::Flac);
assert!(result.is_err());
let err_msg = result.unwrap_err();
assert!(
err_msg.contains("could not probe") || err_msg.contains("no decodable audio track"),
"unexpected error: {err_msg}"
);
}
#[test]
#[cfg(feature = "symphonia-codecs")]
fn test_decode_with_symphonia_wav_succeeds() {
let sample_rate = 44100u32;
let channels: u16 = 1;
let bits_per_sample: u16 = 16;
let bytes_per_sample = (bits_per_sample / 8) as u32;
let block_align: u16 = channels * (bits_per_sample / 8);
let byte_rate = sample_rate * block_align as u32;
let num_samples = 4410usize; let data_size = num_samples as u32 * bytes_per_sample;
let riff_size = 36 + data_size;
let mut wav = Vec::new();
wav.extend_from_slice(b"RIFF");
wav.extend_from_slice(&riff_size.to_le_bytes());
wav.extend_from_slice(b"WAVE");
wav.extend_from_slice(b"fmt ");
wav.extend_from_slice(&16u32.to_le_bytes()); wav.extend_from_slice(&1u16.to_le_bytes()); wav.extend_from_slice(&channels.to_le_bytes());
wav.extend_from_slice(&sample_rate.to_le_bytes());
wav.extend_from_slice(&byte_rate.to_le_bytes());
wav.extend_from_slice(&block_align.to_le_bytes());
wav.extend_from_slice(&bits_per_sample.to_le_bytes());
wav.extend_from_slice(b"data");
wav.extend_from_slice(&data_size.to_le_bytes());
for i in 0..num_samples {
let t = i as f64 / sample_rate as f64;
let sample = (t * 440.0 * 2.0 * std::f64::consts::PI).sin();
let int_val = (sample * 32767.0) as i16;
wav.extend_from_slice(&int_val.to_le_bytes());
}
let result = decode_with_symphonia(&wav, AudioFormat::Wav);
if let Err(ref e) = result {
assert!(e.contains("Symphonia"), "Unexpected error: {}", e);
return;
}
let buf = result.unwrap();
assert_eq!(buf.sample_rate, 44100);
assert_eq!(buf.channels(), 1);
assert!(!buf.samples.is_empty());
let max_sample = buf.samples.iter().map(|&s| s.abs()).fold(0.0f32, f32::max);
assert!(max_sample > 0.0, "Expected non-zero audio samples");
}
#[test]
#[cfg(feature = "symphonia-codecs")]
fn test_decode_flac_symphonia_path_rejects_invalid_data() {
assert!(decode_flac(b"fLaC").is_err());
}
#[test]
#[cfg(feature = "symphonia-codecs")]
fn test_classify_symphonia_error_policy() {
use std::io::{Error as IoError, ErrorKind};
use symphonia::core::errors::Error;
assert_eq!(
classify_symphonia_error(&Error::IoError(IoError::new(ErrorKind::UnexpectedEof, ""))),
RecoverableError::EndOfStream,
"UnexpectedEof means the stream ended cleanly, not a failure"
);
assert_eq!(
classify_symphonia_error(&Error::IoError(IoError::new(ErrorKind::Interrupted, ""))),
RecoverableError::Skip,
"Interrupted is transient and must be retried"
);
assert_eq!(
classify_symphonia_error(&Error::ResetRequired),
RecoverableError::Skip,
"ResetRequired is recoverable"
);
assert!(matches!(
classify_symphonia_error(&Error::DecodeError("corrupt frame")),
RecoverableError::Unrecoverable(_)
));
assert!(matches!(
classify_symphonia_error(&Error::IoError(IoError::new(ErrorKind::Other, ""))),
RecoverableError::Unrecoverable(_)
));
}
}