diff --git a/crates/lanspread-peer/src/network.rs b/crates/lanspread-peer/src/network.rs index 13260a3..434b8b1 100644 --- a/crates/lanspread-peer/src/network.rs +++ b/crates/lanspread-peer/src/network.rs @@ -6,7 +6,8 @@ use std::{ time::Duration, }; -use futures::{SinkExt, StreamExt}; +use bytes::Bytes; +use futures::SinkExt; use if_addrs::{IfAddr, Interface, get_if_addrs}; use lanspread_proto::{ ChangeHint, @@ -18,8 +19,9 @@ use lanspread_proto::{ Request, Response, }; +use tokio::io::{AsyncRead, AsyncReadExt as _}; use tokio_util::{ - codec::{FramedRead, FramedWrite, LengthDelimitedCodec}, + codec::{FramedWrite, LengthDelimitedCodec}, sync::CancellationToken, }; @@ -72,25 +74,13 @@ async fn send_oneway_request( let mut conn = connector.connect(endpoint).await?; let stream = conn.open_bidirectional_stream().await?; - let (rx, tx) = stream.split(); - let mut framed_rx = FramedRead::new(rx, control_codec()); + let (mut rx, tx) = stream.split(); let mut framed_tx = FramedWrite::new(tx, control_codec()); framed_tx.send(request.encode()?).await?; framed_tx.close().await?; - match framed_rx.next().await { - None => {} - Some(Ok(frame)) => { - let response = Response::decode(frame.freeze())?; - eyre::bail!( - "peer {} unexpectedly responded to one-way request: {response:?}", - endpoint.addr - ); - } - Some(Err(error)) => return Err(error.into()), - } - Ok(()) + require_oneway_response_eof(&mut rx, endpoint.addr).await }) .await } @@ -124,27 +114,13 @@ async fn exchange_request( let mut conn = connector.connect(endpoint).await?; let stream = conn.open_bidirectional_stream().await?; - let (rx, tx) = stream.split(); - let mut framed_rx = FramedRead::new(rx, control_codec()); + let (mut rx, tx) = stream.split(); let mut framed_tx = FramedWrite::new(tx, control_codec()); framed_tx.send(request.encode()?).await?; framed_tx.close().await?; - let frame = framed_rx - .next() - .await - .ok_or_else(|| eyre::eyre!("peer {} returned no control response", endpoint.addr))??; - let response = Response::decode(frame.freeze())?; - - match framed_rx.next().await { - None => Ok(response), - Some(Ok(_)) => eyre::bail!( - "peer {} returned more than one control response frame", - endpoint.addr - ), - Some(Err(error)) => Err(error.into()), - } + read_control_response(&mut rx, connector, endpoint.addr).await }) .await } @@ -185,6 +161,106 @@ fn control_codec() -> LengthDelimitedCodec { .new_codec() } +const CONTROL_FRAME_LENGTH_BYTES: usize = size_of::(); +const CONTROL_RESPONSE_READ_CHUNK_BYTES: usize = 64 * 1024; + +/// Reads one response without allowing its untrusted length prefix to reserve +/// the full frame independently of every other outbound request. +async fn read_control_response( + reader: &mut R, + connector: &QuicConnector, + peer_addr: SocketAddr, +) -> eyre::Result +where + R: AsyncRead + Unpin, +{ + let mut length_bytes = [0; CONTROL_FRAME_LENGTH_BYTES]; + let mut length_bytes_read = 0; + while length_bytes_read < length_bytes.len() { + let read = reader.read(&mut length_bytes[length_bytes_read..]).await?; + if read == 0 { + if length_bytes_read == 0 { + eyre::bail!("peer {peer_addr} returned no control response"); + } + eyre::bail!( + "peer {peer_addr} returned a truncated control response length prefix ({length_bytes_read} of {} bytes)", + length_bytes.len() + ); + } + length_bytes_read += read; + } + + let frame_len_u32 = u32::from_be_bytes(length_bytes); + let frame_len = usize::try_from(frame_len_u32) + .map_err(|_| eyre::eyre!("control response length does not fit usize"))?; + if frame_len > MAX_CONTROL_FRAME_BYTES { + eyre::bail!( + "peer {peer_addr} declared a {frame_len}-byte control response; maximum is {MAX_CONTROL_FRAME_BYTES} bytes" + ); + } + + let _response_bytes = connector + .reserve_control_response_bytes(frame_len_u32) + .await?; + let mut frame = Vec::new(); + while frame.len() < frame_len { + if frame.len() == frame.capacity() { + let requested_capacity = frame + .capacity() + .saturating_mul(2) + .max(CONTROL_RESPONSE_READ_CHUNK_BYTES) + .min(frame_len); + frame + .try_reserve_exact(requested_capacity - frame.len()) + .map_err(|error| { + eyre::eyre!("failed to reserve control response bytes: {error}") + })?; + } + + let chunk_start = frame.len(); + let chunk_end = frame.capacity().min(frame_len); + frame.resize(chunk_end, 0); + let mut body_bytes_read = chunk_start; + while body_bytes_read < chunk_end { + let read = reader.read(&mut frame[body_bytes_read..chunk_end]).await?; + if read == 0 { + frame.truncate(body_bytes_read); + eyre::bail!( + "peer {peer_addr} returned a truncated control response body ({} of {frame_len} bytes)", + frame.len() + ); + } + body_bytes_read += read; + } + } + + require_control_response_eof(reader, peer_addr).await?; + let response = Response::decode(Bytes::from(frame))?; + Ok(response) +} + +async fn require_control_response_eof(reader: &mut R, peer_addr: SocketAddr) -> eyre::Result<()> +where + R: AsyncRead + Unpin, +{ + let mut trailing = [0]; + if reader.read(&mut trailing).await? != 0 { + eyre::bail!("peer {peer_addr} returned trailing bytes after one control response frame"); + } + Ok(()) +} + +async fn require_oneway_response_eof(reader: &mut R, peer_addr: SocketAddr) -> eyre::Result<()> +where + R: AsyncRead + Unpin, +{ + let mut unexpected = [0]; + if reader.read(&mut unexpected).await? != 0 { + eyre::bail!("peer {peer_addr} unexpectedly responded to a one-way request"); + } + Ok(()) +} + async fn run_short_request( peer_addr: SocketAddr, cancellation: &CancellationToken, @@ -311,10 +387,20 @@ mod tests { time::Duration, }; - use tokio::sync::oneshot; + use lanspread_proto::{ + ControlErrorCode, + ControlMessage as _, + MAX_CONTROL_FRAME_BYTES, + Response, + }; + use tokio::{ + io::{AsyncWriteExt as _, DuplexStream, duplex}, + sync::oneshot, + }; use tokio_util::sync::CancellationToken; - use super::run_network_operation; + use super::{read_control_response, require_oneway_response_eof, run_network_operation}; + use crate::quic_runtime::QuicConnector; struct DropProbe(Arc); @@ -328,6 +414,102 @@ mod tests { SocketAddr::from(([127, 0, 0, 1], 12345)) } + async fn closed_reader(bytes: &[u8]) -> DuplexStream { + let (mut writer, reader) = duplex(bytes.len().saturating_add(1)); + writer + .write_all(bytes) + .await + .expect("test response should fit its duplex buffer"); + writer + .shutdown() + .await + .expect("test response writer should close"); + reader + } + + fn framed_response(body: &[u8]) -> Vec { + let mut frame = u32::try_from(body.len()) + .expect("test response length should fit u32") + .to_be_bytes() + .to_vec(); + frame.extend_from_slice(body); + frame + } + + #[tokio::test] + async fn bounded_response_reader_accepts_one_complete_frame() { + let encoded = Response::Error(ControlErrorCode::InvalidRequest) + .encode() + .expect("test response should encode"); + let mut reader = closed_reader(&framed_response(&encoded)).await; + let connector = QuicConnector::unavailable(); + + assert!(matches!( + read_control_response(&mut reader, &connector, peer_addr()) + .await + .expect("one complete response should decode"), + Response::Error(ControlErrorCode::InvalidRequest) + )); + } + + #[tokio::test] + async fn bounded_response_reader_rejects_oversized_prefix_before_body() { + let declared = + u32::try_from(MAX_CONTROL_FRAME_BYTES + 1).expect("control frame bound should fit u32"); + let mut reader = closed_reader(&declared.to_be_bytes()).await; + let connector = QuicConnector::unavailable(); + + let error = read_control_response(&mut reader, &connector, peer_addr()) + .await + .expect_err("oversized response must be rejected"); + assert!(format!("{error:#}").contains("maximum is")); + } + + #[tokio::test] + async fn bounded_response_reader_rejects_truncated_body() { + let mut response = 10_u32.to_be_bytes().to_vec(); + response.extend_from_slice(b"short"); + let mut reader = closed_reader(&response).await; + let connector = QuicConnector::unavailable(); + + let error = read_control_response(&mut reader, &connector, peer_addr()) + .await + .expect_err("truncated response must be rejected"); + assert!(format!("{error:#}").contains("truncated control response body (5 of 10 bytes)")); + } + + #[tokio::test] + async fn bounded_response_reader_rejects_trailing_frame_bytes() { + let encoded = Response::Error(ControlErrorCode::InvalidRequest) + .encode() + .expect("test response should encode"); + let mut response = framed_response(&encoded); + response.extend_from_slice(&1_u32.to_be_bytes()); + response.push(b'x'); + let mut reader = closed_reader(&response).await; + let connector = QuicConnector::unavailable(); + + let error = read_control_response(&mut reader, &connector, peer_addr()) + .await + .expect_err("trailing response frame must be rejected"); + assert!(format!("{error:#}").contains("trailing bytes")); + } + + #[tokio::test] + async fn one_way_response_path_requires_immediate_eof() { + let mut empty = closed_reader(&[]).await; + require_oneway_response_eof(&mut empty, peer_addr()) + .await + .expect("empty response side should be accepted"); + + let mut response = closed_reader(&[0]).await; + assert!( + require_oneway_response_eof(&mut response, peer_addr()) + .await + .is_err() + ); + } + #[tokio::test] async fn cancellation_drops_operation_scope_before_returning() { let cancellation = CancellationToken::new(); diff --git a/crates/lanspread-peer/src/quic_runtime.rs b/crates/lanspread-peer/src/quic_runtime.rs index 90d1527..b5e2883 100644 --- a/crates/lanspread-peer/src/quic_runtime.rs +++ b/crates/lanspread-peer/src/quic_runtime.rs @@ -33,7 +33,10 @@ use s2n_quic_core::{ path::mtu, time::{Clock, Timestamp}, }; -use tokio::task::{JoinError, JoinHandle}; +use tokio::{ + sync::{OwnedSemaphorePermit, Semaphore}, + task::{JoinError, JoinHandle}, +}; use tokio_util::sync::{CancellationToken, WaitForCancellationFutureOwned}; use crate::{ @@ -62,6 +65,12 @@ pub(crate) const SERVER_CONTROL_RECEIVE_WINDOW: u64 = MAX_REQUEST_FRAME_BYTES as /// open request stream may hold its full per-stream allowance. pub(crate) const SERVER_CONNECTION_RECEIVE_WINDOW: u64 = SERVER_CONTROL_RECEIVE_WINDOW * MAX_OPEN_BIDIRECTIONAL_STREAMS; +/// Aggregate application buffer budget for concurrently received control responses. +/// +/// Response frames may legitimately approach 8 MiB, but discovery and refresh +/// tasks must not each reserve that allowance independently after reading only a +/// length prefix. All connector clones in one network generation share this pool. +pub(crate) const MAX_INFLIGHT_CONTROL_RESPONSE_BYTES: usize = 32 * 1024 * 1024; pub(crate) fn quic_client_limits() -> eyre::Result { Ok(Limits::default() @@ -349,6 +358,7 @@ fn join_result(result: Result<(), JoinError>, expected_abort: bool) -> eyre::Res #[derive(Clone, Debug)] pub(crate) struct QuicConnector { client: Option, + control_response_bytes: Arc, } impl QuicConnector { @@ -362,9 +372,22 @@ impl QuicConnector { Ok(PeerConnection::new(client.connect(connect).await?)) } + pub(crate) async fn reserve_control_response_bytes( + &self, + byte_count: u32, + ) -> eyre::Result { + Arc::clone(&self.control_response_bytes) + .acquire_many_owned(byte_count) + .await + .map_err(|_| eyre::eyre!("control-response byte budget was closed")) + } + #[cfg(test)] pub(crate) fn unavailable() -> Self { - Self { client: None } + Self { + client: None, + control_response_bytes: Arc::new(Semaphore::new(MAX_INFLIGHT_CONTROL_RESPONSE_BYTES)), + } } } @@ -427,6 +450,7 @@ pub(crate) fn start_quic_client() -> eyre::Result<(QuicClientRuntime, QuicConnec let endpoint = endpoint_control.take_started()?; let connector = QuicConnector { client: Some(client.clone()), + control_response_bytes: Arc::new(Semaphore::new(MAX_INFLIGHT_CONTROL_RESPONSE_BYTES)), }; Ok((QuicClientRuntime { client, endpoint }, connector)) @@ -468,10 +492,17 @@ impl Drop for PeerConnection { mod tests { use std::{sync::Arc, time::Duration}; + use lanspread_proto::MAX_CONTROL_FRAME_BYTES; use tokio::sync::Notify; use tokio_util::sync::CancellationToken; - use super::{EndpointTask, EndpointTaskControl, start_quic_client}; + use super::{ + EndpointTask, + EndpointTaskControl, + MAX_INFLIGHT_CONTROL_RESPONSE_BYTES, + QuicConnector, + start_quic_client, + }; #[test] fn endpoint_control_rejects_take_before_provider_start() { @@ -490,6 +521,46 @@ mod tests { .expect("client endpoint should join cleanly"); } + #[tokio::test] + async fn connector_clones_share_and_release_the_control_response_budget() { + assert_eq!( + MAX_INFLIGHT_CONTROL_RESPONSE_BYTES % MAX_CONTROL_FRAME_BYTES, + 0 + ); + let max_frame = + u32::try_from(MAX_CONTROL_FRAME_BYTES).expect("control frame bound should fit u32"); + let connector = QuicConnector::unavailable(); + let connector_clone = connector.clone(); + let mut permits = Vec::new(); + for _ in 0..MAX_INFLIGHT_CONTROL_RESPONSE_BYTES / MAX_CONTROL_FRAME_BYTES { + permits.push( + connector + .reserve_control_response_bytes(max_frame) + .await + .expect("configured aggregate budget should admit its exact capacity"), + ); + } + + let blocked = connector_clone.reserve_control_response_bytes(1); + tokio::pin!(blocked); + tokio::select! { + biased; + result = blocked.as_mut() => { + panic!("response budget exceeded its configured capacity: {result:?}"); + } + () = tokio::task::yield_now() => {} + } + + let released = permits + .pop() + .expect("the exact-capacity fixture should hold at least one permit"); + drop(released); + let reacquired = blocked + .await + .expect("dropping a permit should release shared response capacity"); + drop(reacquired); + } + #[tokio::test] async fn endpoint_owner_signals_and_joins_its_task() { let stop = CancellationToken::new();