const CONTENT_TYPE_HANDSHAKE: u8 = 0x16;
const HANDSHAKE_TYPE_CLIENT_HELLO: u8 = 0x01;
const EXTENSION_SERVER_NAME: u16 = 0x0000;
const NAME_TYPE_HOST_NAME: u8 = 0x00;
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum SniPeek {
Found(String),
Absent,
Incomplete,
NotTls,
}
struct Reader<'a> {
buf: &'a [u8],
pos: usize,
}
impl<'a> Reader<'a> {
fn new(buf: &'a [u8]) -> Self {
Self { buf, pos: 0 }
}
fn remaining(&self) -> usize {
self.buf.len().saturating_sub(self.pos)
}
fn u8(&mut self) -> Option<u8> {
let b = *self.buf.get(self.pos)?;
self.pos += 1;
Some(b)
}
fn u16(&mut self) -> Option<u16> {
let hi = self.u8()? as u16;
let lo = self.u8()? as u16;
Some((hi << 8) | lo)
}
fn u24(&mut self) -> Option<u32> {
let a = self.u8()? as u32;
let b = self.u8()? as u32;
let c = self.u8()? as u32;
Some((a << 16) | (b << 8) | c)
}
fn bytes(&mut self, n: usize) -> Option<&'a [u8]> {
let end = self.pos.checked_add(n)?;
let slice = self.buf.get(self.pos..end)?;
self.pos = end;
Some(slice)
}
fn skip_prefixed(&mut self, len_bytes: usize) -> Option<()> {
let len = match len_bytes {
1 => self.u8()? as usize,
2 => self.u16()? as usize,
_ => return None,
};
self.bytes(len).map(|_| ())
}
}
pub fn parse_sni(buf: &[u8]) -> SniPeek {
let mut records = Reader::new(buf);
let mut handshake: Vec<u8> = Vec::new();
loop {
if records.remaining() == 0 {
break;
}
if records.remaining() < 5 {
if !handshake.is_empty() {
break;
}
return match records.u8() {
Some(CONTENT_TYPE_HANDSHAKE) => SniPeek::Incomplete,
Some(_) => SniPeek::NotTls,
None => SniPeek::Incomplete,
};
}
let Some(content_type) = records.u8() else {
return SniPeek::Incomplete;
};
if content_type != CONTENT_TYPE_HANDSHAKE {
if !handshake.is_empty() {
break;
}
return SniPeek::NotTls;
}
if records.u16().is_none() {
return SniPeek::Incomplete;
}
let Some(len) = records.u16() else {
return SniPeek::Incomplete;
};
match records.bytes(len as usize) {
Some(payload) => handshake.extend_from_slice(payload),
None => break,
}
}
parse_handshake(&handshake)
}
fn parse_handshake(handshake: &[u8]) -> SniPeek {
let mut r = Reader::new(handshake);
match r.u8() {
Some(HANDSHAKE_TYPE_CLIENT_HELLO) => {}
Some(_) => return SniPeek::NotTls,
None => return SniPeek::Incomplete,
}
let Some(body_len) = r.u24() else {
return SniPeek::Incomplete;
};
if r.remaining() < body_len as usize {
return SniPeek::Incomplete;
}
if r.bytes(34).is_none() {
return SniPeek::Incomplete;
}
if r.skip_prefixed(1).is_none() || r.skip_prefixed(2).is_none() || r.skip_prefixed(1).is_none()
{
return SniPeek::Incomplete;
}
if r.remaining() == 0 {
return SniPeek::Absent;
}
let Some(ext_total) = r.u16() else {
return SniPeek::Incomplete;
};
let Some(ext_bytes) = r.bytes(ext_total as usize) else {
return SniPeek::Incomplete;
};
let mut ext = Reader::new(ext_bytes);
while ext.remaining() > 0 {
let (Some(ext_type), Some(ext_len)) = (ext.u16(), ext.u16()) else {
return SniPeek::Incomplete;
};
let Some(body) = ext.bytes(ext_len as usize) else {
return SniPeek::Incomplete;
};
if ext_type == EXTENSION_SERVER_NAME {
return parse_server_name_list(body);
}
}
SniPeek::Absent
}
fn parse_server_name_list(body: &[u8]) -> SniPeek {
let mut r = Reader::new(body);
let Some(list_len) = r.u16() else {
return SniPeek::Incomplete;
};
let Some(list) = r.bytes(list_len as usize) else {
return SniPeek::Incomplete;
};
let mut names = Reader::new(list);
while names.remaining() > 0 {
let (Some(name_type), Some(name_len)) = (names.u8(), names.u16()) else {
return SniPeek::Incomplete;
};
let Some(name) = names.bytes(name_len as usize) else {
return SniPeek::Incomplete;
};
if name_type != NAME_TYPE_HOST_NAME {
continue;
}
return match std::str::from_utf8(name) {
Ok(host) if is_routable_host_name(host) => {
SniPeek::Found(host.trim_end_matches('.').to_ascii_lowercase())
}
_ => SniPeek::Absent,
};
}
SniPeek::Absent
}
fn is_routable_host_name(host: &str) -> bool {
!host.is_empty()
&& host.len() <= 253
&& host
.bytes()
.all(|b| b.is_ascii_alphanumeric() || matches!(b, b'-' | b'.' | b'_'))
}
#[cfg(test)]
mod tests {
use super::*;
fn server_name_ext(host: &str) -> Vec<u8> {
let mut entry = vec![NAME_TYPE_HOST_NAME];
entry.extend_from_slice(&(host.len() as u16).to_be_bytes());
entry.extend_from_slice(host.as_bytes());
let mut body = (entry.len() as u16).to_be_bytes().to_vec();
body.extend_from_slice(&entry);
body
}
fn extension(ext_type: u16, body: &[u8]) -> Vec<u8> {
let mut out = ext_type.to_be_bytes().to_vec();
out.extend_from_slice(&(body.len() as u16).to_be_bytes());
out.extend_from_slice(body);
out
}
fn client_hello(extensions: Vec<Vec<u8>>) -> Vec<u8> {
let mut body = Vec::new();
body.extend_from_slice(&[0x03, 0x03]); body.extend_from_slice(&[0x11; 32]); body.push(0); body.extend_from_slice(&[0x00, 0x02, 0x13, 0x01]); body.extend_from_slice(&[0x01, 0x00]);
let ext_bytes: Vec<u8> = extensions.concat();
body.extend_from_slice(&(ext_bytes.len() as u16).to_be_bytes());
body.extend_from_slice(&ext_bytes);
let mut msg = vec![HANDSHAKE_TYPE_CLIENT_HELLO];
let len = body.len() as u32;
msg.extend_from_slice(&[(len >> 16) as u8, (len >> 8) as u8, len as u8]);
msg.extend_from_slice(&body);
msg
}
fn record(payload: &[u8]) -> Vec<u8> {
let mut out = vec![CONTENT_TYPE_HANDSHAKE, 0x03, 0x01];
out.extend_from_slice(&(payload.len() as u16).to_be_bytes());
out.extend_from_slice(payload);
out
}
fn records_of(payload: &[u8], chunk: usize) -> Vec<u8> {
payload.chunks(chunk).flat_map(record).collect()
}
#[test]
fn test_parse_sni_finds_hostname() {
let hello = client_hello(vec![extension(
EXTENSION_SERVER_NAME,
&server_name_ext("api.localhost"),
)]);
assert_eq!(
parse_sni(&record(&hello)),
SniPeek::Found("api.localhost".to_string())
);
}
#[test]
fn test_parse_sni_lowercases_and_strips_root_dot() {
let hello = client_hello(vec![extension(
EXTENSION_SERVER_NAME,
&server_name_ext("API.LocalHost."),
)]);
assert_eq!(
parse_sni(&record(&hello)),
SniPeek::Found("api.localhost".to_string())
);
}
#[test]
fn test_parse_sni_skips_other_extensions() {
let hello = client_hello(vec![
extension(0x002b, &[0x02, 0x03, 0x04]), extension(EXTENSION_SERVER_NAME, &server_name_ext("app.localhost")),
extension(0x0010, &[0x00, 0x03, 0x02, b'h', b'2']), ]);
assert_eq!(
parse_sni(&record(&hello)),
SniPeek::Found("app.localhost".to_string())
);
}
#[test]
fn test_parse_sni_absent_without_extension() {
let hello = client_hello(vec![extension(0x002b, &[0x02, 0x03, 0x04])]);
assert_eq!(parse_sni(&record(&hello)), SniPeek::Absent);
}
#[test]
fn test_parse_sni_absent_with_empty_extension_list() {
let hello = client_hello(vec![]);
assert_eq!(parse_sni(&record(&hello)), SniPeek::Absent);
}
#[test]
fn test_parse_sni_absent_when_extensions_omitted_entirely() {
let full = client_hello(vec![]);
let mut body = full[4..].to_vec();
body.truncate(body.len() - 2);
let mut msg = vec![HANDSHAKE_TYPE_CLIENT_HELLO];
let len = body.len() as u32;
msg.extend_from_slice(&[(len >> 16) as u8, (len >> 8) as u8, len as u8]);
msg.extend_from_slice(&body);
assert_eq!(parse_sni(&record(&msg)), SniPeek::Absent);
}
#[test]
fn test_parse_sni_across_fragmented_records() {
let hello = client_hello(vec![extension(
EXTENSION_SERVER_NAME,
&server_name_ext("split.localhost"),
)]);
for chunk in [1usize, 5, 17, 64] {
assert_eq!(
parse_sni(&records_of(&hello, chunk)),
SniPeek::Found("split.localhost".to_string()),
"fragmented into {chunk}-byte records"
);
}
}
#[test]
fn test_parse_sni_incomplete_prefixes_of_a_fragmented_hello() {
let hello = client_hello(vec![extension(
EXTENSION_SERVER_NAME,
&server_name_ext("split.localhost"),
)]);
let wire = records_of(&hello, 7);
for n in 1..wire.len() {
assert_eq!(
parse_sni(&wire[..n]),
SniPeek::Incomplete,
"prefix of {n} bytes should be incomplete"
);
}
assert_eq!(
parse_sni(&wire),
SniPeek::Found("split.localhost".to_string())
);
}
#[test]
fn test_parse_sni_incomplete_on_truncated_record() {
let hello = client_hello(vec![extension(
EXTENSION_SERVER_NAME,
&server_name_ext("api.localhost"),
)]);
let wire = record(&hello);
assert_eq!(parse_sni(&wire[..wire.len() - 10]), SniPeek::Incomplete);
assert_eq!(parse_sni(&[]), SniPeek::Incomplete);
assert_eq!(parse_sni(&[CONTENT_TYPE_HANDSHAKE]), SniPeek::Incomplete);
}
#[test]
fn test_parse_sni_finds_hostname_before_a_non_handshake_record() {
let hello = client_hello(vec![extension(
EXTENSION_SERVER_NAME,
&server_name_ext("api.localhost"),
)]);
let mut wire = record(&hello);
wire.extend_from_slice(&[0x14, 0x03, 0x03, 0x00, 0x01, 0x01]);
assert_eq!(
parse_sni(&wire),
SniPeek::Found("api.localhost".to_string())
);
wire.extend_from_slice(&[0x17, 0x03, 0x03, 0x00, 0x03, 0xaa, 0xbb, 0xcc]);
assert_eq!(
parse_sni(&wire),
SniPeek::Found("api.localhost".to_string())
);
let mut fragmented = records_of(&hello, 9);
fragmented.extend_from_slice(&[0x14, 0x03, 0x03, 0x00, 0x01, 0x01]);
assert_eq!(
parse_sni(&fragmented),
SniPeek::Found("api.localhost".to_string())
);
}
#[test]
fn test_parse_sni_ignores_a_partial_trailing_record_header() {
let hello = client_hello(vec![extension(
EXTENSION_SERVER_NAME,
&server_name_ext("api.localhost"),
)]);
let mut wire = record(&hello);
wire.extend_from_slice(&[0x14, 0x03]);
assert_eq!(
parse_sni(&wire),
SniPeek::Found("api.localhost".to_string())
);
}
#[test]
fn test_parse_sni_rejects_non_tls() {
assert_eq!(parse_sni(b"GET / HTTP/1.1\r\n"), SniPeek::NotTls);
assert_eq!(
parse_sni(&[0x15, 0x03, 0x01, 0x00, 0x02, 0x01, 0x00]),
SniPeek::NotTls
);
assert_eq!(parse_sni(b"G"), SniPeek::NotTls);
}
#[test]
fn test_parse_sni_rejects_non_client_hello_handshake() {
let mut msg = vec![0x02, 0x00, 0x00, 0x02, 0x03, 0x03];
msg = record(&msg);
assert_eq!(parse_sni(&msg), SniPeek::NotTls);
}
#[test]
fn test_parse_sni_absent_for_non_ascii_hostname() {
let hello = client_hello(vec![extension(
EXTENSION_SERVER_NAME,
&server_name_ext("café.localhost"),
)]);
assert_eq!(parse_sni(&record(&hello)), SniPeek::Absent);
}
#[test]
fn test_parse_sni_absent_for_control_characters() {
for host in [
"api.localhost\n",
"api\r\nlocalhost",
"api.localhost\u{0}",
"api\tlocalhost",
"api localhost",
"api.localhost\u{7f}",
"\u{1b}[31mapi.localhost",
] {
let hello = client_hello(vec![extension(
EXTENSION_SERVER_NAME,
&server_name_ext(host),
)]);
assert_eq!(
parse_sni(&record(&hello)),
SniPeek::Absent,
"{host:?} must not be routable"
);
}
}
#[test]
fn test_parse_sni_accepts_host_name_characters() {
for host in [
"api.localhost",
"api-2.my_project.localhost",
"API.LOCALHOST",
] {
let hello = client_hello(vec![extension(
EXTENSION_SERVER_NAME,
&server_name_ext(host),
)]);
assert_eq!(
parse_sni(&record(&hello)),
SniPeek::Found(host.to_ascii_lowercase()),
"{host:?} must stay routable"
);
}
}
#[test]
fn test_parse_sni_absent_for_oversized_hostname() {
let host = format!("{}.localhost", "a".repeat(250));
let hello = client_hello(vec![extension(
EXTENSION_SERVER_NAME,
&server_name_ext(&host),
)]);
assert_eq!(parse_sni(&record(&hello)), SniPeek::Absent);
}
#[test]
fn test_parse_sni_absent_for_empty_hostname() {
let hello = client_hello(vec![extension(EXTENSION_SERVER_NAME, &server_name_ext(""))]);
assert_eq!(parse_sni(&record(&hello)), SniPeek::Absent);
}
#[test]
fn test_parse_sni_skips_unknown_name_types() {
let mut entries = vec![0x7f, 0x00, 0x02, 0xaa, 0xbb];
let host = "second.localhost";
entries.push(NAME_TYPE_HOST_NAME);
entries.extend_from_slice(&(host.len() as u16).to_be_bytes());
entries.extend_from_slice(host.as_bytes());
let mut body = (entries.len() as u16).to_be_bytes().to_vec();
body.extend_from_slice(&entries);
let hello = client_hello(vec![extension(EXTENSION_SERVER_NAME, &body)]);
assert_eq!(
parse_sni(&record(&hello)),
SniPeek::Found("second.localhost".to_string())
);
}
#[test]
fn test_parse_sni_does_not_panic_on_arbitrary_bytes() {
let hello = client_hello(vec![extension(
EXTENSION_SERVER_NAME,
&server_name_ext("api.localhost"),
)]);
let wire = record(&hello);
for i in 0..wire.len() {
for corruption in [0x00u8, 0xff, 0x7f] {
let mut bad = wire.clone();
bad[i] = corruption;
let _ = parse_sni(&bad);
}
}
}
}