Skip to main content

rtc_interceptor/nack/
generator.rs

1//! NACK Generator Interceptor - Generates NACK requests for missing packets.
2
3use super::receive_log::ReceiveLog;
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::{Duration, Instant};
12
13/// Builder for the NackGeneratorInterceptor.
14///
15/// # Example
16///
17/// ```
18/// use rtc_interceptor::{Registry, NackGeneratorBuilder};
19/// use std::time::Duration;
20///
21/// let chain = Registry::new()
22///     .with(NackGeneratorBuilder::new()
23///         .with_size(512)
24///         .with_interval(Duration::from_millis(100))
25///         .with_skip_last_n(2)
26///         .build())
27///     .build();
28/// ```
29pub struct NackGeneratorBuilder<P> {
30    /// Size of the receive log (must be power of 2: 64, 128, ..., 32768).
31    size: u16,
32    /// Interval between NACK generation cycles.
33    interval: Duration,
34    /// Number of most recent packets to skip when generating NACKs.
35    skip_last_n: u16,
36    /// Maximum number of NACKs to send per missing packet (0 = unlimited).
37    max_nacks_per_packet: u16,
38    _phantom: PhantomData<P>,
39}
40
41impl<P> Default for NackGeneratorBuilder<P> {
42    fn default() -> Self {
43        Self {
44            size: 512,
45            interval: Duration::from_millis(100),
46            skip_last_n: 0,
47            max_nacks_per_packet: 0,
48            _phantom: PhantomData,
49        }
50    }
51}
52
53impl<P> NackGeneratorBuilder<P> {
54    /// Create a new builder with default settings.
55    pub fn new() -> Self {
56        Self::default()
57    }
58
59    /// Set the size of the receive log.
60    ///
61    /// Size must be a power of 2 between 64 and 32768 (inclusive).
62    pub fn with_size(mut self, size: u16) -> Self {
63        self.size = size;
64        self
65    }
66
67    /// Set the interval between NACK generation cycles.
68    pub fn with_interval(mut self, interval: Duration) -> Self {
69        self.interval = interval;
70        self
71    }
72
73    /// Set the number of most recent packets to skip when generating NACKs.
74    ///
75    /// This helps avoid generating NACKs for packets that are simply delayed
76    /// and haven't arrived yet.
77    pub fn with_skip_last_n(mut self, skip_last_n: u16) -> Self {
78        self.skip_last_n = skip_last_n;
79        self
80    }
81
82    /// Set the maximum number of NACKs to send per missing packet.
83    ///
84    /// Set to 0 (default) for unlimited NACKs.
85    pub fn with_max_nacks_per_packet(mut self, max: u16) -> Self {
86        self.max_nacks_per_packet = max;
87        self
88    }
89
90    /// Build the interceptor factory function.
91    pub fn build(self) -> impl FnOnce(P) -> NackGeneratorInterceptor<P> {
92        move |inner| {
93            NackGeneratorInterceptor::new(
94                inner,
95                self.size,
96                self.interval,
97                self.skip_last_n,
98                self.max_nacks_per_packet,
99            )
100        }
101    }
102}
103
104/// Interceptor that generates NACK requests for missing RTP packets.
105///
106/// This interceptor monitors incoming RTP packets on remote streams,
107/// tracks which sequence numbers have been received, and periodically
108/// generates RTCP TransportLayerNack packets for missing sequences.
109#[derive(Interceptor)]
110pub struct NackGeneratorInterceptor<P> {
111    #[next]
112    inner: P,
113
114    /// Configuration
115    size: u16,
116    interval: Duration,
117    skip_last_n: u16,
118    max_nacks_per_packet: u16,
119
120    /// Next timeout for NACK generation
121    next_timeout: Option<Instant>,
122
123    /// Sender SSRC for NACK packets
124    sender_ssrc: u32,
125
126    /// Receive logs per remote stream SSRC
127    receive_logs: HashMap<u32, ReceiveLog>,
128
129    /// NACK count per (SSRC, sequence number) for max_nacks_per_packet limiting
130    nack_counts: HashMap<u32, HashMap<u16, u16>>,
131
132    /// Queue for outgoing NACK packets
133    write_queue: VecDeque<TaggedPacket>,
134}
135
136impl<P> NackGeneratorInterceptor<P> {
137    fn new(
138        inner: P,
139        size: u16,
140        interval: Duration,
141        skip_last_n: u16,
142        max_nacks_per_packet: u16,
143    ) -> Self {
144        Self {
145            inner,
146            size,
147            interval,
148            skip_last_n,
149            max_nacks_per_packet,
150            next_timeout: None,
151            sender_ssrc: rand::random(),
152            receive_logs: HashMap::new(),
153            nack_counts: HashMap::new(),
154            write_queue: VecDeque::new(),
155        }
156    }
157
158    /// Generate NACKs for all streams with missing packets.
159    fn generate_nacks(&mut self, now: Instant) {
160        for (&ssrc, receive_log) in &self.receive_logs {
161            let missing = receive_log.missing_seq_numbers(self.skip_last_n);
162            if missing.is_empty() {
163                // Clear nack counts for this SSRC if no missing packets
164                self.nack_counts.remove(&ssrc);
165                continue;
166            }
167
168            // Initialize nack counts for this SSRC if needed
169            let nack_count = self.nack_counts.entry(ssrc).or_default();
170
171            // Filter by max_nacks_per_packet if configured
172            let filtered: Vec<u16> = if self.max_nacks_per_packet > 0 {
173                missing
174                    .iter()
175                    .filter(|&&seq| {
176                        let count = nack_count.entry(seq).or_insert(0);
177                        if *count < self.max_nacks_per_packet {
178                            *count += 1;
179                            true
180                        } else {
181                            false
182                        }
183                    })
184                    .copied()
185                    .collect()
186            } else {
187                missing.clone()
188            };
189
190            if filtered.is_empty() {
191                continue;
192            }
193
194            // Clean up nack counts for packets no longer missing
195            nack_count.retain(|seq, _| missing.contains(seq));
196
197            // Create NACK packet
198            let nack = rtcp::transport_feedbacks::transport_layer_nack::TransportLayerNack {
199                sender_ssrc: self.sender_ssrc,
200                media_ssrc: ssrc,
201                nacks: rtcp::transport_feedbacks::transport_layer_nack::nack_pairs_from_sequence_numbers(
202                    &filtered,
203                ),
204            };
205
206            self.write_queue.push_back(TaggedPacket {
207                now,
208                transport: TransportContext::default(),
209                message: Packet::Rtcp(vec![Box::new(nack)]),
210            });
211        }
212    }
213}
214
215#[interceptor]
216impl<P: Interceptor> NackGeneratorInterceptor<P> {
217    #[overrides]
218    fn handle_read(&mut self, msg: TaggedPacket) -> Result<(), Self::Error> {
219        // Track incoming RTP packets
220        if let Packet::Rtp(ref rtp_packet) = msg.message
221            && let Some(receive_log) = self.receive_logs.get_mut(&rtp_packet.header.ssrc)
222        {
223            receive_log.add(rtp_packet.header.sequence_number);
224
225            // Arm the NACK timer from the first tracked packet's instant. `None` means nothing is
226            // scheduled, so the interceptor asks for no wake-up until a stream is actually flowing.
227            if self.next_timeout.is_none() {
228                self.next_timeout = Some(msg.now + self.interval);
229            }
230        }
231
232        self.inner.handle_read(msg)
233    }
234
235    #[overrides]
236    fn poll_write(&mut self) -> Option<Self::Wout> {
237        // First drain generated NACK packets
238        if let Some(pkt) = self.write_queue.pop_front() {
239            return Some(pkt);
240        }
241        self.inner.poll_write()
242    }
243
244    #[overrides]
245    fn handle_timeout(&mut self, now: Self::Time) -> Result<(), Self::Error> {
246        if let Some(next_timeout) = self.next_timeout
247            && now >= next_timeout
248        {
249            self.next_timeout = Some(now + self.interval);
250            self.generate_nacks(now);
251        }
252
253        self.inner.handle_timeout(now)
254    }
255
256    #[overrides]
257    fn poll_timeout(&mut self) -> Option<Self::Time> {
258        match (self.next_timeout, self.inner.poll_timeout()) {
259            (Some(a), Some(b)) => Some(a.min(b)),
260            (Some(a), None) => Some(a),
261            (None, Some(b)) => Some(b),
262            (None, None) => None,
263        }
264    }
265
266    #[overrides]
267    fn bind_remote_stream(&mut self, info: &StreamInfo) {
268        if stream_supports_nack(info)
269            && let Some(receive_log) = ReceiveLog::new(self.size)
270        {
271            self.receive_logs.insert(info.ssrc, receive_log);
272        }
273        self.inner.bind_remote_stream(info);
274    }
275
276    #[overrides]
277    fn unbind_remote_stream(&mut self, info: &StreamInfo) {
278        self.receive_logs.remove(&info.ssrc);
279        self.nack_counts.remove(&info.ssrc);
280        self.inner.unbind_remote_stream(info);
281    }
282}
283
284#[cfg(test)]
285mod tests {
286    use super::*;
287    use crate::Registry;
288    use crate::stream_info::RTCPFeedback;
289    use sansio::Protocol;
290
291    fn make_rtp_packet(ssrc: u32, seq: u16) -> TaggedPacket {
292        TaggedPacket {
293            now: Instant::now(),
294            transport: Default::default(),
295            message: Packet::Rtp(rtp::Packet {
296                header: rtp::header::Header {
297                    ssrc,
298                    sequence_number: seq,
299                    ..Default::default()
300                },
301                ..Default::default()
302            }),
303        }
304    }
305
306    #[test]
307    fn test_nack_generator_builder_defaults() {
308        let chain = Registry::new()
309            .with(NackGeneratorBuilder::default().build())
310            .build();
311
312        assert_eq!(chain.size, 512);
313        assert_eq!(chain.interval, Duration::from_millis(100));
314        assert_eq!(chain.skip_last_n, 0);
315        assert_eq!(chain.max_nacks_per_packet, 0);
316    }
317
318    #[test]
319    fn test_nack_generator_builder_custom() {
320        let chain = Registry::new()
321            .with(
322                NackGeneratorBuilder::new()
323                    .with_size(1024)
324                    .with_interval(Duration::from_millis(50))
325                    .with_skip_last_n(3)
326                    .with_max_nacks_per_packet(5)
327                    .build(),
328            )
329            .build();
330
331        assert_eq!(chain.size, 1024);
332        assert_eq!(chain.interval, Duration::from_millis(50));
333        assert_eq!(chain.skip_last_n, 3);
334        assert_eq!(chain.max_nacks_per_packet, 5);
335    }
336
337    #[test]
338    fn test_nack_generator_no_nack_without_binding() {
339        let mut chain = Registry::new()
340            .with(
341                NackGeneratorBuilder::new()
342                    .with_interval(Duration::from_millis(100))
343                    .build(),
344            )
345            .build();
346
347        let now = Instant::now();
348
349        // Receive packets without binding stream (no receive log)
350        chain.handle_read(make_rtp_packet(12345, 0)).unwrap();
351        chain.handle_read(make_rtp_packet(12345, 2)).unwrap(); // Gap at 1
352
353        // Trigger timeout
354        let later = now + Duration::from_millis(200);
355        chain.handle_timeout(later).unwrap();
356
357        // No NACK should be generated (stream not bound)
358        assert!(chain.poll_write().is_none());
359    }
360
361    #[test]
362    fn test_nack_generator_generates_nack() {
363        let mut chain = Registry::new()
364            .with(
365                NackGeneratorBuilder::new()
366                    .with_size(64)
367                    .with_interval(Duration::from_millis(100))
368                    .build(),
369            )
370            .build();
371
372        // Bind remote stream with NACK support
373        let info = StreamInfo {
374            ssrc: 12345,
375            clock_rate: 90000,
376            rtcp_feedback: vec![RTCPFeedback {
377                typ: "nack".to_string(),
378                parameter: "".to_string(),
379            }],
380            ..Default::default()
381        };
382        chain.bind_remote_stream(&info);
383
384        let base_time = Instant::now();
385
386        // Receive packets with gap
387        let mut pkt = make_rtp_packet(12345, 10);
388        pkt.now = base_time;
389        chain.handle_read(pkt).unwrap();
390
391        let mut pkt = make_rtp_packet(12345, 12); // Gap at 11
392        pkt.now = base_time;
393        chain.handle_read(pkt).unwrap();
394
395        chain.poll_read();
396
397        // Trigger timeout
398        let later = base_time + Duration::from_millis(200);
399        chain.handle_timeout(later).unwrap();
400
401        // Should generate NACK for seq 11
402        let nack_pkt = chain.poll_write();
403        assert!(nack_pkt.is_some());
404
405        if let Some(tagged) = nack_pkt {
406            if let Packet::Rtcp(rtcp_packets) = tagged.message {
407                assert_eq!(rtcp_packets.len(), 1);
408                let nack = rtcp_packets[0]
409                    .as_any()
410                    .downcast_ref::<rtcp::transport_feedbacks::transport_layer_nack::TransportLayerNack>()
411                    .expect("Expected TransportLayerNack");
412                assert_eq!(nack.media_ssrc, 12345);
413                assert!(!nack.nacks.is_empty());
414            } else {
415                panic!("Expected RTCP packet");
416            }
417        }
418    }
419
420    #[test]
421    fn test_nack_generator_skip_last_n() {
422        let mut chain = Registry::new()
423            .with(
424                NackGeneratorBuilder::new()
425                    .with_size(64)
426                    .with_interval(Duration::from_millis(100))
427                    .with_skip_last_n(2)
428                    .build(),
429            )
430            .build();
431
432        let info = StreamInfo {
433            ssrc: 12345,
434            clock_rate: 90000,
435            rtcp_feedback: vec![RTCPFeedback {
436                typ: "nack".to_string(),
437                parameter: "".to_string(),
438            }],
439            ..Default::default()
440        };
441        chain.bind_remote_stream(&info);
442
443        let base_time = Instant::now();
444
445        // Receive: 10, 11, 12, 14, 16, 18 (gaps at 13, 15, 17)
446        for seq in [10u16, 11, 12, 14, 16, 18] {
447            let mut pkt = make_rtp_packet(12345, seq);
448            pkt.now = base_time;
449            chain.handle_read(pkt).unwrap();
450        }
451
452        // Trigger timeout
453        let later = base_time + Duration::from_millis(200);
454        chain.handle_timeout(later).unwrap();
455
456        // With skip_last_n=2, should only NACK for 13, 15 (not 17)
457        let nack_pkt = chain.poll_write();
458        assert!(nack_pkt.is_some());
459
460        if let Some(tagged) = nack_pkt
461            && let Packet::Rtcp(rtcp_packets) = tagged.message
462        {
463            let nack = rtcp_packets[0]
464                .as_any()
465                .downcast_ref::<rtcp::transport_feedbacks::transport_layer_nack::TransportLayerNack>()
466                .expect("Expected TransportLayerNack");
467
468            // Get all nacked sequence numbers
469            let mut nacked_seqs = Vec::new();
470            for nack_pair in &nack.nacks {
471                nacked_seqs.push(nack_pair.packet_id);
472                for i in 0..16 {
473                    if nack_pair.lost_packets & (1 << i) != 0 {
474                        nacked_seqs.push(nack_pair.packet_id.wrapping_add(i + 1));
475                    }
476                }
477            }
478
479            // Should contain 13, 15 but not 17
480            assert!(nacked_seqs.contains(&13));
481            assert!(nacked_seqs.contains(&15));
482            assert!(!nacked_seqs.contains(&17));
483        }
484    }
485
486    #[test]
487    fn test_nack_generator_unbind_removes_stream() {
488        let mut chain = Registry::new()
489            .with(
490                NackGeneratorBuilder::new()
491                    .with_size(64)
492                    .with_interval(Duration::from_millis(100))
493                    .build(),
494            )
495            .build();
496
497        let info = StreamInfo {
498            ssrc: 12345,
499            clock_rate: 90000,
500            rtcp_feedback: vec![RTCPFeedback {
501                typ: "nack".to_string(),
502                parameter: "".to_string(),
503            }],
504            ..Default::default()
505        };
506
507        chain.bind_remote_stream(&info);
508        assert!(chain.receive_logs.contains_key(&12345));
509
510        chain.unbind_remote_stream(&info);
511        assert!(!chain.receive_logs.contains_key(&12345));
512        assert!(!chain.nack_counts.contains_key(&12345));
513    }
514
515    #[test]
516    fn test_nack_generator_no_nack_support() {
517        let mut chain = Registry::new()
518            .with(
519                NackGeneratorBuilder::new()
520                    .with_size(64)
521                    .with_interval(Duration::from_millis(100))
522                    .build(),
523            )
524            .build();
525
526        // Bind stream without NACK support
527        let info = StreamInfo {
528            ssrc: 12345,
529            clock_rate: 90000,
530            rtcp_feedback: vec![], // No NACK support
531            ..Default::default()
532        };
533        chain.bind_remote_stream(&info);
534
535        // Should not create receive log
536        assert!(!chain.receive_logs.contains_key(&12345));
537    }
538}