use std::{
sync::{
Arc, Mutex,
atomic::{AtomicU64, Ordering},
},
time::Duration,
};
use ffmpeg_next as ffmpeg;
use crate::pp_log::{PpLog, pp_error, pp_info};
use crossbeam_channel::{Receiver, RecvTimeoutError, Sender, TrySendError, select};
use str0m::{
change::{SdpAnswer, SdpOffer},
format::Codec,
media::{Direction, MediaKind, Mid},
};
use crate::{
buffer::MediaBuffer,
bus::{Bus, BusEvent},
contract::{InputContract, OutputContract, PortContract},
control::{
ControlMsg, ControlReceiver, RequestKind, apply_finish, apply_one, drain_control,
wait_out_pause,
},
element::{Element, ElementType, Sink, Source, SourceElement, element_pp_log},
error::Result,
pad::SrcPad,
};
use super::{
command::{Command, TrackId, WebRtcError},
stream_info::{WebRtcStreamInfo, annex_b_nalus, str0m_codec},
};
fn packet_kind(kind: MediaKind) -> crate::contract::MediaKind {
match kind {
MediaKind::Audio => crate::contract::MediaKind::AudioPacket,
MediaKind::Video => crate::contract::MediaKind::VideoPacket,
}
}
const START_CODE: [u8; 4] = [0, 0, 0, 1];
fn starts_with_start_code(data: &[u8]) -> bool {
data.starts_with(&START_CODE) || data.starts_with(&START_CODE[1..])
}
fn annex_b_codec(codec: Codec) -> bool {
matches!(codec, Codec::H264 | Codec::H265 | Codec::H266)
}
fn carries_parameter_sets(payload: &[u8], codec: Codec) -> Option<bool> {
let (sps, pps) = match codec {
Codec::H264 => (7, 8),
Codec::H265 => (33, 34),
Codec::H266 => (15, 16),
_ => return None,
};
let nal_type = |nalu: &&[u8]| match codec {
Codec::H264 => nalu.first().map(|byte| byte & 0x1f),
Codec::H266 => nalu.get(1).map(|byte| byte >> 3),
_ => nalu.first().map(|byte| (byte >> 1) & 0x3f),
};
let present: Vec<u8> = annex_b_nalus(payload).iter().filter_map(nal_type).collect();
Some(present.contains(&sps) && present.contains(&pps))
}
fn avcc_parameter_sets(config: &[u8]) -> Option<(Option<Vec<u8>>, usize)> {
const HEADER: usize = 6;
if config.len() < HEADER || config[0] != 1 || starts_with_start_code(config) {
return None;
}
let nal_length_size = (config[4] & 0x03) as usize + 1;
let mut parameter_sets = Vec::new();
let mut offset = HEADER - 1;
for count_mask in [0x1f_u8, 0xff] {
let count = config.get(offset)? & count_mask;
offset += 1;
for _ in 0..count {
let length =
u16::from_be_bytes([*config.get(offset)?, *config.get(offset + 1)?]) as usize;
offset += 2;
let end = offset.checked_add(length)?;
if end > config.len() || length == 0 {
return None;
}
parameter_sets.extend_from_slice(&START_CODE);
parameter_sets.extend_from_slice(&config[offset..end]);
offset = end;
}
}
Some((
(!parameter_sets.is_empty()).then_some(parameter_sets),
nal_length_size,
))
}
fn length_prefixed_to_annex_b(payload: &[u8], nal_length_size: usize) -> Option<Vec<u8>> {
let mut annex_b = Vec::with_capacity(payload.len() + START_CODE.len());
let mut offset = 0;
while offset < payload.len() {
let prefix = payload.get(offset..offset + nal_length_size)?;
let length = prefix
.iter()
.fold(0usize, |value, byte| (value << 8) | usize::from(*byte));
offset += nal_length_size;
let end = offset.checked_add(length)?;
if length == 0 || end > payload.len() {
return None;
}
annex_b.extend_from_slice(&START_CODE);
annex_b.extend_from_slice(&payload[offset..end]);
offset = end;
}
(!annex_b.is_empty()).then_some(annex_b)
}
fn rewritten_packet(packet: &ffmpeg::Packet, payload: &[u8]) -> MediaBuffer {
let mut rewritten = ffmpeg::Packet::copy(payload);
rewritten.set_time_base(packet.time_base());
rewritten.set_pts(packet.pts());
rewritten.set_dts(packet.dts());
rewritten.set_stream(packet.stream());
rewritten.set_flags(packet.flags());
rewritten.set_duration(packet.duration());
MediaBuffer::Packet(Arc::new(rewritten))
}
pub enum TrackEndpoints {
Send(WebRtcTrackSink),
Recv(WebRtcTrackSource),
SendRecv(WebRtcTrackSink, WebRtcTrackSource),
Inactive,
}
pub struct AttachedTrack {
pub id: TrackId,
pub mid: Mid,
pub kind: MediaKind,
pub endpoints: TrackEndpoints,
}
#[derive(Clone)]
pub struct WebRtcHandle {
pub(super) next_id: Arc<AtomicU64>,
pub(super) command_tx: Sender<Command>,
pub(super) new_track_rx: Receiver<AttachedTrack>,
}
impl WebRtcHandle {
pub fn add_track(
&self,
kind: MediaKind,
direction: Direction,
codec: Codec,
) -> Result<TrackId> {
let id = TrackId(self.next_id.fetch_add(1, Ordering::Relaxed));
self.command_tx
.send(Command::AddTrack(id, kind, direction, codec))
.map_err(|_| WebRtcError::Closed)?;
Ok(id)
}
pub fn next_track(&self) -> Result<AttachedTrack> {
self.new_track_rx
.recv()
.map_err(|_| WebRtcError::Closed.into())
}
pub fn set_answer(&self, answer: SdpAnswer) {
let _ = self.command_tx.send(Command::SetAnswer(answer));
}
pub fn accept_remote_offer(&self, offer: SdpOffer) -> Result<SdpAnswer> {
let (reply_tx, reply_rx) = crossbeam_channel::bounded(0);
self.command_tx
.send(Command::AcceptOffer(offer, reply_tx))
.map_err(|_| WebRtcError::Closed)?;
reply_rx
.recv()
.map_err(|_| WebRtcError::Closed)?
.map_err(Into::into)
}
}
pub struct WebRtcTrackSink {
pp_log: PpLog,
id: TrackId,
kind: MediaKind,
codec: Option<Codec>,
negotiated_codecs: Arc<Mutex<Vec<Codec>>>,
command_tx: Sender<Command>,
timestamp_offset: Option<i64>,
parameter_sets: Option<Vec<u8>>,
nal_length_size: Option<usize>,
bitstream_checked: bool,
parameter_sets_checked: bool,
}
impl WebRtcTrackSink {
pub(super) fn new(
id: TrackId,
kind: MediaKind,
codec: Option<Codec>,
negotiated_codecs: Arc<Mutex<Vec<Codec>>>,
command_tx: Sender<Command>,
) -> Self {
Self {
id,
kind,
codec,
negotiated_codecs,
command_tx,
timestamp_offset: None,
parameter_sets: None,
nal_length_size: None,
bitstream_checked: false,
parameter_sets_checked: false,
pp_log: element_pp_log(
ElementType::WebRtcPeer,
&format!("webrtc-track-{}", id.0),
None,
),
}
}
pub fn negotiated_codecs(&self) -> Vec<Codec> {
self.negotiated_codecs.lock().unwrap().clone()
}
pub fn set_source_parameters(&mut self, parameters: &ffmpeg::codec::Parameters) -> Result<()> {
let id = parameters.id();
let codec = str0m_codec(id).ok_or(WebRtcError::SourceCodecUnsupported(id))?;
let negotiated = self.negotiated_codecs();
if !negotiated.contains(&codec) {
return Err(WebRtcError::OutboundCodecNotNegotiated {
track_id: self.id,
codec,
negotiated,
}
.into());
}
let bytes = unsafe {
let raw = parameters.as_ptr();
let size = usize::try_from((*raw).extradata_size).unwrap_or(0);
match ((*raw).extradata.is_null() || size == 0).then_some(()) {
Some(()) => Vec::new(),
None => std::slice::from_raw_parts((*raw).extradata, size).to_vec(),
}
};
let (parameter_sets, nal_length_size) = match () {
_ if !annex_b_codec(codec) => (None, None),
_ if bytes.is_empty() => (None, None),
_ if starts_with_start_code(&bytes) => (Some(bytes), None),
_ if codec == Codec::H264 => match avcc_parameter_sets(&bytes) {
Some((annex_b, length_size)) => (annex_b, Some(length_size)),
None => {
return Err(WebRtcError::ParameterSetsNotSupported {
track_id: self.id,
codec,
}
.into());
}
},
_ => {
return Err(WebRtcError::ParameterSetsNotSupported {
track_id: self.id,
codec,
}
.into());
}
};
self.codec = Some(codec);
self.parameter_sets = parameter_sets;
self.nal_length_size = nal_length_size;
self.forget_what_was_checked();
Ok(())
}
pub fn set_codec(&mut self, codec: Codec) -> Result<()> {
let negotiated = self.negotiated_codecs();
if !negotiated.contains(&codec) {
return Err(WebRtcError::OutboundCodecNotNegotiated {
track_id: self.id,
codec,
negotiated,
}
.into());
}
self.codec = Some(codec);
self.parameter_sets = None;
self.nal_length_size = None;
self.forget_what_was_checked();
Ok(())
}
fn forget_what_was_checked(&mut self) {
self.bitstream_checked = false;
self.parameter_sets_checked = false;
}
}
impl Element for WebRtcTrackSink {
fn name(&self) -> Arc<str> {
format!("webrtc-track-{}", self.id.0).into()
}
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 Sink for WebRtcTrackSink {
fn input_contract(&self) -> InputContract {
InputContract::Fixed(PortContract::packet(packet_kind(self.kind)))
}
fn consume(&mut self, buf: MediaBuffer) -> Result<()> {
if !matches!(buf, MediaBuffer::Packet(_) | MediaBuffer::Eos) {
let kind = match buf {
MediaBuffer::Video(_) => "Video",
MediaBuffer::Audio(_) => "Audio",
MediaBuffer::Packet(_) | MediaBuffer::Eos => unreachable!("matched above"),
};
pp_error!(self, "unsupported buffer: {kind}");
return Err(WebRtcError::UnsupportedBuffer(kind).into());
}
if matches!(buf, MediaBuffer::Packet(_)) && self.codec.is_none() {
pp_error!(self, "outbound codec is not declared");
return Err(WebRtcError::OutboundCodecNotDeclared(self.id).into());
}
if let MediaBuffer::Packet(_) = &buf {
let codec = self.codec.expect("checked above");
let negotiated = self.negotiated_codecs();
if !negotiated.contains(&codec) {
pp_error!(self, "outbound codec {codec:?} is not negotiated");
return Err(WebRtcError::OutboundCodecNotNegotiated {
track_id: self.id,
codec,
negotiated,
}
.into());
}
}
let buf = self.prepare_bitstream(buf)?;
let buf = self.prepend_parameter_sets(buf);
let buf = self.normalize_packet_timestamp(buf)?;
match self
.command_tx
.try_send(Command::Push(self.id, self.codec, buf))
{
Ok(()) | Err(TrySendError::Full(_)) => Ok(()),
Err(TrySendError::Disconnected(_)) => {
pp_error!(self, "WebRtcPeer::run gone — track is dead");
Err(WebRtcError::Closed.into())
}
}
}
fn control(&mut self, _msg: ControlMsg) -> Result<()> {
Ok(())
}
}
impl WebRtcTrackSink {
fn prepend_parameter_sets(&self, buf: MediaBuffer) -> MediaBuffer {
let Some(headers) = self.parameter_sets.as_deref() else {
return buf;
};
let MediaBuffer::Packet(packet) = &buf else {
return buf;
};
let Some(payload) = packet.data() else {
return buf;
};
if !packet.is_key() || payload.starts_with(headers) {
return buf;
}
let mut joined = Vec::with_capacity(headers.len() + payload.len());
joined.extend_from_slice(headers);
joined.extend_from_slice(payload);
rewritten_packet(packet, &joined)
}
fn prepare_bitstream(&mut self, buf: MediaBuffer) -> Result<MediaBuffer> {
let Some(codec) = self.codec.filter(|codec| annex_b_codec(*codec)) else {
return Ok(buf);
};
let MediaBuffer::Packet(packet) = &buf else {
return Ok(buf);
};
let Some(payload) = packet.data() else {
return Ok(buf);
};
let converted = match self.nal_length_size {
Some(_) if starts_with_start_code(payload) => None,
Some(nal_length_size) => {
let Some(annex_b) = length_prefixed_to_annex_b(payload, nal_length_size) else {
pp_error!(
self,
"outbound packet is not a valid length-prefixed access unit"
);
return Err(WebRtcError::MalformedLengthPrefixedPacket(self.id).into());
};
Some(annex_b)
}
None => {
if !self.bitstream_checked {
if !starts_with_start_code(payload) {
pp_error!(self, "outbound packet is not Annex-B");
return Err(WebRtcError::NotAnnexB(self.id).into());
}
self.bitstream_checked = true;
}
None
}
};
if packet.is_key() && self.parameter_sets.is_none() && !self.parameter_sets_checked {
let outgoing = converted.as_deref().unwrap_or(payload);
if carries_parameter_sets(outgoing, codec) == Some(false) {
pp_error!(self, "outbound keyframe carries no parameter sets");
return Err(WebRtcError::MissingParameterSets(self.id).into());
}
self.parameter_sets_checked = true;
}
match converted {
Some(annex_b) => Ok(rewritten_packet(packet, &annex_b)),
None => Ok(buf),
}
}
fn normalize_packet_timestamp(&mut self, buf: MediaBuffer) -> Result<MediaBuffer> {
let MediaBuffer::Packet(packet) = buf else {
return Ok(buf);
};
let Some(pts) = packet.pts() else {
return Ok(MediaBuffer::Packet(packet));
};
let offset = match self.timestamp_offset {
Some(offset) => offset,
None if pts < 0 => {
pts.checked_neg()
.ok_or(WebRtcError::PacketTimestampNormalizationOverflow {
value: pts,
offset: 0,
})?
}
None => 0,
};
self.timestamp_offset = Some(offset);
if offset == 0 {
return Ok(MediaBuffer::Packet(packet));
}
let shifted = |value: i64| {
value
.checked_add(offset)
.ok_or(WebRtcError::PacketTimestampNormalizationOverflow { value, offset })
};
let mut normalized = (*packet).clone();
normalized.set_pts(Some(shifted(pts)?));
normalized.set_dts(packet.dts().map(shifted).transpose()?);
Ok(MediaBuffer::Packet(Arc::new(normalized)))
}
}
pub struct WebRtcTrackSource {
id: TrackId,
pp_log: PpLog,
name: Arc<str>,
pad: SrcPad,
data_rx: Receiver<MediaBuffer>,
codec: Arc<Mutex<Option<Codec>>>,
negotiated_codecs: Arc<Mutex<Vec<Codec>>>,
stream_info: Mutex<StreamInfoState>,
}
struct StreamInfoState {
rx: Receiver<WebRtcStreamInfo>,
cached: Option<WebRtcStreamInfo>,
}
impl WebRtcTrackSource {
pub(super) fn new(
id: TrackId,
kind: MediaKind,
name: impl Into<String>,
data_rx: Receiver<MediaBuffer>,
codec: Arc<Mutex<Option<Codec>>>,
negotiated_codecs: Arc<Mutex<Vec<Codec>>>,
stream_info_rx: Receiver<WebRtcStreamInfo>,
) -> Self {
let name: Arc<str> = name.into().into();
let pp_log = element_pp_log(ElementType::WebRtcPeer, &name, None);
let pad = SrcPad::with_contract(
format!("{name}_src"),
OutputContract::Fixed(PortContract::packet(packet_kind(kind))),
);
Self {
id,
name,
pp_log,
pad,
data_rx,
codec,
negotiated_codecs,
stream_info: Mutex::new(StreamInfoState {
rx: stream_info_rx,
cached: None,
}),
}
}
pub fn wait_stream_info(&self, timeout: Duration) -> Result<WebRtcStreamInfo> {
let mut state = self.stream_info.lock().unwrap();
if let Some(info) = &state.cached {
return Ok(info.clone());
}
match state.rx.recv_timeout(timeout) {
Ok(info) => {
state.cached = Some(info.clone());
Ok(info)
}
Err(RecvTimeoutError::Timeout) => Err(WebRtcError::StreamInfoTimeout {
track_id: self.id,
timeout,
}
.into()),
Err(RecvTimeoutError::Disconnected) => Err(WebRtcError::Closed.into()),
}
}
pub fn negotiated_codecs(&self) -> Vec<Codec> {
self.negotiated_codecs.lock().unwrap().clone()
}
pub fn codec(&self) -> Option<Codec> {
*self.codec.lock().unwrap()
}
}
impl Element for WebRtcTrackSource {
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 Source for WebRtcTrackSource {
fn src_pads(&mut self) -> &mut [SrcPad] {
std::slice::from_mut(&mut self.pad)
}
}
impl SourceElement for WebRtcTrackSource {
fn is_live(&self) -> bool {
true
}
fn is_seekable(&self) -> bool {
false
}
fn run(&mut self, control: &ControlReceiver, bus: &Bus) -> Result<()> {
pp_info!(self, "started");
loop {
if drain_control(control, self, bus)?.stopped {
pp_info!(self, "stopped");
return Ok(());
}
select! {
recv(control.rx) -> req => {
match req {
Ok(req) => {
match req.kind {
RequestKind::Finish => {
apply_finish(self, bus, &req.ack);
pp_info!(self, "finished");
return Ok(());
}
RequestKind::Control(msg) => {
if apply_one(self, bus, &msg, &req.ack)? {
pp_info!(self, "stopped");
return Ok(());
}
if msg == ControlMsg::Pause
&& wait_out_pause(control, self, bus)?
{
pp_info!(self, "stopped");
return Ok(());
}
}
}
}
Err(_) => {
pp_info!(self, "run: control channel gone, ending");
return Ok(());
}
}
}
recv(self.data_rx) -> buf => {
match buf {
Ok(buf) if buf.is_eos() => {
pp_info!(self, "event=eos phase=source_received");
break;
}
Ok(buf) => {
if let Err(error) = self.pad.push(buf) {
bus.post(
&self.pp_log,
BusEvent::Error {
element_type: ElementType::WebRtcPeer,
name: self.name.clone(),
error,
},
);
}
}
Err(_) => {
pp_info!(self, "run: WebRtcPeer gone, ending");
break;
}
}
}
}
}
while let Some((_msg, ack)) = control.try_recv() {
let _ = ack.send(());
}
self.pad.push_eos(&self.pp_log)
}
fn seek(&mut self, target: Duration) -> Result<Duration> {
Ok(target)
}
}