use crate::audio::format::AudioFormat;
use crate::audio::samples::AudioBuffer;
#[cfg(feature = "video-codecs")]
use super::ffmpeg_encode;
pub fn encode(buffer: &AudioBuffer, format: AudioFormat) -> Result<Vec<u8>, String> {
match format {
AudioFormat::Wav => encode_wav(buffer),
AudioFormat::Pcm => encode_pcm(buffer),
#[cfg(feature = "video-codecs")]
AudioFormat::Mp3
| AudioFormat::Flac
| AudioFormat::Ogg
| AudioFormat::Aac
| AudioFormat::Opus => ffmpeg_encode(buffer, format),
#[cfg(not(feature = "video-codecs"))]
AudioFormat::Mp3
| AudioFormat::Flac
| AudioFormat::Ogg
| AudioFormat::Aac
| AudioFormat::Opus => {
Err(format!("encoding {:?} requires the `video-codecs` feature (FFmpeg)", format))
}
AudioFormat::Unknown => Err("Cannot encode to Unknown format".into()),
}
}
fn encode_pcm(buffer: &AudioBuffer) -> Result<Vec<u8>, String> {
if buffer.sample_rate == 0 {
return Err("Sample rate must be > 0".into());
}
let mut out = Vec::with_capacity(buffer.samples.len() * 4);
for &s in &buffer.samples {
out.extend_from_slice(&s.to_le_bytes());
}
Ok(out)
}
fn encode_wav(buffer: &AudioBuffer) -> Result<Vec<u8>, String> {
if buffer.sample_rate == 0 {
return Err("Sample rate must be > 0".into());
}
let channels = buffer.channels();
if channels == 0 {
return Err("Channel count must be > 0".into());
}
let bits_per_sample: u16 = 16;
let bytes_per_sample: u64 = (bits_per_sample / 8) as u64;
let channels_usize = channels as usize;
if !buffer.samples.len().is_multiple_of(channels_usize) {
return Err(format!(
"WAV encoding requires a whole number of frames: {} samples is not a multiple of the
{channels}-channel frame size (the buffer ends with {} sample(s) of an incomplete
final frame)",
buffer.samples.len(),
buffer.samples.len() % channels_usize,
));
}
let block_align = channels as u64 * bytes_per_sample;
if block_align > u16::MAX as u64 {
return Err(format!(
"WAV block align {block_align} (channels {channels} × {bytes_per_sample} bytes) \
exceeds the 16-bit header field"
));
}
let byte_rate = buffer.sample_rate as u64 * block_align;
if byte_rate > u32::MAX as u64 {
return Err(format!(
"WAV byte rate {byte_rate} (sample_rate {} × block_align {block_align}) exceeds the \
32-bit header field; the sample rate is too high for WAV output",
buffer.sample_rate
));
}
let sample_count = buffer.samples.len() as u64;
let data_size = sample_count.checked_mul(bytes_per_sample).ok_or_else(|| {
format!("WAV data size overflows: {sample_count} samples × {bytes_per_sample} bytes")
})?;
if data_size > u32::MAX as u64 {
return Err(format!(
"WAV data chunk of {data_size} bytes exceeds the 32-bit size field; the buffer is \
too large for WAV output"
));
}
let file_size = 36u64 + data_size;
if file_size > u32::MAX as u64 {
return Err(format!(
"WAV RIFF size {file_size} exceeds the 32-bit header field; the buffer is too large \
for WAV output"
));
}
let data_size = data_size as u32;
let file_size = file_size as u32;
let byte_rate = byte_rate as u32;
let block_align = block_align as u16;
let channels_field = channels as u16;
let mut out = Vec::with_capacity(44 + data_size as usize);
out.extend_from_slice(b"RIFF");
out.extend_from_slice(&file_size.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(&channels_field.to_le_bytes());
out.extend_from_slice(&buffer.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(&bits_per_sample.to_le_bytes());
out.extend_from_slice(b"data");
out.extend_from_slice(&data_size.to_le_bytes());
for &sample in &buffer.samples {
let clamped = sample.clamp(-1.0, 1.0);
let int_val = (clamped * 32767.0) as i16;
out.extend_from_slice(&int_val.to_le_bytes());
}
Ok(out)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_encode_wav() {
let buf = AudioBuffer::new(44100, vec![0.0, 0.5, -0.5, 1.0, -1.0], 1);
let wav = encode_wav(&buf).unwrap();
assert!(wav.starts_with(b"RIFF"));
assert!(wav.len() > 44);
}
#[test]
fn test_encode_wav_empty_buffer() {
let buf = AudioBuffer::new(44100, vec![], 1);
let wav = encode_wav(&buf).unwrap();
assert!(wav.starts_with(b"RIFF"));
}
#[test]
fn test_encode_pcm() {
let buf = AudioBuffer::new(44100, vec![0.0, 0.5, 1.0], 1);
let pcm = encode(&buf, AudioFormat::Pcm).unwrap();
assert_eq!(pcm.len(), 3 * 4); }
#[test]
#[cfg(feature = "video-codecs")]
fn test_encode_all_formats_succeed() {
let samples: Vec<f32> = (0..16384)
.map(|i| {
let phase = (i as f32 / 44100.0 * 440.0 * 2.0 * std::f32::consts::PI).sin();
phase * 0.5
})
.collect();
let buf = AudioBuffer::new(44100, samples, 2);
for format in &[
AudioFormat::Mp3,
AudioFormat::Flac,
AudioFormat::Ogg,
AudioFormat::Aac,
AudioFormat::Opus,
] {
let result = encode(&buf, *format);
assert!(result.is_ok(), "Encoding to {:?} should succeed: {:?}", format, result);
let data = result.unwrap();
assert!(!data.is_empty(), "Encoded {:?} data should not be empty", format);
}
}
#[test]
#[cfg(not(feature = "video-codecs"))]
fn test_encode_compressed_formats_report_missing_feature() {
let samples: Vec<f32> = (0..16384)
.map(|i| {
let phase = (i as f32 / 44100.0 * 440.0 * 2.0 * std::f32::consts::PI).sin();
phase * 0.5
})
.collect();
let buf = AudioBuffer::new(44100, samples, 2);
for format in &[
AudioFormat::Mp3,
AudioFormat::Flac,
AudioFormat::Ogg,
AudioFormat::Aac,
AudioFormat::Opus,
] {
let err = encode(&buf, *format)
.expect_err("encoding a compressed format without `video-codecs` must fail");
assert!(
err.contains("video-codecs"),
"error for {:?} should name the `video-codecs` feature, got: {err}",
format
);
}
}
#[test]
fn test_encode_unknown_returns_error() {
let buf = AudioBuffer::new(44100, vec![], 1);
assert!(encode(&buf, AudioFormat::Unknown).is_err());
}
#[test]
fn encode_wav_rejects_incomplete_final_frame() {
let buf = AudioBuffer::new(44100, vec![0.1, 0.2, 0.3], 2);
let err = encode_wav(&buf).unwrap_err();
assert!(err.contains("whole number of frames"), "got: {err}");
}
#[test]
fn encoded_wav_roundtrips_through_the_library_decoder() {
for channels in 1u8..=4 {
let samples: Vec<f32> = (0..(channels as usize * 8)).map(|i| i as f32 * 0.01).collect();
let buf = AudioBuffer::new(44100, samples, channels);
let wav = encode_wav(&buf).expect("aligned buffer must encode");
let decoded = crate::audio::decoder::decode(&wav)
.unwrap_or_else(|e| panic!("{channels}-channel WAV failed to decode: {e}"));
assert_eq!(decoded.channels(), channels);
assert_eq!(decoded.samples.len(), buf.samples.len());
}
}
#[test]
fn test_encode_wav_rejects_unrepresentable_byte_rate() {
let buf = AudioBuffer::new(u32::MAX, vec![0.0, 0.0], 2);
let err = encode_wav(&buf).unwrap_err();
assert!(err.contains("byte rate"), "got: {err}");
}
#[test]
fn test_encode_wav_rejects_byte_rate_boundary() {
let sample_rate = u32::MAX / 2 + 1; let buf = AudioBuffer::new(sample_rate, vec![0.0, 0.0], 2);
assert!(encode_wav(&buf).is_err());
}
#[test]
fn test_encode_wav_header_is_consistent() {
let buf = AudioBuffer::new(44100, vec![0.0, 0.5, -0.5, 1.0], 1);
let wav = encode_wav(&buf).unwrap();
let data_size = u32::from_le_bytes([wav[40], wav[41], wav[42], wav[43]]);
assert_eq!(data_size, 4 * 2);
assert_eq!(wav.len(), 44 + data_size as usize);
let riff_size = u32::from_le_bytes([wav[4], wav[5], wav[6], wav[7]]);
assert_eq!(riff_size, 36 + data_size);
let byte_rate = u32::from_le_bytes([wav[28], wav[29], wav[30], wav[31]]);
assert_eq!(byte_rate, 44100 * 2);
}
}