--- a/src/association/mod.rs +++ b/src/association/mod.rs @@ -66,6 +66,8 @@ #[cfg(test)] mod association_test; +#[cfg(test)] +mod receive_limits_test; /// Reasons why an association might be lost #[non_exhaustive] @@ -200,6 +202,7 @@ handshake_completed: bool, max_send_message_size: u32, max_receive_message_size: u32, + receive_limits: Option, inflight_queue_length: usize, will_send_shutdown: bool, bytes_received: usize, @@ -323,6 +326,7 @@ handshake_completed: false, max_send_message_size: 0, max_receive_message_size: 0, + receive_limits: None, inflight_queue_length: 0, will_send_shutdown: false, bytes_received: 0, @@ -443,6 +447,7 @@ max_receive_buffer_size: config.max_receive_buffer_size(), max_send_message_size: config.max_send_message_size(), max_receive_message_size: config.max_receive_message_size(), + receive_limits: config.receive_limits(), my_max_num_outbound_streams: config.max_num_outbound_streams(), my_max_num_inbound_streams: config.max_num_inbound_streams(), max_payload_size, @@ -2363,6 +2368,20 @@ self.stats.inc_datas(); let can_push = self.payload_queue.can_push(d, self.peer_last_tsn); + if can_push { + self.check_receive_limits(d)?; + } + // A small fragment may otherwise retain a whole large packet through + // Bytes::slice. Compact only when enforcing the hard resource policy; + // subsequent queue clones share this one bounded payload allocation. + let mut compact; + let d = if can_push && self.receive_limits.is_some() { + compact = d.clone(); + compact.user_data = Bytes::copy_from_slice(&d.user_data); + &compact + } else { + d + }; let mut stream_handle_data = false; let mut defer_stream_data = false; if can_push && self.data_is_above_pending_reset(d) { @@ -3371,6 +3390,12 @@ accept: bool, default_payload_type: PayloadProtocolIdentifier, ) -> Option> { + if self + .receive_limits + .is_some_and(|limits| self.streams.len() >= limits.max_streams()) + { + return None; + } let s = StreamState::new( self.side, stream_identifier, @@ -3410,7 +3435,51 @@ } } + /// Count each retained TSN once, including DATA already delivered to the + /// application but still held behind a gap in the association payload queue. + fn retained_receive_data(&self) -> (usize, usize) { + let mut bytes = self.payload_queue.get_num_bytes(); + let mut chunks = self.payload_queue.len(); + for stream in self.streams.values() { + for chunk in stream.reassembly_queue.chunks() { + if self.payload_queue.get(chunk.tsn).is_none() { + bytes = bytes.saturating_add(chunk.user_data.len()); + chunks = chunks.saturating_add(1); + } + } + } + for chunk in self.deferred_reset_data.values() { + if self.payload_queue.get(chunk.tsn).is_none() { + bytes = bytes.saturating_add(chunk.user_data.len()); + chunks = chunks.saturating_add(1); + } + } + (bytes, chunks) + } + + fn check_receive_limits(&self, data: &ChunkPayloadData) -> Result<()> { + let Some(limits) = self.receive_limits else { + return Ok(()); + }; + let (bytes, chunks) = self.retained_receive_data(); + if data.user_data.len() > (limits.max_buffered_bytes() as usize).saturating_sub(bytes) + || chunks >= limits.max_buffered_chunks() + || (!self.streams.contains_key(&data.stream_identifier) + && self.streams.len() >= limits.max_streams()) + { + return Err(Error::ErrReceiveLimitExceeded); + } + Ok(()) + } + pub(crate) fn get_my_receiver_window_credit(&self) -> u32 { + if let Some(limits) = self.receive_limits { + let (bytes, chunks) = self.retained_receive_data(); + if chunks >= limits.max_buffered_chunks() { + return 0; + } + return self.max_receive_buffer_size.saturating_sub(bytes as u32); + } let mut bytes_queued = 0; for s in self.streams.values() { bytes_queued += s.get_num_bytes_in_reassembly_queue() as u32; --- /dev/null +++ b/src/association/receive_limits_test.rs @@ -0,0 +1,231 @@ +use super::*; +use crate::ReceiveLimits; +use crate::packet::PartialDecode; + +fn association(message: u32, bytes: u32, chunks: usize, streams: usize) -> Association { + Association { + state: AssociationState::Established, + peer_last_tsn: 0, + my_next_tsn: 1, + source_port: 5000, + destination_port: 5000, + max_receive_buffer_size: bytes, + max_receive_message_size: message, + receive_limits: Some(ReceiveLimits::new(message, bytes, chunks, streams)), + max_payload_size: 1200, + ..Default::default() + } +} + +fn data(tsn: u32, stream: u16, ssn: u16, len: usize) -> ChunkPayloadData { + ChunkPayloadData { + tsn, + stream_identifier: stream, + stream_sequence_number: ssn, + beginning_fragment: true, + ending_fragment: true, + payload_type: PayloadProtocolIdentifier::Binary, + user_data: Bytes::from(vec![42; len]), + ..Default::default() + } +} + +fn receive(a: &mut Association, chunk: ChunkPayloadData) { + let packet = a.create_packet(vec![Box::new(chunk)]).marshal().unwrap(); + let partial = PartialDecode::unmarshal(&packet).unwrap(); + a.handle_event(AssociationEvent(AssociationEventInner::Datagram( + Transmit { + now: Instant::now(), + remote: a.remote_addr, + ecn: None, + local_ip: None, + payload: Payload::PartialDecode(partial), + }, + ))); +} + +#[test] +fn policy_is_opt_in_and_preserves_high_stream_identifiers() { + assert!(TransportConfig::default().receive_limits().is_none()); + let limits = ReceiveLimits::new(8192, 32768, 64, 8); + let config = TransportConfig::default().with_receive_limits(limits); + assert_eq!(config.max_receive_buffer_size(), 32768); + assert_eq!(config.max_receive_message_size(), 8192); + assert_eq!(config.max_num_inbound_streams(), u16::MAX); + let overridden = config + .with_max_receive_message_size(u32::MAX) + .with_max_receive_buffer_size(u32::MAX); + assert_eq!(overridden.max_receive_message_size(), 8192); + assert_eq!(overridden.max_receive_buffer_size(), 32768); + assert_eq!( + overridden + .with_max_receive_message_size(1024) + .max_receive_message_size(), + 1024 + ); + let mut a = association(8192, 32768, 64, 8); + for (n, stream) in [0, 64, 257, 32768, 65530, 65534, 4000, 5000] + .into_iter() + .enumerate() + { + a.handle_data(&data(n as u32 + 1, stream, 0, 1)).unwrap(); + assert!(a.stream(stream).unwrap().read().unwrap().is_some()); + } + assert_eq!(a.streams.len(), 8); + assert_eq!( + a.handle_data(&data(9, 6000, 0, 1)).unwrap_err(), + Error::ErrReceiveLimitExceeded + ); + assert_eq!(a.streams.len(), 8); + assert!( + a.open_stream(65000, PayloadProtocolIdentifier::Binary) + .is_err() + ); +} + +#[test] +fn fragmented_message_limit_is_checked_before_reassembly_grows() { + for unordered in [false, true] { + let mut a = association(8192, 32768, 64, 8); + for n in 0..8 { + let mut chunk = data(n + 1, 65000, 0, 1024); + chunk.unordered = unordered; + chunk.beginning_fragment = n == 0; + chunk.ending_fragment = false; + a.handle_data(&chunk).unwrap(); + } + assert_eq!(a.retained_receive_data(), (8192, 8)); + let mut tail = data(9, 65000, 0, 1); + tail.unordered = unordered; + tail.beginning_fragment = false; + assert_eq!( + a.handle_data(&tail).unwrap_err(), + Error::ErrInboundPacketTooLarge + ); + assert_eq!(a.retained_receive_data(), (8192, 8)); + } +} + +#[test] +fn all_receive_queues_share_the_byte_and_fragment_budget() { + let mut a = association(16, 32, 64, 8); + // Complete data on an unordered stream can be consumed while a missing + // earlier TSN keeps the same payload alive in the association queue. + let mut first = data(2, 0, 0, 16); + first.unordered = true; + a.handle_data(&first).unwrap(); + assert_eq!(a.retained_receive_data(), (16, 1)); + drop(a.stream(0).unwrap().read().unwrap().unwrap()); + assert_eq!(a.get_my_receiver_window_credit(), 16); + assert_eq!(a.retained_receive_data(), (16, 1)); + let mut second = data(3, 65000, 0, 16); + second.ending_fragment = false; + a.handle_data(&second).unwrap(); + assert_eq!(a.retained_receive_data(), (32, 2)); + // Missing TSNs are not allowed to bypass the hard resource budget. + assert_eq!( + a.handle_data(&data(1, 0, 1, 1)).unwrap_err(), + Error::ErrReceiveLimitExceeded + ); + assert_eq!(a.retained_receive_data(), (32, 2)); +} + +#[test] +fn missing_tsn_can_fill_a_gap_and_restore_receive_credit_within_the_budget() { + let mut a = association(16, 32, 64, 8); + let mut first = data(2, 0, 0, 8); + first.unordered = true; + a.handle_data(&first).unwrap(); + drop(a.stream(0).unwrap().read().unwrap().unwrap()); + assert_eq!(a.get_my_receiver_window_credit(), 24); + let mut partial = data(3, 65000, 0, 8); + partial.ending_fragment = false; + a.handle_data(&partial).unwrap(); + assert_eq!(a.get_my_receiver_window_credit(), 16); + a.handle_data(&data(1, 257, 0, 8)).unwrap(); + assert_eq!(a.peer_last_tsn, 3); + assert!(a.payload_queue.is_empty()); + drop(a.stream(257).unwrap().read().unwrap().unwrap()); + assert_eq!(a.get_my_receiver_window_credit(), 24); + let mut tail = data(4, 65000, 0, 8); + tail.beginning_fragment = false; + a.handle_data(&tail).unwrap(); + assert_eq!(a.stream(65000).unwrap().read().unwrap().unwrap().len(), 16); + assert_eq!(a.retained_receive_data(), (0, 0)); + assert_eq!(a.get_my_receiver_window_credit(), 32); + assert!(!a.is_closed()); +} + +#[test] +fn tiny_incomplete_messages_cannot_exhaust_chunk_metadata() { + let mut a = association(8192, 32768, 64, 8); + for n in 0..64 { + let mut chunk = data(n + 1, 0, n as u16, 1); + chunk.ending_fragment = false; + a.handle_data(&chunk).unwrap(); + } + assert_eq!(a.retained_receive_data(), (64, 64)); + assert_eq!(a.get_my_receiver_window_credit(), 0); + assert_eq!( + a.handle_data(&data(65, 0, 64, 1)).unwrap_err(), + Error::ErrReceiveLimitExceeded + ); + assert_eq!(a.retained_receive_data(), (64, 64)); +} + +#[test] +fn reset_deferred_data_is_counted_before_new_generation_is_readable() { + let mut a = association(8, 16, 64, 8); + a.handle_data(&data(1, 0, 0, 8)).unwrap(); + let reset: Box = Box::new(ParamOutgoingResetRequest { + reconfig_request_sequence_number: 7, + reconfig_response_sequence_number: u32::MAX, + sender_last_tsn: 1, + stream_identifiers: vec![0], + }); + a.handle_reconfig_param(&reset, &mut vec![]).unwrap(); + assert!(a.retiring_streams.contains_key(&0)); + a.handle_data(&data(2, 0, 0, 8)).unwrap(); + assert!(a.deferred_reset_data.contains_key(&2)); + assert_eq!(a.retained_receive_data(), (16, 2)); + assert_eq!( + a.handle_data(&data(3, 0, 1, 1)).unwrap_err(), + Error::ErrReceiveLimitExceeded + ); + assert_eq!(a.retained_receive_data(), (16, 2)); + drop(a.stream(0).unwrap().read().unwrap().unwrap()); + assert!(a.retained_receive_data().0 <= 8); +} + +#[test] +fn compact_payloads_do_not_keep_unrelated_packet_bytes_alive() { + let mut a = association(8, 16, 64, 8); + let packet = Bytes::from(vec![42; 8192]); + let mut chunk = data(2, 0, 0, 1); + chunk.user_data = packet.slice(4000..4001); + a.handle_data(&chunk).unwrap(); + let retained = &a.payload_queue.get(2).unwrap().user_data; + assert_eq!(retained.as_ref(), chunk.user_data.as_ref()); + assert_ne!(retained.as_ptr(), chunk.user_data.as_ptr()); + let queued = &a.streams.get(&0).unwrap().reassembly_queue.ordered[0].chunks[0]; + assert_eq!(retained.as_ptr(), queued.user_data.as_ptr()); +} + +#[test] +fn overflow_on_the_wire_closes_the_association_and_fresh_peer_can_receive() { + let mut a = association(8, 8, 2, 8); + receive(&mut a, data(1, 0, 0, 8)); + assert!(!a.is_closed()); + receive(&mut a, data(2, 0, 1, 1)); + assert!(a.is_closed()); + assert!(core::iter::from_fn(|| a.poll()).any(|e| matches!(e, Event::AssociationLost { .. }))); + assert!(a.streams.is_empty()); + assert!(a.deferred_reset_data.is_empty()); + let mut fresh = association(8, 8, 2, 8); + receive(&mut fresh, data(1, 65000, 0, 8)); + assert!(!fresh.is_closed()); + assert_eq!( + fresh.stream(65000).unwrap().read().unwrap().unwrap().len(), + 8 + ); +} --- a/src/config.rs +++ b/src/config.rs @@ -22,11 +22,72 @@ // Default max retransmit value (RFC 4960 Section 15) const DEFAULT_MAX_INIT_RETRANS: usize = 8; +/// Optional hard limits on retained inbound DATA state. +/// +/// These limits apply before adding a new fragment, including data retained for +/// missing TSNs and stream resets. They do not bound all association heap use, +/// allocator overhead, or control-chunk state. Exceeding a limit closes the +/// association; the limits are a resource policy, not SCTP flow control. +#[derive(Debug, Clone, Copy)] +pub struct ReceiveLimits { + max_message_size: u32, + max_buffered_bytes: u32, + max_buffered_chunks: usize, + max_streams: usize, +} + +impl ReceiveLimits { + /// Construct a receive resource policy. + /// + /// Stream count is the number of live stream states, independent of stream + /// identifier values. It includes locally opened streams. + /// + /// # Panics + /// + /// Panics if any limit is zero or a message cannot fit in the byte budget. + pub fn new( + max_message_size: u32, + max_buffered_bytes: u32, + max_buffered_chunks: usize, + max_streams: usize, + ) -> Self { + assert!(max_message_size > 0 && max_message_size <= max_buffered_bytes); + assert!(max_buffered_chunks > 0 && max_streams > 0); + Self { + max_message_size, + max_buffered_bytes, + max_buffered_chunks, + max_streams, + } + } + + /// Maximum size of an individual received message, enforced in reassembly. + pub fn max_message_size(self) -> u32 { + self.max_message_size + } + + /// Maximum retained DATA payload bytes across receive queues. + pub fn max_buffered_bytes(self) -> u32 { + self.max_buffered_bytes + } + + /// Maximum retained DATA fragments across receive queues. + pub fn max_buffered_chunks(self) -> usize { + self.max_buffered_chunks + } + + /// Maximum live stream states; this does not limit identifier values. + pub fn max_streams(self) -> usize { + self.max_streams + } +} + /// Config collects the arguments to create_association construction into /// a single structure #[derive(Debug)] pub struct TransportConfig { max_receive_buffer_size: u32, + receive_limits: Option, max_num_outbound_streams: u16, max_num_inbound_streams: u16, @@ -65,6 +126,7 @@ fn default() -> Self { TransportConfig { max_receive_buffer_size: INITIAL_RECV_BUF_SIZE, + receive_limits: None, max_send_message_size: DEFAULT_MAX_MESSAGE_SIZE, max_receive_message_size: DEFAULT_MAX_MESSAGE_SIZE, max_num_outbound_streams: u16::MAX, @@ -79,6 +141,21 @@ } impl TransportConfig { + /// Set an optional hard policy for retained inbound DATA state. + /// + /// Also sets the advertised receive window and per-message limit. The hard + /// policy is disabled by default, preserving ordinary SCTP window behavior. + pub fn with_receive_limits(mut self, limits: ReceiveLimits) -> Self { + self.max_receive_buffer_size = limits.max_buffered_bytes; + self.max_receive_message_size = limits.max_message_size; + self.receive_limits = Some(limits); + self + } + + pub(crate) fn receive_limits(&self) -> Option { + self.receive_limits + } + pub fn with_max_receive_buffer_size(mut self, value: u32) -> Self { self.max_receive_buffer_size = value; self @@ -111,7 +188,10 @@ } pub(crate) fn max_receive_buffer_size(&self) -> u32 { - self.max_receive_buffer_size + self.receive_limits + .map_or(self.max_receive_buffer_size, |limits| { + self.max_receive_buffer_size.min(limits.max_buffered_bytes) + }) } pub(crate) fn max_send_message_size(&self) -> u32 { @@ -119,7 +199,10 @@ } pub(crate) fn max_receive_message_size(&self) -> u32 { - self.max_receive_message_size + self.receive_limits + .map_or(self.max_receive_message_size, |limits| { + self.max_receive_message_size.min(limits.max_message_size) + }) } pub(crate) fn max_num_outbound_streams(&self) -> u16 { --- a/src/error.rs +++ b/src/error.rs @@ -217,6 +217,8 @@ ErrOutboundPacketTooLarge, #[error("inbound packet larger than maximum message size")] ErrInboundPacketTooLarge, + #[error("inbound DATA resource limit exceeded")] + ErrReceiveLimitExceeded, #[error("Stream closed")] ErrStreamClosed, #[error("Stream not existed")] --- a/src/lib.rs +++ b/src/lib.rs @@ -75,8 +75,8 @@ mod config; pub use crate::config::{ - ClientConfig, DEFAULT_SCTP_PORT, EndpointConfig, MAX_SNAP_INIT_BYTES, ServerConfig, - TransportConfig, generate_snap_token, + ClientConfig, DEFAULT_SCTP_PORT, EndpointConfig, MAX_SNAP_INIT_BYTES, ReceiveLimits, + ServerConfig, TransportConfig, generate_snap_token, }; mod endpoint; --- a/src/packet.rs +++ b/src/packet.rs @@ -317,7 +317,21 @@ } pub(crate) fn marshal(&self) -> Result { - let mut buf = BytesMut::with_capacity(PACKET_HEADER_SIZE); + // Chunk lengths are known before writing. Reserve the complete padded + // packet once instead of growing from the common header for every send. + let capacity = self + .chunks + .iter() + .try_fold(PACKET_HEADER_SIZE, |total, chunk| { + let length = CHUNK_HEADER_SIZE + .checked_add(chunk.value_length()) + .ok_or(Error::ErrOutboundPacketTooLarge)?; + total + .checked_add(length) + .and_then(|n| n.checked_add(get_padding_size(length))) + .ok_or(Error::ErrOutboundPacketTooLarge) + })?; + let mut buf = BytesMut::with_capacity(capacity); self.marshal_to(&mut buf)?; Ok(buf.freeze()) } --- a/src/queue/reassembly_queue.rs +++ b/src/queue/reassembly_queue.rs @@ -223,6 +223,15 @@ n_bytes: 0, max_message_size, } + } + + /// Every retained DATA fragment, complete or still awaiting reassembly. + pub(crate) fn chunks(&self) -> impl Iterator { + self.ordered + .iter() + .chain(&self.unordered) + .flat_map(|message| &message.chunks) + .chain(&self.unordered_chunks) } pub(crate) fn push(&mut self, chunk: ChunkPayloadData) -> Result {