mod arq;
#[cfg(feature = "srt-encrypt")]
mod crypto;
mod egress;
mod handshake;
#[cfg(feature = "srt-encrypt")]
mod keymaterial;
mod packet;
pub use egress::SrtCaller;
pub use handshake::{HandshakeType, SrtHandshake};
pub use packet::{ControlType, SrtPacket};
pub use crate::protocol::tsdemux::{TsDemuxer, TsPayload, TsTrackKind};
use crate::inbound::{InboundProtocol, IngestContext, PublishSession};
use crate::{CodecId, MediaFrame, Result, StreamKey};
use async_trait::async_trait;
use std::net::SocketAddr;
use std::time::{Duration, Instant};
use tokio_util::sync::CancellationToken;
use tracing::{debug, info, warn};
use arq::Receiver;
const ARQ_WINDOW: usize = 1024;
const CONTROL_INTERVAL: Duration = Duration::from_millis(10);
#[derive(Debug, Clone)]
pub struct SrtHandler {
bind: SocketAddr,
key: StreamKey,
#[cfg(feature = "srt-encrypt")]
passphrase: Option<String>,
}
impl SrtHandler {
pub fn new(bind: SocketAddr, key: StreamKey) -> Self {
Self {
bind,
key,
#[cfg(feature = "srt-encrypt")]
passphrase: None,
}
}
#[cfg(feature = "srt-encrypt")]
pub fn with_passphrase(mut self, passphrase: impl Into<String>) -> Self {
self.passphrase = Some(passphrase.into());
self
}
#[cfg(feature = "srt-encrypt")]
fn answer(
&self,
datagram: &[u8],
km: &mut Option<keymaterial::KeyMaterial>,
) -> Option<Vec<u8>> {
match &self.passphrase {
Some(pass) => {
let (reply, got) = handshake::respond_with_km(datagram, pass.as_bytes())?;
if got.is_some() {
*km = got;
}
Some(reply)
}
None => handshake::respond(datagram),
}
}
#[cfg(not(feature = "srt-encrypt"))]
fn answer(&self, datagram: &[u8]) -> Option<Vec<u8>> {
handshake::respond(datagram)
}
}
fn publish_payload(demux: &mut TsDemuxer, sess: &PublishSession, payload: &[u8]) -> Result<()> {
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;
}
sess.publish_frame(frame)?;
}
Ok(())
}
#[async_trait]
impl InboundProtocol for SrtHandler {
fn name(&self) -> &'static str {
"srt"
}
async fn serve(&self, ctx: IngestContext, shutdown: CancellationToken) -> Result<()> {
use tokio::net::UdpSocket;
let socket = UdpSocket::bind(self.bind).await?;
info!(bind = %self.bind, "srt listener bound");
let mut buf = vec![0u8; 1500];
let mut peer: Option<SocketAddr> = None;
let mut caller_id: u32 = 0;
let mut session: Option<PublishSession> = None;
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);
#[cfg(feature = "srt-encrypt")]
let mut key_material: Option<keymaterial::KeyMaterial> = None;
loop {
tokio::select! {
_ = shutdown.cancelled() => break,
_ = control.tick() => {
if let Some(p) = peer {
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), p).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), p).await;
}
if let Some(sess) = session.as_ref() {
for ready in receiver.relieve() {
publish_payload(&mut demux, sess, &ready)?;
}
}
}
}
r = socket.recv_from(&mut buf) => {
let (n, from) = match r {
Ok(v) => v,
Err(e) => { warn!(error = %e, "srt recv failed"); continue; }
};
let datagram = &buf[..n];
let Some(pkt) = SrtPacket::parse(datagram) else { continue; };
match pkt {
SrtPacket::Control { control_type, .. } => {
if control_type == ControlType::Handshake {
#[cfg(feature = "srt-encrypt")]
let reply = self.answer(datagram, &mut key_material);
#[cfg(not(feature = "srt-encrypt"))]
let reply = self.answer(datagram);
if let Some(reply) = reply {
let _ = socket.send_to(&reply, from).await;
peer = Some(from);
if let Some(h) = SrtHandshake::parse(datagram) {
caller_id = h.socket_id;
}
debug!(%from, "srt handshake answered");
}
}
}
SrtPacket::Data { sequence, key_flag, payload_offset, .. } => {
if peer != Some(from) {
continue; }
if session.is_none() {
session = Some(ctx.open_publish(self.key.clone()).await?);
}
let sess = session.as_ref().unwrap();
let payload = datagram[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) {
publish_payload(&mut demux, sess, &ready)?;
}
}
}
}
}
}
if let Some(sess) = session {
sess.finish().await?;
}
Ok(())
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn handler_reports_name_and_key() {
let h = SrtHandler::new(
"127.0.0.1:9000".parse().unwrap(),
StreamKey::new("live", "feed"),
);
assert_eq!(h.name(), "srt");
assert_eq!(h.key.stream_id.as_str(), "feed");
}
use std::collections::HashSet;
use std::time::Duration;
use tokio::net::UdpSocket;
use tokio::sync::mpsc;
use tokio::time::timeout;
#[tokio::test]
async fn arq_recovers_dropped_packets() {
let listener = UdpSocket::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
let (tx, rx) = mpsc::channel(32);
let shutdown = CancellationToken::new();
let caller_task = tokio::spawn(SrtCaller::new(addr).run(rx, shutdown.clone()));
let mut buf = [0u8; 2048];
let (n, peer) = listener.recv_from(&mut buf).await.unwrap();
let dest = SrtHandshake::parse(&buf[..n]).unwrap().socket_id;
listener
.send_to(&handshake::respond(&buf[..n]).unwrap(), peer)
.await
.unwrap();
let (n, peer) = listener.recv_from(&mut buf).await.unwrap();
listener
.send_to(&handshake::respond(&buf[..n]).unwrap(), peer)
.await
.unwrap();
let count = 12usize;
for i in 0..count {
tx.send(bytes::Bytes::from(vec![i as u8; 100]))
.await
.unwrap();
}
let mut receiver = Receiver::new(64);
let mut delivered: Vec<Vec<u8>> = Vec::new();
let mut dropped_once: HashSet<u32> = HashSet::new();
while delivered.len() < count {
let (n, from) = timeout(Duration::from_secs(5), listener.recv_from(&mut buf))
.await
.expect("packet within timeout")
.unwrap();
let SrtPacket::Data {
sequence,
payload_offset,
..
} = SrtPacket::parse(&buf[..n]).unwrap()
else {
continue;
};
if (sequence == 3 || sequence == 8) && dropped_once.insert(sequence) {
let nak = arq::build_nak(&[(sequence, sequence)], 0, dest);
listener.send_to(&nak, from).await.unwrap();
continue;
}
let payload = buf[payload_offset..n].to_vec();
delivered.extend(receiver.push(sequence, payload));
}
shutdown.cancel();
let _ = caller_task.await;
let expected: Vec<Vec<u8>> = (0..count).map(|i| vec![i as u8; 100]).collect();
assert_eq!(delivered, expected, "every payload recovered, in order");
assert_eq!(dropped_once.len(), 2, "both losses were actually injected");
}
#[cfg(feature = "srt-encrypt")]
#[tokio::test]
async fn encrypted_caller_payload_decrypts_with_recovered_km() {
let pass = "swordfish-correct-horse";
let listener = UdpSocket::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
let (tx, rx) = mpsc::channel(8);
let shutdown = CancellationToken::new();
let caller_task = tokio::spawn(
SrtCaller::new(addr)
.with_passphrase(pass)
.run(rx, shutdown.clone()),
);
let mut buf = [0u8; 2048];
let (n, peer) = listener.recv_from(&mut buf).await.unwrap();
let (reply, none) = handshake::respond_with_km(&buf[..n], pass.as_bytes()).unwrap();
assert!(none.is_none(), "induction carries no key material");
listener.send_to(&reply, peer).await.unwrap();
let (n, peer) = listener.recv_from(&mut buf).await.unwrap();
let (reply, km) = handshake::respond_with_km(&buf[..n], pass.as_bytes()).unwrap();
let km = km.expect("conclusion KMREQ yields key material");
assert!(handshake::respond_with_km(&buf[..n], b"wrong").is_none());
listener.send_to(&reply, peer).await.unwrap();
let plain = vec![0x47u8; TS_BYTES_PER_DATAGRAM_TEST];
tx.send(bytes::Bytes::from(plain.clone())).await.unwrap();
let (n, _) = timeout(Duration::from_secs(5), listener.recv_from(&mut buf))
.await
.expect("data packet")
.unwrap();
let SrtPacket::Data {
sequence,
key_flag,
payload_offset,
..
} = SrtPacket::parse(&buf[..n]).unwrap()
else {
panic!("expected a data packet");
};
assert_eq!(key_flag, 1, "payload flagged as even-key encrypted");
let mut wire = buf[payload_offset..n].to_vec();
assert_ne!(wire, plain, "payload is ciphertext on the wire");
km.transform(sequence, &mut wire);
assert_eq!(wire, plain, "recovered key material decrypts the payload");
shutdown.cancel();
let _ = caller_task.await;
}
#[cfg(feature = "srt-encrypt")]
const TS_BYTES_PER_DATAGRAM_TEST: usize = 7 * 188;
}