Skip to main content

rtc_interceptor/twcc/
sender.rs

1//! Tags outgoing RTP packets with transport-wide sequence numbers.
2
3use super::stream_supports_twcc;
4use crate::Interceptor;
5use crate::stream_info::StreamInfo;
6use crate::{Packet, TaggedPacket};
7use sansio::Protocol;
8use shared::error::Error;
9use shared::marshal::Marshal;
10use std::collections::HashMap;
11use std::collections::VecDeque;
12use std::time::Instant;
13
14/// Builder for the [`TwccSenderInterceptor`].
15///
16/// # Example
17///
18/// ```
19/// use rtc_interceptor::{Slot, Registry, TwccSenderBuilder};
20///
21/// let chain = Registry::new()
22///     .with(Slot::TwccSender, TwccSenderBuilder::new().build())
23///     .build();
24/// ```
25#[derive(Default)]
26pub struct TwccSenderBuilder {
27    /// The first transport-wide sequence number to hand out.
28    initial_sequence_number: u16,
29}
30
31impl TwccSenderBuilder {
32    /// Create a new builder with default settings.
33    pub fn new() -> Self {
34        Self::default()
35    }
36
37    /// Set the first transport-wide sequence number to hand out.
38    ///
39    /// Defaults to zero. The counter is shared across every local stream and wraps, so the
40    /// starting point only matters to a test that wants to pin the numbers it asserts on, or to
41    /// a session resuming a numbering it had already begun.
42    pub fn with_initial_sequence_number(mut self, sequence_number: u16) -> Self {
43        self.initial_sequence_number = sequence_number;
44        self
45    }
46
47    /// Build the interceptor.
48    pub fn build(self) -> TwccSenderInterceptor {
49        TwccSenderInterceptor::new(self.initial_sequence_number)
50    }
51}
52
53/// Per-stream state.
54struct LocalStream {
55    /// Header extension ID for transport-wide CC.
56    hdr_ext_id: u8,
57}
58
59/// Numbers every departing RTP packet so the remote can report on it
60/// ([`draft-holmer-rmcat-transport-wide-cc-extensions-01`]).
61///
62/// # Where this belongs in the chain
63///
64/// **Between the pacer and the send history**, near the wire. Numbering identifies a
65/// *transmission*, not a packet, so it has to happen after the pacer has decided what actually
66/// leaves and before the history records it — a retransmission is a separate transmission and gets
67/// its own number.
68///
69/// Under the nested chain this could not hold: a retransmission left from the NACK responder's own
70/// queue and never reached the tagger at all.
71///
72/// [`draft-holmer-rmcat-transport-wide-cc-extensions-01`]: http://www.ietf.org/id/draft-holmer-rmcat-transport-wide-cc-extensions-01
73#[derive(Default)]
74pub struct TwccSenderInterceptor {
75    /// Transport-wide sequence number counter, shared across all streams.
76    next_sequence_number: u16,
77    streams: HashMap<u32, LocalStream>,
78    /// Inbound packets ready for the next interceptor.
79    read_queue: VecDeque<TaggedPacket>,
80    /// Outbound packets ready for the next interceptor: what passed through, plus
81    /// anything this one generated.
82    write_queue: VecDeque<TaggedPacket>,
83}
84
85impl TwccSenderInterceptor {
86    /// A tagger with no streams bound yet.
87    fn new(initial_sequence_number: u16) -> Self {
88        Self {
89            next_sequence_number: initial_sequence_number,
90            ..Default::default()
91        }
92    }
93}
94
95impl Protocol<TaggedPacket, TaggedPacket, ()> for TwccSenderInterceptor {
96    type Rout = TaggedPacket;
97    type Wout = TaggedPacket;
98    type Eout = ();
99    type Error = Error;
100    type Time = Instant;
101
102    fn handle_read(&mut self, msg: TaggedPacket) -> Result<(), Self::Error> {
103        self.read_queue.push_back(msg);
104        Ok(())
105    }
106
107    fn poll_read(&mut self) -> Option<Self::Rout> {
108        self.read_queue.pop_front()
109    }
110
111    fn handle_write(&mut self, mut msg: TaggedPacket) -> Result<(), Self::Error> {
112        if let Packet::Rtp(ref mut rtp_packet) = msg.message.packet
113            && let Some(stream) = self.streams.get(&rtp_packet.header.ssrc)
114        {
115            let seq = self.next_sequence_number;
116            self.next_sequence_number = self.next_sequence_number.wrapping_add(1);
117
118            let tcc_ext = rtp::extension::transport_cc_extension::TransportCcExtension {
119                transport_sequence: seq,
120            };
121            if let Ok(ext_data) = tcc_ext.marshal() {
122                let _ = rtp_packet
123                    .header
124                    .set_extension(stream.hdr_ext_id, ext_data.freeze());
125            }
126        }
127        self.write_queue.push_back(msg);
128        Ok(())
129    }
130
131    fn poll_write(&mut self) -> Option<Self::Wout> {
132        self.write_queue.pop_front()
133    }
134
135    fn handle_timeout(&mut self, _now: Instant) -> Result<(), Self::Error> {
136        Ok(())
137    }
138
139    fn poll_timeout(&mut self) -> Option<Self::Time> {
140        None
141    }
142}
143
144impl Interceptor for TwccSenderInterceptor {
145    fn bind_local_stream(&mut self, info: &StreamInfo) {
146        if let Some(hdr_ext_id) = stream_supports_twcc(info)
147            && hdr_ext_id != 0
148        {
149            self.streams.insert(info.ssrc, LocalStream { hdr_ext_id });
150        }
151    }
152
153    fn unbind_local_stream(&mut self, info: &StreamInfo) {
154        self.streams.remove(&info.ssrc);
155    }
156
157    fn bind_remote_stream(&mut self, _info: &StreamInfo) {}
158
159    fn unbind_remote_stream(&mut self, _info: &StreamInfo) {}
160}
161
162#[cfg(test)]
163mod tests {
164    use super::*;
165    use crate::AttributedPacket;
166    use crate::chain::Chain;
167    use crate::stream_info::RTPHeaderExtension;
168    use sansio::Protocol;
169    use shared::TransportContext;
170    use shared::marshal::Unmarshal;
171    use std::time::Instant;
172
173    const TWCC_URI: &str =
174        "http://www.ietf.org/id/draft-holmer-rmcat-transport-wide-cc-extensions-01";
175
176    fn stream_info(ssrc: u32, ext_id: u16) -> StreamInfo {
177        StreamInfo {
178            ssrc,
179            rtp_header_extensions: vec![RTPHeaderExtension {
180                uri: TWCC_URI.to_owned(),
181                id: ext_id,
182            }],
183            ..Default::default()
184        }
185    }
186
187    fn packet(sequence_number: u16, ssrc: u32) -> TaggedPacket {
188        TaggedPacket {
189            now: Instant::now(),
190            transport: TransportContext::default(),
191            message: AttributedPacket::new(Packet::Rtp(rtp::Packet {
192                header: rtp::header::Header {
193                    version: 2,
194                    sequence_number,
195                    ssrc,
196                    ..Default::default()
197                },
198                payload: vec![0u8; 4].into(),
199            })),
200        }
201    }
202
203    fn tag_of(msg: &TaggedPacket, ext_id: u8) -> Option<u16> {
204        let Packet::Rtp(rtp) = &msg.message.packet else {
205            return None;
206        };
207        let data = rtp.header.get_extension(ext_id)?;
208        rtp::extension::transport_cc_extension::TransportCcExtension::unmarshal(&mut data.as_ref())
209            .ok()
210            .map(|e| e.transport_sequence)
211    }
212
213    fn chain() -> Chain {
214        let mut chain = Chain::new(vec![Box::new(TwccSenderBuilder::new().build())]);
215        chain.bind_local_stream(&stream_info(1, 5));
216        chain
217    }
218
219    #[test]
220    fn a_bound_stream_is_numbered_consecutively() {
221        let mut chain = chain();
222        let mut tags = Vec::new();
223        for sequence_number in 0..3 {
224            chain.handle_write(packet(sequence_number, 1)).unwrap();
225            while let Some(out) = chain.poll_write() {
226                tags.push(tag_of(&out, 5));
227            }
228        }
229        assert_eq!(vec![Some(0), Some(1), Some(2)], tags);
230    }
231
232    #[test]
233    fn an_unbound_stream_is_left_alone() {
234        let mut chain = chain();
235        chain.handle_write(packet(0, 999)).unwrap();
236        let out = chain.poll_write().expect("passes through");
237        assert_eq!(None, tag_of(&out, 5), "no extension added");
238    }
239
240    #[test]
241    fn unbinding_stops_the_numbering() {
242        let mut chain = chain();
243        chain.unbind_local_stream(&stream_info(1, 5));
244        chain.handle_write(packet(0, 1)).unwrap();
245        let out = chain.poll_write().expect("passes through");
246        assert_eq!(None, tag_of(&out, 5));
247    }
248
249    /// The counter is transport-wide: one sequence across every stream, not one per SSRC.
250    #[test]
251    fn the_counter_is_shared_across_streams() {
252        let mut chain = chain();
253        chain.bind_local_stream(&stream_info(2, 5));
254
255        let mut tags = Vec::new();
256        for ssrc in [1, 2, 1] {
257            chain.handle_write(packet(0, ssrc)).unwrap();
258            while let Some(out) = chain.poll_write() {
259                tags.push(tag_of(&out, 5));
260            }
261        }
262        assert_eq!(vec![Some(0), Some(1), Some(2)], tags);
263    }
264
265    /// An interceptor wireward of the tagger sees the number, which is what lets a send history key on it.
266    #[test]
267    fn a_stage_closer_to_the_wire_sees_the_tag() {
268        #[derive(Default)]
269        struct Recorder {
270            seen: Vec<Option<u16>>,
271            write_queue: VecDeque<TaggedPacket>,
272        }
273        impl Protocol<TaggedPacket, TaggedPacket, ()> for Recorder {
274            type Rout = TaggedPacket;
275            type Wout = TaggedPacket;
276            type Eout = ();
277            type Error = Error;
278            type Time = Instant;
279
280            fn handle_write(&mut self, msg: TaggedPacket) -> Result<(), Self::Error> {
281                self.seen.push(tag_of(&msg, 5));
282                self.write_queue.push_back(msg);
283                Ok(())
284            }
285
286            fn poll_write(&mut self) -> Option<Self::Wout> {
287                self.write_queue.pop_front()
288            }
289
290            fn handle_read(&mut self, _msg: TaggedPacket) -> Result<(), Self::Error> {
291                Ok(())
292            }
293
294            fn poll_read(&mut self) -> Option<Self::Rout> {
295                None
296            }
297
298            fn handle_timeout(&mut self, _now: Instant) -> Result<(), Self::Error> {
299                Ok(())
300            }
301
302            fn poll_timeout(&mut self) -> Option<Self::Time> {
303                None
304            }
305        }
306        impl Interceptor for Recorder {
307            fn bind_local_stream(&mut self, _info: &StreamInfo) {}
308            fn unbind_local_stream(&mut self, _info: &StreamInfo) {}
309            fn bind_remote_stream(&mut self, _info: &StreamInfo) {}
310            fn unbind_remote_stream(&mut self, _info: &StreamInfo) {}
311        }
312        // index 0 = wireward of the tagger at index 1, so it acts *after* on the write walk.
313        let mut chain = Chain::new(vec![
314            Box::new(Recorder::default()),
315            Box::new(TwccSenderBuilder::new().build()),
316        ]);
317        chain.bind_local_stream(&stream_info(1, 5));
318
319        chain.handle_write(packet(0, 1)).unwrap();
320        let out = chain.poll_write().expect("reaches the driver");
321        assert_eq!(Some(0), tag_of(&out, 5), "the packet leaves tagged");
322    }
323}