mirror of
https://github.com/n0-computer/noq.git
synced 2026-10-03 00:29:07 +00:00
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:
committed by
Dirkjan Ochtman
parent
0af674127f
commit
736f87bcdc
@@ -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!(
|
||||
|
||||
@@ -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
@@ -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
|
||||
|
||||
@@ -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};
|
||||
|
||||
@@ -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
|
||||
));
|
||||
}
|
||||
|
||||
@@ -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
@@ -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
@@ -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);
|
||||
|
||||
@@ -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())
|
||||
}
|
||||
}
|
||||
@@ -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
@@ -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 {
|
||||
|
||||
Reference in New Issue
Block a user