1use crate::report::receiver_stream::ReceiverStream;
4use crate::stream_info::StreamInfo;
5use crate::{Interceptor, Packet, TaggedPacket, interceptor};
6use shared::TransportContext;
7use shared::error::Error;
8use std::collections::{HashMap, VecDeque};
9use std::marker::PhantomData;
10use std::time::{Duration, Instant};
11
12pub struct ReceiverReportBuilder<P> {
31 interval: Duration,
33 _phantom: PhantomData<P>,
34}
35
36impl<P> Default for ReceiverReportBuilder<P> {
37 fn default() -> Self {
38 Self {
39 interval: Duration::from_secs(1),
40 _phantom: PhantomData,
41 }
42 }
43}
44
45impl<P> ReceiverReportBuilder<P> {
46 pub fn new() -> Self {
50 Self::default()
51 }
52
53 pub fn with_interval(mut self, interval: Duration) -> Self {
69 self.interval = interval;
70 self
71 }
72
73 pub fn build(self) -> impl FnOnce(P) -> ReceiverReportInterceptor<P> {
86 move |inner| ReceiverReportInterceptor::new(inner, self.interval)
87 }
88}
89
90#[derive(Interceptor)]
109pub struct ReceiverReportInterceptor<P> {
110 #[next]
111 inner: P,
112
113 interval: Duration,
114 eto: Instant,
115
116 streams: HashMap<u32, ReceiverStream>,
117
118 read_queue: VecDeque<TaggedPacket>,
119 write_queue: VecDeque<TaggedPacket>,
120}
121
122impl<P> ReceiverReportInterceptor<P> {
123 fn new(inner: P, interval: Duration) -> Self {
125 Self {
126 inner,
127
128 interval,
129 eto: Instant::now(),
130
131 streams: HashMap::new(),
132
133 read_queue: VecDeque::new(),
134 write_queue: VecDeque::new(),
135 }
136 }
137
138 fn process_rtp(&mut self, now: Instant, ssrc: u32, seq: u16, timestamp: u32) {
140 let stream = self.streams.entry(ssrc).or_insert_with(|| {
142 ReceiverStream::new(ssrc, 90000)
144 });
145
146 let pkt = rtp::packet::Packet {
148 header: rtp::header::Header {
149 ssrc,
150 sequence_number: seq,
151 timestamp,
152 ..Default::default()
153 },
154 ..Default::default()
155 };
156
157 stream.process_rtp(now, &pkt);
158 }
159
160 fn process_sender_report(&mut self, now: Instant, sr: &rtcp::sender_report::SenderReport) {
162 if let Some(stream) = self.streams.get_mut(&sr.ssrc) {
163 stream.process_sender_report(now, sr);
164 }
165 }
166
167 fn generate_reports(&mut self, now: Instant) -> Vec<rtcp::receiver_report::ReceiverReport> {
169 self.streams
170 .values_mut()
171 .map(|stream| stream.generate_report(now))
172 .collect()
173 }
174
175 fn register_stream(&mut self, ssrc: u32, clock_rate: u32) {
177 self.streams
178 .entry(ssrc)
179 .or_insert_with(|| ReceiverStream::new(ssrc, clock_rate));
180 }
181}
182
183#[interceptor]
184impl<P: Interceptor> ReceiverReportInterceptor<P> {
185 #[overrides]
186 fn handle_read(&mut self, msg: TaggedPacket) -> Result<(), Self::Error> {
187 if let Packet::Rtcp(rtcp_packets) = &msg.message {
188 for rtcp_packet in rtcp_packets {
189 if let Some(sr) = rtcp_packet
190 .as_any()
191 .downcast_ref::<rtcp::sender_report::SenderReport>()
192 && let Some(stream) = self.streams.get_mut(&sr.ssrc)
193 {
194 stream.process_sender_report(msg.now, sr);
195 }
196 }
197 } else if let Packet::Rtp(rtp_packet) = &msg.message
198 && let Some(stream) = self.streams.get_mut(&rtp_packet.header.ssrc)
199 {
200 stream.process_rtp(msg.now, rtp_packet);
201 }
202
203 self.inner.handle_read(msg)
204 }
205
206 #[overrides]
207 fn poll_write(&mut self) -> Option<Self::Wout> {
208 if let Some(pkt) = self.write_queue.pop_front() {
210 return Some(pkt);
211 }
212 self.inner.poll_write()
213 }
214
215 #[overrides]
216 fn handle_timeout(&mut self, now: Self::Time) -> Result<(), Self::Error> {
217 if self.eto <= now {
218 self.eto = now + self.interval;
219
220 for stream in self.streams.values_mut() {
221 let rr = stream.generate_report(now);
222 self.write_queue.push_back(TaggedPacket {
223 now,
224 transport: TransportContext::default(),
225 message: Packet::Rtcp(vec![Box::new(rr)]),
226 });
227 }
228 }
229
230 self.inner.handle_timeout(now)
231 }
232
233 #[overrides]
234 fn poll_timeout(&mut self) -> Option<Self::Time> {
235 if let Some(eto) = self.inner.poll_timeout()
236 && eto < self.eto
237 {
238 Some(eto)
239 } else {
240 Some(self.eto)
241 }
242 }
243
244 #[overrides]
245 fn bind_remote_stream(&mut self, info: &StreamInfo) {
246 let stream = ReceiverStream::new(info.ssrc, info.clock_rate);
247 self.streams.insert(info.ssrc, stream);
248
249 self.inner.bind_remote_stream(info);
250 }
251
252 #[overrides]
253 fn unbind_remote_stream(&mut self, info: &StreamInfo) {
254 self.streams.remove(&info.ssrc);
255
256 self.inner.unbind_remote_stream(info);
257 }
258}
259
260#[cfg(test)]
261mod tests {
262 use super::*;
263 use crate::Registry;
264 use sansio::Protocol;
265
266 fn dummy_rtp_packet() -> TaggedPacket {
267 TaggedPacket {
268 now: Instant::now(),
269 transport: Default::default(),
270 message: crate::Packet::Rtp(rtp::Packet::default()),
271 }
272 }
273
274 #[test]
275 fn test_receiver_report_builder_default() {
276 let chain = Registry::new()
278 .with(ReceiverReportBuilder::default().build())
279 .build();
280
281 assert_eq!(chain.interval, Duration::from_secs(1));
282 assert!(chain.streams.is_empty());
283 }
284
285 #[test]
286 fn test_receiver_report_builder_with_custom_interval() {
287 let chain = Registry::new()
289 .with(
290 ReceiverReportBuilder::default()
291 .with_interval(Duration::from_millis(500))
292 .build(),
293 )
294 .build();
295
296 assert_eq!(chain.interval, Duration::from_millis(500));
297 }
298
299 #[test]
300 fn test_receiver_report_chain_handle_read_write() {
301 let mut chain = Registry::new()
303 .with(ReceiverReportBuilder::default().build())
304 .build();
305
306 let pkt = dummy_rtp_packet();
308 chain.handle_read(pkt).unwrap();
309 assert!(chain.poll_read().is_some());
310
311 let pkt2 = dummy_rtp_packet();
313 let pkt2_message = pkt2.message.clone();
314 chain.handle_write(pkt2).unwrap();
315 assert_eq!(chain.poll_write().unwrap().message, pkt2_message);
316 }
317
318 #[test]
319 fn test_register_stream() {
320 let mut chain = Registry::new()
321 .with(ReceiverReportBuilder::default().build())
322 .build();
323
324 chain.register_stream(12345, 48000);
325 assert!(chain.streams.contains_key(&12345));
326 }
327
328 #[test]
329 fn test_process_rtp() {
330 let mut chain = Registry::new()
331 .with(ReceiverReportBuilder::default().build())
332 .build();
333
334 let now = Instant::now();
335 chain.process_rtp(now, 12345, 1, 1000);
336
337 assert!(chain.streams.contains_key(&12345));
338 }
339
340 #[test]
341 fn test_generate_reports() {
342 let mut chain = Registry::new()
343 .with(ReceiverReportBuilder::default().build())
344 .build();
345
346 let now = Instant::now();
347 chain.process_rtp(now, 12345, 1, 1000);
348 chain.process_rtp(now, 12345, 2, 2000);
349
350 let reports = chain.generate_reports(now);
351 assert_eq!(reports.len(), 1);
352 }
353
354 #[test]
355 fn test_chained_interceptors() {
356 use crate::report::sender::SenderReportBuilder;
357
358 let mut chain = Registry::new()
360 .with(ReceiverReportBuilder::default().build())
361 .with(
362 SenderReportBuilder::default()
363 .with_interval(Duration::from_millis(250))
364 .build(),
365 )
366 .build();
367
368 let pkt = dummy_rtp_packet();
370 chain.handle_read(pkt).unwrap();
371 assert!(chain.poll_read().is_some());
372
373 let pkt2 = dummy_rtp_packet();
374 let pkt2_message = pkt2.message.clone();
375 chain.handle_write(pkt2).unwrap();
376 assert_eq!(chain.poll_write().unwrap().message, pkt2_message);
377 }
378
379 #[test]
380 fn test_receiver_report_generation_on_timeout() {
381 let mut chain = Registry::new()
384 .with(
385 ReceiverReportBuilder::default()
386 .with_interval(Duration::from_secs(1))
387 .build(),
388 )
389 .build();
390
391 let info = StreamInfo {
393 ssrc: 123456,
394 clock_rate: 90000,
395 ..Default::default()
396 };
397 chain.bind_remote_stream(&info);
398
399 let base_time = Instant::now();
400
401 for i in 0..10u16 {
403 let pkt = TaggedPacket {
404 now: base_time,
405 transport: Default::default(),
406 message: Packet::Rtp(rtp::Packet {
407 header: rtp::header::Header {
408 ssrc: 123456,
409 sequence_number: i,
410 timestamp: i as u32 * 3000,
411 ..Default::default()
412 },
413 ..Default::default()
414 }),
415 };
416 chain.handle_read(pkt).unwrap();
417 chain.poll_read();
418 }
419
420 chain.handle_timeout(base_time).unwrap();
422
423 while chain.poll_write().is_some() {}
425
426 let later_time = base_time + Duration::from_secs(2);
428 chain.handle_timeout(later_time).unwrap();
429
430 let report = chain.poll_write();
432 assert!(report.is_some());
433
434 if let Some(tagged) = report {
435 if let Packet::Rtcp(rtcp_packets) = tagged.message {
436 assert_eq!(rtcp_packets.len(), 1);
437 let rr = rtcp_packets[0]
438 .as_any()
439 .downcast_ref::<rtcp::receiver_report::ReceiverReport>()
440 .expect("Expected ReceiverReport");
441 assert_eq!(rr.reports.len(), 1);
442 assert_eq!(rr.reports[0].ssrc, 123456);
443 assert_eq!(rr.reports[0].last_sequence_number, 9);
444 assert_eq!(rr.reports[0].fraction_lost, 0);
445 assert_eq!(rr.reports[0].total_lost, 0);
446 } else {
447 panic!("Expected RTCP packet");
448 }
449 }
450 }
451
452 #[test]
453 fn test_receiver_report_with_packet_loss() {
454 let mut chain = Registry::new()
456 .with(
457 ReceiverReportBuilder::default()
458 .with_interval(Duration::from_secs(1))
459 .build(),
460 )
461 .build();
462
463 let info = StreamInfo {
464 ssrc: 123456,
465 clock_rate: 90000,
466 ..Default::default()
467 };
468 chain.bind_remote_stream(&info);
469
470 let base_time = Instant::now();
471
472 let pkt = TaggedPacket {
474 now: base_time,
475 transport: Default::default(),
476 message: Packet::Rtp(rtp::Packet {
477 header: rtp::header::Header {
478 ssrc: 123456,
479 sequence_number: 1,
480 timestamp: 3000,
481 ..Default::default()
482 },
483 ..Default::default()
484 }),
485 };
486 chain.handle_read(pkt).unwrap();
487 chain.poll_read();
488
489 let pkt = TaggedPacket {
491 now: base_time,
492 transport: Default::default(),
493 message: Packet::Rtp(rtp::Packet {
494 header: rtp::header::Header {
495 ssrc: 123456,
496 sequence_number: 3,
497 timestamp: 9000,
498 ..Default::default()
499 },
500 ..Default::default()
501 }),
502 };
503 chain.handle_read(pkt).unwrap();
504 chain.poll_read();
505
506 let later_time = base_time + Duration::from_secs(2);
508 chain.handle_timeout(later_time).unwrap();
509
510 let report = chain.poll_write();
511 assert!(report.is_some());
512
513 if let Some(tagged) = report {
514 if let Packet::Rtcp(rtcp_packets) = tagged.message {
515 let rr = rtcp_packets[0]
516 .as_any()
517 .downcast_ref::<rtcp::receiver_report::ReceiverReport>()
518 .expect("Expected ReceiverReport");
519 assert_eq!(rr.reports[0].last_sequence_number, 3);
520 assert_eq!(rr.reports[0].total_lost, 1);
522 assert_eq!(rr.reports[0].fraction_lost, (256u32 * 1 / 3) as u8);
524 } else {
525 panic!("Expected RTCP packet");
526 }
527 }
528 }
529
530 #[test]
531 fn test_receiver_report_with_sender_report() {
532 let mut chain = Registry::new()
534 .with(
535 ReceiverReportBuilder::default()
536 .with_interval(Duration::from_secs(1))
537 .build(),
538 )
539 .build();
540
541 let info = StreamInfo {
542 ssrc: 123456,
543 clock_rate: 90000,
544 ..Default::default()
545 };
546 chain.bind_remote_stream(&info);
547
548 let base_time = Instant::now();
549
550 let pkt = TaggedPacket {
552 now: base_time,
553 transport: Default::default(),
554 message: Packet::Rtp(rtp::Packet {
555 header: rtp::header::Header {
556 ssrc: 123456,
557 sequence_number: 1,
558 timestamp: 3000,
559 ..Default::default()
560 },
561 ..Default::default()
562 }),
563 };
564 chain.handle_read(pkt).unwrap();
565 chain.poll_read();
566
567 let sr = rtcp::sender_report::SenderReport {
569 ssrc: 123456,
570 ntp_time: 0x1234_5678_0000_0000,
571 rtp_time: 3000,
572 packet_count: 100,
573 octet_count: 10000,
574 ..Default::default()
575 };
576 let sr_pkt = TaggedPacket {
577 now: base_time,
578 transport: Default::default(),
579 message: Packet::Rtcp(vec![Box::new(sr)]),
580 };
581 chain.handle_read(sr_pkt).unwrap();
582
583 let later_time = base_time + Duration::from_secs(1);
585 chain.handle_timeout(later_time).unwrap();
586
587 let report = chain.poll_write();
588 assert!(report.is_some());
589
590 if let Some(tagged) = report {
591 if let Packet::Rtcp(rtcp_packets) = tagged.message {
592 let rr = rtcp_packets[0]
593 .as_any()
594 .downcast_ref::<rtcp::receiver_report::ReceiverReport>()
595 .expect("Expected ReceiverReport");
596 assert_eq!(rr.reports[0].delay, 65536);
598 assert_eq!(rr.reports[0].last_sender_report, 0x5678_0000);
600 } else {
601 panic!("Expected RTCP packet");
602 }
603 }
604 }
605
606 #[test]
607 fn test_receiver_report_multiple_streams() {
608 let mut chain = Registry::new()
610 .with(
611 ReceiverReportBuilder::default()
612 .with_interval(Duration::from_secs(1))
613 .build(),
614 )
615 .build();
616
617 let info1 = StreamInfo {
618 ssrc: 111111,
619 clock_rate: 90000,
620 ..Default::default()
621 };
622 let info2 = StreamInfo {
623 ssrc: 222222,
624 clock_rate: 48000,
625 ..Default::default()
626 };
627 chain.bind_remote_stream(&info1);
628 chain.bind_remote_stream(&info2);
629
630 let base_time = Instant::now();
631
632 for i in 0..5u16 {
634 let pkt = TaggedPacket {
635 now: base_time,
636 transport: Default::default(),
637 message: Packet::Rtp(rtp::Packet {
638 header: rtp::header::Header {
639 ssrc: 111111,
640 sequence_number: i,
641 timestamp: i as u32 * 3000,
642 ..Default::default()
643 },
644 ..Default::default()
645 }),
646 };
647 chain.handle_read(pkt).unwrap();
648 chain.poll_read();
649 }
650
651 let pkt = TaggedPacket {
653 now: base_time,
654 transport: Default::default(),
655 message: Packet::Rtp(rtp::Packet {
656 header: rtp::header::Header {
657 ssrc: 222222,
658 sequence_number: 0,
659 timestamp: 0,
660 ..Default::default()
661 },
662 ..Default::default()
663 }),
664 };
665 chain.handle_read(pkt).unwrap();
666 chain.poll_read();
667
668 let pkt = TaggedPacket {
669 now: base_time,
670 transport: Default::default(),
671 message: Packet::Rtp(rtp::Packet {
672 header: rtp::header::Header {
673 ssrc: 222222,
674 sequence_number: 5, timestamp: 5 * 960,
676 ..Default::default()
677 },
678 ..Default::default()
679 }),
680 };
681 chain.handle_read(pkt).unwrap();
682 chain.poll_read();
683
684 let later_time = base_time + Duration::from_secs(2);
686 chain.handle_timeout(later_time).unwrap();
687
688 let mut ssrcs = vec![];
690 let mut total_lost = vec![];
691
692 while let Some(tagged) = chain.poll_write() {
693 if let Packet::Rtcp(rtcp_packets) = tagged.message {
694 for rtcp_pkt in rtcp_packets {
695 if let Some(rr) = rtcp_pkt
696 .as_any()
697 .downcast_ref::<rtcp::receiver_report::ReceiverReport>()
698 {
699 for report in &rr.reports {
700 ssrcs.push(report.ssrc);
701 total_lost.push(report.total_lost);
702 }
703 }
704 }
705 }
706 }
707
708 assert_eq!(ssrcs.len(), 2);
709 assert!(ssrcs.contains(&111111));
710 assert!(ssrcs.contains(&222222));
711
712 let idx1 = ssrcs.iter().position(|&s| s == 111111).unwrap();
714 assert_eq!(total_lost[idx1], 0);
715
716 let idx2 = ssrcs.iter().position(|&s| s == 222222).unwrap();
718 assert_eq!(total_lost[idx2], 4);
719 }
720
721 #[test]
722 fn test_receiver_report_unbind_stream() {
723 let mut chain = Registry::new()
725 .with(
726 ReceiverReportBuilder::default()
727 .with_interval(Duration::from_secs(1))
728 .build(),
729 )
730 .build();
731
732 let info = StreamInfo {
733 ssrc: 123456,
734 clock_rate: 90000,
735 ..Default::default()
736 };
737 chain.bind_remote_stream(&info);
738
739 let base_time = Instant::now();
740
741 let pkt = TaggedPacket {
743 now: base_time,
744 transport: Default::default(),
745 message: Packet::Rtp(rtp::Packet {
746 header: rtp::header::Header {
747 ssrc: 123456,
748 sequence_number: 0,
749 timestamp: 0,
750 ..Default::default()
751 },
752 ..Default::default()
753 }),
754 };
755 chain.handle_read(pkt).unwrap();
756 chain.poll_read();
757
758 chain.unbind_remote_stream(&info);
760
761 let later_time = base_time + Duration::from_secs(2);
763 chain.handle_timeout(later_time).unwrap();
764
765 assert!(chain.poll_write().is_none());
767 }
768
769 #[test]
770 fn test_receiver_report_sequence_wrap() {
771 let mut chain = Registry::new()
773 .with(
774 ReceiverReportBuilder::default()
775 .with_interval(Duration::from_secs(1))
776 .build(),
777 )
778 .build();
779
780 let info = StreamInfo {
781 ssrc: 123456,
782 clock_rate: 90000,
783 ..Default::default()
784 };
785 chain.bind_remote_stream(&info);
786
787 let base_time = Instant::now();
788
789 let pkt = TaggedPacket {
791 now: base_time,
792 transport: Default::default(),
793 message: Packet::Rtp(rtp::Packet {
794 header: rtp::header::Header {
795 ssrc: 123456,
796 sequence_number: 0xffff,
797 timestamp: 0,
798 ..Default::default()
799 },
800 ..Default::default()
801 }),
802 };
803 chain.handle_read(pkt).unwrap();
804 chain.poll_read();
805
806 let pkt = TaggedPacket {
808 now: base_time,
809 transport: Default::default(),
810 message: Packet::Rtp(rtp::Packet {
811 header: rtp::header::Header {
812 ssrc: 123456,
813 sequence_number: 0x00,
814 timestamp: 3000,
815 ..Default::default()
816 },
817 ..Default::default()
818 }),
819 };
820 chain.handle_read(pkt).unwrap();
821 chain.poll_read();
822
823 let later_time = base_time + Duration::from_secs(2);
825 chain.handle_timeout(later_time).unwrap();
826
827 let report = chain.poll_write();
828 assert!(report.is_some());
829
830 if let Some(tagged) = report {
831 if let Packet::Rtcp(rtcp_packets) = tagged.message {
832 let rr = rtcp_packets[0]
833 .as_any()
834 .downcast_ref::<rtcp::receiver_report::ReceiverReport>()
835 .expect("Expected ReceiverReport");
836 assert_eq!(rr.reports[0].last_sequence_number, 1 << 16);
838 assert_eq!(rr.reports[0].fraction_lost, 0);
839 assert_eq!(rr.reports[0].total_lost, 0);
840 } else {
841 panic!("Expected RTCP packet");
842 }
843 }
844 }
845}