1use 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
14const DEFAULT_INTERVAL: Duration = Duration::from_millis(100);
16
17pub struct TwccReceiverBuilder<P> {
32 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 pub fn new() -> Self {
49 Self::default()
50 }
51
52 pub fn with_interval(mut self, interval: Duration) -> Self {
54 self.interval = interval;
55 self
56 }
57
58 pub fn build(self) -> impl FnOnce(P) -> TwccReceiverInterceptor<P> {
60 move |inner| TwccReceiverInterceptor::new(inner, self.interval)
61 }
62}
63
64struct RemoteStream {
66 hdr_ext_id: u8,
68}
69
70#[derive(Interceptor)]
75pub struct TwccReceiverInterceptor<P> {
76 #[next]
77 inner: P,
78
79 interval: Duration,
81
82 start_time: Option<Instant>,
84
85 recorder: Option<Recorder>,
87
88 streams: HashMap<u32, RemoteStream>,
90
91 write_queue: VecDeque<TaggedPacket>,
93
94 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 if let Packet::Rtp(ref rtp_packet) = msg.message
133 && let Some(stream) = self.streams.get(&rtp_packet.header.ssrc)
134 {
135 if self.recorder.is_none() {
137 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 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 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 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 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 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 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 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 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 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 let timeout_time = start + Duration::from_millis(150);
338 chain.handle_timeout(timeout_time).unwrap();
339
340 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 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 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 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 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 let timeout = chain.poll_timeout();
428 assert!(timeout.is_some());
429 }
430}