use crate::chain::Chain;
use crate::noop::NoopInterceptor;
use crate::{BoxedInterceptor, Interceptor};
use log::warn;
use std::collections::BTreeMap;
#[derive(Copy, Clone, Debug)]
#[non_exhaustive]
#[repr(usize)]
pub enum Slot {
CongestionControl = 1_000,
TwccSender = 2_000,
Pacer = 3_000,
NackResponder = 4_000,
FecEncoder = 5_000,
FecDecoder = 6_000,
NackGenerator = 7_000,
TwccReceiver = 8_000,
Rfc8888 = 9_000,
ReceiverReport = 10_000,
SenderReport = 11_000,
IntervalPli = 12_000,
JitterBuffer = 13_000,
Custom(usize),
}
impl Slot {
pub const fn slot(self) -> usize {
match self {
Self::CongestionControl => 1_000,
Self::TwccSender => 2_000,
Self::Pacer => 3_000,
Self::NackResponder => 4_000,
Self::FecEncoder => 5_000,
Self::FecDecoder => 6_000,
Self::NackGenerator => 7_000,
Self::TwccReceiver => 8_000,
Self::Rfc8888 => 9_000,
Self::ReceiverReport => 10_000,
Self::SenderReport => 11_000,
Self::IntervalPli => 12_000,
Self::JitterBuffer => 13_000,
Self::Custom(position) => position,
}
}
}
impl From<usize> for Slot {
fn from(position: usize) -> Self {
Self::Custom(position)
}
}
impl From<Slot> for usize {
fn from(slot: Slot) -> Self {
slot.slot()
}
}
impl PartialEq for Slot {
fn eq(&self, other: &Self) -> bool {
self.slot() == other.slot()
}
}
impl Eq for Slot {}
impl PartialOrd for Slot {
fn partial_cmp(&self, other: &Self) -> Option<std::cmp::Ordering> {
Some(self.cmp(other))
}
}
impl Ord for Slot {
fn cmp(&self, other: &Self) -> std::cmp::Ordering {
self.slot().cmp(&other.slot())
}
}
impl std::hash::Hash for Slot {
fn hash<H: std::hash::Hasher>(&self, state: &mut H) {
self.slot().hash(state);
}
}
#[derive(Default)]
pub struct Registry {
interceptors: BTreeMap<Slot, BoxedInterceptor>,
names: BTreeMap<Slot, String>,
}
fn short_type_name<T: ?Sized>() -> String {
let full = std::any::type_name::<T>();
let mut out = String::with_capacity(full.len());
let mut segment = String::new();
let flush = |segment: &mut String, out: &mut String| {
out.push_str(segment.rsplit("::").next().unwrap_or(segment));
segment.clear();
};
for ch in full.chars() {
if ch.is_alphanumeric() || ch == '_' || ch == ':' {
segment.push(ch);
} else {
flush(&mut segment, &mut out);
out.push(ch);
}
}
flush(&mut segment, &mut out);
out
}
impl Registry {
pub fn new() -> Self {
Self::default()
}
pub fn with<T: Interceptor + 'static>(mut self, slot: Slot, interceptor: T) -> Self {
let name = short_type_name::<T>();
if let Some(displaced) = self.names.insert(slot, name.clone()) {
warn!("{slot:?} already held {displaced}; {name} replaced it");
}
self.interceptors.insert(slot, Box::new(interceptor));
self
}
pub fn slots(&self) -> Vec<(Slot, String)> {
self.names
.iter()
.map(|(slot, name)| (*slot, name.clone()))
.collect()
}
pub fn build(self) -> impl Interceptor {
let mut interceptors: Vec<BoxedInterceptor> = self.interceptors.into_values().collect();
interceptors.push(Box::new(NoopInterceptor::new()));
Chain::new(interceptors)
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::StreamInfo;
use crate::{AttributedPacket, Packet, TaggedPacket};
use sansio::Protocol;
use shared::TransportContext;
use shared::error::Error;
use std::collections::VecDeque;
use std::sync::{Arc, Mutex};
use std::time::Instant;
#[derive(Clone, Default)]
struct Log(Arc<Mutex<Vec<&'static str>>>);
struct Marker {
name: &'static str,
log: Log,
read_queue: VecDeque<TaggedPacket>,
write_queue: VecDeque<TaggedPacket>,
}
impl Marker {
fn new(name: &'static str, log: Log) -> Self {
Self {
name,
log,
read_queue: VecDeque::new(),
write_queue: VecDeque::new(),
}
}
}
impl Protocol<TaggedPacket, TaggedPacket, ()> for Marker {
type Rout = TaggedPacket;
type Wout = TaggedPacket;
type Eout = ();
type Error = Error;
type Time = Instant;
fn handle_read(&mut self, msg: TaggedPacket) -> Result<(), Self::Error> {
self.log.0.lock().unwrap().push(self.name);
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.log.0.lock().unwrap().push(self.name);
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> {
Ok(())
}
fn poll_timeout(&mut self) -> Option<Self::Time> {
None
}
}
impl Interceptor for Marker {
fn bind_local_stream(&mut self, _info: &StreamInfo) {}
fn unbind_local_stream(&mut self, _info: &StreamInfo) {}
fn bind_remote_stream(&mut self, _info: &StreamInfo) {}
fn unbind_remote_stream(&mut self, _info: &StreamInfo) {}
}
fn packet() -> TaggedPacket {
TaggedPacket {
now: Instant::now(),
transport: TransportContext::default(),
message: AttributedPacket::new(Packet::Rtp(rtp::Packet::default())),
}
}
fn chain(log: &Log) -> impl Interceptor {
Registry::new()
.with(Slot::TwccSender, Marker::new("wire", log.clone()))
.with(Slot::NackGenerator, Marker::new("middle", log.clone()))
.with(Slot::JitterBuffer, Marker::new("app", log.clone()))
.build()
}
#[test]
fn call_order_does_not_decide_chain_order() {
let log = Log::default();
let mut chain = Registry::new()
.with(Slot::JitterBuffer, Marker::new("app", log.clone()))
.with(Slot::TwccSender, Marker::new("wire", log.clone()))
.with(Slot::NackGenerator, Marker::new("middle", log.clone()))
.build();
chain.handle_read(packet()).unwrap();
while chain.poll_read().is_some() {}
assert_eq!(vec!["wire", "middle", "app"], *log.0.lock().unwrap());
}
#[test]
fn a_slot_holds_one_interceptor() {
let log = Log::default();
let mut chain = Registry::new()
.with(Slot::NackGenerator, Marker::new("first", log.clone()))
.with(Slot::NackGenerator, Marker::new("second", log.clone()))
.build();
chain.handle_read(packet()).unwrap();
while chain.poll_read().is_some() {}
assert_eq!(
vec!["second"],
*log.0.lock().unwrap(),
"the later one claimed the slot; the earlier one is not in the chain"
);
}
#[test]
fn read_runs_in_the_order_stages_were_added() {
let log = Log::default();
let mut chain = chain(&log);
chain.handle_read(packet()).unwrap();
while chain.poll_read().is_some() {}
assert_eq!(vec!["wire", "middle", "app"], *log.0.lock().unwrap());
}
#[test]
fn write_runs_in_reverse() {
let log = Log::default();
let mut chain = chain(&log);
chain.handle_write(packet()).unwrap();
while chain.poll_write().is_some() {}
assert_eq!(vec!["app", "middle", "wire"], *log.0.lock().unwrap());
}
#[test]
fn a_registry_with_nothing_added_still_has_the_terminus() {
let mut chain = Registry::new().build();
chain
.handle_read(TaggedPacket {
now: Instant::now(),
transport: TransportContext::default(),
message: AttributedPacket::new(Packet::Rtcp(vec![])),
})
.unwrap();
assert!(
chain.poll_read().is_none(),
"inbound RTCP stops before the application"
);
}
#[test]
fn the_terminus_is_application_most() {
let log = Log::default();
let mut chain = Registry::new()
.with(Slot::TwccSender, Marker::new("wire", log.clone()))
.build();
chain
.handle_read(TaggedPacket {
now: Instant::now(),
transport: TransportContext::default(),
message: AttributedPacket::new(Packet::Rtcp(vec![])),
})
.unwrap();
assert_eq!(
vec!["wire"],
*log.0.lock().unwrap(),
"the stage saw the RTCP packet; the terminus dropped it afterwards"
);
assert!(chain.poll_read().is_none());
}
#[test]
fn a_custom_slot_sits_where_its_number_says() {
let log = Log::default();
let mut chain = Registry::new()
.with(Slot::FecDecoder, Marker::new("fec", log.clone()))
.with(Slot::NackGenerator, Marker::new("nack", log.clone()))
.with(Slot::from(6_500), Marker::new("mine", log.clone()))
.build();
chain.handle_read(packet()).unwrap();
while chain.poll_read().is_some() {}
assert_eq!(
vec!["fec", "mine", "nack"],
*log.0.lock().unwrap(),
"6_500 belongs after the FEC decoder at 6_000 and before the NACK generator at 7_000"
);
}
#[test]
fn equality_and_ordering_both_follow_the_position() {
assert_eq!(Slot::TwccSender, Slot::from(2_000));
assert_eq!(
std::cmp::Ordering::Equal,
Slot::TwccSender.cmp(&Slot::from(2_000))
);
assert!(Slot::from(1_500) > Slot::CongestionControl);
assert!(Slot::from(1_500) < Slot::TwccSender);
assert!(
Slot::from(20_000) > Slot::JitterBuffer,
"a position past every named slot sorts past them, not by declaration order"
);
}
#[test]
fn the_named_slots_are_spaced_by_a_thousand() {
let named = [
Slot::CongestionControl,
Slot::TwccSender,
Slot::Pacer,
Slot::NackResponder,
Slot::FecEncoder,
Slot::FecDecoder,
Slot::NackGenerator,
Slot::TwccReceiver,
Slot::Rfc8888,
Slot::ReceiverReport,
Slot::SenderReport,
Slot::IntervalPli,
Slot::JitterBuffer,
];
for pair in named.windows(2) {
assert_eq!(
1_000,
pair[1].slot() - pair[0].slot(),
"{:?} and {:?} must stay a thousand apart",
pair[0],
pair[1]
);
}
}
#[test]
fn slots_carry_the_interceptor_names() {
let log = Log::default();
let registry = Registry::new()
.with(Slot::JitterBuffer, Marker::new("app", log.clone()))
.with(Slot::TwccSender, crate::TwccSenderBuilder::new().build());
assert_eq!(
vec![
(Slot::TwccSender, "TwccSenderInterceptor".to_owned()),
(Slot::JitterBuffer, "Marker".to_owned()),
],
registry.slots(),
"names come back with their slots, sorted wire-to-application"
);
}
#[test]
fn names_are_stripped_of_their_module_path() {
let registry = Registry::new().with(
Slot::CongestionControl,
crate::CongestionControlBuilder::new(crate::ConstantBitrate::new(1_000_000.0)).build(),
);
let (_, name) = ®istry.slots()[0];
assert!(
!name.contains("::"),
"a module path leaked into the name: {name}"
);
assert_eq!(
"CongestionControlInterceptor<ConstantBitrate>", name,
"the generic argument is shortened too, and kept — it is what tells two \
congestion controllers apart"
);
}
}