use broadcast_common::{Parse, Serialize};
use crate::RtmpError;
use crate::chunk::Message;
type Result<T> = core::result::Result<T, RtmpError>;
pub mod msg_type {
pub const SET_CHUNK_SIZE: u8 = 1;
pub const ABORT: u8 = 2;
pub const ACKNOWLEDGEMENT: u8 = 3;
pub const USER_CONTROL: u8 = 4;
pub const WINDOW_ACK_SIZE: u8 = 5;
pub const SET_PEER_BANDWIDTH: u8 = 6;
pub const AUDIO: u8 = 8;
pub const VIDEO: u8 = 9;
pub const DATA_AMF3: u8 = 15;
pub const COMMAND_AMF3: u8 = 17;
pub const DATA_AMF0: u8 = 18;
pub const COMMAND_AMF0: u8 = 20;
pub const AGGREGATE: u8 = 22;
}
pub const CONTROL_CHUNK_STREAM_ID: u32 = 2;
pub const CONTROL_MESSAGE_STREAM_ID: u32 = 0;
const U32_LEN: usize = 4;
const SET_PEER_BANDWIDTH_LEN: usize = U32_LEN + 1;
const SET_CHUNK_SIZE_RESERVED_MASK: u32 = 0x8000_0000;
const SET_CHUNK_SIZE_VALUE_MASK: u32 = 0x7FFF_FFFF;
fn read_u32_be(b: &[u8]) -> u32 {
u32::from_be_bytes([b[0], b[1], b[2], b[3]])
}
fn need_u32(bytes: &[u8], what: &'static str) -> Result<u32> {
if bytes.len() < U32_LEN {
return Err(RtmpError::BufferTooShort {
need: U32_LEN,
have: bytes.len(),
what,
});
}
Ok(read_u32_be(bytes))
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
#[non_exhaustive]
pub enum LimitType {
Hard,
Soft,
Dynamic,
}
impl LimitType {
#[must_use]
pub fn name(&self) -> &'static str {
match self {
LimitType::Hard => "hard",
LimitType::Soft => "soft",
LimitType::Dynamic => "dynamic",
}
}
pub const fn from_u8(v: u8) -> core::result::Result<Self, RtmpError> {
match v {
0 => Ok(LimitType::Hard),
1 => Ok(LimitType::Soft),
2 => Ok(LimitType::Dynamic),
_ => Err(RtmpError::Malformed {
what: "set peer bandwidth limit type (must be 0..=2)",
}),
}
}
#[must_use]
pub const fn to_u8(self) -> u8 {
match self {
LimitType::Hard => 0,
LimitType::Soft => 1,
LimitType::Dynamic => 2,
}
}
}
broadcast_common::impl_spec_display!(LimitType);
#[non_exhaustive]
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum ProtocolControl {
SetChunkSize(u32),
Abort {
chunk_stream_id: u32,
},
Acknowledgement(u32),
WindowAckSize(u32),
SetPeerBandwidth {
ack_window_size: u32,
limit_type: LimitType,
},
}
impl ProtocolControl {
#[must_use]
pub fn name(&self) -> &'static str {
match self {
ProtocolControl::SetChunkSize(_) => "set chunk size",
ProtocolControl::Abort { .. } => "abort message",
ProtocolControl::Acknowledgement(_) => "acknowledgement",
ProtocolControl::WindowAckSize(_) => "window acknowledgement size",
ProtocolControl::SetPeerBandwidth { .. } => "set peer bandwidth",
}
}
#[must_use]
pub fn message_type_id(&self) -> u8 {
match self {
ProtocolControl::SetChunkSize(_) => msg_type::SET_CHUNK_SIZE,
ProtocolControl::Abort { .. } => msg_type::ABORT,
ProtocolControl::Acknowledgement(_) => msg_type::ACKNOWLEDGEMENT,
ProtocolControl::WindowAckSize(_) => msg_type::WINDOW_ACK_SIZE,
ProtocolControl::SetPeerBandwidth { .. } => msg_type::SET_PEER_BANDWIDTH,
}
}
pub fn from_message(message: &Message) -> Result<Option<Self>> {
Self::from_payload(message.message_type_id, &message.payload)
}
pub fn from_payload(message_type_id: u8, payload: &[u8]) -> Result<Option<Self>> {
match message_type_id {
msg_type::SET_CHUNK_SIZE => {
let raw = need_u32(payload, "set chunk size payload")?;
if raw & SET_CHUNK_SIZE_RESERVED_MASK != 0 {
return Err(RtmpError::Malformed {
what: "set chunk size reserved top bit (must be 0)",
});
}
let size = raw & SET_CHUNK_SIZE_VALUE_MASK;
if size == 0 {
return Err(RtmpError::Malformed {
what: "set chunk size value (must be >= 1)",
});
}
Ok(Some(ProtocolControl::SetChunkSize(size)))
}
msg_type::ABORT => {
let chunk_stream_id = need_u32(payload, "abort message payload")?;
Ok(Some(ProtocolControl::Abort { chunk_stream_id }))
}
msg_type::ACKNOWLEDGEMENT => {
let sequence_number = need_u32(payload, "acknowledgement payload")?;
Ok(Some(ProtocolControl::Acknowledgement(sequence_number)))
}
msg_type::WINDOW_ACK_SIZE => {
let window = need_u32(payload, "window acknowledgement size payload")?;
Ok(Some(ProtocolControl::WindowAckSize(window)))
}
msg_type::SET_PEER_BANDWIDTH => {
if payload.len() < SET_PEER_BANDWIDTH_LEN {
return Err(RtmpError::BufferTooShort {
need: SET_PEER_BANDWIDTH_LEN,
have: payload.len(),
what: "set peer bandwidth payload",
});
}
let ack_window_size = read_u32_be(&payload[0..U32_LEN]);
let limit_type = LimitType::from_u8(payload[U32_LEN])?;
Ok(Some(ProtocolControl::SetPeerBandwidth {
ack_window_size,
limit_type,
}))
}
_ => Ok(None),
}
}
#[must_use]
pub fn to_message(&self) -> Message {
Message {
chunk_stream_id: CONTROL_CHUNK_STREAM_ID,
timestamp: 0,
message_type_id: self.message_type_id(),
message_stream_id: CONTROL_MESSAGE_STREAM_ID,
payload: self.to_bytes(),
}
}
}
broadcast_common::impl_spec_display!(ProtocolControl);
impl Serialize for ProtocolControl {
type Error = RtmpError;
fn serialized_len(&self) -> usize {
match self {
ProtocolControl::SetChunkSize(_)
| ProtocolControl::Abort { .. }
| ProtocolControl::Acknowledgement(_)
| ProtocolControl::WindowAckSize(_) => U32_LEN,
ProtocolControl::SetPeerBandwidth { .. } => SET_PEER_BANDWIDTH_LEN,
}
}
fn serialize_into(&self, buf: &mut [u8]) -> Result<usize> {
let written = self.serialized_len();
if buf.len() < written {
return Err(RtmpError::BufferTooShort {
need: written,
have: buf.len(),
what: "protocol control payload output",
});
}
match *self {
ProtocolControl::SetChunkSize(size) => {
if size == 0 || size & SET_CHUNK_SIZE_RESERVED_MASK != 0 {
return Err(RtmpError::Malformed {
what: "set chunk size value (must be 1..=0x7FFF_FFFF)",
});
}
buf[0..U32_LEN].copy_from_slice(&size.to_be_bytes());
}
ProtocolControl::Abort { chunk_stream_id } => {
buf[0..U32_LEN].copy_from_slice(&chunk_stream_id.to_be_bytes());
}
ProtocolControl::Acknowledgement(sequence_number) => {
buf[0..U32_LEN].copy_from_slice(&sequence_number.to_be_bytes());
}
ProtocolControl::WindowAckSize(window) => {
buf[0..U32_LEN].copy_from_slice(&window.to_be_bytes());
}
ProtocolControl::SetPeerBandwidth {
ack_window_size,
limit_type,
} => {
buf[0..U32_LEN].copy_from_slice(&ack_window_size.to_be_bytes());
buf[U32_LEN] = limit_type.to_u8();
}
}
Ok(written)
}
}
const EVENT_TYPE_LEN: usize = 2;
mod event_type {
pub const STREAM_BEGIN: u16 = 0;
pub const STREAM_EOF: u16 = 1;
pub const STREAM_DRY: u16 = 2;
pub const SET_BUFFER_LENGTH: u16 = 3;
pub const STREAM_IS_RECORDED: u16 = 4;
pub const PING_REQUEST: u16 = 6;
pub const PING_RESPONSE: u16 = 7;
}
#[non_exhaustive]
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum UserControl {
StreamBegin(u32),
StreamEof(u32),
StreamDry(u32),
SetBufferLength {
stream_id: u32,
buffer_ms: u32,
},
StreamIsRecorded(u32),
PingRequest(u32),
PingResponse(u32),
}
impl UserControl {
#[must_use]
pub fn name(&self) -> &'static str {
match self {
UserControl::StreamBegin(_) => "stream begin",
UserControl::StreamEof(_) => "stream eof",
UserControl::StreamDry(_) => "stream dry",
UserControl::SetBufferLength { .. } => "set buffer length",
UserControl::StreamIsRecorded(_) => "stream is recorded",
UserControl::PingRequest(_) => "ping request",
UserControl::PingResponse(_) => "ping response",
}
}
#[must_use]
pub fn event_type(&self) -> u16 {
match self {
UserControl::StreamBegin(_) => event_type::STREAM_BEGIN,
UserControl::StreamEof(_) => event_type::STREAM_EOF,
UserControl::StreamDry(_) => event_type::STREAM_DRY,
UserControl::SetBufferLength { .. } => event_type::SET_BUFFER_LENGTH,
UserControl::StreamIsRecorded(_) => event_type::STREAM_IS_RECORDED,
UserControl::PingRequest(_) => event_type::PING_REQUEST,
UserControl::PingResponse(_) => event_type::PING_RESPONSE,
}
}
#[must_use]
pub fn to_message(&self) -> Message {
Message {
chunk_stream_id: CONTROL_CHUNK_STREAM_ID,
timestamp: 0,
message_type_id: msg_type::USER_CONTROL,
message_stream_id: CONTROL_MESSAGE_STREAM_ID,
payload: self.to_bytes(),
}
}
}
broadcast_common::impl_spec_display!(UserControl);
impl<'a> Parse<'a> for UserControl {
type Error = RtmpError;
fn parse(bytes: &'a [u8]) -> Result<Self> {
if bytes.len() < EVENT_TYPE_LEN {
return Err(RtmpError::BufferTooShort {
need: EVENT_TYPE_LEN,
have: bytes.len(),
what: "user control event type",
});
}
let event = u16::from_be_bytes([bytes[0], bytes[1]]);
let data = &bytes[EVENT_TYPE_LEN..];
match event {
event_type::STREAM_BEGIN => Ok(UserControl::StreamBegin(need_u32(
data,
"stream begin event data",
)?)),
event_type::STREAM_EOF => Ok(UserControl::StreamEof(need_u32(
data,
"stream eof event data",
)?)),
event_type::STREAM_DRY => Ok(UserControl::StreamDry(need_u32(
data,
"stream dry event data",
)?)),
event_type::SET_BUFFER_LENGTH => {
if data.len() < 2 * U32_LEN {
return Err(RtmpError::BufferTooShort {
need: 2 * U32_LEN,
have: data.len(),
what: "set buffer length event data",
});
}
Ok(UserControl::SetBufferLength {
stream_id: read_u32_be(&data[0..U32_LEN]),
buffer_ms: read_u32_be(&data[U32_LEN..2 * U32_LEN]),
})
}
event_type::STREAM_IS_RECORDED => Ok(UserControl::StreamIsRecorded(need_u32(
data,
"stream is recorded event data",
)?)),
event_type::PING_REQUEST => Ok(UserControl::PingRequest(need_u32(
data,
"ping request event data",
)?)),
event_type::PING_RESPONSE => Ok(UserControl::PingResponse(need_u32(
data,
"ping response event data",
)?)),
_ => Err(RtmpError::Unsupported {
what: "user control event type (unrecognised)",
}),
}
}
}
impl Serialize for UserControl {
type Error = RtmpError;
fn serialized_len(&self) -> usize {
let data_len = match self {
UserControl::SetBufferLength { .. } => 2 * U32_LEN,
_ => U32_LEN,
};
EVENT_TYPE_LEN + data_len
}
fn serialize_into(&self, buf: &mut [u8]) -> Result<usize> {
let written = self.serialized_len();
if buf.len() < written {
return Err(RtmpError::BufferTooShort {
need: written,
have: buf.len(),
what: "user control event output",
});
}
buf[0..EVENT_TYPE_LEN].copy_from_slice(&self.event_type().to_be_bytes());
let data = &mut buf[EVENT_TYPE_LEN..written];
match *self {
UserControl::StreamBegin(stream_id)
| UserControl::StreamEof(stream_id)
| UserControl::StreamDry(stream_id)
| UserControl::StreamIsRecorded(stream_id)
| UserControl::PingRequest(stream_id)
| UserControl::PingResponse(stream_id) => {
data[0..U32_LEN].copy_from_slice(&stream_id.to_be_bytes());
}
UserControl::SetBufferLength {
stream_id,
buffer_ms,
} => {
data[0..U32_LEN].copy_from_slice(&stream_id.to_be_bytes());
data[U32_LEN..2 * U32_LEN].copy_from_slice(&buffer_ms.to_be_bytes());
}
}
Ok(written)
}
}
#[cfg(test)]
mod tests {
use super::*;
fn message(message_type_id: u8, payload: Vec<u8>) -> Message {
Message {
chunk_stream_id: CONTROL_CHUNK_STREAM_ID,
timestamp: 0,
message_type_id,
message_stream_id: CONTROL_MESSAGE_STREAM_ID,
payload,
}
}
#[test]
fn limit_type_round_trip_and_name() {
for (byte, lt, name) in [
(0u8, LimitType::Hard, "hard"),
(1, LimitType::Soft, "soft"),
(2, LimitType::Dynamic, "dynamic"),
] {
let parsed = LimitType::from_u8(byte).unwrap();
assert_eq!(parsed, lt);
assert_eq!(parsed.to_u8(), byte);
assert_eq!(parsed.name(), name);
assert_eq!(parsed.to_string(), name);
}
}
#[test]
fn limit_type_out_of_range_is_malformed() {
assert!(matches!(
LimitType::from_u8(3),
Err(RtmpError::Malformed { .. })
));
}
fn protocol_control_round_trip(pc: ProtocolControl) {
let bytes = pc.to_bytes();
let parsed = ProtocolControl::from_payload(pc.message_type_id(), &bytes)
.unwrap()
.expect("known protocol control type id");
assert_eq!(parsed, pc);
let msg = message(pc.message_type_id(), bytes.clone());
let via_message = ProtocolControl::from_message(&msg).unwrap().unwrap();
assert_eq!(via_message, pc);
assert_eq!(via_message.to_bytes(), bytes);
}
#[test]
fn set_chunk_size_round_trips() {
protocol_control_round_trip(ProtocolControl::SetChunkSize(4096));
}
#[test]
fn abort_round_trips() {
protocol_control_round_trip(ProtocolControl::Abort { chunk_stream_id: 7 });
}
#[test]
fn acknowledgement_round_trips() {
protocol_control_round_trip(ProtocolControl::Acknowledgement(1_048_576));
}
#[test]
fn window_ack_size_round_trips() {
protocol_control_round_trip(ProtocolControl::WindowAckSize(2_500_000));
}
#[test]
fn set_peer_bandwidth_round_trips_every_limit_type() {
for limit_type in [LimitType::Hard, LimitType::Soft, LimitType::Dynamic] {
protocol_control_round_trip(ProtocolControl::SetPeerBandwidth {
ack_window_size: 2_500_000,
limit_type,
});
}
}
#[test]
fn set_chunk_size_reserved_top_bit_rejected_on_parse() {
let bytes = 0x8000_1000u32.to_be_bytes().to_vec();
assert!(matches!(
ProtocolControl::from_payload(msg_type::SET_CHUNK_SIZE, &bytes),
Err(RtmpError::Malformed { .. })
));
}
#[test]
fn set_chunk_size_zero_rejected() {
let bytes = 0u32.to_be_bytes().to_vec();
assert!(matches!(
ProtocolControl::from_payload(msg_type::SET_CHUNK_SIZE, &bytes),
Err(RtmpError::Malformed { .. })
));
assert!(matches!(
ProtocolControl::SetChunkSize(0).serialize_into(&mut [0u8; 4]),
Err(RtmpError::Malformed { .. })
));
}
#[test]
fn set_chunk_size_serialize_layout_matches_spec() {
let bytes = ProtocolControl::SetChunkSize(1).to_bytes();
assert_eq!(bytes, vec![0x00, 0x00, 0x00, 0x01]);
}
#[test]
fn set_peer_bandwidth_serialize_layout_matches_spec() {
let bytes = ProtocolControl::SetPeerBandwidth {
ack_window_size: 0x0002_5000,
limit_type: LimitType::Dynamic,
}
.to_bytes();
assert_eq!(bytes, vec![0x00, 0x02, 0x50, 0x00, 0x02]);
}
#[test]
fn set_peer_bandwidth_wrong_limit_type_mapping_would_fail() {
assert_eq!(LimitType::Hard.to_u8(), 0);
assert_eq!(LimitType::Dynamic.to_u8(), 2);
assert_ne!(LimitType::Hard.to_u8(), LimitType::Dynamic.to_u8());
}
#[test]
fn from_message_none_for_non_control_type_id() {
let msg = message(msg_type::AUDIO, vec![0u8; 4]);
assert!(ProtocolControl::from_message(&msg).unwrap().is_none());
}
#[test]
fn from_message_some_for_control_type_id() {
let msg = message(
msg_type::WINDOW_ACK_SIZE,
1_000_000u32.to_be_bytes().to_vec(),
);
assert!(ProtocolControl::from_message(&msg).unwrap().is_some());
}
#[test]
fn to_message_uses_control_csid_and_stream_id() {
let msg = ProtocolControl::SetChunkSize(4096).to_message();
assert_eq!(msg.chunk_stream_id, CONTROL_CHUNK_STREAM_ID);
assert_eq!(msg.message_stream_id, CONTROL_MESSAGE_STREAM_ID);
assert_eq!(msg.message_type_id, msg_type::SET_CHUNK_SIZE);
}
#[test]
fn protocol_control_display_matches_name() {
assert_eq!(
ProtocolControl::Acknowledgement(1).to_string(),
ProtocolControl::Acknowledgement(1).name()
);
}
fn user_control_round_trip(uc: UserControl) {
let bytes = uc.to_bytes();
let parsed = UserControl::parse(&bytes).unwrap();
assert_eq!(parsed, uc);
assert_eq!(parsed.to_bytes(), bytes);
}
#[test]
fn stream_begin_round_trips() {
user_control_round_trip(UserControl::StreamBegin(1));
}
#[test]
fn stream_begin_serialize_layout_matches_spec() {
let bytes = UserControl::StreamBegin(1).to_bytes();
assert_eq!(bytes, vec![0x00, 0x00, 0x00, 0x00, 0x00, 0x01]);
}
#[test]
fn stream_eof_round_trips() {
user_control_round_trip(UserControl::StreamEof(1));
}
#[test]
fn stream_dry_round_trips() {
user_control_round_trip(UserControl::StreamDry(1));
}
#[test]
fn set_buffer_length_round_trips() {
user_control_round_trip(UserControl::SetBufferLength {
stream_id: 1,
buffer_ms: 3000,
});
}
#[test]
fn stream_is_recorded_round_trips() {
user_control_round_trip(UserControl::StreamIsRecorded(1));
}
#[test]
fn ping_request_round_trips() {
user_control_round_trip(UserControl::PingRequest(0x1234_5678));
}
#[test]
fn ping_response_round_trips() {
user_control_round_trip(UserControl::PingResponse(0x1234_5678));
}
#[test]
fn unrecognised_event_type_is_unsupported() {
let bytes = [0x00, 0x05, 0x00, 0x00, 0x00, 0x01];
assert!(matches!(
UserControl::parse(&bytes),
Err(RtmpError::Unsupported { .. })
));
}
#[test]
fn user_control_event_type_wrong_mapping_would_fail() {
assert_eq!(UserControl::StreamBegin(0).event_type(), 0);
assert_eq!(UserControl::StreamEof(0).event_type(), 1);
}
#[test]
fn user_control_display_matches_name() {
assert_eq!(
UserControl::StreamBegin(1).to_string(),
UserControl::StreamBegin(1).name()
);
}
#[test]
fn to_message_uses_control_csid_and_user_control_type_id() {
let msg = UserControl::StreamBegin(1).to_message();
assert_eq!(msg.chunk_stream_id, CONTROL_CHUNK_STREAM_ID);
assert_eq!(msg.message_stream_id, CONTROL_MESSAGE_STREAM_ID);
assert_eq!(msg.message_type_id, msg_type::USER_CONTROL);
}
}