use std::borrow::Cow;
use nom::Err as NomErr;
use super::RejectReason;
const TLS_CONTENT_TYPE_HANDSHAKE: u8 = 22;
const TLS_HANDSHAKE_TYPE_CLIENT_HELLO: u8 = 1;
const MAX_TLS_RECORD_LEN: usize = 16384;
const EXT_SERVER_NAME: u16 = 0x0000;
const EXT_ALPN: u16 = 0x0010;
const EXT_ENCRYPTED_CLIENT_HELLO: u16 = 0xfe0d;
pub(super) enum ParseOutcome {
NeedMore,
Reject(RejectReason),
ClientHello {
sni: Option<String>,
alpn: Vec<Vec<u8>>,
ech_present: bool,
},
}
struct RecordHeader {
content_type: u8,
length: u16,
}
fn record_header(i: &[u8]) -> nom::IResult<&[u8], RecordHeader> {
let in_len = i.len();
let (i, content_type) = nom::number::streaming::be_u8(i)?;
let (i, _legacy_record_version) = nom::number::streaming::be_u16(i)?;
let (i, length) = nom::number::streaming::be_u16(i)?;
debug_assert!(i.len() <= in_len, "record_header must not grow its input");
debug_assert_eq!(
in_len - i.len(),
5,
"TLS record header is exactly 5 bytes (type + legacy_version + length)"
);
Ok((
i,
RecordHeader {
content_type,
length,
},
))
}
fn take_record_content(i: &[u8], length: u16) -> nom::IResult<&[u8], &[u8]> {
nom::bytes::streaming::take(length)(i)
}
fn append_fragment<'a>(acc: Cow<'a, [u8]>, fragment: &'a [u8]) -> Cow<'a, [u8]> {
match acc {
Cow::Borrowed([]) => Cow::Borrowed(fragment),
Cow::Borrowed(b) => {
let mut owned = Vec::with_capacity(b.len() + fragment.len());
owned.extend_from_slice(b);
owned.extend_from_slice(fragment);
Cow::Owned(owned)
}
Cow::Owned(mut owned) => {
owned.extend_from_slice(fragment);
Cow::Owned(owned)
}
}
}
pub(super) fn parse_client_hello(buf: &[u8]) -> ParseOutcome {
let mut pos = 0usize;
let mut handshake: Cow<'_, [u8]> = Cow::Borrowed(&[][..]);
let mut records_seen: u32 = 0;
let mut msg_type_checked = false;
loop {
let remaining = &buf[pos..];
let (after_header, header) = match record_header(remaining) {
Ok(v) => v,
Err(NomErr::Incomplete(_)) => return ParseOutcome::NeedMore,
Err(_) => return ParseOutcome::Reject(RejectReason::MalformedRecord),
};
records_seen += 1;
if header.content_type != TLS_CONTENT_TYPE_HANDSHAKE {
let reason = if records_seen == 1 {
RejectReason::NotTls
} else {
RejectReason::MalformedRecord
};
debug_assert!(
matches!(reason, RejectReason::NotTls) == (records_seen == 1),
"only the first non-handshake record may reject as NotTls"
);
return ParseOutcome::Reject(reason);
}
if header.length as usize > MAX_TLS_RECORD_LEN {
return ParseOutcome::Reject(RejectReason::MalformedRecord);
}
let (after_content, content) = match take_record_content(after_header, header.length) {
Ok(v) => v,
Err(NomErr::Incomplete(_)) => return ParseOutcome::NeedMore,
Err(_) => return ParseOutcome::Reject(RejectReason::MalformedRecord),
};
let consumed_so_far = buf.len() - after_content.len();
debug_assert!(
consumed_so_far <= buf.len(),
"record walk cursor must never exceed the buffer length"
);
debug_assert!(
consumed_so_far > pos,
"each successfully parsed record must strictly advance the cursor"
);
pos = consumed_so_far;
handshake = append_fragment(handshake, content);
if !msg_type_checked && let Some(&msg_type) = handshake.first() {
if msg_type != TLS_HANDSHAKE_TYPE_CLIENT_HELLO {
return ParseOutcome::Reject(RejectReason::MalformedHandshake);
}
msg_type_checked = true;
}
if handshake.len() >= 4 {
let hs_len = u32::from_be_bytes([0, handshake[1], handshake[2], handshake[3]]) as usize;
if handshake.len() >= 4 + hs_len {
let body = &handshake[4..4 + hs_len];
return parse_client_hello_body(body);
}
}
}
}
fn parse_client_hello_body(body: &[u8]) -> ParseOutcome {
match parse_client_hello_fields(body) {
Ok((sni, alpn, ech_present)) => ParseOutcome::ClientHello {
sni,
alpn,
ech_present,
},
Err(reason) => ParseOutcome::Reject(reason),
}
}
fn malformed(_: NomErr<nom::error::Error<&[u8]>>) -> RejectReason {
RejectReason::MalformedHandshake
}
type ClientHelloFields = (Option<String>, Vec<Vec<u8>>, bool);
fn parse_client_hello_fields(body: &[u8]) -> Result<ClientHelloFields, RejectReason> {
use nom::{
bytes::complete::take,
number::complete::{be_u8, be_u16},
};
let in_len = body.len();
let (i, _legacy_version) = be_u16(body).map_err(malformed)?;
let (i, _random) = take(32usize)(i).map_err(malformed)?;
let (i, session_id_len) = be_u8(i).map_err(malformed)?;
let (i, _session_id) = take(session_id_len)(i).map_err(malformed)?;
let (i, cipher_suites_len) = be_u16(i).map_err(malformed)?;
let (i, _cipher_suites) = take(cipher_suites_len)(i).map_err(malformed)?;
let (i, compression_len) = be_u8(i).map_err(malformed)?;
let (i, _compression_methods) = take(compression_len)(i).map_err(malformed)?;
debug_assert!(
i.len() <= in_len,
"parse_client_hello_fields must not grow its input"
);
debug_assert!(
in_len - i.len() >= 2 + 32 + 1 + 1 + 2,
"legacy_version (2) + random (32) + the session_id (1), cipher_suites (2) and compression_methods (1) length prefixes must be consumed"
);
if i.is_empty() {
return Ok((None, Vec::new(), false));
}
let (i, ext_total_len) = be_u16(i).map_err(malformed)?;
let (_, ext_block) = take(ext_total_len)(i).map_err(malformed)?;
parse_extensions(ext_block)
}
fn parse_extensions(mut ext_block: &[u8]) -> Result<ClientHelloFields, RejectReason> {
use nom::{bytes::complete::take, number::complete::be_u16};
let mut sni: Option<String> = None;
let mut alpn: Vec<Vec<u8>> = Vec::new();
let mut ech_present = false;
let mut seen_sni = false;
let mut seen_alpn = false;
while !ext_block.is_empty() {
let in_len = ext_block.len();
let (i, ext_type) = be_u16(ext_block).map_err(malformed)?;
let (i, ext_len) = be_u16(i).map_err(malformed)?;
let (i, data) = take(ext_len)(i).map_err(malformed)?;
debug_assert_eq!(
in_len - i.len(),
4 + ext_len as usize,
"an extension record consumes exactly 4 + declared_len bytes"
);
match ext_type {
EXT_SERVER_NAME => {
if seen_sni {
return Err(RejectReason::MalformedHandshake);
}
seen_sni = true;
sni = parse_server_name_extension(data)?;
}
EXT_ALPN => {
if seen_alpn {
return Err(RejectReason::MalformedHandshake);
}
seen_alpn = true;
alpn = parse_alpn_extension(data)?;
}
EXT_ENCRYPTED_CLIENT_HELLO => ech_present = true,
_ => {}
}
ext_block = i;
}
debug_assert!(
seen_sni || sni.is_none(),
"sni can only be extracted when a server_name extension was seen"
);
debug_assert!(
seen_alpn || alpn.is_empty(),
"alpn can only be non-empty when an alpn extension was seen"
);
Ok((sni, alpn, ech_present))
}
fn parse_server_name_extension(data: &[u8]) -> Result<Option<String>, RejectReason> {
use nom::{
bytes::complete::take,
number::complete::{be_u8, be_u16},
};
if data.is_empty() {
return Ok(None);
}
let (i, list_len) = be_u16(data).map_err(malformed)?;
let (_, mut entries) = take(list_len)(i).map_err(malformed)?;
let mut host_name: Option<&[u8]> = None;
let mut host_name_entries: u32 = 0;
while !entries.is_empty() {
let (i, name_type) = be_u8(entries).map_err(malformed)?;
let (i, name_len) = be_u16(i).map_err(malformed)?;
let (i, name) = take(name_len)(i).map_err(malformed)?;
if name_type == 0 {
host_name_entries += 1;
if host_name_entries > 1 {
return Err(RejectReason::MalformedHandshake);
}
host_name = Some(name);
}
entries = i;
}
debug_assert!(
entries.is_empty(),
"server_name_list walk must fully drain the declared list"
);
debug_assert!(
host_name_entries <= 1,
"at most one host_name entry may survive the walk -- a second must reject early"
);
match host_name {
Some(name) => match std::str::from_utf8(name) {
Ok(s) => Ok(Some(s.to_owned())),
Err(_) => Err(RejectReason::MalformedHandshake),
},
None => Ok(None),
}
}
fn parse_alpn_extension(data: &[u8]) -> Result<Vec<Vec<u8>>, RejectReason> {
use nom::{
bytes::complete::take,
number::complete::{be_u8, be_u16},
};
let (i, list_len) = be_u16(data).map_err(malformed)?;
let (_, mut entries) = take(list_len)(i).map_err(malformed)?;
let mut protocols = Vec::new();
while !entries.is_empty() {
let (i, name_len) = be_u8(entries).map_err(malformed)?;
if name_len == 0 {
return Err(RejectReason::MalformedHandshake);
}
let (i, name) = take(name_len)(i).map_err(malformed)?;
protocols.push(name.to_vec());
entries = i;
}
if protocols.is_empty() {
return Err(RejectReason::MalformedHandshake);
}
debug_assert!(
protocols.len() <= list_len as usize / 2,
"post-fix, each ALPN entry costs >= 2 wire bytes, halving the pre-fix entry-count bound"
);
debug_assert!(
entries.is_empty(),
"ALPN protocol_name_list walk must fully drain the declared list"
);
Ok(protocols)
}
#[cfg(test)]
pub(super) fn encode_extension(ext_type: u16, data: &[u8]) -> Vec<u8> {
let mut out = Vec::with_capacity(4 + data.len());
out.extend_from_slice(&ext_type.to_be_bytes());
out.extend_from_slice(&(data.len() as u16).to_be_bytes());
out.extend_from_slice(data);
out
}
#[cfg(test)]
pub(super) fn encode_sni_extension(host: &str) -> Vec<u8> {
let mut name_list = vec![0u8]; name_list.extend_from_slice(&(host.len() as u16).to_be_bytes());
name_list.extend_from_slice(host.as_bytes());
let mut data = Vec::new();
data.extend_from_slice(&(name_list.len() as u16).to_be_bytes());
data.extend_from_slice(&name_list);
encode_extension(EXT_SERVER_NAME, &data)
}
#[cfg(test)]
pub(super) fn encode_alpn_extension(protocols: &[&[u8]]) -> Vec<u8> {
let mut list = Vec::new();
for p in protocols {
list.push(p.len() as u8);
list.extend_from_slice(p);
}
let mut data = Vec::new();
data.extend_from_slice(&(list.len() as u16).to_be_bytes());
data.extend_from_slice(&list);
encode_extension(EXT_ALPN, &data)
}
#[cfg(test)]
pub(super) fn encode_grease_extension() -> Vec<u8> {
encode_extension(0x0a0a, &[0x00])
}
#[cfg(test)]
pub(super) fn build_client_hello_body(extra_extensions: &[Vec<u8>]) -> Vec<u8> {
let mut body = Vec::new();
body.extend_from_slice(&[0x03, 0x03]); body.extend_from_slice(&[0u8; 32]); body.push(0); body.extend_from_slice(&[0x00, 0x02, 0x13, 0x01]); body.push(1); body.push(0);
let mut ext_block = Vec::new();
for ext in extra_extensions {
ext_block.extend_from_slice(ext);
}
body.extend_from_slice(&(ext_block.len() as u16).to_be_bytes());
body.extend_from_slice(&ext_block);
body
}
#[cfg(test)]
pub(super) fn wrap_handshake(body: &[u8]) -> Vec<u8> {
let mut hs = Vec::with_capacity(4 + body.len());
hs.push(TLS_HANDSHAKE_TYPE_CLIENT_HELLO);
let len = body.len() as u32;
hs.extend_from_slice(&len.to_be_bytes()[1..4]);
hs.extend_from_slice(body);
hs
}
#[cfg(test)]
pub(super) fn wrap_record(content_type: u8, payload: &[u8]) -> Vec<u8> {
let mut rec = Vec::with_capacity(5 + payload.len());
rec.push(content_type);
rec.extend_from_slice(&[0x03, 0x03]); rec.extend_from_slice(&(payload.len() as u16).to_be_bytes());
rec.extend_from_slice(payload);
rec
}
#[cfg(test)]
pub(super) fn build_client_hello_wire(extra_extensions: &[Vec<u8>]) -> Vec<u8> {
wrap_record(
TLS_CONTENT_TYPE_HANDSHAKE,
&wrap_handshake(&build_client_hello_body(extra_extensions)),
)
}
#[cfg(test)]
pub(super) fn split_into_records(wire: &[u8], chunk_count: usize) -> Vec<u8> {
assert!(chunk_count >= 1, "test helper requires at least one chunk");
let payload = &wire[5..];
let chunk_len = payload.len().div_ceil(chunk_count);
let mut out = Vec::new();
for chunk in payload.chunks(chunk_len.max(1)) {
out.extend_from_slice(&wrap_record(TLS_CONTENT_TYPE_HANDSHAKE, chunk));
}
out
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn empty_input_needs_more() {
assert!(matches!(parse_client_hello(&[]), ParseOutcome::NeedMore));
}
#[test]
fn non_handshake_first_record_is_not_tls() {
let record = wrap_record(23, &[0u8; 4]);
assert!(matches!(
parse_client_hello(&record),
ParseOutcome::Reject(RejectReason::NotTls)
));
}
#[test]
fn oversized_record_length_is_malformed() {
let mut record = vec![TLS_CONTENT_TYPE_HANDSHAKE, 0x03, 0x03];
record.extend_from_slice(&(MAX_TLS_RECORD_LEN as u16 + 1).to_be_bytes());
assert!(matches!(
parse_client_hello(&record),
ParseOutcome::Reject(RejectReason::MalformedRecord)
));
}
#[test]
fn wrong_handshake_msg_type_is_malformed_handshake() {
let mut hs = vec![2u8]; hs.extend_from_slice(&[0, 0, 1]); hs.push(0);
let record = wrap_record(TLS_CONTENT_TYPE_HANDSHAKE, &hs);
assert!(matches!(
parse_client_hello(&record),
ParseOutcome::Reject(RejectReason::MalformedHandshake)
));
}
#[test]
fn one_byte_drip_needs_more_until_complete() {
let wire = build_client_hello_wire(&[encode_sni_extension("example.com")]);
for i in 0..wire.len() {
assert!(
matches!(parse_client_hello(&wire[..i]), ParseOutcome::NeedMore),
"prefix of length {i} (of {}) must NeedMore",
wire.len()
);
}
match parse_client_hello(&wire) {
ParseOutcome::ClientHello { sni, .. } => {
assert_eq!(sni.as_deref(), Some("example.com"));
}
_ => panic!("expected a complete ClientHello, got a different outcome"),
}
}
#[test]
fn byte_replay_buffer_is_untouched_by_parsing() {
let wire = build_client_hello_wire(&[
encode_sni_extension("example.com"),
encode_alpn_extension(&[b"h2", b"http/1.1"]),
]);
let before = wire.clone();
let _ = parse_client_hello(&wire);
assert_eq!(wire, before, "parsing must never mutate the input buffer");
}
#[test]
fn multi_record_client_hello_reassembles() {
let wire = build_client_hello_wire(&[
encode_sni_extension("split.example.com"),
encode_alpn_extension(&[b"h2"]),
]);
for chunks in 1..=6usize {
let split = split_into_records(&wire, chunks);
match parse_client_hello(&split) {
ParseOutcome::ClientHello { sni, alpn, .. } => {
assert_eq!(sni.as_deref(), Some("split.example.com"));
assert_eq!(alpn, vec![b"h2".to_vec()]);
}
_ => panic!("chunk count {chunks} must reassemble to a complete ClientHello"),
}
}
}
#[test]
fn non_handshake_record_mid_reassembly_is_malformed_record() {
let wire = build_client_hello_wire(&[encode_sni_extension("example.com")]);
let split = split_into_records(&wire, 3);
let first_len = u16::from_be_bytes([split[3], split[4]]) as usize;
let second_record_offset = 5 + first_len;
let mut corrupted = split.clone();
corrupted[second_record_offset] = 23; assert!(matches!(
parse_client_hello(&corrupted),
ParseOutcome::Reject(RejectReason::MalformedRecord)
));
}
#[test]
fn grease_extensions_are_skipped_by_length() {
let wire = build_client_hello_wire(&[
encode_grease_extension(),
encode_sni_extension("grease.example.com"),
encode_alpn_extension(&[b"h2"]),
encode_grease_extension(),
]);
match parse_client_hello(&wire) {
ParseOutcome::ClientHello {
sni,
alpn,
ech_present,
} => {
assert_eq!(sni.as_deref(), Some("grease.example.com"));
assert_eq!(alpn, vec![b"h2".to_vec()]);
assert!(!ech_present);
}
_ => panic!("GREASE-laden ClientHello must still parse"),
}
}
#[test]
fn alpn_preserves_client_offer_order() {
let wire = build_client_hello_wire(&[encode_alpn_extension(&[b"http/1.1", b"h2", b"foo"])]);
match parse_client_hello(&wire) {
ParseOutcome::ClientHello { alpn, .. } => {
assert_eq!(
alpn,
vec![b"http/1.1".to_vec(), b"h2".to_vec(), b"foo".to_vec()]
);
}
_ => panic!("expected a complete ClientHello"),
}
}
#[test]
fn no_sni_extension_yields_none() {
let wire = build_client_hello_wire(&[encode_alpn_extension(&[b"h2"])]);
match parse_client_hello(&wire) {
ParseOutcome::ClientHello { sni, .. } => assert_eq!(sni, None),
_ => panic!("expected a complete ClientHello"),
}
}
#[test]
fn ech_extension_presence_is_flagged() {
let wire = build_client_hello_wire(&[encode_extension(0xfe0d, &[0x00, 0x01, 0x02])]);
match parse_client_hello(&wire) {
ParseOutcome::ClientHello {
sni, ech_present, ..
} => {
assert_eq!(sni, None);
assert!(ech_present);
}
_ => panic!("expected a complete ClientHello"),
}
}
#[test]
fn no_extensions_block_yields_no_sni_no_alpn() {
let wire = build_client_hello_wire(&[]);
match parse_client_hello(&wire) {
ParseOutcome::ClientHello {
sni,
alpn,
ech_present,
} => {
assert_eq!(sni, None);
assert!(alpn.is_empty());
assert!(!ech_present);
}
_ => panic!("expected a complete ClientHello"),
}
}
#[test]
fn sni_list_len_overflowing_extension_body_is_malformed_not_panic() {
let mut data = Vec::new();
data.extend_from_slice(&512u16.to_be_bytes()); data.push(0); let wire = build_client_hello_wire(&[encode_extension(EXT_SERVER_NAME, &data)]);
assert!(matches!(
parse_client_hello(&wire),
ParseOutcome::Reject(RejectReason::MalformedHandshake)
));
}
#[test]
fn sni_name_len_overflowing_list_is_malformed_not_panic() {
let mut name_list = vec![0u8]; name_list.extend_from_slice(&512u16.to_be_bytes()); name_list.push(b'a'); let mut data = Vec::new();
data.extend_from_slice(&(name_list.len() as u16).to_be_bytes()); data.extend_from_slice(&name_list);
let wire = build_client_hello_wire(&[encode_extension(EXT_SERVER_NAME, &data)]);
assert!(matches!(
parse_client_hello(&wire),
ParseOutcome::Reject(RejectReason::MalformedHandshake)
));
}
#[test]
fn alpn_list_len_overflowing_extension_body_is_malformed_not_panic() {
let mut data = Vec::new();
data.extend_from_slice(&512u16.to_be_bytes()); data.push(2); let wire = build_client_hello_wire(&[encode_extension(EXT_ALPN, &data)]);
assert!(matches!(
parse_client_hello(&wire),
ParseOutcome::Reject(RejectReason::MalformedHandshake)
));
}
#[test]
fn extension_len_overflowing_block_is_malformed_not_panic() {
let mut lying_ext = Vec::new();
lying_ext.extend_from_slice(&EXT_SERVER_NAME.to_be_bytes());
lying_ext.extend_from_slice(&512u16.to_be_bytes()); lying_ext.extend_from_slice(&[0x00, 0x00]);
let wire = build_client_hello_wire(&[lying_ext]);
assert!(matches!(
parse_client_hello(&wire),
ParseOutcome::Reject(RejectReason::MalformedHandshake)
));
}
#[test]
fn duplicate_server_name_extension_is_rejected() {
let wire = build_client_hello_wire(&[
encode_sni_extension("first.example.com"),
encode_sni_extension("second.example.com"),
]);
assert!(matches!(
parse_client_hello(&wire),
ParseOutcome::Reject(RejectReason::MalformedHandshake)
));
}
#[test]
fn duplicate_alpn_extension_is_rejected() {
let wire = build_client_hello_wire(&[
encode_alpn_extension(&[b"h2"]),
encode_alpn_extension(&[b"http/1.1"]),
]);
assert!(matches!(
parse_client_hello(&wire),
ParseOutcome::Reject(RejectReason::MalformedHandshake)
));
}
#[test]
fn server_name_list_with_two_host_name_entries_is_rejected() {
let mut name_list = Vec::new();
name_list.push(0u8); name_list.extend_from_slice(&(b"first.example.com".len() as u16).to_be_bytes());
name_list.extend_from_slice(b"first.example.com");
name_list.push(0u8); name_list.extend_from_slice(&(b"second.example.com".len() as u16).to_be_bytes());
name_list.extend_from_slice(b"second.example.com");
let mut data = Vec::new();
data.extend_from_slice(&(name_list.len() as u16).to_be_bytes());
data.extend_from_slice(&name_list);
let wire = build_client_hello_wire(&[encode_extension(EXT_SERVER_NAME, &data)]);
assert!(matches!(
parse_client_hello(&wire),
ParseOutcome::Reject(RejectReason::MalformedHandshake)
));
}
#[test]
fn alpn_extension_with_zero_length_name_is_rejected() {
let wire = build_client_hello_wire(&[encode_alpn_extension(&[b"h2", b"", b"http/1.1"])]);
assert!(matches!(
parse_client_hello(&wire),
ParseOutcome::Reject(RejectReason::MalformedHandshake)
));
}
#[test]
fn alpn_extension_with_empty_protocol_list_is_rejected() {
let wire = build_client_hello_wire(&[encode_alpn_extension(&[])]);
assert!(matches!(
parse_client_hello(&wire),
ParseOutcome::Reject(RejectReason::MalformedHandshake)
));
}
#[test]
fn single_sni_and_alpn_extension_still_parses_unchanged() {
let wire = build_client_hello_wire(&[
encode_sni_extension("example.com"),
encode_alpn_extension(&[b"h2", b"http/1.1"]),
]);
match parse_client_hello(&wire) {
ParseOutcome::ClientHello {
sni,
alpn,
ech_present,
} => {
assert_eq!(sni.as_deref(), Some("example.com"));
assert_eq!(alpn, vec![b"h2".to_vec(), b"http/1.1".to_vec()]);
assert!(!ech_present);
}
_ => panic!("expected a complete ClientHello"),
}
}
}