Skip to main content

rtc_interceptor/flexfec/draft03/
decoder.rs

1//! FlexFEC draft-03 recovery: rebuilding a lost media packet from a repair packet.
2
3use super::encoder::BASE_RTP_HEADER_SIZE;
4use shared::marshal::{Marshal, MarshalSize, Unmarshal};
5
6/// Recovered media packets kept for matching against later repair packets.
7const MAX_MEDIA_PACKETS: usize = 100;
8
9/// Repair packets held while waiting for the media they protect.
10const MAX_FEC_PACKETS: usize = 100;
11
12/// Recovered packets retained after pruning.
13const RETAINED_RECOVERED_PACKETS: usize = 192;
14
15/// Sequence distance beyond which a held repair packet is considered stale.
16const STALE_SEQUENCE_DISTANCE: u16 = 0x3FFF;
17
18/// Why a repair packet could not be used.
19#[derive(Debug, Clone, Copy, PartialEq, Eq)]
20pub enum ParseError {
21    /// Shorter than the fields it claims to carry.
22    Truncated,
23    /// The retransmission bit is set; this scheme does not carry retransmissions.
24    RetransmissionBitSet,
25    /// The F bit selects the inflexible generator matrix, which draft-03 does not define here.
26    InflexibleGeneratorMatrix,
27    /// Draft-03 protects exactly one stream per repair packet.
28    MultipleSsrcProtection,
29    /// The last packet mask did not set its k-bit, so the header never terminates.
30    UnterminatedPacketMask,
31}
32
33/// A repair packet's header, parsed.
34#[derive(Debug, Clone)]
35struct RepairHeader {
36    protected_ssrc: u32,
37    sequence_number_base: u16,
38    /// Sequence numbers this repair packet protects, derived from the packet masks.
39    protected_sequence_numbers: Vec<u16>,
40    /// Offset of the repair payload within the packet payload.
41    payload_offset: usize,
42}
43
44/// A media packet a repair packet protects, and the copy of it we hold, if any.
45#[derive(Debug, Clone)]
46struct ProtectedPacket {
47    sequence_number: u16,
48    packet: Option<rtp::Packet>,
49}
50
51/// A repair packet and what it is waiting for.
52#[derive(Debug, Clone)]
53struct RepairState {
54    packet: rtp::Packet,
55    header: RepairHeader,
56    protected: Vec<ProtectedPacket>,
57}
58
59impl RepairState {
60    fn missing(&self) -> usize {
61        self.protected
62            .iter()
63            .filter(|protected| protected.packet.is_none())
64            .count()
65    }
66}
67
68/// Recovers media packets lost from a stream protected by FlexFEC draft-03.
69///
70/// Feed it every packet of both streams — media and repair — and it returns whatever it was able
71/// to rebuild. A repair packet recovers exactly one loss among the packets it covers, which is why
72/// the encoder interleaves: consecutive losses land under different repair packets.
73///
74/// # Difference from upstream
75///
76/// `pion/interceptor` stores each protected packet as a pointer into the recovered-packet slice,
77/// then `append`s to and sorts that same slice. Appending can reallocate and sorting reorders, so
78/// those pointers can dangle or come to refer to a different packet. Copies are held here instead:
79/// the borrow checker would not permit the upstream shape, and the bug it prevents is real.
80#[derive(Debug)]
81pub struct FlexFec03Decoder {
82    repair_ssrc: u32,
83    media_ssrc: u32,
84    /// Media packets seen or recovered, ordered oldest first.
85    recovered: Vec<rtp::Packet>,
86    repair_packets: Vec<RepairState>,
87}
88
89impl FlexFec03Decoder {
90    /// A decoder for `media_ssrc`, protected by the repair stream on `repair_ssrc`.
91    pub fn new(repair_ssrc: u32, media_ssrc: u32) -> Self {
92        Self {
93            repair_ssrc,
94            media_ssrc,
95            recovered: Vec::new(),
96            repair_packets: Vec::new(),
97        }
98    }
99
100    /// Offer a packet from either stream, returning any media packets it made recoverable.
101    ///
102    /// Recovery can cascade: rebuilding one packet may complete another repair packet that was
103    /// waiting on two losses, so this repeats until nothing further can be recovered.
104    pub fn decode(&mut self, packet: rtp::Packet) -> Vec<rtp::Packet> {
105        self.reset_on_large_discontinuity(&packet);
106        self.insert(packet);
107        self.attempt_recovery()
108    }
109
110    /// Media packets currently held, whether received or recovered.
111    pub fn recovered_len(&self) -> usize {
112        self.recovered.len()
113    }
114
115    /// Repair packets currently held while waiting for the media they protect.
116    pub fn pending_repair_packets(&self) -> usize {
117        self.repair_packets.len()
118    }
119
120    /// A jump far beyond the window means the stream moved on; holding the old state would match
121    /// new packets against repair packets that can never complete.
122    fn reset_on_large_discontinuity(&mut self, packet: &rtp::Packet) {
123        if self.recovered.len() < MAX_MEDIA_PACKETS {
124            return;
125        }
126        let Some(newest) = self.recovered.last() else {
127            return;
128        };
129        if newest.header.ssrc != packet.header.ssrc {
130            return;
131        }
132        if sequence_distance(packet.header.sequence_number, newest.header.sequence_number)
133            > MAX_MEDIA_PACKETS as u16
134        {
135            self.recovered.clear();
136            self.repair_packets.clear();
137        }
138    }
139
140    fn insert(&mut self, packet: rtp::Packet) {
141        if packet.header.ssrc == self.repair_ssrc {
142            self.prune_stale_repair_packets(packet.header.sequence_number);
143            self.insert_repair_packet(packet);
144        } else if packet.header.ssrc == self.media_ssrc {
145            self.insert_media_packet(packet);
146        }
147        self.discard_old_recovered_packets();
148    }
149
150    /// Drop repair packets whose sequence numbers are far from what is arriving now: they protect
151    /// media that will never be offered again.
152    fn prune_stale_repair_packets(&mut self, sequence_number: u16) {
153        let repair_ssrc_distance = |state: &RepairState| {
154            sequence_distance(sequence_number, state.packet.header.sequence_number)
155        };
156        self.repair_packets
157            .retain(|state| repair_ssrc_distance(state) <= STALE_SEQUENCE_DISTANCE);
158    }
159
160    fn insert_media_packet(&mut self, packet: rtp::Packet) {
161        if self
162            .recovered
163            .iter()
164            .any(|held| held.header.sequence_number == packet.header.sequence_number)
165        {
166            return;
167        }
168        self.record_recovered(packet);
169    }
170
171    fn insert_repair_packet(&mut self, packet: rtp::Packet) {
172        if self
173            .repair_packets
174            .iter()
175            .any(|state| state.packet.header.sequence_number == packet.header.sequence_number)
176        {
177            return;
178        }
179
180        let Ok(header) = parse_repair_header(&packet.payload) else {
181            return;
182        };
183        if header.protected_ssrc != self.media_ssrc {
184            // Protecting a stream this decoder knows nothing about.
185            return;
186        }
187        if header.protected_sequence_numbers.is_empty() {
188            return;
189        }
190
191        let protected = header
192            .protected_sequence_numbers
193            .iter()
194            .map(|&sequence_number| ProtectedPacket {
195                sequence_number,
196                packet: self
197                    .recovered
198                    .iter()
199                    .find(|held| held.header.sequence_number == sequence_number)
200                    .cloned(),
201            })
202            .collect();
203
204        self.repair_packets.push(RepairState {
205            packet,
206            header,
207            protected,
208        });
209        self.repair_packets.sort_by(|a, b| {
210            sequence_order(
211                a.packet.header.sequence_number,
212                b.packet.header.sequence_number,
213            )
214        });
215        if self.repair_packets.len() > MAX_FEC_PACKETS {
216            self.repair_packets.remove(0);
217        }
218    }
219
220    /// Record a media packet and tell every repair packet waiting for it.
221    fn record_recovered(&mut self, packet: rtp::Packet) {
222        for state in &mut self.repair_packets {
223            for protected in &mut state.protected {
224                if protected.sequence_number == packet.header.sequence_number {
225                    protected.packet = Some(packet.clone());
226                }
227            }
228        }
229        self.recovered.push(packet);
230        self.recovered
231            .sort_by(|a, b| sequence_order(a.header.sequence_number, b.header.sequence_number));
232    }
233
234    fn attempt_recovery(&mut self) -> Vec<rtp::Packet> {
235        let mut recovered_now = Vec::new();
236
237        // Each pass takes its repair packet *out* of the list before using it, which bounds the
238        // loop by the list length whatever happens next. A repair packet is spent once it has
239        // produced its recovery, and one that cannot be used must not be retried.
240        //
241        // Removing it is not merely tidy. Leaving it in and relying on `record_recovered` to
242        // clear its missing count does not terminate: if the recovered sequence number did not
243        // match the slot that was waiting on it, the same repair packet would be selected again
244        // forever.
245        while let Some(index) = self
246            .repair_packets
247            .iter()
248            .position(|state| state.missing() == 1)
249        {
250            let state = self.repair_packets.remove(index);
251            let Some(packet) = self.recover(&state) else {
252                continue;
253            };
254
255            recovered_now.push(packet.clone());
256            self.record_recovered(packet);
257            self.discard_old_recovered_packets();
258        }
259
260        recovered_now
261    }
262
263    /// Rebuild the one missing packet of `state` by XORing the others back out of the repair data.
264    fn recover(&self, state: &RepairState) -> Option<rtp::Packet> {
265        let repair_payload = state.packet.payload.get(state.header.payload_offset..)?;
266
267        // The recovery fields occupy the first 8 bytes; the RTP header is 12.
268        let mut header = vec![0u8; BASE_RTP_HEADER_SIZE];
269        header[..8].copy_from_slice(state.packet.payload.get(..8)?);
270
271        let mut missing_sequence_number = 0u16;
272        for protected in &state.protected {
273            let Some(packet) = &protected.packet else {
274                missing_sequence_number = protected.sequence_number;
275                continue;
276            };
277
278            let mut marshalled = vec![0u8; packet.header.marshal_size()];
279            packet.header.marshal_to(&mut marshalled).ok()?;
280            // Bytes 2..4 of a media header are its sequence number; in the recovery fields that
281            // position carries the payload length instead, so substitute before XORing.
282            let payload_length = (packet.marshal_size() - BASE_RTP_HEADER_SIZE) as u16;
283            marshalled[2..4].copy_from_slice(&payload_length.to_be_bytes());
284
285            for index in 0..8 {
286                header[index] ^= marshalled[index];
287            }
288        }
289
290        // Version 2, and no padding: neither is recovered, both are known.
291        header[0] |= 0x80;
292        header[0] &= 0xBF;
293
294        let payload_length = u16::from_be_bytes([header[2], header[3]]) as usize;
295        if repair_payload.len() < payload_length {
296            return None;
297        }
298        header[2..4].copy_from_slice(&missing_sequence_number.to_be_bytes());
299        header[8..12].copy_from_slice(&self.media_ssrc.to_be_bytes());
300
301        let mut payload = repair_payload[..payload_length].to_vec();
302        for protected in &state.protected {
303            let Some(packet) = &protected.packet else {
304                continue;
305            };
306            let mut marshalled = vec![0u8; packet.marshal_size()];
307            packet.marshal_to(&mut marshalled).ok()?;
308            for (target, &source) in payload.iter_mut().zip(&marshalled[BASE_RTP_HEADER_SIZE..]) {
309                *target ^= source;
310            }
311        }
312
313        header.extend_from_slice(&payload);
314        let mut buffer = header.as_slice();
315        rtp::Packet::unmarshal(&mut buffer).ok()
316    }
317
318    fn discard_old_recovered_packets(&mut self) {
319        if self.recovered.len() > RETAINED_RECOVERED_PACKETS {
320            let excess = self.recovered.len() - RETAINED_RECOVERED_PACKETS;
321            self.recovered.drain(..excess);
322        }
323    }
324}
325
326/// Read the packet masks and the fields needed to recover from a repair payload.
327fn parse_repair_header(data: &[u8]) -> Result<RepairHeader, ParseError> {
328    if data.len() < 20 {
329        return Err(ParseError::Truncated);
330    }
331    if data[0] & 0x80 != 0 {
332        return Err(ParseError::RetransmissionBitSet);
333    }
334    if data[0] & 0x40 != 0 {
335        return Err(ParseError::InflexibleGeneratorMatrix);
336    }
337    if data[8] != 1 {
338        return Err(ParseError::MultipleSsrcProtection);
339    }
340
341    let protected_ssrc = u32::from_be_bytes([data[12], data[13], data[14], data[15]]);
342    let sequence_number_base = u16::from_be_bytes([data[16], data[17]]);
343
344    let mut protected_sequence_numbers = Vec::new();
345    let mask0 = u16::from_be_bytes([data[18], data[19]]) & 0x7FFF;
346    append_mask(
347        &mut protected_sequence_numbers,
348        u64::from(mask0),
349        15,
350        sequence_number_base,
351    );
352
353    if data[18] & 0x80 != 0 {
354        return Ok(RepairHeader {
355            protected_ssrc,
356            sequence_number_base,
357            protected_sequence_numbers,
358            payload_offset: 20,
359        });
360    }
361
362    if data.len() < 24 {
363        return Err(ParseError::Truncated);
364    }
365    let mask1 = u32::from_be_bytes([data[20], data[21], data[22], data[23]]) & 0x7FFF_FFFF;
366    append_mask(
367        &mut protected_sequence_numbers,
368        u64::from(mask1),
369        31,
370        sequence_number_base.wrapping_add(15),
371    );
372
373    if data[20] & 0x80 != 0 {
374        return Ok(RepairHeader {
375            protected_ssrc,
376            sequence_number_base,
377            protected_sequence_numbers,
378            payload_offset: 24,
379        });
380    }
381
382    if data.len() < 32 {
383        return Err(ParseError::Truncated);
384    }
385    let mut mask2_bytes = [0u8; 8];
386    mask2_bytes.copy_from_slice(&data[24..32]);
387    let mask2 = u64::from_be_bytes(mask2_bytes) & 0x7FFF_FFFF_FFFF_FFFF;
388    append_mask(
389        &mut protected_sequence_numbers,
390        mask2,
391        63,
392        sequence_number_base.wrapping_add(46),
393    );
394
395    if data[24] & 0x80 == 0 {
396        // Nothing marks the end of the mask list, so the payload offset is unknown.
397        return Err(ParseError::UnterminatedPacketMask);
398    }
399
400    Ok(RepairHeader {
401        protected_ssrc,
402        sequence_number_base,
403        protected_sequence_numbers,
404        payload_offset: 32,
405    })
406}
407
408/// Expand one packet mask into the sequence numbers it names, most significant bit first.
409fn append_mask(out: &mut Vec<u16>, mask: u64, bit_count: u16, base: u16) {
410    for bit in 0..bit_count {
411        if (mask >> (bit_count - 1 - bit)) & 1 == 1 {
412            out.push(base.wrapping_add(bit));
413        }
414    }
415}
416
417/// Whether `value` is later than `previous` in the wrapping 16-bit sequence space.
418fn is_newer(previous: u16, value: u16) -> bool {
419    const HALF: u16 = 0x8000;
420    let forward = value.wrapping_sub(previous);
421    if forward == HALF {
422        return value > previous;
423    }
424    value != previous && forward < HALF
425}
426
427/// Order two sequence numbers oldest first, respecting the wrap.
428fn sequence_order(a: u16, b: u16) -> std::cmp::Ordering {
429    if a == b {
430        std::cmp::Ordering::Equal
431    } else if is_newer(a, b) {
432        std::cmp::Ordering::Less
433    } else {
434        std::cmp::Ordering::Greater
435    }
436}
437
438/// The shorter distance between two sequence numbers in either direction.
439fn sequence_distance(a: u16, b: u16) -> u16 {
440    a.wrapping_sub(b).min(b.wrapping_sub(a))
441}