diff --git a/crates/lanspread-peer/src/network.rs b/crates/lanspread-peer/src/network.rs index 434b8b1..aca8483 100644 --- a/crates/lanspread-peer/src/network.rs +++ b/crates/lanspread-peer/src/network.rs @@ -200,7 +200,7 @@ where } let _response_bytes = connector - .reserve_control_response_bytes(frame_len_u32) + .reserve_control_response_bytes(peer_addr.ip(), frame_len_u32) .await?; let mut frame = Vec::new(); while frame.len() < frame_len { diff --git a/crates/lanspread-peer/src/quic_runtime.rs b/crates/lanspread-peer/src/quic_runtime.rs index b5e2883..716c2e4 100644 --- a/crates/lanspread-peer/src/quic_runtime.rs +++ b/crates/lanspread-peer/src/quic_runtime.rs @@ -6,12 +6,13 @@ //! still running. use std::{ + collections::HashMap, future::Future as _, io, - net::SocketAddr, + net::{IpAddr, SocketAddr}, ops::{Deref, DerefMut}, pin::Pin, - sync::{Arc, Mutex}, + sync::{Arc, Mutex, Weak}, task::{Context, Poll}, }; @@ -71,6 +72,10 @@ pub(crate) const SERVER_CONNECTION_RECEIVE_WINDOW: u64 = /// 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; +/// One endpoint IP may reserve at most one maximal response. This prevents +/// same-origin large-prefix waiters from entering the fair global semaphore +/// queue ahead of small responses from other LAN hosts. +pub(crate) const MAX_INFLIGHT_CONTROL_RESPONSE_BYTES_PER_IP: usize = MAX_CONTROL_FRAME_BYTES; pub(crate) fn quic_client_limits() -> eyre::Result { Ok(Limits::default() @@ -358,7 +363,42 @@ fn join_result(result: Result<(), JoinError>, expected_abort: bool) -> eyre::Res #[derive(Clone, Debug)] pub(crate) struct QuicConnector { client: Option, - control_response_bytes: Arc, + control_response_bytes: Arc, +} + +#[derive(Debug)] +struct ControlResponseBudget { + global: Arc, + per_ip: Mutex>>, +} + +impl ControlResponseBudget { + fn new() -> Self { + Self { + global: Arc::new(Semaphore::new(MAX_INFLIGHT_CONTROL_RESPONSE_BYTES)), + per_ip: Mutex::new(HashMap::new()), + } + } + + fn origin(&self, ip: IpAddr) -> Arc { + let mut per_ip = self + .per_ip + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner); + per_ip.retain(|_, budget| budget.strong_count() > 0); + if let Some(budget) = per_ip.get(&ip).and_then(Weak::upgrade) { + return budget; + } + let budget = Arc::new(Semaphore::new(MAX_INFLIGHT_CONTROL_RESPONSE_BYTES_PER_IP)); + per_ip.insert(ip, Arc::downgrade(&budget)); + budget + } +} + +#[derive(Debug)] +pub(crate) struct ControlResponsePermit { + _origin: OwnedSemaphorePermit, + _global: OwnedSemaphorePermit, } impl QuicConnector { @@ -374,19 +414,30 @@ impl QuicConnector { pub(crate) async fn reserve_control_response_bytes( &self, + peer_ip: IpAddr, byte_count: u32, - ) -> eyre::Result { - Arc::clone(&self.control_response_bytes) + ) -> eyre::Result { + let origin = self + .control_response_bytes + .origin(peer_ip) .acquire_many_owned(byte_count) .await - .map_err(|_| eyre::eyre!("control-response byte budget was closed")) + .map_err(|_| eyre::eyre!("per-origin control-response byte budget was closed"))?; + let global = Arc::clone(&self.control_response_bytes.global) + .acquire_many_owned(byte_count) + .await + .map_err(|_| eyre::eyre!("global control-response byte budget was closed"))?; + Ok(ControlResponsePermit { + _origin: origin, + _global: global, + }) } #[cfg(test)] pub(crate) fn unavailable() -> Self { Self { client: None, - control_response_bytes: Arc::new(Semaphore::new(MAX_INFLIGHT_CONTROL_RESPONSE_BYTES)), + control_response_bytes: Arc::new(ControlResponseBudget::new()), } } } @@ -450,7 +501,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)), + control_response_bytes: Arc::new(ControlResponseBudget::new()), }; Ok((QuicClientRuntime { client, endpoint }, connector)) @@ -490,7 +541,7 @@ impl Drop for PeerConnection { #[cfg(test)] mod tests { - use std::{sync::Arc, time::Duration}; + use std::{net::IpAddr, sync::Arc, time::Duration}; use lanspread_proto::MAX_CONTROL_FRAME_BYTES; use tokio::sync::Notify; @@ -532,16 +583,23 @@ mod tests { 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 { + for index in 0..MAX_INFLIGHT_CONTROL_RESPONSE_BYTES / MAX_CONTROL_FRAME_BYTES { + let source_ip = IpAddr::V4(std::net::Ipv4Addr::new( + 192, + 0, + 2, + u8::try_from(index + 1).expect("test source index should fit u8"), + )); permits.push( connector - .reserve_control_response_bytes(max_frame) + .reserve_control_response_bytes(source_ip, max_frame) .await .expect("configured aggregate budget should admit its exact capacity"), ); } - let blocked = connector_clone.reserve_control_response_bytes(1); + let blocked = connector_clone + .reserve_control_response_bytes(IpAddr::V4(std::net::Ipv4Addr::new(192, 0, 2, 100)), 1); tokio::pin!(blocked); tokio::select! { biased; @@ -561,6 +619,41 @@ mod tests { drop(reacquired); } + #[tokio::test] + async fn same_origin_large_waiter_cannot_block_another_origins_small_response() { + let connector = QuicConnector::unavailable(); + let same_origin = IpAddr::V4(std::net::Ipv4Addr::new(192, 0, 2, 10)); + let other_origin = IpAddr::V4(std::net::Ipv4Addr::new(192, 0, 2, 11)); + let max_frame = + u32::try_from(MAX_CONTROL_FRAME_BYTES).expect("control frame bound should fit u32"); + let held = connector + .reserve_control_response_bytes(same_origin, max_frame) + .await + .expect("one maximal response should fit its origin budget"); + + let blocked_same_origin = connector.reserve_control_response_bytes(same_origin, 1); + tokio::pin!(blocked_same_origin); + tokio::select! { + biased; + result = blocked_same_origin.as_mut() => { + panic!("same origin exceeded its response budget: {result:?}"); + } + () = tokio::task::yield_now() => {} + } + + let other = connector + .reserve_control_response_bytes(other_origin, 1) + .await + .expect("same-origin waiter must not enter the global queue first"); + drop(other); + drop(held); + drop( + blocked_same_origin + .await + .expect("releasing the origin budget should wake its waiter"), + ); + } + #[tokio::test] async fn endpoint_owner_signals_and_joins_its_task() { let stop = CancellationToken::new();