mod args;
#[cfg(feature = "capture")]
mod devices;
mod hls;
mod moq;
#[cfg(feature = "play")]
mod play;
mod publish;
mod rtc;
mod rtmp;
mod srt;
mod subscribe;
#[cfg(feature = "transcode")]
mod transcode;
mod web;
use args::{Command, Export, ExportSink, Import, ImportSource, Invocation, MoqSide};
use hang::moq_net;
use publish::Publish;
use subscribe::{Subscribe, SubscribeArgs};
use anyhow::Context;
use tokio::task::JoinSet;
#[cfg(feature = "jemalloc")]
#[global_allocator]
static ALLOC: moq_native::jemalloc::tikv_jemallocator::Jemalloc = moq_native::jemalloc::tikv_jemallocator::Jemalloc;
#[derive(Clone)]
struct Net {
#[cfg(feature = "iroh")]
iroh: Option<moq_native::iroh::Endpoint>,
}
impl Net {
fn client(&self, config: moq_native::ClientConfig) -> anyhow::Result<moq_native::Client> {
let client = config.init()?;
#[cfg(feature = "iroh")]
let client = match self.iroh.clone() {
Some(iroh) => client.with_iroh(iroh),
None => client,
};
Ok(client)
}
fn server(&self, config: moq_native::ServerConfig) -> anyhow::Result<moq_native::Server> {
let server = config.init()?;
#[cfg(feature = "iroh")]
let server = match self.iroh.clone() {
Some(iroh) => server.with_iroh(iroh),
None => server,
};
Ok(server)
}
}
#[tokio::main]
async fn main() -> anyhow::Result<()> {
rustls::crypto::aws_lc_rs::default_provider()
.install_default()
.expect("failed to install default crypto provider");
let cli = Invocation::parse();
cli.log.init()?;
cli.validate()?;
let mut stages = cli.stages;
if stages.len() == 1 {
match stages.remove(0) {
Command::Token(token) => {
cli.moq.reject("token")?;
return token.run();
}
#[cfg(feature = "capture")]
Command::Devices => {
cli.moq.reject("devices")?;
return devices::run().await;
}
other => stages.push(other),
}
}
cli.moq.validate()?;
let net = Net {
#[cfg(feature = "iroh")]
iroh: cli.moq.iroh.clone().bind(&cli.moq.client.quic).await?,
};
#[cfg(feature = "jemalloc")]
let jemalloc = moq_native::jemalloc::run();
#[cfg(not(feature = "jemalloc"))]
let jemalloc = std::future::pending::<anyhow::Result<()>>();
let run = async move {
if stages.len() == 1 && !stages[0].is_stageable() {
match stages.remove(0) {
#[cfg(feature = "play")]
Command::Play(args) => return run_play(cli.moq, args, net).await,
#[cfg(feature = "transcode")]
Command::Transcode(args) => return transcode::run(cli.moq, args, net).await,
_ => unreachable!("the local verbs returned before the transport was bound"),
}
}
run_stages(cli.moq, stages, net).await
};
tokio::select! {
result = run => result,
Err(err) = jemalloc => Err(err).context("jemalloc profiler failed"),
}
}
#[derive(Clone, Copy, Default)]
struct Directions {
publish: bool,
consume: bool,
}
impl Directions {
fn of(stages: &[Command]) -> Self {
Self {
publish: stages.iter().any(|stage| matches!(stage, Command::Import(_))),
consume: stages.iter().any(|stage| matches!(stage, Command::Export(_))),
}
}
}
fn spawn_moq(
moq: &MoqSide,
net: &Net,
origin: &moq_net::origin::Producer,
directions: Directions,
tasks: &mut JoinSet<anyhow::Result<()>>,
) -> anyhow::Result<Option<moq_net::bandwidth::Consumer>> {
let mut bandwidth = None;
if let Some(url) = moq.client.connect.clone() {
let mut client = net.client(moq.client.clone())?;
if directions.publish {
client = client.with_publisher(origin.consume());
}
if directions.consume {
let linger = origin.clone().with_linger(moq.client.backoff.linger());
client = client.with_subscriber(linger);
}
let reconnect = client.reconnect(url);
moq::notify_ready();
bandwidth = Some(reconnect.send_bandwidth());
tasks.spawn(async move { Ok(reconnect.closed().await?) });
}
if let Some(web_bind) = moq.server.bind.clone() {
let server = net.server(moq.server.clone())?;
let certificates = server.certificates();
moq::notify_ready();
let origin = origin.clone();
tasks.spawn(async move {
let _: () = match directions {
Directions {
publish: true,
consume: true,
} => server.serve_both(origin.consume(), origin).await?,
Directions {
publish: true,
consume: false,
} => server.serve_publish(origin.consume()).await?,
Directions {
publish: false,
consume: true,
} => server.serve_consume(origin).await?,
Directions {
publish: false,
consume: false,
} => unreachable!("a stage always needs a direction"),
};
Ok(())
});
tasks.spawn(async move { web::run_web(&web_bind, certificates).await });
}
Ok(bandwidth)
}
#[cfg(feature = "play")]
async fn run_play(moq: MoqSide, args: play::Args, net: Net) -> anyhow::Result<()> {
args.validate()?;
let origin = moq.origin()?;
let name = moq.broadcast.clone().unwrap_or_default();
let mut tasks: JoinSet<anyhow::Result<()>> = JoinSet::new();
let directions = Directions {
consume: true,
..Default::default()
};
spawn_moq(&moq, &net, &origin, directions, &mut tasks)?;
play::run(origin.consume(), name, args, tasks)
}
async fn run_stages(moq: MoqSide, stages: Vec<Command>, net: Net) -> anyhow::Result<()> {
let origin = moq.origin()?;
let mut tasks: JoinSet<anyhow::Result<()>> = JoinSet::new();
let mut locals: Vec<Publish> = Vec::new();
let bandwidth = spawn_moq(&moq, &net, &origin, Directions::of(&stages), &mut tasks)?;
let mut stdin = None;
let mut stdout = None;
for stage in stages {
let name = stage.broadcast(&moq);
match stage {
Command::Import(import) => {
if import.source.stdin_format().is_some() {
claim("stdin", &mut stdin, &name)?;
}
if let Some(publish) = spawn_import(&origin, import, name, bandwidth.clone(), &mut tasks)? {
locals.push(publish);
}
}
Command::Export(export) => {
if export.sink.stdout().is_some() {
claim("stdout", &mut stdout, &name)?;
}
spawn_export(&origin, export, name, &mut tasks)?;
}
other => unreachable!("`{}` is not a stage", other.name()),
}
}
if locals.is_empty() {
return drive(tasks).await;
}
let local = tokio::task::LocalSet::new();
supervise(&local, locals.into_iter().map(Publish::run), &mut tasks);
local.run_until(drive(tasks)).await
}
fn supervise<F>(
local: &tokio::task::LocalSet,
pipelines: impl IntoIterator<Item = F>,
tasks: &mut JoinSet<anyhow::Result<()>>,
) where
F: std::future::Future<Output = anyhow::Result<()>> + 'static,
{
for pipeline in pipelines {
let pipeline = local.spawn_local(pipeline);
tasks.spawn(async move { pipeline.await.context("pipeline panicked")? });
}
}
fn claim(stream: &str, held: &mut Option<String>, name: &str) -> anyhow::Result<()> {
if let Some(first) = held {
anyhow::bail!(
"only one stage can use {stream}, but both `{}` and `{}` do",
display_name(first),
display_name(name),
);
}
*held = Some(name.to_string());
Ok(())
}
fn display_name(name: &str) -> &str {
if name.is_empty() { "<root>" } else { name }
}
fn spawn_import(
origin: &moq_net::origin::Producer,
import: Import,
name: String,
bandwidth: Option<moq_net::bandwidth::Consumer>,
tasks: &mut JoinSet<anyhow::Result<()>>,
) -> anyhow::Result<Option<Publish>> {
#[cfg(not(feature = "capture"))]
let _ = bandwidth;
if let ImportSource::Rtc(rtc) = &import.source
&& rtc.connect.is_some()
{
reject_listener_cors(&rtc.cors, "import rtc")?;
}
anyhow::ensure!(
import.latency_max.is_none() || import.source.honors_latency_max(),
"--latency-max is not supported for this source yet; it applies to the stdin container \
formats, hls, and capture"
);
let mut local = None;
if let Some(format) = import.source.stdin_format() {
warn_if_missing_format(&name);
let broadcast = origin
.create_broadcast(&name, moq_net::broadcast::Route::new().with_announce(true))
.context("failed to create broadcast")?;
local = Some(Publish::new(broadcast, &format, import.latency_max)?);
} else {
match import.source {
ImportSource::Hls(hls) => {
warn_if_missing_format(&name);
let origin = origin.clone();
let latency_max = import.latency_max;
tasks.spawn(async move { hls::import(&origin, name, hls.playlist, latency_max).await });
}
ImportSource::Rtmp(rtmp) => {
if let Some(addr) = rtmp.listen {
let name = require_broadcast(name, "import rtmp --listen")?;
tasks.spawn(rtmp::listen_import(origin.clone(), addr, name));
} else if let Some(url) = rtmp.connect {
tasks.spawn(rtmp::connect_import(origin.clone(), url, name));
}
}
ImportSource::Srt(srt) => {
if let Some(addr) = srt.listen {
let name = require_broadcast(name, "import srt --listen")?;
tasks.spawn(srt::listen_import(origin.clone(), addr, name, srt.latency));
} else if let Some(url) = srt.connect {
tasks.spawn(srt::connect_import(origin.clone(), url, name, srt.latency));
}
}
ImportSource::Rtc(rtc) => {
if let Some(addr) = rtc.listen {
let name = require_broadcast(name, "import rtc --listen")?;
tasks.spawn(rtc::listen_import(
origin.clone(),
addr,
rtc.udp_bind,
rtc.public_addr,
rtc.cors,
name,
));
} else if let Some(url) = rtc.connect {
tasks.spawn(rtc::connect_import(origin.clone(), url, name));
}
}
#[cfg(feature = "capture")]
ImportSource::Capture(capture) => {
warn_if_missing_format(&name);
let broadcast = origin
.create_broadcast(&name, moq_net::broadcast::Route::new().with_announce(true))
.context("failed to create broadcast")?;
local = Some(Publish::capture(broadcast, &capture, bandwidth, import.latency_max)?);
}
_ => unreachable!("container formats are handled by stdin_format above"),
}
}
Ok(local)
}
fn spawn_export(
origin: &moq_net::origin::Producer,
export: Export,
name: String,
tasks: &mut JoinSet<anyhow::Result<()>>,
) -> anyhow::Result<()> {
if let ExportSink::Rtc(rtc) = &export.sink
&& rtc.connect.is_some()
{
reject_listener_cors(&rtc.cors, "export rtc")?;
}
if let Some((format, max_latency, fragment_duration)) = export.sink.stdout() {
let args = SubscribeArgs {
format,
max_latency,
fragment_duration,
catalog: export.catalog_format,
select: export.select,
};
let consumer = origin.consume();
tasks.spawn(async move { run_stdout(consumer, name, args).await });
} else {
match export.sink {
ExportSink::Hls(args) => {
let name = require_broadcast(name, "export hls")?;
tasks.spawn(hls::export(origin.consume(), args, name));
}
ExportSink::Rtmp(rtmp) => {
if let Some(addr) = rtmp.endpoint.listen {
let name = require_broadcast(name, "export rtmp --listen")?;
tasks.spawn(rtmp::listen_export(origin.consume(), addr, name, rtmp.latency_max));
} else if let Some(url) = rtmp.endpoint.connect {
tasks.spawn(rtmp::connect_export(origin.consume(), url, name, rtmp.latency_max));
}
}
ExportSink::Srt(srt) => {
if let Some(addr) = srt.listen {
let name = require_broadcast(name, "export srt --listen")?;
tasks.spawn(srt::listen_export(origin.consume(), addr, name, srt.latency));
} else if let Some(url) = srt.connect {
tasks.spawn(srt::connect_export(origin.consume(), url, name, srt.latency));
}
}
ExportSink::Rtc(rtc) => {
if let Some(addr) = rtc.listen {
let name = require_broadcast(name, "export rtc --listen")?;
tasks.spawn(rtc::listen_export(
origin.consume(),
addr,
rtc.udp_bind,
rtc.public_addr,
rtc.cors,
name,
));
} else if let Some(url) = rtc.connect {
tasks.spawn(rtc::connect_export(origin.consume(), url, name));
}
}
_ => unreachable!("container formats are handled by stdout_format above"),
}
}
Ok(())
}
async fn run_stdout(consumer: moq_net::origin::Consumer, name: String, args: SubscribeArgs) -> anyhow::Result<()> {
let catalog = args.catalog_format(&name);
consumer
.announced_broadcast(&name)
.await
.ok_or_else(|| anyhow::anyhow!("origin closed before broadcast `{name}` was announced"))?;
let source = moq_mux::Source::new(consumer, &name);
Subscribe::new(source, catalog, args).run().await
}
async fn drive(mut tasks: JoinSet<anyhow::Result<()>>) -> anyhow::Result<()> {
tasks.spawn(async {
let _ = tokio::signal::ctrl_c().await;
Ok(())
});
while let Some(res) = tasks.join_next().await {
match res {
Ok(Ok(())) => return Ok(()),
Ok(Err(err)) => return Err(err),
Err(err) if err.is_cancelled() => continue,
Err(err) => return Err(err.into()),
}
}
Ok(())
}
fn require_broadcast(name: String, endpoint: &str) -> anyhow::Result<String> {
anyhow::ensure!(
!name.is_empty(),
"`{endpoint}` requires a broadcast: pass --broadcast <name>"
);
Ok(name)
}
fn warn_if_missing_format(name: &str) {
if !name.is_empty() && moq_mux::catalog::CatalogFormat::detect(name).is_none() {
tracing::warn!(
name,
"You should append .hang to your broadcast name to make the catalog format explicit."
);
}
}
fn reject_listener_cors(cors: &crate::web::Cors, endpoint: &str) -> anyhow::Result<()> {
anyhow::ensure!(
cors.origin.is_empty(),
"`--cors-origin` only applies to `{endpoint} --listen`"
);
Ok(())
}
#[cfg(test)]
mod tests {
use super::*;
use std::future::Future;
use std::pin::Pin;
type Pipeline = Pin<Box<dyn Future<Output = anyhow::Result<()>>>>;
#[tokio::test]
async fn a_panicking_pipeline_ends_the_process() {
let local = tokio::task::LocalSet::new();
let mut tasks = JoinSet::new();
let pipelines: Vec<Pipeline> = vec![
Box::pin(async { panic!("pipeline died") }),
Box::pin(std::future::pending()),
];
supervise(&local, pipelines, &mut tasks);
let err = local.run_until(drive(tasks)).await.unwrap_err();
assert!(err.to_string().contains("pipeline panicked"), "{err}");
}
#[tokio::test]
async fn a_finished_pipeline_ends_the_process() {
let local = tokio::task::LocalSet::new();
let mut tasks = JoinSet::new();
let pipelines: Vec<Pipeline> = vec![Box::pin(async { Ok(()) }), Box::pin(std::future::pending())];
supervise(&local, pipelines, &mut tasks);
local.run_until(drive(tasks)).await.unwrap();
}
}