Skip to main content

rtc_interceptor/report/
sender.rs

1//! Sender Report Interceptor - Filters hop-by-hop RTCP feedback.
2
3use super::sender_stream::SenderStream;
4use crate::stream_info::StreamInfo;
5use crate::{Interceptor, Packet, TaggedPacket, interceptor};
6use rtcp::header::PacketType;
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 SenderReportInterceptor.
14///
15/// # Example
16///
17/// ```
18/// use rtc_interceptor::{Registry, SenderReportBuilder};
19/// use std::time::Duration;
20///
21/// // With default interval (1 second)
22/// let chain = Registry::new()
23///     .with(SenderReportBuilder::new().build())
24///     .build();
25///
26/// // With custom interval
27/// let chain = Registry::new()
28///     .with(SenderReportBuilder::new().with_interval(Duration::from_millis(500)).build())
29///     .build();
30///
31/// // With use_latest_packet enabled
32/// let chain = Registry::new()
33///     .with(SenderReportBuilder::new().with_use_latest_packet().build())
34///     .build();
35/// ```
36pub struct SenderReportBuilder<P> {
37    /// Interval between sender reports.
38    interval: Duration,
39    /// Whether to always use the latest packet, even if out-of-order.
40    use_latest_packet: bool,
41    _phantom: PhantomData<P>,
42}
43
44impl<P> Default for SenderReportBuilder<P> {
45    fn default() -> Self {
46        Self {
47            interval: Duration::from_secs(1),
48            use_latest_packet: false,
49            _phantom: PhantomData,
50        }
51    }
52}
53
54impl<P> SenderReportBuilder<P> {
55    /// Create a new builder with default settings.
56    ///
57    /// Default interval is 1 second.
58    pub fn new() -> Self {
59        Self::default()
60    }
61
62    /// Set a custom interval between sender reports.
63    ///
64    /// # Example
65    ///
66    /// ```
67    /// use rtc_interceptor::{Registry, SenderReportBuilder};
68    /// use std::time::Duration;
69    ///
70    /// // The builder is generic over the next layer, so its type is pinned by `with`.
71    /// let registry = Registry::new().with(
72    ///     SenderReportBuilder::new()
73    ///         .with_interval(Duration::from_millis(500))
74    ///         .build(),
75    /// );
76    /// ```
77    pub fn with_interval(mut self, interval: Duration) -> Self {
78        self.interval = interval;
79        self
80    }
81
82    /// Enable always using the latest packet for timestamp tracking,
83    /// even if it appears to be out-of-order based on sequence numbers.
84    ///
85    /// By default (disabled), only in-order packets update the RTP↔NTP
86    /// timestamp correlation. This prevents out-of-order packets from
87    /// corrupting the timestamp mapping.
88    ///
89    /// Enable this option when:
90    /// - Packets are guaranteed to arrive in order
91    /// - The application reorders packets before the interceptor
92    /// - You want the sender report to always reflect the most recent packet
93    ///
94    /// # Example
95    ///
96    /// ```
97    /// use rtc_interceptor::{Registry, SenderReportBuilder};
98    ///
99    /// let registry =
100    ///     Registry::new().with(SenderReportBuilder::new().with_use_latest_packet().build());
101    /// ```
102    pub fn with_use_latest_packet(mut self) -> Self {
103        self.use_latest_packet = true;
104        self
105    }
106
107    /// Create a builder function for use with Registry.
108    ///
109    /// This returns a closure that can be passed to `Registry::with()`.
110    ///
111    /// # Example
112    ///
113    /// ```
114    /// use rtc_interceptor::{Registry, SenderReportBuilder};
115    ///
116    /// let registry = Registry::new()
117    ///     .with(SenderReportBuilder::new().build());
118    /// ```
119    pub fn build(self) -> impl FnOnce(P) -> SenderReportInterceptor<P> {
120        move |inner| SenderReportInterceptor::new(inner, self.interval, self.use_latest_packet)
121    }
122}
123
124/// Interceptor that filters hop-by-hop RTCP reports.
125///
126/// This interceptor filters out RTCP Receiver Reports and Transport-Specific
127/// Feedback, which are hop-by-hop reports that should not be forwarded
128/// end-to-end.
129///
130/// # Type Parameters
131///
132/// - `P`: The inner protocol being wrapped
133///
134/// # Example
135///
136/// ```
137/// use rtc_interceptor::{Registry, SenderReportBuilder};
138///
139/// let chain = Registry::new()
140///     .with(SenderReportBuilder::new().build())
141///     .build();
142/// ```
143#[derive(Interceptor)]
144pub struct SenderReportInterceptor<P> {
145    #[next]
146    inner: P,
147
148    interval: Duration,
149    eto: Instant,
150
151    /// Whether to always use the latest packet, even if out-of-order.
152    use_latest_packet: bool,
153
154    streams: HashMap<u32, SenderStream>,
155
156    read_queue: VecDeque<TaggedPacket>,
157    write_queue: VecDeque<TaggedPacket>,
158}
159
160impl<P> SenderReportInterceptor<P> {
161    /// Create a new SenderReportInterceptor.
162    fn new(inner: P, interval: Duration, use_latest_packet: bool) -> Self {
163        Self {
164            inner,
165
166            interval,
167            eto: Instant::now(),
168
169            use_latest_packet,
170
171            streams: HashMap::new(),
172
173            read_queue: VecDeque::new(),
174            write_queue: VecDeque::new(),
175        }
176    }
177
178    /// Check if an RTCP packet type should be filtered.
179    ///
180    /// Returns `true` for hop-by-hop report types that should not be forwarded:
181    /// - Receiver Report (201)
182    /// - Transport-Specific Feedback (205)
183    fn should_filter(packet_type: PacketType) -> bool {
184        packet_type == PacketType::ReceiverReport
185            || (packet_type == PacketType::TransportSpecificFeedback)
186    }
187
188    /// Get a reference to the inner protocol.
189    fn inner(&self) -> &P {
190        &self.inner
191    }
192
193    /// Get a mutable reference to the inner protocol.
194    fn inner_mut(&mut self) -> &mut P {
195        &mut self.inner
196    }
197}
198
199#[interceptor]
200impl<P: Interceptor> SenderReportInterceptor<P> {
201    #[overrides]
202    fn handle_write(&mut self, msg: TaggedPacket) -> Result<(), Self::Error> {
203        if let Packet::Rtp(rtp_packet) = &msg.message
204            && let Some(stream) = self.streams.get_mut(&rtp_packet.header.ssrc)
205        {
206            stream.process_rtp(msg.now, rtp_packet);
207        }
208
209        self.inner.handle_write(msg)
210    }
211
212    #[overrides]
213    fn poll_write(&mut self) -> Option<Self::Wout> {
214        // First drain generated RTCP reports
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 handle_timeout(&mut self, now: Self::Time) -> Result<(), Self::Error> {
223        if self.eto <= now {
224            self.eto = now + self.interval;
225
226            for stream in self.streams.values_mut() {
227                let rr = stream.generate_report(now);
228                self.write_queue.push_back(TaggedPacket {
229                    now,
230                    transport: TransportContext::default(),
231                    message: Packet::Rtcp(vec![Box::new(rr)]),
232                });
233            }
234        }
235
236        self.inner.handle_timeout(now)
237    }
238
239    #[overrides]
240    fn poll_timeout(&mut self) -> Option<Self::Time> {
241        if let Some(eto) = self.inner.poll_timeout()
242            && eto < self.eto
243        {
244            Some(eto)
245        } else {
246            Some(self.eto)
247        }
248    }
249
250    #[overrides]
251    fn bind_local_stream(&mut self, info: &StreamInfo) {
252        let stream = SenderStream::new(info.ssrc, info.clock_rate, self.use_latest_packet);
253        self.streams.insert(info.ssrc, stream);
254
255        self.inner.bind_local_stream(info);
256    }
257
258    #[overrides]
259    fn unbind_local_stream(&mut self, info: &StreamInfo) {
260        self.streams.remove(&info.ssrc);
261
262        self.inner.unbind_local_stream(info);
263    }
264}
265
266#[cfg(test)]
267mod tests {
268    use super::*;
269    use crate::{NoopInterceptor, Registry};
270    use sansio::Protocol;
271
272    fn dummy_rtp_packet() -> TaggedPacket {
273        TaggedPacket {
274            now: Instant::now(),
275            transport: Default::default(),
276            message: crate::Packet::Rtp(rtp::Packet::default()),
277        }
278    }
279
280    #[test]
281    fn test_sender_report_builder_default() {
282        // Build with default interval (1 second)
283        let chain = Registry::new()
284            .with(SenderReportBuilder::default().build())
285            .build();
286
287        assert_eq!(chain.interval, Duration::from_secs(1));
288    }
289
290    #[test]
291    fn test_sender_report_builder_with_custom_interval() {
292        // Build with custom interval
293        let chain = Registry::new()
294            .with(
295                SenderReportBuilder::default()
296                    .with_interval(Duration::from_millis(500))
297                    .build(),
298            )
299            .build();
300
301        assert_eq!(chain.interval, Duration::from_millis(500));
302    }
303
304    #[test]
305    fn test_sender_report_chain_handle_read_write() {
306        // Build a chain and test packet flow
307        let mut chain = Registry::new()
308            .with(SenderReportBuilder::default().build())
309            .build();
310
311        // Test read path
312        let pkt = dummy_rtp_packet();
313        chain.handle_read(pkt).unwrap();
314        assert!(chain.poll_read().is_some());
315
316        // Test write path
317        let pkt2 = dummy_rtp_packet();
318        let pkt2_message = pkt2.message.clone();
319        chain.handle_write(pkt2).unwrap();
320        assert_eq!(chain.poll_write().unwrap().message, pkt2_message);
321    }
322
323    #[test]
324    fn test_should_filter() {
325        // Receiver Report (RR) - should filter
326        assert!(SenderReportInterceptor::<NoopInterceptor>::should_filter(
327            PacketType::ReceiverReport
328        ));
329
330        // Transport-Specific Feedback - should filter
331        assert!(SenderReportInterceptor::<NoopInterceptor>::should_filter(
332            PacketType::TransportSpecificFeedback
333        ));
334
335        // Sender Report (SR) - should NOT filter
336        assert!(!SenderReportInterceptor::<NoopInterceptor>::should_filter(
337            PacketType::SenderReport
338        ));
339
340        // Source Description (SDES) - should NOT filter
341        assert!(!SenderReportInterceptor::<NoopInterceptor>::should_filter(
342            PacketType::SourceDescription
343        ));
344
345        // Goodbye (BYE) - should NOT filter
346        assert!(!SenderReportInterceptor::<NoopInterceptor>::should_filter(
347            PacketType::Goodbye
348        ));
349    }
350
351    #[test]
352    fn test_inner_access() {
353        let mut chain = Registry::new()
354            .with(SenderReportBuilder::default().build())
355            .build();
356
357        // Test immutable access
358        let _ = chain.inner();
359
360        // Test mutable access - can modify inner
361        let pkt = dummy_rtp_packet();
362        let pkt_message = pkt.message.clone();
363        chain.inner_mut().handle_write(pkt).unwrap();
364        assert_eq!(chain.inner_mut().poll_write().unwrap().message, pkt_message);
365    }
366
367    #[test]
368    fn test_use_latest_packet_option() {
369        // Build with use_latest_packet enabled
370        let chain = Registry::new()
371            .with(
372                SenderReportBuilder::default()
373                    .with_use_latest_packet()
374                    .build(),
375            )
376            .build();
377
378        assert!(chain.use_latest_packet);
379
380        // Build without use_latest_packet (default)
381        let chain_default = Registry::new()
382            .with(SenderReportBuilder::default().build())
383            .build();
384
385        assert!(!chain_default.use_latest_packet);
386    }
387
388    #[test]
389    fn test_use_latest_packet_combined_options() {
390        // Test combining multiple options
391        let chain = Registry::new()
392            .with(
393                SenderReportBuilder::default()
394                    .with_interval(Duration::from_millis(250))
395                    .with_use_latest_packet()
396                    .build(),
397            )
398            .build();
399
400        assert_eq!(chain.interval, Duration::from_millis(250));
401        assert!(chain.use_latest_packet);
402    }
403
404    #[test]
405    fn test_sender_report_generation_on_timeout() {
406        // Port of pion's TestSenderInterceptor - tests full timeout/report cycle
407        // No ticker mocking needed - sans-I/O pattern lets us control time directly
408        let mut chain = Registry::new()
409            .with(
410                SenderReportBuilder::default()
411                    .with_interval(Duration::from_secs(1))
412                    .build(),
413            )
414            .build();
415
416        // Bind a local stream
417        let info = StreamInfo {
418            ssrc: 123456,
419            clock_rate: 90000,
420            ..Default::default()
421        };
422        chain.bind_local_stream(&info);
423
424        let base_time = Instant::now();
425
426        // Send some RTP packets through the write path
427        for i in 0..5u16 {
428            let pkt = TaggedPacket {
429                now: base_time,
430                transport: Default::default(),
431                message: Packet::Rtp(rtp::Packet {
432                    header: rtp::header::Header {
433                        ssrc: 123456,
434                        sequence_number: i,
435                        timestamp: i as u32 * 3000,
436                        ..Default::default()
437                    },
438                    payload: vec![0u8; 100].into(),
439                    ..Default::default()
440                }),
441            };
442            chain.handle_write(pkt).unwrap();
443            // Drain the write queue
444            chain.poll_write();
445        }
446
447        // First timeout triggers report generation (eto was set at construction)
448        chain.handle_timeout(base_time).unwrap();
449
450        // Drain any reports from initial timeout
451        while chain.poll_write().is_some() {}
452
453        // Advance time past the interval
454        let later_time = base_time + Duration::from_secs(2);
455        chain.handle_timeout(later_time).unwrap();
456
457        // Now a sender report should be generated
458        let report = chain.poll_write();
459        assert!(report.is_some());
460
461        if let Some(tagged) = report {
462            if let Packet::Rtcp(rtcp_packets) = tagged.message {
463                assert_eq!(rtcp_packets.len(), 1);
464                let sr = rtcp_packets[0]
465                    .as_any()
466                    .downcast_ref::<rtcp::sender_report::SenderReport>()
467                    .expect("Expected SenderReport");
468                assert_eq!(sr.ssrc, 123456);
469                assert_eq!(sr.packet_count, 5);
470                assert_eq!(sr.octet_count, 500);
471            } else {
472                panic!("Expected RTCP packet");
473            }
474        }
475    }
476
477    #[test]
478    fn test_sender_report_multiple_streams() {
479        // Test that multiple streams each generate their own sender reports
480        let mut chain = Registry::new()
481            .with(
482                SenderReportBuilder::default()
483                    .with_interval(Duration::from_secs(1))
484                    .build(),
485            )
486            .build();
487
488        // Bind two local streams
489        let info1 = StreamInfo {
490            ssrc: 111111,
491            clock_rate: 90000,
492            ..Default::default()
493        };
494        let info2 = StreamInfo {
495            ssrc: 222222,
496            clock_rate: 48000,
497            ..Default::default()
498        };
499        chain.bind_local_stream(&info1);
500        chain.bind_local_stream(&info2);
501
502        let base_time = Instant::now();
503
504        // Send packets on stream 1
505        for i in 0..3u16 {
506            let pkt = TaggedPacket {
507                now: base_time,
508                transport: Default::default(),
509                message: Packet::Rtp(rtp::Packet {
510                    header: rtp::header::Header {
511                        ssrc: 111111,
512                        sequence_number: i,
513                        timestamp: i as u32 * 3000,
514                        ..Default::default()
515                    },
516                    payload: vec![0u8; 50].into(),
517                    ..Default::default()
518                }),
519            };
520            chain.handle_write(pkt).unwrap();
521            chain.poll_write();
522        }
523
524        // Send packets on stream 2
525        for i in 0..7u16 {
526            let pkt = TaggedPacket {
527                now: base_time,
528                transport: Default::default(),
529                message: Packet::Rtp(rtp::Packet {
530                    header: rtp::header::Header {
531                        ssrc: 222222,
532                        sequence_number: i,
533                        timestamp: i as u32 * 960,
534                        ..Default::default()
535                    },
536                    payload: vec![0u8; 200].into(),
537                    ..Default::default()
538                }),
539            };
540            chain.handle_write(pkt).unwrap();
541            chain.poll_write();
542        }
543
544        // Trigger timeout
545        let later_time = base_time + Duration::from_secs(2);
546        chain.handle_timeout(later_time).unwrap();
547
548        // Should get two sender reports
549        let mut ssrcs = vec![];
550        let mut packet_counts = vec![];
551        let mut octet_counts = vec![];
552
553        while let Some(tagged) = chain.poll_write() {
554            if let Packet::Rtcp(rtcp_packets) = tagged.message {
555                for rtcp_pkt in rtcp_packets {
556                    if let Some(sr) = rtcp_pkt
557                        .as_any()
558                        .downcast_ref::<rtcp::sender_report::SenderReport>()
559                    {
560                        ssrcs.push(sr.ssrc);
561                        packet_counts.push(sr.packet_count);
562                        octet_counts.push(sr.octet_count);
563                    }
564                }
565            }
566        }
567
568        assert_eq!(ssrcs.len(), 2);
569        assert!(ssrcs.contains(&111111));
570        assert!(ssrcs.contains(&222222));
571
572        // Find stream 1's report
573        let idx1 = ssrcs.iter().position(|&s| s == 111111).unwrap();
574        assert_eq!(packet_counts[idx1], 3);
575        assert_eq!(octet_counts[idx1], 150);
576
577        // Find stream 2's report
578        let idx2 = ssrcs.iter().position(|&s| s == 222222).unwrap();
579        assert_eq!(packet_counts[idx2], 7);
580        assert_eq!(octet_counts[idx2], 1400);
581    }
582
583    #[test]
584    fn test_sender_report_unbind_stream() {
585        // Test that unbinding a stream stops generating reports for it
586        let mut chain = Registry::new()
587            .with(
588                SenderReportBuilder::default()
589                    .with_interval(Duration::from_secs(1))
590                    .build(),
591            )
592            .build();
593
594        let info = StreamInfo {
595            ssrc: 123456,
596            clock_rate: 90000,
597            ..Default::default()
598        };
599        chain.bind_local_stream(&info);
600
601        let base_time = Instant::now();
602
603        // Send some packets
604        let pkt = TaggedPacket {
605            now: base_time,
606            transport: Default::default(),
607            message: Packet::Rtp(rtp::Packet {
608                header: rtp::header::Header {
609                    ssrc: 123456,
610                    sequence_number: 0,
611                    timestamp: 0,
612                    ..Default::default()
613                },
614                payload: vec![0u8; 100].into(),
615                ..Default::default()
616            }),
617        };
618        chain.handle_write(pkt).unwrap();
619        chain.poll_write();
620
621        // Unbind the stream
622        chain.unbind_local_stream(&info);
623
624        // Trigger timeout
625        let later_time = base_time + Duration::from_secs(2);
626        chain.handle_timeout(later_time).unwrap();
627
628        // No report should be generated (stream was unbound)
629        assert!(chain.poll_write().is_none());
630    }
631
632    #[test]
633    fn test_poll_timeout_returns_earliest() {
634        // Test that poll_timeout returns the earliest timeout
635        let mut chain = Registry::new()
636            .with(
637                SenderReportBuilder::default()
638                    .with_interval(Duration::from_secs(5))
639                    .build(),
640            )
641            .build();
642
643        // The interceptor should return its own eto
644        let timeout = chain.poll_timeout();
645        assert!(timeout.is_some());
646    }
647}