use crate::Interceptor;
use crate::{Packet, StreamInfo, TaggedPacket};
use sansio::Protocol;
use shared::error::Error;
use std::collections::VecDeque;
use std::time::Instant;
#[derive(Default)]
pub struct NoopInterceptor {
rtcp_readable: bool,
read_queue: VecDeque<TaggedPacket>,
write_queue: VecDeque<TaggedPacket>,
}
impl NoopInterceptor {
pub fn new(rtcp_readable: bool) -> Self {
Self {
rtcp_readable,
..Default::default()
}
}
}
impl Protocol<TaggedPacket, TaggedPacket, ()> for NoopInterceptor {
type Rout = TaggedPacket;
type Wout = TaggedPacket;
type Eout = ();
type Error = Error;
type Time = Instant;
fn handle_read(&mut self, msg: TaggedPacket) -> Result<(), Self::Error> {
let keep = match msg.message.packet {
Packet::Rtp(_) => true,
Packet::Rtcp(_) => self.rtcp_readable,
};
if keep {
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<Self::Wout> {
self.write_queue.pop_front()
}
}
impl Interceptor for NoopInterceptor {
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) {}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::{AttributedPacket, Registry, StreamInfo};
use sansio::Protocol;
use shared::TransportContext;
use shared::error::Error;
use std::collections::VecDeque;
use std::time::Instant;
fn packet(message: Packet) -> TaggedPacket {
TaggedPacket {
now: Instant::now(),
transport: TransportContext::default(),
message: AttributedPacket::new(message),
}
}
#[test]
fn inbound_rtcp_does_not_reach_the_application() {
let mut chain = Registry::new().build();
chain.handle_read(packet(Packet::Rtcp(vec![]))).unwrap();
assert!(chain.poll_read().is_none());
}
#[test]
fn inbound_rtp_passes_through() {
let mut chain = Registry::new().build();
chain
.handle_read(packet(Packet::Rtp(rtp::Packet::default())))
.unwrap();
assert!(chain.poll_read().is_some());
}
#[test]
fn outbound_rtcp_is_not_affected() {
let mut chain = Registry::new().build();
chain.handle_write(packet(Packet::Rtcp(vec![]))).unwrap();
assert!(chain.poll_write().is_some());
}
#[test]
fn stages_before_it_still_see_inbound_rtcp() {
#[derive(Default)]
struct Counter {
seen: std::sync::Arc<std::sync::atomic::AtomicUsize>,
read_queue: VecDeque<TaggedPacket>,
write_queue: VecDeque<TaggedPacket>,
}
impl Protocol<TaggedPacket, TaggedPacket, ()> for Counter {
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 matches!(msg.message.packet, Packet::Rtcp(_)) {
self.seen.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
}
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<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 Counter {
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) {}
}
let counter = Counter::default();
let seen = counter.seen.clone();
let mut chain = Registry::new().with(counter).build();
chain.handle_read(packet(Packet::Rtcp(vec![]))).unwrap();
assert_eq!(1, seen.load(std::sync::atomic::Ordering::Relaxed));
assert!(chain.poll_read().is_none(), "but it stops at the terminus");
}
}