use crate::audio::decode::Decoder;
use crate::error::Error;
use flac_codec::decode::FlacStreamReader;
use std::sync::Arc;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
struct StreamParams {
sample_rate: u32,
channels: u8,
bits_per_sample: u32,
}
#[derive(Clone)]
pub struct FlacDecoder {
expected: Option<StreamParams>,
}
impl FlacDecoder {
pub fn new() -> Self {
Self { expected: None }
}
pub fn with_header(header: &[u8]) -> Result<Self, Error> {
let info = flac_codec::metadata::read_info(header)
.map_err(|e| Error::Protocol(format!("Invalid FLAC codec header: {e}")))?;
Ok(Self {
expected: Some(StreamParams {
sample_rate: info.sample_rate,
channels: info.channels.get(),
bits_per_sample: info.bits_per_sample.into(),
}),
})
}
}
impl Default for FlacDecoder {
fn default() -> Self {
Self::new()
}
}
impl Decoder for FlacDecoder {
fn decode(&self, data: &[u8]) -> Result<Arc<[i32]>, Error> {
if data.is_empty() {
return Err(Error::Protocol("Empty FLAC data".to_string()));
}
let mut remaining: &[u8] = data;
let mut out: Vec<i32> = Vec::new();
while !remaining.is_empty() {
if remaining.len() < 2 || remaining[0] != 0xFF || remaining[1] >> 1 != 0b1111100 {
return Err(Error::Protocol(
"FLAC chunk does not start at a frame boundary".to_string(),
));
}
let mut reader = FlacStreamReader::new(&mut remaining);
let frame = reader
.read()
.map_err(|e| Error::Protocol(format!("FLAC decode error: {e}")))?;
if let Some(expected) = self.expected {
let actual = StreamParams {
sample_rate: frame.sample_rate,
channels: frame.channels,
bits_per_sample: frame.bits_per_sample,
};
if actual != expected {
return Err(Error::Protocol(format!(
"FLAC frame format {actual:?} does not match stream header {expected:?}"
)));
}
}
let shift = 32u32.saturating_sub(frame.bits_per_sample);
out.extend(frame.samples.iter().map(|s| s.wrapping_shl(shift)));
}
Ok(Arc::from(out.into_boxed_slice()))
}
}
#[cfg(test)]
mod tests {
use super::*;
use flac_codec::encode::{FlacStreamWriter, Options};
fn encode_frame(
sample_rate: u32,
channels: u8,
bits_per_sample: u32,
samples: &[i32],
) -> Vec<u8> {
let mut buf = Vec::new();
let mut writer = FlacStreamWriter::new(&mut buf, Options::default());
writer
.write(sample_rate, channels, bits_per_sample, samples)
.unwrap();
drop(writer); buf
}
fn test_signal_16(frames: usize) -> Vec<i32> {
(0..frames)
.flat_map(|i| {
let left = ((i as i32 * 373) % 32767) - 16384;
let right = -(((i as i32 * 151) % 32767) - 16384);
[left, right]
})
.collect()
}
#[test]
fn round_trip_single_frame_16bit() {
let samples = test_signal_16(1024);
let chunk = encode_frame(48000, 2, 16, &samples);
let decoder = FlacDecoder::new();
let decoded = decoder.decode(&chunk).unwrap();
assert_eq!(decoded.len(), samples.len());
for (d, s) in decoded.iter().zip(samples.iter()) {
assert_eq!(*d, s << 16);
}
}
#[test]
fn round_trip_multiple_frames_in_one_chunk() {
let a = test_signal_16(512);
let b = test_signal_16(400);
let mut chunk = encode_frame(48000, 2, 16, &a);
chunk.extend(encode_frame(48000, 2, 16, &b));
let decoder = FlacDecoder::new();
let decoded = decoder.decode(&chunk).unwrap();
assert_eq!(decoded.len(), a.len() + b.len());
let expected: Vec<i32> = a.iter().chain(b.iter()).map(|s| s << 16).collect();
assert_eq!(decoded.as_ref(), expected.as_slice());
}
#[test]
fn chunk_per_frame_stream() {
let decoder = FlacDecoder::new();
for len in [512usize, 256, 1024] {
let samples = test_signal_16(len);
let chunk = encode_frame(44100, 2, 16, &samples);
let decoded = decoder.decode(&chunk).unwrap();
assert_eq!(decoded.len(), samples.len());
assert_eq!(decoded[0], samples[0] << 16);
}
}
#[test]
fn scales_24bit_to_full_i32() {
let samples: Vec<i32> = vec![8_388_607, -8_388_608, 0, 1, -1, 4096, -4096, 42];
let chunk = encode_frame(96000, 1, 24, &samples);
let decoder = FlacDecoder::new();
let decoded = decoder.decode(&chunk).unwrap();
assert_eq!(decoded.len(), samples.len());
for (d, s) in decoded.iter().zip(samples.iter()) {
assert_eq!(*d, s << 8);
}
}
#[test]
fn empty_data_is_error() {
let decoder = FlacDecoder::new();
assert!(decoder.decode(&[]).is_err());
}
#[test]
fn garbage_data_is_error() {
let decoder = FlacDecoder::new();
let garbage: Vec<u8> = (0..1024u32).map(|i| (i % 251) as u8).collect();
assert!(decoder.decode(&garbage).is_err());
}
#[test]
fn fake_sync_garbage_is_error() {
let decoder = FlacDecoder::new();
let mut garbage = vec![0xFF, 0xF8];
garbage.extend((0..1024u32).map(|i| (i % 251) as u8));
assert!(decoder.decode(&garbage).is_err());
}
#[test]
fn leading_garbage_before_valid_frame_is_error() {
let samples = test_signal_16(256);
let frame = encode_frame(48000, 2, 16, &samples);
let mut chunk = vec![0x00, 0x01, 0x02, 0x03];
chunk.extend(&frame);
let decoder = FlacDecoder::new();
assert!(decoder.decode(&chunk).is_err());
}
#[test]
fn garbage_between_frames_is_error() {
let a = test_signal_16(256);
let b = test_signal_16(256);
let mut chunk = encode_frame(48000, 2, 16, &a);
chunk.extend_from_slice(&[0x13, 0x37]);
chunk.extend(encode_frame(48000, 2, 16, &b));
let decoder = FlacDecoder::new();
assert!(decoder.decode(&chunk).is_err());
}
#[test]
fn scales_8bit_to_full_i32() {
let samples: Vec<i32> = vec![127, -128, 0, 1, -1, 64, -64];
let chunk = encode_frame(48000, 1, 8, &samples);
let decoder = FlacDecoder::new();
let decoded = decoder.decode(&chunk).unwrap();
assert_eq!(decoded.len(), samples.len());
for (d, s) in decoded.iter().zip(samples.iter()) {
assert_eq!(*d, s << 24);
}
}
#[test]
fn truncated_frame_is_error() {
let samples = test_signal_16(1024);
let chunk = encode_frame(48000, 2, 16, &samples);
let truncated = &chunk[..chunk.len() - 16];
let decoder = FlacDecoder::new();
assert!(decoder.decode(truncated).is_err());
}
#[test]
fn trailing_garbage_after_valid_frame_is_error() {
let samples = test_signal_16(256);
let mut chunk = encode_frame(48000, 2, 16, &samples);
chunk.extend_from_slice(&[0xDE, 0xAD, 0xBE, 0xEF]);
let decoder = FlacDecoder::new();
assert!(decoder.decode(&chunk).is_err());
}
fn make_header(sample_rate: u32, channels: u8, bits_per_sample: u32) -> Vec<u8> {
use flac_codec::encode::FlacSampleWriter;
use std::io::Cursor;
let mut flac = Cursor::new(Vec::new());
let mut writer = FlacSampleWriter::new(
&mut flac,
Options::default(),
sample_rate,
bits_per_sample,
channels,
None,
)
.unwrap();
writer.write(&vec![0i32; 16 * channels as usize]).unwrap();
writer.finalize().unwrap();
let bytes = flac.into_inner();
bytes[..42].to_vec()
}
#[test]
fn with_header_accepts_matching_frames() {
let header = make_header(48000, 2, 16);
let decoder = FlacDecoder::with_header(&header).unwrap();
let samples = test_signal_16(512);
let chunk = encode_frame(48000, 2, 16, &samples);
let decoded = decoder.decode(&chunk).unwrap();
assert_eq!(decoded.len(), samples.len());
}
#[test]
fn with_header_rejects_mismatched_frames() {
let header = make_header(44100, 2, 16);
let decoder = FlacDecoder::with_header(&header).unwrap();
let samples = test_signal_16(512);
let chunk = encode_frame(48000, 2, 16, &samples);
assert!(decoder.decode(&chunk).is_err());
}
#[test]
fn invalid_header_is_error() {
assert!(FlacDecoder::with_header(b"not a flac header").is_err());
assert!(FlacDecoder::with_header(&[]).is_err());
}
}