Skip to main content

rtc_interceptor/twcc/
receiver.rs

1//! TWCC Receiver Interceptor - tracks incoming packets and generates feedback.
2
3use super::recorder::Recorder;
4use super::stream_supports_twcc;
5use crate::stream_info::StreamInfo;
6use crate::{Interceptor, Packet, TaggedPacket, interceptor};
7use shared::TransportContext;
8use shared::error::Error;
9use shared::marshal::Unmarshal;
10use std::collections::{HashMap, VecDeque};
11use std::marker::PhantomData;
12use std::time::{Duration, Instant};
13
14/// Default interval for sending TWCC feedback.
15const DEFAULT_INTERVAL: Duration = Duration::from_millis(100);
16
17/// Builder for the TwccReceiverInterceptor.
18///
19/// # Example
20///
21/// ```
22/// use rtc_interceptor::{Registry, TwccReceiverBuilder};
23/// use std::time::Duration;
24///
25/// let chain = Registry::new()
26///     .with(TwccReceiverBuilder::new()
27///         .with_interval(Duration::from_millis(100))
28///         .build())
29///     .build();
30/// ```
31pub struct TwccReceiverBuilder<P> {
32    /// Interval between feedback reports.
33    interval: Duration,
34    _phantom: PhantomData<P>,
35}
36
37impl<P> Default for TwccReceiverBuilder<P> {
38    fn default() -> Self {
39        Self {
40            interval: DEFAULT_INTERVAL,
41            _phantom: PhantomData,
42        }
43    }
44}
45
46impl<P> TwccReceiverBuilder<P> {
47    /// Create a new builder with default settings.
48    pub fn new() -> Self {
49        Self::default()
50    }
51
52    /// Set the interval between feedback reports.
53    pub fn with_interval(mut self, interval: Duration) -> Self {
54        self.interval = interval;
55        self
56    }
57
58    /// Build the interceptor factory function.
59    pub fn build(self) -> impl FnOnce(P) -> TwccReceiverInterceptor<P> {
60        move |inner| TwccReceiverInterceptor::new(inner, self.interval)
61    }
62}
63
64/// Per-stream state for the receiver.
65struct RemoteStream {
66    /// Header extension ID for transport-wide CC.
67    hdr_ext_id: u8,
68}
69
70/// Interceptor that tracks incoming RTP packets and generates TWCC feedback.
71///
72/// This interceptor examines incoming RTP packets for transport-wide CC sequence
73/// numbers and periodically generates TransportLayerCC feedback packets.
74#[derive(Interceptor)]
75pub struct TwccReceiverInterceptor<P> {
76    #[next]
77    inner: P,
78
79    /// Configuration
80    interval: Duration,
81
82    /// Start time for calculating arrival times.
83    start_time: Option<Instant>,
84
85    /// TWCC recorder for building feedback.
86    recorder: Option<Recorder>,
87
88    /// Remote stream state per SSRC.
89    streams: HashMap<u32, RemoteStream>,
90
91    /// Queue for feedback packets.
92    write_queue: VecDeque<TaggedPacket>,
93
94    /// Next timeout for sending feedback.
95    next_timeout: Option<Instant>,
96}
97
98impl<P> TwccReceiverInterceptor<P> {
99    fn new(inner: P, interval: Duration) -> Self {
100        Self {
101            inner,
102            interval,
103            start_time: None,
104            recorder: None,
105            streams: HashMap::new(),
106            write_queue: VecDeque::new(),
107            next_timeout: None,
108        }
109    }
110
111    fn generate_feedback(&mut self, now: Instant) {
112        let Some(recorder) = self.recorder.as_mut() else {
113            return;
114        };
115
116        let packets = recorder.build_feedback_packet();
117        for pkt in packets {
118            self.write_queue.push_back(TaggedPacket {
119                now,
120                transport: TransportContext::default(),
121                message: Packet::Rtcp(vec![pkt]),
122            });
123        }
124    }
125}
126
127#[interceptor]
128impl<P: Interceptor> TwccReceiverInterceptor<P> {
129    #[overrides]
130    fn handle_read(&mut self, msg: TaggedPacket) -> Result<(), Self::Error> {
131        // Process incoming RTP packets with TWCC extension
132        if let Packet::Rtp(ref rtp_packet) = msg.message
133            && let Some(stream) = self.streams.get(&rtp_packet.header.ssrc)
134        {
135            // Initialize recorder on first packet
136            if self.recorder.is_none() {
137                // Use a random sender SSRC for feedback
138                self.recorder = Some(Recorder::new(rand::random()));
139                self.start_time = Some(msg.now);
140                self.next_timeout = Some(msg.now + self.interval);
141            }
142
143            // Extract transport CC sequence number
144            if let Some(ext_data) = rtp_packet.header.get_extension(stream.hdr_ext_id)
145                && let Ok(tcc) =
146                    rtp::extension::transport_cc_extension::TransportCcExtension::unmarshal(
147                        &mut ext_data.as_ref(),
148                    )
149            {
150                // Calculate arrival time in microseconds since start
151                let arrival_time = self
152                    .start_time
153                    .map(|start| msg.now.duration_since(start).as_micros() as i64)
154                    .unwrap_or(0);
155
156                if let Some(recorder) = self.recorder.as_mut() {
157                    recorder.record(rtp_packet.header.ssrc, tcc.transport_sequence, arrival_time);
158                }
159            }
160        }
161
162        self.inner.handle_read(msg)
163    }
164
165    #[overrides]
166    fn poll_write(&mut self) -> Option<Self::Wout> {
167        // First drain feedback packets
168        if let Some(pkt) = self.write_queue.pop_front() {
169            return Some(pkt);
170        }
171        self.inner.poll_write()
172    }
173
174    #[overrides]
175    fn handle_timeout(&mut self, now: Self::Time) -> Result<(), Self::Error> {
176        // Check if we need to send feedback
177        if let Some(timeout) = self.next_timeout
178            && now >= timeout
179        {
180            self.generate_feedback(now);
181            self.next_timeout = Some(now + self.interval);
182        }
183        self.inner.handle_timeout(now)
184    }
185
186    #[overrides]
187    fn poll_timeout(&mut self) -> Option<Self::Time> {
188        match (self.next_timeout, self.inner.poll_timeout()) {
189            (Some(a), Some(b)) => Some(a.min(b)),
190            (Some(a), None) => Some(a),
191            (None, Some(b)) => Some(b),
192            (None, None) => None,
193        }
194    }
195
196    #[overrides]
197    fn bind_remote_stream(&mut self, info: &StreamInfo) {
198        if let Some(hdr_ext_id) = stream_supports_twcc(info) {
199            // Don't track if ID is 0 (invalid)
200            if hdr_ext_id != 0 {
201                self.streams.insert(info.ssrc, RemoteStream { hdr_ext_id });
202            }
203        }
204        self.inner.bind_remote_stream(info);
205    }
206
207    #[overrides]
208    fn unbind_remote_stream(&mut self, info: &StreamInfo) {
209        self.streams.remove(&info.ssrc);
210        self.inner.unbind_remote_stream(info);
211    }
212}
213
214#[cfg(test)]
215mod tests {
216    use super::*;
217    use crate::Registry;
218    use crate::stream_info::RTPHeaderExtension;
219    use sansio::Protocol;
220    use shared::marshal::Marshal;
221
222    fn make_rtp_packet_with_twcc(
223        ssrc: u32,
224        seq: u16,
225        twcc_seq: u16,
226        hdr_ext_id: u8,
227    ) -> rtp::Packet {
228        let mut pkt = rtp::Packet {
229            header: rtp::header::Header {
230                ssrc,
231                sequence_number: seq,
232                ..Default::default()
233            },
234            payload: vec![].into(),
235        };
236
237        let tcc_ext = rtp::extension::transport_cc_extension::TransportCcExtension {
238            transport_sequence: twcc_seq,
239        };
240        if let Ok(ext_data) = tcc_ext.marshal() {
241            let _ = pkt.header.set_extension(hdr_ext_id, ext_data.freeze());
242        }
243
244        pkt
245    }
246
247    #[test]
248    fn test_twcc_receiver_builder_defaults() {
249        let chain = Registry::new()
250            .with(TwccReceiverBuilder::default().build())
251            .build();
252
253        assert_eq!(chain.interval, DEFAULT_INTERVAL);
254        assert!(chain.recorder.is_none());
255    }
256
257    #[test]
258    fn test_twcc_receiver_builder_custom_interval() {
259        let chain = Registry::new()
260            .with(
261                TwccReceiverBuilder::new()
262                    .with_interval(Duration::from_millis(50))
263                    .build(),
264            )
265            .build();
266
267        assert_eq!(chain.interval, Duration::from_millis(50));
268    }
269
270    #[test]
271    fn test_twcc_receiver_records_packets() {
272        let mut chain = Registry::new()
273            .with(TwccReceiverBuilder::new().build())
274            .build();
275
276        // Bind remote stream with TWCC support
277        let info = StreamInfo {
278            ssrc: 12345,
279            rtp_header_extensions: vec![RTPHeaderExtension {
280                uri: super::super::TRANSPORT_CC_URI.to_string(),
281                id: 5,
282            }],
283            ..Default::default()
284        };
285        chain.bind_remote_stream(&info);
286
287        let now = Instant::now();
288
289        // Receive RTP packet with TWCC extension
290        let rtp = make_rtp_packet_with_twcc(12345, 1, 0, 5);
291        let pkt = TaggedPacket {
292            now,
293            transport: Default::default(),
294            message: Packet::Rtp(rtp),
295        };
296        chain.handle_read(pkt).unwrap();
297
298        // Recorder should be initialized
299        assert!(chain.recorder.is_some());
300        assert!(chain.next_timeout.is_some());
301    }
302
303    #[test]
304    fn test_twcc_receiver_generates_feedback_on_timeout() {
305        let mut chain = Registry::new()
306            .with(
307                TwccReceiverBuilder::new()
308                    .with_interval(Duration::from_millis(100))
309                    .build(),
310            )
311            .build();
312
313        let info = StreamInfo {
314            ssrc: 12345,
315            rtp_header_extensions: vec![RTPHeaderExtension {
316                uri: super::super::TRANSPORT_CC_URI.to_string(),
317                id: 5,
318            }],
319            ..Default::default()
320        };
321        chain.bind_remote_stream(&info);
322
323        let start = Instant::now();
324
325        // Receive some packets
326        for i in 0..5u16 {
327            let rtp = make_rtp_packet_with_twcc(12345, i, i, 5);
328            let pkt = TaggedPacket {
329                now: start + Duration::from_millis(i as u64 * 10),
330                transport: Default::default(),
331                message: Packet::Rtp(rtp),
332            };
333            chain.handle_read(pkt).unwrap();
334        }
335
336        // Trigger timeout
337        let timeout_time = start + Duration::from_millis(150);
338        chain.handle_timeout(timeout_time).unwrap();
339
340        // Should have feedback packet
341        let feedback = chain.poll_write();
342        assert!(feedback.is_some());
343
344        if let Some(tagged) = feedback {
345            if let Packet::Rtcp(rtcp_packets) = tagged.message {
346                assert!(!rtcp_packets.is_empty());
347            } else {
348                panic!("Expected RTCP packet");
349            }
350        }
351    }
352
353    #[test]
354    fn test_twcc_receiver_no_feedback_without_binding() {
355        let mut chain = Registry::new()
356            .with(TwccReceiverBuilder::new().build())
357            .build();
358
359        let now = Instant::now();
360
361        // Receive packet without binding (no TWCC tracking)
362        let rtp = make_rtp_packet_with_twcc(12345, 1, 0, 5);
363        let pkt = TaggedPacket {
364            now,
365            transport: Default::default(),
366            message: Packet::Rtp(rtp),
367        };
368        chain.handle_read(pkt).unwrap();
369
370        // Recorder should not be initialized
371        assert!(chain.recorder.is_none());
372    }
373
374    #[test]
375    fn test_twcc_receiver_unbind_removes_stream() {
376        let mut chain = Registry::new()
377            .with(TwccReceiverBuilder::new().build())
378            .build();
379
380        let info = StreamInfo {
381            ssrc: 12345,
382            rtp_header_extensions: vec![RTPHeaderExtension {
383                uri: super::super::TRANSPORT_CC_URI.to_string(),
384                id: 5,
385            }],
386            ..Default::default()
387        };
388
389        chain.bind_remote_stream(&info);
390        assert!(chain.streams.contains_key(&12345));
391
392        chain.unbind_remote_stream(&info);
393        assert!(!chain.streams.contains_key(&12345));
394    }
395
396    #[test]
397    fn test_twcc_receiver_poll_timeout() {
398        let mut chain = Registry::new()
399            .with(TwccReceiverBuilder::new().build())
400            .build();
401
402        // No timeout initially
403        assert!(chain.poll_timeout().is_none());
404
405        let info = StreamInfo {
406            ssrc: 12345,
407            rtp_header_extensions: vec![RTPHeaderExtension {
408                uri: super::super::TRANSPORT_CC_URI.to_string(),
409                id: 5,
410            }],
411            ..Default::default()
412        };
413        chain.bind_remote_stream(&info);
414
415        let now = Instant::now();
416
417        // Receive a packet to initialize recorder
418        let rtp = make_rtp_packet_with_twcc(12345, 1, 0, 5);
419        let pkt = TaggedPacket {
420            now,
421            transport: Default::default(),
422            message: Packet::Rtp(rtp),
423        };
424        chain.handle_read(pkt).unwrap();
425
426        // Should have timeout now
427        let timeout = chain.poll_timeout();
428        assert!(timeout.is_some());
429    }
430}