use super::estimator::BandwidthEstimator;
use crate::Interceptor;
use crate::rtpfb::convert::{convert_ccfb, convert_twcc};
use crate::rtpfb::history::History;
use crate::stream_info::StreamInfo;
use crate::twcc::stream_supports_twcc;
use crate::{Attribute, Packet, TaggedPacket};
use sansio::Protocol;
use shared::error::Error;
use shared::marshal::{MarshalSize, Unmarshal};
use std::collections::{HashMap, VecDeque};
use std::time::{Duration, Instant};
pub const DEFAULT_PRUNE_HORIZON: Duration = Duration::from_secs(2);
struct LocalStream {
hdr_ext_id: u8,
}
pub struct CongestionControlBuilder<E: BandwidthEstimator> {
estimator: E,
prune_horizon: Duration,
}
impl<E: BandwidthEstimator> CongestionControlBuilder<E> {
pub fn new(estimator: E) -> Self {
Self {
estimator,
prune_horizon: DEFAULT_PRUNE_HORIZON,
}
}
pub fn with_prune_horizon(mut self, prune_horizon: Duration) -> Self {
self.prune_horizon = prune_horizon;
self
}
pub fn build(self) -> CongestionControlInterceptor<E> {
CongestionControlInterceptor {
last_target: self.estimator.target_bitrate(),
estimator: self.estimator,
prune_horizon: self.prune_horizon,
history: History::new(),
streams: HashMap::new(),
read_queue: VecDeque::new(),
write_queue: VecDeque::new(),
}
}
}
pub struct CongestionControlInterceptor<E: BandwidthEstimator> {
estimator: E,
history: History,
streams: HashMap<u32, LocalStream>,
prune_horizon: Duration,
last_target: f64,
read_queue: VecDeque<TaggedPacket>,
write_queue: VecDeque<TaggedPacket>,
}
impl<E: BandwidthEstimator> CongestionControlInterceptor<E> {
pub fn estimator(&self) -> &E {
&self.estimator
}
pub fn outstanding(&self) -> usize {
self.history.len()
}
fn twcc_sequence_number(&self, rtp_packet: &rtp::Packet) -> Option<u16> {
let stream = self.streams.get(&rtp_packet.header.ssrc)?;
let mut extension = rtp_packet.header.get_extension(stream.hdr_ext_id)?;
rtp::extension::transport_cc_extension::TransportCcExtension::unmarshal(&mut extension)
.ok()
.map(|extension| extension.transport_sequence)
}
#[allow(clippy::borrowed_box)]
fn ingest(&mut self, now: Instant, rtcp_packet: &Box<dyn rtcp::Packet>) -> bool {
let payload = rtcp_packet.as_any();
if let Some(feedback) = payload
.downcast_ref::<rtcp::transport_feedbacks::transport_layer_cc::TransportLayerCc>(
) {
for acknowledgement in convert_twcc(feedback) {
self.history.on_twcc_feedback(now, acknowledgement);
}
return true;
}
if let Some(feedback) = payload
.downcast_ref::<rtcp::transport_feedbacks::cc_feedback_report::CcFeedbackReport>(
) {
let (_report_delay, per_stream) = convert_ccfb(feedback);
for (ssrc, acknowledgements) in per_stream {
for acknowledgement in acknowledgements {
self.history.on_ccfb_feedback(now, ssrc, acknowledgement);
}
}
return true;
}
false
}
}
impl<E: BandwidthEstimator> Protocol<TaggedPacket, TaggedPacket, ()>
for CongestionControlInterceptor<E>
{
type Rout = TaggedPacket;
type Wout = TaggedPacket;
type Eout = ();
type Error = Error;
type Time = Instant;
fn handle_read(&mut self, mut msg: TaggedPacket) -> Result<(), Self::Error> {
let mut reported = false;
if let Packet::Rtcp(ref rtcp_packets) = msg.message.packet {
let feedback: Vec<_> = rtcp_packets.to_vec();
for rtcp_packet in &feedback {
reported |= self.ingest(msg.now, rtcp_packet);
}
}
if reported {
let reports = self.history.take_reports();
self.estimator.on_reports(msg.now, &reports);
let target = self.estimator.target_bitrate();
if target != self.last_target {
self.last_target = target;
msg.message.add(Attribute::TargetBitrateChanged {
bits_per_second: target,
});
}
}
self.read_queue.push_back(msg);
Ok(())
}
fn poll_read(&mut self) -> Option<Self::Rout> {
self.read_queue.pop_front()
}
fn handle_write(&mut self, msg: TaggedPacket) -> Result<(), Self::Error> {
if let Packet::Rtp(ref rtp_packet) = msg.message.packet {
let twcc_sequence_number = self.twcc_sequence_number(rtp_packet);
if self.streams.contains_key(&rtp_packet.header.ssrc) {
self.history.add_outgoing(
rtp_packet.header.ssrc,
rtp_packet.header.sequence_number,
twcc_sequence_number.is_some(),
twcc_sequence_number.unwrap_or_default(),
rtp_packet.marshal_size(),
msg.now,
);
}
}
self.write_queue.push_back(msg);
Ok(())
}
fn poll_write(&mut self) -> Option<Self::Wout> {
self.write_queue.pop_front()
}
fn handle_timeout(&mut self, now: Instant) -> Result<(), Self::Error> {
self.history
.prune_before(now.checked_sub(self.prune_horizon).unwrap_or(now));
self.estimator.handle_timeout(now);
Ok(())
}
fn poll_timeout(&mut self) -> Option<Self::Time> {
self.estimator.poll_timeout()
}
}
impl<E: BandwidthEstimator> Interceptor for CongestionControlInterceptor<E> {
fn bind_local_stream(&mut self, info: &StreamInfo) {
let hdr_ext_id = stream_supports_twcc(info).unwrap_or_default();
self.streams.insert(info.ssrc, LocalStream { hdr_ext_id });
}
fn unbind_local_stream(&mut self, info: &StreamInfo) {
self.streams.remove(&info.ssrc);
}
fn bind_remote_stream(&mut self, _info: &StreamInfo) {}
fn unbind_remote_stream(&mut self, _info: &StreamInfo) {}
}