use std::sync::Arc;
use std::time::{Duration, Instant};
use dimpl::{Config, Dtls};
use crate::common::*;
#[test]
#[cfg(feature = "rcgen")]
fn dtls13_handshake_with_small_mtu() {
use dimpl::certificate::generate_self_signed_certificate;
let _ = env_logger::try_init();
let client_cert = generate_self_signed_certificate().expect("gen client cert");
let server_cert = generate_self_signed_certificate().expect("gen server cert");
let config = dtls13_config_with_mtu(200);
let mut now = Instant::now();
let mut client = Dtls::new_13(Arc::clone(&config), client_cert, now);
client.set_active(true);
let mut server = Dtls::new_13(config, server_cert, now);
server.set_active(false);
let mut client_connected = false;
let mut server_connected = false;
let mut max_packet_size = 0usize;
for _ in 0..40 {
client.handle_timeout(now).expect("client timeout");
server.handle_timeout(now).expect("server timeout");
let client_out = drain_outputs(&mut client);
let server_out = drain_outputs(&mut server);
for p in &client_out.packets {
if p.len() > max_packet_size {
max_packet_size = p.len();
}
}
if client_out.connected {
client_connected = true;
}
if server_out.connected {
server_connected = true;
}
deliver_packets(&client_out.packets, &mut server);
deliver_packets(&server_out.packets, &mut client);
if client_connected && server_connected {
break;
}
now += Duration::from_millis(10);
}
assert!(client_connected, "Client should connect with small MTU");
assert!(server_connected, "Server should connect with small MTU");
assert!(
max_packet_size <= 200,
"Packets should respect MTU: max was {}",
max_packet_size
);
}
#[test]
#[cfg(feature = "rcgen")]
fn dtls13_large_application_data_fragmented() {
use dimpl::certificate::generate_self_signed_certificate;
let _ = env_logger::try_init();
let client_cert = generate_self_signed_certificate().expect("gen client cert");
let server_cert = generate_self_signed_certificate().expect("gen server cert");
let config = dtls13_config_with_mtu(300);
let mut now = Instant::now();
let mut client = Dtls::new_13(Arc::clone(&config), client_cert, now);
client.set_active(true);
let mut server = Dtls::new_13(config, server_cert, now);
server.set_active(false);
for _ in 0..40 {
client.handle_timeout(now).expect("client timeout");
server.handle_timeout(now).expect("server timeout");
let client_out = drain_outputs(&mut client);
let server_out = drain_outputs(&mut server);
deliver_packets(&client_out.packets, &mut server);
deliver_packets(&server_out.packets, &mut client);
if client_out.connected && server_out.connected {
break;
}
now += Duration::from_millis(10);
}
let large_data = vec![0xABu8; 1000];
client
.send_application_data(&large_data)
.expect("client send large data");
let mut server_received: Vec<u8> = Vec::new();
let mut _packet_count = 0;
for _ in 0..20 {
let client_out = drain_outputs(&mut client);
_packet_count += client_out.packets.len();
deliver_packets(&client_out.packets, &mut server);
let server_out = drain_outputs(&mut server);
for data in server_out.app_data {
server_received.extend_from_slice(&data);
}
if server_received.len() >= large_data.len() {
break;
}
now += Duration::from_millis(10);
}
assert_eq!(
server_received, large_data,
"Large data should be received correctly"
);
}
#[test]
#[cfg(feature = "rcgen")]
fn dtls13_fragmentation_during_hrr() {
use dimpl::certificate::generate_self_signed_certificate;
let _ = env_logger::try_init();
let client_cert = generate_self_signed_certificate().expect("gen client cert");
let server_cert = generate_self_signed_certificate().expect("gen server cert");
let config = dtls13_config_with_mtu(200);
let mut now = Instant::now();
let mut client = Dtls::new_13(Arc::clone(&config), client_cert, now);
client.set_active(true);
let mut server = Dtls::new_13(config, server_cert, now);
server.set_active(false);
let mut client_connected = false;
let mut server_connected = false;
let mut max_packet_size = 0usize;
let mut saw_hrr = false;
let mut flight_count = 0;
for _ in 0..40 {
client.handle_timeout(now).expect("client timeout");
server.handle_timeout(now).expect("server timeout");
let client_out = drain_outputs(&mut client);
let server_out = drain_outputs(&mut server);
if !client_out.packets.is_empty() {
flight_count += 1;
}
if !server_out.packets.is_empty() && !client_connected && flight_count <= 2 {
saw_hrr = true;
}
for p in &client_out.packets {
if p.len() > max_packet_size {
max_packet_size = p.len();
}
}
for p in &server_out.packets {
if p.len() > max_packet_size {
max_packet_size = p.len();
}
}
if client_out.connected {
client_connected = true;
}
if server_out.connected {
server_connected = true;
}
deliver_packets(&client_out.packets, &mut server);
deliver_packets(&server_out.packets, &mut client);
if client_connected && server_connected {
break;
}
now += Duration::from_millis(10);
}
assert!(
client_connected,
"Client should connect with HRR and small MTU"
);
assert!(
server_connected,
"Server should connect with HRR and small MTU"
);
assert!(
max_packet_size <= 200,
"Packets should respect MTU: max was {}",
max_packet_size
);
assert!(
saw_hrr || flight_count >= 2,
"Should have seen HRR or multiple client flights"
);
}
#[test]
#[cfg(feature = "rcgen")]
fn dtls13_fragmented_handshake_with_packet_loss() {
use dimpl::certificate::generate_self_signed_certificate;
let _ = env_logger::try_init();
let client_cert = generate_self_signed_certificate().expect("gen client cert");
let server_cert = generate_self_signed_certificate().expect("gen server cert");
let config = Arc::new(
Config::builder()
.mtu(200)
.flight_retries(8)
.handshake_timeout(Duration::from_secs(120))
.build()
.expect("Failed to build DTLS 1.3 config"),
);
let mut now = Instant::now();
let mut client = Dtls::new_13(Arc::clone(&config), client_cert, now);
client.set_active(true);
let mut server = Dtls::new_13(config, server_cert, now);
server.set_active(false);
let mut client_connected = false;
let mut server_connected = false;
let mut dropped_client = 0;
let mut dropped_server = 0;
let mut prev_client_count = 0usize;
let mut client_drop_armed = false;
let mut prev_server_count = 0usize;
let mut server_drop_armed = false;
for i in 0..120 {
client.handle_timeout(now).expect("client timeout");
server.handle_timeout(now).expect("server timeout");
let client_out = drain_outputs(&mut client);
let server_out = drain_outputs(&mut server);
if client_out.connected {
client_connected = true;
}
if server_out.connected {
server_connected = true;
}
if !client_out.packets.is_empty() && client_out.packets.len() != prev_client_count {
client_drop_armed = true;
prev_client_count = client_out.packets.len();
}
if !client_out.packets.is_empty() && client_drop_armed && client_out.packets.len() > 1 {
client_drop_armed = false;
dropped_client += 1;
for p in &client_out.packets[1..] {
let _ = server.handle_packet(p);
}
} else {
deliver_packets(&client_out.packets, &mut server);
}
if !server_out.packets.is_empty() && server_out.packets.len() != prev_server_count {
server_drop_armed = true;
prev_server_count = server_out.packets.len();
}
if !server_out.packets.is_empty() && server_drop_armed && server_out.packets.len() > 1 {
server_drop_armed = false;
dropped_server += 1;
for p in &server_out.packets[1..] {
let _ = client.handle_packet(p);
}
} else {
deliver_packets(&server_out.packets, &mut client);
}
if client_connected && server_connected {
break;
}
if i % 5 == 4 {
now += Duration::from_secs(2);
} else {
now += Duration::from_millis(10);
}
}
assert!(
client_connected,
"Client should connect despite fragmented packet loss"
);
assert!(
server_connected,
"Server should connect despite fragmented packet loss"
);
assert!(
dropped_client > 0,
"Should have dropped at least one client packet"
);
assert!(
dropped_server > 0,
"Should have dropped at least one server packet"
);
}
#[test]
#[cfg(feature = "rcgen")]
fn dtls13_overlapping_fragments_reassembled_successfully() {
use dimpl::certificate::generate_self_signed_certificate;
let _ = env_logger::try_init();
let client_cert = generate_self_signed_certificate().expect("gen client cert");
let server_cert = generate_self_signed_certificate().expect("gen server cert");
let config = dtls13_config_with_mtu(150);
let mut now = Instant::now();
let mut client = Dtls::new_13(Arc::clone(&config), client_cert, now);
client.set_active(true);
let mut server = Dtls::new_13(config, server_cert, now);
server.set_active(false);
client.handle_timeout(now).expect("client timeout");
let client_out = drain_outputs(&mut client);
let handshake_packets: Vec<usize> = client_out
.packets
.iter()
.enumerate()
.filter(|(_, p)| p.len() > 25 && p[0] == 0x16 && p[3] == 0x00 && p[4] == 0x00)
.map(|(i, _)| i)
.collect();
assert!(
handshake_packets.len() >= 2,
"ClientHello should be fragmented into at least 2 packets, got {}",
handshake_packets.len()
);
let mut modified_packets = client_out.packets.clone();
let first_idx = handshake_packets[0];
let second_idx = handshake_packets[1];
let first_frag_len = {
let p = &modified_packets[first_idx];
((p[22] as usize) << 16) | ((p[23] as usize) << 8) | (p[24] as usize)
};
let overlap_data: Vec<u8> = {
let p = &modified_packets[first_idx];
let body_end = 25 + first_frag_len;
p[body_end - 10..body_end].to_vec()
};
let new_packet = {
let packet = &modified_packets[second_idx];
let orig_offset =
((packet[19] as u32) << 16) | ((packet[20] as u32) << 8) | (packet[21] as u32);
assert!(
orig_offset >= 10,
"Second fragment offset should be >= 10 for overlap test, got {}",
orig_offset
);
let orig_frag_len =
((packet[22] as u32) << 16) | ((packet[23] as u32) << 8) | (packet[24] as u32);
let orig_record_len = u16::from_be_bytes([packet[11], packet[12]]);
let new_offset = orig_offset - 10;
let new_frag_len = orig_frag_len + 10;
let new_record_len = orig_record_len + 10;
let mut p = Vec::with_capacity(packet.len() + 10);
p.extend_from_slice(&packet[..19]);
p.push((new_offset >> 16) as u8);
p.push((new_offset >> 8) as u8);
p.push(new_offset as u8);
p.push((new_frag_len >> 16) as u8);
p.push((new_frag_len >> 8) as u8);
p.push(new_frag_len as u8);
p.extend_from_slice(&overlap_data);
p.extend_from_slice(&packet[25..]);
p[11] = (new_record_len >> 8) as u8;
p[12] = new_record_len as u8;
p
};
modified_packets[second_idx] = new_packet;
deliver_packets(&modified_packets, &mut server);
let mut client_connected = false;
let mut server_connected = false;
for _ in 0..40 {
now += Duration::from_millis(10);
client.handle_timeout(now).expect("client timeout");
match server.handle_timeout(now) {
Ok(()) => {}
Err(e) => panic!("Server should not error with overlapping fragments: {}", e),
}
let client_out = drain_outputs(&mut client);
let server_out = drain_outputs(&mut server);
if client_out.connected {
client_connected = true;
}
if server_out.connected {
server_connected = true;
}
deliver_packets(&client_out.packets, &mut server);
deliver_packets(&server_out.packets, &mut client);
if client_connected && server_connected {
break;
}
}
assert!(
server_connected,
"Server should reassemble overlapping fragments and complete the handshake"
);
}