use std::net::{TcpStream, ToSocketAddrs};
use std::os::unix::io::IntoRawFd;
use std::sync::{Mutex, mpsc};
use std::time::{Duration, Instant};
use crate::buffer::Buffer;
use crate::chunk::reader::{ChunkMessage, chunk_read_owned};
use crate::chunk::state::{ChunkRegistry, DEFAULT_MAX_MSG_LENGTH};
use crate::chunk::writer::chunk_write;
use crate::ertmp::multitrack_media::foreach_track;
use crate::handshake::{self, Handshake};
use crate::media::{is_on_metadata_payload, populate_av_frame, populate_multitrack_frame};
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 MAX_AGGREGATE_SUBTAGS: usize = 4096;
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_INBOUND_PING_RESPONSES: usize = 8;
const INBOUND_PING_WINDOW: Duration = Duration::from_secs(1);
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::Timeout)?;
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,
inbound_ping_window_start: Option<Instant>,
inbound_ping_responses: usize,
}
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,
inbound_ping_window_start: None,
inbound_ping_responses: 0,
}
}
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,
None,
)?;
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)?;
self.try_flush_send_buffer()?;
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();
match chunk_read_owned(&mut self.recv_buffer, &mut self.chunk_reg, &mut msg) {
Ok((1, payload)) if msg.is_complete => {
*messages_processed += 1;
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);
}
} else if msg.msg_type_id == msg_dispatch::RTMP_MSG_USER_CONTROL {
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(cb) = self.on_frame_cb {
let frame_type = if msg.msg_type_id == msg_dispatch::RTMP_MSG_AUDIO {
FrameType::Audio
} else {
FrameType::Video
};
self.deliver_av_frame_cb(cb, frame_type, msg.timestamp, payload);
}
} else if msg.msg_type_id == msg_dispatch::RTMP_MSG_AMF0_DATA
|| msg.msg_type_id == msg_dispatch::RTMP_MSG_AMF3_DATA
{
let data_payload = if msg.msg_type_id == msg_dispatch::RTMP_MSG_AMF3_DATA
&& !payload.is_empty()
&& payload[0] == 0x00
{
payload[1..].to_vec()
} else {
payload
};
if let Some(cb) = self.on_frame_cb {
self.deliver_script_frame_cb(cb, msg.timestamp, &data_payload);
}
} else if msg.msg_type_id == msg_dispatch::RTMP_MSG_AGGREGATE {
self.handle_aggregate_message(msg.timestamp, &payload)?;
}
}
Ok(_) => break,
Err(_) => return Err(ErrorCode::Chunk),
}
}
Ok(())
}
fn handle_aggregate_message(&mut self, base_timestamp: u32, payload: &[u8]) -> Result<()> {
let mut pos = 0usize;
let mut have_base = false;
let mut sub_base_ts: u32 = 0;
let mut subtags = 0usize;
while pos + 11 <= payload.len() {
if subtags >= MAX_AGGREGATE_SUBTAGS {
return Err(ErrorCode::Protocol);
}
subtags += 1;
let tag_type = payload[pos];
let data_size = ((payload[pos + 1] as u32) << 16)
| ((payload[pos + 2] as u32) << 8)
| (payload[pos + 3] as u32);
let ts = ((payload[pos + 4] as u32) << 16)
| ((payload[pos + 5] as u32) << 8)
| (payload[pos + 6] as u32)
| ((payload[pos + 7] as u32) << 24);
let body = pos + 11;
let data_size = data_size as usize;
if body + data_size > payload.len() {
return Err(ErrorCode::Protocol);
}
if !have_base {
sub_base_ts = ts;
have_base = true;
}
let out_ts = base_timestamp.wrapping_add(ts.wrapping_sub(sub_base_ts));
let tag_payload = &payload[body..body + data_size];
if let Some(cb) = self.on_frame_cb {
match tag_type {
msg_dispatch::RTMP_MSG_AUDIO => {
self.deliver_av_frame_cb(
cb,
FrameType::Audio,
out_ts,
tag_payload.to_vec(),
);
}
msg_dispatch::RTMP_MSG_VIDEO => {
self.deliver_av_frame_cb(
cb,
FrameType::Video,
out_ts,
tag_payload.to_vec(),
);
}
msg_dispatch::RTMP_MSG_AMF0_DATA => {
self.deliver_script_frame_cb(cb, out_ts, tag_payload);
}
_ => {
pos = body + data_size + 4;
continue;
}
}
}
pos = body + data_size + 4;
}
Ok(())
}
fn deliver_av_frame_cb(
&mut self,
cb: fn(&Frame),
frame_type: FrameType,
timestamp: u32,
payload: Vec<u8>,
) {
let had_multitrack = foreach_track(frame_type, &payload, |track| {
self.invoke_multitrack_on_frame_cb(
cb,
frame_type,
timestamp,
track.track_id,
track.fourcc,
track.packet_type,
track.video_frame_type,
track.payload,
);
});
if !had_multitrack {
self.invoke_on_frame_cb(cb, frame_type, timestamp, u8::MAX, &payload);
}
}
fn invoke_multitrack_on_frame_cb(
&mut self,
cb: fn(&Frame),
frame_type: FrameType,
timestamp: u32,
track_id: u8,
fourcc: [u8; 4],
packet_type: u8,
video_frame_type: u8,
payload: &[u8],
) {
self.frame_cb_scratch.clear();
self.frame_cb_scratch.extend_from_slice(payload);
let mut frame = Frame {
frame_type,
timestamp,
size: self.frame_cb_scratch.len() as u32,
data: self.frame_cb_scratch.as_ptr(),
track_id,
..Default::default()
};
populate_multitrack_frame(&mut frame, fourcc, packet_type, video_frame_type);
cb(&frame);
}
fn invoke_on_frame_cb(
&mut self,
cb: fn(&Frame),
frame_type: FrameType,
timestamp: u32,
track_id: u8,
payload: &[u8],
) {
self.frame_cb_scratch.clear();
self.frame_cb_scratch.extend_from_slice(payload);
let mut frame = Frame {
frame_type,
timestamp,
size: self.frame_cb_scratch.len() as u32,
data: self.frame_cb_scratch.as_ptr(),
track_id,
..Default::default()
};
populate_av_frame(&mut frame, &self.frame_cb_scratch);
cb(&frame);
}
fn deliver_script_frame_cb(&mut self, cb: fn(&Frame), timestamp: u32, payload: &[u8]) {
let is_metadata = u8::from(is_on_metadata_payload(payload));
self.frame_cb_scratch.clear();
self.frame_cb_scratch.extend_from_slice(payload);
let frame = Frame {
frame_type: FrameType::Script,
timestamp,
size: self.frame_cb_scratch.len() as u32,
data: self.frame_cb_scratch.as_ptr(),
is_metadata,
..Default::default()
};
cb(&frame);
}
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);
}
if self.send_buffer.available() > 0 {
Ok(poll_again)
} else {
self.send_buffer.reset();
Ok(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 = ((payload[0] as u16) << 8) | (payload[1] as u16);
let (event_type, param1, param2) = if event_type == control::UCTRL_SET_BUFFER_LENGTH {
control::read_user_control(payload, true)?
} else {
let (ty, p1, _) = control::read_user_control(payload, false)?;
(ty, p1, None)
};
match event_type {
control::UCTRL_PING_REQUEST => {
let now = Instant::now();
if let Some(start) = self.inbound_ping_window_start {
if now.duration_since(start) >= INBOUND_PING_WINDOW {
self.inbound_ping_window_start = Some(now);
self.inbound_ping_responses = 0;
}
} else {
self.inbound_ping_window_start = Some(now);
}
if self.inbound_ping_responses >= MAX_INBOUND_PING_RESPONSES {
return Err(ErrorCode::Protocol);
}
self.inbound_ping_responses += 1;
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())?;
}
control::UCTRL_STREAM_BEGIN | control::UCTRL_STREAM_EOF => {}
control::UCTRL_SET_BUFFER_LENGTH => {
let _ = param2;
}
_ => {}
}
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;
self.inbound_ping_window_start = None;
self.inbound_ping_responses = 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();
match chunk_read_owned(&mut self.recv_buffer, &mut self.chunk_reg, &mut msg) {
Ok((1, payload)) if msg.is_complete => {
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 try_flush_send_buffer_shrinks_after_full_drain() {
use crate::buffer::BUFFER_RESET_CAPACITY;
use std::os::unix::io::IntoRawFd;
use std::os::unix::net::UnixStream;
let (client_end, _peer) = UnixStream::pair().unwrap();
client_end.set_nonblocking(true).unwrap();
let mut client = Client::new();
client.transport = Some(Transport::new_plain(client_end.into_raw_fd()));
let big = vec![0u8; BUFFER_RESET_CAPACITY * 4];
client.send_buffer.write(&big).unwrap();
assert!(client.send_buffer.capacity() > BUFFER_RESET_CAPACITY);
client.try_flush_send_buffer().unwrap();
assert_eq!(client.send_buffer.available(), 0);
assert!(
client.send_buffer.capacity() <= BUFFER_RESET_CAPACITY,
"send_buffer should shrink back to {BUFFER_RESET_CAPACITY} after a full flush, got {}",
client.send_buffer.capacity()
);
}
#[test]
fn frame_cb_scratch_retains_payload_after_delivery() {
let mut client = Client::new();
let video_payload = [0x17u8, 0x01, 0x02, 0x03];
let mut wire = Buffer::new();
let mut cmsg = ChunkMessage::default();
cmsg.csid = 6;
cmsg.fmt = 0;
cmsg.msg_length = video_payload.len() as u32;
cmsg.msg_type_id = msg_dispatch::RTMP_MSG_VIDEO;
cmsg.msg_stream_id = 1;
chunk_write(&mut wire, &cmsg, &video_payload, video_payload.len(), 128).unwrap();
client.recv_buffer.write(wire.peek()).unwrap();
client.on_frame_cb = Some(|_| {});
let mut messages_processed = 0;
client
.drain_ready_messages(&mut messages_processed)
.unwrap();
assert_eq!(client.frame_cb_scratch.as_slice(), &video_payload[..]);
}
#[test]
fn script_callbacks_only_mark_on_metadata_events() {
use std::sync::{LazyLock, Mutex};
static FLAGS: LazyLock<Mutex<Vec<u8>>> = LazyLock::new(|| Mutex::new(Vec::new()));
let mut client = Client::new();
FLAGS.lock().unwrap().clear();
let mut cue_point = Buffer::new();
crate::amf::amf0::write_string(&mut cue_point, "onCuePoint").unwrap();
client.deliver_script_frame_cb(
|frame| FLAGS.lock().unwrap().push(frame.is_metadata),
10,
cue_point.as_slice(),
);
let mut metadata = Buffer::new();
crate::amf::amf0::write_string(&mut metadata, "@setDataFrame").unwrap();
crate::amf::amf0::write_string(&mut metadata, "onMetaData").unwrap();
client.deliver_script_frame_cb(
|frame| FLAGS.lock().unwrap().push(frame.is_metadata),
20,
metadata.as_slice(),
);
assert_eq!(*FLAGS.lock().unwrap(), vec![0, 1]);
}
#[test]
fn drain_ready_messages_splits_multitrack_video() {
use std::sync::{LazyLock, Mutex};
static SEEN: LazyLock<Mutex<Vec<(u8, Vec<u8>)>>> = LazyLock::new(|| Mutex::new(Vec::new()));
let payload = vec![
0x86, 0x10, b'a', b'v', b'c', b'1', 0x00, 0x00, 0x00, 0x03, 0xAA, 0xBB, 0xCC, 0x01,
0x00, 0x00, 0x02, 0xDD, 0xEE,
];
let mut wire = Buffer::new();
let mut cmsg = ChunkMessage::default();
cmsg.csid = 6;
cmsg.fmt = 0;
cmsg.msg_length = payload.len() as u32;
cmsg.msg_type_id = msg_dispatch::RTMP_MSG_VIDEO;
cmsg.msg_stream_id = 1;
chunk_write(&mut wire, &cmsg, &payload, payload.len(), 128).unwrap();
let mut client = Client::new();
client.recv_buffer.write(wire.peek()).unwrap();
SEEN.lock().unwrap().clear();
client.on_frame_cb = Some(|frame| {
let data =
unsafe { std::slice::from_raw_parts(frame.data, frame.size as usize).to_vec() };
SEEN.lock().unwrap().push((frame.track_id, data));
});
let mut messages_processed = 0;
client
.drain_ready_messages(&mut messages_processed)
.unwrap();
let seen = SEEN.lock().unwrap().clone();
assert_eq!(seen.len(), 2);
assert_eq!(seen[0].0, 0);
assert_eq!(seen[0].1, vec![0xAA, 0xBB, 0xCC]);
assert_eq!(seen[1].0, 1);
assert_eq!(seen[1].1, vec![0xDD, 0xEE]);
}
#[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_rate_limit_rejects_flood() {
use std::os::unix::io::IntoRawFd;
use std::os::unix::net::UnixStream;
let (client_end, _peer) = UnixStream::pair().unwrap();
client_end.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);
for i in 0..MAX_INBOUND_PING_RESPONSES {
ping.reset();
control::write_user_control_ping_request(&mut ping, i as u32).unwrap();
client.handle_user_control(ping.as_slice()).unwrap();
}
ping.reset();
control::write_user_control_ping_request(&mut ping, 99).unwrap();
assert_eq!(
client.handle_user_control(ping.as_slice()).unwrap_err(),
ErrorCode::Protocol
);
}
#[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)
);
}
}