diff --git a/src/client/mod.rs b/src/client/mod.rs index 427b350..d200ccc 100644 --- a/src/client/mod.rs +++ b/src/client/mod.rs @@ -11,7 +11,7 @@ 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::ertmp::multitrack_media::{foreach_track, is_multitrack_container}; use crate::handshake::{self, Handshake}; use crate::media::{is_on_metadata_payload, populate_av_frame, populate_multitrack_frame}; use crate::message::command; @@ -526,7 +526,7 @@ impl Client { } else { FrameType::Video }; - self.deliver_av_frame_cb(cb, frame_type, msg.timestamp, payload); + 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 @@ -594,7 +594,7 @@ impl Client { FrameType::Audio, out_ts, tag_payload.to_vec(), - ); + )?; } msg_dispatch::RTMP_MSG_VIDEO => { self.deliver_av_frame_cb( @@ -602,7 +602,7 @@ impl Client { FrameType::Video, out_ts, tag_payload.to_vec(), - ); + )?; } msg_dispatch::RTMP_MSG_AMF0_DATA => { self.deliver_script_frame_cb(cb, out_ts, tag_payload); @@ -627,8 +627,9 @@ impl Client { frame_type: FrameType, timestamp: u32, payload: Vec, - ) { - let had_multitrack = foreach_track(frame_type, &payload, |track| { + ) -> Result<()> { + let is_multitrack = is_multitrack_container(frame_type, &payload); + let parsed_multitrack = foreach_track(frame_type, &payload, |track| { self.invoke_multitrack_on_frame_cb( cb, frame_type, @@ -640,9 +641,13 @@ impl Client { track.payload, ); }); - if !had_multitrack { + if is_multitrack && !parsed_multitrack { + return Err(ErrorCode::Protocol); + } + if !is_multitrack { self.invoke_on_frame_cb(cb, frame_type, timestamp, u8::MAX, &payload); } + Ok(()) } fn invoke_multitrack_on_frame_cb( @@ -1321,6 +1326,37 @@ mod tests { assert_eq!(seen[1].1, vec![0xDD, 0xEE]); } + #[test] + fn drain_ready_messages_rejects_oversized_multitrack_video() { + let mut payload = vec![0x86, 0x10, b'a', b'v', b'c', b'1']; + for id in 0..=crate::ertmp::multitrack_media::MAX_MULTITRACK_SUBTRACKS { + payload.push(id as u8); + payload.extend_from_slice(&[0x00, 0x00, 0x00]); + } + + 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; + let chunk_size = payload.len(); + chunk_write(&mut wire, &cmsg, &payload, payload.len(), chunk_size).unwrap(); + + let mut client = Client::new(); + client.chunk_reg.set_all_chunk_size(chunk_size as u32); + client.recv_buffer.write(wire.peek()).unwrap(); + client.on_frame_cb = Some(|_| panic!("invalid multitrack must not reach callback")); + + let mut messages_processed = 0; + assert_eq!( + client.drain_ready_messages(&mut messages_processed), + Err(ErrorCode::Protocol) + ); + assert!(client.frame_cb_scratch.is_empty()); + } + #[test] fn poll_drains_leftover_messages_before_enforcing_staging_cap() { use std::io::Write; diff --git a/src/ertmp/multitrack_media.rs b/src/ertmp/multitrack_media.rs index 2c8292c..bf9698f 100644 --- a/src/ertmp/multitrack_media.rs +++ b/src/ertmp/multitrack_media.rs @@ -7,6 +7,9 @@ use crate::types::FrameType; pub const ERTMP_AUDIO_PACKET_TYPE_MULTITRACK: u8 = 5; pub const ERTMP_VIDEO_PACKET_TYPE_MULTITRACK: u8 = 6; +/// Cap sub-tracks unpacked from a single multitrack container (mirrors +/// `message::message::MAX_AGGREGATE_SUBTAGS` and `session::conn::MAX_AGGREGATE_SUBTAGS`). +pub const MAX_MULTITRACK_SUBTRACKS: usize = 4096; #[derive(Debug, Clone, Copy, PartialEq, Eq)] #[repr(u8)] @@ -106,6 +109,9 @@ pub fn foreach_track( if pos + track_size > payload.len() { return false; } + if tracks.len() >= MAX_MULTITRACK_SUBTRACKS { + return false; + } tracks.push(MediaTrackSlice { track_id, packet_type: inner_packet_type, @@ -217,4 +223,30 @@ mod tests { let payload = vec![0x96, 0x13, b'a', b'v', b'c', b'1', 0, 0, 0, 1, 0xAA]; assert!(multitrack_has_keyframe(&payload)); } + + fn build_many_tracks_zero_payload_message(track_count: usize) -> Vec { + let mut payload = vec![0x86, 0x10, b'a', b'v', b'c', b'1']; + for id in 0..track_count { + payload.push(id as u8); + payload.extend_from_slice(&[0x00, 0x00, 0x00]); + } + payload + } + + #[test] + fn rejects_multitrack_messages_with_too_many_subtracks() { + let at_limit = build_many_tracks_zero_payload_message(MAX_MULTITRACK_SUBTRACKS); + let mut calls = 0; + assert!(foreach_track(FrameType::Video, &at_limit, |_| calls += 1)); + assert_eq!(calls, MAX_MULTITRACK_SUBTRACKS); + + let over_limit = build_many_tracks_zero_payload_message(MAX_MULTITRACK_SUBTRACKS + 1); + let mut over_calls = 0; + assert!(!foreach_track( + FrameType::Video, + &over_limit, + |_| over_calls += 1 + )); + assert_eq!(over_calls, 0); + } } diff --git a/src/session/conn.rs b/src/session/conn.rs index ebf007b..8af81c8 100644 --- a/src/session/conn.rs +++ b/src/session/conn.rs @@ -7,7 +7,7 @@ use crate::chunk::reader::{ChunkMessage, chunk_read_owned}; use crate::chunk::state::{ChunkRegistry, DEFAULT_CHUNK_SIZE, DEFAULT_MAX_MSG_LENGTH}; use crate::chunk::writer::chunk_write; use crate::ertmp::connect_amf::{negotiate_caps, write_negotiated_caps}; -use crate::ertmp::multitrack_media::{first_track_fourcc, foreach_track}; +use crate::ertmp::multitrack_media::{first_track_fourcc, foreach_track, is_multitrack_container}; use crate::handshake::{self, Handshake, HandshakeState}; use crate::media::{ is_on_metadata_payload, normalize_modex_payload, populate_av_frame, populate_multitrack_frame, @@ -471,8 +471,10 @@ impl Conn { .media_bytes_received .saturating_add(payload.len() as u64); - if let Some(cb) = self.on_frame_cb { - let had_multitrack = foreach_track(frame_type, parse_payload, |track| { + let is_multitrack = is_multitrack_container(frame_type, parse_payload); + let cb = self.on_frame_cb; + let parsed_multitrack = foreach_track(frame_type, parse_payload, |track| { + if let Some(cb) = cb { self.invoke_multitrack_on_frame_cb( cb, frame_type, @@ -483,12 +485,16 @@ impl Conn { track.video_frame_type, track.payload, ); - }); - if !had_multitrack { + } + }); + if is_multitrack && !parsed_multitrack { + return Err(ErrorCode::Protocol); + } + if !is_multitrack { + if let Some(cb) = cb { self.invoke_on_frame_cb(cb, frame_type, timestamp, u8::MAX, parse_payload); } } - if self .queue_relay_frame(frame_type, timestamp, payload, parse_payload) .is_err()