use std::collections::HashMap;
use std::net::{IpAddr, Ipv4Addr, Ipv6Addr, SocketAddr};
use std::sync::Arc;
use std::time::{Duration, Instant, SystemTime, UNIX_EPOCH};
use str0m::{Candidate, Event, IceConnectionState, Input, Output, Rtc, net::Receive};
use tokio::net::UdpSocket;
use tokio::sync::mpsc;
use crate::egress::{EgressClock, EgressSource, WriteRequest};
use crate::{Error, Result, codec};
pub(crate) type Packet = (Vec<u8>, SocketAddr);
pub(crate) const SESSION_INBOX: usize = 256;
const EGRESS_SEND_BUFFER_VIDEO: usize = 3000;
const ICE_ESTABLISH_TIMEOUT: Duration = Duration::from_secs(30);
pub trait MediaSink: Send {
fn on_track(
&mut self,
mid: str0m::media::Mid,
kind: str0m::media::MediaKind,
codec: str0m::format::Codec,
audio_params: Option<(u32, u32)>,
) -> Result<()>;
fn on_frame(&mut self, mid: str0m::media::Mid, frame: codec::Frame) -> Result<()>;
fn abort(&mut self, err: moq_net::Error);
}
#[non_exhaustive]
pub enum MediaRole {
Ingest(Box<dyn MediaSink>),
Egress(Box<EgressSource>),
}
pub struct Session {
rtc: Rtc,
socket: Arc<UdpSocket>,
locals: Vec<SocketAddr>,
inbound: mpsc::Receiver<Packet>,
role: MediaRole,
writes_rx: Option<mpsc::Receiver<WriteRequest>>,
ingest_clock: IngestClock,
egress_clock: EgressClock,
}
impl Session {
pub fn ingest(
rtc: Rtc,
socket: Arc<UdpSocket>,
locals: Vec<SocketAddr>,
inbound: mpsc::Receiver<Packet>,
sink: Box<dyn MediaSink>,
) -> Self {
Self {
rtc,
socket,
locals,
inbound,
role: MediaRole::Ingest(sink),
writes_rx: None,
ingest_clock: IngestClock::default(),
egress_clock: EgressClock::default(),
}
}
pub fn egress(
rtc: Rtc,
socket: Arc<UdpSocket>,
locals: Vec<SocketAddr>,
inbound: mpsc::Receiver<Packet>,
mut source: EgressSource,
) -> Self {
let writes_rx = source.take_writes();
Self {
rtc,
socket,
locals,
inbound,
role: MediaRole::Egress(Box::new(source)),
writes_rx: Some(writes_rx),
ingest_clock: IngestClock::default(),
egress_clock: EgressClock::default(),
}
}
pub async fn run(mut self) -> Result<()> {
let result = self.run_loop().await;
if let Err(err) = &result
&& !matches!(err, Error::SessionClosed | Error::IceTimeout)
&& let MediaRole::Ingest(sink) = &mut self.role
{
sink.abort(moq_net::Error::Transport(err.to_string()));
}
result
}
async fn run_loop(&mut self) -> Result<()> {
let started = Instant::now();
let mut connected = false;
let socket_v6 = self.socket.local_addr().map_err(Error::Io)?.is_ipv6();
loop {
if !self.rtc.is_alive() {
return Err(Error::SessionClosed);
}
if !connected && started.elapsed() >= ICE_ESTABLISH_TIMEOUT {
return Err(Error::IceTimeout);
}
let timeout = match self.rtc.poll_output().map_err(Error::Rtc)? {
Output::Timeout(t) => t,
Output::Transmit(t) => {
let dst = crate::net::to_family(t.destination, socket_v6);
if let Err(err) = self.socket.send_to(&t.contents, dst).await {
tracing::warn!(%err, %dst, "send failed");
}
continue;
}
Output::Event(event) => {
if let Event::IceConnectionStateChange(state) = &event {
connected |= state.is_connected();
}
self.handle_event(event)?;
continue;
}
};
let now = Instant::now();
let mut duration = timeout.saturating_duration_since(now);
if !connected {
duration = duration.min(ICE_ESTABLISH_TIMEOUT.saturating_sub(started.elapsed()));
}
if duration.is_zero() {
self.rtc.handle_input(Input::Timeout(now)).map_err(Error::Rtc)?;
continue;
}
tokio::select! {
biased;
Some(req) = async {
match self.writes_rx.as_mut() {
Some(rx) => rx.recv().await,
None => std::future::pending::<Option<WriteRequest>>().await,
}
} => {
let now = Instant::now();
let wallclock = self.egress_clock.wallclock(req.time, now);
crate::egress::dispatch(&mut self.rtc, req, wallclock);
}
packet = self.inbound.recv() => {
match packet {
Some((data, src)) => {
let now = Instant::now();
let local = pick_local(&self.locals, src);
let recv = Receive::new(str0m::net::Protocol::Udp, src, local, &data)
.map_err(Error::RtcInput)?;
self.rtc.handle_input(Input::Receive(now, recv)).map_err(Error::Rtc)?;
}
None => return Err(Error::SessionClosed),
}
}
_ = tokio::time::sleep(duration) => {
self.rtc
.handle_input(Input::Timeout(Instant::now()))
.map_err(Error::Rtc)?;
}
}
}
}
fn handle_event(&mut self, event: Event) -> Result<()> {
match event {
Event::IceConnectionStateChange(state) => {
tracing::debug!(?state, "ice state");
if state == IceConnectionState::Disconnected {
return Err(Error::SessionClosed);
}
}
Event::MediaAdded(added) => self.handle_media_added(added)?,
Event::MediaData(data) => {
if let MediaRole::Ingest(sink) = &mut self.role {
let media_us = media_time_to_micros(&data.time);
let timestamp_us = self.ingest_clock.normalize(data.mid, data.network_time, media_us);
sink.on_frame(
data.mid,
codec::Frame {
timestamp_us,
payload: bytes::Bytes::from_owner(data.data),
},
)?;
}
}
Event::SenderFeedback(feedback) => {
if matches!(&self.role, MediaRole::Ingest(_)) {
self.ingest_clock.observe(feedback.mid, feedback.sender_info);
}
}
Event::KeyframeRequest(req) => {
tracing::debug!(?req, "keyframe request from peer");
}
_ => {}
}
Ok(())
}
fn handle_media_added(&mut self, added: str0m::media::MediaAdded) -> Result<()> {
let pt = self.rtc.media(added.mid).and_then(|m| m.remote_pts().first().copied());
let params = pt.and_then(|pt| self.rtc.codec_config().params().iter().find(|p| p.pt() == pt).copied());
let params = match params {
Some(p) => p,
None => {
tracing::warn!(?added.mid, "no codec params for media; ignoring");
return Ok(());
}
};
let spec = params.spec();
let codec = spec.codec;
match &mut self.role {
MediaRole::Ingest(sink) => {
let audio_params = if codec.is_audio() {
Some((spec.clock_rate.get(), spec.channels.unwrap_or(1) as u32))
} else {
None
};
sink.on_track(added.mid, added.kind, codec, audio_params)?;
}
MediaRole::Egress(source) => {
source.on_track(added.mid, codec, params.pt(), spec.clock_rate)?;
}
}
Ok(())
}
}
#[derive(Default)]
pub(crate) struct IngestClock {
arrival_epoch: Option<Instant>,
ntp_epoch_us: Option<i128>,
tracks: HashMap<str0m::media::Mid, IngestTrackClock>,
}
impl IngestClock {
fn observe(&mut self, mid: str0m::media::Mid, sender: str0m::rtp::rtcp::SenderInfo) {
self.tracks.entry(mid).or_default().sender = Some(SenderAnchor::new(sender));
self.establish_ntp_epoch();
}
fn normalize(&mut self, mid: str0m::media::Mid, arrival: Instant, media_us: u64) -> u64 {
let epoch = *self.arrival_epoch.get_or_insert(arrival);
let track = self.tracks.entry(mid).or_default();
let offset = *track.arrival_offset_us.get_or_insert_with(|| {
let wall_us = if arrival >= epoch {
arrival.duration_since(epoch).as_micros() as i64
} else {
-(epoch.duration_since(arrival).as_micros() as i64)
};
wall_us as i128 - media_us as i128
});
let fallback = to_u64(media_us as i128 + offset);
let previous = track.last_output_us;
track.last_media_us = Some(media_us);
track.last_output_us = Some(fallback);
self.establish_ntp_epoch();
let mapped = self
.ntp_epoch_us
.zip(self.tracks.get(&mid).and_then(|track| track.sender))
.map(|(epoch, sender)| to_u64(sender.capture_time_us(media_us) - epoch));
let output = match mapped {
Some(mapped) => mapped.max(previous.map_or(fallback, |last| last.saturating_add(1))),
None => fallback,
};
self.tracks.get_mut(&mid).expect("track was inserted").last_output_us = Some(output);
output
}
fn establish_ntp_epoch(&mut self) {
if self.ntp_epoch_us.is_some() || self.tracks.len() < 2 {
return;
}
let mut epoch = i128::MAX;
for track in self.tracks.values() {
let (Some(sender), Some(media_us), Some(output_us)) =
(track.sender, track.last_media_us, track.last_output_us)
else {
return;
};
epoch = epoch.min(sender.capture_time_us(media_us) - output_us as i128);
}
self.ntp_epoch_us = Some(epoch);
}
}
#[derive(Default)]
struct IngestTrackClock {
arrival_offset_us: Option<i128>,
sender: Option<SenderAnchor>,
last_media_us: Option<u64>,
last_output_us: Option<u64>,
}
#[derive(Clone, Copy)]
struct SenderAnchor {
ntp_us: i128,
rtp_us: i128,
}
impl SenderAnchor {
fn new(sender: str0m::rtp::rtcp::SenderInfo) -> Self {
Self {
ntp_us: system_time_to_micros(sender.ntp_time),
rtp_us: media_time_to_micros(&sender.rtp_time) as i128,
}
}
fn capture_time_us(self, media_us: u64) -> i128 {
self.ntp_us + media_us as i128 - self.rtp_us
}
}
fn system_time_to_micros(time: SystemTime) -> i128 {
match time.duration_since(UNIX_EPOCH) {
Ok(duration) => duration.as_micros() as i128,
Err(err) => -(err.duration().as_micros() as i128),
}
}
fn to_u64(value: i128) -> u64 {
value.clamp(0, u64::MAX as i128) as u64
}
pub(crate) fn log_session_end(role: &str, result: &Result<()>) {
match result {
Ok(()) | Err(Error::SessionClosed) => tracing::debug!(role, "session ended"),
Err(Error::IceTimeout) => tracing::debug!(role, "session ended: ICE never connected"),
Err(err) => tracing::warn!(%err, role, "session ended"),
}
}
fn pick_local(locals: &[SocketAddr], src: SocketAddr) -> SocketAddr {
locals
.iter()
.find(|l| l.is_ipv4() == src.is_ipv4())
.copied()
.unwrap_or(locals[0])
}
fn media_time_to_micros(time: &str0m::media::MediaTime) -> u64 {
let numer = time.numer() as i128;
let denom = time.denom() as i128;
if denom == 0 {
return 0;
}
let micros = (numer.saturating_mul(1_000_000)) / denom;
micros.max(0) as u64
}
pub(crate) struct Bridges {
inner: HashMap<str0m::media::Mid, Box<dyn codec::Bridge>>,
}
impl Bridges {
pub fn new() -> Self {
Self { inner: HashMap::new() }
}
pub fn insert(&mut self, mid: str0m::media::Mid, bridge: Box<dyn codec::Bridge>) {
self.inner.insert(mid, bridge);
}
pub fn push(&mut self, mid: str0m::media::Mid, frame: codec::Frame) -> Result<()> {
if let Some(bridge) = self.inner.get_mut(&mid) {
bridge.push(frame)?;
}
Ok(())
}
pub fn abort(&mut self, err: moq_net::Error) {
for bridge in std::mem::take(&mut self.inner).into_values() {
bridge.abort(err.clone());
}
}
}
pub fn rtc_config_with_codecs(codecs: &[str0m::format::Codec]) -> str0m::RtcConfig {
use str0m::format::Codec;
let mut config = str0m::RtcConfig::new()
.clear_codecs()
.set_send_buffer_video(EGRESS_SEND_BUFFER_VIDEO);
for c in codecs {
config = match c {
Codec::Opus => config.enable_opus(true),
Codec::H264 => config.enable_h264(true),
Codec::H265 => config.enable_h265(true),
Codec::Vp8 => config.enable_vp8(true),
Codec::Vp9 => config.enable_vp9(true),
Codec::Av1 => config.enable_av1(true),
_ => config,
};
}
config
}
pub fn rtc_with_codecs(codecs: &[str0m::format::Codec]) -> Rtc {
rtc_config_with_codecs(codecs).build(std::time::Instant::now())
}
pub async fn bind_udp(advertise: &[SocketAddr]) -> Result<(Arc<UdpSocket>, Vec<SocketAddr>)> {
let socket = UdpSocket::bind(("0.0.0.0", 0)).await?;
let local = socket.local_addr()?;
let candidates = advertised_candidates(advertise, local)?;
Ok((Arc::new(socket), candidates))
}
pub(crate) fn advertised_candidates(advertise: &[SocketAddr], local: SocketAddr) -> Result<Vec<SocketAddr>> {
let port = local.port();
let candidates = if advertise.is_empty() {
let ip = match local.ip() {
IpAddr::V4(ip) if ip.is_unspecified() => IpAddr::V4(Ipv4Addr::LOCALHOST),
IpAddr::V6(ip) if ip.is_unspecified() => IpAddr::V6(Ipv6Addr::LOCALHOST),
ip => ip,
};
let candidate = SocketAddr::new(ip, port);
if candidate != local {
tracing::info!(bound = %local, advertised = %candidate, "webrtc udp bind is unspecified, advertising loopback ICE candidate");
}
vec![candidate]
} else {
advertise.iter().map(|addr| SocketAddr::new(addr.ip(), port)).collect()
};
for addr in &candidates {
Candidate::host(*addr, "udp").map_err(str0m::RtcError::from)?;
}
Ok(candidates)
}
pub fn spawn_socket_reader(socket: Arc<UdpSocket>) -> mpsc::Receiver<Packet> {
let (tx, rx) = mpsc::channel(SESSION_INBOX);
tokio::spawn(async move {
let mut buf = vec![0u8; 65_535];
loop {
match socket.recv_from(&mut buf).await {
Ok((len, src)) => {
let src = crate::net::canonical(src);
if let Err(mpsc::error::TrySendError::Closed(_)) = tx.try_send((buf[..len].to_vec(), src)) {
break;
}
}
Err(err) => {
tracing::warn!(%err, "webrtc client socket recv failed");
break;
}
}
}
});
rx
}
#[cfg(test)]
mod tests {
use std::time::{Duration, UNIX_EPOCH};
use str0m::media::Mid;
use str0m::rtp::Ssrc;
use str0m::rtp::rtcp::SenderInfo;
use super::*;
#[test]
fn advertised_candidates_use_loopback_for_unspecified_ipv4() {
let local: SocketAddr = "0.0.0.0:4444".parse().unwrap();
let candidates = advertised_candidates(&[], local).unwrap();
assert_eq!(candidates, vec!["127.0.0.1:4444".parse().unwrap()]);
}
#[test]
fn advertised_candidates_use_loopback_for_unspecified_ipv6() {
let local: SocketAddr = "[::]:4444".parse().unwrap();
let candidates = advertised_candidates(&[], local).unwrap();
assert_eq!(candidates, vec!["[::1]:4444".parse().unwrap()]);
}
#[test]
fn advertised_candidates_keep_bound_address_when_specific() {
let local: SocketAddr = "127.0.0.1:4444".parse().unwrap();
assert_eq!(advertised_candidates(&[], local).unwrap(), vec![local]);
}
#[test]
fn advertised_candidates_reuse_bound_port_for_configured_addresses() {
let local: SocketAddr = "0.0.0.0:4444".parse().unwrap();
let advertised = vec!["127.0.0.1:1000".parse().unwrap(), "[::1]:2000".parse().unwrap()];
assert_eq!(
advertised_candidates(&advertised, local).unwrap(),
vec!["127.0.0.1:4444".parse().unwrap(), "[::1]:4444".parse().unwrap()]
);
}
#[test]
fn advertised_candidates_reject_configured_unspecified_addresses() {
let local: SocketAddr = "127.0.0.1:4444".parse().unwrap();
let advertised = vec!["0.0.0.0:1000".parse().unwrap()];
assert!(advertised_candidates(&advertised, local).is_err());
}
#[test]
fn pick_local_matches_address_family() {
let v4: SocketAddr = "1.2.3.4:5000".parse().unwrap();
let v6: SocketAddr = "[2001:db8::1]:5000".parse().unwrap();
let locals = vec![v4, v6];
let src_v4: SocketAddr = "9.9.9.9:1".parse().unwrap();
let src_v6: SocketAddr = "[2001:db8::2]:1".parse().unwrap();
assert_eq!(pick_local(&locals, src_v4), v4);
assert_eq!(pick_local(&locals, src_v6), v6);
assert_eq!(pick_local(&[v4], src_v6), v4);
}
#[test]
fn ingest_clock_rebases_first_frame_to_zero() {
let mut clock = IngestClock::default();
let mid = Mid::from("0");
let t0 = Instant::now();
assert_eq!(clock.normalize(mid, t0, 5_000_000_000), 0);
}
#[test]
fn ingest_clock_tracks_rtp_delta_within_track() {
let mut clock = IngestClock::default();
let mid = Mid::from("0");
let t0 = Instant::now();
assert_eq!(clock.normalize(mid, t0, 5_000_000_000), 0);
let arrival = t0 + Duration::from_millis(17); assert_eq!(clock.normalize(mid, arrival, 5_000_020_000), 20_000);
}
#[test]
fn ingest_clock_keeps_tracks_in_sync_via_arrival() {
let mut clock = IngestClock::default();
let audio = Mid::from("0");
let video = Mid::from("1");
let t0 = Instant::now();
assert_eq!(clock.normalize(audio, t0, 1_000_000_000), 0);
let video_arrival = t0 + Duration::from_millis(5);
assert_eq!(clock.normalize(video, video_arrival, 8_000_000_000), 5_000);
assert_eq!(
clock.normalize(video, video_arrival + Duration::from_millis(33), 8_000_033_000),
38_000
);
}
#[test]
fn ingest_clock_handles_track_arriving_before_epoch() {
let mut clock = IngestClock::default();
let audio = Mid::from("0");
let video = Mid::from("1");
let t0 = Instant::now();
assert_eq!(clock.normalize(audio, t0, 1_000_000), 0);
let video_arrival = t0 - Duration::from_millis(5);
assert_eq!(clock.normalize(video, video_arrival, 8_000_000), 0);
assert_eq!(
clock.normalize(video, video_arrival + Duration::from_millis(33), 8_033_000),
28_000
);
}
#[test]
fn ingest_clock_replaces_arrival_jitter_with_sender_report_sync() {
let mut clock = IngestClock::default();
let audio = Mid::from("0");
let video = Mid::from("1");
let t0 = Instant::now();
let audio_base = 1_000_000_000;
let video_base = 8_000_000_000;
assert_eq!(clock.normalize(audio, t0, audio_base), 0);
assert_eq!(
clock.normalize(video, t0 + Duration::from_millis(50), video_base),
50_000
);
assert_eq!(
clock.normalize(audio, t0 + Duration::from_secs(1), audio_base + 1_000_000),
1_000_000
);
assert_eq!(
clock.normalize(video, t0 + Duration::from_millis(1_050), video_base + 1_000_000,),
1_050_000
);
let report_time = UNIX_EPOCH + Duration::from_secs(1_700_000_001);
clock.observe(audio, sender_info(1, report_time, audio_base + 1_000_000));
clock.observe(video, sender_info(2, report_time, video_base + 1_000_000));
let audio_time = clock.normalize(audio, t0 + Duration::from_millis(1_020), audio_base + 1_020_000);
let video_time = clock.normalize(video, t0 + Duration::from_millis(1_070), video_base + 1_020_000);
assert_eq!(audio_time, video_time);
assert_eq!(audio_time, 1_070_000);
}
fn sender_info(ssrc: u32, ntp_time: SystemTime, rtp_us: u64) -> SenderInfo {
SenderInfo {
ssrc: Ssrc::from(ssrc),
ntp_time,
rtp_time: str0m::media::MediaTime::from_micros(rtp_us),
sender_packet_count: 0,
sender_octet_count: 0,
}
}
}