use std::{
collections::HashMap,
net::UdpSocket,
sync::{
Arc, Mutex,
atomic::{AtomicU64, Ordering},
},
time::{Duration, Instant},
};
use crate::pp_log::{PpLog, pp_error, pp_info, pp_warn};
use crossbeam_channel::{Receiver, Sender, TrySendError, bounded, unbounded};
use ffmpeg_next as ffmpeg;
use str0m::{
Event, Input, Output, Rtc,
change::{SdpOffer, SdpPendingOffer},
format::Codec,
media::{Direction, MediaKind, MediaTime, Mid},
net::{Protocol, Receive},
};
use crate::{
buffer::MediaBuffer,
bus::{Bus, BusEvent},
driver::{Driver, StopReceiver},
element::{Element, ElementType, element_pp_log},
error::Result,
time::{InvalidTimeBase, MediaTimestamp},
};
use super::{
command::{Command, TrackId, TrackOutState, WebRtcError},
stream_info::{StreamInfoProbe, WebRtcStreamInfo},
track::{AttachedTrack, TrackEndpoints, WebRtcHandle, WebRtcTrackSink, WebRtcTrackSource},
};
const POLL_INTERVAL: Duration = Duration::from_millis(20);
const CHANNEL_CAPACITY: usize = 128;
pub struct WebRtcPeer {
pp_log: PpLog,
name: Arc<str>,
rtc: Rtc,
socket: UdpSocket,
tracks_in: HashMap<Mid, TrackInState>,
tracks_out: HashMap<TrackId, TrackOutState>,
track_codec: HashMap<TrackId, Codec>,
negotiated_codecs: HashMap<Mid, Arc<Mutex<Vec<Codec>>>>,
track_direction: HashMap<Mid, Direction>,
remote_closed: bool,
pending: Option<(SdpPendingOffer, Vec<TrackId>)>,
next_id: Arc<AtomicU64>,
command_tx: Sender<Command>,
command_rx: Receiver<Command>,
new_track_tx: Sender<AttachedTrack>,
on_offer: Box<dyn FnMut(SdpOffer) + Send>,
on_keyframe_request: Box<dyn FnMut(TrackId) + Send>,
}
struct TrackInState {
data_tx: Sender<MediaBuffer>,
codec: Arc<Mutex<Option<Codec>>>,
stream_info_tx: Sender<WebRtcStreamInfo>,
stream_info_probe: StreamInfoProbe,
stream_info_sent: bool,
}
impl WebRtcPeer {
pub fn new(
name: impl Into<String>,
rtc: Rtc,
socket: UdpSocket,
on_offer: impl FnMut(SdpOffer) + Send + 'static,
on_keyframe_request: impl FnMut(TrackId) + Send + 'static,
) -> (Self, WebRtcHandle) {
let name: Arc<str> = name.into().into();
let pp_log = element_pp_log(ElementType::WebRtcPeer, &name, None);
pp_info!(
pp_log: &pp_log,
"created: local_addr={:?}",
socket.local_addr()
);
let (command_tx, command_rx) = bounded(CHANNEL_CAPACITY);
let (new_track_tx, new_track_rx) = unbounded();
let next_id = Arc::new(AtomicU64::new(0));
(
Self {
name,
pp_log,
rtc,
socket,
tracks_in: HashMap::new(),
tracks_out: HashMap::new(),
track_codec: HashMap::new(),
negotiated_codecs: HashMap::new(),
track_direction: HashMap::new(),
remote_closed: false,
pending: None,
next_id: next_id.clone(),
command_tx: command_tx.clone(),
command_rx,
new_track_tx,
on_offer: Box::new(on_offer),
on_keyframe_request: Box::new(on_keyframe_request),
},
WebRtcHandle {
next_id,
command_tx,
new_track_rx,
},
)
}
pub(super) fn attach_track(
&mut self,
id: TrackId,
mid: Mid,
kind: MediaKind,
direction: Direction,
) {
pp_info!(
self,
"track attached: id={id:?}, mid={mid}, kind={kind:?}, direction={direction:?}"
);
self.track_direction.insert(mid, direction);
let negotiated_codecs = Arc::new(Mutex::new(self.codecs_for_mid(mid)));
self.negotiated_codecs
.insert(mid, negotiated_codecs.clone());
let outbound_codec = self.track_codec.remove(&id);
let sink = direction.is_sending().then(|| {
WebRtcTrackSink::new(
id,
kind,
outbound_codec,
negotiated_codecs.clone(),
self.command_tx.clone(),
)
});
let source = direction.is_receiving().then(|| {
let (tx, rx) = bounded(CHANNEL_CAPACITY);
let (stream_info_tx, stream_info_rx) = bounded(1);
let codec = Arc::new(Mutex::new(None));
self.tracks_in.insert(
mid,
TrackInState {
data_tx: tx,
codec: codec.clone(),
stream_info_tx,
stream_info_probe: StreamInfoProbe::new(),
stream_info_sent: false,
},
);
WebRtcTrackSource::new(
id,
kind,
format!("webrtc-track-{}-in", id.0),
rx,
codec,
negotiated_codecs.clone(),
stream_info_rx,
)
});
let endpoints = match (sink, source) {
(Some(sink), Some(source)) => TrackEndpoints::SendRecv(sink, source),
(Some(sink), None) => TrackEndpoints::Send(sink),
(None, Some(source)) => TrackEndpoints::Recv(source),
(None, None) => TrackEndpoints::Inactive,
};
let _ = self.new_track_tx.send(AttachedTrack {
id,
mid,
kind,
endpoints,
});
}
fn codecs_for_mid(&mut self, mid: Mid) -> Vec<Codec> {
let Some(writer) = self.rtc.writer(mid) else {
return Vec::new();
};
let mut codecs = Vec::new();
for codec in writer.payload_params().map(|params| params.spec().codec) {
if !codecs.contains(&codec) {
codecs.push(codec);
}
}
codecs
}
fn refresh_codecs(&mut self, mid: Mid) {
let codecs = self.codecs_for_mid(mid);
if let Some(shared) = self.negotiated_codecs.get(&mid) {
*shared.lock().unwrap() = codecs;
}
}
fn apply_command(&mut self, cmd: Command, bus: &Bus) -> Result<()> {
match cmd {
Command::AddTrack(id, kind, direction, codec) => {
pp_info!(
self,
"add_track requested: id={id:?}, kind={kind:?}, direction={direction:?}, codec={codec:?}"
);
self.tracks_out
.insert(id, TrackOutState::ToOpen(kind, direction));
self.track_codec.insert(id, codec);
}
Command::Push(id, codec, buf) => {
if let Err(error) = self.write_track(id, codec, buf) {
bus.post(
&self.pp_log,
BusEvent::Error {
element_type: ElementType::WebRtcPeer,
name: self.name.clone(),
error,
},
);
}
}
Command::SetAnswer(answer) => {
let Some((pending, ids)) = self.pending.take() else {
return Ok(());
};
self.rtc
.sdp_api()
.accept_answer(pending, answer)
.inspect_err(|error| pp_error!(self, "accept_answer failed: {error}"))
.map_err(WebRtcError::from)?;
pp_info!(self, "renegotiation complete: {} track(s)", ids.len());
for id in ids {
if let Some(state @ TrackOutState::Negotiating(_)) = self.tracks_out.get(&id) {
let mid = state.mid().expect("Negotiating always carries a Mid");
self.tracks_out.insert(id, TrackOutState::Open(mid));
self.refresh_codecs(mid);
}
}
}
Command::AcceptOffer(offer, reply) => {
let result = self
.rtc
.sdp_api()
.accept_offer(offer)
.inspect_err(|error| pp_error!(self, "accept_offer failed: {error}"))
.map_err(WebRtcError::from);
if result.is_ok() {
pp_info!(self, "accepted remote offer");
}
let _ = reply.send(result);
}
}
Ok(())
}
fn negotiate_if_needed(&mut self) {
if self.pending.is_some() {
return;
}
let to_open: Vec<TrackId> = self
.tracks_out
.iter()
.filter(|(_, s)| matches!(s, TrackOutState::ToOpen(..)))
.map(|(id, _)| *id)
.collect();
if to_open.is_empty() {
return;
}
let mut newly_negotiating = Vec::with_capacity(to_open.len());
let mut api = self.rtc.sdp_api();
for &id in &to_open {
let Some(TrackOutState::ToOpen(kind, direction)) = self.tracks_out.get(&id) else {
continue;
};
let (kind, direction) = (*kind, *direction);
let mid = api.add_media(kind, direction, None, None, None);
self.tracks_out.insert(id, TrackOutState::Negotiating(mid));
newly_negotiating.push((id, mid, kind, direction));
}
if let Some((offer, pending)) = api.apply() {
pp_info!(self, "renegotiation started: {} track(s)", to_open.len());
self.pending = Some((pending, to_open));
(self.on_offer)(offer);
}
for (id, mid, kind, direction) in newly_negotiating {
self.attach_track(id, mid, kind, direction);
}
}
fn write_track(&mut self, id: TrackId, codec: Option<Codec>, buf: MediaBuffer) -> Result<()> {
let Some(TrackOutState::Open(mid)) = self.tracks_out.get(&id) else {
return Ok(());
};
let MediaBuffer::Packet(packet) = buf else {
return Ok(()); };
let Some(writer) = self.rtc.writer(*mid) else {
return Ok(());
};
let codec = codec.ok_or(WebRtcError::OutboundCodecNotDeclared(id))?;
let mut negotiated = Vec::new();
let mut pt = None;
for params in writer.payload_params() {
let candidate = params.spec().codec;
if !negotiated.contains(&candidate) {
negotiated.push(candidate);
}
if candidate == codec && pt.is_none() {
pt = Some(params.pt());
}
}
let Some(pt) = pt else {
return Err(WebRtcError::OutboundCodecNotNegotiated {
track_id: id,
codec,
negotiated,
}
.into());
};
let data = packet.data().unwrap_or(&[]).to_vec();
let rtp_time = packet_rtp_time(&packet)?;
writer
.write(pt, Instant::now(), rtp_time, data)
.inspect_err(|error| pp_error!(self, "writer.write failed: {error}"))
.map_err(WebRtcError::from)?;
Ok(())
}
fn drive_until_timeout(&mut self, bus: &Bus) -> Result<Instant> {
loop {
let output = self
.rtc
.poll_output()
.inspect_err(|error| pp_error!(self, "poll_output failed: {error}"))
.map_err(WebRtcError::from)?;
match output {
Output::Timeout(deadline) => return Ok(deadline),
Output::Transmit(t) => {
let _ = self.socket.send_to(&t.contents, t.destination);
}
Output::Event(event) => self.handle_event(event, bus),
}
}
}
fn close_connection(&mut self, bus: &Bus) {
if !self.rtc.is_alive() {
return; }
if let Err(error) = self.rtc.close() {
pp_warn!(self, "close failed, ending without notifying: {error}");
return;
}
if let Err(error) = self.drive_until_timeout(bus) {
pp_warn!(self, "draining close_notify failed: {error}");
return;
}
pp_info!(self, "event=close phase=completed outcome=ok");
}
pub(super) fn handle_event(&mut self, event: Event, bus: &Bus) {
match event {
Event::MediaAdded(added) => {
let id = TrackId(self.next_id.fetch_add(1, Ordering::Relaxed));
self.tracks_out.insert(id, TrackOutState::Open(added.mid));
self.attach_track(id, added.mid, added.kind, added.direction);
}
Event::MediaData(data) => {
if let Some(track) = self.tracks_in.get_mut(&data.mid) {
let codec = data.params.spec();
track.codec.lock().unwrap().replace(codec.codec);
if !track.stream_info_sent
&& let Some(info) = track.stream_info_probe.observe(codec, &data.data)
{
track.stream_info_sent = track.stream_info_tx.try_send(info).is_ok();
}
let mut packet = ffmpeg::Packet::copy(&data.data);
packet.set_time_base(ffmpeg::Rational::new(1, data.time.denom() as i32));
let pts = data.time.numer() as i64;
packet.set_pts(Some(pts));
packet.set_dts(Some(pts));
if data.is_keyframe() {
let flags = packet.flags() | ffmpeg::codec::packet::Flags::KEY;
packet.set_flags(flags);
}
match track
.data_tx
.try_send(MediaBuffer::Packet(Arc::new(packet)))
{
Ok(()) => {}
Err(TrySendError::Full(_)) => {
bus.post(
&self.pp_log,
BusEvent::Dropped {
element_type: ElementType::WebRtcPeer,
name: self.name.clone(),
},
);
}
Err(TrySendError::Disconnected(_)) => {
self.tracks_in.remove(&data.mid);
}
}
}
}
Event::KeyframeRequest(req) => {
if let Some((&id, _)) = self
.tracks_out
.iter()
.find(|(_, s)| s.mid() == Some(req.mid))
{
pp_info!(self, "keyframe requested: id={id:?}, mid={}", req.mid);
(self.on_keyframe_request)(id);
}
}
Event::Connected => {
pp_info!(self, "ICE+DTLS connected");
}
Event::IceConnectionStateChange(state) => {
pp_info!(self, "ICE connection state: {state:?}");
}
Event::MediaChanged(changed) => {
pp_info!(
self,
"media changed: mid={}, direction={:?}",
changed.mid,
changed.direction
);
self.refresh_codecs(changed.mid);
if let Some(&from) = self.track_direction.get(&changed.mid)
&& from != changed.direction
{
self.track_direction.insert(changed.mid, changed.direction);
bus.post(
&self.pp_log,
BusEvent::Error {
element_type: ElementType::WebRtcPeer,
name: self.name.clone(),
error: WebRtcError::DirectionChanged {
mid: changed.mid,
from,
to: changed.direction,
}
.into(),
},
);
}
}
Event::Closed => {
pp_info!(self, "event=close phase=remote_received outcome=ok");
self.remote_closed = true;
}
_ => {}
}
}
}
pub(super) fn packet_rtp_time(
packet: &ffmpeg::Packet,
) -> std::result::Result<MediaTime, WebRtcError> {
let pts = packet.pts().ok_or(WebRtcError::MissingPacketPts)?;
let timestamp = MediaTimestamp::try_new(pts, packet.time_base()).map_err(
|InvalidTimeBase {
numerator,
denominator,
}| WebRtcError::InvalidPacketTimeBase {
numerator,
denominator,
},
)?;
to_str0m_media_time(timestamp)
}
fn to_str0m_media_time(timestamp: MediaTimestamp) -> std::result::Result<MediaTime, WebRtcError> {
let time_base = timestamp.time_base().get();
let numerator = time_base.numerator();
let denominator = time_base.denominator();
let pts = u64::try_from(timestamp.pts())
.map_err(|_| WebRtcError::NegativePacketPts(timestamp.pts()))?;
let frequency = str0m::media::Frequency::new(denominator as u32).ok_or(
WebRtcError::InvalidPacketTimeBase {
numerator,
denominator,
},
)?;
let numer = pts
.checked_mul(numerator as u64)
.ok_or(WebRtcError::PacketTimestampOverflow {
pts,
numerator,
denominator,
})?;
Ok(MediaTime::new(numer, frequency))
}
impl Element for WebRtcPeer {
fn name(&self) -> Arc<str> {
self.name.clone()
}
fn element_type(&self) -> ElementType {
ElementType::WebRtcPeer
}
fn pp_log(&self) -> &PpLog {
&self.pp_log
}
fn pp_log_mut(&mut self) -> &mut PpLog {
&mut self.pp_log
}
}
impl Driver for WebRtcPeer {
fn run(&mut self, stop: &StopReceiver, bus: &Bus) -> Result<()> {
pp_info!(self, "started");
let mut buf = vec![0u8; 2000];
loop {
while let Ok(cmd) = self.command_rx.try_recv() {
self.apply_command(cmd, bus)?;
}
self.negotiate_if_needed();
let deadline = self.drive_until_timeout(bus)?;
if self.remote_closed || !self.rtc.is_alive() || stop.is_stopped() {
pp_info!(
self,
"stopped rtc_alive={} remote_closed={}",
self.rtc.is_alive(),
self.remote_closed
);
if !self.remote_closed {
self.close_connection(bus);
}
self.tracks_in.clear();
return Ok(());
}
let wait = deadline
.saturating_duration_since(Instant::now())
.min(POLL_INTERVAL)
.max(Duration::from_millis(1));
self.socket
.set_read_timeout(Some(wait))
.inspect_err(|error| pp_error!(self, "set_read_timeout failed: {error}"))
.map_err(WebRtcError::from)?;
match self.socket.recv_from(&mut buf) {
Ok((n, source)) => {
let Ok(contents) = buf[..n].try_into() else {
continue; };
let destination = self
.socket
.local_addr()
.inspect_err(|error| pp_error!(self, "local_addr failed: {error}"))
.map_err(WebRtcError::from)?;
self.rtc
.handle_input(Input::Receive(
Instant::now(),
Receive {
proto: Protocol::Udp,
source,
destination,
contents,
},
))
.inspect_err(|error| {
pp_error!(self, "handle_input(Receive) failed: {error}")
})
.map_err(WebRtcError::from)?;
}
Err(e)
if matches!(
e.kind(),
std::io::ErrorKind::WouldBlock | std::io::ErrorKind::TimedOut
) =>
{
self.rtc
.handle_input(Input::Timeout(Instant::now()))
.inspect_err(|error| {
pp_error!(self, "handle_input(Timeout) failed: {error}")
})
.map_err(WebRtcError::from)?;
}
Err(e) => {
pp_error!(self, "recv_from failed: {e}");
return Err(WebRtcError::from(e).into());
}
}
}
}
}