mod rtmp;
mod rtsp;
mod srt;
use std::sync::Arc;
use std::sync::atomic::{AtomicU32, AtomicU64, Ordering};
use std::time::{Duration, Instant};
use media_plane::trunk::{SampleCursorItem, Trunk};
use tokio::time::sleep;
use tokio_util::sync::CancellationToken;
use crate::config::PushFormat;
use broadcast_common::Package;
use transmux::TsMux;
use transmux::ir::{Media, Sample, Track, TrackSpec};
pub use rtmp::{RtmpTransport, RtmpTransportConfig};
pub use rtsp::{RtspTransport, RtspTransportConfig};
pub use srt::{SrtTransport, SrtTransportConfig};
const STREAM_TYPE_PRIVATE: u8 = 0x06;
#[non_exhaustive]
#[derive(Debug, thiserror::Error)]
pub enum SendMediaError {
#[error("push mux failed: {0}")]
Mux(String),
#[error("push send failed: {0}")]
Transport(Box<dyn std::error::Error + Send + Sync>),
}
#[async_trait::async_trait]
pub trait PushTransport: Send + 'static {
type Config: Send + Sync + Clone;
type Error: std::error::Error + Send + Sync + 'static;
async fn connect(url: &str, config: &Self::Config) -> Result<Self, Self::Error>
where
Self: Sized;
async fn send(&mut self, data: &[u8]) -> Result<(), Self::Error>;
async fn setup(&mut self, _tracks: &[TrackSpec]) -> Result<(), Self::Error> {
Ok(())
}
async fn send_media(&mut self, media: &Media) -> Result<u64, SendMediaError> {
let bytes = TsMux::new()
.package(media)
.map_err(|e| SendMediaError::Mux(e.to_string()))?;
self.send(&bytes)
.await
.map_err(|e| SendMediaError::Transport(Box::new(e)))?;
Ok(bytes.len() as u64)
}
fn close(&mut self);
}
#[derive(Debug, Clone, PartialEq, Eq)]
#[non_exhaustive]
pub enum ReconnectState {
Ready,
Backoff { resume_at: Instant },
Failed,
}
#[derive(Debug, Clone)]
pub struct ReconnectEngine {
state: ReconnectState,
attempt: u32,
policy: crate::config::ReconnectPolicy,
}
impl ReconnectEngine {
pub fn new(policy: crate::config::ReconnectPolicy) -> Self {
Self {
state: ReconnectState::Ready,
attempt: 0,
policy,
}
}
pub fn should_connect(&self) -> bool {
match self.state {
ReconnectState::Ready => true,
ReconnectState::Failed => false,
ReconnectState::Backoff { resume_at } => Instant::now() >= resume_at,
}
}
pub fn on_connect(&mut self) {
self.attempt = 0;
self.state = ReconnectState::Ready;
}
pub fn on_disconnect(&mut self) {
self.attempt = self.attempt.saturating_add(1);
if let Some(max) = self.policy.max_attempts {
if self.attempt >= max {
self.state = ReconnectState::Failed;
return;
}
}
let backoff = self.policy.backoff_for(self.attempt.saturating_sub(1));
self.state = ReconnectState::Backoff {
resume_at: Instant::now() + backoff,
};
}
pub fn is_failed(&self) -> bool {
self.state == ReconnectState::Failed
}
pub fn time_until_retry(&self) -> Duration {
match self.state {
ReconnectState::Ready | ReconnectState::Failed => Duration::ZERO,
ReconnectState::Backoff { resume_at } => {
resume_at.saturating_duration_since(Instant::now())
}
}
}
pub fn state(&self) -> &ReconnectState {
&self.state
}
}
#[derive(Debug, Default)]
pub struct PushMetrics {
samples_sent: AtomicU64,
samples_dropped: AtomicU64,
bytes_sent: AtomicU64,
reconnect_count: AtomicU32,
}
impl PushMetrics {
pub fn samples_sent(&self) -> u64 {
self.samples_sent.load(Ordering::Relaxed)
}
pub fn samples_dropped(&self) -> u64 {
self.samples_dropped.load(Ordering::Relaxed)
}
pub fn bytes_sent(&self) -> u64 {
self.bytes_sent.load(Ordering::Relaxed)
}
pub fn reconnect_count(&self) -> u32 {
self.reconnect_count.load(Ordering::Relaxed)
}
}
pub async fn drive_push<T: PushTransport>(
trunk: Arc<Trunk>,
url: String,
config: T::Config,
_format: PushFormat,
reconnect: crate::config::ReconnectPolicy,
cancel: CancellationToken,
) -> PushMetrics {
let metrics = PushMetrics::default();
let mut transport: Option<T> = None;
let mut engine = ReconnectEngine::new(reconnect);
let mut cursor = trunk.subscribe();
let _ = _format;
loop {
if cancel.is_cancelled() {
if let Some(t) = transport.as_mut() {
t.close();
}
return metrics;
}
if let Some(listener) = trunk.listen() {
listener.wait_deadline(Instant::now() + Duration::from_millis(250));
}
let mut drained: Vec<(u32, Vec<Sample>)> = Vec::new();
while let Some(item) = cursor.poll() {
match item {
SampleCursorItem::Timed { track_id, sample }
| SampleCursorItem::Sparse { track_id, sample } => {
match drained.iter_mut().find(|(id, _)| *id == track_id) {
Some((_, samples)) => samples.push(sample),
None => drained.push((track_id, vec![sample])),
}
}
SampleCursorItem::Lagged { skipped } | SampleCursorItem::Degraded { skipped } => {
metrics
.samples_dropped
.fetch_add(skipped, Ordering::Relaxed);
}
_ => {}
}
}
if engine.is_failed() {
return metrics;
}
if engine.should_connect() && transport.is_none() {
match T::connect(&url, &config).await {
Ok(mut conn) => {
engine.on_connect();
metrics.reconnect_count.fetch_add(1, Ordering::Relaxed);
let tracks = trunk.tracks();
if let Err(e) = conn.setup(&tracks).await {
tracing::warn!(%url, error = %e, "push setup failed; closing");
conn.close();
engine.on_disconnect();
continue;
}
transport = Some(conn);
}
Err(e) => {
tracing::warn!(%url, error = %e, "push connect failed; backing off");
engine.on_disconnect();
}
}
}
if drained.is_empty() {
continue;
}
if transport.is_some() {
let media = media_from_samples(&trunk.tracks(), &drained);
match transport.as_mut().unwrap().send_media(&media).await {
Ok(sent_bytes) => {
metrics.bytes_sent.fetch_add(sent_bytes, Ordering::Relaxed);
let sent: u64 = drained.iter().map(|(_, s)| s.len() as u64).sum();
metrics.samples_sent.fetch_add(sent, Ordering::Relaxed);
}
Err(SendMediaError::Mux(msg)) => {
tracing::error!(error = %msg, "push mux failed; dropping batch");
let dropped: u64 = drained.iter().map(|(_, s)| s.len() as u64).sum();
metrics
.samples_dropped
.fetch_add(dropped, Ordering::Relaxed);
}
Err(SendMediaError::Transport(e)) => {
tracing::warn!(%url, error = %e, "push send failed; reconnecting");
transport.as_mut().unwrap().close();
transport = None;
engine.on_disconnect();
}
}
} else {
let dropped: u64 = drained.iter().map(|(_, s)| s.len() as u64).sum();
metrics
.samples_dropped
.fetch_add(dropped, Ordering::Relaxed);
let wait = engine.time_until_retry();
if !wait.is_zero() {
sleep(wait).await;
}
}
}
}
fn media_from_samples(tracks: &[TrackSpec], drained: &[(u32, Vec<Sample>)]) -> Media {
let mut built: Vec<Track> = Vec::new();
for (id, samples) in drained {
let spec = tracks
.iter()
.find(|s| s.track_id == *id)
.cloned()
.unwrap_or_else(|| {
TrackSpec::new(
*id,
90_000,
transmux::CodecConfig::Data {
stream_type: STREAM_TYPE_PRIVATE, descriptors: Vec::new(),
carriage: transmux::ir::DataCarriage::Pes,
},
)
});
built.push(Track::new(spec, samples.clone()));
}
let timescale = tracks.iter().next().map(|t| t.timescale).unwrap_or(90_000);
Media::new(built, timescale)
}
#[cfg(test)]
mod tests {
use super::*;
use crate::config::ReconnectPolicy;
fn policy(max_attempts: Option<u32>) -> ReconnectPolicy {
ReconnectPolicy {
initial_backoff_ms: 1_000,
max_backoff_ms: 30_000,
max_attempts,
}
}
#[test]
fn backoff_doubles_each_attempt_capped_at_max() {
let eng = ReconnectEngine::new(policy(None));
assert_eq!(eng.policy.backoff_for(0), Duration::from_millis(1_000));
assert_eq!(eng.policy.backoff_for(1), Duration::from_millis(2_000));
assert_eq!(eng.policy.backoff_for(2), Duration::from_millis(4_000));
assert_eq!(eng.policy.backoff_for(20), Duration::from_millis(30_000));
assert_eq!(eng.policy.backoff_for(30), Duration::from_millis(30_000));
}
#[test]
fn max_attempts_triggers_failed() {
let mut eng = ReconnectEngine::new(policy(Some(3)));
assert!(!eng.is_failed());
for _ in 0..3 {
eng.on_disconnect();
}
assert!(
eng.is_failed(),
"exceeding max_attempts must fail the engine"
);
}
#[test]
fn on_connect_resets_attempt_counter() {
let mut eng = ReconnectEngine::new(policy(Some(2)));
eng.on_disconnect();
eng.on_disconnect();
assert!(eng.is_failed());
eng.on_connect();
assert!(!eng.is_failed());
assert!(eng.should_connect());
}
#[test]
fn should_connect_respects_backoff_timing() {
let mut eng = ReconnectEngine::new(policy(None));
assert!(eng.should_connect());
eng.on_disconnect();
assert!(eng.time_until_retry() > Duration::ZERO);
assert!(!eng.should_connect() || eng.time_until_retry().is_zero());
assert!(!eng.is_failed());
}
}