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 eto: 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 eto: Instant::now(),
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
209 self.inner.handle_write(msg)
210 }
211
212 #[overrides]
213 fn poll_write(&mut self) -> Option<Self::Wout> {
214 if let Some(pkt) = self.write_queue.pop_front() {
216 return Some(pkt);
217 }
218 self.inner.poll_write()
219 }
220
221 #[overrides]
222 fn handle_timeout(&mut self, now: Self::Time) -> Result<(), Self::Error> {
223 if self.eto <= now {
224 self.eto = now + self.interval;
225
226 for stream in self.streams.values_mut() {
227 let rr = stream.generate_report(now);
228 self.write_queue.push_back(TaggedPacket {
229 now,
230 transport: TransportContext::default(),
231 message: Packet::Rtcp(vec![Box::new(rr)]),
232 });
233 }
234 }
235
236 self.inner.handle_timeout(now)
237 }
238
239 #[overrides]
240 fn poll_timeout(&mut self) -> Option<Self::Time> {
241 if let Some(eto) = self.inner.poll_timeout()
242 && eto < self.eto
243 {
244 Some(eto)
245 } else {
246 Some(self.eto)
247 }
248 }
249
250 #[overrides]
251 fn bind_local_stream(&mut self, info: &StreamInfo) {
252 let stream = SenderStream::new(info.ssrc, info.clock_rate, self.use_latest_packet);
253 self.streams.insert(info.ssrc, stream);
254
255 self.inner.bind_local_stream(info);
256 }
257
258 #[overrides]
259 fn unbind_local_stream(&mut self, info: &StreamInfo) {
260 self.streams.remove(&info.ssrc);
261
262 self.inner.unbind_local_stream(info);
263 }
264}
265
266#[cfg(test)]
267mod tests {
268 use super::*;
269 use crate::{NoopInterceptor, Registry};
270 use sansio::Protocol;
271
272 fn dummy_rtp_packet() -> TaggedPacket {
273 TaggedPacket {
274 now: Instant::now(),
275 transport: Default::default(),
276 message: crate::Packet::Rtp(rtp::Packet::default()),
277 }
278 }
279
280 #[test]
281 fn test_sender_report_builder_default() {
282 let chain = Registry::new()
284 .with(SenderReportBuilder::default().build())
285 .build();
286
287 assert_eq!(chain.interval, Duration::from_secs(1));
288 }
289
290 #[test]
291 fn test_sender_report_builder_with_custom_interval() {
292 let chain = Registry::new()
294 .with(
295 SenderReportBuilder::default()
296 .with_interval(Duration::from_millis(500))
297 .build(),
298 )
299 .build();
300
301 assert_eq!(chain.interval, Duration::from_millis(500));
302 }
303
304 #[test]
305 fn test_sender_report_chain_handle_read_write() {
306 let mut chain = Registry::new()
308 .with(SenderReportBuilder::default().build())
309 .build();
310
311 let pkt = dummy_rtp_packet();
313 chain.handle_read(pkt).unwrap();
314 assert!(chain.poll_read().is_some());
315
316 let pkt2 = dummy_rtp_packet();
318 let pkt2_message = pkt2.message.clone();
319 chain.handle_write(pkt2).unwrap();
320 assert_eq!(chain.poll_write().unwrap().message, pkt2_message);
321 }
322
323 #[test]
324 fn test_should_filter() {
325 assert!(SenderReportInterceptor::<NoopInterceptor>::should_filter(
327 PacketType::ReceiverReport
328 ));
329
330 assert!(SenderReportInterceptor::<NoopInterceptor>::should_filter(
332 PacketType::TransportSpecificFeedback
333 ));
334
335 assert!(!SenderReportInterceptor::<NoopInterceptor>::should_filter(
337 PacketType::SenderReport
338 ));
339
340 assert!(!SenderReportInterceptor::<NoopInterceptor>::should_filter(
342 PacketType::SourceDescription
343 ));
344
345 assert!(!SenderReportInterceptor::<NoopInterceptor>::should_filter(
347 PacketType::Goodbye
348 ));
349 }
350
351 #[test]
352 fn test_inner_access() {
353 let mut chain = Registry::new()
354 .with(SenderReportBuilder::default().build())
355 .build();
356
357 let _ = chain.inner();
359
360 let pkt = dummy_rtp_packet();
362 let pkt_message = pkt.message.clone();
363 chain.inner_mut().handle_write(pkt).unwrap();
364 assert_eq!(chain.inner_mut().poll_write().unwrap().message, pkt_message);
365 }
366
367 #[test]
368 fn test_use_latest_packet_option() {
369 let chain = Registry::new()
371 .with(
372 SenderReportBuilder::default()
373 .with_use_latest_packet()
374 .build(),
375 )
376 .build();
377
378 assert!(chain.use_latest_packet);
379
380 let chain_default = Registry::new()
382 .with(SenderReportBuilder::default().build())
383 .build();
384
385 assert!(!chain_default.use_latest_packet);
386 }
387
388 #[test]
389 fn test_use_latest_packet_combined_options() {
390 let chain = Registry::new()
392 .with(
393 SenderReportBuilder::default()
394 .with_interval(Duration::from_millis(250))
395 .with_use_latest_packet()
396 .build(),
397 )
398 .build();
399
400 assert_eq!(chain.interval, Duration::from_millis(250));
401 assert!(chain.use_latest_packet);
402 }
403
404 #[test]
405 fn test_sender_report_generation_on_timeout() {
406 let mut chain = Registry::new()
409 .with(
410 SenderReportBuilder::default()
411 .with_interval(Duration::from_secs(1))
412 .build(),
413 )
414 .build();
415
416 let info = StreamInfo {
418 ssrc: 123456,
419 clock_rate: 90000,
420 ..Default::default()
421 };
422 chain.bind_local_stream(&info);
423
424 let base_time = Instant::now();
425
426 for i in 0..5u16 {
428 let pkt = TaggedPacket {
429 now: base_time,
430 transport: Default::default(),
431 message: Packet::Rtp(rtp::Packet {
432 header: rtp::header::Header {
433 ssrc: 123456,
434 sequence_number: i,
435 timestamp: i as u32 * 3000,
436 ..Default::default()
437 },
438 payload: vec![0u8; 100].into(),
439 ..Default::default()
440 }),
441 };
442 chain.handle_write(pkt).unwrap();
443 chain.poll_write();
445 }
446
447 chain.handle_timeout(base_time).unwrap();
449
450 while chain.poll_write().is_some() {}
452
453 let later_time = base_time + Duration::from_secs(2);
455 chain.handle_timeout(later_time).unwrap();
456
457 let report = chain.poll_write();
459 assert!(report.is_some());
460
461 if let Some(tagged) = report {
462 if let Packet::Rtcp(rtcp_packets) = tagged.message {
463 assert_eq!(rtcp_packets.len(), 1);
464 let sr = rtcp_packets[0]
465 .as_any()
466 .downcast_ref::<rtcp::sender_report::SenderReport>()
467 .expect("Expected SenderReport");
468 assert_eq!(sr.ssrc, 123456);
469 assert_eq!(sr.packet_count, 5);
470 assert_eq!(sr.octet_count, 500);
471 } else {
472 panic!("Expected RTCP packet");
473 }
474 }
475 }
476
477 #[test]
478 fn test_sender_report_multiple_streams() {
479 let mut chain = Registry::new()
481 .with(
482 SenderReportBuilder::default()
483 .with_interval(Duration::from_secs(1))
484 .build(),
485 )
486 .build();
487
488 let info1 = StreamInfo {
490 ssrc: 111111,
491 clock_rate: 90000,
492 ..Default::default()
493 };
494 let info2 = StreamInfo {
495 ssrc: 222222,
496 clock_rate: 48000,
497 ..Default::default()
498 };
499 chain.bind_local_stream(&info1);
500 chain.bind_local_stream(&info2);
501
502 let base_time = Instant::now();
503
504 for i in 0..3u16 {
506 let pkt = TaggedPacket {
507 now: base_time,
508 transport: Default::default(),
509 message: Packet::Rtp(rtp::Packet {
510 header: rtp::header::Header {
511 ssrc: 111111,
512 sequence_number: i,
513 timestamp: i as u32 * 3000,
514 ..Default::default()
515 },
516 payload: vec![0u8; 50].into(),
517 ..Default::default()
518 }),
519 };
520 chain.handle_write(pkt).unwrap();
521 chain.poll_write();
522 }
523
524 for i in 0..7u16 {
526 let pkt = TaggedPacket {
527 now: base_time,
528 transport: Default::default(),
529 message: Packet::Rtp(rtp::Packet {
530 header: rtp::header::Header {
531 ssrc: 222222,
532 sequence_number: i,
533 timestamp: i as u32 * 960,
534 ..Default::default()
535 },
536 payload: vec![0u8; 200].into(),
537 ..Default::default()
538 }),
539 };
540 chain.handle_write(pkt).unwrap();
541 chain.poll_write();
542 }
543
544 let later_time = base_time + Duration::from_secs(2);
546 chain.handle_timeout(later_time).unwrap();
547
548 let mut ssrcs = vec![];
550 let mut packet_counts = vec![];
551 let mut octet_counts = vec![];
552
553 while let Some(tagged) = chain.poll_write() {
554 if let Packet::Rtcp(rtcp_packets) = tagged.message {
555 for rtcp_pkt in rtcp_packets {
556 if let Some(sr) = rtcp_pkt
557 .as_any()
558 .downcast_ref::<rtcp::sender_report::SenderReport>()
559 {
560 ssrcs.push(sr.ssrc);
561 packet_counts.push(sr.packet_count);
562 octet_counts.push(sr.octet_count);
563 }
564 }
565 }
566 }
567
568 assert_eq!(ssrcs.len(), 2);
569 assert!(ssrcs.contains(&111111));
570 assert!(ssrcs.contains(&222222));
571
572 let idx1 = ssrcs.iter().position(|&s| s == 111111).unwrap();
574 assert_eq!(packet_counts[idx1], 3);
575 assert_eq!(octet_counts[idx1], 150);
576
577 let idx2 = ssrcs.iter().position(|&s| s == 222222).unwrap();
579 assert_eq!(packet_counts[idx2], 7);
580 assert_eq!(octet_counts[idx2], 1400);
581 }
582
583 #[test]
584 fn test_sender_report_unbind_stream() {
585 let mut chain = Registry::new()
587 .with(
588 SenderReportBuilder::default()
589 .with_interval(Duration::from_secs(1))
590 .build(),
591 )
592 .build();
593
594 let info = StreamInfo {
595 ssrc: 123456,
596 clock_rate: 90000,
597 ..Default::default()
598 };
599 chain.bind_local_stream(&info);
600
601 let base_time = Instant::now();
602
603 let pkt = TaggedPacket {
605 now: base_time,
606 transport: Default::default(),
607 message: Packet::Rtp(rtp::Packet {
608 header: rtp::header::Header {
609 ssrc: 123456,
610 sequence_number: 0,
611 timestamp: 0,
612 ..Default::default()
613 },
614 payload: vec![0u8; 100].into(),
615 ..Default::default()
616 }),
617 };
618 chain.handle_write(pkt).unwrap();
619 chain.poll_write();
620
621 chain.unbind_local_stream(&info);
623
624 let later_time = base_time + Duration::from_secs(2);
626 chain.handle_timeout(later_time).unwrap();
627
628 assert!(chain.poll_write().is_none());
630 }
631
632 #[test]
633 fn test_poll_timeout_returns_earliest() {
634 let mut chain = Registry::new()
636 .with(
637 SenderReportBuilder::default()
638 .with_interval(Duration::from_secs(5))
639 .build(),
640 )
641 .build();
642
643 let timeout = chain.poll_timeout();
645 assert!(timeout.is_some());
646 }
647}