use std::io;
use std::net::SocketAddr;
use tokio::io::{AsyncReadExt, AsyncWriteExt};
use tokio::net::{TcpListener, TcpStream, ToSocketAddrs};
use crate::server::{ServerConfig, ServerEvent, ServerSession};
const READ_CHUNK: usize = 8192;
fn io_err(e: crate::RtmpError) -> io::Error {
io::Error::new(io::ErrorKind::InvalidData, e)
}
#[derive(Debug)]
pub struct AsyncRtmpServer {
listener: TcpListener,
config: ServerConfig,
}
impl AsyncRtmpServer {
pub async fn bind<A: ToSocketAddrs>(addr: A, config: ServerConfig) -> io::Result<Self> {
let listener = TcpListener::bind(addr).await?;
Ok(Self { listener, config })
}
pub async fn accept(&self) -> io::Result<RtmpConnection> {
let (stream, _peer) = self.listener.accept().await?;
Ok(RtmpConnection::new(
stream,
ServerSession::new(self.config.clone()),
))
}
pub fn local_addr(&self) -> io::Result<SocketAddr> {
self.listener.local_addr()
}
}
#[derive(Debug)]
pub struct RtmpConnection {
stream: TcpStream,
session: ServerSession,
closed: bool,
}
impl RtmpConnection {
fn new(stream: TcpStream, session: ServerSession) -> Self {
Self {
stream,
session,
closed: false,
}
}
pub fn peer_addr(&self) -> io::Result<SocketAddr> {
self.stream.peer_addr()
}
pub async fn next_events(&mut self) -> io::Result<Option<Vec<ServerEvent>>> {
if self.closed {
return Ok(None);
}
let mut chunk = [0u8; READ_CHUNK];
let n = self.stream.read(&mut chunk).await?;
if n == 0 {
self.closed = true;
return Ok(None);
}
let (reply, events) = match self.session.handle_data(&chunk[..n]) {
Ok(v) => v,
Err(e) => {
self.closed = true;
return Err(io_err(e));
}
};
if !reply.is_empty() {
self.stream.write_all(&reply).await?;
}
if events.iter().any(|e| matches!(e, ServerEvent::Eof)) {
self.closed = true;
}
Ok(Some(events))
}
}
#[cfg(test)]
mod tests {
use super::*;
const FIXTURE: &str = concat!(
env!("CARGO_MANIFEST_DIR"),
"/tests/fixtures/obs-publish.bin"
);
#[tokio::test]
async fn loopback_replay_of_real_publish_reaches_connected_publish_media() {
let fixture = std::fs::read(FIXTURE).expect("read tests/fixtures/obs-publish.bin");
let server = AsyncRtmpServer::bind("127.0.0.1:0", ServerConfig::default())
.await
.expect("bind ephemeral loopback port");
let addr = server.local_addr().expect("local_addr");
let client = tokio::spawn(async move {
let mut stream = TcpStream::connect(addr).await.expect("connect loopback");
stream
.write_all(&fixture)
.await
.expect("write fixture bytes");
let mut sink = [0u8; READ_CHUNK];
let mut replied_bytes = 0usize;
loop {
match stream.read(&mut sink).await {
Ok(0) | Err(_) => break,
Ok(n) => replied_bytes += n,
}
}
replied_bytes
});
let mut conn = server.accept().await.expect("accept the client connection");
let mut events = Vec::new();
while let Some(batch) = conn
.next_events()
.await
.expect("next_events must not error")
{
events.extend(batch);
}
drop(conn);
let replied_bytes = client.await.expect("client task must not panic");
const HANDSHAKE_REPLY_LEN: usize = 1 + 1536 + 1536;
assert!(
replied_bytes >= HANDSHAKE_REPLY_LEN,
"next_events must write the session's reply bytes back to the socket \
(expected at least the {HANDSHAKE_REPLY_LEN}-byte S0+S1+S2 handshake reply), \
got {replied_bytes} bytes"
);
assert!(
events
.iter()
.any(|e| matches!(e, ServerEvent::Connected { app } if app == "live")),
"must emit Connected{{app: \"live\"}} over the real socket; got {events:?}"
);
assert!(
events.iter().any(
|e| matches!(e, ServerEvent::Publish { stream_key, .. } if stream_key == "testkey")
),
"must emit Publish{{stream_key: \"testkey\", ..}} over the real socket; got {events:?}"
);
let media_count = events
.iter()
.filter(|e| matches!(e, ServerEvent::Media { .. }))
.count();
assert!(
media_count >= 1,
"must emit at least one Media event over the real socket, got {media_count}"
);
}
}