Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 2 additions & 0 deletions CHANGELOG.md
Original file line number Diff line number Diff line change
@@ -1,5 +1,7 @@
# Unreleased

* Start handshake clocks at ClientHello and honor timing settings across DTLS 1.2/1.3 and Auto #161

# 0.7.3

* Fix DTLS 1.2 ClientHello retransmissions #160
Expand Down
17 changes: 14 additions & 3 deletions README.md
Original file line number Diff line number Diff line change
Expand Up @@ -102,12 +102,15 @@ use std::time::Instant;
use dimpl::{certificate, Config, Dtls, Output};

// Stub I/O to keep the example focused on the state machine
enum Event { Udp(Vec<u8>), Timer(Instant) }
fn wait_next_event(_next_wake: Option<Instant>) -> Event { Event::Udp(Vec::new()) }
enum Event { Udp(Vec<u8>, Instant), Timer(Instant) }
fn wait_next_event(_next_wake: Option<Instant>) -> Event {
Event::Udp(Vec::new(), Instant::now())
}
fn send_udp(_bytes: &[u8]) {}

fn example_event_loop(mut dtls: Dtls) -> Result<(), dimpl::Error> {
let mut next_wake: Option<Instant> = None;
let mut received_packet: Option<Vec<u8>> = None;
loop {
// Drain engine output until we have to wait for I/O or a timer
let mut out_buf = vec![0u8; 2048];
Expand Down Expand Up @@ -138,9 +141,17 @@ fn example_event_loop(mut dtls: Dtls) -> Result<(), dimpl::Error> {
}
}

if let Some(packet) = received_packet.take() {
dtls.handle_packet(&packet)?;
continue;
}

// Block waiting for either UDP input or the scheduled timeout
match wait_next_event(next_wake) {
Event::Udp(pkt) => dtls.handle_packet(&pkt)?,
Event::Udp(pkt, now) => {
dtls.handle_timeout(now)?;
received_packet = Some(pkt);
}
Event::Timer(now) => dtls.handle_timeout(now)?,
}
}
Expand Down
140 changes: 111 additions & 29 deletions src/auto.rs
Original file line number Diff line number Diff line change
Expand Up @@ -16,7 +16,7 @@
/// and falls back to DTLS 1.2 via [`Error::Dtls12Fallback`] if the
/// reassembled ClientHello does not offer DTLS 1.3.
use std::sync::Arc;
use std::time::{Duration, Instant};
use std::time::Instant;

use arrayvec::ArrayVec;

Expand All @@ -28,6 +28,7 @@ use crate::dtls13::message::Random;
use crate::dtls13::message::SignatureAlgorithmsExtension;
use crate::dtls13::message::SupportedGroupsExtension;
use crate::dtls13::message::UseSrtpExtension;
use crate::timer::HandshakeTimers;
use crate::types::NamedGroup;
use crate::{Config, CryptoError, DtlsCertificate, Error, Output, SeededRng, TimeoutError};
// Extension type constants
Expand Down Expand Up @@ -265,10 +266,8 @@ pub(crate) struct ClientPending {
needs_send: bool,
/// Last time handle_timeout was called.
last_now: Instant,
/// When to retransmit the wire_packet.
retransmit_at: Option<Instant>,
/// How many retransmits have occurred.
retransmit_count: usize,
timers: HandshakeTimers,
rng: SeededRng,
}

