Skip to main content

rtc_interceptor/nack/
responder.rs

1//! NACK Responder Interceptor - Responds to NACK requests by retransmitting packets.
2
3use super::send_buffer::SendBuffer;
4use super::stream_supports_nack;
5use crate::stream_info::StreamInfo;
6use crate::{Interceptor, Packet, TaggedPacket, interceptor};
7use shared::TransportContext;
8use shared::error::Error;
9use std::collections::{HashMap, VecDeque};
10use std::marker::PhantomData;
11use std::time::Instant;
12
13/// Builder for the NackResponderInterceptor.
14///
15/// # Example
16///
17/// ```
18/// use rtc_interceptor::{Registry, NackResponderBuilder};
19///
20/// let chain = Registry::new()
21///     .with(NackResponderBuilder::new()
22///         .with_size(1024)
23///         .build())
24///     .build();
25/// ```
26pub struct NackResponderBuilder<P> {
27    /// Size of the send buffer (must be power of 2: 1, 2, 4, ..., 32768).
28    size: u16,
29    _phantom: PhantomData<P>,
30}
31
32impl<P> Default for NackResponderBuilder<P> {
33    fn default() -> Self {
34        Self {
35            size: 1024,
36            _phantom: PhantomData,
37        }
38    }
39}
40
41impl<P> NackResponderBuilder<P> {
42    /// Create a new builder with default settings.
43    pub fn new() -> Self {
44        Self::default()
45    }
46
47    /// Set the size of the send buffer.
48    ///
49    /// Size must be a power of 2 between 1 and 32768 (inclusive).
50    /// Larger buffers can retransmit older packets but use more memory.
51    pub fn with_size(mut self, size: u16) -> Self {
52        self.size = size;
53        self
54    }
55
56    /// Build the interceptor factory function.
57    pub fn build(self) -> impl FnOnce(P) -> NackResponderInterceptor<P> {
58        move |inner| NackResponderInterceptor::new(inner, self.size)
59    }
60}
61
62/// Per-stream state for the responder.
63struct LocalStream {
64    /// Buffer of sent packets for retransmission.
65    send_buffer: SendBuffer,
66    /// RTX SSRC for RFC4588 retransmission (if configured).
67    ssrc_rtx: Option<u32>,
68    /// RTX payload type for RFC4588 retransmission (if configured).
69    payload_type_rtx: Option<u8>,
70    /// Sequence number counter for RTX packets.
71    rtx_sequence_number: u16,
72}
73
74/// Interceptor that responds to NACK requests by retransmitting packets.
75///
76/// This interceptor buffers outgoing RTP packets on local streams and
77/// retransmits them when RTCP TransportLayerNack packets are received.
78#[derive(Interceptor)]
79pub struct NackResponderInterceptor<P> {
80    #[next]
81    inner: P,
82
83    /// Configuration
84    size: u16,
85
86    /// Send buffers per local stream SSRC
87    streams: HashMap<u32, LocalStream>,
88
89    /// Queue for retransmitted packets
90    write_queue: VecDeque<TaggedPacket>,
91}
92
93impl<P> NackResponderInterceptor<P> {
94    fn new(inner: P, size: u16) -> Self {
95        Self {
96            inner,
97            size,
98            streams: HashMap::new(),
99            write_queue: VecDeque::new(),
100        }
101    }
102
103    /// Handle a NACK request by queuing retransmissions.
104    fn handle_nack(
105        &mut self,
106        now: Instant,
107        nack: &rtcp::transport_feedbacks::transport_layer_nack::TransportLayerNack,
108    ) {
109        // Collect sequence numbers to retransmit
110        let mut seqs_to_retransmit = Vec::new();
111
112        for nack_pair in &nack.nacks {
113            // Check the base packet ID
114            seqs_to_retransmit.push(nack_pair.packet_id);
115
116            // Check each bit in lost_packets bitmap
117            for i in 0..16 {
118                if nack_pair.lost_packets & (1 << i) != 0 {
119                    let seq = nack_pair.packet_id.wrapping_add(i + 1);
120                    seqs_to_retransmit.push(seq);
121                }
122            }
123        }
124
125        let Some(stream) = self.streams.get_mut(&nack.media_ssrc) else {
126            return;
127        };
128
129        // Queue retransmissions
130        for seq in seqs_to_retransmit {
131            let Some(original_packet) = stream.send_buffer.get(seq) else {
132                continue;
133            };
134
135            let packet = if let (Some(ssrc_rtx), Some(pt_rtx)) =
136                (stream.ssrc_rtx, stream.payload_type_rtx)
137            {
138                // RFC4588: Create RTX packet
139                // - Use RTX SSRC and payload type
140                // - Prepend original sequence number (2 bytes big-endian) to payload
141                // - Use separate RTX sequence number counter
142                let original_seq = original_packet.header.sequence_number;
143                let mut rtx_payload = Vec::with_capacity(2 + original_packet.payload.len());
144                rtx_payload.extend_from_slice(&original_seq.to_be_bytes());
145                rtx_payload.extend_from_slice(&original_packet.payload);
146
147                let rtx_seq = stream.rtx_sequence_number;
148                stream.rtx_sequence_number = stream.rtx_sequence_number.wrapping_add(1);
149
150                rtp::Packet {
151                    header: rtp::header::Header {
152                        // Not left to `..Default::default()`: the default
153                        // header is version 0, and receivers discard
154                        // version != 2 before examining anything else, so a
155                        // defaulted RTX packet is dropped on arrival and the
156                        // NACKed gap never repairs.
157                        version: 2,
158                        ssrc: ssrc_rtx,
159                        payload_type: pt_rtx,
160                        sequence_number: rtx_seq,
161                        timestamp: original_packet.header.timestamp,
162                        marker: original_packet.header.marker,
163                        ..Default::default()
164                    },
165                    payload: rtx_payload.into(),
166                }
167            } else {
168                // No RTX: retransmit original packet as-is
169                original_packet.clone()
170            };
171
172            self.write_queue.push_back(TaggedPacket {
173                now,
174                transport: TransportContext::default(),
175                message: Packet::Rtp(packet),
176            });
177        }
178    }
179}
180
181#[interceptor]
182impl<P: Interceptor> NackResponderInterceptor<P> {
183    #[overrides]
184    fn handle_read(&mut self, msg: TaggedPacket) -> Result<(), Self::Error> {
185        // Process NACK packets
186        if let Packet::Rtcp(ref rtcp_packets) = msg.message {
187            for rtcp_packet in rtcp_packets {
188                if let Some(nack) = rtcp_packet
189                    .as_any()
190                    .downcast_ref::<rtcp::transport_feedbacks::transport_layer_nack::TransportLayerNack>()
191                {
192                    self.handle_nack(msg.now, nack);
193                }
194            }
195        }
196
197        self.inner.handle_read(msg)
198    }
199
200    #[overrides]
201    fn handle_write(&mut self, msg: TaggedPacket) -> Result<(), Self::Error> {
202        // Buffer outgoing RTP packets
203        if let Packet::Rtp(ref rtp_packet) = msg.message
204            && let Some(stream) = self.streams.get_mut(&rtp_packet.header.ssrc)
205        {
206            stream.send_buffer.add(rtp_packet.clone());
207        }
208
209        self.inner.handle_write(msg)
210    }
211
212    #[overrides]
213    fn poll_write(&mut self) -> Option<Self::Wout> {
214        // First drain retransmitted packets
215        if let Some(pkt) = self.write_queue.pop_front() {
216            return Some(pkt);
217        }
218        self.inner.poll_write()
219    }
220
221    #[overrides]
222    fn bind_local_stream(&mut self, info: &StreamInfo) {
223        if stream_supports_nack(info)
224            && let Some(send_buffer) = SendBuffer::new(self.size)
225        {
226            self.streams.insert(
227                info.ssrc,
228                LocalStream {
229                    send_buffer,
230                    ssrc_rtx: info.ssrc_rtx,
231                    payload_type_rtx: info.payload_type_rtx,
232                    rtx_sequence_number: 0,
233                },
234            );
235        }
236        self.inner.bind_local_stream(info);
237    }
238
239    #[overrides]
240    fn unbind_local_stream(&mut self, info: &StreamInfo) {
241        self.streams.remove(&info.ssrc);
242        self.inner.unbind_local_stream(info);
243    }
244}
245
246#[cfg(test)]
247mod tests {
248    use super::*;
249    use crate::Registry;
250    use crate::stream_info::RTCPFeedback;
251    use sansio::Protocol;
252
253    fn make_rtp_packet(ssrc: u32, seq: u16, payload: &[u8]) -> TaggedPacket {
254        TaggedPacket {
255            now: Instant::now(),
256            transport: Default::default(),
257            message: Packet::Rtp(rtp::Packet {
258                header: rtp::header::Header {
259                    ssrc,
260                    sequence_number: seq,
261                    ..Default::default()
262                },
263                payload: payload.to_vec().into(),
264            }),
265        }
266    }
267
268    fn make_nack_packet(sender_ssrc: u32, media_ssrc: u32, nacks: Vec<(u16, u16)>) -> TaggedPacket {
269        let nack_pairs: Vec<rtcp::transport_feedbacks::transport_layer_nack::NackPair> = nacks
270            .into_iter()
271            .map(|(packet_id, lost_packets)| {
272                rtcp::transport_feedbacks::transport_layer_nack::NackPair {
273                    packet_id,
274                    lost_packets,
275                }
276            })
277            .collect();
278
279        TaggedPacket {
280            now: Instant::now(),
281            transport: Default::default(),
282            message: Packet::Rtcp(vec![Box::new(
283                rtcp::transport_feedbacks::transport_layer_nack::TransportLayerNack {
284                    sender_ssrc,
285                    media_ssrc,
286                    nacks: nack_pairs,
287                },
288            )]),
289        }
290    }
291
292    #[test]
293    fn test_nack_responder_builder_defaults() {
294        let chain = Registry::new()
295            .with(NackResponderBuilder::default().build())
296            .build();
297
298        assert_eq!(chain.size, 1024);
299    }
300
301    #[test]
302    fn test_nack_responder_builder_custom() {
303        let chain = Registry::new()
304            .with(NackResponderBuilder::new().with_size(2048).build())
305            .build();
306
307        assert_eq!(chain.size, 2048);
308    }
309
310    #[test]
311    fn test_nack_responder_retransmits_packet() {
312        let mut chain = Registry::new()
313            .with(NackResponderBuilder::new().with_size(8).build())
314            .build();
315
316        // Bind local stream with NACK support
317        let info = StreamInfo {
318            ssrc: 12345,
319            clock_rate: 90000,
320            rtcp_feedback: vec![RTCPFeedback {
321                typ: "nack".to_string(),
322                parameter: "".to_string(),
323            }],
324            ..Default::default()
325        };
326        chain.bind_local_stream(&info);
327
328        let now = Instant::now();
329
330        // Send packets 10, 11, 12, 14, 15 (missing 13)
331        for seq in [10u16, 11, 12, 14, 15] {
332            let mut pkt = make_rtp_packet(12345, seq, &[seq as u8]);
333            pkt.now = now;
334            chain.handle_write(pkt).unwrap();
335            chain.poll_write(); // Drain normal write
336        }
337
338        // Receive NACK for 11, 12, 13, 15
339        // nack_pair: packet_id=11, lost_packets=0b1011 means 11, 12, 13, 15
340        let mut nack = make_nack_packet(999, 12345, vec![(11, 0b1011)]);
341        nack.now = now;
342        chain.handle_read(nack).unwrap();
343
344        // Should retransmit 11, 12, 15 (13 was never sent)
345        let mut retransmitted = Vec::new();
346        while let Some(pkt) = chain.poll_write() {
347            if let Packet::Rtp(rtp) = pkt.message {
348                retransmitted.push(rtp.header.sequence_number);
349            }
350        }
351
352        assert!(retransmitted.contains(&11));
353        assert!(retransmitted.contains(&12));
354        assert!(!retransmitted.contains(&13)); // Never sent
355        assert!(retransmitted.contains(&15));
356    }
357
358    #[test]
359    fn test_nack_responder_no_retransmit_without_binding() {
360        let mut chain = Registry::new()
361            .with(NackResponderBuilder::new().with_size(8).build())
362            .build();
363
364        let now = Instant::now();
365
366        // Send packets without binding stream (no buffer)
367        for seq in [10u16, 11, 12] {
368            let mut pkt = make_rtp_packet(12345, seq, &[seq as u8]);
369            pkt.now = now;
370            chain.handle_write(pkt).unwrap();
371            chain.poll_write();
372        }
373
374        // Receive NACK
375        let mut nack = make_nack_packet(999, 12345, vec![(11, 0)]);
376        nack.now = now;
377        chain.handle_read(nack).unwrap();
378
379        // No retransmissions (stream not bound)
380        assert!(chain.poll_write().is_none());
381    }
382
383    #[test]
384    fn test_nack_responder_no_retransmit_expired_packet() {
385        let mut chain = Registry::new()
386            .with(NackResponderBuilder::new().with_size(8).build())
387            .build();
388
389        let info = StreamInfo {
390            ssrc: 12345,
391            clock_rate: 90000,
392            rtcp_feedback: vec![RTCPFeedback {
393                typ: "nack".to_string(),
394                parameter: "".to_string(),
395            }],
396            ..Default::default()
397        };
398        chain.bind_local_stream(&info);
399
400        let now = Instant::now();
401
402        // Send packets 0-15 (buffer size is 8, so 0-7 will be pushed out)
403        for seq in 0..16u16 {
404            let mut pkt = make_rtp_packet(12345, seq, &[seq as u8]);
405            pkt.now = now;
406            chain.handle_write(pkt).unwrap();
407            chain.poll_write();
408        }
409
410        // Request retransmit of seq 0 (should be expired from buffer)
411        let mut nack = make_nack_packet(999, 12345, vec![(0, 0)]);
412        nack.now = now;
413        chain.handle_read(nack).unwrap();
414
415        // No retransmission (packet too old)
416        assert!(chain.poll_write().is_none());
417
418        // But seq 10 should still be available
419        let mut nack = make_nack_packet(999, 12345, vec![(10, 0)]);
420        nack.now = now;
421        chain.handle_read(nack).unwrap();
422
423        let pkt = chain.poll_write();
424        assert!(pkt.is_some());
425        if let Some(tagged) = pkt
426            && let Packet::Rtp(rtp) = tagged.message
427        {
428            assert_eq!(rtp.header.sequence_number, 10);
429        }
430    }
431
432    #[test]
433    fn test_nack_responder_unbind_removes_stream() {
434        let mut chain = Registry::new()
435            .with(NackResponderBuilder::new().with_size(8).build())
436            .build();
437
438        let info = StreamInfo {
439            ssrc: 12345,
440            clock_rate: 90000,
441            rtcp_feedback: vec![RTCPFeedback {
442                typ: "nack".to_string(),
443                parameter: "".to_string(),
444            }],
445            ..Default::default()
446        };
447
448        chain.bind_local_stream(&info);
449        assert!(chain.streams.contains_key(&12345));
450
451        chain.unbind_local_stream(&info);
452        assert!(!chain.streams.contains_key(&12345));
453    }
454
455    #[test]
456    fn test_nack_responder_no_nack_support() {
457        let mut chain = Registry::new()
458            .with(NackResponderBuilder::new().with_size(8).build())
459            .build();
460
461        // Bind stream without NACK support
462        let info = StreamInfo {
463            ssrc: 12345,
464            clock_rate: 90000,
465            rtcp_feedback: vec![], // No NACK support
466            ..Default::default()
467        };
468        chain.bind_local_stream(&info);
469
470        // Should not create send buffer
471        assert!(!chain.streams.contains_key(&12345));
472    }
473
474    #[test]
475    fn test_nack_responder_passthrough() {
476        let mut chain = Registry::new()
477            .with(NackResponderBuilder::new().with_size(8).build())
478            .build();
479
480        let now = Instant::now();
481
482        // RTP packets should pass through
483        let mut pkt = make_rtp_packet(12345, 1, &[1]);
484        pkt.now = now;
485        chain.handle_write(pkt).unwrap();
486        let out = chain.poll_write();
487        assert!(out.is_some());
488
489        // RTCP packets should pass through to read
490        let mut nack = make_nack_packet(999, 12345, vec![(1, 0)]);
491        nack.now = now;
492        chain.handle_read(nack).unwrap();
493        let out = chain.poll_read();
494        assert!(out.is_none());
495    }
496
497    #[test]
498    fn test_nack_responder_rfc4588_rtx() {
499        let mut chain = Registry::new()
500            .with(NackResponderBuilder::new().with_size(8).build())
501            .build();
502
503        // Bind local stream with NACK support AND RTX configured
504        let info = StreamInfo {
505            ssrc: 1,
506            ssrc_rtx: Some(2), // RTX SSRC
507            payload_type: 96,
508            payload_type_rtx: Some(97), // RTX payload type
509            clock_rate: 90000,
510            rtcp_feedback: vec![RTCPFeedback {
511                typ: "nack".to_string(),
512                parameter: "".to_string(),
513            }],
514            ..Default::default()
515        };
516        chain.bind_local_stream(&info);
517
518        let now = Instant::now();
519
520        // Send packets 10, 11, 12, 14, 15 (missing 13)
521        for seq in [10u16, 11, 12, 14, 15] {
522            let mut pkt = make_rtp_packet(1, seq, &[seq as u8]);
523            pkt.now = now;
524            chain.handle_write(pkt).unwrap();
525            chain.poll_write(); // Drain normal write
526        }
527
528        // Receive NACK for 11, 12, 13, 15
529        // nack_pair: packet_id=11, lost_packets=0b1011 means 11, 12, 13, 15
530        let mut nack = make_nack_packet(999, 1, vec![(11, 0b1011)]);
531        nack.now = now;
532        chain.handle_read(nack).unwrap();
533
534        // Should retransmit 11, 12, 15 (13 was never sent) using RTX format
535        let mut rtx_seq = 0u16;
536        for expected_original_seq in [11u16, 12, 15] {
537            let pkt = chain.poll_write();
538            assert!(
539                pkt.is_some(),
540                "Expected RTX packet for seq {}",
541                expected_original_seq
542            );
543
544            if let Some(tagged) = pkt {
545                if let Packet::Rtp(rtp) = tagged.message {
546                    // Verify RTX SSRC
547                    assert_eq!(rtp.header.ssrc, 2, "RTX packet should use RTX SSRC");
548                    // Verify RTX payload type
549                    assert_eq!(
550                        rtp.header.payload_type, 97,
551                        "RTX packet should use RTX payload type"
552                    );
553                    // Verify RTX sequence number (increments separately)
554                    assert_eq!(
555                        rtp.header.sequence_number, rtx_seq,
556                        "RTX seq should be {}",
557                        rtx_seq
558                    );
559                    rtx_seq += 1;
560
561                    // Verify payload: first 2 bytes should be original sequence number (big-endian)
562                    assert!(
563                        rtp.payload.len() >= 2,
564                        "RTX payload should have at least 2 bytes"
565                    );
566                    let original_seq_from_payload =
567                        u16::from_be_bytes([rtp.payload[0], rtp.payload[1]]);
568                    assert_eq!(
569                        original_seq_from_payload, expected_original_seq,
570                        "RTX payload should contain original seq"
571                    );
572
573                    // Verify original payload follows
574                    assert_eq!(
575                        rtp.payload[2..],
576                        [expected_original_seq as u8],
577                        "Original payload should follow seq number"
578                    );
579                } else {
580                    panic!("Expected RTP packet");
581                }
582            }
583        }
584
585        // No more packets
586        assert!(chain.poll_write().is_none());
587    }
588}