use std::{
collections::HashMap,
net::UdpSocket,
sync::{
Arc, Mutex,
atomic::{AtomicU64, Ordering},
},
time::{Duration, Instant},
};
use crate::pp_log::{PpLog, pp_error, pp_info};
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::{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},
track::{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,
#[allow(clippy::type_complexity)]
tracks_in: HashMap<Mid, (Sender<MediaBuffer>, Arc<Mutex<Option<Codec>>>)>,
tracks_out: HashMap<TrackId, TrackOutState>,
track_codec: HashMap<TrackId, Codec>,
pending: Option<(SdpPendingOffer, Vec<TrackId>)>,
next_id: Arc<AtomicU64>,
command_tx: Sender<Command>,
command_rx: Receiver<Command>,
new_track_tx: Sender<(TrackId, Mid, MediaKind, WebRtcTrackSink, WebRtcTrackSource)>,
on_offer: Box<dyn FnMut(SdpOffer) + Send>,
on_keyframe_request: Box<dyn FnMut(TrackId) + Send>,
}
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(),
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,
},
)
}
fn attach_track(&mut self, id: TrackId, mid: Mid, kind: MediaKind) {
pp_info!(self, "track attached: id={id:?}, mid={mid}, kind={kind:?}");
let reply = WebRtcTrackSink::new(id, self.command_tx.clone());
let (tx, rx) = bounded(CHANNEL_CAPACITY);
let codec = Arc::new(Mutex::new(None));
self.tracks_in.insert(mid, (tx, codec.clone()));
let source = WebRtcTrackSource::new(format!("webrtc-track-{}-in", id.0), rx, codec);
let _ = self.new_track_tx.send((id, mid, kind, reply, source));
}
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, buf) => {
if let Err(error) = self.write_track(id, 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));
}
}
}
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));
}
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) in newly_negotiating {
self.attach_track(id, mid, kind);
}
}
fn write_track(&mut self, id: TrackId, 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 pt = match self.track_codec.get(&id) {
Some(&codec) => writer.payload_params().find(|p| p.spec().codec == codec),
None => writer.payload_params().next(),
}
.map(|p| p.pt());
let Some(pt) = pt else {
return Ok(()); };
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 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);
}
Event::MediaData(data) => {
if let Some((tx, codec)) = self.tracks_in.get(&data.mid) {
*codec.lock().unwrap() = Some(data.params.spec().codec);
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 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
);
}
_ => {}
}
}
}
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.rtc.is_alive() || stop.is_stopped() {
pp_info!(self, "stopped rtc_alive={}", self.rtc.is_alive());
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());
}
}
}
}
}