use std::collections::HashMap;
use broadcast_common::{Parse, Serialize};
use crate::RtmpError;
type Result<T> = core::result::Result<T, RtmpError>;
pub const DEFAULT_CHUNK_SIZE: u32 = 128;
pub const MAX_CHUNK_SIZE: u32 = 16 * 1024 * 1024;
pub const MAX_MESSAGE_LEN: u32 = 8 * 1024 * 1024;
pub const MAX_CSIDS: usize = 64;
pub const EXTENDED_TIMESTAMP_MARKER: u32 = 0x00FF_FFFF;
const U24_LEN: usize = 3;
const EXTENDED_TIMESTAMP_LEN: usize = 4;
const TYPE0_LEN: usize = 11;
const TYPE1_LEN: usize = 7;
const TYPE2_LEN: usize = 3;
const TYPE3_LEN: usize = 0;
const BASIC_HEADER_CSID_OFFSET: u32 = 64;
const BASIC_HEADER_2BYTE_MARKER: u8 = 0;
const BASIC_HEADER_3BYTE_MARKER: u8 = 1;
const BASIC_HEADER_FMT_SHIFT: u8 = 6;
const BASIC_HEADER_MARKER_MASK: u8 = 0x3F;
const BASIC_HEADER_1BYTE_MIN_CSID: u32 = 2;
const BASIC_HEADER_1BYTE_MAX_CSID: u32 = 63;
const BASIC_HEADER_2BYTE_MIN_CSID: u32 = 64;
const BASIC_HEADER_2BYTE_MAX_CSID: u32 = 319;
const BASIC_HEADER_3BYTE_MIN_CSID: u32 = 320;
const BASIC_HEADER_3BYTE_MAX_CSID: u32 = 65599;
fn read_u24_be(b: &[u8]) -> u32 {
(u32::from(b[0]) << 16) | (u32::from(b[1]) << 8) | u32::from(b[2])
}
fn write_u24_be(v: u32, buf: &mut [u8]) {
buf[0] = (v >> 16) as u8;
buf[1] = (v >> 8) as u8;
buf[2] = v as u8;
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
#[non_exhaustive]
pub enum Fmt {
Type0,
Type1,
Type2,
Type3,
}
impl Fmt {
#[must_use]
pub fn name(&self) -> &'static str {
match self {
Fmt::Type0 => "type 0",
Fmt::Type1 => "type 1",
Fmt::Type2 => "type 2",
Fmt::Type3 => "type 3",
}
}
pub const fn from_bits(bits: u8) -> core::result::Result<Self, RtmpError> {
match bits {
0 => Ok(Fmt::Type0),
1 => Ok(Fmt::Type1),
2 => Ok(Fmt::Type2),
3 => Ok(Fmt::Type3),
_ => Err(RtmpError::Malformed {
what: "chunk fmt (must be 0..=3)",
}),
}
}
#[must_use]
pub const fn to_bits(self) -> u8 {
match self {
Fmt::Type0 => 0,
Fmt::Type1 => 1,
Fmt::Type2 => 2,
Fmt::Type3 => 3,
}
}
}
broadcast_common::impl_spec_display!(Fmt);
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum BasicHeaderForm {
One,
Two,
Three,
}
fn basic_header_form(csid: u32) -> Result<BasicHeaderForm> {
match csid {
BASIC_HEADER_1BYTE_MIN_CSID..=BASIC_HEADER_1BYTE_MAX_CSID => Ok(BasicHeaderForm::One),
BASIC_HEADER_2BYTE_MIN_CSID..=BASIC_HEADER_2BYTE_MAX_CSID => Ok(BasicHeaderForm::Two),
BASIC_HEADER_3BYTE_MIN_CSID..=BASIC_HEADER_3BYTE_MAX_CSID => Ok(BasicHeaderForm::Three),
_ => Err(RtmpError::Malformed {
what: "chunk stream id (must be 2..=65599)",
}),
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct BasicHeader {
pub fmt: Fmt,
pub chunk_stream_id: u32,
}
impl<'a> Parse<'a> for BasicHeader {
type Error = RtmpError;
fn parse(bytes: &'a [u8]) -> Result<Self> {
if bytes.is_empty() {
return Err(RtmpError::BufferTooShort {
need: 1,
have: 0,
what: "chunk basic header",
});
}
let byte0 = bytes[0];
let fmt = Fmt::from_bits((byte0 >> BASIC_HEADER_FMT_SHIFT) & 0x03)?;
let marker = byte0 & BASIC_HEADER_MARKER_MASK;
let chunk_stream_id = match marker {
BASIC_HEADER_2BYTE_MARKER => {
if bytes.len() < 2 {
return Err(RtmpError::BufferTooShort {
need: 2,
have: bytes.len(),
what: "chunk basic header (2-byte form)",
});
}
u32::from(bytes[1]) + BASIC_HEADER_CSID_OFFSET
}
BASIC_HEADER_3BYTE_MARKER => {
if bytes.len() < 3 {
return Err(RtmpError::BufferTooShort {
need: 3,
have: bytes.len(),
what: "chunk basic header (3-byte form)",
});
}
u32::from(bytes[1]) + u32::from(bytes[2]) * 256 + BASIC_HEADER_CSID_OFFSET
}
literal => u32::from(literal),
};
Ok(BasicHeader {
fmt,
chunk_stream_id,
})
}
}
impl Serialize for BasicHeader {
type Error = RtmpError;
fn serialized_len(&self) -> usize {
match basic_header_form(self.chunk_stream_id) {
Ok(BasicHeaderForm::One) => 1,
Ok(BasicHeaderForm::Two) => 2,
Ok(BasicHeaderForm::Three) => 3,
Err(_) => 3,
}
}
fn serialize_into(&self, buf: &mut [u8]) -> Result<usize> {
let form = basic_header_form(self.chunk_stream_id)?;
let fmt_bits = self.fmt.to_bits() << BASIC_HEADER_FMT_SHIFT;
match form {
BasicHeaderForm::One => {
if buf.is_empty() {
return Err(RtmpError::BufferTooShort {
need: 1,
have: 0,
what: "chunk basic header output (1-byte form)",
});
}
buf[0] = fmt_bits | (self.chunk_stream_id as u8);
Ok(1)
}
BasicHeaderForm::Two => {
if buf.len() < 2 {
return Err(RtmpError::BufferTooShort {
need: 2,
have: buf.len(),
what: "chunk basic header output (2-byte form)",
});
}
buf[0] = fmt_bits | BASIC_HEADER_2BYTE_MARKER;
buf[1] = (self.chunk_stream_id - BASIC_HEADER_CSID_OFFSET) as u8;
Ok(2)
}
BasicHeaderForm::Three => {
if buf.len() < 3 {
return Err(RtmpError::BufferTooShort {
need: 3,
have: buf.len(),
what: "chunk basic header output (3-byte form)",
});
}
buf[0] = fmt_bits | BASIC_HEADER_3BYTE_MARKER;
let rel = self.chunk_stream_id - BASIC_HEADER_CSID_OFFSET;
buf[1] = rel as u8;
buf[2] = (rel >> 8) as u8;
Ok(3)
}
}
}
}
fn needs_extended_timestamp(field: u32) -> bool {
field >= EXTENDED_TIMESTAMP_MARKER
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
#[non_exhaustive]
pub enum MessageHeader {
Type0 {
timestamp: u32,
message_length: u32,
message_type_id: u8,
message_stream_id: u32,
},
Type1 {
timestamp_delta: u32,
message_length: u32,
message_type_id: u8,
},
Type2 {
timestamp_delta: u32,
},
Type3,
}
impl MessageHeader {
pub fn parse(fmt: Fmt, bytes: &[u8]) -> Result<(Self, usize)> {
match fmt {
Fmt::Type0 => {
if bytes.len() < TYPE0_LEN {
return Err(RtmpError::BufferTooShort {
need: TYPE0_LEN,
have: bytes.len(),
what: "type 0 message header",
});
}
let raw_timestamp = read_u24_be(&bytes[0..U24_LEN]);
let message_length = read_u24_be(&bytes[U24_LEN..2 * U24_LEN]);
let message_type_id = bytes[2 * U24_LEN];
let message_stream_id =
u32::from_le_bytes([bytes[7], bytes[8], bytes[9], bytes[10]]);
let (timestamp, consumed) = resolve_extended(raw_timestamp, bytes, TYPE0_LEN)?;
Ok((
MessageHeader::Type0 {
timestamp,
message_length,
message_type_id,
message_stream_id,
},
consumed,
))
}
Fmt::Type1 => {
if bytes.len() < TYPE1_LEN {
return Err(RtmpError::BufferTooShort {
need: TYPE1_LEN,
have: bytes.len(),
what: "type 1 message header",
});
}
let raw_delta = read_u24_be(&bytes[0..U24_LEN]);
let message_length = read_u24_be(&bytes[U24_LEN..2 * U24_LEN]);
let message_type_id = bytes[2 * U24_LEN];
let (timestamp_delta, consumed) = resolve_extended(raw_delta, bytes, TYPE1_LEN)?;
Ok((
MessageHeader::Type1 {
timestamp_delta,
message_length,
message_type_id,
},
consumed,
))
}
Fmt::Type2 => {
if bytes.len() < TYPE2_LEN {
return Err(RtmpError::BufferTooShort {
need: TYPE2_LEN,
have: bytes.len(),
what: "type 2 message header",
});
}
let raw_delta = read_u24_be(&bytes[0..U24_LEN]);
let (timestamp_delta, consumed) = resolve_extended(raw_delta, bytes, TYPE2_LEN)?;
Ok((MessageHeader::Type2 { timestamp_delta }, consumed))
}
Fmt::Type3 => Ok((MessageHeader::Type3, TYPE3_LEN)),
}
}
}
fn resolve_extended(raw: u32, bytes: &[u8], fixed_len: usize) -> Result<(u32, usize)> {
if raw == EXTENDED_TIMESTAMP_MARKER {
let need = fixed_len + EXTENDED_TIMESTAMP_LEN;
if bytes.len() < need {
return Err(RtmpError::BufferTooShort {
need,
have: bytes.len(),
what: "extended timestamp",
});
}
let ext = u32::from_be_bytes([
bytes[fixed_len],
bytes[fixed_len + 1],
bytes[fixed_len + 2],
bytes[fixed_len + 3],
]);
Ok((ext, need))
} else {
Ok((raw, fixed_len))
}
}
impl Serialize for MessageHeader {
type Error = RtmpError;
fn serialized_len(&self) -> usize {
match self {
MessageHeader::Type0 { timestamp, .. } => TYPE0_LEN + extended_len(*timestamp),
MessageHeader::Type1 {
timestamp_delta, ..
} => TYPE1_LEN + extended_len(*timestamp_delta),
MessageHeader::Type2 { timestamp_delta } => TYPE2_LEN + extended_len(*timestamp_delta),
MessageHeader::Type3 => TYPE3_LEN,
}
}
fn serialize_into(&self, buf: &mut [u8]) -> Result<usize> {
match *self {
MessageHeader::Type0 {
timestamp,
message_length,
message_type_id,
message_stream_id,
} => {
let extended = needs_extended_timestamp(timestamp);
let written = TYPE0_LEN + if extended { EXTENDED_TIMESTAMP_LEN } else { 0 };
if buf.len() < written {
return Err(RtmpError::BufferTooShort {
need: written,
have: buf.len(),
what: "type 0 message header output",
});
}
let field = if extended {
EXTENDED_TIMESTAMP_MARKER
} else {
timestamp
};
write_u24_be(field, &mut buf[0..U24_LEN]);
write_u24_be(message_length, &mut buf[U24_LEN..2 * U24_LEN]);
buf[2 * U24_LEN] = message_type_id;
buf[7..11].copy_from_slice(&message_stream_id.to_le_bytes());
if extended {
buf[11..15].copy_from_slice(×tamp.to_be_bytes());
}
Ok(written)
}
MessageHeader::Type1 {
timestamp_delta,
message_length,
message_type_id,
} => {
let extended = needs_extended_timestamp(timestamp_delta);
let written = TYPE1_LEN + if extended { EXTENDED_TIMESTAMP_LEN } else { 0 };
if buf.len() < written {
return Err(RtmpError::BufferTooShort {
need: written,
have: buf.len(),
what: "type 1 message header output",
});
}
let field = if extended {
EXTENDED_TIMESTAMP_MARKER
} else {
timestamp_delta
};
write_u24_be(field, &mut buf[0..U24_LEN]);
write_u24_be(message_length, &mut buf[U24_LEN..2 * U24_LEN]);
buf[2 * U24_LEN] = message_type_id;
if extended {
buf[7..11].copy_from_slice(×tamp_delta.to_be_bytes());
}
Ok(written)
}
MessageHeader::Type2 { timestamp_delta } => {
let extended = needs_extended_timestamp(timestamp_delta);
let written = TYPE2_LEN + if extended { EXTENDED_TIMESTAMP_LEN } else { 0 };
if buf.len() < written {
return Err(RtmpError::BufferTooShort {
need: written,
have: buf.len(),
what: "type 2 message header output",
});
}
let field = if extended {
EXTENDED_TIMESTAMP_MARKER
} else {
timestamp_delta
};
write_u24_be(field, &mut buf[0..U24_LEN]);
if extended {
buf[3..7].copy_from_slice(×tamp_delta.to_be_bytes());
}
Ok(written)
}
MessageHeader::Type3 => Ok(0),
}
}
}
fn extended_len(field: u32) -> usize {
if needs_extended_timestamp(field) {
EXTENDED_TIMESTAMP_LEN
} else {
0
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct Message {
pub chunk_stream_id: u32,
pub timestamp: u32,
pub message_type_id: u8,
pub message_stream_id: u32,
pub payload: Vec<u8>,
}
#[derive(Debug, Clone, Default)]
struct CsidState {
timestamp: u32,
timestamp_delta: u32,
message_length: u32,
message_type_id: u8,
message_stream_id: u32,
extended: bool,
initialized: bool,
in_progress: bool,
payload: Vec<u8>,
}
#[derive(Debug)]
pub struct ChunkAssembler {
chunk_size: u32,
csids: HashMap<u32, CsidState>,
pending: Vec<u8>,
}
impl Default for ChunkAssembler {
fn default() -> Self {
Self::new()
}
}
impl ChunkAssembler {
#[must_use]
pub fn new() -> Self {
Self {
chunk_size: DEFAULT_CHUNK_SIZE,
csids: HashMap::new(),
pending: Vec::new(),
}
}
pub fn set_chunk_size(&mut self, n: u32) {
self.chunk_size = n.clamp(1, MAX_CHUNK_SIZE);
}
pub fn push(&mut self, input: &[u8]) -> Result<Vec<Message>> {
self.feed(input);
let mut out = Vec::new();
while let Some(msg) = self.next_message()? {
out.push(msg);
}
Ok(out)
}
pub(crate) fn feed(&mut self, input: &[u8]) {
self.pending.extend_from_slice(input);
}
pub(crate) fn next_message(&mut self) -> Result<Option<Message>> {
loop {
match Self::try_parse_one(&self.pending, &self.csids, self.chunk_size) {
Ok(Some(parsed)) => {
self.pending.drain(..parsed.consumed);
let state = self.csids.entry(parsed.csid).or_default();
state.timestamp = parsed.timestamp;
state.timestamp_delta = parsed.timestamp_delta;
state.message_length = parsed.message_length;
state.message_type_id = parsed.message_type_id;
state.message_stream_id = parsed.message_stream_id;
state.extended = parsed.extended;
state.initialized = true;
if parsed.payload.len() as u32 == parsed.message_length {
state.payload.clear();
state.in_progress = false;
return Ok(Some(Message {
chunk_stream_id: parsed.csid,
timestamp: parsed.timestamp,
message_type_id: parsed.message_type_id,
message_stream_id: parsed.message_stream_id,
payload: parsed.payload,
}));
}
state.payload = parsed.payload;
state.in_progress = true;
}
Ok(None) => return Ok(None),
Err(e) => return Err(e),
}
}
}
fn try_parse_one(
buf: &[u8],
states: &HashMap<u32, CsidState>,
chunk_size: u32,
) -> Result<Option<ParsedChunk>> {
let bh = match BasicHeader::parse(buf) {
Ok(bh) => bh,
Err(RtmpError::BufferTooShort { .. }) => return Ok(None),
Err(e) => return Err(e),
};
let marker = buf[0] & BASIC_HEADER_MARKER_MASK;
let header_len = match marker {
BASIC_HEADER_2BYTE_MARKER => 2,
BASIC_HEADER_3BYTE_MARKER => 3,
_ => 1,
};
let existing = states.get(&bh.chunk_stream_id);
if existing.is_none() && states.len() >= MAX_CSIDS {
return Err(RtmpError::Malformed {
what: "too many concurrent chunk stream ids (csid flood)",
});
}
let rest = &buf[header_len..];
let (mh, mh_consumed) = match MessageHeader::parse(bh.fmt, rest) {
Ok(v) => v,
Err(RtmpError::BufferTooShort { .. }) => return Ok(None),
Err(e) => return Err(e),
};
let mut consumed = header_len + mh_consumed;
let (resolved, starts_new) = match (bh.fmt, mh) {
(
Fmt::Type0,
MessageHeader::Type0 {
timestamp,
message_length,
message_type_id,
message_stream_id,
},
) => {
let used_extended = mh_consumed > TYPE0_LEN;
(
ResolvedHeader {
timestamp_delta: timestamp,
timestamp,
message_length,
message_type_id,
message_stream_id,
extended: used_extended,
},
true,
)
}
(
Fmt::Type1,
MessageHeader::Type1 {
timestamp_delta,
message_length,
message_type_id,
},
) => {
let existing = existing.ok_or(RtmpError::Malformed {
what: "type 1 chunk header on a csid with no prior chunk to inherit from",
})?;
let used_extended = mh_consumed > TYPE1_LEN;
(
ResolvedHeader {
timestamp: existing.timestamp.wrapping_add(timestamp_delta),
timestamp_delta,
message_length,
message_type_id,
message_stream_id: existing.message_stream_id,
extended: used_extended,
},
true,
)
}
(Fmt::Type2, MessageHeader::Type2 { timestamp_delta }) => {
let existing = existing.ok_or(RtmpError::Malformed {
what: "type 2 chunk header on a csid with no prior chunk to inherit from",
})?;
let used_extended = mh_consumed > TYPE2_LEN;
(
ResolvedHeader {
timestamp: existing.timestamp.wrapping_add(timestamp_delta),
timestamp_delta,
message_length: existing.message_length,
message_type_id: existing.message_type_id,
message_stream_id: existing.message_stream_id,
extended: used_extended,
},
true,
)
}
(Fmt::Type3, MessageHeader::Type3) => {
let existing = existing.ok_or(RtmpError::Malformed {
what: "type 3 chunk header on a csid with no prior chunk to inherit from",
})?;
let continuation = existing.in_progress;
if existing.extended {
if buf.len() < consumed + EXTENDED_TIMESTAMP_LEN {
return Ok(None);
}
consumed += EXTENDED_TIMESTAMP_LEN;
}
if continuation {
(
ResolvedHeader {
timestamp: existing.timestamp,
timestamp_delta: existing.timestamp_delta,
message_length: existing.message_length,
message_type_id: existing.message_type_id,
message_stream_id: existing.message_stream_id,
extended: existing.extended,
},
false,
)
} else {
(
ResolvedHeader {
timestamp: existing.timestamp.wrapping_add(existing.timestamp_delta),
timestamp_delta: existing.timestamp_delta,
message_length: existing.message_length,
message_type_id: existing.message_type_id,
message_stream_id: existing.message_stream_id,
extended: existing.extended,
},
true,
)
}
}
_ => unreachable!("MessageHeader::parse always returns the variant for its Fmt"),
};
if resolved.message_length > MAX_MESSAGE_LEN {
return Err(RtmpError::Malformed {
what: "message length exceeds the maximum accepted message size",
});
}
let already_accumulated = if starts_new {
0
} else {
existing.map(|s| s.payload.len()).unwrap_or(0)
};
let remaining_needed =
(resolved.message_length as usize).saturating_sub(already_accumulated);
let take = (chunk_size as usize).min(remaining_needed);
if buf.len() < consumed + take {
return Ok(None);
}
let mut payload = if starts_new {
Vec::new()
} else {
existing.map(|s| s.payload.clone()).unwrap_or_default()
};
payload.extend_from_slice(&buf[consumed..consumed + take]);
consumed += take;
Ok(Some(ParsedChunk {
csid: bh.chunk_stream_id,
consumed,
timestamp: resolved.timestamp,
timestamp_delta: resolved.timestamp_delta,
message_length: resolved.message_length,
message_type_id: resolved.message_type_id,
message_stream_id: resolved.message_stream_id,
extended: resolved.extended,
payload,
}))
}
}
struct ResolvedHeader {
timestamp: u32,
timestamp_delta: u32,
message_length: u32,
message_type_id: u8,
message_stream_id: u32,
extended: bool,
}
struct ParsedChunk {
csid: u32,
consumed: usize,
timestamp: u32,
timestamp_delta: u32,
message_length: u32,
message_type_id: u8,
message_stream_id: u32,
extended: bool,
payload: Vec<u8>,
}
#[derive(Debug, Clone)]
pub struct ChunkWriter {
chunk_size: u32,
}
impl Default for ChunkWriter {
fn default() -> Self {
Self::new()
}
}
impl ChunkWriter {
#[must_use]
pub fn new() -> Self {
Self {
chunk_size: DEFAULT_CHUNK_SIZE,
}
}
pub fn set_chunk_size(&mut self, n: u32) {
self.chunk_size = n.min(MAX_CHUNK_SIZE);
}
#[must_use]
pub fn write(&mut self, msg: &Message) -> Vec<u8> {
let chunk_size = (self.chunk_size as usize).max(1);
let message_length = msg.payload.len() as u32;
let extended = needs_extended_timestamp(msg.timestamp);
let mut out = Vec::with_capacity(TYPE0_LEN + msg.payload.len() + 16);
let bh0 = BasicHeader {
fmt: Fmt::Type0,
chunk_stream_id: msg.chunk_stream_id,
};
let mh0 = MessageHeader::Type0 {
timestamp: msg.timestamp,
message_length,
message_type_id: msg.message_type_id,
message_stream_id: msg.message_stream_id,
};
write_serialized(&mut out, &bh0);
write_serialized(&mut out, &mh0);
let mut offset = 0usize;
let take0 = chunk_size.min(msg.payload.len());
out.extend_from_slice(&msg.payload[offset..offset + take0]);
offset += take0;
while offset < msg.payload.len() {
let bh = BasicHeader {
fmt: Fmt::Type3,
chunk_stream_id: msg.chunk_stream_id,
};
write_serialized(&mut out, &bh);
if extended {
out.extend_from_slice(&msg.timestamp.to_be_bytes());
}
let take = chunk_size.min(msg.payload.len() - offset);
out.extend_from_slice(&msg.payload[offset..offset + take]);
offset += take;
}
out
}
}
fn write_serialized<T: Serialize<Error = RtmpError>>(out: &mut Vec<u8>, item: &T) {
let len = item.serialized_len();
let start = out.len();
out.resize(start + len, 0);
let n = item
.serialize_into(&mut out[start..])
.expect("valid chunk_stream_id (2..=65599) is a ChunkWriter::write precondition");
out.truncate(start + n);
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn u24_round_trip_zero() {
let mut buf = [0xFFu8; U24_LEN];
write_u24_be(0, &mut buf);
assert_eq!(buf, [0, 0, 0]);
assert_eq!(read_u24_be(&buf), 0);
}
#[test]
fn u24_round_trip_max() {
let mut buf = [0u8; U24_LEN];
write_u24_be(0x00FF_FFFF, &mut buf);
assert_eq!(buf, [0xFF, 0xFF, 0xFF]);
assert_eq!(read_u24_be(&buf), 0x00FF_FFFF);
}
#[test]
fn u24_round_trip_mid_value() {
let mut buf = [0u8; U24_LEN];
write_u24_be(0x0012_3456, &mut buf);
assert_eq!(buf, [0x12, 0x34, 0x56]);
assert_eq!(read_u24_be(&buf), 0x0012_3456);
}
#[test]
fn fmt_from_bits_round_trip() {
for (bits, fmt) in [
(0u8, Fmt::Type0),
(1, Fmt::Type1),
(2, Fmt::Type2),
(3, Fmt::Type3),
] {
let parsed = Fmt::from_bits(bits).unwrap();
assert_eq!(parsed, fmt);
assert_eq!(parsed.to_bits(), bits);
}
}
#[test]
fn fmt_from_bits_out_of_range_is_malformed() {
assert!(matches!(
Fmt::from_bits(4),
Err(RtmpError::Malformed { .. })
));
}
#[test]
fn fmt_display_matches_name() {
assert_eq!(Fmt::Type0.to_string(), "type 0");
assert_eq!(Fmt::Type3.to_string(), "type 3");
}
#[test]
fn basic_header_one_byte_form_round_trip_build_serialize_parse() {
for csid in [BASIC_HEADER_1BYTE_MIN_CSID, 5, BASIC_HEADER_1BYTE_MAX_CSID] {
let bh = BasicHeader {
fmt: Fmt::Type1,
chunk_stream_id: csid,
};
let mut buf = [0u8; 1];
let n = bh.serialize_into(&mut buf).unwrap();
assert_eq!(n, 1, "csid {csid} must use the 1-byte form");
let parsed = BasicHeader::parse(&buf).unwrap();
assert_eq!(parsed, bh);
}
}
#[test]
fn basic_header_one_byte_form_parse_serialize_byte_identical() {
let bytes = [0x45u8];
let bh = BasicHeader::parse(&bytes).unwrap();
assert_eq!(bh.fmt, Fmt::Type1);
assert_eq!(bh.chunk_stream_id, 5);
let mut buf = [0u8; 1];
bh.serialize_into(&mut buf).unwrap();
assert_eq!(buf, bytes);
}
#[test]
fn basic_header_two_byte_form_round_trip_boundaries() {
for csid in [
BASIC_HEADER_2BYTE_MIN_CSID,
200,
BASIC_HEADER_2BYTE_MAX_CSID,
] {
let bh = BasicHeader {
fmt: Fmt::Type2,
chunk_stream_id: csid,
};
let mut buf = [0u8; 2];
let n = bh.serialize_into(&mut buf).unwrap();
assert_eq!(n, 2, "csid {csid} must use the minimal 2-byte form");
let parsed = BasicHeader::parse(&buf).unwrap();
assert_eq!(parsed, bh);
}
}
#[test]
fn basic_header_two_byte_form_parse_serialize_byte_identical() {
let bytes = [0b10_000000u8, 0x00];
let bh = BasicHeader::parse(&bytes).unwrap();
assert_eq!(bh.fmt, Fmt::Type2);
assert_eq!(bh.chunk_stream_id, 64);
let mut buf = [0u8; 2];
bh.serialize_into(&mut buf).unwrap();
assert_eq!(buf, bytes);
}
#[test]
fn basic_header_three_byte_form_round_trip_boundaries() {
for csid in [
BASIC_HEADER_3BYTE_MIN_CSID,
40000,
BASIC_HEADER_3BYTE_MAX_CSID,
] {
let bh = BasicHeader {
fmt: Fmt::Type3,
chunk_stream_id: csid,
};
let mut buf = [0u8; 3];
let n = bh.serialize_into(&mut buf).unwrap();
assert_eq!(n, 3, "csid {csid} must use the 3-byte form");
let parsed = BasicHeader::parse(&buf).unwrap();
assert_eq!(parsed, bh);
}
}
#[test]
fn basic_header_three_byte_form_parse_serialize_byte_identical() {
let bytes = [0b00_000001u8, 0xFF, 0xFF];
let bh = BasicHeader::parse(&bytes).unwrap();
assert_eq!(bh.fmt, Fmt::Type0);
assert_eq!(bh.chunk_stream_id, BASIC_HEADER_3BYTE_MAX_CSID);
let mut buf = [0u8; 3];
bh.serialize_into(&mut buf).unwrap();
assert_eq!(buf, bytes);
}
#[test]
fn basic_header_2byte_and_3byte_csid_are_little_endian() {
let bh = BasicHeader {
fmt: Fmt::Type0,
chunk_stream_id: 64 + 0x0102,
};
let mut buf = [0u8; 3];
bh.serialize_into(&mut buf).unwrap();
assert_eq!(buf[1], 0x02, "low byte of csid-64 must come first");
assert_eq!(buf[2], 0x01, "high byte of csid-64 must come second");
assert_eq!(BasicHeader::parse(&buf).unwrap(), bh);
}
#[test]
fn basic_header_csid_zero_is_malformed_on_serialize() {
let bh = BasicHeader {
fmt: Fmt::Type0,
chunk_stream_id: 0,
};
let mut buf = [0u8; 3];
assert!(matches!(
bh.serialize_into(&mut buf),
Err(RtmpError::Malformed { .. })
));
}
#[test]
fn basic_header_csid_one_is_malformed_on_serialize() {
let bh = BasicHeader {
fmt: Fmt::Type0,
chunk_stream_id: 1,
};
let mut buf = [0u8; 3];
assert!(matches!(
bh.serialize_into(&mut buf),
Err(RtmpError::Malformed { .. })
));
}
#[test]
fn basic_header_csid_above_max_is_malformed_on_serialize() {
let bh = BasicHeader {
fmt: Fmt::Type0,
chunk_stream_id: BASIC_HEADER_3BYTE_MAX_CSID + 1,
};
let mut buf = [0u8; 3];
assert!(matches!(
bh.serialize_into(&mut buf),
Err(RtmpError::Malformed { .. })
));
}
#[test]
fn basic_header_empty_input_is_buffer_too_short() {
assert!(matches!(
BasicHeader::parse(&[]),
Err(RtmpError::BufferTooShort {
need: 1,
have: 0,
..
})
));
}
#[test]
fn basic_header_truncated_two_byte_form_is_buffer_too_short() {
let bytes = [0b00_000000u8]; assert!(matches!(
BasicHeader::parse(&bytes),
Err(RtmpError::BufferTooShort {
need: 2,
have: 1,
..
})
));
}
#[test]
fn basic_header_truncated_three_byte_form_is_buffer_too_short() {
let bytes = [0b00_000001u8, 0xAB]; assert!(matches!(
BasicHeader::parse(&bytes),
Err(RtmpError::BufferTooShort {
need: 3,
have: 2,
..
})
));
}
#[test]
fn type0_round_trip_build_serialize_parse_no_extended() {
let mh = MessageHeader::Type0 {
timestamp: 0x0011_2233,
message_length: 0x0004_5566,
message_type_id: 0x09,
message_stream_id: 0xAABB_CCDD,
};
let mut buf = [0u8; TYPE0_LEN];
let n = mh.serialize_into(&mut buf).unwrap();
assert_eq!(n, TYPE0_LEN);
let (parsed, consumed) = MessageHeader::parse(Fmt::Type0, &buf).unwrap();
assert_eq!(consumed, TYPE0_LEN);
assert_eq!(parsed, mh);
}
#[test]
fn type0_parse_serialize_byte_identical_no_extended_and_le_stream_id() {
let bytes: [u8; TYPE0_LEN] = [
0x00, 0x11, 0x22, 0x33, 0x44, 0x55, 0x09, 0xDD, 0xCC, 0xBB, 0xAA, ];
let (mh, consumed) = MessageHeader::parse(Fmt::Type0, &bytes).unwrap();
assert_eq!(consumed, TYPE0_LEN);
assert_eq!(
mh,
MessageHeader::Type0 {
timestamp: 0x0000_1122,
message_length: 0x0033_4455,
message_type_id: 0x09,
message_stream_id: 0xAABB_CCDD,
}
);
let mut buf = [0u8; TYPE0_LEN];
mh.serialize_into(&mut buf).unwrap();
assert_eq!(
buf, bytes,
"byte-identical round trip, LE stream id included"
);
}
#[test]
fn type0_extended_timestamp_parse_serialize_byte_identical() {
let bytes: [u8; TYPE0_LEN + 4] = [
0xFF, 0xFF, 0xFF, 0x00, 0x00, 0x10, 0x08, 0x01, 0x00, 0x00, 0x00, 0x01, 0x02, 0x03, 0x04, ];
let (mh, consumed) = MessageHeader::parse(Fmt::Type0, &bytes).unwrap();
assert_eq!(consumed, TYPE0_LEN + 4);
assert_eq!(
mh,
MessageHeader::Type0 {
timestamp: 0x0102_0304,
message_length: 0x0000_0010,
message_type_id: 0x08,
message_stream_id: 1,
}
);
let mut buf = [0u8; TYPE0_LEN + 4];
let n = mh.serialize_into(&mut buf).unwrap();
assert_eq!(n, TYPE0_LEN + 4);
assert_eq!(
buf, bytes,
"extended timestamp path must round-trip byte-identically"
);
}
#[test]
fn type0_timestamp_exactly_at_marker_boundary_uses_extended_path() {
let mh = MessageHeader::Type0 {
timestamp: EXTENDED_TIMESTAMP_MARKER,
message_length: 10,
message_type_id: 1,
message_stream_id: 0,
};
assert_eq!(mh.serialized_len(), TYPE0_LEN + 4);
let mut buf = [0u8; TYPE0_LEN + 4];
let n = mh.serialize_into(&mut buf).unwrap();
assert_eq!(n, TYPE0_LEN + 4);
assert_eq!(
&buf[0..3],
[0xFF, 0xFF, 0xFF],
"24-bit field must be the sentinel"
);
assert_eq!(
&buf[11..15],
&EXTENDED_TIMESTAMP_MARKER.to_be_bytes()[..],
"extended field carries the real value"
);
let (parsed, consumed) = MessageHeader::parse(Fmt::Type0, &buf).unwrap();
assert_eq!(consumed, TYPE0_LEN + 4);
assert_eq!(parsed, mh);
}
#[test]
fn type1_round_trip_build_serialize_parse_no_extended() {
let mh = MessageHeader::Type1 {
timestamp_delta: 20,
message_length: 32,
message_type_id: 8,
};
let mut buf = [0u8; TYPE1_LEN];
let n = mh.serialize_into(&mut buf).unwrap();
assert_eq!(n, TYPE1_LEN);
let (parsed, consumed) = MessageHeader::parse(Fmt::Type1, &buf).unwrap();
assert_eq!(consumed, TYPE1_LEN);
assert_eq!(parsed, mh);
}
#[test]
fn type1_extended_timestamp_parse_serialize_byte_identical() {
let bytes: [u8; TYPE1_LEN + 4] = [
0xFF, 0xFF, 0xFF, 0x00, 0x00, 0x20, 0x09, 0x0A, 0x0B, 0x0C, 0x0D, ];
let (mh, consumed) = MessageHeader::parse(Fmt::Type1, &bytes).unwrap();
assert_eq!(consumed, TYPE1_LEN + 4);
assert_eq!(
mh,
MessageHeader::Type1 {
timestamp_delta: 0x0A0B_0C0D,
message_length: 0x0000_0020,
message_type_id: 0x09,
}
);
let mut buf = [0u8; TYPE1_LEN + 4];
mh.serialize_into(&mut buf).unwrap();
assert_eq!(buf, bytes);
}
#[test]
fn type2_round_trip_build_serialize_parse_no_extended() {
let mh = MessageHeader::Type2 {
timestamp_delta: 20,
};
let mut buf = [0u8; TYPE2_LEN];
let n = mh.serialize_into(&mut buf).unwrap();
assert_eq!(n, TYPE2_LEN);
let (parsed, consumed) = MessageHeader::parse(Fmt::Type2, &buf).unwrap();
assert_eq!(consumed, TYPE2_LEN);
assert_eq!(parsed, mh);
}
#[test]
fn type2_extended_timestamp_parse_serialize_byte_identical() {
let bytes: [u8; TYPE2_LEN + 4] = [
0xFF, 0xFF, 0xFF, 0x11, 0x22, 0x33, 0x44, ];
let (mh, consumed) = MessageHeader::parse(Fmt::Type2, &bytes).unwrap();
assert_eq!(consumed, TYPE2_LEN + 4);
assert_eq!(
mh,
MessageHeader::Type2 {
timestamp_delta: 0x1122_3344,
}
);
let mut buf = [0u8; TYPE2_LEN + 4];
mh.serialize_into(&mut buf).unwrap();
assert_eq!(buf, bytes);
}
#[test]
fn type3_round_trip_is_zero_bytes() {
let mh = MessageHeader::Type3;
assert_eq!(mh.serialized_len(), 0);
let mut buf: [u8; 0] = [];
let n = mh.serialize_into(&mut buf).unwrap();
assert_eq!(n, 0);
let (parsed, consumed) = MessageHeader::parse(Fmt::Type3, &[]).unwrap();
assert_eq!(consumed, 0);
assert_eq!(parsed, MessageHeader::Type3);
}
#[test]
fn type0_truncated_input_is_buffer_too_short() {
let bytes = [0u8; TYPE0_LEN - 1];
assert!(matches!(
MessageHeader::parse(Fmt::Type0, &bytes),
Err(RtmpError::BufferTooShort {
need: TYPE0_LEN,
..
})
));
}
#[test]
fn type0_extended_marker_but_truncated_extended_field_is_buffer_too_short() {
let mut bytes = [0u8; TYPE0_LEN + 2]; bytes[0] = 0xFF;
bytes[1] = 0xFF;
bytes[2] = 0xFF;
assert!(matches!(
MessageHeader::parse(Fmt::Type0, &bytes),
Err(RtmpError::BufferTooShort {
need,
..
}) if need == TYPE0_LEN + 4
));
}
#[test]
fn message_stream_id_le_differs_from_be_for_asymmetric_value() {
let v: u32 = 0xAABB_CCDD;
assert_ne!(v.to_le_bytes(), v.to_be_bytes());
}
fn msg(csid: u32, timestamp: u32, type_id: u8, stream_id: u32, payload: Vec<u8>) -> Message {
Message {
chunk_stream_id: csid,
timestamp,
message_type_id: type_id,
message_stream_id: stream_id,
payload,
}
}
#[test]
fn writer_assembler_round_trip_small_message_single_chunk() {
let original = msg(4, 1000, 9, 1, vec![0xAB; 50]);
let mut writer = ChunkWriter::new();
let bytes = writer.write(&original);
let mut assembler = ChunkAssembler::new();
let out = assembler.push(&bytes).unwrap();
assert_eq!(out.len(), 1, "one message must come back out");
assert_eq!(out[0], original);
}
#[test]
fn writer_assembler_round_trip_message_larger_than_chunk_size() {
let original = msg(6, 5000, 9, 42, (0u8..=255).cycle().take(300).collect());
let mut writer = ChunkWriter::new();
let bytes = writer.write(&original);
let expected_len = 1 + TYPE0_LEN + 128 + (1 + 128) + (1 + 44);
assert_eq!(bytes.len(), expected_len);
let mut assembler = ChunkAssembler::new();
let out = assembler.push(&bytes).unwrap();
assert_eq!(
out.len(),
1,
"the 3 chunks must reassemble into ONE message"
);
assert_eq!(out[0], original);
assert_eq!(out[0].payload.len(), 300);
}
#[test]
fn assembler_multi_chunk_payload_reassembled_in_order() {
let mut assembler = ChunkAssembler::new();
assembler.set_chunk_size(4);
let bh0 = BasicHeader {
fmt: Fmt::Type0,
chunk_stream_id: 3,
};
let mh0 = MessageHeader::Type0 {
timestamp: 0,
message_length: 10,
message_type_id: 8,
message_stream_id: 0,
};
let mut input = Vec::new();
write_serialized(&mut input, &bh0);
write_serialized(&mut input, &mh0);
input.extend_from_slice(&[1, 2, 3, 4]);
let bh3 = BasicHeader {
fmt: Fmt::Type3,
chunk_stream_id: 3,
};
write_serialized(&mut input, &bh3);
input.extend_from_slice(&[5, 6, 7, 8]);
write_serialized(&mut input, &bh3);
input.extend_from_slice(&[9, 10]);
let out = assembler.push(&input).unwrap();
assert_eq!(out.len(), 1);
assert_eq!(out[0].payload, vec![1, 2, 3, 4, 5, 6, 7, 8, 9, 10]);
}
#[test]
fn assembler_header_inheritance_type0_type1_type2_type3() {
let mut assembler = ChunkAssembler::new();
assembler.set_chunk_size(5);
let csid = 5;
let mut input = Vec::new();
write_serialized(
&mut input,
&BasicHeader {
fmt: Fmt::Type0,
chunk_stream_id: csid,
},
);
write_serialized(
&mut input,
&MessageHeader::Type0 {
timestamp: 1000,
message_length: 5,
message_type_id: 8,
message_stream_id: 7,
},
);
input.extend_from_slice(&[0; 5]);
write_serialized(
&mut input,
&BasicHeader {
fmt: Fmt::Type1,
chunk_stream_id: csid,
},
);
write_serialized(
&mut input,
&MessageHeader::Type1 {
timestamp_delta: 20,
message_length: 5,
message_type_id: 8,
},
);
input.extend_from_slice(&[1; 5]);
write_serialized(
&mut input,
&BasicHeader {
fmt: Fmt::Type2,
chunk_stream_id: csid,
},
);
write_serialized(
&mut input,
&MessageHeader::Type2 {
timestamp_delta: 30,
},
);
input.extend_from_slice(&[2; 5]);
write_serialized(
&mut input,
&BasicHeader {
fmt: Fmt::Type3,
chunk_stream_id: csid,
},
);
input.extend_from_slice(&[3; 5]);
let out = assembler.push(&input).unwrap();
assert_eq!(out.len(), 4);
assert_eq!(out[0].timestamp, 1000);
assert_eq!(out[0].message_stream_id, 7);
assert_eq!(out[0].message_type_id, 8);
assert_eq!(out[0].payload, vec![0; 5]);
assert_eq!(out[1].timestamp, 1020, "fmt1: 1000 + delta 20");
assert_eq!(out[1].message_stream_id, 7, "fmt1 inherits stream id");
assert_eq!(out[1].message_type_id, 8);
assert_eq!(out[1].payload, vec![1; 5]);
assert_eq!(out[2].timestamp, 1050, "fmt2: 1020 + delta 30");
assert_eq!(out[2].message_stream_id, 7, "fmt2 inherits stream id");
assert_eq!(out[2].message_type_id, 8, "fmt2 inherits type id");
assert_eq!(out[2].payload, vec![2; 5], "fmt2 inherits message length");
assert_eq!(
out[3].timestamp, 1080,
"fmt3 (new message) inherits fmt2's delta 30: 1050 + 30"
);
assert_eq!(out[3].message_stream_id, 7, "fmt3 inherits stream id");
assert_eq!(out[3].message_type_id, 8, "fmt3 inherits type id");
assert_eq!(out[3].payload, vec![3; 5], "fmt3 inherits message length");
}
#[test]
fn assembler_mid_stream_set_chunk_size_changes_split_boundary() {
let mut assembler = ChunkAssembler::new();
let csid = 7;
let first = msg(csid, 100, 8, 1, vec![0xAA; 10]);
let mut writer = ChunkWriter::new();
let first_bytes = writer.write(&first);
let out = assembler.push(&first_bytes).unwrap();
assert_eq!(out, vec![first]);
assembler.set_chunk_size(4);
let bh0 = BasicHeader {
fmt: Fmt::Type0,
chunk_stream_id: csid,
};
let mh0 = MessageHeader::Type0 {
timestamp: 200,
message_length: 10,
message_type_id: 8,
message_stream_id: 1,
};
let mut chunk1 = Vec::new();
write_serialized(&mut chunk1, &bh0);
write_serialized(&mut chunk1, &mh0);
chunk1.extend_from_slice(&[1, 2, 3, 4]);
let out = assembler.push(&chunk1).unwrap();
assert!(
out.is_empty(),
"only 4 of 10 payload bytes arrived, message must not complete yet"
);
let bh3 = BasicHeader {
fmt: Fmt::Type3,
chunk_stream_id: csid,
};
let mut chunk2 = Vec::new();
write_serialized(&mut chunk2, &bh3);
chunk2.extend_from_slice(&[5, 6, 7, 8]);
let out = assembler.push(&chunk2).unwrap();
assert!(out.is_empty(), "8 of 10 payload bytes, still incomplete");
let mut chunk3 = Vec::new();
write_serialized(&mut chunk3, &bh3);
chunk3.extend_from_slice(&[9, 10]);
let out = assembler.push(&chunk3).unwrap();
assert_eq!(out.len(), 1);
assert_eq!(out[0].payload, vec![1, 2, 3, 4, 5, 6, 7, 8, 9, 10]);
}
#[test]
fn writer_assembler_round_trip_extended_timestamp_split_across_chunks() {
let original = msg(8, EXTENDED_TIMESTAMP_MARKER + 12345, 9, 2, vec![0x7E; 300]);
let mut writer = ChunkWriter::new();
let bytes = writer.write(&original);
let bh_len = 1; let first_header_len = bh_len + TYPE0_LEN + EXTENDED_TIMESTAMP_LEN;
assert_eq!(&bytes[bh_len..bh_len + 3], [0xFF, 0xFF, 0xFF]);
let ext_offset = bh_len + TYPE0_LEN;
assert_eq!(
&bytes[ext_offset..ext_offset + 4],
&original.timestamp.to_be_bytes()
);
let first_payload_take = 128usize;
let second_chunk_start = first_header_len + first_payload_take;
assert_eq!(bytes[second_chunk_start] >> 6, Fmt::Type3.to_bits());
let second_ext_offset = second_chunk_start + 1;
assert_eq!(
&bytes[second_ext_offset..second_ext_offset + 4],
&original.timestamp.to_be_bytes(),
"fmt3 continuation must carry the same extended timestamp"
);
let mut assembler = ChunkAssembler::new();
let out = assembler.push(&bytes).unwrap();
assert_eq!(out.len(), 1);
assert_eq!(out[0], original);
assert_eq!(out[0].timestamp, EXTENDED_TIMESTAMP_MARKER + 12345);
}
#[test]
fn assembler_partial_feed_split_mid_header_no_drop_or_duplicate() {
let original = msg(9, 42, 8, 3, vec![0x11; 200]);
let mut writer = ChunkWriter::new();
let bytes = writer.write(&original);
let split_at = 4;
assert!(split_at < 1 + TYPE0_LEN);
let mut assembler = ChunkAssembler::new();
let out1 = assembler.push(&bytes[..split_at]).unwrap();
assert!(out1.is_empty(), "partial header must not error or complete");
let out2 = assembler.push(&bytes[split_at..]).unwrap();
assert_eq!(out2.len(), 1, "message must complete exactly once");
assert_eq!(out2[0], original);
}
#[test]
fn assembler_partial_feed_split_mid_payload_no_drop_or_duplicate() {
let original = msg(10, 42, 8, 3, vec![0x22; 300]);
let mut writer = ChunkWriter::new();
let bytes = writer.write(&original);
let split_at = 1 + TYPE0_LEN + 60;
let mut assembler = ChunkAssembler::new();
let out1 = assembler.push(&bytes[..split_at]).unwrap();
assert!(out1.is_empty());
let out2 = assembler.push(&bytes[split_at..]).unwrap();
assert_eq!(out2.len(), 1);
assert_eq!(out2[0], original);
}
#[test]
fn assembler_partial_feed_byte_at_a_time_never_drops_or_duplicates() {
let original = msg(11, 7, 9, 4, vec![0x33; 260]);
let mut writer = ChunkWriter::new();
let bytes = writer.write(&original);
let mut assembler = ChunkAssembler::new();
let mut collected = Vec::new();
for b in &bytes {
collected.extend(assembler.push(std::slice::from_ref(b)).unwrap());
}
assert_eq!(collected.len(), 1);
assert_eq!(collected[0], original);
}
#[test]
fn assembler_type1_on_unseen_csid_is_malformed() {
let mut assembler = ChunkAssembler::new();
let mut input = Vec::new();
write_serialized(
&mut input,
&BasicHeader {
fmt: Fmt::Type1,
chunk_stream_id: 20,
},
);
write_serialized(
&mut input,
&MessageHeader::Type1 {
timestamp_delta: 5,
message_length: 3,
message_type_id: 1,
},
);
input.extend_from_slice(&[0, 0, 0]);
assert!(matches!(
assembler.push(&input),
Err(RtmpError::Malformed { .. })
));
}
#[test]
fn assembler_type3_on_unseen_csid_is_malformed() {
let mut assembler = ChunkAssembler::new();
let mut input = Vec::new();
write_serialized(
&mut input,
&BasicHeader {
fmt: Fmt::Type3,
chunk_stream_id: 21,
},
);
assert!(matches!(
assembler.push(&input),
Err(RtmpError::Malformed { .. })
));
}
#[test]
fn assembler_truncated_input_never_panics_across_many_split_points() {
let original = msg(12, 99, 8, 5, vec![0x44; 400]);
let mut writer = ChunkWriter::new();
let bytes = writer.write(&original);
for split in 0..=bytes.len() {
let mut assembler = ChunkAssembler::new();
let first = assembler.push(&bytes[..split]);
let Ok(first_msgs) = first else {
continue;
};
let second = assembler.push(&bytes[split..]).unwrap();
let mut all = first_msgs;
all.extend(second);
assert_eq!(all, vec![original.clone()]);
}
}
#[test]
fn writer_default_chunk_size_matches_assembler_default() {
assert_eq!(ChunkWriter::new().chunk_size, DEFAULT_CHUNK_SIZE);
assert_eq!(ChunkAssembler::new().chunk_size, DEFAULT_CHUNK_SIZE);
}
#[test]
fn assembler_type3_immediately_after_type0_uses_type0_timestamp_as_implied_delta() {
let mut assembler = ChunkAssembler::new();
let csid = 40;
let mut input = Vec::new();
write_serialized(
&mut input,
&BasicHeader {
fmt: Fmt::Type0,
chunk_stream_id: csid,
},
);
write_serialized(
&mut input,
&MessageHeader::Type0 {
timestamp: 1000,
message_length: 5,
message_type_id: 8,
message_stream_id: 2,
},
);
input.extend_from_slice(&[1, 2, 3, 4, 5]);
write_serialized(
&mut input,
&BasicHeader {
fmt: Fmt::Type3,
chunk_stream_id: csid,
},
);
input.extend_from_slice(&[9, 9, 9, 9, 9]);
let out = assembler.push(&input).unwrap();
assert_eq!(out.len(), 2);
assert_eq!(out[0].timestamp, 1000);
assert_eq!(
out[1].timestamp, 2000,
"fmt3 immediately after fmt0 implies delta == the fmt0's own timestamp (1000), not 0: 1000 + 1000"
);
}
#[test]
fn assembler_type2_on_unseen_csid_is_malformed() {
let mut assembler = ChunkAssembler::new();
let mut input = Vec::new();
write_serialized(
&mut input,
&BasicHeader {
fmt: Fmt::Type2,
chunk_stream_id: 22,
},
);
write_serialized(&mut input, &MessageHeader::Type2 { timestamp_delta: 5 });
assert!(matches!(
assembler.push(&input),
Err(RtmpError::Malformed { .. })
));
}
#[test]
fn assembler_set_chunk_size_zero_is_floored_to_one() {
let mut assembler = ChunkAssembler::new();
assembler.set_chunk_size(0);
assert_eq!(assembler.chunk_size, 1, "floored at 1, not stuck at 0");
let csid = 41;
let mut input = Vec::new();
write_serialized(
&mut input,
&BasicHeader {
fmt: Fmt::Type0,
chunk_stream_id: csid,
},
);
write_serialized(
&mut input,
&MessageHeader::Type0 {
timestamp: 1,
message_length: 3,
message_type_id: 8,
message_stream_id: 0,
},
);
input.push(0xAA);
write_serialized(
&mut input,
&BasicHeader {
fmt: Fmt::Type3,
chunk_stream_id: csid,
},
);
input.push(0xBB);
write_serialized(
&mut input,
&BasicHeader {
fmt: Fmt::Type3,
chunk_stream_id: csid,
},
);
input.push(0xCC);
let out = assembler.push(&input).unwrap();
assert_eq!(out.len(), 1);
assert_eq!(out[0].payload, vec![0xAA, 0xBB, 0xCC]);
}
fn single_chunk(csid: u32, payload: &[u8]) -> Vec<u8> {
let mut input = Vec::new();
write_serialized(
&mut input,
&BasicHeader {
fmt: Fmt::Type0,
chunk_stream_id: csid,
},
);
write_serialized(
&mut input,
&MessageHeader::Type0 {
timestamp: 0,
message_length: payload.len() as u32,
message_type_id: 9,
message_stream_id: 1,
},
);
input.extend_from_slice(payload);
input
}
#[test]
fn oversized_message_length_header_is_rejected_without_allocating() {
let mut assembler = ChunkAssembler::new();
let mut input = Vec::new();
write_serialized(
&mut input,
&BasicHeader {
fmt: Fmt::Type0,
chunk_stream_id: 4,
},
);
write_serialized(
&mut input,
&MessageHeader::Type0 {
timestamp: 0,
message_length: 0x00FF_FFFF, message_type_id: 9,
message_stream_id: 1,
},
);
input.extend(std::iter::repeat_n(0u8, DEFAULT_CHUNK_SIZE as usize));
let err = assembler.push(&input).expect_err(
"a message_length beyond MAX_MESSAGE_LEN must be rejected before any \
message_length-sized buffer is allocated",
);
assert!(matches!(err, RtmpError::Malformed { .. }));
}
#[test]
fn message_length_at_the_cap_is_accepted() {
let mut assembler = ChunkAssembler::new();
let mut input = Vec::new();
write_serialized(
&mut input,
&BasicHeader {
fmt: Fmt::Type0,
chunk_stream_id: 4,
},
);
write_serialized(
&mut input,
&MessageHeader::Type0 {
timestamp: 0,
message_length: MAX_MESSAGE_LEN,
message_type_id: 9,
message_stream_id: 1,
},
);
input.extend(std::iter::repeat_n(0u8, DEFAULT_CHUNK_SIZE as usize));
assert!(assembler.push(&input).is_ok());
}
#[test]
fn csid_flood_beyond_max_csids_is_rejected() {
let mut assembler = ChunkAssembler::new();
for i in 0..MAX_CSIDS {
let csid = BASIC_HEADER_1BYTE_MIN_CSID + i as u32;
let out = assembler
.push(&single_chunk(csid, &[0xAB]))
.unwrap_or_else(|e| panic!("csid {csid} (#{i}, within the bound) rejected: {e}"));
assert_eq!(out.len(), 1);
}
let flood_csid = BASIC_HEADER_1BYTE_MIN_CSID + MAX_CSIDS as u32;
let err = assembler
.push(&single_chunk(flood_csid, &[0xCD]))
.expect_err("a new csid beyond MAX_CSIDS must be rejected, not silently accepted");
assert!(matches!(err, RtmpError::Malformed { .. }));
}
#[test]
fn csid_flood_cap_does_not_count_repeats_of_the_same_csid() {
let mut assembler = ChunkAssembler::new();
for i in 0..(MAX_CSIDS * 4) {
let out = assembler
.push(&single_chunk(BASIC_HEADER_1BYTE_MIN_CSID, &[i as u8]))
.expect("repeated use of a single already-known csid must never be rejected");
assert_eq!(out.len(), 1);
}
}
#[test]
fn set_chunk_size_is_capped_at_max_chunk_size() {
let mut assembler = ChunkAssembler::new();
assembler.set_chunk_size(u32::MAX);
assert_eq!(assembler.chunk_size, MAX_CHUNK_SIZE);
let mut writer = ChunkWriter::new();
writer.set_chunk_size(u32::MAX);
assert_eq!(writer.chunk_size, MAX_CHUNK_SIZE);
}
}