use super::recorder::CcFeedbackRecorder;
use crate::stream_info::StreamInfo;
use crate::{Interceptor, Packet, TaggedPacket, interceptor};
use rtcp::transport_feedbacks::cc_feedback_report::Ecn;
use shared::TransportContext;
use shared::error::Error;
use shared::time::SystemInstant;
use std::collections::HashSet;
use std::collections::VecDeque;
use std::marker::PhantomData;
use std::time::{Duration, Instant};
pub const DEFAULT_INTERVAL: Duration = Duration::from_millis(100);
pub const DEFAULT_MAX_REPORT_SIZE: usize = 1200;
pub struct Rfc8888Builder<P> {
interval: Duration,
max_report_size: usize,
sender_ssrc: u32,
_phantom: PhantomData<P>,
}
impl<P> Default for Rfc8888Builder<P> {
fn default() -> Self {
Self {
interval: DEFAULT_INTERVAL,
max_report_size: DEFAULT_MAX_REPORT_SIZE,
sender_ssrc: 0,
_phantom: PhantomData,
}
}
}
impl<P> Rfc8888Builder<P> {
pub fn new() -> Self {
Self::default()
}
pub fn with_interval(mut self, interval: Duration) -> Self {
self.interval = interval;
self
}
pub fn with_max_report_size(mut self, max_report_size: usize) -> Self {
self.max_report_size = max_report_size;
self
}
pub fn with_sender_ssrc(mut self, sender_ssrc: u32) -> Self {
self.sender_ssrc = sender_ssrc;
self
}
pub fn build(self) -> impl FnOnce(P) -> Rfc8888Interceptor<P> {
move |inner| {
Rfc8888Interceptor::new(inner, self.interval, self.max_report_size, self.sender_ssrc)
}
}
}
#[derive(Interceptor)]
pub struct Rfc8888Interceptor<P> {
#[next]
inner: P,
interval: Duration,
max_report_size: usize,
sender_ssrc: u32,
recorder: CcFeedbackRecorder,
streams: HashSet<u32>,
next_timeout: Option<Instant>,
epoch: Option<SystemInstant>,
write_queue: VecDeque<TaggedPacket>,
}
impl<P> Rfc8888Interceptor<P> {
fn new(inner: P, interval: Duration, max_report_size: usize, sender_ssrc: u32) -> Self {
Self {
inner,
interval,
max_report_size,
sender_ssrc,
recorder: CcFeedbackRecorder::new(),
streams: HashSet::new(),
next_timeout: None,
epoch: None,
write_queue: VecDeque::new(),
}
}
fn report_timestamp(&mut self, now: Instant) -> u32 {
let epoch = self.epoch.get_or_insert_with(|| SystemInstant::now(now));
(epoch.ntp(now) >> 16) as u32
}
fn arm(&mut self, now: Instant) {
if self.next_timeout.is_none() && !self.streams.is_empty() && !self.interval.is_zero() {
self.next_timeout = Some(now + self.interval);
}
}
}
#[interceptor]
impl<P: Interceptor> Rfc8888Interceptor<P> {
#[overrides]
fn bind_remote_stream(&mut self, info: &StreamInfo) {
self.streams.insert(info.ssrc);
self.inner.bind_remote_stream(info);
}
#[overrides]
fn unbind_remote_stream(&mut self, info: &StreamInfo) {
self.streams.remove(&info.ssrc);
self.recorder.remove_stream(info.ssrc);
if self.streams.is_empty() {
self.next_timeout = None;
}
self.inner.unbind_remote_stream(info);
}
#[overrides]
fn handle_read(&mut self, msg: TaggedPacket) -> Result<(), Self::Error> {
if let Packet::Rtp(rtp_packet) = &msg.message
&& self.streams.contains(&rtp_packet.header.ssrc)
{
self.recorder.add_packet(
msg.now,
rtp_packet.header.ssrc,
rtp_packet.header.sequence_number,
Ecn::NotEct,
);
self.arm(msg.now);
}
self.inner.handle_read(msg)
}
#[overrides]
fn handle_timeout(&mut self, now: Self::Time) -> Result<(), Self::Error> {
self.arm(now);
if let Some(next_timeout) = self.next_timeout
&& now >= next_timeout
{
self.next_timeout = Some(now + self.interval);
if !self.recorder.is_empty() {
let report_timestamp = self.report_timestamp(now);
let report = self.recorder.build_report(
now,
self.sender_ssrc,
report_timestamp,
self.max_report_size,
);
if !report.report_blocks.is_empty() {
self.write_queue.push_back(TaggedPacket {
now,
transport: TransportContext::default(),
message: Packet::Rtcp(vec![Box::new(report)]),
});
}
}
}
self.inner.handle_timeout(now)
}
#[overrides]
fn poll_timeout(&mut self) -> Option<Self::Time> {
match (self.next_timeout, self.inner.poll_timeout()) {
(Some(mine), Some(theirs)) => Some(mine.min(theirs)),
(mine, theirs) => mine.or(theirs),
}
}
#[overrides]
fn poll_write(&mut self) -> Option<Self::Wout> {
if let Some(packet) = self.write_queue.pop_front() {
return Some(packet);
}
self.inner.poll_write()
}
}