use std::net::{TcpStream, ToSocketAddrs};
use std::os::unix::io::IntoRawFd;
use std::sync::{mpsc, Mutex};
use std::time::{Duration, Instant};
use crate::buffer::Buffer;
use crate::chunk::reader::{chunk_read, ChunkMessage};
use crate::chunk::state::{ChunkRegistry, DEFAULT_MAX_MSG_LENGTH};
use crate::chunk::writer::chunk_write;
use crate::handshake::{self, Handshake};
use crate::message::command;
use crate::message::control;
use crate::message::message as msg_dispatch;
use crate::net;
use crate::transport::Transport;
use crate::types::*;
const HANDSHAKE_SIZE: usize = 1536;
const RECV_POLL_TIMEOUT_MS: i32 = 10_000;
pub const MAX_CLIENT_FRAME_BYTES: usize = DEFAULT_MAX_MSG_LENGTH as usize;
const MAX_MESSAGES_PER_POLL: usize = 256;
const MAX_RECV_BYTES_PER_POLL: usize = 256 * 1024;
const MAX_RECV_BYTES_PER_COMMAND_WAIT: usize = 256 * 1024;
const TCP_CONNECT_TIMEOUT_SECS: u64 = 10;
const MAX_RECV_BUFFER_PAYLOAD_BYTES: usize = 2 * DEFAULT_MAX_MSG_LENGTH as usize;
const MIN_PRACTICAL_CHUNK_SIZE: usize = 128;
const MAX_CHUNK_HEADER_OVERHEAD_BYTES: usize = 7;
const FIRST_CHUNK_EXTRA_OVERHEAD_BYTES: usize = 11;
const MAX_RECV_BUFFER_BYTES: usize = MAX_RECV_BUFFER_PAYLOAD_BYTES
+ (MAX_RECV_BUFFER_PAYLOAD_BYTES / MIN_PRACTICAL_CHUNK_SIZE + 1) * MAX_CHUNK_HEADER_OVERHEAD_BYTES
+ 2 * FIRST_CHUNK_EXTRA_OVERHEAD_BYTES;
const MAX_DNS_QUEUE_DEPTH: usize = 32;
fn resolve_socket_addrs(
host: &str,
port: u16,
deadline: Instant,
) -> Result<Vec<std::net::SocketAddr>> {
struct DnsJob {
host: String,
port: u16,
reply: mpsc::Sender<std::result::Result<Vec<std::net::SocketAddr>, ()>>,
}
static DNS_TX: Mutex<Option<mpsc::SyncSender<DnsJob>>> = Mutex::new(None);
let tx = {
let mut guard = DNS_TX.lock().map_err(|_| ErrorCode::Internal)?;
if let Some(tx) = guard.as_ref() {
tx.clone()
} else {
let (job_tx, job_rx) = mpsc::sync_channel::<DnsJob>(MAX_DNS_QUEUE_DEPTH);
std::thread::Builder::new()
.name("lrtmp2-dns".into())
.spawn(move || {
while let Ok(job) = job_rx.recv() {
let result = (job.host.as_str(), job.port)
.to_socket_addrs()
.map(|iter| iter.collect::<Vec<_>>())
.map_err(|_| ());
let _ = job.reply.send(result);
}
})
.map_err(|_| ErrorCode::Internal)?;
*guard = Some(job_tx.clone());
job_tx
}
};
let (reply_tx, reply_rx) = mpsc::channel();
tx.try_send(DnsJob {
host: host.to_string(),
port,
reply: reply_tx,
})
.map_err(|_| ErrorCode::Internal)?;
let remaining = deadline.saturating_duration_since(Instant::now());
match reply_rx.recv_timeout(remaining) {
Ok(Ok(addrs)) if !addrs.is_empty() => Ok(addrs),
Ok(Ok(_)) | Ok(Err(_)) => Err(ErrorCode::Io),
Err(mpsc::RecvTimeoutError::Timeout) => Err(ErrorCode::Timeout),
Err(mpsc::RecvTimeoutError::Disconnected) => Err(ErrorCode::Io),
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
#[repr(C)]
pub enum ClientState {
Disconnected = 0,
Handshaking,
Connected,
AppConnected,
StreamCreated,
Publishing,
Playing,
}
pub struct Client {
pub client_fd: i32,
pub transport: Option<Transport>,
pub handshake: Handshake,
pub state: ClientState,
pub send_buffer: Buffer,
pub recv_buffer: Buffer,
pub chunk_reg: ChunkRegistry,
pub stream_id: u32,
pub app: String,
pub stream_key: String,
pub on_frame_cb: Option<fn(&Frame)>,
frame_cb_scratch: Vec<u8>,
tls_ca_file: Option<String>,
tls_insecure: bool,
}
impl Client {
pub fn new() -> Self {
Self {
client_fd: -1,
transport: None,
handshake: Handshake::default(),
state: ClientState::Disconnected,
send_buffer: Buffer::new(),
recv_buffer: Buffer::new(),
chunk_reg: ChunkRegistry::new(),
stream_id: 0,
app: String::new(),
stream_key: String::new(),
on_frame_cb: None,
frame_cb_scratch: Vec::new(),
tls_ca_file: None,
tls_insecure: false,
}
}
pub fn set_tls_client_config(&mut self, ca_file: Option<String>, insecure: bool) {
self.tls_ca_file = ca_file;
self.tls_insecure = insecure;
}
pub fn connect(&mut self, url: &str) -> Result<()> {
let (use_tls, host, port, app, stream_key) = parse_rtmp_url(url)?;
if use_tls && !crate::transport::tls_available() {
return Err(ErrorCode::Unsupported);
}
self.reset_session_state();
let deadline = Instant::now() + Duration::from_secs(TCP_CONNECT_TIMEOUT_SECS);
let addrs = resolve_socket_addrs(&host, port, deadline)?;
let mut last_err_was_timeout = false;
let mut stream = None;
for addr in addrs {
let remaining = deadline.saturating_duration_since(Instant::now());
if remaining.is_zero() {
last_err_was_timeout = true;
break;
}
match TcpStream::connect_timeout(&addr, remaining) {
Ok(s) => {
stream = Some(s);
break;
}
Err(e) => last_err_was_timeout = e.kind() == std::io::ErrorKind::TimedOut,
}
}
let stream = stream.ok_or(if last_err_was_timeout {
ErrorCode::Timeout
} else {
ErrorCode::Io
})?;
let mut transport = if use_tls {
Transport::connect_tls(
stream,
&host,
self.tls_ca_file.as_deref(),
self.tls_insecure,
)?
} else {
Transport::new_plain(stream.into_raw_fd())
};
self.state = ClientState::Handshaking;
if let Err(e) = self.do_handshake(&mut transport) {
return Err(e);
}
self.client_fd = transport.fd();
self.transport = Some(transport);
self.app = app.clone();
self.stream_key = stream_key;
self.state = ClientState::Connected;
if let Err(e) = self.do_amf_connect(&app, &host, port, use_tls) {
self.reset_session_state();
return Err(e);
}
Ok(())
}
pub fn publish(&mut self) -> Result<()> {
if self.state != ClientState::AppConnected {
return Err(ErrorCode::Protocol);
}
let mut amf = Buffer::with_capacity(256);
command::build_publish(&mut amf, &self.stream_key, "live")?;
self.send_command_msg(self.stream_id, amf.as_slice())?;
let mut status = self.wait_for_command("onStatus")?;
command::read_onstatus(&mut status)?;
self.state = ClientState::Publishing;
Ok(())
}
fn do_amf_connect(&mut self, app: &str, host: &str, port: u16, use_tls: bool) -> Result<()> {
let scheme = if use_tls { "rtmps" } else { "rtmp" };
let tc_url = format!("{scheme}://{host}:{port}/{app}");
let mut connect_amf = Buffer::with_capacity(512);
command::build_connect(&mut connect_amf, app, &tc_url, "", "", "FMLE/3.0", 0, 0)?;
self.send_command_msg(0, connect_amf.as_slice())?;
let mut result = self.wait_for_command("_result")?;
command::read_connect_result(&mut result)?;
let mut create_stream_amf = Buffer::with_capacity(64);
command::build_create_stream(&mut create_stream_amf, 2.0)?;
self.send_command_msg(0, create_stream_amf.as_slice())?;
let mut create_result = self.wait_for_command("_result")?;
let (_txn, stream_id) = command::read_create_stream_result(&mut create_result)?;
self.stream_id = stream_id as u32;
self.state = ClientState::AppConnected;
Ok(())
}
pub fn play(&mut self) -> Result<()> {
if self.state != ClientState::AppConnected {
return Err(ErrorCode::Protocol);
}
let mut amf = Buffer::with_capacity(256);
command::build_play(&mut amf, &self.stream_key)?;
self.send_command_msg(self.stream_id, amf.as_slice())?;
let mut status = self.wait_for_command("onStatus")?;
command::read_onstatus(&mut status)?;
self.state = ClientState::Playing;
Ok(())
}
pub fn send_frame(&mut self, frame: &Frame) -> Result<()> {
if self.state != ClientState::Publishing {
return Err(ErrorCode::Protocol);
}
let payload = self.frame_payload_slice(frame)?;
self.send_frame_payload(frame.frame_type, frame.timestamp, payload)
}
pub fn send_frame_payload(
&mut self,
frame_type: FrameType,
timestamp: u32,
payload: &[u8],
) -> Result<()> {
if self.state != ClientState::Publishing {
return Err(ErrorCode::Protocol);
}
self.try_flush_send_buffer()?;
self.service_inbound(0)?;
if payload.len() > MAX_CLIENT_FRAME_BYTES {
return Err(ErrorCode::Protocol);
}
let mut cmsg = ChunkMessage::default();
cmsg.timestamp = timestamp;
cmsg.msg_length = payload.len() as u32;
cmsg.msg_stream_id = self.stream_id;
if frame_type == FrameType::Audio {
cmsg.csid = 4;
cmsg.msg_type_id = 0x08; } else {
cmsg.csid = 6;
cmsg.msg_type_id = 0x09; }
cmsg.fmt = 0;
chunk_write(
&mut self.send_buffer,
&cmsg,
payload,
payload.len(),
128,
)?;
let data = self.send_buffer.peek().to_vec();
if let Some(ref mut transport) = self.transport {
transport.send(&data)?;
}
self.send_buffer.reset();
Ok(())
}
pub fn poll(&mut self, timeout_ms: i32) -> Result<()> {
if self.state == ClientState::Publishing {
let send_poll_again = self.try_flush_send_buffer()?;
if self.send_buffer.available() > 0 {
if let Some(t) = self.transport.as_ref() {
let again = send_poll_again.unwrap_or(2);
poll_for_transport_direction(t.fd(), again, timeout_ms)?;
}
self.try_flush_send_buffer()?;
}
let inbound_timeout = if self.send_buffer.available() > 0 {
0
} else {
timeout_ms
};
self.service_inbound(inbound_timeout)?;
self.try_flush_send_buffer()?;
return Ok(());
}
if self.state != ClientState::Playing {
return Err(ErrorCode::Protocol);
}
self.try_flush_send_buffer()?;
let (poll_fd, has_buffered_tls_data) = {
let Some(t) = self.transport.as_ref() else {
return Err(ErrorCode::Internal);
};
(t.fd(), t.pending() > 0)
};
let mut messages_processed = 0usize;
self.drain_ready_messages(&mut messages_processed)?;
if messages_processed == 0 && !has_buffered_tls_data {
let mut pfd = libc::pollfd {
fd: poll_fd,
events: libc::POLLIN,
revents: 0,
};
unsafe { libc::poll(&mut pfd, 1, timeout_ms) };
}
let mut buf = [0u8; 65536];
let mut bytes_drained = 0usize;
loop {
if bytes_drained >= MAX_RECV_BYTES_PER_POLL {
break;
}
let (n, again) = {
let Some(t) = self.transport.as_mut() else {
return Err(ErrorCode::Internal);
};
let mut again = 0i32;
let n = t.recv(&mut buf, &mut again);
(n, again)
};
if n > 0 {
let chunk_len = n as usize;
if self
.recv_buffer
.available()
.saturating_add(chunk_len)
> MAX_RECV_BUFFER_BYTES
{
return Err(ErrorCode::Protocol);
}
self.recv_buffer
.write(&buf[..chunk_len])
.map_err(|_| ErrorCode::Internal)?;
bytes_drained += chunk_len;
} else if n == 0 {
return Err(ErrorCode::Io);
} else if again == 2 {
let mut wpfd = libc::pollfd {
fd: poll_fd,
events: libc::POLLOUT,
revents: 0,
};
let rc = unsafe { libc::poll(&mut wpfd, 1, timeout_ms) };
if rc <= 0 {
break;
}
} else {
break;
}
}
self.drain_ready_messages(&mut messages_processed)?;
self.try_flush_send_buffer()?;
Ok(())
}
fn drain_ready_messages(&mut self, messages_processed: &mut usize) -> Result<()> {
loop {
if *messages_processed >= MAX_MESSAGES_PER_POLL {
break;
}
let mut msg = ChunkMessage::default();
let mut payload_ptr: *const u8 = std::ptr::null();
let mut payload_len = 0;
match chunk_read(
&mut self.recv_buffer,
&mut self.chunk_reg,
None,
&mut msg,
&mut payload_ptr,
&mut payload_len,
) {
Ok(1) if msg.is_complete => {
*messages_processed += 1;
if msg.msg_type_id == msg_dispatch::RTMP_MSG_SET_CHUNK_SIZE {
let payload = if payload_ptr.is_null() || payload_len == 0 {
&[][..]
} else {
unsafe { std::slice::from_raw_parts(payload_ptr, payload_len) }
};
if let Ok(cs) = control::read_set_chunk_size(payload) {
self.chunk_reg.set_all_chunk_size(cs);
}
} else if msg.msg_type_id == msg_dispatch::RTMP_MSG_USER_CONTROL {
let payload = if payload_ptr.is_null() || payload_len == 0 {
&[][..]
} else {
unsafe { std::slice::from_raw_parts(payload_ptr, payload_len) }
};
self.handle_user_control(payload)?;
} else if msg.msg_type_id == msg_dispatch::RTMP_MSG_AUDIO
|| msg.msg_type_id == msg_dispatch::RTMP_MSG_VIDEO
{
if let Some(ref cb) = self.on_frame_cb {
self.frame_cb_scratch.clear();
if !payload_ptr.is_null() && payload_len > 0 {
self.frame_cb_scratch.extend_from_slice(unsafe {
std::slice::from_raw_parts(payload_ptr, payload_len)
});
}
let frame = Frame {
frame_type: if msg.msg_type_id == msg_dispatch::RTMP_MSG_AUDIO {
FrameType::Audio
} else {
FrameType::Video
},
timestamp: msg.timestamp,
size: self.frame_cb_scratch.len() as u32,
data: self.frame_cb_scratch.as_ptr(),
..Default::default()
};
cb(&frame);
}
}
}
Ok(_) => break,
Err(_) => return Err(ErrorCode::Chunk),
}
}
Ok(())
}
fn queue_user_control_message(&mut self, payload: &[u8]) -> Result<()> {
let mut cmsg = ChunkMessage::default();
cmsg.csid = 2;
cmsg.fmt = 0;
cmsg.msg_length = payload.len() as u32;
cmsg.msg_type_id = msg_dispatch::RTMP_MSG_USER_CONTROL;
cmsg.msg_stream_id = 0;
chunk_write(
&mut self.send_buffer,
&cmsg,
payload,
payload.len(),
128,
)?;
Ok(())
}
fn try_flush_send_buffer(&mut self) -> Result<Option<i32>> {
let mut poll_again = None;
while self.send_buffer.available() > 0 {
let Some(ref mut transport) = self.transport else {
break;
};
let pending = self.send_buffer.peek();
let mut again = 0i32;
let n = transport.try_send(pending, &mut again)?;
if n == 0 {
if again != 0 {
poll_again = Some(again);
}
break;
}
self.send_buffer.drain(n);
}
Ok(if self.send_buffer.available() > 0 {
poll_again
} else {
None
})
}
fn send_user_control_message(&mut self, payload: &[u8]) -> Result<()> {
self.queue_user_control_message(payload)?;
let data = self.send_buffer.peek().to_vec();
if let Some(ref mut transport) = self.transport {
transport.send(&data)?;
}
self.send_buffer.reset();
Ok(())
}
fn send_user_control_message_nonblocking(&mut self, payload: &[u8]) -> Result<()> {
self.queue_user_control_message(payload)?;
self.try_flush_send_buffer()?;
Ok(())
}
fn handle_user_control(&mut self, payload: &[u8]) -> Result<()> {
if payload.len() < 6 {
return Ok(());
}
let (event_type, param1, _) = control::read_user_control(payload, false)?;
if event_type == control::UCTRL_PING_REQUEST {
let mut buf = Buffer::with_capacity(6);
control::write_user_control_ping_response(&mut buf, param1)?;
self.send_user_control_message_nonblocking(buf.as_slice())?;
}
Ok(())
}
fn service_inbound(&mut self, timeout_ms: i32) -> Result<()> {
let Some(t) = self.transport.as_ref() else {
return Ok(());
};
let poll_fd = t.fd();
let has_buffered_tls_data = t.pending() > 0;
let mut messages_processed = 0usize;
self.drain_ready_messages(&mut messages_processed)?;
if messages_processed == 0 && !has_buffered_tls_data {
let mut pfd = libc::pollfd {
fd: poll_fd,
events: libc::POLLIN,
revents: 0,
};
let rc = unsafe { libc::poll(&mut pfd, 1, timeout_ms) };
if rc <= 0 {
return Ok(());
}
}
let mut buf = [0u8; 4096];
let mut bytes_drained = 0usize;
loop {
if bytes_drained >= MAX_RECV_BYTES_PER_POLL {
break;
}
if messages_processed >= MAX_MESSAGES_PER_POLL {
break;
}
let (n, again) = {
let Some(t) = self.transport.as_mut() else {
return Ok(());
};
let mut again = 0i32;
let n = t.recv(&mut buf, &mut again);
(n, again)
};
if n > 0 {
let chunk_len = n as usize;
if self
.recv_buffer
.available()
.saturating_add(chunk_len)
> MAX_RECV_BUFFER_BYTES
{
return Err(ErrorCode::Protocol);
}
self
.recv_buffer
.write(&buf[..chunk_len])
.map_err(|_| ErrorCode::Internal)?;
bytes_drained += chunk_len;
self.drain_ready_messages(&mut messages_processed)?;
} else if n == 0 {
return Err(ErrorCode::Io);
} else if again == 2 {
break;
} else {
break;
}
}
Ok(())
}
fn frame_payload_slice<'a>(&self, frame: &'a Frame) -> Result<&'a [u8]> {
if frame.size == 0 {
return Ok(&[]);
}
if frame.data.is_null() {
return Err(ErrorCode::Internal);
}
let len = frame.size as usize;
if len > MAX_CLIENT_FRAME_BYTES {
return Err(ErrorCode::Protocol);
}
Ok(unsafe { std::slice::from_raw_parts(frame.data, len) })
}
fn reset_session_state(&mut self) {
self.transport = None;
self.client_fd = -1;
self.recv_buffer.reset();
self.send_buffer.reset();
self.chunk_reg.destroy();
self.chunk_reg.init();
handshake::client_init(&mut self.handshake);
self.state = ClientState::Disconnected;
self.stream_id = 0;
}
fn do_handshake(&mut self, transport: &mut Transport) -> Result<()> {
handshake::client_init(&mut self.handshake);
handshake::client_generate_c0c1(&mut self.handshake)?;
let c0c1 = self.handshake.out.peek().to_vec();
transport.send(&c0c1)?;
self.handshake.out.reset();
let s0s1 = read_exact(transport, 1 + HANDSHAKE_SIZE)?;
let mut buf = Buffer::new();
buf.write(&s0s1).map_err(|_| ErrorCode::Internal)?;
handshake::client_read_s0(&mut self.handshake, &mut buf)?;
handshake::client_read_s1(&mut self.handshake, &mut buf)?;
let c2 = self.handshake.out.peek().to_vec();
transport.send(&c2)?;
self.handshake.out.reset();
let s2 = read_exact(transport, HANDSHAKE_SIZE)?;
let mut buf2 = Buffer::new();
buf2.write(&s2).map_err(|_| ErrorCode::Internal)?;
handshake::client_read_s2(&mut self.handshake, &mut buf2)?;
Ok(())
}
fn send_command_msg(&mut self, msg_stream_id: u32, amf_data: &[u8]) -> Result<()> {
let mut cmsg = ChunkMessage::default();
cmsg.csid = 3;
cmsg.fmt = 0;
cmsg.msg_length = amf_data.len() as u32;
cmsg.msg_type_id = 0x14; cmsg.msg_stream_id = msg_stream_id;
chunk_write(&mut self.send_buffer, &cmsg, amf_data, amf_data.len(), 128)?;
let data = self.send_buffer.peek().to_vec();
if let Some(ref mut transport) = self.transport {
transport.send(&data)?;
}
self.send_buffer.reset();
Ok(())
}
fn wait_for_command(&mut self, want: &str) -> Result<Buffer> {
let mut recv_budget = MAX_RECV_BYTES_PER_COMMAND_WAIT;
for _ in 0..64 {
let (msg, payload) = self.recv_message(&mut recv_budget)?;
if msg.msg_type_id != msg_dispatch::RTMP_MSG_AMF0_COMMAND {
continue;
}
let mut buf = Buffer::from_slice(&payload);
let mut name_buf = [0u8; 64];
if command::peek_name(&mut buf, &mut name_buf).is_err() {
continue;
}
let name = std::str::from_utf8(&name_buf)
.unwrap_or("")
.trim_end_matches('\0');
if name == want {
return Ok(buf);
}
}
Err(ErrorCode::Timeout)
}
fn recv_message(&mut self, recv_budget: &mut usize) -> Result<(ChunkMessage, Vec<u8>)> {
loop {
let mut msg = ChunkMessage::default();
let mut payload_ptr: *const u8 = std::ptr::null();
let mut payload_len = 0;
match chunk_read(
&mut self.recv_buffer,
&mut self.chunk_reg,
None,
&mut msg,
&mut payload_ptr,
&mut payload_len,
) {
Ok(1) if msg.is_complete => {
let payload = if payload_ptr.is_null() || payload_len == 0 {
Vec::new()
} else {
unsafe { std::slice::from_raw_parts(payload_ptr, payload_len) }.to_vec()
};
if msg.msg_type_id == msg_dispatch::RTMP_MSG_SET_CHUNK_SIZE {
if let Ok(cs) = control::read_set_chunk_size(&payload) {
self.chunk_reg.set_all_chunk_size(cs);
}
continue;
}
if msg.msg_type_id == msg_dispatch::RTMP_MSG_USER_CONTROL {
self.handle_user_control(&payload)?;
continue;
}
return Ok((msg, payload));
}
Ok(_) => {}
Err(_) => return Err(ErrorCode::Chunk),
}
if *recv_budget == 0 {
return Err(ErrorCode::Timeout);
}
let mut tmp = [0u8; 4096];
let read_cap = tmp.len().min(*recv_budget);
let (n, again, t_fd) = {
let t = self.transport.as_mut().ok_or(ErrorCode::Internal)?;
let mut again = 0i32;
let n = t.recv(&mut tmp[..read_cap], &mut again);
(n, again, t.fd())
};
if n > 0 {
let chunk_len = n as usize;
if self
.recv_buffer
.available()
.saturating_add(chunk_len)
> MAX_RECV_BUFFER_BYTES
{
return Err(ErrorCode::Protocol);
}
*recv_budget -= chunk_len;
self.recv_buffer
.write(&tmp[..chunk_len])
.map_err(|_| ErrorCode::Internal)?;
} else if n == 0 {
return Err(ErrorCode::Io);
} else if again != 0 {
poll_for_transport_direction(t_fd, again, RECV_POLL_TIMEOUT_MS)?;
} else {
return Err(ErrorCode::Io);
}
}
}
}
fn poll_for_transport_direction(fd: i32, again: i32, timeout_ms: i32) -> Result<()> {
let events = if again == 2 {
libc::POLLOUT
} else {
libc::POLLIN
};
loop {
let mut pfd = libc::pollfd { fd, events, revents: 0 };
let rc = unsafe { libc::poll(&mut pfd, 1, timeout_ms) };
if rc == 0 {
return Err(ErrorCode::Timeout);
}
if rc < 0 {
if std::io::Error::last_os_error().raw_os_error() == Some(libc::EINTR) {
continue;
}
return Err(ErrorCode::Io);
}
return Ok(());
}
}
fn read_exact(transport: &mut Transport, n: usize) -> Result<Vec<u8>> {
let mut out = vec![0u8; n];
let mut got = 0;
while got < n {
let mut again = 0i32;
let r = transport.recv(&mut out[got..], &mut again);
if r > 0 {
got += r as usize;
} else if r == 0 {
return Err(ErrorCode::Io);
} else if again != 0 {
poll_for_transport_direction(transport.fd(), again, RECV_POLL_TIMEOUT_MS)?;
} else {
return Err(ErrorCode::Io);
}
}
Ok(out)
}
fn parse_rtmp_url(url: &str) -> Result<(bool, String, u16, String, String)> {
let (use_tls, rest, default_port) = if let Some(rest) = url.strip_prefix("rtmps://") {
(true, rest, "443")
} else if let Some(rest) = url.strip_prefix("rtmp://") {
(false, rest, "1935")
} else {
return Err(ErrorCode::Internal);
};
let (authority, path) = match rest.find('/') {
Some(i) => (&rest[..i], &rest[i + 1..]),
None => (rest, ""),
};
let mut host = String::new();
let mut port_str = String::new();
net::split_host_port(authority, &mut host, &mut port_str, default_port)?;
let port: u16 = port_str.parse().map_err(|_| ErrorCode::Internal)?;
let mut parts = path.splitn(2, '/');
let app = parts.next().unwrap_or("").to_string();
let stream_key = parts.next().unwrap_or("").to_string();
if app.is_empty() || stream_key.is_empty() {
return Err(ErrorCode::Internal);
}
Ok((use_tls, host, port, app, stream_key))
}
impl Default for Client {
fn default() -> Self {
Self::new()
}
}
impl Drop for Client {
fn drop(&mut self) {
if self.transport.is_none() && self.client_fd >= 0 {
unsafe {
libc::close(self.client_fd);
}
}
}
}
#[cfg(test)]
mod tests {
use super::*;
fn rtmp_user_control_ping_chunk(token: u32) -> Vec<u8> {
let mut payload = Buffer::with_capacity(6);
control::write_user_control_ping_request(&mut payload, token).unwrap();
let payload_len = payload.available();
let mut wire = Buffer::new();
let mut cmsg = ChunkMessage::default();
cmsg.csid = 2;
cmsg.fmt = 0;
cmsg.msg_length = payload_len as u32;
cmsg.msg_type_id = msg_dispatch::RTMP_MSG_USER_CONTROL;
cmsg.msg_stream_id = 0;
chunk_write(
&mut wire,
&cmsg,
payload.as_slice(),
payload_len,
128,
)
.unwrap();
wire.peek().to_vec()
}
#[test]
fn recv_budget_is_at_least_one_socket_read() {
assert!(MAX_RECV_BYTES_PER_POLL >= 65536);
}
#[test]
fn recv_buffer_staging_cap_covers_two_max_messages_at_min_chunk_size() {
let payload = 2 * DEFAULT_MAX_MSG_LENGTH as usize;
let chunks = payload.div_ceil(MIN_PRACTICAL_CHUNK_SIZE);
let worst_case_wire_bytes =
payload + chunks * MAX_CHUNK_HEADER_OVERHEAD_BYTES + 2 * FIRST_CHUNK_EXTRA_OVERHEAD_BYTES;
assert!(MAX_RECV_BUFFER_BYTES >= worst_case_wire_bytes);
}
#[test]
fn poll_rejects_recv_buffer_growth_past_staging_cap() {
use std::io::Write;
use std::os::unix::io::IntoRawFd;
use std::os::unix::net::UnixStream;
let (client_end, mut peer) = UnixStream::pair().unwrap();
client_end.set_nonblocking(true).unwrap();
peer.set_nonblocking(true).unwrap();
peer.write_all(&[0x01, 0x02, 0x03]).unwrap();
let mut client = Client::new();
client.state = ClientState::Playing;
client.transport = Some(Transport::new_plain(client_end.into_raw_fd()));
client
.recv_buffer
.write(&vec![0u8; MAX_RECV_BUFFER_BYTES * 2])
.unwrap();
assert_eq!(client.poll(0), Err(ErrorCode::Protocol));
}
#[test]
fn poll_drains_leftover_messages_before_enforcing_staging_cap() {
use std::io::Write;
use std::os::unix::io::IntoRawFd;
use std::os::unix::net::UnixStream;
let (client_end, mut peer) = UnixStream::pair().unwrap();
client_end.set_nonblocking(true).unwrap();
peer.set_nonblocking(true).unwrap();
peer.write_all(&[0x01, 0x02, 0x03]).unwrap();
let mut client = Client::new();
client.state = ClientState::Playing;
client.transport = Some(Transport::new_plain(client_end.into_raw_fd()));
let msg_count = MAX_RECV_BUFFER_BYTES / 13;
client
.recv_buffer
.write(&vec![0u8; msg_count * 13])
.unwrap();
assert_eq!(client.poll(0), Ok(()));
}
#[test]
fn poll_does_not_block_on_socket_readiness_when_messages_already_staged() {
use std::io::Write;
use std::os::unix::io::IntoRawFd;
use std::os::unix::net::UnixStream;
use std::time::Instant;
let (client_end, peer) = UnixStream::pair().unwrap();
client_end.set_nonblocking(true).unwrap();
let _peer = peer;
let mut client = Client::new();
client.state = ClientState::Playing;
client.transport = Some(Transport::new_plain(client_end.into_raw_fd()));
client.recv_buffer.write(&[0u8; 13]).unwrap();
let start = Instant::now();
assert_eq!(client.poll(5_000), Ok(()));
assert!(
start.elapsed() < Duration::from_millis(1_000),
"poll() blocked on socket readiness instead of draining the staged message first"
);
}
#[test]
fn command_wait_recv_budget_bounds_connect_handshake_amplification() {
assert!(MAX_RECV_BYTES_PER_COMMAND_WAIT < 64 * 4 * 1024 * 1024);
assert!(MAX_RECV_BYTES_PER_COMMAND_WAIT >= 65536);
}
#[test]
fn tcp_connect_timeout_is_bounded() {
assert!(TCP_CONNECT_TIMEOUT_SECS > 0);
assert!(TCP_CONNECT_TIMEOUT_SECS <= 30);
}
#[test]
fn tls_client_config_defaults_to_verified() {
let client = Client::new();
assert_eq!(client.tls_ca_file, None);
assert!(!client.tls_insecure);
}
#[test]
fn tls_client_config_is_stored() {
let mut client = Client::new();
client.set_tls_client_config(Some("/etc/ca.pem".to_string()), true);
assert_eq!(client.tls_ca_file.as_deref(), Some("/etc/ca.pem"));
assert!(client.tls_insecure);
}
#[test]
fn connect_refused_reports_io_not_timeout() {
let listener = std::net::TcpListener::bind("127.0.0.1:0").unwrap();
let port = listener.local_addr().unwrap().port();
drop(listener);
let mut client = Client::new();
let err = client
.connect(&format!("rtmp://127.0.0.1:{port}/live/stream"))
.unwrap_err();
assert_eq!(err, ErrorCode::Io);
}
#[test]
fn inbound_ping_requests_are_answered() {
use std::io::Read;
use std::os::unix::io::IntoRawFd;
use std::os::unix::net::UnixStream;
let (client_end, mut peer) = UnixStream::pair().unwrap();
client_end.set_nonblocking(true).unwrap();
peer.set_nonblocking(true).unwrap();
let mut client = Client::new();
client.state = ClientState::AppConnected;
client.transport = Some(Transport::new_plain(client_end.into_raw_fd()));
let mut ping = Buffer::with_capacity(6);
control::write_user_control_ping_request(&mut ping, 99).unwrap();
client.handle_user_control(ping.as_slice()).unwrap();
let mut out = [0u8; 256];
let n = peer.read(&mut out).unwrap();
assert!(n > 0);
let ping_response = control::UCTRL_PING_RESPONSE.to_be_bytes();
assert!(
out[..n].windows(2).any(|w| w == ping_response),
"peer should receive a UserControl ping response"
);
}
#[test]
fn send_frame_payload_services_inbound_pings() {
use std::io::{Read, Write};
use std::os::unix::io::IntoRawFd;
use std::os::unix::net::UnixStream;
let (client_end, mut peer) = UnixStream::pair().unwrap();
client_end.set_nonblocking(true).unwrap();
let mut client = Client::new();
client.chunk_reg.init();
client.state = ClientState::Publishing;
client.stream_id = 1;
client.transport = Some(Transport::new_plain(client_end.into_raw_fd()));
peer.write_all(&rtmp_user_control_ping_chunk(77)).unwrap();
client
.send_frame_payload(FrameType::Video, 0, &[0x17, 0x00])
.unwrap();
let mut out = [0u8; 512];
let n = peer.read(&mut out).unwrap();
assert!(n > 0);
let ping_response = control::UCTRL_PING_RESPONSE.to_be_bytes();
assert!(
out[..n].windows(2).any(|w| w == ping_response),
"send_frame_payload should answer inbound pings before sending media"
);
}
#[test]
fn publishing_poll_services_inbound_pings() {
use std::io::{Read, Write};
use std::os::unix::io::IntoRawFd;
use std::os::unix::net::UnixStream;
let (client_end, mut peer) = UnixStream::pair().unwrap();
client_end.set_nonblocking(true).unwrap();
let mut client = Client::new();
client.chunk_reg.init();
client.state = ClientState::Publishing;
client.stream_id = 1;
client.transport = Some(Transport::new_plain(client_end.into_raw_fd()));
peer.write_all(&rtmp_user_control_ping_chunk(88)).unwrap();
client.poll(0).unwrap();
let mut out = [0u8; 512];
let n = peer.read(&mut out).unwrap();
assert!(n > 0);
let ping_response = control::UCTRL_PING_RESPONSE.to_be_bytes();
assert!(
out[..n].windows(2).any(|w| w == ping_response),
"publishing poll should answer inbound pings for idle publishers"
);
}
#[test]
fn parse_rtmp_url_defaults_to_plaintext_and_port_1935() {
let (use_tls, host, port, app, stream_key) =
parse_rtmp_url("rtmp://example.com/live/streamkey").unwrap();
assert!(!use_tls);
assert_eq!(host, "example.com");
assert_eq!(port, 1935);
assert_eq!(app, "live");
assert_eq!(stream_key, "streamkey");
}
#[test]
fn parse_rtmp_url_rtmps_defaults_to_tls_and_port_443() {
let (use_tls, host, port, app, stream_key) =
parse_rtmp_url("rtmps://example.com/live/streamkey").unwrap();
assert!(use_tls);
assert_eq!(host, "example.com");
assert_eq!(port, 443);
assert_eq!(app, "live");
assert_eq!(stream_key, "streamkey");
}
#[test]
fn parse_rtmp_url_rtmps_respects_explicit_port() {
let (use_tls, host, port, _app, _stream_key) =
parse_rtmp_url("rtmps://example.com:1935/live/streamkey").unwrap();
assert!(use_tls);
assert_eq!(host, "example.com");
assert_eq!(port, 1935);
}
#[test]
fn parse_rtmp_url_rejects_unknown_scheme() {
assert_eq!(
parse_rtmp_url("http://example.com/live/streamkey"),
Err(ErrorCode::Internal)
);
}
#[test]
fn parse_rtmp_url_rejects_missing_stream_key() {
assert_eq!(
parse_rtmp_url("rtmp://example.com/live"),
Err(ErrorCode::Internal)
);
}
}