rtc_interceptor/flexfec/draft03/
decoder.rs1use super::encoder::BASE_RTP_HEADER_SIZE;
4use shared::marshal::{Marshal, MarshalSize, Unmarshal};
5
6const MAX_MEDIA_PACKETS: usize = 100;
8
9const MAX_FEC_PACKETS: usize = 100;
11
12const RETAINED_RECOVERED_PACKETS: usize = 192;
14
15const STALE_SEQUENCE_DISTANCE: u16 = 0x3FFF;
17
18#[derive(Debug, Clone, Copy, PartialEq, Eq)]
20pub enum ParseError {
21 Truncated,
23 RetransmissionBitSet,
25 InflexibleGeneratorMatrix,
27 MultipleSsrcProtection,
29 UnterminatedPacketMask,
31}
32
33#[derive(Debug, Clone)]
35struct RepairHeader {
36 protected_ssrc: u32,
37 sequence_number_base: u16,
38 protected_sequence_numbers: Vec<u16>,
40 payload_offset: usize,
42}
43
44#[derive(Debug, Clone)]
46struct ProtectedPacket {
47 sequence_number: u16,
48 packet: Option<rtp::Packet>,
49}
50
51#[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#[derive(Debug)]
81pub struct FlexFec03Decoder {
82 repair_ssrc: u32,
83 media_ssrc: u32,
84 recovered: Vec<rtp::Packet>,
86 repair_packets: Vec<RepairState>,
87}
88
89impl FlexFec03Decoder {
90 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 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 pub fn recovered_len(&self) -> usize {
112 self.recovered.len()
113 }
114
115 pub fn pending_repair_packets(&self) -> usize {
117 self.repair_packets.len()
118 }
119
120 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 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 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 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 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 fn recover(&self, state: &RepairState) -> Option<rtp::Packet> {
265 let repair_payload = state.packet.payload.get(state.header.payload_offset..)?;
266
267 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 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 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
326fn 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 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
408fn 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
417fn 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
427fn 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
438fn sequence_distance(a: u16, b: u16) -> u16 {
440 a.wrapping_sub(b).min(b.wrapping_sub(a))
441}