use std::mem;
use std::sync::Arc;
use std::sync::atomic::AtomicBool;
use std::time::{Duration, Instant};
use arrayvec::ArrayVec;
use super::queue::{QueueRx, QueueTx};
use crate::buffer::{Buf, BufferPool, TmpBuf};
use crate::crypto::Aad;
use crate::crypto::Cipher;
use crate::crypto::HmacProvider;
use crate::crypto::Nonce;
use crate::crypto::SigningKey;
use crate::crypto::SupportedDtls13CipherSuite;
use crate::crypto::SupportedKxGroup;
use crate::crypto::prf_hkdf;
use crate::dtls13::incoming::{Incoming, Record, RecordHandler};
use crate::dtls13::message::Body;
use crate::dtls13::message::ContentType;
use crate::dtls13::message::Dtls13CipherSuite;
use crate::dtls13::message::Dtls13Record;
use crate::dtls13::message::Handshake;
use crate::dtls13::message::Header;
use crate::dtls13::message::KeyUpdateRequest;
use crate::dtls13::message::MessageType;
use crate::dtls13::message::Sequence;
use crate::timer::ExponentialBackoff;
use crate::types::{HashAlgorithm, Random};
use crate::window::ReplayWindow;
use crate::{Config, DtlsCertificate, Error, InternalError, Output, SeededRng};
const MAX_DEFRAGMENT_PACKETS: usize = 50;
const MAX_SEQUENCE_NUMBER: u64 = (1u64 << 48) - 1;
pub struct Engine {
config: Arc<Config>,
certificate: DtlsCertificate,
rng: SeededRng,
buffers_free: BufferPool,
sequence_epoch_0: Sequence,
queue_rx: QueueRx,
queue_tx: QueueTx,
cipher_suite: Option<Dtls13CipherSuite>,
hs_send_keys: Option<EpochKeys>,
hs_recv_keys: Option<EpochKeys>,
hs_expected_recv_seq: u64,
app_send_epoch: u16,
hs_send_seq: u64,
app_send_seq: u64,
app_send_keys: Option<EpochKeys>,
prev_app_send_keys: Option<EpochKeys>,
prev_app_send_epoch: u16,
prev_app_send_seq: u64,
app_recv_keys: ArrayVec<RecvEpochEntry, 4>,
peer_encryption_enabled: bool,
signing_key: Box<dyn SigningKey>,
is_client: bool,
peer_handshake_seq_no: u16,
next_handshake_seq_no: u16,
pub(crate) transcript: Buf,
hs_replay: ReplayWindow,
received_record_numbers: ArrayVec<(u64, u64), 32>,
handshake_ack_deadline: Option<Instant>,
datagram_sealed: bool,
flight_saved_records: ArrayVec<Entry, 12>,
flight_backoff: ExponentialBackoff,
flight_timeout: Timeout,
connect_timeout: Timeout,
release_app_data: bool,
exporter_master_secret: Option<Buf>,
app_send_record_count: u64,
aead_encryption_threshold: u64,
needs_key_update: bool,
key_update_in_flight: bool,
close_notify_sequence: Option<Sequence>,
close_notify_reported: bool,
}
struct EpochKeys {
cipher: Box<dyn Cipher>,
iv: [u8; 12],
traffic_secret: Buf,
sn_key: Buf,
}
struct RecvEpochEntry {
epoch: u16,
keys: EpochKeys,
expected_recv_seq: u64,
replay: ReplayWindow,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum Timeout {
Disabled,
Unarmed,
Armed(Instant),
}
#[derive(Debug)]
struct Entry {
content_type: ContentType,
epoch: u16,
send_seq: u64,
fragment: Buf,
acked: bool,
}
enum PollOutput<'a> {
Data(&'a [u8]),
BufferTooSmall { needed: usize },
None(&'a mut [u8]),
}
impl Engine {
pub fn new(config: Arc<Config>, certificate: DtlsCertificate) -> Self {
let mut rng = SeededRng::new(config.rng_seed());
let flight_backoff =
ExponentialBackoff::new(config.flight_start_rto(), config.flight_retries(), &mut rng);
let signing_key = config
.crypto_provider()
.key_provider
.load_private_key(&certificate.private_key)
.expect("Failed to load private key");
let aead_encryption_threshold =
jittered_aead_threshold(config.aead_encryption_limit(), &mut rng);
Self {
config,
certificate,
rng,
buffers_free: BufferPool::default(),
sequence_epoch_0: Sequence::new(0),
queue_rx: QueueRx::new(),
queue_tx: QueueTx::new(),
cipher_suite: None,
hs_send_keys: None,
hs_recv_keys: None,
hs_expected_recv_seq: 0,
app_send_epoch: 3,
hs_send_seq: 0,
app_send_seq: 0,
app_send_keys: None,
prev_app_send_keys: None,
prev_app_send_epoch: 0,
prev_app_send_seq: 0,
app_recv_keys: ArrayVec::new(),
peer_encryption_enabled: false,
signing_key,
is_client: false,
peer_handshake_seq_no: 0,
next_handshake_seq_no: 0,
transcript: Buf::new(),
hs_replay: ReplayWindow::new(),
received_record_numbers: ArrayVec::new(),
handshake_ack_deadline: None,
datagram_sealed: false,
flight_saved_records: ArrayVec::new(),
flight_backoff,
flight_timeout: Timeout::Unarmed,
connect_timeout: Timeout::Unarmed,
release_app_data: false,
exporter_master_secret: None,
app_send_record_count: 0,
aead_encryption_threshold,
needs_key_update: false,
key_update_in_flight: false,
close_notify_sequence: None,
close_notify_reported: false,
}
}
pub fn into_fallback(self) -> (Arc<Config>, DtlsCertificate) {
(self.config, self.certificate)
}
pub fn set_client(&mut self, is_client: bool) {
self.is_client = is_client;
}
pub fn inject_hybrid_client_hello(&mut self, transcript_bytes: &[u8]) {
self.transcript.extend_from_slice(transcript_bytes);
self.next_handshake_seq_no = 1;
if self.sequence_epoch_0.sequence_number < MAX_SEQUENCE_NUMBER {
self.sequence_epoch_0.sequence_number += 1;
}
}
pub fn config(&self) -> &Config {
&self.config
}
pub fn cipher_suite(&self) -> Option<Dtls13CipherSuite> {
self.cipher_suite
}
pub fn set_cipher_suite(&mut self, cipher_suite: Dtls13CipherSuite) {
self.cipher_suite = Some(cipher_suite);
}
pub fn app_send_epoch(&self) -> u16 {
self.app_send_epoch
}
pub fn is_key_update_in_flight(&self) -> bool {
self.key_update_in_flight
}
pub fn needs_key_update(&mut self) -> bool {
if self.needs_key_update {
self.needs_key_update = false;
true
} else {
false
}
}
pub fn is_cipher_suite_allowed(&self, suite: Dtls13CipherSuite) -> bool {
self.config
.dtls13_cipher_suites()
.any(|cs| cs.suite() == suite)
}
pub fn certificate_der(&self) -> &[u8] {
&self.certificate.certificate
}
pub fn signing_key(&mut self) -> &mut dyn SigningKey {
&mut *self.signing_key
}
pub fn parse_packet(&mut self, packet: &[u8]) -> Result<(), InternalError> {
let cs = self.cipher_suite;
let incoming = Incoming::parse_packet(packet, self, cs)?;
if let Some(incoming) = incoming {
self.insert_incoming(incoming)?;
}
Ok(())
}
fn insert_incoming(&mut self, incoming: Incoming) -> Result<(), Error> {
if self.queue_rx.len() >= self.config.max_queue_rx() {
warn!(
"Receive queue full (max {}): {:?}",
self.config.max_queue_rx(),
self.queue_rx
);
return Err(Error::ReceiveQueueFull);
}
if incoming.first().first_handshake().is_some() {
self.insert_incoming_handshake(incoming)
} else {
self.insert_incoming_non_handshake(incoming)
}
}
fn insert_incoming_handshake(&mut self, incoming: Incoming) -> Result<(), Error> {
let first_record = incoming.first();
let handshake = first_record
.first_handshake()
.expect("caller ensures handshake");
let key_current = (
handshake.header.message_seq,
handshake.header.fragment_offset,
);
let maybe_dupe_seq = incoming
.records()
.iter()
.filter_map(|r| r.first_handshake())
.filter_map(|h| h.dupe_triggers_resend())
.next();
if let Some(dupe_seq) = maybe_dupe_seq {
if dupe_seq < self.peer_handshake_seq_no {
self.flight_resend("dupe triggers resend")?;
}
}
if handshake.header.message_seq < self.peer_handshake_seq_no {
return Ok(());
}
if self.release_app_data
&& handshake.header.message_seq >= self.peer_handshake_seq_no
&& handshake.header.msg_type != MessageType::KeyUpdate
{
return Err(Error::RenegotiationAttempt);
}
let search_result = self.queue_rx.binary_search_by(|item| {
let key_other = item
.first()
.first_handshake()
.as_ref()
.map(|h| (h.header.message_seq, h.header.fragment_offset))
.unwrap_or((u16::MAX, u32::MAX));
key_other.cmp(&key_current)
});
match search_result {
Err(index) => {
for record in incoming.records().iter() {
let seq = record.record().sequence;
if seq.epoch >= 2 && record.record().content_type == ContentType::Handshake {
let _ = self
.received_record_numbers
.try_push((seq.epoch as u64, seq.sequence_number));
}
}
self.queue_rx.insert(index, incoming);
}
Ok(index) => {
let existing = &self.queue_rx[index];
let should_replace = existing.first().is_handled() || {
let existing_corrupt = existing
.first()
.first_handshake()
.map(|h| h.header.length != h.header.fragment_length)
.unwrap_or(false);
let incoming_ok = incoming
.first()
.first_handshake()
.map(|h| h.header.length == h.header.fragment_length)
.unwrap_or(false);
existing_corrupt && incoming_ok
};
if should_replace {
for record in incoming.records().iter() {
let seq = record.record().sequence;
if seq.epoch >= 2 && record.record().content_type == ContentType::Handshake
{
let _ = self
.received_record_numbers
.try_push((seq.epoch as u64, seq.sequence_number));
}
}
self.queue_rx[index] = incoming;
}
}
}
Ok(())
}
fn insert_incoming_non_handshake(&mut self, incoming: Incoming) -> Result<(), Error> {
let first = incoming.first();
let seq_current = first.record().sequence;
let search_result = self
.queue_rx
.binary_search_by_key(&seq_current, |item| item.first().record().sequence);
match search_result {
Err(index) => self.queue_rx.insert(index, incoming),
Ok(_) => {
}
}
Ok(())
}
pub fn handle_timeout(&mut self, now: Instant) -> Result<(), Error> {
if self.connect_timeout == Timeout::Unarmed {
debug!(
"Connect timeout in: {:.03}s",
self.config.handshake_timeout().as_secs_f32()
);
let timeout = now + self.config.handshake_timeout();
self.connect_timeout = Timeout::Armed(timeout);
}
if self.flight_timeout == Timeout::Unarmed {
debug!(
"Flight timeout in: {:.03}s",
self.flight_backoff.rto().as_secs_f32()
);
let timeout = now + self.flight_backoff.rto();
self.flight_timeout = Timeout::Armed(timeout);
}
if let Timeout::Armed(connect_timeout) = self.connect_timeout {
if now >= connect_timeout {
return Err(Error::Timeout(crate::TimeoutError::Connect));
}
}
let Timeout::Armed(flight_timeout) = self.flight_timeout else {
return Ok(());
};
if now >= flight_timeout {
if self.flight_backoff.can_retry() {
self.flight_backoff.attempt(&mut self.rng);
debug!(
"Re-arm flight timeout due to resend in {}",
self.flight_backoff.rto().as_secs_f32()
);
let timeout = now + self.flight_backoff.rto();
self.flight_timeout = Timeout::Armed(timeout);
self.flight_resend("flight timeout")?;
} else {
return Err(Error::Timeout(crate::TimeoutError::Handshake));
}
}
self.maybe_schedule_handshake_ack(now);
self.maybe_flush_handshake_ack(now)?;
Ok(())
}
pub fn poll_output<'a>(&mut self, buf: &'a mut [u8], now: Instant) -> Output<'a> {
self.purge_handled_queue_rx();
let buf = match self.poll_app_data(buf) {
PollOutput::Data(p) => return Output::ApplicationData(p),
PollOutput::BufferTooSmall { needed } => return Output::BufferTooSmall { needed },
PollOutput::None(b) => b,
};
self.maybe_schedule_handshake_ack(now);
match self.poll_packet_tx(buf) {
PollOutput::Data(p) => return Output::Packet(p),
PollOutput::BufferTooSmall { needed } => return Output::BufferTooSmall { needed },
PollOutput::None(_) => {}
}
if self.close_notify_sequence.is_some() && !self.close_notify_reported {
self.close_notify_reported = true;
return Output::CloseNotify;
}
let next_timeout = self.poll_timeout(now);
Output::Timeout(next_timeout)
}
fn poll_app_data<'a>(&mut self, buf: &'a mut [u8]) -> PollOutput<'a> {
if !self.release_app_data {
return PollOutput::None(buf);
}
let mut unhandled = self
.queue_rx
.iter()
.flat_map(|i| i.records().iter())
.filter(|r| r.record().content_type == ContentType::ApplicationData)
.skip_while(|r| r.is_handled());
let Some(next) = unhandled.next() else {
return PollOutput::None(buf);
};
let record_buffer = next.buffer();
let fragment = next.record().fragment(record_buffer);
let len = fragment.len();
if len > buf.len() {
return PollOutput::BufferTooSmall { needed: len };
}
buf[..len].copy_from_slice(fragment);
next.set_handled();
PollOutput::Data(&buf[..len])
}
fn purge_handled_queue_rx(&mut self) {
while let Some(peek) = self.queue_rx.front() {
let fully_handled = peek.records().iter().all(|r| r.is_handled());
if fully_handled {
let incoming = self.queue_rx.pop_front().unwrap();
incoming
.into_records()
.for_each(|r| self.buffers_free.push(r.into_buffer()));
} else {
break;
}
}
}
fn poll_packet_tx<'a>(&mut self, buf: &'a mut [u8]) -> PollOutput<'a> {
let Some(p) = self.queue_tx.front() else {
return PollOutput::None(buf);
};
if p.len() > buf.len() {
return PollOutput::BufferTooSmall { needed: p.len() };
}
let p = self
.queue_tx
.pop_front()
.expect("queue front checked above");
let len = p.len();
buf[..len].copy_from_slice(&p);
PollOutput::Data(&buf[..len])
}
fn seal_current_datagram(&mut self) {
self.datagram_sealed = true;
}
fn poll_timeout(&self, now: Instant) -> Instant {
if self.connect_timeout == Timeout::Disabled
&& self.flight_timeout == Timeout::Disabled
&& self.handshake_ack_deadline.is_none()
{
const DISTANT_FUTURE: Duration = Duration::from_secs(10 * 365 * 24 * 60 * 60);
return now + DISTANT_FUTURE;
}
let mut timeout = match (self.connect_timeout, self.flight_timeout) {
(Timeout::Armed(c), Timeout::Armed(f)) => {
if c < f {
c
} else {
f
}
}
(Timeout::Armed(c), _) => c,
(_, Timeout::Armed(f)) => f,
_ => now + Duration::from_secs(10 * 365 * 24 * 60 * 60),
};
if let Some(deadline) = self.handshake_ack_deadline {
if deadline < timeout {
timeout = deadline;
}
}
timeout
}
pub fn flight_begin(&mut self, flight_no: u8) {
debug!("Begin flight {}", flight_no);
self.flight_backoff.reset(&mut self.rng);
self.flight_clear_resends();
self.flight_timeout = Timeout::Unarmed;
}
pub fn flight_stop_resend_timers(&mut self) {
debug!("Stop connect and flight timeouts");
self.flight_timeout = Timeout::Disabled;
self.connect_timeout = Timeout::Disabled;
}
fn flight_clear_resends(&mut self) {
for entry in self.flight_saved_records.drain(..) {
self.buffers_free.push(entry.fragment);
}
}
pub fn flight_resend(&mut self, reason: &str) -> Result<(), Error> {
debug!("Resending flight due to {}", reason);
let mut records = mem::take(&mut self.flight_saved_records);
self.seal_current_datagram();
for entry in &mut records {
if entry.acked {
continue;
}
if entry.epoch == 0 {
let new_seq = self.sequence_epoch_0.sequence_number;
self.create_plaintext_record(entry.content_type, false, |fragment| {
fragment.extend_from_slice(&entry.fragment);
})?;
entry.send_seq = new_seq;
} else {
let new_seq = if entry.epoch == 2 {
self.hs_send_seq
} else if self.prev_app_send_keys.is_some()
&& entry.epoch == self.prev_app_send_epoch
{
self.prev_app_send_seq
} else {
self.app_send_seq
};
self.create_ciphertext_record(
entry.content_type,
entry.epoch,
false,
|fragment| {
fragment.extend_from_slice(&entry.fragment);
},
)?;
entry.send_seq = new_seq;
}
}
self.flight_saved_records = records;
Ok(())
}
pub fn has_complete_handshake(&mut self, wanted: MessageType) -> bool {
self.has_complete_handshake_with_seq(wanted, self.peer_handshake_seq_no)
}
fn has_complete_handshake_with_seq(&mut self, wanted: MessageType, expected_seq: u16) -> bool {
let mut skip_handled = self
.queue_rx
.iter()
.flat_map(|i| i.records().iter())
.skip_while(|r| r.is_handled())
.take(MAX_DEFRAGMENT_PACKETS)
.flat_map(|r| r.handshakes().iter())
.skip_while(|h| h.is_handled())
.peekable();
let maybe_first_handshake = skip_handled.peek();
let Some(first) = maybe_first_handshake else {
return false;
};
if first.header.message_seq != expected_seq {
return false;
}
if first.header.msg_type != wanted {
return false;
}
let wanted_seq = first.header.message_seq;
let wanted_length = first.header.length;
let mut last_fragment_end = 0;
for h in skip_handled {
if wanted_seq != h.header.message_seq {
continue;
}
if h.header.fragment_offset > last_fragment_end {
return false;
}
let end = h.header.fragment_offset + h.header.fragment_length;
if end > last_fragment_end {
last_fragment_end = end;
}
if last_fragment_end == wanted_length {
return true;
}
}
false
}
pub fn next_handshake(
&mut self,
wanted: MessageType,
defragment_buffer: &mut Buf,
) -> Result<Option<Handshake>, InternalError> {
self.next_handshake_with_options(wanted, defragment_buffer, false)
}
pub(crate) fn next_client_hello_for_auto_sense(
&mut self,
defragment_buffer: &mut Buf,
) -> Result<Option<Handshake>, InternalError> {
self.next_handshake_with_options(MessageType::ClientHello, defragment_buffer, true)
}
fn next_handshake_with_options(
&mut self,
wanted: MessageType,
defragment_buffer: &mut Buf,
allow_unknown_client_hello_suites: bool,
) -> Result<Option<Handshake>, InternalError> {
if !self.has_complete_handshake(wanted) {
return Ok(None);
}
let iter = self
.queue_rx
.iter()
.flat_map(|i| i.records().iter())
.skip_while(|r| r.is_handled())
.flat_map(|r| r.handshakes().iter().map(move |h| (h, r.buffer())))
.skip_while(|(h, _)| h.is_handled());
let handshake = if allow_unknown_client_hello_suites {
Handshake::defragment_allow_unknown_client_hello_suites(
iter,
defragment_buffer,
self.cipher_suite,
Some(&mut self.transcript),
)
} else {
Handshake::defragment(
iter,
defragment_buffer,
self.cipher_suite,
Some(&mut self.transcript),
)
}?;
Ok(Some(handshake))
}
pub fn next_handshake_no_transcript(
&mut self,
wanted: MessageType,
defragment_buffer: &mut Buf,
) -> Result<Option<Handshake>, InternalError> {
if !self.has_complete_handshake(wanted) {
return Ok(None);
}
let iter = self
.queue_rx
.iter()
.flat_map(|i| i.records().iter())
.skip_while(|r| r.is_handled())
.flat_map(|r| r.handshakes().iter().map(move |h| (h, r.buffer())))
.skip_while(|(h, _)| h.is_handled());
let handshake = Handshake::defragment(
iter,
defragment_buffer,
self.cipher_suite,
None, )?;
Ok(Some(handshake))
}
pub fn advance_peer_handshake_seq(&mut self) {
self.peer_handshake_seq_no += 1;
}
pub fn create_plaintext_record<F>(
&mut self,
content_type: ContentType,
save_fragment: bool,
f: F,
) -> Result<(), Error>
where
F: FnOnce(&mut Buf),
{
let mut fragment = self.buffers_free.pop();
f(&mut fragment);
let current_seq = self.sequence_epoch_0.sequence_number;
if save_fragment {
let mut clone = self.buffers_free.pop();
clone.extend_from_slice(&fragment);
self.flight_saved_records.push(Entry {
content_type,
epoch: 0,
send_seq: current_seq,
fragment: clone,
acked: false,
});
}
let record_wire_len = Dtls13Record::PLAINTEXT_HEADER_LEN + fragment.len();
let can_append = !self.datagram_sealed
&& self
.queue_tx
.back()
.map(|b| b.len() + record_wire_len <= self.config.mtu())
.unwrap_or(false);
if !can_append && self.queue_tx.len() >= self.config.max_queue_tx() {
warn!(
"Transmit queue full (max {}): {:?}",
self.config.max_queue_tx(),
self.queue_tx
);
return Err(Error::TransmitQueueFull);
}
let sequence = self.sequence_epoch_0;
let record = Dtls13Record {
content_type,
sequence,
length: fragment.len() as u16,
fragment_range: 0..fragment.len(),
};
if self.sequence_epoch_0.sequence_number >= MAX_SEQUENCE_NUMBER {
return Err(Error::CryptoError(
crate::CryptoError::Epoch0SequenceNumberExhausted,
));
}
self.sequence_epoch_0.sequence_number += 1;
if can_append {
let last = self.queue_tx.back_mut().unwrap();
record.serialize(&fragment, last);
} else {
self.datagram_sealed = false;
let mut buffer = self.buffers_free.pop();
buffer.clear();
record.serialize(&fragment, &mut buffer);
self.queue_tx.push_back(buffer);
}
self.buffers_free.push(fragment);
Ok(())
}
pub fn create_ciphertext_record<F>(
&mut self,
content_type: ContentType,
epoch: u16,
save_fragment: bool,
f: F,
) -> Result<(), Error>
where
F: FnOnce(&mut Buf),
{
let mut fragment = self.buffers_free.pop();
f(&mut fragment);
let seq = if epoch == 2 {
self.hs_send_seq
} else if self.prev_app_send_keys.is_some() && epoch == self.prev_app_send_epoch {
self.prev_app_send_seq
} else {
self.app_send_seq
};
if save_fragment {
let mut clone = self.buffers_free.pop();
clone.extend_from_slice(&fragment);
self.flight_saved_records.push(Entry {
content_type,
epoch,
send_seq: seq,
fragment: clone,
acked: false,
});
}
fragment.push(content_type.as_u8());
let suite = self.suite_provider();
let tag_len = suite.tag_len();
let keys = if epoch == 2 {
self.hs_send_keys.as_mut()
} else if self.prev_app_send_keys.is_some() && epoch == self.prev_app_send_epoch {
self.prev_app_send_keys.as_mut()
} else {
self.app_send_keys.as_mut()
};
let Some(keys) = keys else {
return Err(Error::CryptoError(
crate::CryptoError::SendKeysNotAvailable { epoch },
));
};
let nonce = Nonce::xor(&keys.iv, seq);
let epoch_bits = (epoch & 0x03) as u8;
let flags: u8 = 0b0010_0000
| 0b0000_1000 | 0b0000_0100 | epoch_bits;
let ciphertext_len = fragment.len() + tag_len;
let mut header_buf = [0u8; 5];
header_buf[0] = flags;
header_buf[1..3].copy_from_slice(&(seq as u16).to_be_bytes());
header_buf[3..5].copy_from_slice(&(ciphertext_len as u16).to_be_bytes());
let aad = Aad::new_dtls13(&header_buf);
let mut sn_key = [0u8; 32];
let sn_key_len = keys.sn_key.len();
sn_key[..sn_key_len].copy_from_slice(&keys.sn_key);
keys.cipher
.encrypt(&mut fragment, aad, nonce)
.map_err(Error::CryptoError)?;
let sn_mask = if fragment.len() >= 16 {
let ciphertext_sample: [u8; 16] = fragment[..16].try_into().unwrap();
suite.encrypt_sn(&sn_key[..sn_key_len], &ciphertext_sample)
} else {
[0u8; 16] };
let record_wire_len = 5 + fragment.len();
let can_append = !self.datagram_sealed
&& self
.queue_tx
.back()
.map(|b| b.len() + record_wire_len <= self.config.mtu())
.unwrap_or(false);
if !can_append && self.queue_tx.len() >= self.config.max_queue_tx() {
warn!(
"Transmit queue full (max {}): {:?}",
self.config.max_queue_tx(),
self.queue_tx
);
return Err(Error::TransmitQueueFull);
}
let record = Dtls13Record {
content_type: ContentType::ApplicationData,
sequence: Sequence {
epoch,
sequence_number: seq,
},
length: fragment.len() as u16,
fragment_range: 0..fragment.len(),
};
if epoch == 2 {
if self.hs_send_seq >= MAX_SEQUENCE_NUMBER {
return Err(Error::CryptoError(
crate::CryptoError::SendSequenceNumberExhausted { epoch },
));
}
self.hs_send_seq += 1;
} else if self.prev_app_send_keys.is_some() && epoch == self.prev_app_send_epoch {
if self.prev_app_send_seq >= MAX_SEQUENCE_NUMBER {
return Err(Error::CryptoError(
crate::CryptoError::SendSequenceNumberExhausted { epoch },
));
}
self.prev_app_send_seq += 1;
} else {
if self.app_send_seq >= MAX_SEQUENCE_NUMBER {
return Err(Error::CryptoError(
crate::CryptoError::SendSequenceNumberExhausted { epoch },
));
}
self.app_send_seq += 1;
if epoch >= 3 {
self.app_send_record_count += 1;
if self.app_send_record_count >= self.aead_encryption_threshold {
self.needs_key_update = true;
}
}
}
if can_append {
let last = self.queue_tx.back_mut().unwrap();
let header_start = last.len();
record.serialize(&fragment, last);
last[header_start + 1] ^= sn_mask[0];
last[header_start + 2] ^= sn_mask[1];
} else {
self.datagram_sealed = false;
let mut buffer = self.buffers_free.pop();
buffer.clear();
record.serialize(&fragment, &mut buffer);
buffer[1] ^= sn_mask[0];
buffer[2] ^= sn_mask[1];
self.queue_tx.push_back(buffer);
}
self.buffers_free.push(fragment);
Ok(())
}
pub fn create_handshake<F>(&mut self, msg_type: MessageType, f: F) -> Result<(), Error>
where
F: FnOnce(&mut Buf, &mut Self) -> Result<(), Error>,
{
let mut body_buffer = self.buffers_free.pop();
f(&mut body_buffer, self)?;
let handshake_header = Header {
msg_type,
length: body_buffer.len() as u32,
message_seq: self.next_handshake_seq_no,
fragment_offset: 0,
fragment_length: body_buffer.len() as u32,
};
self.transcript.push(msg_type.as_u8());
self.transcript
.extend_from_slice(&handshake_header.length.to_be_bytes()[1..]);
self.transcript
.extend_from_slice(&body_buffer[..handshake_header.length as usize]);
self.next_handshake_seq_no += 1;
let epoch = epoch_for_message(msg_type);
let total_len = body_buffer.len();
let mut offset: usize = 0;
let handshake_header_len = 12usize;
let tag_len = if epoch >= 2 {
self.suite_provider().tag_len()
} else {
0
};
let protection_overhead = if epoch >= 2 { tag_len + 1 } else { 0 };
while offset < total_len || (total_len == 0 && offset == 0) {
let already_used_in_current = self.queue_tx.back().map(|b| b.len()).unwrap_or(0);
let available_in_current = self.config.mtu().saturating_sub(already_used_in_current);
let record_header_len = if epoch == 0 {
Dtls13Record::PLAINTEXT_HEADER_LEN
} else {
5 };
let fixed_overhead = record_header_len + handshake_header_len + protection_overhead;
let available_for_body = if available_in_current > fixed_overhead {
available_in_current - fixed_overhead
} else {
self.config.mtu().saturating_sub(fixed_overhead)
};
let remaining_body_bytes = total_len.saturating_sub(offset);
let chunk_len = if total_len == 0 {
0
} else {
remaining_body_bytes.min(available_for_body)
};
let frag_range = if chunk_len == 0 {
0..0
} else {
offset..offset + chunk_len
};
let frag_handshake = Handshake {
header: Header {
msg_type,
length: handshake_header.length,
message_seq: handshake_header.message_seq,
fragment_offset: offset as u32,
fragment_length: chunk_len as u32,
},
body: Body::Fragment(frag_range),
handled: AtomicBool::new(false),
};
if epoch == 0 {
self.create_plaintext_record(ContentType::Handshake, true, |fragment| {
frag_handshake.serialize(&body_buffer, fragment);
})?;
} else {
self.create_ciphertext_record(ContentType::Handshake, epoch, true, |fragment| {
frag_handshake.serialize(&body_buffer, fragment);
})?;
}
if total_len == 0 {
break;
}
offset += chunk_len;
}
self.buffers_free.push(body_buffer);
Ok(())
}
pub fn release_application_data(&mut self) {
self.release_app_data = true;
self.hs_recv_keys = None;
}
pub fn release_application_data_retaining_handshake_keys(&mut self) {
self.release_app_data = true;
}
pub fn close_notify_received(&self) -> bool {
self.close_notify_sequence.is_some()
}
pub fn cancel_flights(&mut self) {
self.flight_saved_records.clear();
self.flight_timeout = Timeout::Disabled;
self.connect_timeout = Timeout::Disabled;
self.handshake_ack_deadline = None;
}
pub fn abort(&mut self) {
self.queue_tx.clear();
self.flight_saved_records.clear();
self.flight_timeout = Timeout::Disabled;
self.connect_timeout = Timeout::Disabled;
self.handshake_ack_deadline = None;
}
pub fn send_ack(&mut self) -> Result<(), Error> {
self.send_ack_inner(false)
}
pub fn send_ack_retransmittable(&mut self) -> Result<(), Error> {
if !self.received_record_numbers.is_empty() {
self.flight_clear_resends();
}
self.send_ack_inner(true)
}
fn send_ack_inner(&mut self, save_fragment: bool) -> Result<(), Error> {
if self.received_record_numbers.is_empty() {
return Ok(());
}
let entries = mem::take(&mut self.received_record_numbers);
let epoch = if self.app_send_keys.is_some() {
self.app_send_epoch
} else {
2
};
self.create_ciphertext_record(ContentType::Ack, epoch, save_fragment, |fragment| {
let len = (entries.len() * 16) as u16;
fragment.extend_from_slice(&len.to_be_bytes());
for &(ep, seq) in &entries {
fragment.extend_from_slice(&ep.to_be_bytes());
fragment.extend_from_slice(&seq.to_be_bytes());
}
})?;
Ok(())
}
fn process_ack(&mut self, ack_data: &[u8]) -> Result<(), Error> {
if ack_data.len() < 2 {
return Ok(());
}
let record_numbers_len = u16::from_be_bytes([ack_data[0], ack_data[1]]) as usize;
let entries_data = &ack_data[2..];
if entries_data.len() != record_numbers_len || record_numbers_len % 16 != 0 {
return Ok(());
}
let num_entries = record_numbers_len / 16;
for i in 0..num_entries {
let offset = i * 16;
if offset + 16 > entries_data.len() {
break;
}
let ack_epoch = u64::from_be_bytes(
entries_data[offset..offset + 8].try_into().unwrap(),
);
let ack_seq = u64::from_be_bytes(
entries_data[offset + 8..offset + 16].try_into().unwrap(),
);
for entry in &mut self.flight_saved_records {
if entry.epoch as u64 == ack_epoch && entry.send_seq == ack_seq {
entry.acked = true;
}
}
}
let has_epoch2 = self.flight_saved_records.iter().any(|e| e.epoch == 2);
let all_epoch2_acked = self
.flight_saved_records
.iter()
.filter(|e| e.epoch == 2)
.all(|e| e.acked);
if has_epoch2 && all_epoch2_acked {
debug!("Handshake flight ACKed; stopping retransmission");
self.flight_timeout = Timeout::Disabled;
self.flight_clear_resends();
}
if self.key_update_in_flight
&& !self.flight_saved_records.is_empty()
&& self.flight_saved_records.iter().all(|e| e.acked)
{
debug!("KeyUpdate ACKed; rotating send keys");
self.update_send_keys()?;
self.prev_app_send_keys = None;
self.key_update_in_flight = false;
self.flight_clear_resends();
self.flight_timeout = Timeout::Disabled;
}
Ok(())
}
fn handshake_in_progress(&self) -> bool {
self.hs_send_keys.is_some() && self.app_send_keys.is_none() && !self.release_app_data
}
fn handshake_ack_help_needed(&self) -> bool {
let mut skip_handled = self
.queue_rx
.iter()
.flat_map(|i| i.records().iter())
.skip_while(|r| r.is_handled())
.take(MAX_DEFRAGMENT_PACKETS)
.flat_map(|r| r.handshakes().iter())
.skip_while(|h| h.is_handled())
.peekable();
let Some(first) = skip_handled.peek() else {
return false;
};
if first.header.message_seq != self.peer_handshake_seq_no {
return true;
}
let wanted_seq = first.header.message_seq;
let wanted_length = first.header.length;
let mut last_fragment_end = 0;
for h in skip_handled {
if wanted_seq != h.header.message_seq {
continue;
}
if h.header.fragment_offset > last_fragment_end {
return true;
}
let end = h.header.fragment_offset + h.header.fragment_length;
if end > last_fragment_end {
last_fragment_end = end;
}
if last_fragment_end == wanted_length {
return false;
}
}
true
}
fn has_gap_in_incoming_handshake(&self) -> bool {
let mut skip_handled = self
.queue_rx
.iter()
.flat_map(|i| i.records().iter())
.skip_while(|r| r.is_handled())
.take(MAX_DEFRAGMENT_PACKETS)
.flat_map(|r| r.handshakes().iter())
.skip_while(|h| h.is_handled())
.peekable();
let Some(first) = skip_handled.peek() else {
return false;
};
if first.header.message_seq != self.peer_handshake_seq_no {
return true;
}
let wanted_seq = first.header.message_seq;
let wanted_length = first.header.length;
let mut last_fragment_end = 0;
for h in skip_handled {
if wanted_seq != h.header.message_seq {
continue;
}
if h.header.fragment_offset > last_fragment_end {
return true;
}
let end = h.header.fragment_offset + h.header.fragment_length;
if end > last_fragment_end {
last_fragment_end = end;
}
if last_fragment_end == wanted_length {
return false;
}
}
false
}
fn maybe_schedule_handshake_ack(&mut self, now: Instant) {
if !self.handshake_in_progress() {
self.handshake_ack_deadline = None;
return;
}
if self.handshake_ack_deadline.is_some() {
return;
}
if !self.handshake_ack_help_needed() {
return;
}
let delay = if self.has_gap_in_incoming_handshake() {
Duration::from_millis(0)
} else {
let rto = self.flight_backoff.rto();
if rto > Duration::from_millis(0) {
rto / 4
} else {
Duration::from_millis(0)
}
};
self.handshake_ack_deadline = Some(now + delay);
}
fn maybe_flush_handshake_ack(&mut self, now: Instant) -> Result<(), Error> {
let Some(deadline) = self.handshake_ack_deadline else {
return Ok(());
};
if now < deadline {
return Ok(());
}
let mut record_numbers = ArrayVec::<(u64, u64), 32>::new();
for incoming in self.queue_rx.iter() {
for r in incoming.records().iter() {
if r.record().sequence.epoch == 2
&& r.record().content_type == ContentType::Handshake
{
let seq = r.record().sequence;
let _ = record_numbers.try_push((seq.epoch as u64, seq.sequence_number));
}
}
}
self.handshake_ack_deadline = None;
if record_numbers.is_empty() {
return Ok(());
}
self.send_handshake_ack_epoch2(&record_numbers)
}
fn send_handshake_ack_epoch2(&mut self, record_numbers: &[(u64, u64)]) -> Result<(), Error> {
if !self.handshake_in_progress() {
return Ok(());
}
self.create_ciphertext_record(ContentType::Ack, 2, false, |fragment| {
let len = (record_numbers.len() * 16) as u16;
fragment.extend_from_slice(&len.to_be_bytes());
for &(epoch, seq) in record_numbers {
fragment.extend_from_slice(&epoch.to_be_bytes());
fragment.extend_from_slice(&seq.to_be_bytes());
}
})?;
Ok(())
}
pub(crate) fn pop_buffer(&mut self) -> Buf {
self.buffers_free.pop()
}
pub(crate) fn push_buffer(&mut self, buf: Buf) {
self.buffers_free.push(buf);
}
fn hmac(&self) -> &dyn HmacProvider {
self.config.crypto_provider().hmac_provider
}
fn hash_algorithm(&self) -> HashAlgorithm {
self.cipher_suite.unwrap().hash_algorithm()
}
fn suite_provider(&self) -> &'static dyn SupportedDtls13CipherSuite {
let suite = self.cipher_suite.unwrap();
*self
.config
.crypto_provider()
.dtls13_cipher_suites
.iter()
.find(|cs| cs.suite() == suite)
.expect("cipher suite not found in provider")
}
pub fn derive_early_secret(&mut self) -> Result<Buf, Error> {
let hash = self.hash_algorithm();
let hash_len = hash.output_len();
let zeros = [0u8; 48];
let zeros = &zeros[..hash_len];
let mut early_secret = self.buffers_free.pop();
prf_hkdf::hkdf_extract(self.hmac(), hash, zeros, zeros, &mut early_secret)
.map_err(Error::CryptoError)?;
Ok(early_secret)
}
pub fn derive_handshake_secrets(
&mut self,
shared_secret: &[u8],
) -> Result<(Buf, Buf, Buf), Error> {
let early_secret = self.derive_early_secret()?;
let hash = self.hash_algorithm();
let hash_len = hash.output_len();
let hmac = self.hmac();
let empty_hash = self.transcript_hash_of(b"");
let mut derived = Buf::new();
prf_hkdf::hkdf_expand_label_dtls13(
hmac,
hash,
&early_secret,
b"derived",
&empty_hash,
&mut derived,
hash_len,
)
.map_err(Error::CryptoError)?;
let mut handshake_secret = Buf::new();
prf_hkdf::hkdf_extract(hmac, hash, &derived, shared_secret, &mut handshake_secret)
.map_err(Error::CryptoError)?;
let mut transcript_hash = Buf::new();
self.transcript_hash(&mut transcript_hash);
let mut c_hs_traffic = Buf::new();
prf_hkdf::hkdf_expand_label_dtls13(
hmac,
hash,
&handshake_secret,
b"c hs traffic",
&transcript_hash,
&mut c_hs_traffic,
hash_len,
)
.map_err(Error::CryptoError)?;
let mut s_hs_traffic = Buf::new();
prf_hkdf::hkdf_expand_label_dtls13(
hmac,
hash,
&handshake_secret,
b"s hs traffic",
&transcript_hash,
&mut s_hs_traffic,
hash_len,
)
.map_err(Error::CryptoError)?;
Ok((c_hs_traffic, s_hs_traffic, handshake_secret))
}
pub fn install_handshake_keys(
&mut self,
client_traffic_secret: &Buf,
server_traffic_secret: &Buf,
) -> Result<(), Error> {
let (send_secret, recv_secret) = if self.is_client {
(client_traffic_secret, server_traffic_secret)
} else {
(server_traffic_secret, client_traffic_secret)
};
self.hs_send_keys = Some(self.derive_epoch_keys(send_secret)?);
self.hs_recv_keys = Some(self.derive_epoch_keys(recv_secret)?);
self.hs_send_seq = 0;
Ok(())
}
pub fn derive_application_secrets(
&mut self,
handshake_secret: &[u8],
) -> Result<(Buf, Buf), Error> {
let hash = self.hash_algorithm();
let hash_len = hash.output_len();
let hmac = self.hmac();
let empty_hash = self.transcript_hash_of(b"");
let mut derived = Buf::new();
prf_hkdf::hkdf_expand_label_dtls13(
hmac,
hash,
handshake_secret,
b"derived",
&empty_hash,
&mut derived,
hash_len,
)
.map_err(Error::CryptoError)?;
let zeros = [0u8; 48];
let zeros = &zeros[..hash_len];
let mut master_secret = Buf::new();
prf_hkdf::hkdf_extract(hmac, hash, &derived, zeros, &mut master_secret)
.map_err(Error::CryptoError)?;
let mut transcript_hash = Buf::new();
self.transcript_hash(&mut transcript_hash);
let mut exp_master = Buf::new();
prf_hkdf::hkdf_expand_label_dtls13(
hmac,
hash,
&master_secret,
b"exp master",
&transcript_hash,
&mut exp_master,
hash_len,
)
.map_err(Error::CryptoError)?;
let mut c_ap_traffic = Buf::new();
prf_hkdf::hkdf_expand_label_dtls13(
hmac,
hash,
&master_secret,
b"c ap traffic",
&transcript_hash,
&mut c_ap_traffic,
hash_len,
)
.map_err(Error::CryptoError)?;
let mut s_ap_traffic = Buf::new();
prf_hkdf::hkdf_expand_label_dtls13(
hmac,
hash,
&master_secret,
b"s ap traffic",
&transcript_hash,
&mut s_ap_traffic,
hash_len,
)
.map_err(Error::CryptoError)?;
self.exporter_master_secret = Some(exp_master);
Ok((c_ap_traffic, s_ap_traffic))
}
pub fn install_application_keys(
&mut self,
client_traffic_secret: &Buf,
server_traffic_secret: &Buf,
) -> Result<(), Error> {
let (send_secret, recv_secret) = if self.is_client {
(client_traffic_secret, server_traffic_secret)
} else {
(server_traffic_secret, client_traffic_secret)
};
self.app_send_keys = Some(self.derive_epoch_keys(send_secret)?);
let recv_keys = self.derive_epoch_keys(recv_secret)?;
self.app_recv_keys.push(RecvEpochEntry {
epoch: 3,
keys: recv_keys,
expected_recv_seq: 0,
replay: ReplayWindow::new(),
});
self.app_send_epoch = 3;
self.app_send_seq = 0;
Ok(())
}
fn derive_next_traffic_secret(&self, current: &Buf) -> Result<Buf, Error> {
let hash = self.hash_algorithm();
let hash_len = hash.output_len();
let hmac = self.hmac();
let mut next = Buf::new();
prf_hkdf::hkdf_expand_label_dtls13(
hmac,
hash,
current,
b"traffic upd",
&[],
&mut next,
hash_len,
)
.map_err(Error::CryptoError)?;
Ok(next)
}
fn update_send_keys(&mut self) -> Result<(), Error> {
let current_keys = self.app_send_keys.take().ok_or(Error::InvalidState(
crate::InvalidStateError::NoCurrentAppSendKeysForKeyUpdate,
))?;
let next_secret = self.derive_next_traffic_secret(¤t_keys.traffic_secret)?;
let new_keys = self.derive_epoch_keys(&next_secret)?;
self.prev_app_send_keys = Some(current_keys);
self.prev_app_send_epoch = self.app_send_epoch;
self.prev_app_send_seq = self.app_send_seq;
self.app_send_keys = Some(new_keys);
self.app_send_epoch += 1;
self.app_send_seq = 0;
self.app_send_record_count = 0;
self.aead_encryption_threshold =
jittered_aead_threshold(self.config.aead_encryption_limit(), &mut self.rng);
self.needs_key_update = false;
debug!("Send keys updated to epoch {}", self.app_send_epoch);
Ok(())
}
pub fn update_recv_keys(&mut self) -> Result<u16, Error> {
let latest = self.app_recv_keys.last().ok_or(Error::InvalidState(
crate::InvalidStateError::NoCurrentAppRecvKeysForKeyUpdate,
))?;
let next_secret = self.derive_next_traffic_secret(&latest.keys.traffic_secret)?;
let new_epoch = latest.epoch + 1;
let new_keys = self.derive_epoch_keys(&next_secret)?;
if self.app_recv_keys.is_full() {
self.app_recv_keys.remove(0);
}
self.app_recv_keys.push(RecvEpochEntry {
epoch: new_epoch,
keys: new_keys,
expected_recv_seq: 0,
replay: ReplayWindow::new(),
});
debug!("Recv keys updated to epoch {}", new_epoch);
Ok(new_epoch)
}
pub fn create_key_update(&mut self, request: KeyUpdateRequest) -> Result<(), Error> {
self.flight_backoff.reset(&mut self.rng);
self.flight_clear_resends();
self.flight_timeout = Timeout::Unarmed;
let msg_seq = self.next_handshake_seq_no;
self.next_handshake_seq_no += 1;
let epoch = self.app_send_epoch;
self.create_ciphertext_record(ContentType::Handshake, epoch, true, |fragment| {
fragment.push(MessageType::KeyUpdate.as_u8());
fragment.extend_from_slice(&1u32.to_be_bytes()[1..]); fragment.extend_from_slice(&msg_seq.to_be_bytes()); fragment.extend_from_slice(&0u32.to_be_bytes()[1..]); fragment.extend_from_slice(&1u32.to_be_bytes()[1..]); fragment.push(request.as_u8());
})?;
self.key_update_in_flight = true;
debug!(
"KeyUpdate sent (request={:?}) on epoch {}, awaiting ACK before rotating send keys",
request, epoch
);
Ok(())
}
pub fn reset_for_hello_retry(&mut self) {
self.hs_send_seq = 0;
self.handshake_ack_deadline = None;
self.queue_rx.retain(|item| !item.first().is_handled());
}
fn derive_epoch_keys(&self, traffic_secret: &Buf) -> Result<EpochKeys, Error> {
let hash = self.hash_algorithm();
let suite = self.suite_provider();
let hmac = self.hmac();
let mut key = Buf::new();
prf_hkdf::hkdf_expand_label_dtls13(
hmac,
hash,
traffic_secret,
b"key",
&[],
&mut key,
suite.key_len(),
)
.map_err(Error::CryptoError)?;
let mut iv_buf = Buf::new();
prf_hkdf::hkdf_expand_label_dtls13(
hmac,
hash,
traffic_secret,
b"iv",
&[],
&mut iv_buf,
suite.iv_len(),
)
.map_err(Error::CryptoError)?;
let mut sn_key = Buf::new();
prf_hkdf::hkdf_expand_label_dtls13(
hmac,
hash,
traffic_secret,
b"sn",
&[],
&mut sn_key,
suite.key_len(),
)
.map_err(Error::CryptoError)?;
let cipher = suite.create_cipher(&key).map_err(Error::CryptoError)?;
let mut iv = [0u8; 12];
iv.copy_from_slice(&iv_buf);
let mut secret = Buf::new();
secret.extend_from_slice(traffic_secret);
Ok(EpochKeys {
cipher,
iv,
traffic_secret: secret,
sn_key,
})
}
pub fn compute_verify_data(&self, traffic_secret: &[u8]) -> Result<Buf, Error> {
let hash = self.hash_algorithm();
let hash_len = hash.output_len();
let hmac = self.hmac();
let mut finished_key = Buf::new();
prf_hkdf::hkdf_expand_label_dtls13(
hmac,
hash,
traffic_secret,
b"finished",
&[],
&mut finished_key,
hash_len,
)
.map_err(Error::CryptoError)?;
let mut transcript_hash = Buf::new();
self.transcript_hash(&mut transcript_hash);
let mut verify_data = Buf::new();
prf_hkdf::hkdf_extract(
hmac,
hash,
&finished_key,
&transcript_hash,
&mut verify_data,
)
.map_err(Error::CryptoError)?;
Ok(verify_data)
}
pub fn transcript_hash(&self, out: &mut Buf) {
let hash = self.hash_algorithm();
let mut ctx = self
.config
.crypto_provider()
.hash_provider
.create_hash(hash);
ctx.update(&self.transcript);
ctx.clone_and_finalize(out);
}
fn transcript_hash_of(&self, data: &[u8]) -> Buf {
let hash = self.hash_algorithm();
let mut ctx = self
.config
.crypto_provider()
.hash_provider
.create_hash(hash);
ctx.update(data);
let mut out = Buf::new();
ctx.clone_and_finalize(&mut out);
out
}
pub fn replace_transcript_with_message_hash(&mut self, split_at: usize) {
let hash = self.hash_algorithm();
let mut hash_ctx = self
.config
.crypto_provider()
.hash_provider
.create_hash(hash);
hash_ctx.update(&self.transcript[..split_at]);
let mut hash_value = Buf::new();
hash_ctx.clone_and_finalize(&mut hash_value);
let mut new_transcript = self.buffers_free.pop();
new_transcript.push(0xFE);
let hash_len = hash_value.len() as u32;
new_transcript.extend_from_slice(&hash_len.to_be_bytes()[1..]);
new_transcript.extend_from_slice(&hash_value);
new_transcript.extend_from_slice(&self.transcript[split_at..]);
let old = mem::replace(&mut self.transcript, new_transcript);
self.buffers_free.push(old);
}
pub fn enable_peer_encryption(&mut self) -> Result<(), InternalError> {
debug!("Peer encryption enabled");
self.peer_encryption_enabled = true;
let maybe_index = self
.queue_rx
.iter()
.position(|i| i.records().iter().any(|r| r.record().sequence.epoch >= 2));
let Some(index) = maybe_index else {
return Ok(());
};
let all = self.queue_rx.split_off(index);
for incoming in all {
let unhandled = incoming.into_records().filter(|r| !r.is_handled());
for record in unhandled {
let buf = record.into_buffer();
self.parse_packet(&buf)?;
self.buffers_free.push(buf);
}
}
Ok(())
}
pub fn find_kx_group(
&self,
group: crate::types::NamedGroup,
) -> Option<&'static dyn SupportedKxGroup> {
self.config.kx_groups().find(|g| g.name() == group)
}
pub fn verify_signature(
&self,
cert_der: &[u8],
data: &[u8],
signature: &[u8],
hash_alg: HashAlgorithm,
sig_alg: crate::types::SignatureAlgorithm,
) -> Result<(), Error> {
self.config
.crypto_provider()
.signature_verification
.verify_signature(cert_der, data, signature, hash_alg, sig_alg)
.map_err(Error::CryptoError)
}
pub fn extract_srtp_keying_material(
&self,
profile: crate::crypto::SrtpProfile,
) -> Result<(ArrayVec<u8, 88>, crate::crypto::SrtpProfile), Error> {
let hash = self.hash_algorithm();
let hash_len = hash.output_len();
let hmac = self.hmac();
let exp_master = self
.exporter_master_secret
.as_ref()
.ok_or(Error::CryptoError(
crate::CryptoError::ExporterMasterSecretNotDerived,
))?;
let total_len = profile.keying_material_len();
let empty_hash = self.transcript_hash_of(b"");
let mut derived = Buf::new();
prf_hkdf::hkdf_expand_label_dtls13(
hmac,
hash,
exp_master,
b"EXTRACTOR-dtls_srtp",
&empty_hash,
&mut derived,
hash_len,
)
.map_err(Error::CryptoError)?;
let context_hash = self.transcript_hash_of(b"");
let mut keying_material_buf = Buf::new();
prf_hkdf::hkdf_expand_label_dtls13(
hmac,
hash,
&derived,
b"exporter",
&context_hash,
&mut keying_material_buf,
total_len,
)
.map_err(Error::CryptoError)?;
let mut keying_material = ArrayVec::new();
for &b in keying_material_buf.iter().take(total_len) {
keying_material.push(b);
}
Ok((keying_material, profile))
}
pub fn random(&mut self) -> Random {
Random::new(&mut self.rng)
}
pub fn random_arr<const N: usize>(&mut self) -> [u8; N] {
self.rng.random()
}
}
fn jittered_aead_threshold(limit: u64, rng: &mut SeededRng) -> u64 {
let quarter = limit / 4;
if quarter == 0 {
return limit;
}
let offset: u64 = rng.random::<u64>() % (quarter + 1);
limit - quarter + offset
}
fn epoch_for_message(msg_type: MessageType) -> u16 {
match msg_type {
MessageType::ClientHello | MessageType::ServerHello => 0,
_ => 2,
}
}
fn reconstruct_sequence(partial: u64, expected: u64, bits: u32) -> u64 {
let mask = (1u64 << bits) - 1;
let window = 1u64 << bits;
let half = window / 2;
let received_partial = partial & mask;
let expected_partial = expected & mask;
let diff = (received_partial as i64) - (expected_partial as i64);
let diff = if diff > half as i64 {
diff - window as i64
} else if diff < -(half as i64) {
diff + window as i64
} else {
diff
};
(expected as i64 + diff).max(0) as u64
}
impl RecordHandler for Engine {
fn classify_record(&mut self, record: Record) -> Result<Option<Record>, Error> {
if let Some(cn_seq) = self.close_notify_sequence {
if record.record().sequence > cn_seq {
self.push_buffer(record.into_buffer());
return Ok(None);
}
}
let epoch = record.record().sequence.epoch;
if epoch == 0
&& self.peer_encryption_enabled
&& matches!(
record.record().content_type,
ContentType::Ack | ContentType::Alert
)
{
self.push_buffer(record.into_buffer());
return Ok(None);
}
match record.record().content_type {
ContentType::Ack => {
let fragment = record.record().fragment(record.buffer());
self.process_ack(fragment)?;
self.push_buffer(record.into_buffer());
Ok(None)
}
ContentType::Alert => {
let description = {
let fragment = record.record().fragment(record.buffer());
fragment.get(1).copied()
};
let sequence = record.record().sequence;
self.push_buffer(record.into_buffer());
match description {
Some(0) => {
self.close_notify_sequence.get_or_insert(sequence);
Ok(None)
}
Some(90) => Ok(None),
Some(description) => {
Err(Error::SecurityError(crate::SecurityError::FatalAlert {
description,
}))
}
None => Ok(None),
}
}
ContentType::ChangeCipherSpec => {
trace!("Discarding CCS record");
self.push_buffer(record.into_buffer());
Ok(None)
}
_ => Ok(Some(record)),
}
}
fn is_peer_encryption_enabled(&self) -> bool {
self.peer_encryption_enabled
}
fn resolve_epoch(&self, epoch_bits: u8) -> u16 {
let epoch_bits = epoch_bits as u16;
let mut best = None;
for entry in &self.app_recv_keys {
if (entry.epoch & 0x03) == epoch_bits {
best = Some(entry.epoch);
}
}
if self.hs_recv_keys.is_some() && (2 & 0x03) == epoch_bits && !self.release_app_data {
return 2;
}
if let Some(epoch) = best {
return epoch;
}
if self.hs_recv_keys.is_some() && (2 & 0x03) == epoch_bits {
return 2;
}
epoch_bits
}
fn resolve_sequence(&self, epoch: u16, seq_bits: u64, s_flag: bool) -> u64 {
let expected = if epoch == 2 {
self.hs_expected_recv_seq
} else {
self.app_recv_keys
.iter()
.find(|e| e.epoch == epoch)
.map(|e| e.expected_recv_seq)
.unwrap_or(0)
};
let bits: u32 = if s_flag { 16 } else { 8 };
reconstruct_sequence(seq_bits, expected, bits)
}
fn replay_check(&self, seq: Sequence) -> bool {
if seq.epoch == 2 {
self.hs_replay.check(seq.sequence_number)
} else {
match self.app_recv_keys.iter().find(|e| e.epoch == seq.epoch) {
Some(entry) => entry.replay.check(seq.sequence_number),
None => false, }
}
}
fn replay_update(&mut self, seq: Sequence) {
if seq.epoch == 2 {
self.hs_replay.update(seq.sequence_number);
} else if let Some(entry) = self.app_recv_keys.iter_mut().find(|e| e.epoch == seq.epoch) {
entry.replay.update(seq.sequence_number);
}
let next = seq.sequence_number + 1;
if seq.epoch == 2 {
if next > self.hs_expected_recv_seq {
self.hs_expected_recv_seq = next;
}
} else {
for entry in &mut self.app_recv_keys {
if entry.epoch == seq.epoch {
if next > entry.expected_recv_seq {
entry.expected_recv_seq = next;
}
break;
}
}
}
}
fn min_protected_fragment_len(&self) -> usize {
self.suite_provider().min_protected_fragment_len()
}
fn decrypt_record(
&mut self,
header: &[u8],
seq: Sequence,
ciphertext: &mut TmpBuf,
) -> Result<(), Error> {
let keys = if seq.epoch == 2 {
self.hs_recv_keys.as_mut()
} else {
self.app_recv_keys
.iter_mut()
.find(|e| e.epoch == seq.epoch)
.map(|e| &mut e.keys)
};
let Some(keys) = keys else {
return Err(Error::CryptoError(
crate::CryptoError::RecvKeysNotAvailable { epoch: seq.epoch },
));
};
let nonce = Nonce::xor(&keys.iv, seq.sequence_number);
let aad = Aad::new_dtls13(header);
keys.cipher
.decrypt(ciphertext, aad, nonce)
.map_err(Error::CryptoError)?;
Ok(())
}
fn decrypt_sequence_number(
&self,
epoch: u16,
seq_bytes: &mut [u8],
ciphertext_sample: &[u8; 16],
) {
let sn_key = if epoch == 2 {
self.hs_recv_keys.as_ref().map(|k| &k.sn_key)
} else {
self.app_recv_keys
.iter()
.find(|e| e.epoch == epoch)
.map(|e| &e.keys.sn_key)
};
let Some(sn_key) = sn_key else {
return; };
let mask = self.suite_provider().encrypt_sn(sn_key, ciphertext_sample);
for (i, byte) in seq_bytes.iter_mut().enumerate() {
*byte ^= mask[i];
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[cfg(feature = "rcgen")]
use crate::certificate::generate_self_signed_certificate;
#[cfg(feature = "rcgen")]
fn test_engine() -> Engine {
let cert = generate_self_signed_certificate().expect("gen cert");
let config = Arc::new(Config::builder().build().expect("build config"));
Engine::new(config, cert)
}
struct PassthroughRecordHandler;
impl RecordHandler for PassthroughRecordHandler {
fn classify_record(&mut self, record: Record) -> Result<Option<Record>, Error> {
Ok(Some(record))
}
fn is_peer_encryption_enabled(&self) -> bool {
true
}
fn resolve_epoch(&self, _epoch_bits: u8) -> u16 {
2
}
fn resolve_sequence(&self, _epoch: u16, seq_bits: u64, _s_flag: bool) -> u64 {
seq_bits
}
fn replay_check(&self, _seq: Sequence) -> bool {
true
}
fn replay_update(&mut self, _seq: Sequence) {}
fn min_protected_fragment_len(&self) -> usize {
0
}
fn decrypt_record(
&mut self,
_header: &[u8],
_seq: Sequence,
_ciphertext: &mut TmpBuf,
) -> Result<(), Error> {
Ok(())
}
fn decrypt_sequence_number(
&self,
_epoch: u16,
_seq_bytes: &mut [u8],
_ciphertext_sample: &[u8; 16],
) {
}
}
fn encrypted_key_update_record(seq: u16) -> Vec<u8> {
let mut fragment = Vec::new();
fragment.push(MessageType::KeyUpdate.as_u8());
fragment.extend_from_slice(&1u32.to_be_bytes()[1..]);
fragment.extend_from_slice(&0u16.to_be_bytes());
fragment.extend_from_slice(&0u32.to_be_bytes()[1..]);
fragment.extend_from_slice(&1u32.to_be_bytes()[1..]);
fragment.push(KeyUpdateRequest::UpdateRequested.as_u8());
fragment.push(ContentType::Handshake.as_u8());
let mut packet = Vec::new();
packet.push(
0b0010_0000
| 0b0000_1000 | 0b0000_0100 | 0b0000_0010, );
packet.extend_from_slice(&seq.to_be_bytes());
packet.extend_from_slice(&(fragment.len() as u16).to_be_bytes());
packet.extend_from_slice(&fragment);
packet
}
fn encrypted_application_data_record(seq: u16, data: &[u8]) -> Vec<u8> {
let mut fragment = Vec::new();
fragment.extend_from_slice(data);
fragment.push(ContentType::ApplicationData.as_u8());
let mut packet = Vec::new();
packet.push(
0b0010_0000
| 0b0000_1000 | 0b0000_0100 | 0b0000_0010, );
packet.extend_from_slice(&seq.to_be_bytes());
packet.extend_from_slice(&(fragment.len() as u16).to_be_bytes());
packet.extend_from_slice(&fragment);
packet
}
fn parsed_key_update(seq: u16) -> Incoming {
Incoming::parse_packet(
&encrypted_key_update_record(seq),
&mut PassthroughRecordHandler,
Some(Dtls13CipherSuite::AES_128_GCM_SHA256),
)
.expect("parse key update packet")
.expect("packet contains a record")
}
fn parsed_key_update_with_app_data(key_update_seq: u16, app_seq: u16) -> Incoming {
let mut packet = encrypted_key_update_record(key_update_seq);
packet.extend_from_slice(&encrypted_application_data_record(app_seq, b"app-data"));
Incoming::parse_packet(
&packet,
&mut PassthroughRecordHandler,
Some(Dtls13CipherSuite::AES_128_GCM_SHA256),
)
.expect("parse coalesced packet")
.expect("packet contains records")
}
#[test]
#[cfg(feature = "rcgen")]
fn epoch_0_sequence_number_rejects_overflow() {
let mut engine = test_engine();
engine.sequence_epoch_0.sequence_number = MAX_SEQUENCE_NUMBER;
let result = engine.create_plaintext_record(ContentType::Handshake, false, |buf| {
buf.extend_from_slice(b"test")
});
assert!(
result.is_err(),
"epoch-0 must reject sequence overflow at MAX_SEQUENCE_NUMBER"
);
}
#[test]
#[cfg(feature = "rcgen")]
fn derive_handshake_secrets_returns_handshake_secret() {
let mut engine = test_engine();
engine.set_cipher_suite(Dtls13CipherSuite::AES_128_GCM_SHA256);
engine
.transcript
.extend_from_slice(b"dummy transcript for test");
let shared_secret = [0x42u8; 32];
let (c_hs_traffic, _s_hs_traffic, handshake_secret) =
engine.derive_handshake_secrets(&shared_secret).unwrap();
let hash = engine.hash_algorithm();
let hash_len = hash.output_len();
let hmac = engine.hmac();
let mut transcript_hash = Buf::new();
engine.transcript_hash(&mut transcript_hash);
let mut c_hs_manual = Buf::new();
prf_hkdf::hkdf_expand_label_dtls13(
hmac,
hash,
&handshake_secret,
b"c hs traffic",
&transcript_hash,
&mut c_hs_manual,
hash_len,
)
.expect("hkdf_expand_label_dtls13");
assert_eq!(
c_hs_manual.as_ref(),
c_hs_traffic.as_ref(),
"handshake_secret from derive_handshake_secrets() must reproduce \
the same traffic secrets"
);
}
#[test]
#[cfg(feature = "rcgen")]
fn derive_early_secret_uses_buffer_pool() {
let mut engine = test_engine();
engine.set_cipher_suite(Dtls13CipherSuite::AES_128_GCM_SHA256);
let mut marked = Buf::new();
marked.extend_from_slice(&[0xAA; 256]);
engine.buffers_free.push(marked);
let early_secret = engine.derive_early_secret().unwrap();
assert!(
early_secret.into_vec().capacity() >= 256,
"derive_early_secret must use the buffer pool, returning a buffer with pooled capacity"
);
}
#[test]
#[cfg(feature = "rcgen")]
fn ack_tracking_full_does_not_panic_on_handshake_replacement() {
let mut engine = test_engine();
let first = parsed_key_update(0);
engine
.insert_incoming(first)
.expect("insert initial key update");
engine.queue_rx[0]
.first()
.first_handshake()
.expect("initial key update handshake")
.set_handled();
engine.received_record_numbers.clear();
for sequence in 0..engine.received_record_numbers.capacity() {
engine.received_record_numbers.push((2, sequence as u64));
}
let replacement = parsed_key_update(1);
let result = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
engine
.insert_incoming(replacement)
.expect("replace handled key update")
}));
assert!(
result.is_ok(),
"full ACK bookkeeping must not panic when a handled handshake is replaced"
);
assert_eq!(engine.queue_rx.len(), 1);
assert_eq!(
engine.queue_rx[0].first().record().sequence.sequence_number,
1
);
assert_eq!(
engine.received_record_numbers.len(),
engine.received_record_numbers.capacity(),
"overflowing ACK bookkeeping should keep existing entries and drop the extra one"
);
}
#[test]
#[cfg(feature = "rcgen")]
fn ack_tracking_ignores_non_handshake_records_in_coalesced_datagram() {
let mut engine = test_engine();
let incoming = parsed_key_update_with_app_data(7, 8);
assert_eq!(incoming.records().len(), 2);
assert_eq!(
incoming.records()[0].record().content_type,
ContentType::Handshake
);
assert_eq!(
incoming.records()[1].record().content_type,
ContentType::ApplicationData
);
engine
.insert_incoming(incoming)
.expect("insert coalesced datagram");
assert_eq!(
engine.received_record_numbers.as_slice(),
&[(2, 7)],
"ACK bookkeeping must include only handshake records"
);
}
#[test]
#[cfg(feature = "rcgen")]
fn malformed_ack_record_number_vector_is_ignored() {
let mut engine = test_engine();
engine.flight_saved_records.push(Entry {
content_type: ContentType::Handshake,
epoch: 2,
send_seq: 7,
fragment: Buf::new(),
acked: false,
});
let mut malformed_ack = Vec::new();
malformed_ack.extend_from_slice(&17u16.to_be_bytes());
malformed_ack.extend_from_slice(&2u64.to_be_bytes());
malformed_ack.extend_from_slice(&7u64.to_be_bytes());
malformed_ack.push(0);
engine.process_ack(&malformed_ack).unwrap();
assert!(
!engine.flight_saved_records[0].acked,
"malformed ACK vector length must not partially acknowledge records"
);
}
}