use std::net::SocketAddr;
use std::sync::Arc;
use std::time::Instant;
use tokio::net::UdpSocket;
use tokio::sync::mpsc;
use tokio_util::sync::CancellationToken;
use tracing::{debug, info, warn};
use super::arq::{self, Receiver, SendBuffer};
use super::packet::{build_data_packet, ControlType, SrtPacket};
use super::{handshake, resolve_streamid, ARQ_WINDOW, CONTROL_INTERVAL};
use crate::bus::PlaybackRegistry;
use crate::inbound::{IngestContext, PublishSession};
use crate::packager::{MpegTsMuxer, Muxer};
use crate::protocol::tsdemux::{TsDemuxer, TsTrackKind};
use crate::{CodecId, MediaFrame, StreamKey};
#[cfg(feature = "srt-encrypt")]
use super::keymaterial::KeyMaterial;
const TS_BYTES_PER_DATAGRAM: usize = 7 * 188;
const SEND_WINDOW: usize = 4096;
const IDLE_TIMEOUT: std::time::Duration = std::time::Duration::from_secs(30);
pub(super) struct ConnConfig {
pub socket: Arc<UdpSocket>,
pub peer: SocketAddr,
pub ctx: IngestContext,
pub playback: Option<Arc<dyn PlaybackRegistry>>,
pub gate: Option<crate::auth::EgressGate>,
pub default_key: StreamKey,
#[cfg(feature = "srt-encrypt")]
pub passphrase: Option<String>,
pub shutdown: CancellationToken,
}
pub(super) fn is_handshake(dg: &[u8]) -> bool {
matches!(
SrtPacket::parse(dg),
Some(SrtPacket::Control {
control_type: ControlType::Handshake,
..
})
)
}
#[derive(Clone, Copy, PartialEq)]
enum Mode {
Publish,
Request,
}
fn streamid_fields(sid: &str) -> (Mode, Option<String>) {
let body = sid.strip_prefix("#!::").unwrap_or(sid);
let mut mode = Mode::Publish;
let mut token = None;
for kv in body.split(',') {
if let Some(m) = kv.strip_prefix("m=") {
mode = if m.eq_ignore_ascii_case("request") {
Mode::Request
} else {
Mode::Publish
};
} else if let Some(t) = kv.strip_prefix("t=").or_else(|| kv.strip_prefix("u=")) {
if !t.is_empty() {
token = Some(t.to_string());
}
}
}
(mode, token)
}
#[cfg(feature = "srt-encrypt")]
fn answer(dg: &[u8], pass: Option<&str>, km: &mut Option<KeyMaterial>) -> Option<Vec<u8>> {
match pass {
Some(p) => {
let (reply, got) = handshake::respond_with_km(dg, p.as_bytes())?;
if got.is_some() {
*km = got;
}
Some(reply)
}
None => handshake::respond(dg),
}
}
#[cfg(not(feature = "srt-encrypt"))]
fn answer(dg: &[u8]) -> Option<Vec<u8>> {
handshake::respond(dg)
}
pub(super) async fn run(cfg: ConnConfig, mut rx: mpsc::Receiver<Vec<u8>>) {
let ConnConfig {
socket,
peer,
ctx,
playback,
gate,
default_key,
#[cfg(feature = "srt-encrypt")]
passphrase,
shutdown,
} = cfg;
let mut caller_id: u32 = 0;
let mut streamid: Option<String> = None;
#[cfg(feature = "srt-encrypt")]
let mut key_material: Option<KeyMaterial> = None;
let mut first_data: Option<Vec<u8>> = None;
'hs: loop {
let dg = tokio::select! {
_ = shutdown.cancelled() => return,
d = tokio::time::timeout(IDLE_TIMEOUT, rx.recv()) => match d {
Ok(Some(d)) => d,
_ => return, }
};
match SrtPacket::parse(&dg) {
Some(SrtPacket::Control {
control_type: ControlType::Handshake,
..
}) => {
#[cfg(feature = "srt-encrypt")]
let reply = answer(&dg, passphrase.as_deref(), &mut key_material);
#[cfg(not(feature = "srt-encrypt"))]
let reply = answer(&dg);
if let Some(reply) = reply {
let _ = socket.send_to(&reply, peer).await;
if let Some(h) = handshake::SrtHandshake::parse(&dg) {
caller_id = h.socket_id;
}
if let Some(sid) = handshake::stream_id(&dg) {
streamid = Some(sid);
break 'hs; }
}
}
Some(SrtPacket::Data { .. }) => {
first_data = Some(dg);
break 'hs;
}
_ => {} }
}
let (mode, token) = match &streamid {
Some(sid) => streamid_fields(sid),
None => (Mode::Publish, None),
};
let key = match &streamid {
Some(sid) => resolve_streamid(&default_key, sid),
None => default_key.clone(),
};
match mode {
Mode::Publish => {
run_publish(
&socket,
peer,
caller_id,
&ctx,
key,
token,
first_data,
&shutdown,
&mut rx,
#[cfg(feature = "srt-encrypt")]
key_material,
)
.await
}
Mode::Request => {
let Some(playback) = playback else {
debug!(%peer, "srt request but egress disabled");
return;
};
let allowed = match gate.as_ref() {
Some(g) => g(key.clone(), token.clone(), Some(peer)).await,
None => true,
};
if !allowed {
debug!(%peer, %key, "srt request denied by egress gate");
return;
}
run_request(
&socket,
peer,
caller_id,
playback,
key,
&shutdown,
&mut rx,
#[cfg(feature = "srt-encrypt")]
key_material,
)
.await
}
}
}
fn publish_payload(demux: &mut TsDemuxer, sess: &PublishSession, payload: &[u8]) -> bool {
for au in demux.push(payload) {
if au.codec == CodecId::Unknown {
continue;
}
let pts = au.pts_ms;
let mut frame = match au.kind {
TsTrackKind::Video => MediaFrame::new_video(pts, pts, au.data, au.codec, au.keyframe),
TsTrackKind::Audio => MediaFrame::new_audio(pts, au.data, au.codec),
};
if au.is_config {
frame.flags |= crate::FrameFlags::CONFIG;
}
if sess.publish_frame(frame).is_err() {
return false;
}
}
true
}
#[allow(clippy::too_many_arguments)]
async fn run_publish(
socket: &Arc<UdpSocket>,
peer: SocketAddr,
caller_id: u32,
ctx: &IngestContext,
key: StreamKey,
token: Option<String>,
first_data: Option<Vec<u8>>,
shutdown: &CancellationToken,
rx: &mut mpsc::Receiver<Vec<u8>>,
#[cfg(feature = "srt-encrypt")] key_material: Option<KeyMaterial>,
) {
let mut creds = crate::auth::Credentials::default();
creds.params.push(("proto".into(), "srt".into()));
creds.addr = Some(peer);
creds.token = token;
let session = match ctx.open_publish_checked(key.clone(), &creds).await {
Ok(s) => s,
Err(e) => {
debug!(%peer, %key, error = %e, "srt publish rejected");
return;
}
};
info!(%peer, %key, "srt publish started");
let mut demux = TsDemuxer::new();
let mut receiver = Receiver::new(ARQ_WINDOW);
let mut ack_no: u32 = 0;
let start = Instant::now();
let mut control = tokio::time::interval(CONTROL_INTERVAL);
control.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Delay);
let mut pending = first_data;
loop {
let dg = if let Some(d) = pending.take() {
d
} else {
tokio::select! {
_ = shutdown.cancelled() => break,
_ = control.tick() => {
let ts = start.elapsed().as_micros() as u32;
if let Some(ack) = receiver.ack_seq() {
let _ = socket.send_to(&arq::build_ack(ack_no, ack, ts, caller_id), peer).await;
ack_no = ack_no.wrapping_add(1);
}
let missing = receiver.missing();
if !missing.is_empty() {
let _ = socket.send_to(&arq::build_nak(&missing, ts, caller_id), peer).await;
}
for ready in receiver.relieve() {
if !publish_payload(&mut demux, &session, &ready) { return; }
}
continue;
}
d = tokio::time::timeout(IDLE_TIMEOUT, rx.recv()) => match d {
Ok(Some(d)) => d,
_ => break,
}
}
};
match SrtPacket::parse(&dg) {
Some(SrtPacket::Control {
control_type: ControlType::Handshake,
..
}) => {
#[cfg(feature = "srt-encrypt")]
let mut km = None;
#[cfg(feature = "srt-encrypt")]
if let Some(reply) = answer(&dg, None, &mut km) {
let _ = socket.send_to(&reply, peer).await;
}
#[cfg(not(feature = "srt-encrypt"))]
if let Some(reply) = answer(&dg) {
let _ = socket.send_to(&reply, peer).await;
}
}
Some(SrtPacket::Data {
sequence,
key_flag,
payload_offset,
..
}) => {
let payload = dg[payload_offset..].to_vec();
#[cfg(feature = "srt-encrypt")]
let payload = {
let mut payload = payload;
if key_flag != 0 {
match &key_material {
Some(km) => km.transform(sequence, &mut payload),
None => continue, }
}
payload
};
#[cfg(not(feature = "srt-encrypt"))]
if key_flag != 0 {
continue; }
for ready in receiver.push(sequence, payload) {
if !publish_payload(&mut demux, &session, &ready) {
return;
}
}
}
_ => {}
}
}
let _ = session.finish().await;
info!(%peer, %key, "srt publish stopped");
}
#[allow(clippy::too_many_arguments)]
async fn run_request(
socket: &Arc<UdpSocket>,
peer: SocketAddr,
caller_id: u32,
playback: Arc<dyn PlaybackRegistry>,
key: StreamKey,
shutdown: &CancellationToken,
rx: &mut mpsc::Receiver<Vec<u8>>,
#[cfg(feature = "srt-encrypt")] key_material: Option<KeyMaterial>,
) {
let handle = match playback.get_stream(&key) {
Ok(h) => h,
Err(e) => {
debug!(%peer, %key, error = %e, "srt request: stream unavailable");
return;
}
};
info!(%peer, %key, "srt request (egress) started");
let mut sub = handle.subscribe_resilient();
let mut mux = MpegTsMuxer::new();
if mux.start_segment().is_err() {
return;
}
let start = Instant::now();
let mut seq: u32 = 0;
let mut msg: u32 = 0;
let mut send_buf = SendBuffer::new(SEND_WINDOW);
#[cfg(feature = "srt-encrypt")]
let key_flag: u8 = if key_material.is_some() { 1 } else { 0 };
#[cfg(not(feature = "srt-encrypt"))]
let key_flag: u8 = 0;
loop {
tokio::select! {
_ = shutdown.cancelled() => break,
dg = rx.recv() => {
let Some(dg) = dg else { break };
match arq::control_type(&dg) {
Some(ControlType::Nak) => {
for pkt in send_buf.retransmit(&arq::parse_nak(&dg)) {
let _ = socket.send_to(&pkt, peer).await;
}
}
Some(ControlType::Ack) => {
if let Some(ack) = arq::parse_ack(&dg) {
send_buf.acknowledge(ack);
}
if let Some(ack_no) = arq::ack_seqno(&dg) {
let ts = start.elapsed().as_micros() as u32;
let _ = socket.send_to(&arq::build_ackack(ack_no, ts, caller_id), peer).await;
}
}
_ => {}
}
}
frame = sub.recv() => {
let Some(frame) = frame else { break }; if mux.write(&frame).is_err() {
break;
}
let Ok(Some(ts_bytes)) = mux.take_partial() else { continue };
for piece in ts_bytes.chunks(TS_BYTES_PER_DATAGRAM) {
let ts = start.elapsed().as_micros() as u32;
let payload = piece.to_vec();
#[cfg(feature = "srt-encrypt")]
let payload = {
let mut payload = payload;
if let Some(km) = &key_material {
km.transform(seq, &mut payload);
}
payload
};
let pkt = build_data_packet(seq, msg, ts, caller_id, key_flag, false, &payload);
if socket.send_to(&pkt, peer).await.is_err() {
warn!(%peer, "srt egress send failed");
return;
}
send_buf.record(seq, pkt);
seq = seq.wrapping_add(1) & 0x7FFF_FFFF;
msg = msg.wrapping_add(1) & 0x03FF_FFFF;
}
}
}
}
info!(%peer, %key, "srt request (egress) stopped");
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn streamid_fields_parse_mode_and_token() {
let (m, t) = streamid_fields("#!::r=live/test,m=publish,t=secret");
assert!(matches!(m, Mode::Publish));
assert_eq!(t.as_deref(), Some("secret"));
let (m, t) = streamid_fields("#!::r=live/test,m=request");
assert!(matches!(m, Mode::Request));
assert_eq!(t, None);
let (m, _) = streamid_fields("live/test");
assert!(matches!(m, Mode::Publish));
}
}