use std::sync::Arc;
use std::time::{Duration, Instant};
use arrayvec::ArrayVec;
use crate::buffer::Buf;
use crate::crypto::ActiveKeyExchange;
use crate::dtls13::message::KeyShareClientHello;
use crate::dtls13::message::KeyShareEntry;
use crate::dtls13::message::Random;
use crate::dtls13::message::SignatureAlgorithmsExtension;
use crate::dtls13::message::SupportedGroupsExtension;
use crate::dtls13::message::UseSrtpExtension;
use crate::types::NamedGroup;
use crate::{Config, DtlsCertificate, Error, Output, SeededRng};
const EXT_SUPPORTED_GROUPS: u16 = 0x000A;
const EXT_EC_POINT_FORMATS: u16 = 0x000B;
const EXT_SIGNATURE_ALGORITHMS: u16 = 0x000D;
const EXT_USE_SRTP: u16 = 0x000E;
const EXT_PADDING: u16 = 0x0015;
const EXT_EXTENDED_MASTER_SECRET: u16 = 0x0017;
const EXT_SUPPORTED_VERSIONS: u16 = 0x002B;
const EXT_KEY_SHARE: u16 = 0x0033;
const EXT_RENEGOTIATION_INFO: u16 = 0xFF01;
pub(crate) struct HybridClientHello {
pub random: Random,
pub active_key_exchange: Box<dyn ActiveKeyExchange>,
pub transcript_bytes: Buf,
pub handshake_fragment: Buf,
}
impl HybridClientHello {
pub fn new(config: &Arc<Config>) -> Result<Self, Error> {
let mut rng = SeededRng::new(config.rng_seed());
let random = Random::new(&mut rng);
let group = config
.kx_groups()
.next()
.ok_or_else(|| Error::CryptoError("No supported key exchange groups".into()))?;
let kx_buf = Buf::new();
let key_exchange = group
.start_exchange(kx_buf)
.map_err(|e| Error::CryptoError(format!("Failed to start key exchange: {}", e)))?;
let mut ch_body = Buf::new();
ch_body.extend_from_slice(&0xFEFDu16.to_be_bytes());
random.serialize(&mut ch_body);
ch_body.push(0);
ch_body.push(0);
let mut suites: ArrayVec<u16, 16> = ArrayVec::new();
for cs in config.dtls13_cipher_suites() {
suites.push(cs.suite().as_u16());
}
for cs in config
.dtls12_cipher_suites()
.filter(|cs| !cs.suite().is_psk())
{
suites.push(cs.suite().as_u16());
}
ch_body.extend_from_slice(&((suites.len() * 2) as u16).to_be_bytes());
for &suite in &suites {
ch_body.extend_from_slice(&suite.to_be_bytes());
}
ch_body.push(1);
ch_body.push(0);
let mut ext_buf = Buf::new();
let mut ext_entries: Vec<(u16, usize, usize)> = Vec::new();
let start = ext_buf.len();
ext_buf.push(4); ext_buf.extend_from_slice(&0xFEFCu16.to_be_bytes()); ext_buf.extend_from_slice(&0xFEFDu16.to_be_bytes()); ext_entries.push((EXT_SUPPORTED_VERSIONS, start, ext_buf.len()));
let start = ext_buf.len();
let groups: ArrayVec<NamedGroup, 4> = config.kx_groups().map(|g| g.name()).collect();
let sg = SupportedGroupsExtension { groups };
sg.serialize(&mut ext_buf);
ext_entries.push((EXT_SUPPORTED_GROUPS, start, ext_buf.len()));
let pub_key = key_exchange.pub_key();
let mut key_data = Buf::new();
let pub_key_start = key_data.len();
key_data.extend_from_slice(pub_key);
let pub_key_end = key_data.len();
let start = ext_buf.len();
let mut entries = ArrayVec::new();
entries.push(KeyShareEntry {
group: group.name(),
key_exchange_range: pub_key_start..pub_key_end,
});
let ks = KeyShareClientHello { entries };
ks.serialize(&key_data, &mut ext_buf);
ext_entries.push((EXT_KEY_SHARE, start, ext_buf.len()));
let start = ext_buf.len();
let sa = SignatureAlgorithmsExtension::default();
sa.serialize(&mut ext_buf);
ext_entries.push((EXT_SIGNATURE_ALGORITHMS, start, ext_buf.len()));
let start = ext_buf.len();
let use_srtp = UseSrtpExtension::default();
use_srtp.serialize(&mut ext_buf);
ext_entries.push((EXT_USE_SRTP, start, ext_buf.len()));
let start = ext_buf.len();
ext_buf.push(1); ext_buf.push(0); ext_entries.push((EXT_EC_POINT_FORMATS, start, ext_buf.len()));
ext_entries.push((EXT_EXTENDED_MASTER_SECRET, ext_buf.len(), ext_buf.len()));
let start = ext_buf.len();
ext_buf.push(0); ext_entries.push((EXT_RENEGOTIATION_INFO, start, ext_buf.len()));
let record_header = 13usize;
let handshake_header = 12usize;
let body_so_far = ch_body.len()
+ 2 + ext_entries.iter().map(|(_, s, e)| 4 + (e - s)).sum::<usize>();
let total_so_far = record_header + handshake_header + body_so_far;
let deficit = config.mtu().saturating_sub(total_so_far);
if deficit >= 4 {
let pad_data_len = deficit - 4; let start = ext_buf.len();
for _ in 0..pad_data_len {
ext_buf.push(0);
}
ext_entries.push((EXT_PADDING, start, ext_buf.len()));
}
let ext_total_len: usize = ext_entries.iter().map(|(_, s, e)| 4 + (e - s)).sum();
ch_body.extend_from_slice(&(ext_total_len as u16).to_be_bytes());
for &(ext_type, start, end) in &ext_entries {
ch_body.extend_from_slice(&ext_type.to_be_bytes());
ch_body.extend_from_slice(&((end - start) as u16).to_be_bytes());
if end > start {
ch_body.extend_from_slice(&ext_buf[start..end]);
}
}
let mut transcript_bytes = Buf::new();
transcript_bytes.push(0x01); let body_len = ch_body.len() as u32;
transcript_bytes.extend_from_slice(&body_len.to_be_bytes()[1..]);
transcript_bytes.extend_from_slice(&ch_body);
let mut handshake_fragment = Buf::new();
handshake_fragment.push(0x01); handshake_fragment.extend_from_slice(&body_len.to_be_bytes()[1..]); handshake_fragment.extend_from_slice(&0u16.to_be_bytes()); handshake_fragment.extend_from_slice(&0u32.to_be_bytes()[1..]); handshake_fragment.extend_from_slice(&body_len.to_be_bytes()[1..]); handshake_fragment.extend_from_slice(&ch_body);
Ok(HybridClientHello {
random,
active_key_exchange: key_exchange,
transcript_bytes,
handshake_fragment,
})
}
pub fn wire_packet(&self) -> Buf {
let mut pkt = Buf::new();
pkt.push(0x16); pkt.extend_from_slice(&0xFEFDu16.to_be_bytes()); pkt.extend_from_slice(&0u16.to_be_bytes()); pkt.extend_from_slice(&[0u8; 6]); pkt.extend_from_slice(&(self.handshake_fragment.len() as u16).to_be_bytes());
pkt.extend_from_slice(&self.handshake_fragment);
pkt
}
}
pub(crate) struct ClientPending {
hybrid: HybridClientHello,
config: Arc<Config>,
certificate: DtlsCertificate,
wire_packet: Buf,
needs_send: bool,
last_now: Instant,
retransmit_at: Option<Instant>,
retransmit_count: usize,
}
impl ClientPending {
pub fn new(
config: Arc<Config>,
certificate: DtlsCertificate,
now: Instant,
) -> Result<Self, Error> {
let hybrid = HybridClientHello::new(&config)?;
let wire_packet = hybrid.wire_packet();
Ok(ClientPending {
hybrid,
config,
certificate,
wire_packet,
needs_send: true,
last_now: now,
retransmit_at: None,
retransmit_count: 0,
})
}
pub fn handle_timeout(&mut self, now: Instant) -> Result<(), Error> {
self.last_now = now;
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("hybrid ClientHello"));
}
self.retransmit_count += 1;
self.needs_send = true;
let shift = self.retransmit_count.min(5) as u32;
let rto = Duration::from_secs(1u64 << shift);
self.retransmit_at = Some(now + rto);
}
}
Ok(())
}
pub fn poll_output<'a>(&mut self, buf: &'a mut [u8]) -> Output<'a> {
if self.needs_send {
let len = self.wire_packet.len();
if buf.len() < len {
let next = self
.retransmit_at
.unwrap_or(self.last_now + Duration::from_secs(1));
return Output::Timeout(next);
}
self.needs_send = false;
buf[..len].copy_from_slice(&self.wire_packet);
return Output::Packet(&buf[..len]);
}
let next = self
.retransmit_at
.unwrap_or(self.last_now + Duration::from_secs(1));
Output::Timeout(next)
}
pub fn into_parts(self) -> (HybridClientHello, Arc<Config>, DtlsCertificate, Instant) {
(self.hybrid, self.config, self.certificate, self.last_now)
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub(crate) enum DetectedVersion {
Dtls12,
Dtls13,
Unknown,
}
pub(crate) fn server_hello_version(packet: &[u8]) -> DetectedVersion {
server_hello_version_inner(packet).unwrap_or(DetectedVersion::Unknown)
}
fn server_hello_version_inner(packet: &[u8]) -> Option<DetectedVersion> {
if packet.len() < 13 {
return None;
}
if packet[0] != 0x16 {
return None;
}
let record_len = u16::from_be_bytes([packet[11], packet[12]]) as usize;
let record_body = packet.get(13..13 + record_len)?;
if record_body.len() < 12 {
return None;
}
let msg_type = record_body[0];
if msg_type == 3 {
return Some(DetectedVersion::Dtls12);
}
if msg_type != 2 {
return None;
}
let fragment_len = ((record_body[9] as usize) << 16)
| ((record_body[10] as usize) << 8)
| (record_body[11] as usize);
let body = record_body.get(12..12 + fragment_len)?;
if body.len() < 35 {
return Some(DetectedVersion::Dtls12);
}
let mut pos = 34;
let sid_len = *body.get(pos)? as usize;
pos += 1 + sid_len;
pos += 2;
pos += 1;
if pos + 2 > body.len() {
return Some(DetectedVersion::Dtls12);
}
let ext_total_len = u16::from_be_bytes([body[pos], body[pos + 1]]) as usize;
pos += 2;
let ext_end = pos + ext_total_len;
if ext_end > body.len() {
return Some(DetectedVersion::Dtls12);
}
while pos + 4 <= ext_end {
let ext_type = u16::from_be_bytes([body[pos], body[pos + 1]]);
let ext_len = u16::from_be_bytes([body[pos + 2], body[pos + 3]]) as usize;
pos += 4;
if ext_type == 0x002B {
if ext_len >= 2 {
let version = u16::from_be_bytes([body[pos], body[pos + 1]]);
if version == 0xFEFC {
return Some(DetectedVersion::Dtls13);
}
}
return Some(DetectedVersion::Dtls12);
}
pos += ext_len;
}
Some(DetectedVersion::Dtls12)
}
#[cfg(test)]
mod tests {
use super::*;
use crate::PskResolver;
use crate::dtls12::message::Dtls12CipherSuite;
fn offered_cipher_suites(hybrid: &HybridClientHello) -> Vec<u16> {
let body = &hybrid.handshake_fragment[12..];
let mut offset = 2 + 32;
let session_id_len = body[offset] as usize;
offset += 1 + session_id_len;
let cookie_len = body[offset] as usize;
offset += 1 + cookie_len;
let suites_len = u16::from_be_bytes([body[offset], body[offset + 1]]) as usize;
offset += 2;
body[offset..offset + suites_len]
.chunks_exact(2)
.map(|chunk| u16::from_be_bytes([chunk[0], chunk[1]]))
.collect()
}
struct DummyResolver;
impl PskResolver for DummyResolver {
fn resolve(&self, _identity: &[u8]) -> Option<Vec<u8>> {
Some(b"0123456789abcdef".to_vec())
}
}
#[test]
fn hello_verify_request_is_dtls12() {
let mut pkt = Vec::new();
pkt.push(0x16); pkt.extend_from_slice(&[0xFE, 0xFD]); pkt.extend_from_slice(&[0x00, 0x00]); pkt.extend_from_slice(&[0x00, 0x00, 0x00, 0x00, 0x00, 0x00]); let len_pos = pkt.len();
pkt.extend_from_slice(&[0x00, 0x00]);
pkt.push(3); pkt.extend_from_slice(&[0x00, 0x00, 0x05]); pkt.extend_from_slice(&[0x00, 0x00]); pkt.extend_from_slice(&[0x00, 0x00, 0x00]); pkt.extend_from_slice(&[0x00, 0x00, 0x05]);
pkt.extend_from_slice(&[0xFE, 0xFD]); pkt.push(0x02); pkt.extend_from_slice(&[0xAA, 0xBB]);
let record_len = (pkt.len() - 13) as u16;
pkt[len_pos] = (record_len >> 8) as u8;
pkt[len_pos + 1] = record_len as u8;
assert_eq!(server_hello_version(&pkt), DetectedVersion::Dtls12);
}
#[test]
fn server_hello_with_supported_versions_is_dtls13() {
let mut pkt = Vec::new();
pkt.push(0x16);
pkt.extend_from_slice(&[0xFE, 0xFD]);
pkt.extend_from_slice(&[0x00, 0x00]);
pkt.extend_from_slice(&[0x00, 0x00, 0x00, 0x00, 0x00, 0x00]);
let len_pos = pkt.len();
pkt.extend_from_slice(&[0x00, 0x00]);
pkt.push(2);
let hs_len_pos = pkt.len();
pkt.extend_from_slice(&[0x00, 0x00, 0x00]); pkt.extend_from_slice(&[0x00, 0x00]); pkt.extend_from_slice(&[0x00, 0x00, 0x00]); let frag_len_pos = pkt.len();
pkt.extend_from_slice(&[0x00, 0x00, 0x00]);
let body_start = pkt.len();
pkt.extend_from_slice(&[0xFE, 0xFD]); pkt.extend_from_slice(&[0u8; 32]); pkt.push(0x00); pkt.extend_from_slice(&[0x13, 0x01]); pkt.push(0x00);
let ext_len_pos = pkt.len();
pkt.extend_from_slice(&[0x00, 0x00]);
let ext_start = pkt.len();
pkt.extend_from_slice(&[0x00, 0x2B]); pkt.extend_from_slice(&[0x00, 0x02]); pkt.extend_from_slice(&[0xFE, 0xFC]);
let ext_total = pkt.len() - ext_start;
pkt[ext_len_pos] = (ext_total >> 8) as u8;
pkt[ext_len_pos + 1] = ext_total as u8;
let body_len = pkt.len() - body_start;
pkt[hs_len_pos] = 0;
pkt[hs_len_pos + 1] = (body_len >> 8) as u8;
pkt[hs_len_pos + 2] = body_len as u8;
pkt[frag_len_pos] = 0;
pkt[frag_len_pos + 1] = (body_len >> 8) as u8;
pkt[frag_len_pos + 2] = body_len as u8;
let record_len = (pkt.len() - 13) as u16;
pkt[len_pos] = (record_len >> 8) as u8;
pkt[len_pos + 1] = record_len as u8;
assert_eq!(server_hello_version(&pkt), DetectedVersion::Dtls13);
}
#[test]
fn server_hello_without_supported_versions_is_dtls12() {
let mut pkt = Vec::new();
pkt.push(0x16);
pkt.extend_from_slice(&[0xFE, 0xFD]);
pkt.extend_from_slice(&[0x00, 0x00]);
pkt.extend_from_slice(&[0x00, 0x00, 0x00, 0x00, 0x00, 0x00]);
let len_pos = pkt.len();
pkt.extend_from_slice(&[0x00, 0x00]);
pkt.push(2);
let hs_len_pos = pkt.len();
pkt.extend_from_slice(&[0x00, 0x00, 0x00]);
pkt.extend_from_slice(&[0x00, 0x00]);
pkt.extend_from_slice(&[0x00, 0x00, 0x00]);
let frag_len_pos = pkt.len();
pkt.extend_from_slice(&[0x00, 0x00, 0x00]);
let body_start = pkt.len();
pkt.extend_from_slice(&[0xFE, 0xFD]); pkt.extend_from_slice(&[0u8; 32]); pkt.push(0x00); pkt.extend_from_slice(&[0xC0, 0x2B]); pkt.push(0x00);
let body_len = pkt.len() - body_start;
pkt[hs_len_pos] = 0;
pkt[hs_len_pos + 1] = (body_len >> 8) as u8;
pkt[hs_len_pos + 2] = body_len as u8;
pkt[frag_len_pos] = 0;
pkt[frag_len_pos + 1] = (body_len >> 8) as u8;
pkt[frag_len_pos + 2] = body_len as u8;
let record_len = (pkt.len() - 13) as u16;
pkt[len_pos] = (record_len >> 8) as u8;
pkt[len_pos + 1] = record_len as u8;
assert_eq!(server_hello_version(&pkt), DetectedVersion::Dtls12);
}
#[test]
fn garbage_packet_is_unknown() {
assert_eq!(
server_hello_version(&[0xFF, 0x00]),
DetectedVersion::Unknown
);
assert_eq!(server_hello_version(&[]), DetectedVersion::Unknown);
}
#[test]
fn too_short_is_unknown() {
let pkt = [
0x17, 0xFE, 0xFD, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x01, 0xFF,
];
assert_eq!(server_hello_version(&pkt), DetectedVersion::Unknown);
}
#[test]
fn hybrid_client_hello_excludes_psk_dtls12_suites() {
let config = Arc::new(
Config::builder()
.with_psk_client(b"identity".to_vec(), Arc::new(DummyResolver))
.build()
.expect("config with PSK should build"),
);
assert!(
config.dtls12_cipher_suites().any(|cs| cs.suite().is_psk()),
"precondition: PSK-enabled config should expose a PSK DTLS 1.2 suite"
);
let hybrid = HybridClientHello::new(&config).expect("hybrid ClientHello should build");
let offered = offered_cipher_suites(&hybrid);
assert!(
!offered.contains(&Dtls12CipherSuite::PSK_AES128_CCM_8.as_u16()),
"auto client must not advertise PSK DTLS 1.2 suites it cannot use after fallback"
);
}
}