diff --git a/perf/src/bin/perf_server.rs b/perf/src/bin/perf_server.rs index 169ca65e4..1069d170f 100644 --- a/perf/src/bin/perf_server.rs +++ b/perf/src/bin/perf_server.rs @@ -123,7 +123,7 @@ async fn run(opt: Opt) -> Result<()> { Ok(()) } -async fn handle(handshake: quinn::Connecting, opt: Arc) -> Result<()> { +async fn handle(handshake: quinn::Incoming, opt: Arc) -> Result<()> { let connection = handshake.await.context("handshake failed")?; debug!("{} connected", connection.remote_address()); tokio::try_join!( diff --git a/quinn-proto/src/config.rs b/quinn-proto/src/config.rs index 9194db374..04f7c430e 100644 --- a/quinn-proto/src/config.rs +++ b/quinn-proto/src/config.rs @@ -741,10 +741,6 @@ pub struct ServerConfig { /// Used to generate one-time AEAD keys to protect handshake tokens pub(crate) token_key: Arc, - /// Whether to require clients to prove ownership of an address before committing resources. - /// - /// Introduces an additional round-trip to the handshake to make denial of service attacks more difficult. - pub(crate) use_retry: bool, /// Microseconds after a stateless retry token was issued for which it's considered valid. pub(crate) retry_token_lifetime: Duration, @@ -769,7 +765,6 @@ impl ServerConfig { crypto, token_key, - use_retry: false, retry_token_lifetime: Duration::from_secs(15), concurrent_connections: 100_000, @@ -790,14 +785,6 @@ impl ServerConfig { self } - /// Whether to require clients to prove ownership of an address before committing resources. - /// - /// Introduces an additional round-trip to the handshake to make denial of service attacks more difficult. - pub fn use_retry(&mut self, value: bool) -> &mut Self { - self.use_retry = value; - self - } - /// Duration after a stateless retry token was issued for which it's considered valid. pub fn retry_token_lifetime(&mut self, value: Duration) -> &mut Self { self.retry_token_lifetime = value; @@ -858,7 +845,6 @@ impl fmt::Debug for ServerConfig { .field("transport", &self.transport) .field("crypto", &"ServerConfig { elided }") .field("token_key", &"[ elided ]") - .field("use_retry", &self.use_retry) .field("retry_token_lifetime", &self.retry_token_lifetime) .field("concurrent_connections", &self.concurrent_connections) .field("migration", &self.migration) diff --git a/quinn-proto/src/endpoint.rs b/quinn-proto/src/endpoint.rs index fd4396922..814803cf7 100644 --- a/quinn-proto/src/endpoint.rs +++ b/quinn-proto/src/endpoint.rs @@ -246,7 +246,7 @@ impl Endpoint { return match first_decode.finish(Some(&*crypto.header.remote)) { Ok(packet) => { - self.handle_first_packet(now, addresses, ecn, packet, remaining, crypto, buf) + self.handle_first_packet(addresses, ecn, packet, remaining, crypto, buf) } Err(e) => { trace!("unable to decode initial packet: {}", e); @@ -412,7 +412,6 @@ impl Endpoint { fn handle_first_packet( &mut self, - now: Instant, addresses: FourTuple, ecn: Option, mut packet: Packet, @@ -478,7 +477,7 @@ impl Endpoint { } }; - let incoming = Incoming { + Some(DatagramEvent::NewConnection(Incoming { addresses, ecn, packet, @@ -490,24 +489,30 @@ impl Endpoint { version, retry_src_cid, orig_dst_cid, - }; - if server_config.use_retry && !incoming.remote_address_validated() { - Some(DatagramEvent::Response(self.retry(incoming, buf))) - } else { - match self.accept(incoming, now, buf) { - Ok((ch, conn)) => Some(DatagramEvent::NewConnection(ch, conn)), - Err((_, response)) => response.map(DatagramEvent::Response), - } - } + })) } /// Attempt to accept this incoming connection (an error may still occur) - fn accept( + pub fn accept( &mut self, incoming: Incoming, now: Instant, buf: &mut BytesMut, ) -> Result<(ConnectionHandle, Connection), (ConnectionError, Option)> { + self.check_connection_limit().map_err(|reason| { + ( + ConnectionError::ConnectionLimitExceeded, + Some(self.initial_close( + incoming.version, + incoming.addresses, + &incoming.crypto, + &incoming.src_cid, + reason, + buf, + )), + ) + })?; + let server_config = self.server_config.as_ref().unwrap().clone(); let ch = ConnectionHandle(self.connections.vacant_key()); @@ -602,8 +607,29 @@ impl Endpoint { Ok(()) } + /// Reject this incoming connection attempt + pub fn reject(&mut self, incoming: Incoming, buf: &mut BytesMut) -> Transmit { + self.initial_close( + incoming.version, + incoming.addresses, + &incoming.crypto, + &incoming.src_cid, + TransportError::CONNECTION_REFUSED(""), + buf, + ) + } + /// Respond with a retry packet, requiring the client to retry with address validation - fn retry(&mut self, incoming: Incoming, buf: &mut BytesMut) -> Transmit { + /// + /// Errors if `incoming.remote_address_validated()` is true. + pub fn retry( + &mut self, + incoming: Incoming, + buf: &mut BytesMut, + ) -> Result { + if incoming.remote_address_validated() { + return Err(RetryError(incoming)); + } let server_config = self.server_config.as_ref().unwrap(); // First Initial @@ -642,13 +668,13 @@ impl Endpoint { )); encode.finish(buf, &*incoming.crypto.header.local, None); - Transmit { + Ok(Transmit { destination: incoming.addresses.remote, ecn: None, size: buf.len(), segment_size: None, src_ip: incoming.addresses.local_ip, - } + }) } fn add_connection( @@ -940,14 +966,14 @@ impl IndexMut for Slab { pub enum DatagramEvent { /// The datagram is redirected to its `Connection` ConnectionEvent(ConnectionHandle, ConnectionEvent), - /// The datagram has resulted in starting a new `Connection` - NewConnection(ConnectionHandle, Connection), + /// The datagram may result in starting a new `Connection` + NewConnection(Incoming), /// Response generated directly by the endpoint Response(Transmit), } /// An incoming connection for which the server has not yet begun its part of the handshake. -struct Incoming { +pub struct Incoming { addresses: FourTuple, ecn: Option, packet: Packet, @@ -962,11 +988,24 @@ struct Incoming { } impl Incoming { + /// The local IP address which was used when the peer established + /// the connection + /// + /// This has the same behavior as [`Connection::local_ip`] + pub fn local_ip(&self) -> Option { + self.addresses.local_ip + } + + /// The peer's UDP address. + pub fn remote_address(&self) -> SocketAddr { + self.addresses.remote + } + /// Whether the socket address that is initiating this connection has been validated. /// /// This means that the sender of the initial packet has proved that they can receive traffic /// sent to `self.remote_address()`. - fn remote_address_validated(&self) -> bool { + pub fn remote_address_validated(&self) -> bool { self.retry_src_cid.is_some() } } @@ -1021,6 +1060,19 @@ pub enum ConnectError { UnsupportedVersion, } +/// Error for attempting to retry an [`Incoming`] which already bears an address +/// validation token from a previous retry +#[derive(Debug, Error)] +#[error("retry() with validated Incoming")] +pub struct RetryError(Incoming); + +impl RetryError { + /// Get the [`Incoming`] + pub fn into_incoming(self) -> Incoming { + self.0 + } +} + /// Reset Tokens which are associated with peer socket addresses /// /// The standard `HashMap` is used since both `SocketAddr` and `ResetToken` are diff --git a/quinn-proto/src/lib.rs b/quinn-proto/src/lib.rs index d271390f0..582f22b00 100644 --- a/quinn-proto/src/lib.rs +++ b/quinn-proto/src/lib.rs @@ -61,7 +61,9 @@ use crate::frame::Frame; pub use crate::frame::{ApplicationClose, ConnectionClose, Datagram}; mod endpoint; -pub use crate::endpoint::{ConnectError, ConnectionHandle, DatagramEvent, Endpoint}; +pub use crate::endpoint::{ + ConnectError, ConnectionHandle, DatagramEvent, Endpoint, Incoming, RetryError, +}; mod shared; pub use crate::shared::{ConnectionEvent, ConnectionId, EcnCodepoint, EndpointEvent}; diff --git a/quinn-proto/src/tests/mod.rs b/quinn-proto/src/tests/mod.rs index 52f17aece..712d49a4b 100644 --- a/quinn-proto/src/tests/mod.rs +++ b/quinn-proto/src/tests/mod.rs @@ -165,13 +165,8 @@ fn draft_version_compat() { #[test] fn stateless_retry() { let _guard = subscribe(); - let mut pair = Pair::new( - Default::default(), - ServerConfig { - use_retry: true, - ..server_config() - }, - ); + let mut pair = Pair::default(); + pair.server.incoming_connection_behavior = IncomingConnectionBehavior::Validate; pair.connect(); } @@ -459,13 +454,8 @@ fn high_latency_handshake() { #[test] fn zero_rtt_happypath() { let _guard = subscribe(); - let mut pair = Pair::new( - Default::default(), - ServerConfig { - use_retry: true, - ..server_config() - }, - ); + let mut pair = Pair::default(); + pair.server.incoming_connection_behavior = IncomingConnectionBehavior::Validate; let config = client_config(); // Establish normal connection @@ -2017,7 +2007,7 @@ fn connect_too_low_mtu() { pair.begin_connect(client_config()); pair.drive(); - pair.server.assert_no_accept() + pair.server.assert_no_accept(); } #[test] @@ -2811,3 +2801,23 @@ fn pure_sender_voluntarily_acks() { let receiver_acks_final = pair.server_conn_mut(server_ch).stats().frame_rx.acks; assert!(receiver_acks_final > receiver_acks_initial); } + +#[test] +fn reject_manually() { + let _guard = subscribe(); + let mut pair = Pair::default(); + pair.server.incoming_connection_behavior = IncomingConnectionBehavior::RejectAll; + + // The server should now reject incoming connections. + let client_ch = pair.begin_connect(client_config()); + pair.drive(); + pair.server.assert_no_accept(); + let client = pair.client.connections.get_mut(&client_ch).unwrap(); + assert!(client.is_closed()); + assert!(matches!( + client.poll(), + Some(Event::ConnectionLost { + reason: ConnectionError::ConnectionClosed(close) + }) if close.error_code == TransportErrorCode::CONNECTION_REFUSED + )); +} diff --git a/quinn-proto/src/tests/util.rs b/quinn-proto/src/tests/util.rs index e3b925306..f30897e82 100644 --- a/quinn-proto/src/tests/util.rs +++ b/quinn-proto/src/tests/util.rs @@ -287,11 +287,19 @@ pub(super) struct TestEndpoint { pub(super) outbound: VecDeque<(Transmit, Bytes)>, delayed: VecDeque<(Transmit, Bytes)>, pub(super) inbound: VecDeque<(Instant, Option, BytesMut)>, - accepted: Option, + accepted: Option>, pub(super) connections: HashMap, conn_events: HashMap>, pub(super) captured_packets: Vec>, pub(super) capture_inbound_packets: bool, + pub(super) incoming_connection_behavior: IncomingConnectionBehavior, +} + +#[derive(Debug, Copy, Clone)] +pub(super) enum IncomingConnectionBehavior { + AcceptAll, + RejectAll, + Validate, } impl TestEndpoint { @@ -318,6 +326,7 @@ impl TestEndpoint { conn_events: HashMap::default(), captured_packets: Vec::new(), capture_inbound_packets: false, + incoming_connection_behavior: IncomingConnectionBehavior::AcceptAll, } } @@ -345,9 +354,22 @@ impl TestEndpoint { .handle(recv_time, remote, None, ecn, packet, &mut buf) { match event { - DatagramEvent::NewConnection(ch, conn) => { - self.connections.insert(ch, conn); - self.accepted = Some(ch); + DatagramEvent::NewConnection(incoming) => { + match self.incoming_connection_behavior { + IncomingConnectionBehavior::AcceptAll => { + let _ = self.try_accept(incoming, now); + } + IncomingConnectionBehavior::RejectAll => { + self.reject(incoming); + } + IncomingConnectionBehavior::Validate => { + if incoming.remote_address_validated() { + let _ = self.try_accept(incoming, now); + } else { + self.retry(incoming); + } + } + } } DatagramEvent::ConnectionEvent(ch, event) => { if self.capture_inbound_packets { @@ -428,8 +450,58 @@ impl TestEndpoint { self.outbound.extend(self.delayed.drain(..)); } + pub(super) fn try_accept( + &mut self, + incoming: Incoming, + now: Instant, + ) -> Result { + let mut buf = BytesMut::new(); + self.endpoint + .accept(incoming, now, &mut buf) + .map(|(ch, conn)| { + self.connections.insert(ch, conn); + self.accepted = Some(Ok(ch)); + ch + }) + .map_err(|(e, transmit)| { + if let Some(transmit) = transmit { + let size = transmit.size; + self.outbound + .extend(split_transmit(transmit, buf.split_to(size).freeze())); + } + self.accepted = Some(Err(e.clone())); + e + }) + } + + pub(super) fn retry(&mut self, incoming: Incoming) { + let mut buf = BytesMut::new(); + let transmit = self.endpoint.retry(incoming, &mut buf).unwrap(); + let size = transmit.size; + self.outbound + .extend(split_transmit(transmit, buf.split_to(size).freeze())); + } + + pub(super) fn reject(&mut self, incoming: Incoming) { + let mut buf = BytesMut::new(); + let transmit = self.endpoint.reject(incoming, &mut buf); + let size = transmit.size; + self.outbound + .extend(split_transmit(transmit, buf.split_to(size).freeze())); + } + pub(super) fn assert_accept(&mut self) -> ConnectionHandle { - self.accepted.take().expect("server didn't connect") + self.accepted + .take() + .expect("server didn't try connecting") + .expect("server experienced error connecting") + } + + pub(super) fn assert_accept_error(&mut self) -> ConnectionError { + self.accepted + .take() + .expect("server didn't try connecting") + .expect_err("server did unexpectedly connect without error") } pub(super) fn assert_no_accept(&self) { diff --git a/quinn/examples/server.rs b/quinn/examples/server.rs index 70e71be77..db887d99a 100644 --- a/quinn/examples/server.rs +++ b/quinn/examples/server.rs @@ -132,9 +132,6 @@ async fn run(options: Opt) -> Result<()> { let mut server_config = quinn::ServerConfig::with_crypto(Arc::new(server_crypto)); let transport_config = Arc::get_mut(&mut server_config.transport).unwrap(); transport_config.max_concurrent_uni_streams(0_u8.into()); - if options.stateless_retry { - server_config.use_retry(true); - } let root = Arc::::from(options.root.clone()); if !root.exists() { @@ -145,19 +142,24 @@ async fn run(options: Opt) -> Result<()> { eprintln!("listening on {}", endpoint.local_addr()?); while let Some(conn) = endpoint.accept().await { - info!("connection incoming"); - let fut = handle_connection(root.clone(), conn); - tokio::spawn(async move { - if let Err(e) = fut.await { - error!("connection failed: {reason}", reason = e.to_string()) - } - }); + if options.stateless_retry && !conn.remote_address_validated() { + info!("requiring connection to validate its address"); + conn.retry().unwrap(); + } else { + info!("accepting connection"); + let fut = handle_connection(root.clone(), conn); + tokio::spawn(async move { + if let Err(e) = fut.await { + error!("connection failed: {reason}", reason = e.to_string()) + } + }); + } } Ok(()) } -async fn handle_connection(root: Arc, conn: quinn::Connecting) -> Result<()> { +async fn handle_connection(root: Arc, conn: quinn::Incoming) -> Result<()> { let connection = conn.await?; let span = info_span!( "connection", diff --git a/quinn/src/endpoint.rs b/quinn/src/endpoint.rs index d24e8517b..d5d370db0 100644 --- a/quinn/src/endpoint.rs +++ b/quinn/src/endpoint.rs @@ -16,7 +16,8 @@ use crate::runtime::{default_runtime, AsyncUdpSocket, Runtime}; use bytes::{Bytes, BytesMut}; use pin_project_lite::pin_project; use proto::{ - self as proto, ClientConfig, ConnectError, ConnectionHandle, DatagramEvent, ServerConfig, + self as proto, ClientConfig, ConnectError, ConnectionError, ConnectionHandle, DatagramEvent, + ServerConfig, }; use rustc_hash::FxHashMap; use tokio::sync::{futures::Notified, mpsc, Notify}; @@ -24,9 +25,9 @@ use tracing::{Instrument, Span}; use udp::{RecvMeta, BATCH_SIZE}; use crate::{ - connection::Connecting, work_limiter::WorkLimiter, ConnectionEvent, EndpointConfig, - EndpointEvent, VarInt, IO_LOOP_BOUND, MAX_TRANSMIT_QUEUE_CONTENTS_LEN, RECV_TIME_BOUND, - SEND_TIME_BOUND, + connection::Connecting, incoming::Incoming, work_limiter::WorkLimiter, ConnectionEvent, + EndpointConfig, EndpointEvent, VarInt, IO_LOOP_BOUND, MAX_INCOMING_CONNECTIONS, + MAX_TRANSMIT_QUEUE_CONTENTS_LEN, RECV_TIME_BOUND, SEND_TIME_BOUND, }; /// A QUIC endpoint. @@ -137,8 +138,10 @@ impl Endpoint { /// Get the next incoming connection attempt from a client /// - /// Yields [`Connecting`] futures that must be `await`ed to obtain the final `Connection`, or - /// `None` if the endpoint is [`close`](Self::close)d. + /// Yields [`Incoming`]s, or `None` if the endpoint is [`close`](Self::close)d. [`Incoming`] + /// can be `await`ed to obtain the final [`Connection`](crate::Connection), or used to e.g. + /// filter connection attempts or force address validation, or converted into an intermediate + /// `Connecting` future which can be used to e.g. send 0.5-RTT data. pub fn accept(&self) -> Accept<'_> { Accept { endpoint: self, @@ -366,12 +369,57 @@ pub(crate) struct EndpointInner { pub(crate) shared: Shared, } +impl EndpointInner { + pub(crate) fn accept( + &self, + incoming: proto::Incoming, + mut response_buffer: BytesMut, + ) -> Result { + let mut state = self.state.lock().unwrap(); + state + .inner + .accept(incoming, Instant::now(), &mut response_buffer) + .map(|(handle, conn)| { + let socket = state.socket.clone(); + let runtime = state.runtime.clone(); + state.connections.insert(handle, conn, socket, runtime) + }) + .map_err(|(e, response)| { + if let Some(transmit) = response { + state.transmit_state.respond(transmit, response_buffer); + } + e + }) + } + + pub(crate) fn reject(&self, incoming: proto::Incoming, mut response_buffer: BytesMut) { + let mut state = self.state.lock().unwrap(); + let transmit = state.inner.reject(incoming, &mut response_buffer); + state.transmit_state.respond(transmit, response_buffer); + } + + pub(crate) fn retry( + &self, + incoming: proto::Incoming, + mut response_buffer: BytesMut, + ) -> Result<(), (proto::RetryError, BytesMut)> { + let mut state = self.state.lock().unwrap(); + match state.inner.retry(incoming, &mut response_buffer) { + Ok(transmit) => { + state.transmit_state.respond(transmit, response_buffer); + Ok(()) + } + Err(e) => Err((e, response_buffer)), + } + } +} + #[derive(Debug)] pub(crate) struct State { socket: Arc, inner: proto::Endpoint, transmit_state: TransmitState, - incoming: VecDeque, + incoming: VecDeque<(proto::Incoming, BytesMut)>, driver: Option, ipv6: bool, connections: ConnectionSet, @@ -423,14 +471,14 @@ impl State { buf, &mut response_buffer, ) { - Some(DatagramEvent::NewConnection(handle, conn)) => { - let conn = self.connections.insert( - handle, - conn, - self.socket.clone(), - self.runtime.clone(), - ); - self.incoming.push_back(conn); + Some(DatagramEvent::NewConnection(incoming)) => { + if self.incoming.len() < MAX_INCOMING_CONNECTIONS { + self.incoming.push_back((incoming, response_buffer)); + } else { + let transmit = + self.inner.reject(incoming, &mut response_buffer); + self.transmit_state.respond(transmit, response_buffer); + } } Some(DatagramEvent::ConnectionEvent(handle, event)) => { // Ignoring errors from dropped connections that haven't yet been cleaned up @@ -661,15 +709,18 @@ pin_project! { } impl<'a> Future for Accept<'a> { - type Output = Option; + type Output = Option; fn poll(self: Pin<&mut Self>, ctx: &mut Context<'_>) -> Poll { let mut this = self.project(); - let endpoint = &mut *this.endpoint.inner.state.lock().unwrap(); + let mut endpoint = this.endpoint.inner.state.lock().unwrap(); if endpoint.driver_lost { return Poll::Ready(None); } - if let Some(conn) = endpoint.incoming.pop_front() { - return Poll::Ready(Some(conn)); + if let Some((incoming, response_buffer)) = endpoint.incoming.pop_front() { + // Release the mutex lock on endpoint so cloning it doesn't deadlock + drop(endpoint); + let incoming = Incoming::new(incoming, this.endpoint.inner.clone(), response_buffer); + return Poll::Ready(Some(incoming)); } if endpoint.connections.close.is_some() { return Poll::Ready(None); diff --git a/quinn/src/incoming.rs b/quinn/src/incoming.rs new file mode 100644 index 000000000..b4ca680d8 --- /dev/null +++ b/quinn/src/incoming.rs @@ -0,0 +1,149 @@ +use std::{ + fmt, + future::{Future, IntoFuture}, + net::{IpAddr, SocketAddr}, + pin::Pin, + task::{Context, Poll}, +}; + +use bytes::BytesMut; +use proto::ConnectionError; +use thiserror::Error; + +use crate::{ + connection::{Connecting, Connection}, + endpoint::EndpointRef, +}; + +/// An incoming connection for which the server has not yet begun its part of the handshake +pub struct Incoming(Option); + +impl Incoming { + pub(crate) fn new( + inner: proto::Incoming, + endpoint: EndpointRef, + response_buffer: BytesMut, + ) -> Self { + Self(Some(State { + inner, + endpoint, + response_buffer, + })) + } + + /// Attempt to accept this incoming connection (an error may still occur) + pub fn accept(mut self) -> Result { + let state = self.0.take().unwrap(); + state.endpoint.accept(state.inner, state.response_buffer) + } + + /// Reject this incoming connection attempt + pub fn reject(mut self) { + let state = self.0.take().unwrap(); + state.endpoint.reject(state.inner, state.response_buffer); + } + + /// Respond with a retry packet, requiring the client to retry with address validation + /// + /// Errors if `remote_address_validated()` is true. + pub fn retry(mut self) -> Result<(), RetryError> { + let state = self.0.take().unwrap(); + state + .endpoint + .retry(state.inner, state.response_buffer) + .map_err(|(e, response_buffer)| { + RetryError(Self(Some(State { + inner: e.into_incoming(), + endpoint: state.endpoint, + response_buffer, + }))) + }) + } + + /// Ignore this incoming connection attempt, not sending any packet in response + pub fn ignore(mut self) { + self.0.take().unwrap(); + } + + /// The local IP address which was used when the peer established + /// the connection + pub fn local_ip(&self) -> Option { + self.0.as_ref().unwrap().inner.local_ip() + } + + /// The peer's UDP address + pub fn remote_address(&self) -> SocketAddr { + self.0.as_ref().unwrap().inner.remote_address() + } + + /// Whether the socket address that is initiating this connection has been validated + /// + /// This means that the sender of the initial packet has proved that they can receive traffic + /// sent to `self.remote_address()`. + pub fn remote_address_validated(&self) -> bool { + self.0.as_ref().unwrap().inner.remote_address_validated() + } +} + +impl Drop for Incoming { + fn drop(&mut self) { + // Implicit reject, similar to Connection's implicit close + if let Some(state) = self.0.take() { + state.endpoint.reject(state.inner, state.response_buffer); + } + } +} + +impl fmt::Debug for Incoming { + fn fmt(&self, f: &mut fmt::Formatter) -> fmt::Result { + let state = self.0.as_ref().unwrap(); + f.debug_struct("Incoming") + .field("inner", &state.inner) + .field("endpoint", &state.endpoint) + // response_buffer is too big and not meaningful enough + .finish_non_exhaustive() + } +} + +struct State { + inner: proto::Incoming, + endpoint: EndpointRef, + response_buffer: BytesMut, +} + +/// Error for attempting to retry an [`Incoming`] which already bears an address +/// validation token from a previous retry +#[derive(Debug, Error)] +#[error("retry() with validated Incoming")] +pub struct RetryError(Incoming); + +impl RetryError { + /// Get the [`Incoming`] + pub fn into_incoming(self) -> Incoming { + self.0 + } +} + +/// Basic adapter to let [`Incoming`] be `await`-ed like a [`Connecting`] +#[derive(Debug)] +pub struct IncomingFuture(Result); + +impl Future for IncomingFuture { + type Output = Result; + + fn poll(mut self: Pin<&mut Self>, cx: &mut Context) -> Poll { + match &mut self.0 { + Ok(ref mut connecting) => Pin::new(connecting).poll(cx), + Err(e) => Poll::Ready(Err(e.clone())), + } + } +} + +impl IntoFuture for Incoming { + type Output = Result; + type IntoFuture = IncomingFuture; + + fn into_future(self) -> Self::IntoFuture { + IncomingFuture(self.accept()) + } +} diff --git a/quinn/src/lib.rs b/quinn/src/lib.rs index f1d8db5cb..cb34a7562 100644 --- a/quinn/src/lib.rs +++ b/quinn/src/lib.rs @@ -54,6 +54,7 @@ macro_rules! ready { mod connection; mod endpoint; +mod incoming; mod mutex; mod recv_stream; mod runtime; @@ -75,6 +76,7 @@ pub use crate::connection::{ UnknownStream, ZeroRttAccepted, }; pub use crate::endpoint::{Accept, Endpoint}; +pub use crate::incoming::{Incoming, IncomingFuture, RetryError}; pub use crate::recv_stream::{ReadError, ReadExactError, ReadToEndError, RecvStream}; #[cfg(feature = "runtime-async-std")] pub use crate::runtime::AsyncStdRuntime; @@ -125,3 +127,10 @@ const SEND_TIME_BOUND: Duration = Duration::from_micros(50); /// generated from the endpoint (retry or initial close) can be dropped when this limit is being execeeded. /// Chose to represent 100 MB of data. const MAX_TRANSMIT_QUEUE_CONTENTS_LEN: usize = 100_000_000; + +/// The maximum number of `IncomingConnection`s we allow to be enqueued at a time before we start +/// rejecting new `IncomingConnection`s automatically. Assuming each `IncomingConnection` accounts +/// for little over 1200 bytes of memory maximum, this should limit an endpoint's incoming +/// connection queue memory consumption to under 100 MiB, a generous amount that still prevents +/// memory exhaustion. +const MAX_INCOMING_CONNECTIONS: usize = 1 << 16; diff --git a/quinn/src/tests.rs b/quinn/src/tests.rs index 05c0ebf5e..97a4e0058 100755 --- a/quinn/src/tests.rs +++ b/quinn/src/tests.rs @@ -159,17 +159,29 @@ fn export_keying_material() { }; runtime.block_on(async move { - let outgoing_conn = endpoint - .connect(endpoint.local_addr().unwrap(), "localhost") - .unwrap() - .await - .expect("connect"); - let incoming_conn = endpoint - .accept() - .await - .expect("endpoint") - .await - .expect("connection"); + let outgoing_conn_fut = tokio::spawn({ + let endpoint = endpoint.clone(); + async move { + endpoint + .connect(endpoint.local_addr().unwrap(), "localhost") + .unwrap() + .await + .expect("connect") + } + }); + let incoming_conn_fut = tokio::spawn({ + let endpoint = endpoint.clone(); + async move { + endpoint + .accept() + .await + .expect("endpoint") + .await + .expect("connection") + } + }); + let outgoing_conn = outgoing_conn_fut.await.unwrap(); + let incoming_conn = incoming_conn_fut.await.unwrap(); let mut i_buf = [0u8; 64]; incoming_conn .export_keying_material(&mut i_buf, b"asdf", b"qwer") @@ -183,70 +195,92 @@ fn export_keying_material() { } #[tokio::test] -async fn accept_after_close() { +async fn ip_blocking() { let _guard = subscribe(); - let endpoint = endpoint(); - - const MSG: &[u8] = b"goodbye!"; - - let sender = endpoint - .connect(endpoint.local_addr().unwrap(), "localhost") - .unwrap() - .await - .expect("connect"); - let mut s = sender.open_uni().await.unwrap(); - s.write_all(MSG).await.unwrap(); - s.finish().await.unwrap(); - sender.close(0u32.into(), b""); - - // Allow some time for the close to be sent and processed - tokio::time::sleep(Duration::from_millis(100)).await; - - // Despite the connection having closed, we should be able to accept it... - let receiver = endpoint - .accept() - .await - .expect("endpoint") - .await - .expect("connection"); - - // ...and read what was sent. - let mut stream = receiver.accept_uni().await.expect("incoming streams"); - let msg = stream - .read_to_end(usize::max_value()) - .await - .expect("read_to_end"); - assert_eq!(msg, MSG); - - // But it's still definitely closed. - assert!(receiver.open_uni().await.is_err()); + let endpoint_factory = EndpointFactory::new(); + let client_1 = endpoint_factory.endpoint(); + let client_1_addr = client_1.local_addr().unwrap(); + let client_2 = endpoint_factory.endpoint(); + let server = endpoint_factory.endpoint(); + let server_addr = server.local_addr().unwrap(); + let server_task = tokio::spawn(async move { + loop { + let accepting = server.accept().await.unwrap(); + if accepting.remote_address() == client_1_addr { + accepting.reject(); + } else if accepting.remote_address_validated() { + accepting.await.expect("connection"); + } else { + accepting.retry().unwrap(); + } + } + }); + tokio::join!( + async move { + let e = client_1 + .connect(server_addr, "localhost") + .unwrap() + .await + .expect_err("server should have blocked this"); + assert!( + matches!(e, crate::ConnectionError::ConnectionClosed(_)), + "wrong error" + ); + }, + async move { + client_2 + .connect(server_addr, "localhost") + .unwrap() + .await + .expect("connect"); + } + ); + server_task.abort(); } /// Construct an endpoint suitable for connecting to itself fn endpoint() -> Endpoint { - endpoint_with_config(TransportConfig::default()) + EndpointFactory::new().endpoint() } fn endpoint_with_config(transport_config: TransportConfig) -> Endpoint { - let cert = rcgen::generate_simple_self_signed(vec!["localhost".into()]).unwrap(); - let key = rustls::PrivateKey(cert.serialize_private_key_der()); - let cert = rustls::Certificate(cert.serialize_der().unwrap()); - let transport_config = Arc::new(transport_config); - let mut server_config = crate::ServerConfig::with_single_cert(vec![cert.clone()], key).unwrap(); - server_config.transport_config(transport_config.clone()); + EndpointFactory::new().endpoint_with_config(transport_config) +} - let mut roots = rustls::RootCertStore::empty(); - roots.add(&cert).unwrap(); - let mut endpoint = Endpoint::server( - server_config, - SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 0), - ) - .unwrap(); - let mut client_config = ClientConfig::with_root_certificates(roots); - client_config.transport_config(transport_config); - endpoint.set_default_client_config(client_config); +/// Constructs endpoints suitable for connecting to themselves and each other +struct EndpointFactory(rcgen::Certificate); - endpoint +impl EndpointFactory { + fn new() -> Self { + Self(rcgen::generate_simple_self_signed(vec!["localhost".into()]).unwrap()) + } + + fn endpoint(&self) -> Endpoint { + self.endpoint_with_config(TransportConfig::default()) + } + + fn endpoint_with_config(&self, transport_config: TransportConfig) -> Endpoint { + let cert = &self.0; + let key = rustls::PrivateKey(cert.serialize_private_key_der()); + let cert = rustls::Certificate(cert.serialize_der().unwrap()); + let transport_config = Arc::new(transport_config); + let mut server_config = + crate::ServerConfig::with_single_cert(vec![cert.clone()], key).unwrap(); + server_config.transport_config(transport_config.clone()); + + let mut roots = rustls::RootCertStore::empty(); + roots.add(&cert).unwrap(); + let mut endpoint = Endpoint::server( + server_config, + SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 0), + ) + .unwrap(); + let mut client_config = ClientConfig::with_root_certificates(roots); + client_config.transport_config(transport_config); + endpoint.set_default_client_config(client_config); + + endpoint + } } #[tokio::test] @@ -259,7 +293,7 @@ async fn zero_rtt() { let endpoint2 = endpoint.clone(); tokio::spawn(async move { for _ in 0..2 { - let incoming = endpoint2.accept().await.unwrap(); + let incoming = endpoint2.accept().await.unwrap().accept().unwrap(); let (connection, established) = incoming.into_0rtt().unwrap_or_else(|_| unreachable!()); let c = connection.clone(); tokio::spawn(async move {