use crate::push::PushTransport;
use srt_runtime::HandshakeConfig;
use srt_runtime::io::SrtSocket;
#[derive(Debug, Clone, Default)]
pub struct SrtTransportConfig {
pub srt_config: HandshakeConfig,
}
pub struct SrtTransport {
socket: Option<SrtSocket>,
}
impl std::fmt::Debug for SrtTransport {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("SrtTransport")
.field("connected", &self.socket.is_some())
.finish()
}
}
#[async_trait::async_trait]
impl PushTransport for SrtTransport {
type Config = SrtTransportConfig;
type Error = srt_runtime::Error;
async fn connect(url: &str, config: &Self::Config) -> Result<Self, Self::Error> {
let addr = url.strip_prefix("srt://").unwrap_or(url);
let socket = SrtSocket::connect(addr, config.srt_config.clone()).await?;
Ok(Self {
socket: Some(socket),
})
}
async fn send(&mut self, data: &[u8]) -> Result<(), Self::Error> {
let socket = self.socket.as_mut().ok_or(srt_runtime::Error::Io {
kind: std::io::ErrorKind::NotConnected,
context: "push send",
})?;
socket.send(data).await
}
fn close(&mut self) {
self.socket = None;
}
}
#[cfg(test)]
mod tests {
use super::*;
use srt_runtime::io::SrtListener;
use std::sync::Arc;
use std::sync::atomic::{AtomicBool, Ordering};
use std::time::Duration;
#[tokio::test]
async fn srt_transport_pushes_bytes_to_a_listener() {
const PAYLOAD: &[u8] = &[0x47, 0x40, 0x00, 0x10, 0xBB, 0xCC, 0xDD, 0xEE, 0xFF, 0x00];
let received: Arc<AtomicBool> = Arc::new(AtomicBool::new(false));
let received_for_task = Arc::clone(&received);
let listener_addr = "127.0.0.1:0".parse::<std::net::SocketAddr>().unwrap();
let mut listener = SrtListener::bind(listener_addr, HandshakeConfig::default())
.await
.expect("bind");
let bound = listener.local_addr().expect("local addr");
let server = tokio::spawn(async move {
let deadline = tokio::time::Instant::now() + Duration::from_secs(30);
let mut sock = match tokio::time::timeout_at(deadline, listener.accept()).await {
Ok(Ok(s)) => s,
Ok(Err(e)) => panic!("accept failed: {e}"),
Err(_) => panic!("accept timed out"),
};
while tokio::time::Instant::now() < deadline {
match tokio::time::timeout(Duration::from_millis(200), sock.recv()).await {
Ok(Ok(Some(data))) if data.as_slice() == PAYLOAD => {
received_for_task.store(true, Ordering::SeqCst);
break;
}
Ok(Ok(Some(_))) => continue,
Ok(Ok(None)) => break,
Ok(Err(e)) => panic!("recv failed: {e}"),
Err(_) => continue,
}
}
});
let cfg = SrtTransportConfig::default();
let transport = tokio::time::timeout(
Duration::from_secs(30),
SrtTransport::connect(&format!("srt://{bound}"), &cfg),
)
.await
.expect("connect must not hang")
.expect("connect");
let mut transport = transport;
transport.send(PAYLOAD).await.expect("send");
let deadline = std::time::Instant::now() + Duration::from_secs(5);
while !received.load(Ordering::SeqCst) && std::time::Instant::now() < deadline {
tokio::time::sleep(Duration::from_millis(50)).await;
}
transport.close();
server.abort();
assert!(
received.load(Ordering::SeqCst),
"downstream SRT listener must receive the pushed bytes"
);
}
}