Skip to main content

rtc_rtp/codec/av1/
depacketizer.rs

1//! AV1 RTP Depacketizer
2//!
3//! Reads AV1 RTP packets and outputs AV1 low overhead bitstream format.
4//! Based on <https://aomediacodec.github.io/av1-rtp-spec/>
5
6use bytes::{BufMut, Bytes, BytesMut};
7
8use crate::codec::av1::leb128::read_leb128;
9use crate::codec::av1::obu::{
10    OBU_HAS_SIZE_BIT, OBU_TYPE_MASK, OBU_TYPE_TEMPORAL_DELIMITER, OBU_TYPE_TILE_LIST,
11};
12use crate::packetizer::Depacketizer;
13use shared::error::{Error, Result};
14
15// AV1 Aggregation Header bit masks
16const AV1_Z_MASK: u8 = 0b1000_0000;
17const AV1_Y_MASK: u8 = 0b0100_0000;
18const AV1_W_MASK: u8 = 0b0011_0000;
19const AV1_N_MASK: u8 = 0b0000_1000;
20
21/// AV1 RTP Depacketizer
22///
23/// Depacketizes AV1 RTP packets into low overhead bitstream format with obu_size fields.
24#[derive(Default, Debug, Clone)]
25pub struct Av1Depacketizer {
26    /// Buffer for fragmented OBU from previous packet
27    buffer: BytesMut,
28    /// Z flag from aggregation header - first OBU is continuation
29    pub z: bool,
30    /// Y flag from aggregation header - last OBU will continue
31    pub y: bool,
32    /// N flag from aggregation header - new coded video sequence
33    pub n: bool,
34}
35
36impl Av1Depacketizer {
37    pub fn new() -> Self {
38        Self::default()
39    }
40}
41
42impl Depacketizer for Av1Depacketizer {
43    /// Depacketize parses an AV1 RTP payload into OBU stream with obu_size_field.
44    ///
45    /// Reference: <https://aomediacodec.github.io/av1-rtp-spec/>
46    fn depacketize(&mut self, payload: &Bytes) -> Result<Bytes> {
47        if payload.len() <= 1 {
48            return Err(Error::ErrShortPacket);
49        }
50
51        // Parse aggregation header
52        // |Z|Y| W |N|-|-|-|
53        let obu_z = (payload[0] & AV1_Z_MASK) != 0;
54        let obu_y = (payload[0] & AV1_Y_MASK) != 0;
55        let obu_count = (payload[0] & AV1_W_MASK) >> 4;
56        let obu_n = (payload[0] & AV1_N_MASK) != 0;
57
58        self.z = obu_z;
59        self.y = obu_y;
60        self.n = obu_n;
61
62        // Clear buffer on new coded video sequence
63        if obu_n {
64            self.buffer.clear();
65        }
66
67        // Clear buffer if Z is not set but we have buffered data
68        if !obu_z && !self.buffer.is_empty() {
69            self.buffer.clear();
70        }
71
72        let mut result = BytesMut::new();
73        let mut offset = 1; // Skip aggregation header
74        let mut obu_offset = 0;
75
76        while offset < payload.len() {
77            let is_first = obu_offset == 0;
78            let is_last = obu_count != 0 && obu_offset == (obu_count - 1) as usize;
79
80            // Read OBU element length
81            let (length_field, is_last) = if obu_count == 0 || !is_last {
82                // W=0 or not last element: length field present
83                let payload_slice = payload.slice(offset..);
84                let (len, n) = read_leb128(&payload_slice);
85                if n == 0 {
86                    return Err(Error::ErrShortPacket);
87                }
88                offset += n;
89
90                // Check if this is actually the last element when W=0
91                let is_last_w0 = obu_count == 0 && offset + len as usize == payload.len();
92                (len as usize, is_last || is_last_w0)
93            } else {
94                // Last element when W != 0: no length field
95                (payload.len() - offset, true)
96            };
97
98            if offset + length_field > payload.len() {
99                return Err(Error::ErrShortPacket);
100            }
101
102            // Build OBU buffer
103            let obu_buffer = if is_first && obu_z {
104                // Continuation of previous packet's OBU
105                if self.buffer.is_empty() {
106                    // Lost first fragment, skip this OBU
107                    if is_last {
108                        break;
109                    }
110                    offset += length_field;
111                    obu_offset += 1;
112                    continue;
113                }
114
115                // Combine buffered data with current fragment
116                let mut combined = std::mem::take(&mut self.buffer);
117                combined.extend_from_slice(&payload[offset..offset + length_field]);
118                combined.freeze()
119            } else {
120                payload.slice(offset..offset + length_field)
121            };
122            offset += length_field;
123
124            // If this is the last OBU and Y flag is set, buffer it for next packet
125            if is_last && obu_y {
126                self.buffer = BytesMut::from(obu_buffer.as_ref());
127                break;
128            }
129
130            // Skip empty OBUs
131            if obu_buffer.is_empty() {
132                if is_last {
133                    break;
134                }
135                obu_offset += 1;
136                continue;
137            }
138
139            // Parse OBU header to check type
140            let obu_type = (obu_buffer[0] & OBU_TYPE_MASK) >> 3;
141
142            // Skip temporal delimiter and tile list OBUs
143            if obu_type == OBU_TYPE_TEMPORAL_DELIMITER || obu_type == OBU_TYPE_TILE_LIST {
144                if is_last {
145                    break;
146                }
147                obu_offset += 1;
148                continue;
149            }
150
151            // Check if OBU has size field
152            let has_size_field = (obu_buffer[0] & OBU_HAS_SIZE_BIT) != 0;
153            let has_extension = (obu_buffer[0] & 0x04) != 0;
154            let header_size = if has_extension { 2 } else { 1 };
155
156            if has_size_field {
157                // OBU already has size field, validate it
158                let payload_slice = obu_buffer.slice(header_size..);
159                let (obu_size, leb_size) = read_leb128(&payload_slice);
160                if leb_size == 0 {
161                    return Err(Error::ErrShortPacket);
162                }
163                let expected_size = header_size + leb_size + obu_size as usize;
164                // When this OBU was reassembled from multiple RTP packets,
165                // `length_field` is the length of the current packet's fragment,
166                // not the complete OBU. Use the reassembled buffer length for
167                // validation in that case.
168                let actual_size = if is_first && obu_z {
169                    obu_buffer.len()
170                } else {
171                    length_field
172                };
173                if actual_size != expected_size {
174                    return Err(Error::ErrShortPacket);
175                }
176                result.extend_from_slice(&obu_buffer);
177            } else {
178                // Add size field to OBU
179                // Set obu_has_size_field bit
180                result.put_u8(obu_buffer[0] | OBU_HAS_SIZE_BIT);
181
182                // Copy extension header if present
183                if has_extension && obu_buffer.len() > 1 {
184                    result.put_u8(obu_buffer[1]);
185                }
186
187                // Write payload size as LEB128
188                let payload_size = obu_buffer.len() - header_size;
189                write_leb128(&mut result, payload_size as u32);
190
191                // Copy OBU payload
192                if header_size < obu_buffer.len() {
193                    result.extend_from_slice(&obu_buffer[header_size..]);
194                }
195            }
196
197            if is_last {
198                break;
199            }
200            obu_offset += 1;
201        }
202
203        // Validate OBU count if W field was set
204        if obu_count != 0 && obu_offset != (obu_count - 1) as usize && !self.y {
205            return Err(Error::ErrShortPacket);
206        }
207
208        Ok(result.freeze())
209    }
210
211    /// Returns true if Z flag is not set (first OBU is not a continuation)
212    fn is_partition_head(&self, payload: &Bytes) -> bool {
213        if payload.is_empty() {
214            return false;
215        }
216        (payload[0] & AV1_Z_MASK) == 0
217    }
218
219    /// Returns true if marker bit is set (end of frame)
220    fn is_partition_tail(&self, marker: bool, _payload: &Bytes) -> bool {
221        marker
222    }
223}
224
225/// Write LEB128 encoded value to buffer
226fn write_leb128(buf: &mut BytesMut, mut value: u32) {
227    loop {
228        let mut byte = (value & 0x7f) as u8;
229        value >>= 7;
230        if value != 0 {
231            byte |= 0x80;
232        }
233        buf.put_u8(byte);
234        if value == 0 {
235            break;
236        }
237    }
238}
239
240#[cfg(test)]
241mod tests {
242    use super::*;
243
244    #[test]
245    fn test_depacketizer_basic() {
246        let mut depacketizer = Av1Depacketizer::new();
247
248        // Simple packet with one OBU element (W=1)
249        // Aggregation header: W=1, no Z, no Y, no N = 0x10
250        // OBU header: type=6 (Frame), no extension, no size = 0x30
251        // Total: aggregation header + OBU header + payload
252        let payload = Bytes::from(vec![
253            0x10, // Aggregation header: W=1
254            0x30, // OBU header: type=6 (Frame), no ext, no size
255            0x01, 0x02, 0x03, // OBU payload
256        ]);
257
258        let result = depacketizer.depacketize(&payload).unwrap();
259        assert!(!result.is_empty());
260        // Should have size field added (OBU_HAS_SIZE_BIT = 0x02)
261        assert_eq!(result[0] & OBU_HAS_SIZE_BIT, OBU_HAS_SIZE_BIT);
262        // Size should be 3 (payload bytes)
263        assert_eq!(result[1], 3);
264        // Payload should follow
265        assert_eq!(&result[2..], &[0x01, 0x02, 0x03]);
266    }
267
268    #[test]
269    fn test_depacketizer_with_w_zero() {
270        let mut depacketizer = Av1Depacketizer::new();
271
272        // Packet with W=0 means each OBU has length prefix
273        // Aggregation header: W=0
274        // Length field (LEB128): 4 bytes
275        // OBU: header + payload
276        let payload = Bytes::from(vec![
277            0x00, // Aggregation header: W=0
278            0x04, // Length field: 4 bytes
279            0x30, // OBU header: type=6 (Frame), no ext, no size
280            0x01, 0x02, 0x03, // OBU payload (3 bytes, total OBU = 4)
281        ]);
282
283        let result = depacketizer.depacketize(&payload).unwrap();
284        assert!(!result.is_empty());
285        // Should have size field added
286        assert_eq!(result[0] & OBU_HAS_SIZE_BIT, OBU_HAS_SIZE_BIT);
287    }
288
289    #[test]
290    fn test_is_partition_head() {
291        let depacketizer = Av1Depacketizer::new();
292
293        // Z=0 means partition head
294        let payload = Bytes::from(vec![0x10, 0x30]);
295        assert!(depacketizer.is_partition_head(&payload));
296
297        // Z=1 means continuation
298        let payload = Bytes::from(vec![0x90, 0x30]);
299        assert!(!depacketizer.is_partition_head(&payload));
300    }
301
302    #[test]
303    fn test_write_leb128() {
304        let mut buf = BytesMut::new();
305
306        // Test small values
307        write_leb128(&mut buf, 0);
308        assert_eq!(buf.as_ref(), &[0x00]);
309
310        buf.clear();
311        write_leb128(&mut buf, 127);
312        assert_eq!(buf.as_ref(), &[0x7f]);
313
314        buf.clear();
315        write_leb128(&mut buf, 128);
316        assert_eq!(buf.as_ref(), &[0x80, 0x01]);
317
318        buf.clear();
319        write_leb128(&mut buf, 16383);
320        assert_eq!(buf.as_ref(), &[0xff, 0x7f]);
321    }
322
323    #[test]
324    fn test_skip_temporal_delimiter() {
325        let mut depacketizer = Av1Depacketizer::new();
326
327        // Packet with temporal delimiter OBU (type=2) which should be skipped
328        let payload = Bytes::from(vec![
329            0x10, // Aggregation header: W=1
330            0x12, // OBU header: type=2 (Temporal Delimiter), no ext, with size
331            0x00, // Size = 0
332        ]);
333
334        let result = depacketizer.depacketize(&payload).unwrap();
335        // Should be empty since temporal delimiter is skipped
336        assert!(result.is_empty());
337    }
338
339    #[test]
340    fn test_fragmented_obu_with_size_field() {
341        let mut depacketizer = Av1Depacketizer::new();
342
343        // Build a Frame OBU that already carries a low-overhead size field.
344        let mut obu = BytesMut::new();
345        obu.put_u8(0x32); // Frame OBU, no extension, has size field
346        let payload = vec![0xAB; 500];
347        write_leb128(&mut obu, payload.len() as u32);
348        obu.extend_from_slice(&payload);
349        let obu = obu.freeze();
350
351        let first_fragment_size = 400;
352        let p1 = Bytes::from_iter(
353            std::iter::once(0x50) // Z=0, Y=1, W=1
354                .chain(obu[..first_fragment_size].iter().copied()),
355        );
356        let p2 = Bytes::from_iter(
357            std::iter::once(0x90) // Z=1, Y=0, W=1
358                .chain(obu[first_fragment_size..].iter().copied()),
359        );
360
361        assert!(depacketizer.depacketize(&p1).unwrap().is_empty());
362        let result = depacketizer.depacketize(&p2).unwrap();
363        assert_eq!(result, obu);
364    }
365}