arcly_stream/protocol/srt/
mod.rs1mod arq;
44mod conn;
45#[cfg(feature = "srt-encrypt")]
46mod crypto;
47mod egress;
48mod handshake;
49#[cfg(feature = "srt-encrypt")]
50mod keymaterial;
51mod packet;
52
53pub use egress::SrtCaller;
54pub use handshake::{HandshakeType, SrtHandshake};
55pub use packet::{ControlType, SrtPacket};
56pub use crate::protocol::tsdemux::{TsDemuxer, TsPayload, TsTrackKind};
59
60use crate::bus::PlaybackRegistry;
61use crate::inbound::{InboundProtocol, IngestContext};
62use crate::{Result, StreamKey};
63use async_trait::async_trait;
64use std::collections::HashMap;
65use std::net::SocketAddr;
66use std::sync::Arc;
67use tokio_util::sync::CancellationToken;
68use tracing::{info, warn};
69
70const ARQ_WINDOW: usize = 1024;
72const CONTROL_INTERVAL: std::time::Duration = std::time::Duration::from_millis(10);
74const CONN_QUEUE: usize = 512;
76
77#[derive(Clone)]
83pub struct SrtHandler {
84 bind: SocketAddr,
85 key: StreamKey,
87 playback: Option<Arc<dyn PlaybackRegistry>>,
90 gate: Option<crate::auth::EgressGate>,
92 #[cfg(feature = "srt-encrypt")]
94 passphrase: Option<String>,
95}
96
97fn resolve_streamid(default: &StreamKey, sid: &str) -> StreamKey {
103 let resource = sid
105 .strip_prefix("#!::")
106 .and_then(|rest| {
107 rest.split(',')
108 .find_map(|kv| kv.strip_prefix("r=").or_else(|| kv.strip_prefix("s=")))
109 })
110 .unwrap_or(sid)
111 .trim_matches('/');
112 match resource.split_once('/') {
113 Some((app, stream)) if !app.is_empty() && !stream.is_empty() => StreamKey::new(app, stream),
114 _ => StreamKey::new(default.app.clone(), resource),
115 }
116}
117
118impl SrtHandler {
119 pub fn new(bind: SocketAddr, key: StreamKey) -> Self {
122 Self {
123 bind,
124 key,
125 playback: None,
126 gate: None,
127 #[cfg(feature = "srt-encrypt")]
128 passphrase: None,
129 }
130 }
131
132 pub fn with_playback(mut self, playback: Arc<dyn PlaybackRegistry>) -> Self {
134 self.playback = Some(playback);
135 self
136 }
137
138 pub fn with_gate(mut self, gate: crate::auth::EgressGate) -> Self {
140 self.gate = Some(gate);
141 self
142 }
143
144 #[cfg(feature = "srt-encrypt")]
147 pub fn with_passphrase(mut self, passphrase: impl Into<String>) -> Self {
148 self.passphrase = Some(passphrase.into());
149 self
150 }
151}
152
153#[async_trait]
154impl InboundProtocol for SrtHandler {
155 fn name(&self) -> &'static str {
156 "srt"
157 }
158
159 async fn serve(&self, ctx: IngestContext, shutdown: CancellationToken) -> Result<()> {
160 use tokio::net::UdpSocket;
161 use tokio::sync::mpsc;
162
163 let socket = Arc::new(UdpSocket::bind(self.bind).await?);
164 info!(bind = %self.bind, "srt server listening");
165
166 let mut conns: HashMap<SocketAddr, mpsc::Sender<Vec<u8>>> = HashMap::new();
169 let mut buf = vec![0u8; 1500];
170
171 loop {
172 tokio::select! {
173 _ = shutdown.cancelled() => break,
174 r = socket.recv_from(&mut buf) => {
175 let (n, from) = match r {
176 Ok(v) => v,
177 Err(e) => { warn!(error = %e, "srt recv failed"); continue; }
178 };
179 let datagram = buf[..n].to_vec();
180
181 if let Some(tx) = conns.get(&from) {
183 if tx.try_send(datagram).is_err() {
184 conns.remove(&from);
187 }
188 continue;
189 }
190
191 if !conn::is_handshake(&datagram) {
193 continue;
194 }
195 let (tx, rx) = mpsc::channel(CONN_QUEUE);
196 let cfg = conn::ConnConfig {
197 socket: socket.clone(),
198 peer: from,
199 ctx: ctx.clone(),
200 playback: self.playback.clone(),
201 gate: self.gate.clone(),
202 default_key: self.key.clone(),
203 #[cfg(feature = "srt-encrypt")]
204 passphrase: self.passphrase.clone(),
205 shutdown: shutdown.clone(),
206 };
207 tokio::spawn(conn::run(cfg, rx));
208 let _ = tx.try_send(datagram);
209 conns.insert(from, tx);
210 }
211 }
212 }
213 Ok(())
214 }
215}
216
217#[cfg(test)]
218mod tests {
219 use super::*;
220
221 #[test]
222 fn resolve_streamid_parses_app_and_stream() {
223 let def = StreamKey::new("live", "srt");
224 let k = resolve_streamid(&def, "show/cam1");
226 assert_eq!(k.app.as_str(), "show");
227 assert_eq!(k.stream_id.as_str(), "cam1");
228 let k = resolve_streamid(&def, "#!::r=show/cam1,m=publish");
230 assert_eq!(k.app.as_str(), "show");
231 assert_eq!(k.stream_id.as_str(), "cam1");
232 let k = resolve_streamid(&def, "cam9");
234 assert_eq!(k.app.as_str(), "live");
235 assert_eq!(k.stream_id.as_str(), "cam9");
236 }
237
238 #[test]
239 fn handler_reports_name_and_key() {
240 let h = SrtHandler::new(
241 "127.0.0.1:9000".parse().unwrap(),
242 StreamKey::new("live", "feed"),
243 );
244 assert_eq!(h.name(), "srt");
245 assert_eq!(h.key.stream_id.as_str(), "feed");
246 }
247
248 use super::arq::Receiver;
249 use super::packet::SrtPacket;
250 use std::collections::HashSet;
251 use std::time::Duration;
252 use tokio::net::UdpSocket;
253 use tokio::sync::mpsc;
254 use tokio::time::timeout;
255
256 #[tokio::test]
260 async fn arq_recovers_dropped_packets() {
261 let listener = UdpSocket::bind("127.0.0.1:0").await.unwrap();
262 let addr = listener.local_addr().unwrap();
263 let (tx, rx) = mpsc::channel(32);
264 let shutdown = CancellationToken::new();
265 let caller_task = tokio::spawn(SrtCaller::new(addr).run(rx, shutdown.clone()));
266
267 let mut buf = [0u8; 2048];
269 let (n, peer) = listener.recv_from(&mut buf).await.unwrap();
270 let dest = SrtHandshake::parse(&buf[..n]).unwrap().socket_id;
271 listener
272 .send_to(&handshake::respond(&buf[..n]).unwrap(), peer)
273 .await
274 .unwrap();
275 let (n, peer) = listener.recv_from(&mut buf).await.unwrap();
276 listener
277 .send_to(&handshake::respond(&buf[..n]).unwrap(), peer)
278 .await
279 .unwrap();
280
281 let count = 12usize;
283 for i in 0..count {
284 tx.send(bytes::Bytes::from(vec![i as u8; 100]))
285 .await
286 .unwrap();
287 }
288
289 let mut receiver = Receiver::new(64);
290 let mut delivered: Vec<Vec<u8>> = Vec::new();
291 let mut dropped_once: HashSet<u32> = HashSet::new();
292 while delivered.len() < count {
293 let (n, from) = timeout(Duration::from_secs(5), listener.recv_from(&mut buf))
294 .await
295 .expect("packet within timeout")
296 .unwrap();
297 let SrtPacket::Data {
298 sequence,
299 payload_offset,
300 ..
301 } = SrtPacket::parse(&buf[..n]).unwrap()
302 else {
303 continue;
304 };
305 if (sequence == 3 || sequence == 8) && dropped_once.insert(sequence) {
307 let nak = arq::build_nak(&[(sequence, sequence)], 0, dest);
308 listener.send_to(&nak, from).await.unwrap();
309 continue;
310 }
311 let payload = buf[payload_offset..n].to_vec();
312 delivered.extend(receiver.push(sequence, payload));
313 }
314
315 shutdown.cancel();
316 let _ = caller_task.await;
317
318 let expected: Vec<Vec<u8>> = (0..count).map(|i| vec![i as u8; 100]).collect();
319 assert_eq!(delivered, expected, "every payload recovered, in order");
320 assert_eq!(dropped_once.len(), 2, "both losses were actually injected");
321 }
322
323 #[cfg(feature = "srt-encrypt")]
326 #[tokio::test]
327 async fn encrypted_caller_payload_decrypts_with_recovered_km() {
328 let pass = "swordfish-correct-horse";
329 let listener = UdpSocket::bind("127.0.0.1:0").await.unwrap();
330 let addr = listener.local_addr().unwrap();
331 let (tx, rx) = mpsc::channel(8);
332 let shutdown = CancellationToken::new();
333 let caller_task = tokio::spawn(
334 SrtCaller::new(addr)
335 .with_passphrase(pass)
336 .run(rx, shutdown.clone()),
337 );
338
339 let mut buf = [0u8; 2048];
340 let (n, peer) = listener.recv_from(&mut buf).await.unwrap();
342 let (reply, none) = handshake::respond_with_km(&buf[..n], pass.as_bytes()).unwrap();
343 assert!(none.is_none(), "induction carries no key material");
344 listener.send_to(&reply, peer).await.unwrap();
345 let (n, peer) = listener.recv_from(&mut buf).await.unwrap();
347 let (reply, km) = handshake::respond_with_km(&buf[..n], pass.as_bytes()).unwrap();
348 let km = km.expect("conclusion KMREQ yields key material");
349 assert!(handshake::respond_with_km(&buf[..n], b"wrong").is_none());
351 listener.send_to(&reply, peer).await.unwrap();
352
353 let plain = vec![0x47u8; TS_BYTES_PER_DATAGRAM_TEST];
354 tx.send(bytes::Bytes::from(plain.clone())).await.unwrap();
355
356 let (n, _) = timeout(Duration::from_secs(5), listener.recv_from(&mut buf))
357 .await
358 .expect("data packet")
359 .unwrap();
360 let SrtPacket::Data {
361 sequence,
362 key_flag,
363 payload_offset,
364 ..
365 } = SrtPacket::parse(&buf[..n]).unwrap()
366 else {
367 panic!("expected a data packet");
368 };
369 assert_eq!(key_flag, 1, "payload flagged as even-key encrypted");
370 let mut wire = buf[payload_offset..n].to_vec();
371 assert_ne!(wire, plain, "payload is ciphertext on the wire");
372 km.transform(sequence, &mut wire);
373 assert_eq!(wire, plain, "recovered key material decrypts the payload");
374
375 shutdown.cancel();
376 let _ = caller_task.await;
377 }
378
379 #[cfg(feature = "srt-encrypt")]
381 const TS_BYTES_PER_DATAGRAM_TEST: usize = 7 * 188;
382}