Skip to main content

na_mpeg2_decoder/audio/
mpa.rs

1use crate::error::{AvError, Result};
2
3use symphonia::core::audio::SampleBuffer;
4use symphonia::core::codecs::{
5    CodecParameters, Decoder, DecoderOptions, CODEC_TYPE_MP1, CODEC_TYPE_MP2, CODEC_TYPE_MP3,
6};
7use symphonia::core::errors::Error as SymphError;
8use symphonia::core::formats::Packet;
9
10#[derive(Clone)]
11pub struct MpaAudioChunk {
12    pub pts_ms: i64,
13    pub sample_rate: u32,
14    pub channels: u16,
15    pub samples: Vec<f32>,
16}
17
18#[derive(Default)]
19pub struct MpaAudioDecoder {
20    buf: Vec<u8>,
21
22    dec: Option<Box<dyn Decoder>>,
23    sample_buf: Option<SampleBuffer<f32>>,
24
25    // Best-effort PTS tracking.
26    next_pts_ms: Option<i64>,
27
28    // Symphonia packet track id (arbitrary but must be consistent).
29    track_id: u32,
30}
31
32impl MpaAudioDecoder {
33    pub fn new() -> Self {
34        Self {
35            buf: Vec::new(),
36            dec: None,
37            sample_buf: None,
38            next_pts_ms: None,
39            track_id: 0,
40        }
41    }
42
43    pub fn push_with<F>(&mut self, data: &[u8], pts_ms: Option<i64>, mut on_chunk: F) -> Result<()>
44    where
45        F: FnMut(MpaAudioChunk),
46    {
47        if let Some(pts) = pts_ms {
48            self.next_pts_ms = Some(pts);
49        }
50
51        self.buf.extend_from_slice(data);
52
53        let mut pos = 0usize;
54        while pos + 4 <= self.buf.len() {
55            let Some(h) = MpaHeader::parse(&self.buf[pos..]) else {
56                pos += 1;
57                continue;
58            };
59
60            if pos + h.frame_len > self.buf.len() {
61                break;
62            }
63
64            // Avoid borrowing self.buf while calling into self (decoder state).
65            let pkt_owned = self.buf[pos..pos + h.frame_len].to_vec();
66            pos += h.frame_len;
67
68            let pts0 = self.next_pts_ms.unwrap_or(0);
69            self.decode_one_packet(&pkt_owned, pts0, h.codec_type, &mut on_chunk)?;
70        }
71
72        if pos > 0 {
73            self.buf.drain(0..pos);
74        }
75
76        Ok(())
77    }
78
79    fn decode_one_packet<F>(
80        &mut self,
81        pkt_bytes: &[u8],
82        pts_ms: i64,
83        codec_type: symphonia::core::codecs::CodecType,
84        on_chunk: &mut F,
85    ) -> Result<()>
86    where
87        F: FnMut(MpaAudioChunk),
88    {
89        if self.dec.is_none() {
90            let mut cp = CodecParameters::new();
91            cp.for_codec(codec_type);
92
93            let dec = symphonia::default::get_codecs()
94                .make(&cp, &DecoderOptions::default())
95                .map_err(AvError::from)?;
96            self.dec = Some(dec);
97        }
98
99        let pkt = Packet::new_from_boxed_slice(
100            self.track_id,
101            0,
102            0,
103            pkt_bytes.to_vec().into_boxed_slice(),
104        );
105
106        let dec = self.dec.as_mut().expect("decoder must be initialized");
107        match dec.decode(&pkt) {
108            Ok(decoded) => {
109                let spec = *decoded.spec();
110                let duration = decoded.capacity();
111                let duration_u64 = duration as u64;
112
113                let sb = match self.sample_buf.as_mut() {
114                    None => {
115                        self.sample_buf = Some(SampleBuffer::<f32>::new(duration_u64, spec));
116                        self.sample_buf.as_mut().unwrap()
117                    }
118                    Some(sb) => {
119                        if sb.capacity() < duration {
120                            *sb = SampleBuffer::<f32>::new(duration_u64, spec);
121                        }
122                        sb
123                    }
124                };
125
126                sb.copy_interleaved_ref(decoded.clone());
127
128                let channels = spec.channels.count() as u16;
129                let samples = sb.samples().to_vec();
130
131                let sample_rate = spec.rate;
132                on_chunk(MpaAudioChunk {
133                    pts_ms,
134                    sample_rate,
135                    channels,
136                    samples,
137                });
138
139                // Advance PTS based on decoded frames.
140                let frames = decoded.frames() as i64;
141                if frames > 0 && sample_rate > 0 {
142                    let dur_ms = (frames * 1000) / (sample_rate as i64);
143                    self.next_pts_ms = Some(pts_ms + dur_ms);
144                }
145            }
146            Err(SymphError::DecodeError(_)) => {
147                // Best-effort: ignore bad frames.
148            }
149            Err(e) => return Err(e.into()),
150        }
151
152        Ok(())
153    }
154}
155
156#[derive(Clone, Copy)]
157struct MpaHeader {
158    frame_len: usize,
159    codec_type: symphonia::core::codecs::CodecType,
160}
161
162impl MpaHeader {
163    fn parse(buf: &[u8]) -> Option<Self> {
164        if buf.len() < 4 {
165            return None;
166        }
167        let b0 = buf[0];
168        let b1 = buf[1];
169        let b2 = buf[2];
170
171        // Sync.
172        if b0 != 0xFF || (b1 & 0xE0) != 0xE0 {
173            return None;
174        }
175
176        let version_id = (b1 >> 3) & 0x03;
177        let layer_id = (b1 >> 1) & 0x03;
178        if version_id == 0x01 || layer_id == 0x00 {
179            return None;
180        }
181
182        let bitrate_idx = (b2 >> 4) & 0x0F;
183        let sr_idx = (b2 >> 2) & 0x03;
184        if bitrate_idx == 0 || bitrate_idx == 0x0F || sr_idx == 0x03 {
185            return None;
186        }
187
188        let padding: u32 = ((b2 >> 1) & 0x01) as u32;
189
190        let (sr, is_v1) = match version_id {
191            0x03 => (SAMPLE_RATES_V1[sr_idx as usize], true),
192            0x02 => (SAMPLE_RATES_V2[sr_idx as usize], false),
193            0x00 => (SAMPLE_RATES_V25[sr_idx as usize], false),
194            _ => return None,
195        };
196
197        let (codec_type, bitrate_kbps, frame_len) = match layer_id {
198            0x03 => {
199                // Layer I
200                let br = if is_v1 {
201                    BITRATES_V1_L1[bitrate_idx as usize]
202                } else {
203                    BITRATES_V2_L1[bitrate_idx as usize]
204                };
205                let fl =
206                    (((12u64 * (br as u64) * 1000u64) / (sr as u64)) + (padding as u64)) * 4u64;
207                (CODEC_TYPE_MP1, br, fl as usize)
208            }
209            0x02 => {
210                // Layer II
211                let br = if is_v1 {
212                    BITRATES_V1_L2[bitrate_idx as usize]
213                } else {
214                    BITRATES_V2_L2L3[bitrate_idx as usize]
215                };
216                let fl = ((144u64 * (br as u64) * 1000u64) / (sr as u64)) + (padding as u64);
217                (CODEC_TYPE_MP2, br, fl as usize)
218            }
219            0x01 => {
220                // Layer III
221                let br = if is_v1 {
222                    BITRATES_V1_L3[bitrate_idx as usize]
223                } else {
224                    BITRATES_V2_L2L3[bitrate_idx as usize]
225                };
226                let coeff: u64 = if is_v1 { 144 } else { 72 };
227                let fl = ((coeff * (br as u64) * 1000u64) / (sr as u64)) + (padding as u64);
228                (CODEC_TYPE_MP3, br, fl as usize)
229            }
230            _ => return None,
231        };
232
233        if bitrate_kbps == 0 || frame_len < 4 {
234            return None;
235        }
236
237        Some(Self {
238            frame_len,
239            codec_type,
240        })
241    }
242}
243
244const SAMPLE_RATES_V1: [u32; 3] = [44100, 48000, 32000];
245const SAMPLE_RATES_V2: [u32; 3] = [22050, 24000, 16000];
246const SAMPLE_RATES_V25: [u32; 3] = [11025, 12000, 8000];
247
248const BITRATES_V1_L1: [u32; 16] = [
249    0, 32, 64, 96, 128, 160, 192, 224, 256, 288, 320, 352, 384, 416, 448, 0,
250];
251const BITRATES_V1_L2: [u32; 16] = [
252    0, 32, 48, 56, 64, 80, 96, 112, 128, 160, 192, 224, 256, 320, 384, 0,
253];
254const BITRATES_V1_L3: [u32; 16] = [
255    0, 32, 40, 48, 56, 64, 80, 96, 112, 128, 160, 192, 224, 256, 320, 0,
256];
257
258const BITRATES_V2_L1: [u32; 16] = [
259    0, 32, 48, 56, 64, 80, 96, 112, 128, 144, 160, 176, 192, 224, 256, 0,
260];
261const BITRATES_V2_L2L3: [u32; 16] = [
262    0, 8, 16, 24, 32, 40, 48, 56, 64, 80, 96, 112, 128, 144, 160, 0,
263];