use crate::rtp::{RtcpPacket, RtpPacket, is_rtcp, marshal_rtcp_packets, parse_rtcp_packets};
use crate::srtp::SrtpSession;
use crate::peer_connection::RtpObserver;
use crate::transports::PacketReceiver;
use crate::transports::ice::conn::IceConn;
use crate::transports::ice::stun::random_u32;
use anyhow::Result;
use async_trait::async_trait;
use bytes::Bytes;
use parking_lot::{Mutex, RwLock};
use std::cell::RefCell;
use std::collections::HashMap;
use std::net::SocketAddr;
use std::sync::Arc;
use std::sync::atomic::{AtomicBool, AtomicU64, AtomicU8, Ordering};
use tokio::sync::mpsc;
use tracing::{debug, trace};
const EXT_ID_NONE: u8 = 0;
#[inline]
fn encode_ext_id(id: Option<u8>) -> u8 {
id.unwrap_or(EXT_ID_NONE)
}
#[inline]
fn decode_ext_id(raw: u8) -> Option<u8> {
if raw == EXT_ID_NONE { None } else { Some(raw) }
}
async fn try_send_with_fallback<T>(
tx: &mpsc::Sender<T>,
value: T,
) -> Result<(), mpsc::error::SendError<T>> {
match tx.try_send(value) {
Ok(()) => Ok(()),
Err(mpsc::error::TrySendError::Full(value)) => tx.send(value).await,
Err(mpsc::error::TrySendError::Closed(value)) => Err(mpsc::error::SendError(value)),
}
}
fn try_send_dropping<T>(
tx: &mpsc::Sender<T>,
value: T,
) -> Result<(), mpsc::error::TrySendError<T>> {
tx.try_send(value)
}
#[derive(Debug, Clone, Copy, Default)]
pub struct RtpRewriteBridgeParams {
pub ssrc_offset: u32,
pub fixed_out_ssrc: Option<u32>,
pub payload_type: Option<u8>,
pub dtmf_payload_type: Option<(u8, u8)>,
pub initial_sequence_number: Option<u16>,
pub initial_timestamp_offset: Option<u32>,
pub strip_extensions: bool,
}
#[derive(Debug, Clone, Copy)]
pub struct RtpRewriteRule {
pub match_payload_type: Option<u8>,
pub fixed_out_ssrc: Option<u32>,
pub ssrc_offset: u32,
pub out_payload_type: Option<u8>,
}
#[derive(Debug, Clone, Copy, Default)]
pub struct RtpRewriteBridgeOptions {
pub strip_extensions: bool,
pub initial_sequence_number: Option<u16>,
pub initial_timestamp_offset: Option<u32>,
}
impl RtpRewriteRule {
pub fn catch_all(params: RtpRewriteBridgeParams) -> Self {
Self {
match_payload_type: None,
fixed_out_ssrc: params.fixed_out_ssrc,
ssrc_offset: params.ssrc_offset,
out_payload_type: params.payload_type,
}
}
pub fn dtmf(src_pt: u8, dst_pt: u8, params: RtpRewriteBridgeParams) -> Self {
Self {
match_payload_type: Some(src_pt),
fixed_out_ssrc: params.fixed_out_ssrc,
ssrc_offset: params.ssrc_offset,
out_payload_type: Some(dst_pt),
}
}
pub fn from_params(params: RtpRewriteBridgeParams) -> Vec<Self> {
let mut rules = vec![Self::catch_all(params)];
if let Some((src_pt, dst_pt)) = params.dtmf_payload_type {
rules.push(Self::dtmf(src_pt, dst_pt, params));
}
rules
}
}
#[derive(Debug, Clone, Copy)]
struct StreamRewriteState {
out_ssrc: u32,
next_sequence_number: u16,
last_source_timestamp: Option<u32>,
timestamp_offset: u32,
}
struct RewriteBridge {
target: Arc<RtpTransport>,
options: RtpRewriteBridgeOptions,
rules: Vec<RtpRewriteRule>,
streams: RefCell<HashMap<u32, StreamRewriteState>>,
}
impl RewriteBridge {
fn new(
target: Arc<RtpTransport>,
options: RtpRewriteBridgeOptions,
rules: Vec<RtpRewriteRule>,
) -> Self {
Self {
target,
options,
rules,
streams: RefCell::new(HashMap::new()),
}
}
fn rule_for(&self, raw_pt: u8) -> Option<RtpRewriteRule> {
self.rules
.iter()
.copied()
.find(|r| r.match_payload_type == Some(raw_pt))
.or_else(|| self.rules.iter().copied().find(|r| r.match_payload_type.is_none()))
}
fn rewrite_packet(&self, packet: &mut RtpPacket) {
let src_ssrc = packet.header.ssrc;
let src_timestamp = packet.header.timestamp;
if self.options.strip_extensions {
packet.header.extension = None;
}
let raw_pt = packet.header.payload_type;
let rule = self.rule_for(raw_pt);
let out_ssrc = match rule {
Some(r) => r
.fixed_out_ssrc
.unwrap_or_else(|| src_ssrc.wrapping_add(r.ssrc_offset)),
None => src_ssrc,
};
let mut streams = self.streams.borrow_mut();
let state = streams
.entry(src_ssrc)
.or_insert_with(|| StreamRewriteState {
out_ssrc,
next_sequence_number: self
.options
.initial_sequence_number
.unwrap_or(random_u32() as u16),
last_source_timestamp: None,
timestamp_offset: self
.options
.initial_timestamp_offset
.unwrap_or_else(random_u32),
});
if let Some(r) = rule {
if let Some(payload_type) = r.out_payload_type {
packet.header.payload_type = payload_type;
}
}
packet.header.ssrc = state.out_ssrc;
if let Some(last_src) = state.last_source_timestamp {
let delta = src_timestamp.wrapping_sub(last_src);
if delta < 0x8000_0000 {
if delta > 900_000 {
state.timestamp_offset = last_src
.wrapping_add(state.timestamp_offset)
.wrapping_add(3000)
.wrapping_sub(src_timestamp);
}
state.last_source_timestamp = Some(src_timestamp);
}
} else {
state.last_source_timestamp = Some(src_timestamp);
}
packet.header.timestamp = src_timestamp.wrapping_add(state.timestamp_offset);
packet.header.sequence_number = state.next_sequence_number;
state.next_sequence_number = state.next_sequence_number.wrapping_add(1);
}
}
#[derive(Default)]
struct ListenerRegistry {
by_ssrc: HashMap<u32, mpsc::Sender<(RtpPacket, SocketAddr)>>,
by_rid: HashMap<String, mpsc::Sender<(RtpPacket, SocketAddr)>>,
by_mid: HashMap<String, mpsc::Sender<(RtpPacket, SocketAddr)>>,
routes: Vec<ListenerRoute>,
}
#[derive(Clone)]
struct ListenerRoute {
mid: Option<String>,
payload_types: Vec<u8>,
tx: mpsc::Sender<(RtpPacket, SocketAddr)>,
provisional: bool,
}
impl ListenerRegistry {
fn route_for_sender_mut(
&mut self,
tx: &mpsc::Sender<(RtpPacket, SocketAddr)>,
) -> &mut ListenerRoute {
if let Some(index) = self
.routes
.iter()
.position(|route| route.tx.same_channel(tx))
{
return &mut self.routes[index];
}
self.routes.retain(|route| !route.tx.is_closed());
self.routes.push(ListenerRoute {
mid: None,
payload_types: Vec::new(),
tx: tx.clone(),
provisional: false,
});
self.routes.last_mut().unwrap()
}
fn register_mid(&mut self, mid: String, tx: mpsc::Sender<(RtpPacket, SocketAddr)>) {
self.by_mid.insert(mid.clone(), tx.clone());
self.route_for_sender_mut(&tx).mid = Some(mid);
}
fn register_payload_types(
&mut self,
payload_types: Vec<u8>,
tx: mpsc::Sender<(RtpPacket, SocketAddr)>,
) {
let route = self.route_for_sender_mut(&tx);
route.payload_types.clear();
for pt in payload_types {
if !route.payload_types.contains(&pt) {
route.payload_types.push(pt);
}
}
}
fn register_payload_type(&mut self, pt: u8, tx: mpsc::Sender<(RtpPacket, SocketAddr)>) {
let route = self.route_for_sender_mut(&tx);
if !route.payload_types.contains(&pt) {
route.payload_types.push(pt);
}
}
fn register_provisional(&mut self, tx: mpsc::Sender<(RtpPacket, SocketAddr)>) {
self.route_for_sender_mut(&tx).provisional = true;
}
fn by_mid(&self, mid: &str) -> Option<mpsc::Sender<(RtpPacket, SocketAddr)>> {
self.by_mid.get(mid).cloned()
}
fn unique_by_pt(&self, pt: u8) -> Option<mpsc::Sender<(RtpPacket, SocketAddr)>> {
let mut selected: Option<&mpsc::Sender<(RtpPacket, SocketAddr)>> = None;
for route in self
.routes
.iter()
.filter(|route| route.payload_types.contains(&pt))
{
if let Some(existing) = selected {
if !existing.same_channel(&route.tx) {
return None;
}
} else {
selected = Some(&route.tx);
}
}
selected.cloned()
}
fn single_provisional(&self) -> Option<mpsc::Sender<(RtpPacket, SocketAddr)>> {
let mut selected: Option<&mpsc::Sender<(RtpPacket, SocketAddr)>> = None;
for route in self.routes.iter().filter(|route| route.provisional) {
if let Some(existing) = selected {
if !existing.same_channel(&route.tx) {
return None;
}
} else {
selected = Some(&route.tx);
}
}
selected.cloned()
}
fn bind_ssrc_route(&mut self, ssrc: u32, tx: mpsc::Sender<(RtpPacket, SocketAddr)>) {
self.by_ssrc.retain(|_, existing| !existing.is_closed());
self.by_ssrc.insert(ssrc, tx);
}
fn remove_sender(&mut self, tx: &mpsc::Sender<(RtpPacket, SocketAddr)>) {
self.by_ssrc
.retain(|_, existing| !existing.same_channel(tx));
self.by_rid.retain(|_, existing| !existing.same_channel(tx));
self.by_mid.retain(|_, existing| !existing.same_channel(tx));
self.routes.retain(|route| !route.tx.same_channel(tx));
}
}
pub struct RtpTransport {
transport: Arc<IceConn>,
srtp_session: Mutex<Option<Arc<Mutex<SrtpSession>>>>,
listeners: Mutex<ListenerRegistry>,
rtcp_listener: Mutex<Option<mpsc::Sender<Vec<RtcpPacket>>>>,
rid_extension_id: AtomicU8,
sdes_mid_extension_id: AtomicU8,
abs_send_time_extension_id: AtomicU8,
rewrite_bridge: Mutex<Option<Box<RewriteBridge>>>,
has_bridge: AtomicBool,
srtp_required: bool,
has_sent_first_packet: AtomicBool,
received_rtp_packets: AtomicU64,
relay_send_failures: AtomicU64,
observers: RwLock<Vec<Arc<dyn RtpObserver>>>,
has_observers: AtomicBool,
}
impl RtpTransport {
pub fn new(transport: Arc<IceConn>, srtp_required: bool) -> Self {
Self::new_with_ssrc_change(transport, srtp_required, false)
}
pub fn new_with_ssrc_change(
transport: Arc<IceConn>,
srtp_required: bool,
_allow_ssrc_change: bool,
) -> Self {
Self {
transport,
srtp_session: Mutex::new(None),
listeners: Mutex::new(ListenerRegistry::default()),
rtcp_listener: Mutex::new(None),
rid_extension_id: AtomicU8::new(EXT_ID_NONE),
sdes_mid_extension_id: AtomicU8::new(EXT_ID_NONE),
abs_send_time_extension_id: AtomicU8::new(EXT_ID_NONE),
rewrite_bridge: Mutex::new(None),
has_bridge: AtomicBool::new(false),
srtp_required,
has_sent_first_packet: AtomicBool::new(false),
received_rtp_packets: AtomicU64::new(0),
relay_send_failures: AtomicU64::new(0),
observers: RwLock::new(Vec::new()),
has_observers: AtomicBool::new(false),
}
}
pub fn received_rtp_packets(&self) -> u64 {
self.received_rtp_packets.load(Ordering::Relaxed)
}
pub fn ice_conn(&self) -> Arc<IceConn> {
self.transport.clone()
}
pub fn start_srtp(&self, srtp_session: SrtpSession) {
let mut session = self.srtp_session.lock();
*session = Some(Arc::new(Mutex::new(srtp_session)));
}
pub fn register_listener_sync(&self, ssrc: u32, tx: mpsc::Sender<(RtpPacket, SocketAddr)>) {
let mut listeners = self.listeners.lock();
listeners.bind_ssrc_route(ssrc, tx);
}
pub fn has_listener(&self, ssrc: u32) -> bool {
let listeners = self.listeners.lock();
listeners.by_ssrc.contains_key(&ssrc)
}
pub fn register_rid_listener(&self, rid: String, tx: mpsc::Sender<(RtpPacket, SocketAddr)>) {
let mut listeners = self.listeners.lock();
listeners.by_rid.retain(|_, existing| !existing.is_closed());
listeners.by_rid.insert(rid, tx);
}
pub fn register_mid_listener(&self, mid: String, tx: mpsc::Sender<(RtpPacket, SocketAddr)>) {
let mut listeners = self.listeners.lock();
listeners.register_mid(mid, tx);
}
pub fn register_pt_listener(&self, pt: u8, tx: mpsc::Sender<(RtpPacket, SocketAddr)>) {
let mut listeners = self.listeners.lock();
listeners.register_payload_type(pt, tx);
}
pub fn register_payload_list_listener(
&self,
payload_types: Vec<u8>,
tx: mpsc::Sender<(RtpPacket, SocketAddr)>,
) {
let mut listeners = self.listeners.lock();
listeners.register_payload_types(payload_types, tx);
}
pub fn register_provisional_listener(&self, tx: mpsc::Sender<(RtpPacket, SocketAddr)>) {
let mut listeners = self.listeners.lock();
listeners.register_provisional(tx);
}
pub fn set_rid_extension_id(&self, id: Option<u8>) {
self.rid_extension_id
.store(encode_ext_id(id), Ordering::Relaxed);
}
pub fn set_sdes_mid_extension_id(&self, id: Option<u8>) {
self.sdes_mid_extension_id
.store(encode_ext_id(id), Ordering::Relaxed);
}
pub fn set_abs_send_time_extension_id(&self, id: Option<u8>) {
self.abs_send_time_extension_id
.store(encode_ext_id(id), Ordering::Relaxed);
}
pub fn remote_addr(&self) -> std::net::SocketAddr {
*self.transport.remote_addr.read()
}
pub fn local_addr(&self) -> std::net::SocketAddr {
self.transport.local_addr()
}
pub fn register_rtcp_listener(&self, tx: mpsc::Sender<Vec<RtcpPacket>>) {
let mut listener = self.rtcp_listener.lock();
*listener = Some(tx);
}
pub fn bridge_rewrite_to(&self, dst: Arc<RtpTransport>, params: RtpRewriteBridgeParams) {
let options = RtpRewriteBridgeOptions {
strip_extensions: params.strip_extensions,
initial_sequence_number: params.initial_sequence_number,
initial_timestamp_offset: params.initial_timestamp_offset,
};
self.bridge_rewrite_rules_to(dst, options, RtpRewriteRule::from_params(params));
}
pub fn bridge_rewrite_rules_to(
&self,
dst: Arc<RtpTransport>,
options: RtpRewriteBridgeOptions,
rules: Vec<RtpRewriteRule>,
) {
*self.rewrite_bridge.lock() = Some(Box::new(RewriteBridge::new(dst, options, rules)));
self.has_bridge.store(true, Ordering::Release);
}
pub fn clear_bridge_rewrite(&self) {
*self.rewrite_bridge.lock() = None;
self.has_bridge.store(false, Ordering::Release);
}
pub fn add_observer(&self, observer: Arc<dyn RtpObserver>) {
self.observers.write().push(observer);
self.has_observers.store(true, Ordering::Release);
}
pub fn clear_observers(&self) {
self.observers.write().clear();
self.has_observers.store(false, Ordering::Release);
}
#[inline]
fn fire_ingress(&self, packet: &RtpPacket, src_addr: SocketAddr) {
if !self.has_observers.load(Ordering::Acquire) {
return;
}
let observers = self.observers.read();
for o in observers.iter() {
o.on_ingress(packet, src_addr);
}
}
#[inline]
fn fire_egress(&self, packet: &RtpPacket) {
if !self.has_observers.load(Ordering::Acquire) {
return;
}
let dst_addr = *self.transport.remote_addr.read();
let observers = self.observers.read();
for o in observers.iter() {
o.on_egress(packet, dst_addr);
}
}
pub async fn send(&self, buf: &[u8]) -> Result<usize> {
let protected = {
let session = self.srtp_session.lock().as_ref().map(|s| s.clone());
match session {
Some(session) => {
let mut packet = RtpPacket::parse(buf)?;
if let Some(id) =
decode_ext_id(self.abs_send_time_extension_id.load(Ordering::Relaxed))
{
let abs_send_time =
crate::rtp::calculate_abs_send_time(std::time::SystemTime::now());
let data = abs_send_time.to_be_bytes();
packet.header.set_extension(id, &data[1..4])?;
}
{
let mut srtp = session.lock();
srtp.protect_rtp(&mut packet)?;
}
packet.marshal()?
}
None => {
if self.srtp_required {
return Err(anyhow::anyhow!("SRTP required but session not ready"));
}
buf.to_vec()
}
}
};
self.transport.send(&protected).await
}
pub async fn send_rtp(&self, mut packet: RtpPacket) -> Result<usize> {
self.fire_egress(&packet);
let is_first = !self.has_sent_first_packet.load(Ordering::Relaxed);
if is_first {
self.has_sent_first_packet.store(true, Ordering::Relaxed);
packet.header.marker = true;
}
if let Some(id) = decode_ext_id(self.abs_send_time_extension_id.load(Ordering::Relaxed)) {
let abs_send_time = crate::rtp::calculate_abs_send_time(std::time::SystemTime::now());
let data = abs_send_time.to_be_bytes();
if let Err(e) = packet.header.set_extension(id, &data[1..4]) {
trace!("RtpTransport: abs-send-time extension skipped: {}", e);
}
}
let protected = {
let session = self.srtp_session.lock().as_ref().map(|s| s.clone());
match session {
Some(session) => {
{
let mut srtp = session.lock();
srtp.protect_rtp(&mut packet)?;
}
packet.marshal()?
}
None => {
if self.srtp_required {
debug!("RtpTransport: SRTP required but session not ready, dropping RTP send");
return Err(anyhow::anyhow!("SRTP required but session not ready"));
}
packet.marshal()?
}
}
};
match self.transport.send(&protected).await {
Ok(n) => {
if is_first {
self.transport.mark_first_outbound();
trace!(
"RtpTransport: first SRTP packet sent ({} bytes)",
protected.len()
);
}
Ok(n)
}
Err(e) => {
debug!(
"RtpTransport: failed to send SRTP packet ({} bytes): {}",
protected.len(),
e
);
Err(e)
}
}
}
pub async fn send_rtcp(&self, packets: &[RtcpPacket]) -> Result<usize> {
let mut raw = marshal_rtcp_packets(packets)?;
let protected = {
let session_guard = self.srtp_session.lock();
if let Some(session) = &*session_guard {
let mut srtp = session.lock();
srtp.protect_rtcp(&mut raw)?;
raw
} else {
if self.srtp_required {
debug!("Failed to send PLI: SRTP required but session not ready");
return Err(anyhow::anyhow!("SRTP required but session not ready"));
}
raw
}
};
self.transport.send_rtcp(&protected).await
}
pub fn send_rtcp_sync(&self, packets: &[RtcpPacket]) {
let Ok(mut raw) = marshal_rtcp_packets(packets) else {
return;
};
{
let session_guard = self.srtp_session.lock();
if let Some(session) = &*session_guard {
if session.lock().protect_rtcp(&mut raw).is_err() {
return;
}
} else if self.srtp_required {
return;
}
}
let _ = self.ice_conn().try_send(&raw);
}
fn try_bridge_rewrite_rtp(
&self,
mut packet: RtpPacket,
marshal_buf: &mut Vec<u8>,
) -> Option<RtpPacket> {
if !self.has_bridge.load(Ordering::Acquire) {
return Some(packet);
}
let target = {
let mut guard = self.rewrite_bridge.lock();
let Some(bridge) = guard.as_mut() else {
return Some(packet);
};
bridge.rewrite_packet(&mut packet);
bridge.target.clone()
};
target.fire_egress(&packet);
{
let session_guard = target.srtp_session.lock();
if let Some(session) = &*session_guard {
if session.lock().protect_rtp(&mut packet).is_err() {
return None; }
} else if target.srtp_required {
return None; }
}
packet.marshal_into(marshal_buf);
if let Err(e) = target.ice_conn().try_send(marshal_buf) {
let relay_failures = self.relay_send_failures.fetch_add(1, Ordering::Relaxed) + 1;
if relay_failures <= 5 || relay_failures % 100 == 0 {
tracing::warn!(
relay_failures,
error = %e,
ssrc = packet.header.ssrc,
pt = packet.header.payload_type,
"RTP rewrite bridge: relay push to destination failed"
);
}
}
None
}
pub fn clear_listeners(&self) -> usize {
let mut count = 0;
{
let mut listeners = self.listeners.lock();
count += listeners.by_ssrc.len();
listeners.by_ssrc.clear();
count += listeners.by_rid.len();
listeners.by_rid.clear();
count += listeners.routes.len();
listeners.routes.clear();
}
{
let mut rtcp_listener = self.rtcp_listener.lock();
if rtcp_listener.is_some() {
*rtcp_listener = None;
count += 1;
}
}
count
}
}
#[async_trait]
impl PacketReceiver for RtpTransport {
async fn receive(&self, packet: Bytes, addr: SocketAddr, marshal_buf: &mut Vec<u8>) {
let is_rtcp_packet = is_rtcp(&packet);
if is_rtcp_packet {
let unprotected: Bytes = {
let session = self.srtp_session.lock().as_ref().map(|s| s.clone());
match session {
Some(session) => {
let mut buf = packet.to_vec();
let mut srtp = session.lock();
match srtp.unprotect_rtcp(&mut buf) {
Ok(()) => Bytes::from(buf),
Err(e) => {
debug!("SRTP unprotect RTCP failed: {}", e);
return;
}
}
}
None => {
if self.srtp_required {
trace!("Dropping packet because SRTP is required but session is not ready");
return;
}
packet
}
}
};
let listener = {
let guard = self.rtcp_listener.lock();
guard.clone()
};
if let Some(tx) = listener {
match parse_rtcp_packets(&unprotected, Some(addr)) {
Ok(packets) => {
if try_send_with_fallback(&tx, packets).await.is_err() {
let mut guard = self.rtcp_listener.lock();
*guard = None;
}
}
Err(e) => {
trace!("RTCP parse failed: {}", e);
}
}
} else {
trace!(
"No RTCP listener, dropping {} bytes from {}",
unprotected.len(),
addr
);
}
} else {
let rtp_packet = {
let session = self.srtp_session.lock().as_ref().map(|s| s.clone());
match session {
Some(session) => {
let mut srtp = session.lock();
match RtpPacket::parse(&packet) {
Ok(mut rtp_packet) => match srtp.unprotect_rtp(&mut rtp_packet) {
Ok(_) => rtp_packet,
Err(_) => return,
},
Err(e) => {
trace!("RTP parse failed: {}", e);
return;
}
}
}
None => {
if self.srtp_required {
trace!("Dropping packet because SRTP is required but session is not ready");
return;
}
match RtpPacket::parse_bytes(packet.clone()) {
Ok(rtp_packet) => rtp_packet,
Err(e) => {
trace!("RTP parse failed: {}", e);
return;
}
}
}
}
};
self.received_rtp_packets.fetch_add(1, Ordering::Relaxed);
self.fire_ingress(&rtp_packet, addr);
let Some(rtp_packet) = self.try_bridge_rewrite_rtp(rtp_packet, marshal_buf) else {
return;
};
let ssrc = rtp_packet.header.ssrc;
let pt = rtp_packet.header.payload_type;
let rid_id = decode_ext_id(self.rid_extension_id.load(Ordering::Relaxed));
let mid_id = decode_ext_id(self.sdes_mid_extension_id.load(Ordering::Relaxed));
let rid_bytes = rid_id.and_then(|id| rtp_packet.header.get_extension(id));
let mid_bytes = mid_id.and_then(|id| rtp_packet.header.get_extension(id));
let listener = {
let mut listeners = self.listeners.lock();
let mut selected = None;
let mut bind_ssrc = false;
if let Some(rid) = &rid_bytes
&& let Ok(rid_str) = std::str::from_utf8(rid)
{
selected = listeners.by_rid.get(rid_str).cloned();
bind_ssrc = selected.is_some();
}
if selected.is_none()
&& let Some(mid) = &mid_bytes
&& let Ok(mid_str) = std::str::from_utf8(mid)
{
selected = listeners.by_mid(mid_str);
bind_ssrc = selected.is_some();
}
if selected.is_none() {
selected = listeners.by_ssrc.get(&ssrc).cloned();
bind_ssrc = false;
}
if selected.is_none() {
selected = listeners.unique_by_pt(pt);
bind_ssrc = selected.is_some();
}
if selected.is_none() {
selected = listeners.single_provisional();
bind_ssrc = false;
}
if let Some(tx) = selected.as_ref()
&& bind_ssrc
{
listeners.bind_ssrc_route(ssrc, tx.clone());
}
selected
};
if let Some(tx) = listener {
match try_send_dropping(&tx, (rtp_packet, addr)) {
Ok(()) => {}
Err(mpsc::error::TrySendError::Full(_)) => {}
Err(mpsc::error::TrySendError::Closed(_)) => {
let mut listeners = self.listeners.lock();
listeners.by_ssrc.remove(&ssrc);
listeners.remove_sender(&tx);
}
}
} else {
trace!(
"No listener found for packet SSRC: {} PT: {} from {}",
ssrc, pt, addr
);
}
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::transports::ice::conn::IceConn;
use tokio::sync::mpsc;
#[tokio::test]
async fn test_specific_listener_isolation() {
use crate::transports::ice::IceSocketWrapper;
use bytes::Bytes;
use tokio::sync::watch;
let (_ice_tx, ice_rx) = watch::channel(None::<IceSocketWrapper>);
let ice_conn = IceConn::new(ice_rx, "127.0.0.1:1234".parse().unwrap(), None);
let transport = RtpTransport::new(ice_conn, false);
let (tx, mut rx) = mpsc::channel(10);
transport.register_listener_sync(100, tx);
let header1 = crate::rtp::RtpHeader::new(0, 1, 0, 100);
let packet1 = crate::rtp::RtpPacket::new(header1, vec![1u8; 160]);
let mut marshal_buf = Vec::new();
transport
.receive(
Bytes::from(packet1.marshal().unwrap()),
"127.0.0.1:5000".parse().unwrap(),
&mut marshal_buf,
)
.await;
let received1 = rx.recv().await.expect("First packet should be received");
assert_eq!(received1.0.header.ssrc, 100);
let header2 = crate::rtp::RtpHeader::new(0, 2, 160, 200);
let packet2 = crate::rtp::RtpPacket::new(header2, vec![2u8; 160]);
transport
.receive(
Bytes::from(packet2.marshal().unwrap()),
"127.0.0.1:5000".parse().unwrap(),
&mut marshal_buf,
)
.await;
tokio::time::timeout(tokio::time::Duration::from_millis(50), rx.recv())
.await
.expect_err(
"Second packet with new SSRC should be dropped when allow_ssrc_change=false",
);
assert!(!transport.has_listener(200));
}
#[tokio::test]
async fn test_provisional_listener_promiscuous_mode() {
use crate::transports::ice::IceSocketWrapper;
use bytes::Bytes;
use tokio::sync::watch;
let (_ice_tx, ice_rx) = watch::channel(None::<IceSocketWrapper>);
let ice_conn = IceConn::new(ice_rx, "127.0.0.1:1234".parse().unwrap(), None);
let transport = RtpTransport::new(ice_conn, false);
let (tx, mut rx) = mpsc::channel(100);
transport.register_provisional_listener(tx);
let addr = "127.0.0.1:5000".parse().unwrap();
let ssrc1 = 1111u32;
let header1 = crate::rtp::RtpHeader::new(0, 1, 0, ssrc1);
let packet1 = crate::rtp::RtpPacket::new(header1, vec![0u8; 160]);
let bytes1 = packet1.marshal().unwrap();
let mut marshal_buf = Vec::new();
transport
.receive(Bytes::from(bytes1), addr, &mut marshal_buf)
.await;
let received1 = rx.recv().await.expect("Should receive packet 1");
assert_eq!(received1.0.header.ssrc, ssrc1);
assert!(
!transport.has_listener(ssrc1),
"SSRC should NOT be bound in promiscuous mode"
);
let ssrc2 = 2222u32;
let header2 = crate::rtp::RtpHeader::new(0, 2, 160, ssrc2);
let packet2 = crate::rtp::RtpPacket::new(header2, vec![1u8; 160]);
let bytes2 = packet2.marshal().unwrap();
transport
.receive(Bytes::from(bytes2), addr, &mut marshal_buf)
.await;
let received2 = rx.recv().await.expect("Should receive packet 2 (new SSRC)");
assert_eq!(received2.0.header.ssrc, ssrc2);
let ssrc3 = 3333u32;
let header3 = crate::rtp::RtpHeader::new(8, 3, 320, ssrc3); let packet3 = crate::rtp::RtpPacket::new(header3, vec![2u8; 160]);
let bytes3 = packet3.marshal().unwrap();
transport
.receive(Bytes::from(bytes3), addr, &mut marshal_buf)
.await;
let received3 = rx
.recv()
.await
.expect("Should receive packet 3 (New PT/SSRC)");
assert_eq!(received3.0.header.ssrc, ssrc3);
assert_eq!(received3.0.header.payload_type, 8);
}
#[tokio::test]
async fn test_ambiguous_payload_type_without_mid_or_ssrc_is_dropped() {
use crate::transports::ice::IceSocketWrapper;
use bytes::Bytes;
use tokio::sync::watch;
let (_ice_tx, ice_rx) = watch::channel(None::<IceSocketWrapper>);
let ice_conn = IceConn::new(ice_rx, "127.0.0.1:1234".parse().unwrap(), None);
let transport = RtpTransport::new(ice_conn, false);
let (audio_tx, mut audio_rx) = mpsc::channel(10);
transport.register_provisional_listener(audio_tx.clone());
transport.register_payload_list_listener(vec![96], audio_tx);
let (video_tx, mut video_rx) = mpsc::channel(10);
transport.register_provisional_listener(video_tx.clone());
transport.register_payload_list_listener(vec![96], video_tx);
let header = crate::rtp::RtpHeader::new(96, 1, 0, 4444);
let packet = crate::rtp::RtpPacket::new(header, vec![0u8; 160]);
let mut marshal_buf = Vec::new();
transport
.receive(
Bytes::from(packet.marshal().unwrap()),
"127.0.0.1:5000".parse().unwrap(),
&mut marshal_buf,
)
.await;
tokio::time::timeout(tokio::time::Duration::from_millis(50), audio_rx.recv())
.await
.expect_err("ambiguous packet should not be routed to audio");
tokio::time::timeout(tokio::time::Duration::from_millis(50), video_rx.recv())
.await
.expect_err("ambiguous packet should not be routed to video");
assert!(!transport.has_listener(4444));
}
#[tokio::test]
async fn test_mid_routes_and_binds_ssrc_when_payload_type_is_ambiguous() {
use crate::transports::ice::IceSocketWrapper;
use bytes::Bytes;
use tokio::sync::watch;
let (_ice_tx, ice_rx) = watch::channel(None::<IceSocketWrapper>);
let ice_conn = IceConn::new(ice_rx, "127.0.0.1:1234".parse().unwrap(), None);
let transport = RtpTransport::new(ice_conn, false);
transport.set_sdes_mid_extension_id(Some(1));
let (audio_tx, mut audio_rx) = mpsc::channel(10);
transport.register_mid_listener("as".to_string(), audio_tx.clone());
transport.register_payload_list_listener(vec![96], audio_tx);
let (video_tx, mut video_rx) = mpsc::channel(10);
transport.register_mid_listener("vs".to_string(), video_tx.clone());
transport.register_payload_list_listener(vec![96], video_tx);
let mut header = crate::rtp::RtpHeader::new(96, 1, 0, 5555);
header.set_extension(1, b"vs").unwrap();
let packet = crate::rtp::RtpPacket::new(header, vec![0u8; 160]);
let mut marshal_buf = Vec::new();
transport
.receive(
Bytes::from(packet.marshal().unwrap()),
"127.0.0.1:5000".parse().unwrap(),
&mut marshal_buf,
)
.await;
let received = video_rx
.recv()
.await
.expect("packet with video MID should route to video");
assert_eq!(received.0.header.ssrc, 5555);
tokio::time::timeout(tokio::time::Duration::from_millis(50), audio_rx.recv())
.await
.expect_err("packet with video MID should not route to audio");
assert!(transport.has_listener(5555));
let header = crate::rtp::RtpHeader::new(96, 2, 160, 5555);
let packet = crate::rtp::RtpPacket::new(header, vec![1u8; 160]);
transport
.receive(
Bytes::from(packet.marshal().unwrap()),
"127.0.0.1:5000".parse().unwrap(),
&mut marshal_buf,
)
.await;
let received = video_rx
.recv()
.await
.expect("bound SSRC should route without MID");
assert_eq!(received.0.header.sequence_number, 2);
}
#[tokio::test]
async fn test_mid_route_overrides_existing_ssrc_mapping() {
use crate::transports::ice::IceSocketWrapper;
use bytes::Bytes;
use tokio::sync::watch;
let (_ice_tx, ice_rx) = watch::channel(None::<IceSocketWrapper>);
let ice_conn = IceConn::new(ice_rx, "127.0.0.1:1234".parse().unwrap(), None);
let transport = RtpTransport::new(ice_conn, false);
transport.set_sdes_mid_extension_id(Some(1));
let (audio_tx, mut audio_rx) = mpsc::channel(10);
transport.register_listener_sync(6666, audio_tx.clone());
transport.register_mid_listener("as".to_string(), audio_tx);
let (video_tx, mut video_rx) = mpsc::channel(10);
transport.register_mid_listener("vs".to_string(), video_tx);
let mut header = crate::rtp::RtpHeader::new(96, 1, 0, 6666);
header.set_extension(1, b"vs").unwrap();
let packet = crate::rtp::RtpPacket::new(header, vec![0u8; 160]);
let mut marshal_buf = Vec::new();
transport
.receive(
Bytes::from(packet.marshal().unwrap()),
"127.0.0.1:5000".parse().unwrap(),
&mut marshal_buf,
)
.await;
let received = video_rx
.recv()
.await
.expect("MID should override stale SSRC mapping");
assert_eq!(received.0.header.ssrc, 6666);
tokio::time::timeout(tokio::time::Duration::from_millis(50), audio_rx.recv())
.await
.expect_err("stale SSRC mapping should not receive the MID packet");
let header = crate::rtp::RtpHeader::new(96, 2, 160, 6666);
let packet = crate::rtp::RtpPacket::new(header, vec![1u8; 160]);
transport
.receive(
Bytes::from(packet.marshal().unwrap()),
"127.0.0.1:5000".parse().unwrap(),
&mut marshal_buf,
)
.await;
let received = video_rx
.recv()
.await
.expect("corrected SSRC mapping should receive packets without MID");
assert_eq!(received.0.header.sequence_number, 2);
}
#[tokio::test]
async fn test_rewrite_bridge_rewrites_packet_fields() {
use crate::transports::ice::IceSocketWrapper;
use tokio::net::UdpSocket;
use tokio::sync::watch;
let src_socket = UdpSocket::bind("127.0.0.1:0").await.unwrap();
let (_src_tx, src_rx) = watch::channel(Some(IceSocketWrapper::Udp(Arc::new(src_socket))));
let src_conn = IceConn::new(src_rx, "127.0.0.1:9".parse().unwrap(), None);
let src_transport = RtpTransport::new(src_conn, false);
let dst_socket = UdpSocket::bind("127.0.0.1:0").await.unwrap();
let (_dst_tx, dst_rx) = watch::channel(Some(IceSocketWrapper::Udp(Arc::new(dst_socket))));
let dst_conn = IceConn::new(dst_rx, "127.0.0.1:9".parse().unwrap(), None);
let dst_transport = Arc::new(RtpTransport::new(dst_conn, false));
src_transport.bridge_rewrite_to(
dst_transport.clone(),
RtpRewriteBridgeParams {
ssrc_offset: 900,
fixed_out_ssrc: None,
payload_type: Some(96),
dtmf_payload_type: None,
initial_sequence_number: Some(32000),
initial_timestamp_offset: Some(12345),
strip_extensions: false,
},
);
let mut guard = src_transport.rewrite_bridge.lock();
let bridge = guard.as_mut().expect("rewrite bridge should be configured");
let mut packet = RtpPacket::new(crate::rtp::RtpHeader::new(0, 7, 1111, 100), vec![1u8; 32]);
bridge.rewrite_packet(&mut packet);
drop(guard);
assert_eq!(packet.header.ssrc, 1000);
assert_eq!(packet.header.payload_type, 96);
assert_eq!(packet.header.sequence_number, 32000);
assert_eq!(packet.header.timestamp, 1111 + 12345);
}
#[tokio::test]
async fn test_rewrite_packet_remaps_dtmf_payload_type() {
use crate::transports::ice::IceSocketWrapper;
use tokio::net::UdpSocket;
use tokio::sync::watch;
let socket = UdpSocket::bind("127.0.0.1:0").await.unwrap();
let (_tx, rx) = watch::channel(Some(IceSocketWrapper::Udp(Arc::new(socket))));
let src_conn = IceConn::new(rx, "127.0.0.1:9".parse().unwrap(), None);
let src_transport = Arc::new(RtpTransport::new(src_conn, false));
let dst_socket = UdpSocket::bind("127.0.0.1:0").await.unwrap();
let (_dst_tx, dst_rx) = watch::channel(Some(IceSocketWrapper::Udp(Arc::new(dst_socket))));
let dst_conn = IceConn::new(dst_rx, "127.0.0.1:9".parse().unwrap(), None);
let dst_transport = Arc::new(RtpTransport::new(dst_conn, false));
src_transport.bridge_rewrite_to(
dst_transport.clone(),
RtpRewriteBridgeParams {
ssrc_offset: 900,
payload_type: Some(96),
dtmf_payload_type: Some((101, 110)),
initial_sequence_number: Some(32000),
initial_timestamp_offset: Some(12345),
fixed_out_ssrc: None,
strip_extensions: false,
},
);
let mut guard = src_transport.rewrite_bridge.lock();
let bridge = guard.as_mut().expect("rewrite bridge should be configured");
let mut audio =
RtpPacket::new(crate::rtp::RtpHeader::new(100, 7, 1111, 1111), vec![1u8; 32]);
bridge.rewrite_packet(&mut audio);
assert_eq!(audio.header.payload_type, 96);
let mut dtmf =
RtpPacket::new(crate::rtp::RtpHeader::new(101, 7, 1111, 1111), vec![1u8; 32]);
bridge.rewrite_packet(&mut dtmf);
assert_eq!(dtmf.header.payload_type, 110);
drop(guard);
}
#[tokio::test]
async fn test_rewrite_rules_route_audio_video_dtmf_to_distinct_ssrc_pt() {
use crate::transports::ice::IceSocketWrapper;
use tokio::net::UdpSocket;
use tokio::sync::watch;
let socket = UdpSocket::bind("127.0.0.1:0").await.unwrap();
let (_tx, rx) = watch::channel(Some(IceSocketWrapper::Udp(Arc::new(socket))));
let src_conn = IceConn::new(rx, "127.0.0.1:9".parse().unwrap(), None);
let src_transport = Arc::new(RtpTransport::new(src_conn, false));
let dst_socket = UdpSocket::bind("127.0.0.1:0").await.unwrap();
let (_dst_tx, dst_rx) = watch::channel(Some(IceSocketWrapper::Udp(Arc::new(dst_socket))));
let dst_conn = IceConn::new(dst_rx, "127.0.0.1:9".parse().unwrap(), None);
let dst_transport = Arc::new(RtpTransport::new(dst_conn, false));
let audio_ssrc = 111u32;
let video_ssrc = 222u32;
let rules = vec![
RtpRewriteRule {
match_payload_type: None,
fixed_out_ssrc: Some(audio_ssrc),
ssrc_offset: 0,
out_payload_type: Some(96),
},
RtpRewriteRule {
match_payload_type: Some(101),
fixed_out_ssrc: Some(audio_ssrc),
ssrc_offset: 0,
out_payload_type: Some(110),
},
RtpRewriteRule {
match_payload_type: Some(98),
fixed_out_ssrc: Some(video_ssrc),
ssrc_offset: 0,
out_payload_type: Some(102),
},
RtpRewriteRule {
match_payload_type: Some(99),
fixed_out_ssrc: Some(video_ssrc),
ssrc_offset: 0,
out_payload_type: Some(103),
},
];
src_transport.bridge_rewrite_rules_to(dst_transport.clone(), Default::default(), rules);
let mut guard = src_transport.rewrite_bridge.lock();
let bridge = guard.as_mut().expect("rewrite bridge should be configured");
let mut audio =
RtpPacket::new(crate::rtp::RtpHeader::new(97, 7, 1111, 1111), vec![1u8; 32]);
bridge.rewrite_packet(&mut audio);
assert_eq!(audio.header.ssrc, audio_ssrc);
assert_eq!(audio.header.payload_type, 96);
let mut dtmf =
RtpPacket::new(crate::rtp::RtpHeader::new(101, 7, 1111, 1111), vec![1u8; 32]);
bridge.rewrite_packet(&mut dtmf);
assert_eq!(dtmf.header.ssrc, audio_ssrc);
assert_eq!(dtmf.header.payload_type, 110);
let mut video =
RtpPacket::new(crate::rtp::RtpHeader::new(98, 7, 2222, 2222), vec![1u8; 32]);
bridge.rewrite_packet(&mut video);
assert_eq!(video.header.ssrc, video_ssrc);
assert_eq!(video.header.payload_type, 102);
let mut rtx =
RtpPacket::new(crate::rtp::RtpHeader::new(99, 7, 3333, 3333), vec![1u8; 32]);
bridge.rewrite_packet(&mut rtx);
assert_eq!(rtx.header.ssrc, video_ssrc);
assert_eq!(rtx.header.payload_type, 103);
let mut video2 =
RtpPacket::new(crate::rtp::RtpHeader::new(98, 8, 2382, 2222), vec![1u8; 32]);
bridge.rewrite_packet(&mut video2);
assert_eq!(video2.header.sequence_number, video.header.sequence_number + 1);
assert_eq!(video2.header.timestamp, video.header.timestamp.wrapping_add(160));
drop(guard);
}
#[tokio::test]
async fn test_rewrite_rules_unmatched_packet_passes_ssrc_through() {
use crate::transports::ice::IceSocketWrapper;
use tokio::net::UdpSocket;
use tokio::sync::watch;
let socket = UdpSocket::bind("127.0.0.1:0").await.unwrap();
let (_tx, rx) = watch::channel(Some(IceSocketWrapper::Udp(Arc::new(socket))));
let src_conn = IceConn::new(rx, "127.0.0.1:9".parse().unwrap(), None);
let src_transport = Arc::new(RtpTransport::new(src_conn, false));
let dst_socket = UdpSocket::bind("127.0.0.1:0").await.unwrap();
let (_dst_tx, dst_rx) = watch::channel(Some(IceSocketWrapper::Udp(Arc::new(dst_socket))));
let dst_conn = IceConn::new(dst_rx, "127.0.0.1:9".parse().unwrap(), None);
let dst_transport = Arc::new(RtpTransport::new(dst_conn, false));
let rules = vec![RtpRewriteRule {
match_payload_type: Some(98),
fixed_out_ssrc: Some(222),
ssrc_offset: 0,
out_payload_type: Some(102),
}];
src_transport.bridge_rewrite_rules_to(dst_transport.clone(), Default::default(), rules);
let mut guard = src_transport.rewrite_bridge.lock();
let bridge = guard.as_mut().expect("rewrite bridge should be configured");
let mut video =
RtpPacket::new(crate::rtp::RtpHeader::new(98, 7, 2222, 2222), vec![1u8; 32]);
bridge.rewrite_packet(&mut video);
assert_eq!(video.header.ssrc, 222);
assert_eq!(video.header.payload_type, 102);
let mut unmatched =
RtpPacket::new(crate::rtp::RtpHeader::new(97, 7, 1111, 1111), vec![1u8; 32]);
bridge.rewrite_packet(&mut unmatched);
assert_eq!(unmatched.header.ssrc, 1111);
assert_eq!(unmatched.header.payload_type, 97);
drop(guard);
}
#[tokio::test]
async fn test_rewrite_rules_legacy_params_equivalent_to_single_rule() {
use crate::transports::ice::IceSocketWrapper;
use tokio::net::UdpSocket;
use tokio::sync::watch;
let socket = UdpSocket::bind("127.0.0.1:0").await.unwrap();
let (_tx, rx) = watch::channel(Some(IceSocketWrapper::Udp(Arc::new(socket))));
let src_conn = IceConn::new(rx, "127.0.0.1:9".parse().unwrap(), None);
let src_transport = Arc::new(RtpTransport::new(src_conn, false));
let dst_socket = UdpSocket::bind("127.0.0.1:0").await.unwrap();
let (_dst_tx, dst_rx) = watch::channel(Some(IceSocketWrapper::Udp(Arc::new(dst_socket))));
let dst_conn = IceConn::new(dst_rx, "127.0.0.1:9".parse().unwrap(), None);
let dst_transport = Arc::new(RtpTransport::new(dst_conn, false));
src_transport.bridge_rewrite_to(
dst_transport.clone(),
RtpRewriteBridgeParams {
ssrc_offset: 900,
fixed_out_ssrc: None,
payload_type: Some(96),
dtmf_payload_type: Some((101, 110)),
initial_sequence_number: Some(32000),
initial_timestamp_offset: Some(12345),
strip_extensions: false,
},
);
let mut guard = src_transport.rewrite_bridge.lock();
let bridge = guard.as_mut().expect("rewrite bridge should be configured");
let mut audio =
RtpPacket::new(crate::rtp::RtpHeader::new(100, 7, 1111, 1111), vec![1u8; 32]);
bridge.rewrite_packet(&mut audio);
assert_eq!(audio.header.ssrc, 1111 + 900);
assert_eq!(audio.header.payload_type, 96);
assert_eq!(audio.header.sequence_number, 32000);
assert_eq!(audio.header.timestamp, 1111 + 12345);
let mut dtmf =
RtpPacket::new(crate::rtp::RtpHeader::new(101, 7, 1111, 1111), vec![1u8; 32]);
bridge.rewrite_packet(&mut dtmf);
assert_eq!(dtmf.header.ssrc, 1111 + 900);
assert_eq!(dtmf.header.payload_type, 110);
drop(guard);
}
#[tokio::test]
async fn test_received_rtp_packets_counter_advances_on_slow_path() {
use crate::transports::ice::IceSocketWrapper;
use tokio::net::UdpSocket;
use tokio::sync::watch;
let socket = UdpSocket::bind("127.0.0.1:0").await.unwrap();
let (_tx, rx) = watch::channel(Some(IceSocketWrapper::Udp(Arc::new(socket))));
let conn = IceConn::new(rx, "127.0.0.1:9".parse().unwrap(), None);
let transport = RtpTransport::new(conn, false);
let mut marshal_buf = Vec::with_capacity(1500);
let addr: SocketAddr = "127.0.0.1:5000".parse().unwrap();
assert_eq!(transport.received_rtp_packets(), 0, "counter starts at zero");
for seq in 1..=3u16 {
let header = crate::rtp::RtpHeader::new(0, seq, 160, 1234);
let packet = crate::rtp::RtpPacket::new(header, vec![1u8; 160]);
transport
.receive(Bytes::from(packet.marshal().unwrap()), addr, &mut marshal_buf)
.await;
}
assert_eq!(
transport.received_rtp_packets(),
3,
"counter must advance by one per accepted inbound RTP packet"
);
}
#[tokio::test]
async fn test_received_rtp_packets_counter_advances_on_fast_path_relay() {
use crate::transports::ice::IceSocketWrapper;
use tokio::net::UdpSocket;
use tokio::sync::watch;
let src_socket = UdpSocket::bind("127.0.0.1:0").await.unwrap();
let (_src_tx, src_rx) = watch::channel(Some(IceSocketWrapper::Udp(Arc::new(src_socket))));
let src_conn = IceConn::new(src_rx, "127.0.0.1:9".parse().unwrap(), None);
let src_transport = Arc::new(RtpTransport::new(src_conn, false));
let ssrc = 4242u32;
let (listener_tx, mut listener_rx) = mpsc::channel::<(RtpPacket, SocketAddr)>(8);
src_transport.register_listener_sync(ssrc, listener_tx);
let dst_socket = UdpSocket::bind("127.0.0.1:0").await.unwrap();
let (_dst_tx, dst_rx) = watch::channel(Some(IceSocketWrapper::Udp(Arc::new(dst_socket))));
let dst_conn = IceConn::new(dst_rx, "127.0.0.1:9".parse().unwrap(), None);
let dst_transport = Arc::new(RtpTransport::new(dst_conn, false));
src_transport.bridge_rewrite_to(
dst_transport.clone(),
RtpRewriteBridgeParams {
ssrc_offset: 0,
fixed_out_ssrc: None,
payload_type: None,
dtmf_payload_type: None,
initial_sequence_number: None,
initial_timestamp_offset: None,
strip_extensions: false,
},
);
assert!(src_transport.has_bridge.load(Ordering::SeqCst));
assert_eq!(src_transport.received_rtp_packets(), 0);
assert_eq!(dst_transport.received_rtp_packets(), 0);
let mut marshal_buf = Vec::with_capacity(1500);
let addr: SocketAddr = "127.0.0.1:5000".parse().unwrap();
for seq in 1..=2u16 {
let header = crate::rtp::RtpHeader::new(0, seq, 160, ssrc);
let packet = crate::rtp::RtpPacket::new(header, vec![1u8; 160]);
src_transport
.receive(
Bytes::from(packet.marshal().unwrap()),
addr,
&mut marshal_buf,
)
.await;
}
assert_eq!(
src_transport.received_rtp_packets(),
2,
"source counter must advance on the fast-path relay"
);
assert_eq!(
dst_transport.received_rtp_packets(),
0,
"relayed packet must not be counted as inbound on the destination"
);
let attempt = tokio::time::timeout(
std::time::Duration::from_millis(150),
listener_rx.recv(),
)
.await;
assert!(
attempt.is_err(),
"listener must NOT receive on the fast-path relay (interceptor path is bypassed)"
);
}
#[tokio::test]
async fn test_routes_pruned_on_closed_tx() {
let (tx1, rx1) = mpsc::channel::<(RtpPacket, SocketAddr)>(8);
let (tx2, rx2) = mpsc::channel::<(RtpPacket, SocketAddr)>(8);
let mut reg = super::ListenerRegistry::default();
reg.register_mid("stream1".into(), tx1.clone());
reg.register_provisional(tx2.clone());
assert_eq!(reg.routes.len(), 2, "two live routes");
drop(rx1);
drop(rx2);
let (tx3, _rx3) = mpsc::channel::<(RtpPacket, SocketAddr)>(8);
reg.register_provisional(tx3);
assert!(
reg.routes.len() <= 1,
"Fix3: routes must be pruned when their tx is closed (len={})",
reg.routes.len()
);
let (tx4, _rx4) = mpsc::channel::<(RtpPacket, SocketAddr)>(8);
let (tx5, _rx5) = mpsc::channel::<(RtpPacket, SocketAddr)>(8);
reg.bind_ssrc_route(111, tx4);
reg.bind_ssrc_route(222, tx5);
assert!(
reg.by_ssrc.len() <= 2,
"Fix3: by_ssrc must not hold stale entries"
);
}
#[tokio::test]
async fn test_routes_do_not_grow_across_transport_replacement() {
use crate::transports::ice::IceSocketWrapper;
use tokio::sync::watch;
let (_tx, rx) = watch::channel::<Option<IceSocketWrapper>>(None);
let conn = crate::transports::ice::conn::IceConn::new(
rx,
"127.0.0.1:0".parse().unwrap(),
None,
);
let transport = Arc::new(super::RtpTransport::new(conn, false));
let ssrc = 1001u32;
for i in 0..5 {
transport.register_listener_sync(ssrc + i, mpsc::channel::<(RtpPacket, SocketAddr)>(8).0);
}
let listeners = transport.listeners.lock();
assert!(
listeners.routes.len() <= 2,
"Fix3: routes must not grow across transport-replacement-like \
register_listener_sync calls (len={})",
listeners.routes.len()
);
assert!(
listeners.by_ssrc.len() <= 3,
"Fix3: by_ssrc must not accumulate stale entries (len={})",
listeners.by_ssrc.len()
);
}
struct CountingObserver {
ingress: std::sync::atomic::AtomicU32,
egress: std::sync::atomic::AtomicU32,
last_pt: std::sync::atomic::AtomicU8,
}
impl crate::peer_connection::RtpObserver for CountingObserver {
fn on_ingress(&self, packet: &RtpPacket, _src_addr: SocketAddr) {
self.ingress
.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
self.last_pt.store(
packet.header.payload_type,
std::sync::atomic::Ordering::Relaxed,
);
}
fn on_egress(&self, packet: &RtpPacket, _dst_addr: SocketAddr) {
self.egress
.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
self.last_pt.store(
packet.header.payload_type,
std::sync::atomic::Ordering::Relaxed,
);
}
}
#[tokio::test]
async fn test_observer_ingress_fires_on_relay_fast_path() {
use crate::transports::ice::IceSocketWrapper;
use tokio::net::UdpSocket;
use tokio::sync::watch;
let src_socket = UdpSocket::bind("127.0.0.1:0").await.unwrap();
let (_src_tx, src_rx) = watch::channel(Some(IceSocketWrapper::Udp(Arc::new(src_socket))));
let src_conn = IceConn::new(src_rx, "127.0.0.1:9".parse().unwrap(), None);
let src_transport = Arc::new(RtpTransport::new(src_conn, false));
let dst_socket = UdpSocket::bind("127.0.0.1:0").await.unwrap();
let (_dst_tx, dst_rx) = watch::channel(Some(IceSocketWrapper::Udp(Arc::new(dst_socket))));
let dst_conn = IceConn::new(dst_rx, "127.0.0.1:9".parse().unwrap(), None);
let dst_transport = Arc::new(RtpTransport::new(dst_conn, false));
src_transport.bridge_rewrite_to(
dst_transport.clone(),
RtpRewriteBridgeParams {
ssrc_offset: 0,
fixed_out_ssrc: None,
payload_type: None,
dtmf_payload_type: None,
initial_sequence_number: None,
initial_timestamp_offset: None,
strip_extensions: false,
},
);
assert!(src_transport.has_bridge.load(Ordering::SeqCst));
let tap = Arc::new(CountingObserver {
ingress: std::sync::atomic::AtomicU32::new(0),
egress: std::sync::atomic::AtomicU32::new(0),
last_pt: std::sync::atomic::AtomicU8::new(0),
});
src_transport.add_observer(tap.clone());
let mut marshal_buf = Vec::with_capacity(1500);
let addr: SocketAddr = "127.0.0.1:5000".parse().unwrap();
for seq in 1..=3u16 {
let header = crate::rtp::RtpHeader::new(8, seq, 160, 4242);
let packet = crate::rtp::RtpPacket::new(header, vec![1u8; 160]);
src_transport
.receive(
Bytes::from(packet.marshal().unwrap()),
addr,
&mut marshal_buf,
)
.await;
}
assert_eq!(
tap.ingress.load(std::sync::atomic::Ordering::Relaxed),
3,
"ingress observer must fire on the relay fast-path"
);
assert_eq!(
tap.last_pt.load(std::sync::atomic::Ordering::Relaxed),
8,
"ingress observer observed the packet's payload type"
);
}
#[tokio::test]
async fn test_observer_egress_fires_on_relay() {
use crate::transports::ice::IceSocketWrapper;
use tokio::net::UdpSocket;
use tokio::sync::watch;
let src_socket = UdpSocket::bind("127.0.0.1:0").await.unwrap();
let (_src_tx, src_rx) = watch::channel(Some(IceSocketWrapper::Udp(Arc::new(src_socket))));
let src_conn = IceConn::new(src_rx, "127.0.0.1:9".parse().unwrap(), None);
let src_transport = Arc::new(RtpTransport::new(src_conn, false));
let dst_socket = UdpSocket::bind("127.0.0.1:0").await.unwrap();
let (_dst_tx, dst_rx) = watch::channel(Some(IceSocketWrapper::Udp(Arc::new(dst_socket))));
let dst_conn = IceConn::new(dst_rx, "127.0.0.1:9".parse().unwrap(), None);
let dst_transport = Arc::new(RtpTransport::new(dst_conn, false));
src_transport.bridge_rewrite_to(
dst_transport.clone(),
RtpRewriteBridgeParams {
ssrc_offset: 0,
fixed_out_ssrc: None,
payload_type: Some(96),
dtmf_payload_type: None,
initial_sequence_number: None,
initial_timestamp_offset: None,
strip_extensions: false,
},
);
let tap = Arc::new(CountingObserver {
ingress: std::sync::atomic::AtomicU32::new(0),
egress: std::sync::atomic::AtomicU32::new(0),
last_pt: std::sync::atomic::AtomicU8::new(0),
});
dst_transport.add_observer(tap.clone());
let mut marshal_buf = Vec::with_capacity(1500);
let addr: SocketAddr = "127.0.0.1:5000".parse().unwrap();
for seq in 1..=3u16 {
let header = crate::rtp::RtpHeader::new(0, seq, 160, 4242);
let packet = crate::rtp::RtpPacket::new(header, vec![1u8; 160]);
src_transport
.receive(
Bytes::from(packet.marshal().unwrap()),
addr,
&mut marshal_buf,
)
.await;
}
assert_eq!(
tap.egress.load(std::sync::atomic::Ordering::Relaxed),
3,
"egress observer on destination must fire for relayed packets"
);
assert_eq!(
tap.last_pt.load(std::sync::atomic::Ordering::Relaxed),
96,
"egress observer saw the rewritten payload type"
);
assert_eq!(
tap.ingress.load(std::sync::atomic::Ordering::Relaxed),
0,
"ingress on destination must NOT fire (relay bypasses dst receive)"
);
}
#[tokio::test]
async fn test_observer_egress_fires_on_normal_send() {
use crate::transports::ice::IceSocketWrapper;
use tokio::net::UdpSocket;
use tokio::sync::watch;
let socket = UdpSocket::bind("127.0.0.1:0").await.unwrap();
let (_tx, rx) = watch::channel(Some(IceSocketWrapper::Udp(Arc::new(socket))));
let conn = IceConn::new(rx, "127.0.0.1:9".parse().unwrap(), None);
let transport = Arc::new(RtpTransport::new(conn, false));
let tap = Arc::new(CountingObserver {
ingress: std::sync::atomic::AtomicU32::new(0),
egress: std::sync::atomic::AtomicU32::new(0),
last_pt: std::sync::atomic::AtomicU8::new(0),
});
transport.add_observer(tap.clone());
for seq in 1..=2u16 {
let header = crate::rtp::RtpHeader::new(0, seq, 160, 7777);
let packet = crate::rtp::RtpPacket::new(header, vec![1u8; 160]);
transport.send_rtp(packet).await.unwrap();
}
assert_eq!(
tap.egress.load(std::sync::atomic::Ordering::Relaxed),
2,
"egress observer must fire on normal send_rtp"
);
assert_eq!(
tap.last_pt.load(std::sync::atomic::Ordering::Relaxed),
0,
"egress observer saw the sent payload type"
);
}
#[tokio::test]
async fn test_no_observer_zero_cost_on_relay() {
use crate::transports::ice::IceSocketWrapper;
use tokio::net::UdpSocket;
use tokio::sync::watch;
let src_socket = UdpSocket::bind("127.0.0.1:0").await.unwrap();
let (_src_tx, src_rx) = watch::channel(Some(IceSocketWrapper::Udp(Arc::new(src_socket))));
let src_conn = IceConn::new(src_rx, "127.0.0.1:9".parse().unwrap(), None);
let src_transport = Arc::new(RtpTransport::new(src_conn, false));
let dst_socket = UdpSocket::bind("127.0.0.1:0").await.unwrap();
let (_dst_tx, dst_rx) = watch::channel(Some(IceSocketWrapper::Udp(Arc::new(dst_socket))));
let dst_conn = IceConn::new(dst_rx, "127.0.0.1:9".parse().unwrap(), None);
let dst_transport = Arc::new(RtpTransport::new(dst_conn, false));
src_transport.bridge_rewrite_to(
dst_transport,
RtpRewriteBridgeParams {
ssrc_offset: 0,
fixed_out_ssrc: None,
payload_type: None,
dtmf_payload_type: None,
initial_sequence_number: None,
initial_timestamp_offset: None,
strip_extensions: false,
},
);
assert!(
!src_transport
.has_observers
.load(std::sync::atomic::Ordering::SeqCst),
"flag must be false when no observer registered"
);
let mut marshal_buf = Vec::with_capacity(1500);
let addr: SocketAddr = "127.0.0.1:5000".parse().unwrap();
let header = crate::rtp::RtpHeader::new(0, 1, 160, 99);
let packet = crate::rtp::RtpPacket::new(header, vec![1u8; 160]);
src_transport
.receive(
Bytes::from(packet.marshal().unwrap()),
addr,
&mut marshal_buf,
)
.await;
assert_eq!(src_transport.received_rtp_packets(), 1);
}
}