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