impl ClientPending {
Expand All @@ -279,37 +278,34 @@ impl ClientPending {
) -> Result<Self, Error> {
let hybrid = HybridClientHello::new(&config)?;
let wire_packet = hybrid.wire_packet();
let mut rng = SeededRng::new(config.rng_seed());
let mut timers = HandshakeTimers::new(&config, &mut rng);
timers.begin_flight(&mut rng);
Ok(ClientPending {
hybrid,
config,
certificate,
wire_packet,
needs_send: true,
last_now: now,
retransmit_at: None,
retransmit_count: 0,
timers,
rng,
})
}

pub fn handle_timeout(&mut self, now: Instant) -> Result<(), Error> {
self.last_now = now;
// Arm initial retransmit timer on first call
if self.retransmit_at.is_none() {
self.retransmit_at = Some(now + Duration::from_secs(1));
return Ok(());
}
if let Some(deadline) = self.retransmit_at {
if now >= deadline {
if self.retransmit_count >= self.config.flight_retries() {
return Err(Error::Timeout(TimeoutError::HybridClientHello));
}
self.retransmit_count += 1;
self.needs_send = true;
// Exponential backoff: 2s, 4s, 8s, ...
let shift = self.retransmit_count.min(5) as u32;
let rto = Duration::from_secs(1u64 << shift);
self.retransmit_at = Some(now + rto);
}
let resend = self
.timers
.handle_timeout(now, &mut self.rng)
.map_err(|error| {
Error::Timeout(match error {
TimeoutError::Handshake => TimeoutError::HybridClientHello,
other => other,
})
})?;
if resend {
self.needs_send = true;
}
Ok(())
}
Expand All @@ -323,16 +319,30 @@ impl ClientPending {
}
self.needs_send = false;
buf[..len].copy_from_slice(&self.wire_packet);
self.timers.start_handshake(self.last_now);
self.timers.flight_sent(self.last_now);
return Output::Packet(&buf[..len]);
}
let next = self
.retransmit_at
.unwrap_or(self.last_now + Duration::from_secs(1));
let next = self.timers.poll_timeout(self.last_now);
Output::Timeout(next)
}

pub fn into_parts(self) -> (HybridClientHello, Arc<Config>, DtlsCertificate, Instant) {
(self.hybrid, self.config, self.certificate, self.last_now)
pub fn into_parts(
self,
) -> (
HybridClientHello,
Arc<Config>,
DtlsCertificate,
Instant,
HandshakeTimers,
) {
(
self.hybrid,
self.config,
self.certificate,
self.last_now,
self.timers,
)
}
}

Expand Down Expand Up @@ -450,6 +460,9 @@ fn server_hello_version_inner(packet: &[u8]) -> Option<DetectedVersion> {

#[cfg(test)]
mod tests {
#[cfg(feature = "rcgen")]
use std::time::Duration;

use super::*;
use crate::PskResolver;
use crate::dtls12::message::Dtls12CipherSuite;
Expand Down Expand Up @@ -481,6 +494,75 @@ mod tests {
}
}

#[test]
#[cfg(feature = "rcgen")]
fn timing_expired_deadline_is_fatal_during_either_client_handoff() {
use crate::certificate::generate_self_signed_certificate;
use crate::{Dtls, Inner};

let now = Instant::now();
let budget = Duration::from_millis(10);
let config = Arc::new(
Config::builder()
.dangerously_set_rng_seed(42)
.handshake_timeout(budget)
.build()
.expect("valid config"),
);
let certificate = generate_self_signed_certificate().expect("certificate");
for dtls13 in [false, true] {
let mut pending = ClientPending::new(config.clone(), certificate.clone(), now)
.expect("pending client");
let mut buffer = [0; 2048];
let Output::Packet(hello) = pending.poll_output(&mut buffer) else {
panic!("expected hybrid ClientHello");
};
let hello = hello.to_vec();
assert!(matches!(
pending.poll_output(&mut buffer),
Output::Timeout(_)
));
let mut server = if dtls13 {
Dtls::new_13(config.clone(), certificate.clone(), now)
} else {
Dtls::new_12(config.clone(), certificate.clone(), now)
};
server.handle_timeout(now).expect("server clock");
assert!(matches!(
server.poll_output(&mut buffer),
Output::Timeout(_)
));
server.handle_packet(&hello).expect("accept ClientHello");
let mut response = None;
loop {
match server.poll_output(&mut buffer) {
Output::Packet(packet) => response = Some(packet.to_vec()),
Output::Timeout(_) => break,
Output::BufferTooSmall { .. } => {
panic!("unexpected server output: buffer too small")
}
Output::Connected => panic!("unexpected server output: connected"),
Output::PeerCert(_) => panic!("unexpected server output: peer certificate"),
Output::KeyingMaterial(_, _) => {
panic!("unexpected server output: keying material")
}
Output::ApplicationData(_) => {
panic!("unexpected server output: application data")
}
Output::CloseNotify => panic!("unexpected server output: close notify"),
}
}
pending.last_now = now + budget;
let mut client = Dtls {
inner: Some(Inner::ClientPending(pending)),
};
assert_eq!(
client.handle_packet(&response.expect("server response")),
Err(Error::Timeout(TimeoutError::Connect))
);
}
}

#[test]
fn hello_verify_request_is_dtls12() {
// Minimal HelloVerifyRequest packet
Expand Down
Loading
Loading