use bytes::{Buf, BufMut, BytesMut};
use std::io;
const FMT_SIZE_MIN: u32 = 16;
const FMT_SIZE_EXTENSIBLE: u32 = 18;
#[derive(Debug, PartialEq)]
pub(crate) struct WavHeader {
pub(crate) size: u32,
pub(crate) fmt_size: u32,
pub(crate) format_tag: u16,
pub(crate) channels: u16,
pub(crate) samples_per_sec: u32,
pub(crate) avg_bytes_per_sec: u32,
pub(crate) block_align: u16,
pub(crate) bits_per_sample: u16,
pub(crate) extension_size: Option<u16>,
pub(crate) extension_fields: Vec<u8>,
pub(crate) pre_fmt_fields: Vec<u8>,
pub(crate) extra_fields: Vec<u8>,
pub(crate) data_size: u32,
}
impl Default for WavHeader {
fn default() -> Self {
WavHeader {
size: 0,
fmt_size: FMT_SIZE_MIN,
format_tag: 1,
channels: 2,
samples_per_sec: 44100,
avg_bytes_per_sec: 176400,
block_align: 4,
bits_per_sample: 16,
extension_size: None,
extension_fields: Vec::new(),
pre_fmt_fields: Vec::new(),
extra_fields: Vec::new(),
data_size: 0,
}
}
}
pub(crate) fn write_wav_header(wav_header: &WavHeader, writer: &mut BytesMut) {
writer.put(&b"RIFF"[..]);
writer.put_u32_le(wav_header.size);
writer.put(&b"WAVE"[..]);
writer.put(&wav_header.pre_fmt_fields[..]);
writer.put(&b"fmt "[..]);
writer.put_u32_le(wav_header.fmt_size);
writer.put_u16_le(wav_header.format_tag);
writer.put_u16_le(wav_header.channels);
writer.put_u32_le(wav_header.samples_per_sec);
writer.put_u32_le(wav_header.avg_bytes_per_sec);
writer.put_u16_le(wav_header.block_align);
writer.put_u16_le(wav_header.bits_per_sample);
if let Some(extension_size) = wav_header.extension_size {
writer.put_u16_le(extension_size);
writer.put(&wav_header.extension_fields[..]);
}
if wav_header.fmt_size % 2 == 1 {
writer.put_u8(0);
}
writer.put(&wav_header.extra_fields[..]);
writer.put(&b"data"[..]);
writer.put_u32_le(wav_header.data_size);
}
struct FmtChunk {
fmt_size: u32,
format_tag: u16,
channels: u16,
samples_per_sec: u32,
avg_bytes_per_sec: u32,
block_align: u16,
bits_per_sample: u16,
extension_size: Option<u16>,
extension_fields: Vec<u8>,
}
pub(crate) fn read_wav_header(reader: &mut BytesMut) -> io::Result<WavHeader> {
reader.expect_bytes(b"RIFF")?;
let size = reader.read_u32_le()?;
reader.expect_bytes(b"WAVE")?;
let mut fmt: Option<FmtChunk> = None;
let mut pre_fmt_fields: Vec<u8> = Vec::new();
let mut extra_fields: Vec<u8> = Vec::new();
loop {
let chunk_id: [u8; 4] = reader.read_bytes()?;
if !is_chunk_id(&chunk_id) {
return Err(invalid_data(format!(
"unexpected wav chunk id {:?} while looking for the data chunk",
String::from_utf8_lossy(&chunk_id)
)));
}
let chunk_size = reader.read_u32_le()?;
if chunk_id == *b"data" {
let fmt = fmt.ok_or_else(|| {
invalid_data("wav has no fmt chunk before the data chunk".to_string())
})?;
return Ok(WavHeader {
size,
fmt_size: fmt.fmt_size,
format_tag: fmt.format_tag,
channels: fmt.channels,
samples_per_sec: fmt.samples_per_sec,
avg_bytes_per_sec: fmt.avg_bytes_per_sec,
block_align: fmt.block_align,
bits_per_sample: fmt.bits_per_sample,
extension_size: fmt.extension_size,
extension_fields: fmt.extension_fields,
pre_fmt_fields,
extra_fields,
data_size: chunk_size,
});
}
if chunk_id == *b"fmt " {
if fmt.is_some() {
return Err(invalid_data("wav has more than one fmt chunk".to_string()));
}
fmt = Some(read_fmt_chunk(reader, chunk_size)?);
continue;
}
let padded_size = (chunk_size as u64).next_multiple_of(2);
let payload = reader.read_bytes_vec(padded_size)?;
let target = if fmt.is_some() {
&mut extra_fields
} else {
&mut pre_fmt_fields
};
target.extend_from_slice(&chunk_id);
target.extend_from_slice(&chunk_size.to_le_bytes());
target.extend_from_slice(&payload);
}
}
fn read_fmt_chunk(reader: &mut BytesMut, fmt_size: u32) -> io::Result<FmtChunk> {
if fmt_size != FMT_SIZE_MIN && fmt_size < FMT_SIZE_EXTENSIBLE {
return Err(invalid_data(format!(
"wav fmt chunk is {fmt_size} bytes, expected {FMT_SIZE_MIN} or at least {FMT_SIZE_EXTENSIBLE}"
)));
}
let mut chunk =
BytesMut::from(&reader.read_bytes_vec((fmt_size as u64).next_multiple_of(2))?[..]);
let format_tag = chunk.read_u16_le()?;
let channels = chunk.read_u16_le()?;
let samples_per_sec = chunk.read_u32_le()?;
let avg_bytes_per_sec = chunk.read_u32_le()?;
let block_align = chunk.read_u16_le()?;
let bits_per_sample = chunk.read_u16_le()?;
let (extension_size, extension_fields) = if fmt_size >= FMT_SIZE_EXTENSIBLE {
let extension_size = chunk.read_u16_le()?;
let extension_fields = chunk.read_bytes_vec((fmt_size - FMT_SIZE_EXTENSIBLE) as u64)?;
(Some(extension_size), extension_fields)
} else if reader.looks_like_stray_extension_size() {
(Some(reader.read_u16_le()?), Vec::new())
} else {
(None, Vec::new())
};
Ok(FmtChunk {
fmt_size,
format_tag,
channels,
samples_per_sec,
avg_bytes_per_sec,
block_align,
bits_per_sample,
extension_size,
extension_fields,
})
}
fn is_chunk_id(bytes: &[u8; 4]) -> bool {
bytes.iter().all(|b| b.is_ascii_graphic() || *b == b' ')
}
fn invalid_data(message: String) -> io::Error {
io::Error::new(io::ErrorKind::InvalidData, message)
}
trait ReadBytesExt {
fn read_bytes_vec(&mut self, n: u64) -> io::Result<Vec<u8>>;
fn read_bytes<const N: usize>(&mut self) -> io::Result<[u8; N]>;
fn read_u16_le(&mut self) -> io::Result<u16>;
fn read_u32_le(&mut self) -> io::Result<u32>;
fn expect_bytes<const N: usize>(&mut self, expected: &[u8; N]) -> io::Result<()>;
fn looks_like_stray_extension_size(&self) -> bool;
}
impl ReadBytesExt for BytesMut {
fn read_bytes_vec(&mut self, n: u64) -> io::Result<Vec<u8>> {
if n > self.remaining() as u64 {
return Err(invalid_data(format!(
"unexpected end of wav data, wanted {n} bytes but only {} left",
self.remaining()
)));
}
Ok(self.split_to(n as usize).to_vec())
}
fn read_bytes<const N: usize>(&mut self) -> io::Result<[u8; N]> {
let mut arr = [0; N];
if self.remaining() < N {
return Err(invalid_data(format!(
"unexpected end of wav data, wanted {N} bytes but only {} left",
self.remaining()
)));
}
self.copy_to_slice(&mut arr);
Ok(arr)
}
fn read_u16_le(&mut self) -> io::Result<u16> {
Ok(u16::from_le_bytes(self.read_bytes()?))
}
fn read_u32_le(&mut self) -> io::Result<u32> {
Ok(u32::from_le_bytes(self.read_bytes()?))
}
fn expect_bytes<const N: usize>(&mut self, expected: &[u8; N]) -> io::Result<()> {
let bytes: [u8; N] = self.read_bytes()?;
if &bytes != expected {
return Err(invalid_data(format!(
"expected {:?} in wav data but found {:?}",
String::from_utf8_lossy(expected),
String::from_utf8_lossy(&bytes)
)));
}
Ok(())
}
fn looks_like_stray_extension_size(&self) -> bool {
if self.remaining() < 6 {
return false;
}
let head: [u8; 4] = self[0..4].try_into().unwrap();
let shifted: [u8; 4] = self[2..6].try_into().unwrap();
!is_chunk_id(&head) && is_chunk_id(&shifted)
}
}
#[cfg(test)]
mod test {
use super::*;
use nom::AsBytes;
use pretty_assertions::assert_eq;
#[test]
fn test_read_write_wav_header() {
let data = include_bytes!("../../testdata/fx_coin_converted.wav");
let mut bytes_mut_in = BytesMut::from(data.as_bytes());
let header_read = read_wav_header(&mut bytes_mut_in).unwrap();
let mut bytes_mut_out = BytesMut::new();
write_wav_header(&header_read, &mut bytes_mut_out);
assert_eq!(data[..78], bytes_mut_out[..78]);
}
#[test]
fn test_write_read_wav_header() {
let header = WavHeader {
size: 120 + 36,
fmt_size: 16,
format_tag: 1,
channels: 1,
samples_per_sec: 44100,
avg_bytes_per_sec: 88200,
block_align: 2,
bits_per_sample: 16,
extension_size: None,
extension_fields: Vec::new(),
pre_fmt_fields: Vec::new(),
extra_fields: Vec::new(),
data_size: 120,
};
let mut bytes_mut = BytesMut::new();
write_wav_header(&header, &mut bytes_mut);
let header_read = read_wav_header(&mut bytes_mut).unwrap();
assert_eq!(header, header_read);
}
#[test]
fn test_write_read_wav_header_pcm_float() {
let header = WavHeader {
size: 120 + 36,
fmt_size: 18,
format_tag: 3,
channels: 1,
samples_per_sec: 44100,
avg_bytes_per_sec: 88200,
block_align: 2,
bits_per_sample: 16,
extension_size: Some(0),
extension_fields: Vec::new(),
pre_fmt_fields: Vec::new(),
extra_fields: Vec::new(),
data_size: 120,
};
let mut bytes_mut = BytesMut::new();
write_wav_header(&header, &mut bytes_mut);
let header_read = read_wav_header(&mut bytes_mut).unwrap();
assert_eq!(header, header_read);
}
#[test]
fn test_read_wav_header_pcm_with_extension_size() {
let header = WavHeader {
size: 40000 + 38,
fmt_size: 18,
format_tag: 1,
channels: 1,
samples_per_sec: 22050,
avg_bytes_per_sec: 44100,
block_align: 2,
bits_per_sample: 16,
extension_size: Some(0),
extension_fields: Vec::new(),
pre_fmt_fields: Vec::new(),
extra_fields: Vec::new(),
data_size: 40000,
};
let mut bytes_mut = BytesMut::new();
write_wav_header(&header, &mut bytes_mut);
let header_read = read_wav_header(&mut bytes_mut).unwrap();
assert_eq!(header, header_read);
}
#[test]
fn test_read_wav_header_fmt_extension_fields() {
let header = WavHeader {
size: 120 + 60,
fmt_size: 40,
format_tag: 0xFFFE,
channels: 1,
samples_per_sec: 22050,
avg_bytes_per_sec: 44100,
block_align: 2,
bits_per_sample: 16,
extension_size: Some(22),
extension_fields: (0u8..22).collect(),
pre_fmt_fields: Vec::new(),
extra_fields: Vec::new(),
data_size: 120,
};
let mut bytes_mut = BytesMut::new();
write_wav_header(&header, &mut bytes_mut);
let header_read = read_wav_header(&mut bytes_mut).unwrap();
assert_eq!(header, header_read);
}
#[test]
fn test_read_wav_header_legacy_vpin_extension_size() {
let mut data = BytesMut::new();
data.put(&b"RIFF"[..]);
data.put_u32_le(156);
data.put(&b"WAVE"[..]);
data.put(&b"fmt "[..]);
data.put_u32_le(16);
data.put_u16_le(3);
data.put_u16_le(1);
data.put_u32_le(44100);
data.put_u32_le(88200);
data.put_u16_le(2);
data.put_u16_le(16);
data.put_u16_le(0);
data.put(&b"data"[..]);
data.put_u32_le(120);
let header_read = read_wav_header(&mut data).unwrap();
assert_eq!(header_read.format_tag, 3);
assert_eq!(header_read.extension_size, Some(0));
assert_eq!(header_read.data_size, 120);
}
#[test]
fn test_read_write_wav_header_chunks_before_fmt() {
let mut data = BytesMut::new();
data.put(&b"RIFF"[..]);
data.put_u32_le(180);
data.put(&b"WAVE"[..]);
data.put(&b"JUNK"[..]);
data.put_u32_le(4);
data.put(&b"\0\0\0\0"[..]);
data.put(&b"bext"[..]);
data.put_u32_le(3);
data.put(&b"abc"[..]);
data.put_u8(0); data.put(&b"fmt "[..]);
data.put_u32_le(16);
data.put_u16_le(1);
data.put_u16_le(1);
data.put_u32_le(22050);
data.put_u32_le(44100);
data.put_u16_le(2);
data.put_u16_le(16);
data.put(&b"data"[..]);
data.put_u32_le(120);
let expected = data.to_vec();
let header_read = read_wav_header(&mut data).unwrap();
assert_eq!(header_read.samples_per_sec, 22050);
assert_eq!(header_read.data_size, 120);
assert_eq!(
header_read.pre_fmt_fields,
b"JUNK\x04\x00\x00\x00\0\0\0\0bext\x03\x00\x00\x00abc\x00"
);
assert!(header_read.extra_fields.is_empty());
let mut written = BytesMut::new();
write_wav_header(&header_read, &mut written);
assert_eq!(expected, written.to_vec());
}
#[test]
fn test_read_wav_header_missing_fmt_chunk() {
let mut data = BytesMut::new();
data.put(&b"RIFF"[..]);
data.put_u32_le(28);
data.put(&b"WAVE"[..]);
data.put(&b"JUNK"[..]);
data.put_u32_le(4);
data.put(&b"\0\0\0\0"[..]);
data.put(&b"data"[..]);
data.put_u32_le(120);
let error = read_wav_header(&mut data).unwrap_err();
assert_eq!(error.kind(), io::ErrorKind::InvalidData);
assert!(error.to_string().contains("no fmt chunk"), "{error}");
}
#[test]
fn test_read_wav_header_odd_sized_chunk() {
let mut data = BytesMut::new();
data.put(&b"RIFF"[..]);
data.put_u32_le(168);
data.put(&b"WAVE"[..]);
data.put(&b"fmt "[..]);
data.put_u32_le(16);
data.put_u16_le(1);
data.put_u16_le(1);
data.put_u32_le(22050);
data.put_u32_le(44100);
data.put_u16_le(2);
data.put_u16_le(16);
data.put(&b"cue "[..]);
data.put_u32_le(3);
data.put(&b"abc"[..]);
data.put_u8(0); data.put(&b"data"[..]);
data.put_u32_le(120);
let header_read = read_wav_header(&mut data).unwrap();
assert_eq!(header_read.data_size, 120);
assert_eq!(header_read.extra_fields, b"cue \x03\x00\x00\x00abc\x00");
}
#[test]
fn test_read_wav_header_bogus_chunk_size() {
let mut data = BytesMut::new();
data.put(&b"RIFF"[..]);
data.put_u32_le(48);
data.put(&b"WAVE"[..]);
data.put(&b"fmt "[..]);
data.put_u32_le(16);
data.put_u16_le(1);
data.put_u16_le(1);
data.put_u32_le(22050);
data.put_u32_le(44100);
data.put_u16_le(2);
data.put_u16_le(16);
data.put(&b"cue "[..]);
data.put_u32_le(u32::MAX);
data.put(&b"data"[..]);
data.put_u32_le(120);
let error = read_wav_header(&mut data).unwrap_err();
assert_eq!(error.kind(), io::ErrorKind::InvalidData);
}
#[test]
fn test_read_wav_header_truncated() {
let mut data = BytesMut::from(&b"RIFFxxxxWAVEfmt "[..]);
let error = read_wav_header(&mut data).unwrap_err();
assert_eq!(error.kind(), io::ErrorKind::InvalidData);
}
#[test]
fn test_read_wav_header_not_a_wav() {
let mut data = BytesMut::from(&b"OggS0123456789abcdef"[..]);
let error = read_wav_header(&mut data).unwrap_err();
assert_eq!(error.kind(), io::ErrorKind::InvalidData);
}
}