use std::{
net::{IpAddr, Ipv4Addr, SocketAddr},
sync::Arc,
};
use bytes::BytesMut;
use crate::{
ClientConfig, Connection, ConnectionHandle, EndpointConfig, Event, ServerConfig,
TransportConfig,
crypto::{
pqc::PqcConfig,
rustls::{QuicClientConfig, QuicServerConfig, configured_provider_with_pqc},
},
endpoint::{DatagramEvent, Endpoint},
};
const MIN_INITIAL_SIZE: usize = 1200;
const SERVER_ADDR: SocketAddr = SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 4433);
#[derive(Debug)]
struct SkipServerVerification;
impl rustls::client::danger::ServerCertVerifier for SkipServerVerification {
fn verify_server_cert(
&self,
_end_entity: &rustls::pki_types::CertificateDer<'_>,
_intermediates: &[rustls::pki_types::CertificateDer<'_>],
_server_name: &rustls::pki_types::ServerName<'_>,
_ocsp_response: &[u8],
_now: rustls::pki_types::UnixTime,
) -> Result<rustls::client::danger::ServerCertVerified, rustls::Error> {
Ok(rustls::client::danger::ServerCertVerified::assertion())
}
fn verify_tls12_signature(
&self,
_message: &[u8],
_cert: &rustls::pki_types::CertificateDer<'_>,
_dss: &rustls::DigitallySignedStruct,
) -> Result<rustls::client::danger::HandshakeSignatureValid, rustls::Error> {
Ok(rustls::client::danger::HandshakeSignatureValid::assertion())
}
fn verify_tls13_signature(
&self,
_message: &[u8],
_cert: &rustls::pki_types::CertificateDer<'_>,
_dss: &rustls::DigitallySignedStruct,
) -> Result<rustls::client::danger::HandshakeSignatureValid, rustls::Error> {
Ok(rustls::client::danger::HandshakeSignatureValid::assertion())
}
fn supported_verify_schemes(&self) -> Vec<rustls::SignatureScheme> {
vec![
rustls::SignatureScheme::RSA_PKCS1_SHA256,
rustls::SignatureScheme::RSA_PKCS1_SHA384,
rustls::SignatureScheme::RSA_PKCS1_SHA512,
rustls::SignatureScheme::ED25519,
rustls::SignatureScheme::ECDSA_NISTP256_SHA256,
rustls::SignatureScheme::ECDSA_NISTP384_SHA384,
rustls::SignatureScheme::RSA_PSS_SHA256,
rustls::SignatureScheme::RSA_PSS_SHA384,
rustls::SignatureScheme::RSA_PSS_SHA512,
]
}
}
fn generate_test_cert() -> (
rustls::pki_types::CertificateDer<'static>,
rustls::pki_types::PrivateKeyDer<'static>,
) {
let cert = rcgen::generate_simple_self_signed(vec!["localhost".to_string()]).unwrap();
let cert_der = cert.cert.into();
let key_der = rustls::pki_types::PrivateKeyDer::Pkcs8(cert.signing_key.serialize_der().into());
(cert_der, key_der)
}
fn server_config() -> Arc<ServerConfig> {
let (cert, key) = generate_test_cert();
let provider = configured_provider_with_pqc(Some(&PqcConfig::default()));
let mut server_crypto = rustls::ServerConfig::builder_with_provider(provider)
.with_protocol_versions(&[&rustls::version::TLS13])
.unwrap()
.with_no_client_auth()
.with_single_cert(vec![cert], key)
.unwrap();
server_crypto.alpn_protocols = vec![b"test".to_vec()];
let mut config =
ServerConfig::with_crypto(Arc::new(QuicServerConfig::try_from(server_crypto).unwrap()));
config.transport_config(Arc::new(TransportConfig::default()));
Arc::new(config)
}
fn client_config() -> ClientConfig {
let provider = configured_provider_with_pqc(Some(&PqcConfig::default()));
let mut client_crypto = rustls::ClientConfig::builder_with_provider(provider)
.with_protocol_versions(&[&rustls::version::TLS13])
.unwrap()
.dangerous()
.with_custom_certificate_verifier(Arc::new(SkipServerVerification))
.with_no_client_auth();
client_crypto.alpn_protocols = vec![b"test".to_vec()];
let mut config =
ClientConfig::new(Arc::new(QuicClientConfig::try_from(client_crypto).unwrap()));
config.transport_config(Arc::new(TransportConfig::default()));
config
}
struct Peer {
endpoint: Endpoint,
conn: Option<Connection>,
handle: Option<ConnectionHandle>,
connected: bool,
}
fn ingress(now: crate::Instant, data: BytesMut, from_addr: SocketAddr, peer: &mut Peer) {
let mut scratch = Vec::new();
match peer
.endpoint
.handle(now, from_addr, None, None, data, &mut scratch)
{
Some(DatagramEvent::ConnectionEvent(_, event)) => {
if let Some(conn) = peer.conn.as_mut() {
conn.handle_event(event);
}
}
Some(DatagramEvent::NewConnection(incoming)) => {
let (handle, conn) = peer
.endpoint
.accept(incoming, now, &mut scratch, None)
.expect("accept first client Initial");
peer.handle = Some(handle);
peer.conn = Some(conn);
}
Some(DatagramEvent::Response(_)) | None => {}
}
}
fn pump(now: crate::Instant, peer: &mut Peer, other: &mut Peer, sizes: &mut Vec<usize>) -> bool {
let mut progress = false;
let handle = peer.handle.unwrap();
while let Some(event) = peer.conn.as_mut().unwrap().poll_endpoint_events() {
if let Some(response) = peer.endpoint.handle_event(handle, event) {
if let Some(conn) = peer.conn.as_mut() {
conn.handle_event(response);
}
}
progress = true;
}
while let Some(event) = peer.conn.as_mut().unwrap().poll() {
if matches!(event, Event::Connected) {
peer.connected = true;
}
progress = true;
}
{
let conn = peer.conn.as_mut().unwrap();
let mut buf = Vec::with_capacity(65_536);
while let Some(transmit) = conn.poll_transmit(now, 1, &mut buf) {
sizes.push(transmit.size);
let datagram = BytesMut::from(&buf[..transmit.size]);
ingress(now, datagram, transmit.destination, other);
buf.clear();
progress = true;
}
}
progress
}
fn drive_to_connected(client: &mut Peer, server: &mut Peer) -> (Vec<usize>, crate::Instant) {
let mut now = crate::Instant::now();
let mut sizes = Vec::new();
for _ in 0..400 {
let mut progress = false;
progress |= pump(now, client, server, &mut sizes);
progress |= pump(now, server, client, &mut sizes);
if client.connected && server.connected {
return (sizes, now);
}
if !progress {
now += crate::Duration::from_millis(25);
for peer in [&mut *client, &mut *server] {
if let Some(conn) = peer.conn.as_mut() {
if conn.poll_timeout().is_some_and(|deadline| deadline <= now) {
conn.handle_timeout(now);
}
}
}
}
}
(sizes, now)
}
fn make_peer_pair() -> (Peer, Peer) {
let server = Peer {
endpoint: Endpoint::new(
Arc::new(EndpointConfig::default()),
Some(server_config()),
false,
None,
),
conn: None,
handle: None,
connected: false,
};
let mut client = Peer {
endpoint: Endpoint::new(Arc::new(EndpointConfig::default()), None, false, None),
conn: None,
handle: None,
connected: false,
};
let (handle, conn) = client
.endpoint
.connect(
crate::Instant::now(),
client_config(),
SERVER_ADDR,
"localhost",
)
.expect("initiate client connection");
client.handle = Some(handle);
client.conn = Some(conn);
(client, server)
}
#[test]
fn pqc_handshake_datagrams_stay_within_1200_bytes() {
let (mut client, mut server) = make_peer_pair();
let (sizes, _) = drive_to_connected(&mut client, &mut server);
assert!(
client.connected && server.connected,
"PQC handshake must complete over a 1200-byte path; got {} datagrams",
sizes.len()
);
assert!(
sizes.len() > 4,
"expected a multi-flight PQC handshake, saw only {} datagrams",
sizes.len()
);
let offenders: Vec<usize> = sizes
.iter()
.copied()
.filter(|&size| size > MIN_INITIAL_SIZE)
.collect();
assert_eq!(
offenders,
Vec::<usize>::new(),
"datagrams exceeded 1200 bytes during PQC handshake"
);
assert_eq!(
sizes.first(),
Some(&MIN_INITIAL_SIZE),
"client's first Initial must be padded to exactly 1200 bytes"
);
assert_eq!(
client.conn.as_ref().unwrap().current_mtu(),
MIN_INITIAL_SIZE as u16
);
assert_eq!(
server.conn.as_ref().unwrap().current_mtu(),
MIN_INITIAL_SIZE as u16
);
}
#[test]
fn pqc_min_initial_size_is_clamped_to_path_mtu() {
use crate::MIN_INITIAL_SIZE;
use crate::transport_parameters::{PqcAlgorithms, TransportParameters};
use super::PqcState;
let mut state = PqcState::new();
let mut params = TransportParameters::default();
params.pqc_algorithms = Some(PqcAlgorithms {
ml_kem_768: true,
ml_dsa_65: true,
});
state.update_from_peer_params(¶ms);
assert!(state.enabled && state.using_pqc && state.handshake_mtu == 4096);
assert_eq!(state.min_initial_size(MIN_INITIAL_SIZE), MIN_INITIAL_SIZE);
assert_eq!(state.min_initial_size(1500), 1500);
assert_eq!(state.min_initial_size(4096), 4096);
assert_eq!(state.min_initial_size(576), MIN_INITIAL_SIZE);
let plain = PqcState::new();
assert_eq!(plain.min_initial_size(1500), MIN_INITIAL_SIZE);
}
#[test]
fn packet_builder_pad_to_is_capped_at_datagram_capacity() {
use super::packet_builder::PacketBuilder;
let (mut client, _server) = make_peer_pair();
let conn = client.conn.as_mut().unwrap();
let mut buf = Vec::with_capacity(MIN_INITIAL_SIZE);
let buf_capacity = MIN_INITIAL_SIZE;
let mut builder = PacketBuilder::new(
crate::Instant::now(),
crate::packet::SpaceId::Initial,
conn.rem_cids.active(),
&mut buf,
buf_capacity,
0,
true,
conn,
)
.expect("build Initial packet");
builder.pad_to(4096);
assert!(
builder.min_size <= builder.max_size,
"pad_to must not raise min_size past max_size ({} > {})",
builder.min_size,
builder.max_size
);
let (len, _padded) = builder.finish(conn, &mut buf);
assert!(
buf.len() <= buf_capacity,
"encrypted datagram {} exceeds buffer capacity {buf_capacity}",
buf.len()
);
assert_eq!(len, buf.len());
assert_eq!(buf.len(), MIN_INITIAL_SIZE);
assert_eq!(buf[0] & 0xc0, 0xc0);
}
#[test]
fn off_path_response_echoes_token_and_validates_peer_path() {
let (mut client, mut server) = make_peer_pair();
let (_, now) = drive_to_connected(&mut client, &mut server);
assert!(client.connected && server.connected);
let candidate = SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 4444);
let token = 0x0123_4567_89ab_cdef;
let client_conn = client.conn.as_mut().unwrap();
client_conn.nat_traversal = None;
assert_ne!(candidate, client_conn.path.remote);
client_conn.path_responses = super::paths::PathResponses::default();
client_conn.path_responses.push(0, token, candidate);
let mut buf = Vec::with_capacity(65_536);
let transmit = client_conn.poll_transmit(now, 1, &mut buf).unwrap();
assert_eq!(transmit.destination, candidate);
assert_eq!(transmit.size, MIN_INITIAL_SIZE);
assert_eq!(transmit.segment_size, None);
let server_conn = server.conn.as_mut().unwrap();
server_conn.path.challenge = Some(token);
server_conn.path.validated = false;
let source = server_conn.path.remote;
ingress(
now,
BytesMut::from(&buf[..transmit.size]),
source,
&mut server,
);
let server_conn = server.conn.as_ref().unwrap();
assert_eq!(
server_conn.path.challenge, None,
"peer must accept echoed token"
);
assert!(server_conn.path.validated);
}
#[test]
fn coordinated_path_challenge_is_padded_and_decoded_by_peer() {
use super::nat_traversal::{CoordinationPhase, NatTraversalState, PunchTarget};
let (mut client, mut server) = make_peer_pair();
let (_, now) = drive_to_connected(&mut client, &mut server);
assert!(client.connected && server.connected);
let token = 0xfedc_ba98_7654_3210;
let client_conn = client.conn.as_mut().unwrap();
let destination = client_conn.path.remote;
let mut nat = NatTraversalState::new(4, crate::Duration::from_secs(10));
nat.prime_passive_coordination_target(
crate::VarInt::from_u32(1),
PunchTarget {
remote_addr: destination,
remote_sequence: crate::VarInt::from_u32(1),
challenge: token,
},
now,
)
.unwrap();
assert_eq!(
nat.get_coordination_phase(),
Some(CoordinationPhase::Preparing)
);
client_conn.nat_traversal = Some(nat);
let now = now + crate::Duration::from_millis(500);
let mut buf = Vec::with_capacity(65_536);
let transmit = client_conn.poll_transmit(now, 1, &mut buf).unwrap();
assert_eq!(transmit.destination, destination);
assert_eq!(transmit.segment_size, None);
assert_eq!(
client_conn
.nat_traversal
.as_ref()
.unwrap()
.get_coordination_phase(),
Some(CoordinationPhase::Validating)
);
let server_conn = server.conn.as_mut().unwrap();
server_conn.path_responses = super::paths::PathResponses::default();
let source = server_conn.path.remote;
ingress(
now,
BytesMut::from(&buf[..transmit.size]),
source,
&mut server,
);
assert_eq!(
server
.conn
.as_mut()
.unwrap()
.path_responses
.pop_on_path(source),
Some(token),
"peer must authenticate and decode the coordinated challenge"
);
assert_eq!(transmit.size, MIN_INITIAL_SIZE);
}