1use 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
13pub struct SenderReportBuilder<P> {
37 interval: Duration,
39 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 pub fn new() -> Self {
59 Self::default()
60 }
61
62 pub fn with_interval(mut self, interval: Duration) -> Self {
78 self.interval = interval;
79 self
80 }
81
82 pub fn with_use_latest_packet(mut self) -> Self {
103 self.use_latest_packet = true;
104 self
105 }
106
107 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#[derive(Interceptor)]
144pub struct SenderReportInterceptor<P> {
145 #[next]
146 inner: P,
147
148 interval: Duration,
149 next_timeout: Option<Instant>,
150
151 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 fn new(inner: P, interval: Duration, use_latest_packet: bool) -> Self {
163 Self {
164 inner,
165
166 interval,
167 next_timeout: None,
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 fn should_filter(packet_type: PacketType) -> bool {
184 packet_type == PacketType::ReceiverReport
185 || (packet_type == PacketType::TransportSpecificFeedback)
186 }
187
188 fn inner(&self) -> &P {
190 &self.inner
191 }
192
193 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 if self.next_timeout.is_none() {
210 self.next_timeout = Some(msg.now + self.interval);
211 }
212 }
213
214 self.inner.handle_write(msg)
215 }
216
217 #[overrides]
218 fn poll_write(&mut self) -> Option<Self::Wout> {
219 if let Some(pkt) = self.write_queue.pop_front() {
221 return Some(pkt);
222 }
223 self.inner.poll_write()
224 }
225
226 #[overrides]
227 fn handle_timeout(&mut self, now: Self::Time) -> Result<(), Self::Error> {
228 if let Some(next_timeout) = self.next_timeout
229 && now >= next_timeout
230 {
231 self.next_timeout = Some(now + self.interval);
232
233 for stream in self.streams.values_mut() {
234 if let Some(rr) = stream.generate_report(now) {
235 self.write_queue.push_back(TaggedPacket {
236 now,
237 transport: TransportContext::default(),
238 message: Packet::Rtcp(vec![Box::new(rr)]),
239 });
240 }
241 }
242 }
243
244 self.inner.handle_timeout(now)
245 }
246
247 #[overrides]
248 fn poll_timeout(&mut self) -> Option<Self::Time> {
249 match (self.next_timeout, self.inner.poll_timeout()) {
250 (Some(a), Some(b)) => Some(a.min(b)),
251 (Some(a), None) => Some(a),
252 (None, Some(b)) => Some(b),
253 (None, None) => None,
254 }
255 }
256
257 #[overrides]
258 fn bind_local_stream(&mut self, info: &StreamInfo) {
259 let stream = SenderStream::new(info.ssrc, info.clock_rate, self.use_latest_packet);
260 self.streams.insert(info.ssrc, stream);
261
262 self.inner.bind_local_stream(info);
263 }
264
265 #[overrides]
266 fn unbind_local_stream(&mut self, info: &StreamInfo) {
267 self.streams.remove(&info.ssrc);
268
269 self.inner.unbind_local_stream(info);
270 }
271}
272
273#[cfg(test)]
274mod tests {
275 use super::*;
276 use crate::{NoopInterceptor, Registry};
277 use sansio::Protocol;
278
279 fn dummy_rtp_packet() -> TaggedPacket {
280 TaggedPacket {
281 now: Instant::now(),
282 transport: Default::default(),
283 message: crate::Packet::Rtp(rtp::Packet::default()),
284 }
285 }
286
287 #[test]
288 fn test_sender_report_builder_default() {
289 let chain = Registry::new()
291 .with(SenderReportBuilder::default().build())
292 .build();
293
294 assert_eq!(chain.interval, Duration::from_secs(1));
295 }
296
297 #[test]
298 fn test_sender_report_builder_with_custom_interval() {
299 let chain = Registry::new()
301 .with(
302 SenderReportBuilder::default()
303 .with_interval(Duration::from_millis(500))
304 .build(),
305 )
306 .build();
307
308 assert_eq!(chain.interval, Duration::from_millis(500));
309 }
310
311 #[test]
312 fn test_sender_report_chain_handle_read_write() {
313 let mut chain = Registry::new()
315 .with(SenderReportBuilder::default().build())
316 .build();
317
318 let pkt = dummy_rtp_packet();
320 chain.handle_read(pkt).unwrap();
321 assert!(chain.poll_read().is_some());
322
323 let pkt2 = dummy_rtp_packet();
325 let pkt2_message = pkt2.message.clone();
326 chain.handle_write(pkt2).unwrap();
327 assert_eq!(chain.poll_write().unwrap().message, pkt2_message);
328 }
329
330 #[test]
331 fn test_should_filter() {
332 assert!(SenderReportInterceptor::<NoopInterceptor>::should_filter(
334 PacketType::ReceiverReport
335 ));
336
337 assert!(SenderReportInterceptor::<NoopInterceptor>::should_filter(
339 PacketType::TransportSpecificFeedback
340 ));
341
342 assert!(!SenderReportInterceptor::<NoopInterceptor>::should_filter(
344 PacketType::SenderReport
345 ));
346
347 assert!(!SenderReportInterceptor::<NoopInterceptor>::should_filter(
349 PacketType::SourceDescription
350 ));
351
352 assert!(!SenderReportInterceptor::<NoopInterceptor>::should_filter(
354 PacketType::Goodbye
355 ));
356 }
357
358 #[test]
359 fn test_inner_access() {
360 let mut chain = Registry::new()
361 .with(SenderReportBuilder::default().build())
362 .build();
363
364 let _ = chain.inner();
366
367 let pkt = dummy_rtp_packet();
369 let pkt_message = pkt.message.clone();
370 chain.inner_mut().handle_write(pkt).unwrap();
371 assert_eq!(chain.inner_mut().poll_write().unwrap().message, pkt_message);
372 }
373
374 #[test]
375 fn test_use_latest_packet_option() {
376 let chain = Registry::new()
378 .with(
379 SenderReportBuilder::default()
380 .with_use_latest_packet()
381 .build(),
382 )
383 .build();
384
385 assert!(chain.use_latest_packet);
386
387 let chain_default = Registry::new()
389 .with(SenderReportBuilder::default().build())
390 .build();
391
392 assert!(!chain_default.use_latest_packet);
393 }
394
395 #[test]
396 fn test_use_latest_packet_combined_options() {
397 let chain = Registry::new()
399 .with(
400 SenderReportBuilder::default()
401 .with_interval(Duration::from_millis(250))
402 .with_use_latest_packet()
403 .build(),
404 )
405 .build();
406
407 assert_eq!(chain.interval, Duration::from_millis(250));
408 assert!(chain.use_latest_packet);
409 }
410
411 #[test]
412 fn test_sender_report_generation_on_timeout() {
413 let mut chain = Registry::new()
416 .with(
417 SenderReportBuilder::default()
418 .with_interval(Duration::from_secs(1))
419 .build(),
420 )
421 .build();
422
423 let info = StreamInfo {
425 ssrc: 123456,
426 clock_rate: 90000,
427 ..Default::default()
428 };
429 chain.bind_local_stream(&info);
430
431 let base_time = Instant::now();
432
433 for i in 0..5u16 {
435 let pkt = TaggedPacket {
436 now: base_time,
437 transport: Default::default(),
438 message: Packet::Rtp(rtp::Packet {
439 header: rtp::header::Header {
440 ssrc: 123456,
441 sequence_number: i,
442 timestamp: i as u32 * 3000,
443 ..Default::default()
444 },
445 payload: vec![0u8; 100].into(),
446 ..Default::default()
447 }),
448 };
449 chain.handle_write(pkt).unwrap();
450 chain.poll_write();
452 }
453
454 chain.handle_timeout(base_time).unwrap();
456
457 while chain.poll_write().is_some() {}
459
460 let later_time = base_time + Duration::from_secs(2);
462 chain.handle_timeout(later_time).unwrap();
463
464 let report = chain.poll_write();
466 assert!(report.is_some());
467
468 if let Some(tagged) = report {
469 if let Packet::Rtcp(rtcp_packets) = tagged.message {
470 assert_eq!(rtcp_packets.len(), 1);
471 let sr = rtcp_packets[0]
472 .as_any()
473 .downcast_ref::<rtcp::sender_report::SenderReport>()
474 .expect("Expected SenderReport");
475 assert_eq!(sr.ssrc, 123456);
476 assert_eq!(sr.packet_count, 5);
477 assert_eq!(sr.octet_count, 500);
478 } else {
479 panic!("Expected RTCP packet");
480 }
481 }
482 }
483
484 #[test]
485 fn test_sender_report_multiple_streams() {
486 let mut chain = Registry::new()
488 .with(
489 SenderReportBuilder::default()
490 .with_interval(Duration::from_secs(1))
491 .build(),
492 )
493 .build();
494
495 let info1 = StreamInfo {
497 ssrc: 111111,
498 clock_rate: 90000,
499 ..Default::default()
500 };
501 let info2 = StreamInfo {
502 ssrc: 222222,
503 clock_rate: 48000,
504 ..Default::default()
505 };
506 chain.bind_local_stream(&info1);
507 chain.bind_local_stream(&info2);
508
509 let base_time = Instant::now();
510
511 for i in 0..3u16 {
513 let pkt = TaggedPacket {
514 now: base_time,
515 transport: Default::default(),
516 message: Packet::Rtp(rtp::Packet {
517 header: rtp::header::Header {
518 ssrc: 111111,
519 sequence_number: i,
520 timestamp: i as u32 * 3000,
521 ..Default::default()
522 },
523 payload: vec![0u8; 50].into(),
524 ..Default::default()
525 }),
526 };
527 chain.handle_write(pkt).unwrap();
528 chain.poll_write();
529 }
530
531 for i in 0..7u16 {
533 let pkt = TaggedPacket {
534 now: base_time,
535 transport: Default::default(),
536 message: Packet::Rtp(rtp::Packet {
537 header: rtp::header::Header {
538 ssrc: 222222,
539 sequence_number: i,
540 timestamp: i as u32 * 960,
541 ..Default::default()
542 },
543 payload: vec![0u8; 200].into(),
544 ..Default::default()
545 }),
546 };
547 chain.handle_write(pkt).unwrap();
548 chain.poll_write();
549 }
550
551 let later_time = base_time + Duration::from_secs(2);
553 chain.handle_timeout(later_time).unwrap();
554
555 let mut ssrcs = vec![];
557 let mut packet_counts = vec![];
558 let mut octet_counts = vec![];
559
560 while let Some(tagged) = chain.poll_write() {
561 if let Packet::Rtcp(rtcp_packets) = tagged.message {
562 for rtcp_pkt in rtcp_packets {
563 if let Some(sr) = rtcp_pkt
564 .as_any()
565 .downcast_ref::<rtcp::sender_report::SenderReport>()
566 {
567 ssrcs.push(sr.ssrc);
568 packet_counts.push(sr.packet_count);
569 octet_counts.push(sr.octet_count);
570 }
571 }
572 }
573 }
574
575 assert_eq!(ssrcs.len(), 2);
576 assert!(ssrcs.contains(&111111));
577 assert!(ssrcs.contains(&222222));
578
579 let idx1 = ssrcs.iter().position(|&s| s == 111111).unwrap();
581 assert_eq!(packet_counts[idx1], 3);
582 assert_eq!(octet_counts[idx1], 150);
583
584 let idx2 = ssrcs.iter().position(|&s| s == 222222).unwrap();
586 assert_eq!(packet_counts[idx2], 7);
587 assert_eq!(octet_counts[idx2], 1400);
588 }
589
590 #[test]
591 fn test_sender_report_unbind_stream() {
592 let mut chain = Registry::new()
594 .with(
595 SenderReportBuilder::default()
596 .with_interval(Duration::from_secs(1))
597 .build(),
598 )
599 .build();
600
601 let info = StreamInfo {
602 ssrc: 123456,
603 clock_rate: 90000,
604 ..Default::default()
605 };
606 chain.bind_local_stream(&info);
607
608 let base_time = Instant::now();
609
610 let pkt = TaggedPacket {
612 now: base_time,
613 transport: Default::default(),
614 message: Packet::Rtp(rtp::Packet {
615 header: rtp::header::Header {
616 ssrc: 123456,
617 sequence_number: 0,
618 timestamp: 0,
619 ..Default::default()
620 },
621 payload: vec![0u8; 100].into(),
622 ..Default::default()
623 }),
624 };
625 chain.handle_write(pkt).unwrap();
626 chain.poll_write();
627
628 chain.unbind_local_stream(&info);
630
631 let later_time = base_time + Duration::from_secs(2);
633 chain.handle_timeout(later_time).unwrap();
634
635 assert!(chain.poll_write().is_none());
637 }
638
639 #[test]
640 fn test_poll_timeout_returns_earliest() {
641 let interval = Duration::from_secs(5);
642 let mut chain = Registry::new()
643 .with(
644 SenderReportBuilder::default()
645 .with_interval(interval)
646 .build(),
647 )
648 .build();
649
650 assert_eq!(
654 chain.poll_timeout(),
655 None,
656 "an idle interceptor must not request a wake-up"
657 );
658
659 let info = StreamInfo {
660 ssrc: 123456,
661 clock_rate: 90000,
662 ..Default::default()
663 };
664 chain.bind_local_stream(&info);
665
666 let base_time = Instant::now();
667 chain
668 .handle_write(TaggedPacket {
669 now: base_time,
670 transport: Default::default(),
671 message: Packet::Rtp(rtp::Packet {
672 header: rtp::header::Header {
673 ssrc: 123456,
674 ..Default::default()
675 },
676 payload: vec![0u8; 100].into(),
677 ..Default::default()
678 }),
679 })
680 .unwrap();
681
682 assert_eq!(chain.poll_timeout(), Some(base_time + interval));
684 }
685}