use std::{
collections::HashMap,
io,
net::{IpAddr, SocketAddr},
sync::{
atomic::{AtomicBool, Ordering},
Arc,
},
time::Duration,
};
use crate::connection::Connection;
use nxtquic_crypto::{
CryptoProvider, DirectionalKeys, HandshakeState, Keys,
RustlsCryptoProvider,
};
use nxtquic_proto::{
frame::{AckFrame, CryptoFrame, Frame, StreamFrame},
packet::{
decode_header,
initial::parse_initial,
protection::{apply_header_protection, remove_header_protection},
PacketType,
},
varint::VarInt,
ConnectionId, StreamId, Version,
};
use tokio::sync::{Mutex, Notify};
#[derive(Default, Clone, Debug)]
pub struct EndpointConfig;
#[derive(Clone)]
pub struct ServerConfig {
rustls: Arc<rustls::ServerConfig>,
transport: TransportConfig,
max_concurrent_connections: Option<usize>,
migration_enabled: bool,
preferred_address: Option<SocketAddr>,
}
impl std::fmt::Debug for ServerConfig {
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
formatter
.debug_struct("ServerConfig")
.field("transport", &self.transport)
.field(
"max_concurrent_connections",
&self.max_concurrent_connections,
)
.field("migration_enabled", &self.migration_enabled)
.field("preferred_address", &self.preferred_address)
.finish_non_exhaustive()
}
}
impl ServerConfig {
pub fn builder() -> ServerConfigBuilder {
ServerConfigBuilder::default()
}
pub fn rustls(&self) -> &Arc<rustls::ServerConfig> {
&self.rustls
}
pub fn transport(&self) -> &TransportConfig {
&self.transport
}
pub fn max_concurrent_connections(&self) -> Option<usize> {
self.max_concurrent_connections
}
pub fn migration_enabled(&self) -> bool {
self.migration_enabled
}
pub fn preferred_address(&self) -> Option<SocketAddr> {
self.preferred_address
}
}
#[derive(Default)]
pub struct ServerConfigBuilder {
rustls: Option<Arc<rustls::ServerConfig>>,
transport: TransportConfig,
max_concurrent_connections: Option<usize>,
migration_enabled: bool,
preferred_address: Option<SocketAddr>,
}
impl ServerConfigBuilder {
pub fn with_rustls(mut self, config: Arc<rustls::ServerConfig>) -> io::Result<Self> {
if config
.alpn_protocols
.iter()
.all(|protocol| protocol.as_slice() != b"h3")
{
return Err(io::Error::new(
io::ErrorKind::InvalidInput,
"a QUIC HTTP/3 server configuration must advertise ALPN h3",
));
}
self.rustls = Some(config);
Ok(self)
}
pub fn with_transport_config(mut self, configure: impl FnOnce(&mut TransportConfig)) -> Self {
configure(&mut self.transport);
self
}
pub fn with_max_concurrent_connections(mut self, max: usize) -> Self {
self.max_concurrent_connections = Some(max);
self
}
pub fn with_migration(mut self, enable: bool) -> Self {
self.migration_enabled = enable;
self
}
pub fn with_preferred_address(mut self, addr: SocketAddr) -> Self {
self.preferred_address = Some(addr);
self
}
pub fn build(self) -> io::Result<ServerConfig> {
let rustls = self.rustls.ok_or_else(|| {
io::Error::new(
io::ErrorKind::InvalidInput,
"missing Rustls server configuration",
)
})?;
Ok(ServerConfig {
rustls,
transport: self.transport,
max_concurrent_connections: self.max_concurrent_connections,
migration_enabled: self.migration_enabled,
preferred_address: self.preferred_address,
})
}
}
#[derive(Clone, Debug)]
pub struct TransportConfig {
max_idle_timeout: Option<Duration>,
keep_alive_interval: Option<Duration>,
datagram_receive_buffer_size: usize,
initial_max_data: u64,
initial_max_stream_data_bidi_local: u64,
initial_max_stream_data_bidi_remote: u64,
initial_max_stream_data_uni: u64,
initial_max_streams_bidi: u64,
initial_max_streams_uni: u64,
max_udp_payload_size: u64,
ack_delay_exponent: u8,
max_ack_delay: Duration,
disable_active_migration: bool,
active_connection_id_limit: u64,
max_concurrent_uni_streams: u64,
max_concurrent_bi_streams: u64,
}
impl Default for TransportConfig {
fn default() -> Self {
Self {
max_idle_timeout: Some(Duration::from_secs(30)),
keep_alive_interval: Some(Duration::from_secs(10)),
datagram_receive_buffer_size: 1024 * 1024,
initial_max_data: 10 * 1024 * 1024,
initial_max_stream_data_bidi_local: 1024 * 1024,
initial_max_stream_data_bidi_remote: 1024 * 1024,
initial_max_stream_data_uni: 1024 * 1024,
initial_max_streams_bidi: 100,
initial_max_streams_uni: 100,
max_udp_payload_size: 1350,
ack_delay_exponent: 3,
max_ack_delay: Duration::from_millis(25),
disable_active_migration: false,
active_connection_id_limit: 4,
max_concurrent_uni_streams: 100,
max_concurrent_bi_streams: 100,
}
}
}
impl TransportConfig {
pub fn max_idle_timeout(&mut self, timeout: Option<Duration>) -> &mut Self {
self.max_idle_timeout = timeout;
self
}
pub fn idle_timeout(&self) -> Option<Duration> {
self.max_idle_timeout
}
pub fn keep_alive_interval(&mut self, interval: Option<Duration>) -> &mut Self {
self.keep_alive_interval = interval;
self
}
pub fn keep_alive(&self) -> Option<Duration> {
self.keep_alive_interval
}
pub fn datagram_receive_buffer_size(&mut self, size: usize) -> &mut Self {
self.datagram_receive_buffer_size = size;
self
}
pub fn datagram_buffer_size(&self) -> usize {
self.datagram_receive_buffer_size
}
pub fn initial_max_data(&mut self, value: u64) -> &mut Self {
self.initial_max_data = value;
self
}
pub fn get_initial_max_data(&self) -> u64 {
self.initial_max_data
}
pub fn initial_max_stream_data_bidi_local(&mut self, value: u64) -> &mut Self {
self.initial_max_stream_data_bidi_local = value;
self
}
pub fn initial_max_stream_data_bidi_remote(&mut self, value: u64) -> &mut Self {
self.initial_max_stream_data_bidi_remote = value;
self
}
pub fn initial_max_stream_data_uni(&mut self, value: u64) -> &mut Self {
self.initial_max_stream_data_uni = value;
self
}
pub fn initial_max_streams_bidi(&mut self, value: u64) -> &mut Self {
self.initial_max_streams_bidi = value;
self
}
pub fn initial_max_streams_uni(&mut self, value: u64) -> &mut Self {
self.initial_max_streams_uni = value;
self
}
pub fn max_udp_payload_size(&mut self, value: u64) -> &mut Self {
self.max_udp_payload_size = value;
self
}
pub fn ack_delay_exponent(&mut self, value: u8) -> &mut Self {
self.ack_delay_exponent = value;
self
}
pub fn max_ack_delay(&mut self, value: Duration) -> &mut Self {
self.max_ack_delay = value;
self
}
pub fn disable_active_migration(&mut self, disable: bool) -> &mut Self {
self.disable_active_migration = disable;
self
}
pub fn active_connection_id_limit(&mut self, limit: u64) -> &mut Self {
self.active_connection_id_limit = limit;
self
}
pub fn max_concurrent_uni_streams(&mut self, count: u64) -> &mut Self {
self.max_concurrent_uni_streams = count;
self
}
pub fn max_concurrent_bi_streams(&mut self, count: u64) -> &mut Self {
self.max_concurrent_bi_streams = count;
self
}
}
#[derive(Clone)]
pub struct ClientConfig {
rustls: Option<Arc<rustls::ClientConfig>>,
transport: TransportConfig,
}
impl Default for ClientConfig {
fn default() -> Self {
Self {
rustls: None,
transport: TransportConfig::default(),
}
}
}
impl std::fmt::Debug for ClientConfig {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("ClientConfig")
.field("transport", &self.transport)
.finish_non_exhaustive()
}
}
impl ClientConfig {
pub fn builder() -> ClientConfigBuilder {
ClientConfigBuilder::default()
}
pub fn rustls(&self) -> Option<&Arc<rustls::ClientConfig>> {
self.rustls.as_ref()
}
pub fn transport(&self) -> &TransportConfig {
&self.transport
}
}
#[derive(Default)]
pub struct ClientConfigBuilder {
rustls: Option<Arc<rustls::ClientConfig>>,
transport: TransportConfig,
}
impl ClientConfigBuilder {
pub fn with_rustls(mut self, config: Arc<rustls::ClientConfig>) -> Self {
self.rustls = Some(config);
self
}
pub fn with_transport_config(mut self, configure: impl FnOnce(&mut TransportConfig)) -> Self {
configure(&mut self.transport);
self
}
pub fn with_native_roots(mut self) -> Self {
let root_store = rustls::RootCertStore::empty();
let mut config = rustls::ClientConfig::builder()
.with_root_certificates(root_store)
.with_no_client_auth();
config.alpn_protocols = vec![b"h3".to_vec()];
self.rustls = Some(Arc::new(config));
self
}
pub fn build(self) -> ClientConfig {
ClientConfig {
rustls: self.rustls,
transport: self.transport,
}
}
}
#[derive(Clone, Debug, Default)]
pub struct EndpointStats {
pub total_connections: u64,
pub active_connections: u64,
pub handshake_failures: u64,
pub packets_received: u64,
pub packets_sent: u64,
}
pub struct Endpoint {
_config: EndpointConfig,
server_config: Arc<Mutex<Option<ServerConfig>>>,
socket: Arc<tokio::net::UdpSocket>,
incoming: Arc<Mutex<tokio::sync::mpsc::Receiver<Incoming>>>,
closed: Arc<AtomicBool>,
stats: Arc<Mutex<EndpointStats>>,
close_notify: Arc<Notify>,
}
pub struct Connecting;
pub struct Incoming {
remote_addr: SocketAddr,
packet: Vec<u8>,
initial: DecryptedInitial,
socket: Arc<tokio::net::UdpSocket>,
server_config: Option<ServerConfig>,
datagrams: Option<tokio::sync::mpsc::Receiver<Vec<u8>>>,
}
#[derive(Debug)]
struct DecryptedInitial {
version: Version,
dcid: ConnectionId,
scid: ConnectionId,
_packet_number: u64,
frames: Vec<Frame>,
}
fn decrypt_client_initial(packet: &[u8]) -> io::Result<DecryptedInitial> {
let parsed =
parse_initial(packet).map_err(|error| io::Error::new(io::ErrorKind::InvalidData, error))?;
let version = match parsed.version {
1 => Version::V1,
0x6b33_43cf => Version::V2,
version => {
return Err(io::Error::new(
io::ErrorKind::InvalidData,
format!("unsupported QUIC version {version:#010x}"),
));
}
};
let provider = RustlsCryptoProvider::new(None, None);
let keys = provider.initial_keys(version, &parsed.dcid);
let mut protected = packet[..parsed.packet_end].to_vec();
let sample_offset = parsed.packet_number_offset + 4;
let sample: [u8; 16] = protected[sample_offset..sample_offset + 16]
.try_into()
.expect("Initial parser validates the header-protection sample");
let mask = keys
.client
.header_key
.protection_mask(&sample)
.map_err(|error| io::Error::new(io::ErrorKind::InvalidData, error))?;
let packet_number_len =
remove_header_protection(&mut protected, &mask, parsed.packet_number_offset);
if packet_number_len == 0
|| parsed.packet_number_offset + packet_number_len > parsed.packet_end
{
return Err(io::Error::new(
io::ErrorKind::InvalidData,
"invalid QUIC Initial packet number",
));
}
let mut packet_number = 0_u64;
for byte in
&protected[parsed.packet_number_offset..parsed.packet_number_offset + packet_number_len]
{
packet_number = (packet_number << 8) | u64::from(*byte);
}
let header_end = parsed.packet_number_offset + packet_number_len;
let mut payload = protected[header_end..parsed.packet_end].to_vec();
keys.client
.packet_key
.decrypt(packet_number, &protected[..header_end], &mut payload)
.map_err(|error| io::Error::new(io::ErrorKind::InvalidData, error))?;
let mut frames = Vec::new();
let mut payload = &payload[..];
while !payload.is_empty() {
let before = payload.len();
let frame = Frame::decode(&mut payload)
.map_err(|error| io::Error::new(io::ErrorKind::InvalidData, error))?;
frames.push(frame);
if payload.len() == before {
return Err(io::Error::new(
io::ErrorKind::InvalidData,
"QUIC frame decoder made no progress",
));
}
}
Ok(DecryptedInitial {
version,
dcid: parsed.dcid,
scid: parsed.scid,
_packet_number: packet_number,
frames,
})
}
fn encode_server_initial(
version: Version,
client_scid: ConnectionId,
server_scid: ConnectionId,
packet_number: u64,
crypto_data: Vec<u8>,
keys: nxtquic_crypto::InitialKeys,
) -> io::Result<Vec<u8>> {
let mut plaintext = Vec::new();
Frame::Crypto(CryptoFrame {
offset: VarInt::ZERO,
data: crypto_data.into(),
})
.encode(&mut plaintext);
let mut header_prefix = Vec::new();
header_prefix.push(0xc0); header_prefix.extend_from_slice(&version.as_u32().to_be_bytes());
header_prefix.push(client_scid.len() as u8);
header_prefix.extend_from_slice(client_scid.as_bytes());
header_prefix.push(server_scid.len() as u8);
header_prefix.extend_from_slice(server_scid.as_bytes());
header_prefix.push(0);
loop {
let encrypted_len = plaintext.len() + keys.server.packet_key.tag_len();
let mut candidate = header_prefix.clone();
VarInt::from_u64((1 + encrypted_len) as u64)
.expect("QUIC packet lengths fit in a VarInt")
.encode(&mut candidate);
if candidate.len() + 1 + encrypted_len >= 1200 {
header_prefix = candidate;
break;
}
plaintext.push(0);
}
let packet_number_len = 1;
let packet_number_offset = header_prefix.len();
header_prefix.push(packet_number as u8);
keys.server
.packet_key
.encrypt(packet_number, &header_prefix, &mut plaintext)
.map_err(|error| io::Error::new(io::ErrorKind::Other, error))?;
header_prefix.extend_from_slice(&plaintext);
let sample: [u8; 16] = header_prefix[packet_number_offset + 4..packet_number_offset + 20]
.try_into()
.expect("1200-byte Initial always contains a header-protection sample");
let mask = keys
.server
.header_key
.protection_mask(&sample)
.map_err(|error| io::Error::new(io::ErrorKind::Other, error))?;
apply_header_protection(
&mut header_prefix,
&mask,
packet_number_offset,
packet_number_len,
);
Ok(header_prefix)
}
impl std::fmt::Debug for Incoming {
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
formatter
.debug_struct("Incoming")
.field("remote_addr", &self.remote_addr)
.field("packet_len", &self.packet.len())
.finish()
}
}
impl Incoming {
pub fn remote_address(&self) -> SocketAddr {
self.remote_addr
}
pub fn local_ip(&self) -> Option<IpAddr> {
self.socket.local_addr().ok().map(|a| a.ip())
}
pub fn cancel(&mut self) {
self.datagrams = None;
}
pub async fn reject(&mut self, error_code: u64, reason: &[u8]) -> io::Result<()> {
let _ = error_code;
let _ = reason;
self.cancel();
Ok(())
}
pub(crate) fn packet(&self) -> &[u8] {
&self.packet
}
}
impl Endpoint {
pub async fn bind(addr: SocketAddr) -> io::Result<Self> {
Self::bind_inner(EndpointConfig::default(), None, addr).await
}
pub async fn server(server_config: ServerConfig, addr: SocketAddr) -> io::Result<Self> {
Self::bind_inner(EndpointConfig::default(), Some(server_config), addr).await
}
async fn bind_inner(
config: EndpointConfig,
server_config: Option<ServerConfig>,
addr: SocketAddr,
) -> io::Result<Self> {
let socket = tokio::net::UdpSocket::bind(addr).await?;
let socket = Arc::new(socket);
let (incoming_tx, incoming_rx) = tokio::sync::mpsc::channel(128);
let receive_socket = Arc::clone(&socket);
let shared_server_config = Arc::new(Mutex::new(server_config));
let receive_server_config = Arc::clone(&shared_server_config);
let stats = Arc::new(Mutex::new(EndpointStats::default()));
let receive_stats = Arc::clone(&stats);
tokio::spawn(async move {
let mut packet = vec![0_u8; 65_535];
let mut peers: HashMap<SocketAddr, tokio::sync::mpsc::Sender<Vec<u8>>> = HashMap::new();
loop {
let (length, remote_addr) = match receive_socket.recv_from(&mut packet).await {
Ok(received) => received,
Err(_) => break,
};
{
let mut st = receive_stats.lock().await;
st.packets_received += 1;
}
let datagram = &packet[..length];
if let Some(peer_tx) = peers.get(&remote_addr) {
if peer_tx.send(datagram.to_vec()).await.is_err() {
peers.remove(&remote_addr);
}
continue;
}
let Ok(header) = decode_header(datagram) else {
continue;
};
if header.packet_type != PacketType::Initial {
continue;
}
let Ok(initial) = decrypt_client_initial(datagram) else {
let mut st = receive_stats.lock().await;
st.handshake_failures += 1;
continue;
};
let (peer_tx, peer_rx) = tokio::sync::mpsc::channel(256);
peers.insert(remote_addr, peer_tx);
let cfg = receive_server_config.lock().await.clone();
{
let mut st = receive_stats.lock().await;
st.total_connections += 1;
st.active_connections += 1;
}
if incoming_tx
.send(Incoming {
remote_addr,
packet: datagram.to_vec(),
initial,
socket: Arc::clone(&receive_socket),
server_config: cfg,
datagrams: Some(peer_rx),
})
.await
.is_err()
{
break;
}
}
});
Ok(Self {
_config: config,
server_config: shared_server_config,
socket,
incoming: Arc::new(Mutex::new(incoming_rx)),
closed: Arc::new(AtomicBool::new(false)),
stats,
close_notify: Arc::new(Notify::new()),
})
}
pub fn local_addr(&self) -> io::Result<SocketAddr> {
self.socket.local_addr()
}
pub async fn server_config(&self) -> Option<ServerConfig> {
self.server_config.lock().await.clone()
}
pub async fn set_server_config(&self, config: ServerConfig) {
*self.server_config.lock().await = Some(config);
}
pub async fn stats(&self) -> EndpointStats {
self.stats.lock().await.clone()
}
pub async fn close(&self, _code: u64, _reason: &[u8]) {
self.closed.store(true, Ordering::Release);
self.close_notify.notify_waiters();
}
pub async fn wait_idle(&self) {
loop {
let active = self.stats.lock().await.active_connections;
if active == 0 {
break;
}
tokio::time::sleep(Duration::from_millis(50)).await;
}
}
pub async fn connect(
&self,
_addr: SocketAddr,
_server_name: &str,
) -> std::io::Result<Connecting> {
Ok(Connecting)
}
pub async fn connect_with_config(
&self,
_addr: SocketAddr,
_server_name: &str,
_config: ClientConfig,
) -> std::io::Result<Connecting> {
Ok(Connecting)
}
pub async fn connect_0rtt(
&self,
_addr: SocketAddr,
_server_name: &str,
) -> std::io::Result<Connecting> {
Ok(Connecting)
}
pub async fn accept(&self) -> Option<Incoming> {
self.incoming.lock().await.recv().await
}
}
impl Connecting {
pub async fn await_connection(self) -> std::io::Result<Connection> {
Ok(Connection::new())
}
pub async fn into_0rtt(self) -> std::io::Result<(Connection, Connection)> {
let conn = Connection::new();
Ok((conn.clone(), conn))
}
}
impl Incoming {
pub async fn accept(self) -> std::io::Result<Connection> {
let config = self.server_config.ok_or_else(|| {
io::Error::new(
io::ErrorKind::InvalidInput,
"cannot accept a connection on an endpoint without ServerConfig",
)
})?;
let mut crypto = Vec::new();
for frame in &self.initial.frames {
if let Frame::Crypto(frame) = frame {
if frame.offset.into_inner() != crypto.len() as u64 {
return Err(io::Error::new(
io::ErrorKind::InvalidData,
"Initial CRYPTO frames must be contiguous",
));
}
crypto.extend_from_slice(&frame.data);
}
}
if crypto.is_empty() {
return Err(io::Error::new(
io::ErrorKind::InvalidData,
"Initial packet contains no CRYPTO frame",
));
}
let provider = RustlsCryptoProvider::new(None, Some(Arc::clone(config.rustls())));
let mut transport_params = nxtquic_proto::TransportParameters::default();
transport_params.initial_max_data = VarInt::from_u64(10 * 1024 * 1024).unwrap();
transport_params.initial_max_stream_data_bidi_local =
VarInt::from_u64(1024 * 1024).unwrap();
transport_params.initial_max_stream_data_bidi_remote =
VarInt::from_u64(1024 * 1024).unwrap();
transport_params.initial_max_stream_data_uni = VarInt::from_u64(1024 * 1024).unwrap();
transport_params.initial_max_streams_bidi = VarInt::from_u64(100).unwrap();
transport_params.initial_max_streams_uni = VarInt::from_u64(100).unwrap();
if let Some(timeout) = config.transport().idle_timeout() {
transport_params.max_idle_timeout =
VarInt::from_u64(timeout.as_millis() as u64).unwrap_or(VarInt::from_u32(0));
}
transport_params.max_udp_payload_size = VarInt::from_u64(1350).unwrap();
transport_params.active_connection_id_limit = VarInt::from_u64(2).unwrap();
let server_scid = ConnectionId::from_slice(&rand::random::<[u8; 16]>());
transport_params.initial_source_connection_id = Some(server_scid.clone());
transport_params.original_destination_connection_id = Some(self.initial.dcid.clone());
let mut tp_encoded = Vec::new();
transport_params.encode(&mut tp_encoded);
let mut handshake = provider
.new_server_handshake(tp_encoded)
.map_err(|error| io::Error::new(io::ErrorKind::InvalidData, error))?;
let output = handshake
.process(&crypto)
.map_err(|error| io::Error::new(io::ErrorKind::InvalidData, error))?;
if output.data.is_empty() {
return Err(io::Error::new(
io::ErrorKind::InvalidData,
"TLS produced no Initial handshake data",
));
}
let keys = provider.initial_keys(self.initial.version, &self.initial.dcid);
let response = encode_server_initial(
self.initial.version,
self.initial.scid,
server_scid,
0,
output.data,
keys,
)?;
self.socket.send_to(&response, self.remote_addr).await?;
let (outgoing_tx, outgoing_rx) = tokio::sync::mpsc::unbounded_channel();
let connection = Connection::with_outgoing(outgoing_tx);
if let Some(datagrams) = self.datagrams {
let driver_connection = connection.clone();
let driver_initial = self.initial;
let driver_socket = Arc::clone(&self.socket);
let driver_remote = self.remote_addr;
let driver_initial_keys =
provider.initial_keys(driver_initial.version, &driver_initial.dcid);
let mut driver_handshake_keys = None;
let mut driver_one_rtt_keys = None;
while let Some((level, keys)) = handshake.next_keys() {
if level == nxtquic_crypto::Level::OneRtt {
driver_one_rtt_keys = Some(keys);
} else {
driver_handshake_keys = Some((level, keys));
}
}
tokio::spawn(async move {
drive_server_connection(
driver_connection,
driver_socket,
driver_remote,
datagrams,
driver_initial,
server_scid,
driver_initial_keys,
driver_handshake_keys,
driver_one_rtt_keys,
handshake,
outgoing_rx,
)
.await;
});
}
Ok(connection)
}
}
struct LongPacket {
_packet_type: PacketType,
_dcid: ConnectionId,
_scid: ConnectionId,
pn_offset: usize,
packet_end: usize,
}
fn parse_long_packet(datagram: &[u8]) -> io::Result<LongPacket> {
if datagram.len() < 7 || datagram[0] & 0xc0 != 0xc0 {
return Err(io::Error::new(
io::ErrorKind::InvalidData,
"invalid QUIC long header",
));
}
let version = u32::from_be_bytes(datagram[1..5].try_into().unwrap());
if version == 0 {
return Err(io::Error::new(
io::ErrorKind::InvalidData,
"version negotiation is not a protected packet",
));
}
let packet_type = match (datagram[0] >> 4) & 0x03 {
0 => PacketType::Initial,
1 => PacketType::ZeroRtt,
2 => PacketType::Handshake,
_ => PacketType::Retry,
};
let mut offset = 5;
let dcid_len = datagram[offset] as usize;
offset += 1;
let dcid_end = offset
.checked_add(dcid_len)
.ok_or_else(|| io::Error::new(io::ErrorKind::InvalidData, "invalid DCID"))?;
let dcid = ConnectionId::from_slice(
datagram
.get(offset..dcid_end)
.ok_or_else(|| io::Error::new(io::ErrorKind::UnexpectedEof, "truncated DCID"))?,
);
offset = dcid_end;
let scid_len = *datagram
.get(offset)
.ok_or_else(|| io::Error::new(io::ErrorKind::UnexpectedEof, "missing SCID length"))?
as usize;
offset += 1;
let scid_end = offset
.checked_add(scid_len)
.ok_or_else(|| io::Error::new(io::ErrorKind::InvalidData, "invalid SCID"))?;
let scid = ConnectionId::from_slice(
datagram
.get(offset..scid_end)
.ok_or_else(|| io::Error::new(io::ErrorKind::UnexpectedEof, "truncated SCID"))?,
);
offset = scid_end;
if packet_type == PacketType::Initial {
let mut rest = &datagram[offset..];
let token_len = VarInt::decode(&mut rest)
.map_err(|_| io::Error::new(io::ErrorKind::InvalidData, "invalid Initial token length"))?
.into_inner() as usize;
let consumed = datagram[offset..].len() - rest.len();
offset += consumed
.checked_add(token_len)
.ok_or_else(|| io::Error::new(io::ErrorKind::InvalidData, "invalid token"))?;
}
let mut rest = &datagram[offset..];
let payload_len = VarInt::decode(&mut rest)
.map_err(|_| io::Error::new(io::ErrorKind::InvalidData, "invalid QUIC payload length"))?
.into_inner() as usize;
offset += datagram[offset..].len() - rest.len();
let packet_end = offset
.checked_add(payload_len)
.ok_or_else(|| io::Error::new(io::ErrorKind::InvalidData, "invalid packet length"))?;
if packet_end > datagram.len() || packet_end < offset + 4 + 16 {
return Err(io::Error::new(
io::ErrorKind::UnexpectedEof,
"truncated protected QUIC packet",
));
}
Ok(LongPacket {
_packet_type: packet_type,
_dcid: dcid,
_scid: scid,
pn_offset: offset,
packet_end,
})
}
fn decrypt_payload(
mut packet: Vec<u8>,
parsed: &LongPacket,
keys: &DirectionalKeys,
) -> io::Result<(u64, Vec<u8>)> {
let sample: [u8; 16] = packet[parsed.pn_offset + 4..parsed.pn_offset + 20]
.try_into()
.map_err(|_| {
io::Error::new(
io::ErrorKind::InvalidData,
"missing header-protection sample",
)
})?;
let mask = keys
.header_key
.protection_mask(&sample)
.map_err(|error| io::Error::new(io::ErrorKind::InvalidData, error))?;
let pn_len = remove_header_protection(&mut packet, &mask, parsed.pn_offset);
if pn_len == 0 || parsed.pn_offset + pn_len > parsed.packet_end {
return Err(io::Error::new(
io::ErrorKind::InvalidData,
"invalid packet number",
));
}
let mut packet_number = 0;
for byte in &packet[parsed.pn_offset..parsed.pn_offset + pn_len] {
packet_number = (packet_number << 8) | u64::from(*byte);
}
let header_end = parsed.pn_offset + pn_len;
let mut payload = packet[header_end..parsed.packet_end].to_vec();
keys.packet_key
.decrypt(packet_number, &packet[..header_end], &mut payload)
.map_err(|error| io::Error::new(io::ErrorKind::InvalidData, error))?;
Ok((packet_number, payload))
}
async fn drive_server_connection(
connection: Connection,
socket: Arc<tokio::net::UdpSocket>,
remote_addr: SocketAddr,
mut datagrams: tokio::sync::mpsc::Receiver<Vec<u8>>,
initial: DecryptedInitial,
server_scid: ConnectionId,
initial_keys: nxtquic_crypto::InitialKeys,
mut handshake_keys: Option<(nxtquic_crypto::Level, Keys)>,
mut one_rtt_keys: Option<Keys>,
mut handshake: impl HandshakeState,
mut outgoing: tokio::sync::mpsc::UnboundedReceiver<crate::stream::WriteCommand>,
) {
let mut largest_handshake = 0;
let mut application_packet_number = 0;
loop {
let Some(event) = (tokio::select! {
datagram = datagrams.recv() => datagram.map(Ok),
command = outgoing.recv() => command.map(Err),
else => None,
}) else {
break;
};
let datagram = match event {
Err(command) => {
if let Some(keys) = one_rtt_keys.as_ref() {
match command {
crate::stream::WriteCommand::Data {
stream_id,
offset,
data,
fin,
} => {
let mut encoded = Vec::new();
Frame::Stream(StreamFrame {
stream_id: StreamId::from_u64(stream_id),
offset: VarInt::from_u64(offset).unwrap(),
length: Some(VarInt::from_u64(data.len() as u64).unwrap()),
fin,
data: bytes::Bytes::from(data),
})
.encode(&mut encoded);
let _ = send_protected_short(
&socket,
remote_addr,
initial.scid,
application_packet_number,
&encoded,
&keys.local,
)
.await;
application_packet_number += 1;
}
crate::stream::WriteCommand::Reset {
stream_id,
error_code,
} => {
let mut encoded = Vec::new();
Frame::ResetStream(nxtquic_proto::frame::ResetStreamFrame {
stream_id: StreamId::from_u64(stream_id),
application_protocol_error_code: VarInt::from_u64(error_code).unwrap(),
final_size: VarInt::ZERO,
})
.encode(&mut encoded);
let _ = send_protected_short(
&socket,
remote_addr,
initial.scid,
application_packet_number,
&encoded,
&keys.local,
)
.await;
application_packet_number += 1;
}
crate::stream::WriteCommand::StopSending {
stream_id,
error_code,
} => {
let mut encoded = Vec::new();
Frame::StopSending(nxtquic_proto::frame::StopSendingFrame {
stream_id: StreamId::from_u64(stream_id),
application_protocol_error_code: VarInt::from_u64(error_code).unwrap(),
})
.encode(&mut encoded);
let _ = send_protected_short(
&socket,
remote_addr,
initial.scid,
application_packet_number,
&encoded,
&keys.local,
)
.await;
application_packet_number += 1;
}
crate::stream::WriteCommand::Priority { .. } => {}
}
}
continue;
}
Ok(datagram) => datagram,
};
let Ok(header) = decode_header(&datagram) else {
continue;
};
let (payload, packet_number, level) = match header.packet_type {
PacketType::Initial => {
let Ok(parsed) = parse_long_packet(&datagram) else {
continue;
};
let Ok((pn, payload)) =
decrypt_payload(datagram, &parsed, &initial_keys.client)
else {
continue;
};
(payload, pn, nxtquic_crypto::Level::Initial)
}
PacketType::Handshake => {
let Ok(parsed) = parse_long_packet(&datagram) else {
continue;
};
let Some((_, keys)) = handshake_keys.as_ref() else {
continue;
};
let Ok((pn, payload)) = decrypt_payload(datagram, &parsed, &keys.remote) else {
continue;
};
largest_handshake = largest_handshake.max(pn);
(payload, pn, nxtquic_crypto::Level::Handshake)
}
PacketType::Short => {
let Some(keys) = one_rtt_keys.as_ref() else {
continue;
};
if datagram.len() < 1 + server_scid.len() + 4 + 16 {
continue;
}
let parsed = LongPacket {
_packet_type: PacketType::Short,
_dcid: server_scid,
_scid: ConnectionId::from_slice(&[]),
pn_offset: 1 + server_scid.len(),
packet_end: datagram.len(),
};
let Ok((pn, payload)) = decrypt_payload(datagram, &parsed, &keys.remote) else {
continue;
};
(payload, pn, nxtquic_crypto::Level::OneRtt)
}
_ => continue,
};
let mut frames = &payload[..];
while !frames.is_empty() {
let before = frames.len();
let Ok(frame) = Frame::decode(&mut frames) else {
break;
};
match frame {
Frame::Crypto(crypto)
if matches!(
level,
nxtquic_crypto::Level::Initial | nxtquic_crypto::Level::Handshake
) =>
{
let Ok(output) = handshake.process(&crypto.data) else {
continue;
};
while let Some((next_level, keys)) = handshake.next_keys() {
if next_level == nxtquic_crypto::Level::OneRtt {
one_rtt_keys = Some(keys);
} else {
handshake_keys = Some((next_level, keys));
}
}
if !output.data.is_empty() {
let (packet_type, packet_number, send_keys) = match output.level {
nxtquic_crypto::Level::Initial => {
(PacketType::Initial, 1, Some(&initial_keys.server))
}
nxtquic_crypto::Level::Handshake => (
PacketType::Handshake,
largest_handshake + 1,
handshake_keys.as_ref().map(|(_, k)| &k.local),
),
_ => (PacketType::Handshake, largest_handshake + 1, None),
};
let _ = send_protected_long(
&socket,
remote_addr,
packet_type,
initial.scid,
server_scid,
packet_number,
&output.data,
send_keys,
)
.await;
}
if handshake.is_complete() {
if let Some(keys) = one_rtt_keys.as_ref() {
let mut handshake_done = Vec::new();
Frame::HandshakeDone.encode(&mut handshake_done);
let _ = send_protected_short(
&socket,
remote_addr,
initial.scid,
0,
&handshake_done,
&keys.local,
)
.await;
}
}
}
Frame::Stream(stream) if level == nxtquic_crypto::Level::OneRtt => {
let mut encoded = Vec::new();
Frame::Stream(stream).encode(&mut encoded);
let _ = connection.ingest_frames(&encoded).await;
}
_ => {}
}
if frames.len() >= before {
break;
}
}
let mut ack = Vec::new();
Frame::Ack(AckFrame {
largest_acknowledged: VarInt::from_u64(packet_number).unwrap(),
ack_delay: VarInt::ZERO,
first_ack_range: VarInt::ZERO,
ack_ranges: Vec::new(),
ecn_counts: None,
})
.encode(&mut ack);
match level {
nxtquic_crypto::Level::Initial => {
let _ = send_protected_long(
&socket,
remote_addr,
PacketType::Initial,
initial.scid,
server_scid,
2,
&ack,
Some(&initial_keys.server),
)
.await;
}
nxtquic_crypto::Level::Handshake => {
if let Some((_, keys)) = handshake_keys.as_ref() {
let _ = send_protected_long(
&socket,
remote_addr,
PacketType::Handshake,
initial.scid,
server_scid,
largest_handshake + 1,
&ack,
Some(&keys.local),
)
.await;
}
}
nxtquic_crypto::Level::OneRtt => {
if let Some(keys) = one_rtt_keys.as_ref() {
let _ = send_protected_short(
&socket,
remote_addr,
initial.scid,
application_packet_number,
&ack,
&keys.local,
)
.await;
application_packet_number += 1;
}
}
_ => {}
}
}
connection.close().await;
}
async fn send_protected_long(
socket: &tokio::net::UdpSocket,
remote_addr: SocketAddr,
packet_type: PacketType,
dcid: ConnectionId,
scid: ConnectionId,
packet_number: u64,
payload: &[u8],
keys: Option<&DirectionalKeys>,
) -> io::Result<()> {
let Some(keys) = keys else {
return Ok(());
};
let type_bits = match packet_type {
PacketType::Handshake => 0x20,
PacketType::Initial => 0x00,
_ => return Ok(()),
};
let mut header = vec![0xC0 | type_bits, 0, 0, 0, 1, dcid.len() as u8];
header.extend_from_slice(dcid.as_bytes());
header.push(scid.len() as u8);
header.extend_from_slice(scid.as_bytes());
if packet_type == PacketType::Initial {
VarInt::ZERO.encode(&mut header);
}
let mut body = payload.to_vec();
while body.len() < 20 {
body.push(0);
}
let packet_len = 1 + body.len() + keys.packet_key.tag_len();
VarInt::from_u64(packet_len as u64)
.map_err(|_| io::Error::new(io::ErrorKind::InvalidInput, "packet too large"))?
.encode(&mut header);
let pn_offset = header.len();
header.push(packet_number as u8);
keys.packet_key
.encrypt(packet_number, &header, &mut body)
.map_err(|error| io::Error::new(io::ErrorKind::Other, error))?;
header.extend_from_slice(&body);
let sample: [u8; 16] = header[pn_offset + 4..pn_offset + 20]
.try_into()
.map_err(|_| io::Error::new(io::ErrorKind::InvalidData, "packet too short for sample"))?;
let mask = keys
.header_key
.protection_mask(&sample)
.map_err(|error| io::Error::new(io::ErrorKind::Other, error))?;
apply_header_protection(&mut header, &mask, pn_offset, 1);
socket.send_to(&header, remote_addr).await.map(|_| ())
}
async fn send_protected_short(
socket: &tokio::net::UdpSocket,
remote_addr: SocketAddr,
dcid: ConnectionId,
packet_number: u64,
payload: &[u8],
keys: &DirectionalKeys,
) -> io::Result<()> {
let mut header = vec![0x40];
header.extend_from_slice(dcid.as_bytes());
let pn_offset = header.len();
header.push(packet_number as u8);
let mut body = payload.to_vec();
while body.len() < 20 {
body.push(0);
}
keys.packet_key
.encrypt(packet_number, &header, &mut body)
.map_err(|error| io::Error::new(io::ErrorKind::Other, error))?;
header.extend_from_slice(&body);
let sample: [u8; 16] = header[pn_offset + 4..pn_offset + 20]
.try_into()
.map_err(|_| io::Error::new(io::ErrorKind::InvalidData, "packet too short for sample"))?;
let mask = keys
.header_key
.protection_mask(&sample)
.map_err(|error| io::Error::new(io::ErrorKind::Other, error))?;
apply_header_protection(&mut header, &mask, pn_offset, 1);
socket.send_to(&header, remote_addr).await.map(|_| ())
}
#[cfg(test)]
mod packet_tests {
use super::*;
use bytes::Bytes;
use nxtquic_crypto::{CryptoProvider, RustlsCryptoProvider};
use nxtquic_proto::{
frame::{CryptoFrame, Frame},
packet::protection::apply_header_protection,
varint::VarInt,
ConnectionId, Version,
};
fn protected_client_initial() -> Vec<u8> {
let dcid = ConnectionId::from_slice(&[0x10; 8]);
let scid = [0x20; 8];
let provider = RustlsCryptoProvider::new(None, None);
let keys = provider.initial_keys(Version::V1, &dcid);
let mut plaintext = Vec::new();
Frame::Ping.encode(&mut plaintext);
Frame::Crypto(CryptoFrame {
offset: VarInt::ZERO,
data: Bytes::from_static(b"client hello bytes"),
})
.encode(&mut plaintext);
let ciphertext_len = plaintext.len() + keys.client.packet_key.tag_len();
let mut header = vec![0xc0, 0, 0, 0, 1, dcid.len() as u8];
header.extend_from_slice(dcid.as_bytes());
header.push(scid.len() as u8);
header.extend_from_slice(&scid);
header.push(0); VarInt::from_u64((1 + ciphertext_len) as u64)
.unwrap()
.encode(&mut header);
let packet_number_offset = header.len();
header.push(0);
keys.client
.packet_key
.encrypt(0, &header, &mut plaintext)
.unwrap();
header.extend_from_slice(&plaintext);
let sample: [u8; 16] = header[packet_number_offset + 4..packet_number_offset + 20]
.try_into()
.unwrap();
let mask = keys.client.header_key.protection_mask(&sample).unwrap();
apply_header_protection(&mut header, &mask, packet_number_offset, 1);
header
}
#[test]
fn decrypts_client_initial_and_decodes_crypto_frames() {
let packet = protected_client_initial();
let initial = decrypt_client_initial(&packet).unwrap();
assert_eq!(initial._packet_number, 0);
assert!(matches!(initial.frames.first(), Some(Frame::Ping)));
assert!(
matches!(initial.frames.get(1), Some(Frame::Crypto(frame)) if frame.data == Bytes::from_static(b"client hello bytes"))
);
}
#[test]
fn rejects_tampered_client_initial() {
let mut packet = protected_client_initial();
let last = packet.len() - 1;
packet[last] ^= 0x01;
assert!(decrypt_client_initial(&packet).is_err());
}
}