use super::recorder::Recorder;
use super::stream_supports_twcc;
use crate::Interceptor;
use crate::stream_info::StreamInfo;
use crate::{AttributedPacket, Packet, TaggedPacket};
use sansio::Protocol;
use shared::TransportContext;
use shared::error::Error;
use shared::marshal::Unmarshal;
use std::collections::{HashMap, VecDeque};
use std::time::{Duration, Instant};
const DEFAULT_INTERVAL: Duration = Duration::from_millis(100);
pub struct TwccReceiverBuilder {
interval: Duration,
}
impl Default for TwccReceiverBuilder {
fn default() -> Self {
Self {
interval: DEFAULT_INTERVAL,
}
}
}
impl TwccReceiverBuilder {
pub fn new() -> Self {
Self::default()
}
pub fn with_interval(mut self, interval: Duration) -> Self {
self.interval = interval;
self
}
pub fn build(self) -> TwccReceiverInterceptor {
TwccReceiverInterceptor::new(self.interval)
}
}
struct RemoteStream {
hdr_ext_id: u8,
}
pub struct TwccReceiverInterceptor {
interval: Duration,
start_time: Option<Instant>,
recorder: Option<Recorder>,
streams: HashMap<u32, RemoteStream>,
write_queue: VecDeque<TaggedPacket>,
next_timeout: Option<Instant>,
read_queue: VecDeque<TaggedPacket>,
}
impl TwccReceiverInterceptor {
fn new(interval: Duration) -> Self {
Self {
read_queue: VecDeque::new(),
interval,
start_time: None,
recorder: None,
streams: HashMap::new(),
write_queue: VecDeque::new(),
next_timeout: None,
}
}
fn generate_feedback(&mut self, now: Instant) {
let Some(recorder) = self.recorder.as_mut() else {
return;
};
let packets = recorder.build_feedback_packet();
for pkt in packets {
self.write_queue.push_back(TaggedPacket {
now,
transport: TransportContext::default(),
message: AttributedPacket::new(Packet::Rtcp(vec![pkt])),
});
}
}
}
impl Protocol<TaggedPacket, TaggedPacket, ()> for TwccReceiverInterceptor {
type Rout = TaggedPacket;
type Wout = TaggedPacket;
type Eout = ();
type Error = Error;
type Time = Instant;
fn handle_read(&mut self, msg: TaggedPacket) -> Result<(), Self::Error> {
if let Packet::Rtp(ref rtp_packet) = msg.message.packet
&& let Some(stream) = self.streams.get(&rtp_packet.header.ssrc)
{
if self.recorder.is_none() {
self.recorder = Some(Recorder::new(rand::random()));
self.start_time = Some(msg.now);
self.next_timeout = Some(msg.now + self.interval);
}
if let Some(ext_data) = rtp_packet.header.get_extension(stream.hdr_ext_id)
&& let Ok(tcc) =
rtp::extension::transport_cc_extension::TransportCcExtension::unmarshal(
&mut ext_data.as_ref(),
)
{
let arrival_time = self
.start_time
.map(|start| msg.now.duration_since(start).as_micros() as i64)
.unwrap_or(0);
if let Some(recorder) = self.recorder.as_mut() {
recorder.record(rtp_packet.header.ssrc, tcc.transport_sequence, arrival_time);
}
}
}
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> {
self.write_queue.push_back(msg);
Ok(())
}
fn poll_write(&mut self) -> Option<TaggedPacket> {
if let Some(pkt) = self.write_queue.pop_front() {
return Some(pkt);
}
None
}
fn handle_timeout(&mut self, now: Instant) -> Result<(), Error> {
if let Some(timeout) = self.next_timeout
&& now >= timeout
{
self.generate_feedback(now);
self.next_timeout = Some(now + self.interval);
}
Ok(())
}
fn poll_timeout(&mut self) -> Option<Instant> {
self.next_timeout
}
}
impl Interceptor for TwccReceiverInterceptor {
fn bind_remote_stream(&mut self, info: &StreamInfo) {
if let Some(hdr_ext_id) = stream_supports_twcc(info) {
if hdr_ext_id != 0 {
self.streams.insert(info.ssrc, RemoteStream { hdr_ext_id });
}
}
}
fn unbind_remote_stream(&mut self, info: &StreamInfo) {
self.streams.remove(&info.ssrc);
}
fn bind_local_stream(&mut self, _info: &StreamInfo) {}
fn unbind_local_stream(&mut self, _info: &StreamInfo) {}
}