Allow accept/reject/retry before handshake begins

This commit removes use_retry from the server config and provides a
public API for the user to manually accept/reject/retry incoming
connections before a handshake begins, and inspect properties such as
an incoming connection's remote address and whether that address is
validated when doing so.

In quinn-proto, Incoming is made public, as well as Endpoint's accept/
reject/retry methods which operate on it. The
DatagramEvent::NewConnection event is modified to return an incoming
but not yet accepted connection.

In quinn, awaiting Endpoint::accept now yields a new
quinn::Incoming type, rather than quinn::Connecting. The new
quinn::Incoming type has all the methods its quinn_proto equivalent has,
as well as an accept method to (fallibly) transition it into a
Connecting, and also reject, retry, and ignore methods.

Furthermore, quinn::Incoming implements IntoFuture with the output type
Result<Connection, ConnectionError>>, which is the same as the Future
output type of Connecting. This lets server code which was
straightforwardly awaiting the result of quinn::Endpoint::accept work
with little to no modification.

The test accept_after_close was removed because the functionality it
was testing for no longer exists.
This commit is contained in:
Phoenix Kahlo
2024-02-15 19:25:32 -06:00
committed by Dirkjan Ochtman
parent 0af674127f
commit 736f87bcdc
11 changed files with 519 additions and 152 deletions
+1 -1
View File
@@ -123,7 +123,7 @@ async fn run(opt: Opt) -> Result<()> {
Ok(())
}
async fn handle(handshake: quinn::Connecting, opt: Arc<Opt>) -> Result<()> {
async fn handle(handshake: quinn::Incoming, opt: Arc<Opt>) -> Result<()> {
let connection = handshake.await.context("handshake failed")?;
debug!("{} connected", connection.remote_address());
tokio::try_join!(
-14
View File
@@ -741,10 +741,6 @@ pub struct ServerConfig {
/// Used to generate one-time AEAD keys to protect handshake tokens
pub(crate) token_key: Arc<dyn HandshakeTokenKey>,
/// 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)
+72 -20
View File
@@ -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<EcnCodepoint>,
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<Transmit>)> {
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<Transmit, RetryError> {
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<ConnectionHandle> for Slab<ConnectionMeta> {
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<EcnCodepoint>,
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<IpAddr> {
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
+3 -1
View File
@@ -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};
+25 -15
View File
@@ -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
));
}
+77 -5
View File
@@ -287,11 +287,19 @@ pub(super) struct TestEndpoint {
pub(super) outbound: VecDeque<(Transmit, Bytes)>,
delayed: VecDeque<(Transmit, Bytes)>,
pub(super) inbound: VecDeque<(Instant, Option<EcnCodepoint>, BytesMut)>,
accepted: Option<ConnectionHandle>,
accepted: Option<Result<ConnectionHandle, ConnectionError>>,
pub(super) connections: HashMap<ConnectionHandle, Connection>,
conn_events: HashMap<ConnectionHandle, VecDeque<ConnectionEvent>>,
pub(super) captured_packets: Vec<Vec<u8>>,
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<ConnectionHandle, ConnectionError> {
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) {
+13 -11
View File
@@ -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::<Path>::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<Path>, conn: quinn::Connecting) -> Result<()> {
async fn handle_connection(root: Arc<Path>, conn: quinn::Incoming) -> Result<()> {
let connection = conn.await?;
let span = info_span!(
"connection",
+70 -19
View File
@@ -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<Connecting, ConnectionError> {
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<dyn AsyncUdpSocket>,
inner: proto::Endpoint,
transmit_state: TransmitState,
incoming: VecDeque<Connecting>,
incoming: VecDeque<(proto::Incoming, BytesMut)>,
driver: Option<Waker>,
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<Connecting>;
type Output = Option<Incoming>;
fn poll(self: Pin<&mut Self>, ctx: &mut Context<'_>) -> Poll<Self::Output> {
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);
+149
View File
@@ -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<State>);
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<Connecting, ConnectionError> {
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<IpAddr> {
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<Connecting, ConnectionError>);
impl Future for IncomingFuture {
type Output = Result<Connection, ConnectionError>;
fn poll(mut self: Pin<&mut Self>, cx: &mut Context) -> Poll<Self::Output> {
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<Connection, ConnectionError>;
type IntoFuture = IncomingFuture;
fn into_future(self) -> Self::IntoFuture {
IncomingFuture(self.accept())
}
}
+9
View File
@@ -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;
+100 -66
View File
@@ -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 {