rtc_interceptor/twcc/
sender.rs1use 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#[derive(Default)]
26pub struct TwccSenderBuilder {
27 initial_sequence_number: u16,
29}
30
31impl TwccSenderBuilder {
32 pub fn new() -> Self {
34 Self::default()
35 }
36
37 pub fn with_initial_sequence_number(mut self, sequence_number: u16) -> Self {
43 self.initial_sequence_number = sequence_number;
44 self
45 }
46
47 pub fn build(self) -> TwccSenderInterceptor {
49 TwccSenderInterceptor::new(self.initial_sequence_number)
50 }
51}
52
53struct LocalStream {
55 hdr_ext_id: u8,
57}
58
59#[derive(Default)]
74pub struct TwccSenderInterceptor {
75 next_sequence_number: u16,
77 streams: HashMap<u32, LocalStream>,
78 read_queue: VecDeque<TaggedPacket>,
80 write_queue: VecDeque<TaggedPacket>,
83}
84
85impl TwccSenderInterceptor {
86 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 #[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 #[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 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}