pub const EXT_HEADER_SIZE: usize = 4;
pub const EXT_MIN_SIZE: usize = 4;
const MAX_EXTENSION_FIELDS: usize = 32;
pub const NTS_UNIQUE_IDENTIFIER: u16 = 0x0104;
pub const NTS_COOKIE: u16 = 0x0204;
pub const NTS_COOKIE_PLACEHOLDER: u16 = 0x0304;
pub const NTS_AUTHENTICATOR: u16 = 0x0404;
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct ExtensionField {
pub field_type: u16,
pub value: Vec<u8>,
}
impl ExtensionField {
pub fn new(field_type: u16, value: Vec<u8>) -> Self {
Self { field_type, value }
}
pub fn unique_identifier(nonce: Vec<u8>) -> Self {
Self::new(NTS_UNIQUE_IDENTIFIER, nonce)
}
pub fn cookie(cookie: Vec<u8>) -> Self {
Self::new(NTS_COOKIE, cookie)
}
pub fn cookie_placeholder(length: usize) -> Self {
Self::new(NTS_COOKIE_PLACEHOLDER, vec![0u8; length])
}
pub fn authenticator(nonce_and_ciphertext: Vec<u8>) -> Self {
Self::new(NTS_AUTHENTICATOR, nonce_and_ciphertext)
}
fn padded_value_len(&self) -> usize {
let len = self.value.len();
(len + 3) & !3
}
pub fn wire_length(&self) -> usize {
EXT_HEADER_SIZE + self.padded_value_len()
}
pub fn serialize(&self) -> Vec<u8> {
let total_len = self.wire_length();
debug_assert!(
total_len <= u16::MAX as usize,
"extension field too large for u16 length: {} bytes",
total_len
);
let wire_len = u16::try_from(total_len).unwrap_or(u16::MAX);
let mut buf = Vec::with_capacity(total_len);
buf.extend_from_slice(&self.field_type.to_be_bytes());
buf.extend_from_slice(&wire_len.to_be_bytes());
buf.extend_from_slice(&self.value);
let padding = self.padded_value_len() - self.value.len();
if padding > 0 {
buf.extend_from_slice(&vec![0u8; padding]);
}
buf
}
pub fn parse(data: &[u8]) -> Result<(Self, usize), ExtensionError> {
if data.len() < EXT_HEADER_SIZE {
return Err(ExtensionError::TooShort {
got: data.len(),
expected: EXT_HEADER_SIZE,
});
}
let field_type = u16::from_be_bytes([data[0], data[1]]);
let field_length = u16::from_be_bytes([data[2], data[3]]) as usize;
if field_length < EXT_HEADER_SIZE {
return Err(ExtensionError::InvalidLength(field_length as u16));
}
if !field_length.is_multiple_of(4) {
return Err(ExtensionError::InvalidLength(field_length as u16));
}
if data.len() < field_length {
return Err(ExtensionError::TooShort {
got: data.len(),
expected: field_length,
});
}
let value_len = field_length - EXT_HEADER_SIZE;
let value = data[EXT_HEADER_SIZE..EXT_HEADER_SIZE + value_len].to_vec();
Ok((Self { field_type, value }, field_length))
}
pub fn parse_all(data: &[u8]) -> Result<Vec<Self>, ExtensionError> {
let mut fields = Vec::new();
let mut offset = 0;
while offset < data.len() {
if data.len() - offset < EXT_HEADER_SIZE {
break;
}
if fields.len() >= MAX_EXTENSION_FIELDS {
break;
}
let (field, consumed) = Self::parse(&data[offset..])?;
fields.push(field);
offset += consumed;
}
Ok(fields)
}
}
#[derive(Debug, thiserror::Error)]
pub enum ExtensionError {
#[error("extension field too short: got {got} bytes, expected at least {expected}")]
TooShort { got: usize, expected: usize },
#[error("invalid extension field length: {0}")]
InvalidLength(u16),
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn roundtrip_unique_identifier() {
let nonce = vec![0x01, 0x02, 0x03, 0x04, 0x05, 0x06, 0x07, 0x08];
let ext = ExtensionField::unique_identifier(nonce.clone());
let bytes = ext.serialize();
assert_eq!(bytes.len(), 12);
assert_eq!(&bytes[0..2], &NTS_UNIQUE_IDENTIFIER.to_be_bytes());
assert_eq!(&bytes[2..4], &12u16.to_be_bytes());
let (parsed, consumed) = ExtensionField::parse(&bytes).unwrap();
assert_eq!(consumed, 12);
assert_eq!(parsed.field_type, NTS_UNIQUE_IDENTIFIER);
assert_eq!(parsed.value, nonce);
}
#[test]
fn roundtrip_with_padding() {
let ext = ExtensionField::new(0x1234, vec![1, 2, 3, 4, 5]);
let bytes = ext.serialize();
assert_eq!(bytes.len(), 12);
assert_eq!(bytes[9], 0);
assert_eq!(bytes[10], 0);
assert_eq!(bytes[11], 0);
let (parsed, consumed) = ExtensionField::parse(&bytes).unwrap();
assert_eq!(consumed, 12);
assert_eq!(parsed.value.len(), 8); assert_eq!(&parsed.value[..5], &[1, 2, 3, 4, 5]);
}
#[test]
fn roundtrip_cookie() {
let cookie = vec![0xDE, 0xAD, 0xBE, 0xEF, 0xCA, 0xFE, 0xBA, 0xBE];
let ext = ExtensionField::cookie(cookie.clone());
let bytes = ext.serialize();
let (parsed, _) = ExtensionField::parse(&bytes).unwrap();
assert_eq!(parsed.field_type, NTS_COOKIE);
assert_eq!(parsed.value, cookie);
}
#[test]
fn roundtrip_cookie_placeholder() {
let ext = ExtensionField::cookie_placeholder(64);
let bytes = ext.serialize();
let (parsed, _) = ExtensionField::parse(&bytes).unwrap();
assert_eq!(parsed.field_type, NTS_COOKIE_PLACEHOLDER);
assert_eq!(parsed.value.len(), 64);
assert!(parsed.value.iter().all(|&b| b == 0));
}
#[test]
fn roundtrip_authenticator() {
let auth_data = vec![0xAA; 48]; let ext = ExtensionField::authenticator(auth_data.clone());
let bytes = ext.serialize();
let (parsed, _) = ExtensionField::parse(&bytes).unwrap();
assert_eq!(parsed.field_type, NTS_AUTHENTICATOR);
assert_eq!(parsed.value, auth_data);
}
#[test]
fn parse_too_short() {
let data = [0u8; 2];
assert!(ExtensionField::parse(&data).is_err());
}
#[test]
fn parse_invalid_length_too_small() {
let data = [0x01, 0x04, 0x00, 0x02];
assert!(ExtensionField::parse(&data).is_err());
}
#[test]
fn parse_invalid_length_not_aligned() {
let data = [0x01, 0x04, 0x00, 0x05, 0x00];
assert!(ExtensionField::parse(&data).is_err());
}
#[test]
fn parse_length_exceeds_data() {
let data = [0x01, 0x04, 0x00, 0x08];
assert!(ExtensionField::parse(&data).is_err());
}
#[test]
fn parse_all_multiple_fields() {
let uid = ExtensionField::unique_identifier(vec![1, 2, 3, 4, 5, 6, 7, 8]);
let cookie = ExtensionField::cookie(vec![0xAA; 16]);
let auth = ExtensionField::authenticator(vec![0xBB; 32]);
let mut data = Vec::new();
data.extend_from_slice(&uid.serialize());
data.extend_from_slice(&cookie.serialize());
data.extend_from_slice(&auth.serialize());
let fields = ExtensionField::parse_all(&data).unwrap();
assert_eq!(fields.len(), 3);
assert_eq!(fields[0].field_type, NTS_UNIQUE_IDENTIFIER);
assert_eq!(fields[1].field_type, NTS_COOKIE);
assert_eq!(fields[2].field_type, NTS_AUTHENTICATOR);
}
#[test]
fn wire_length_calculation() {
let ext = ExtensionField::new(0, vec![0; 8]);
assert_eq!(ext.wire_length(), 12);
let ext = ExtensionField::new(0, vec![0; 5]);
assert_eq!(ext.wire_length(), 12);
let ext = ExtensionField::new(0, vec![0; 1]);
assert_eq!(ext.wire_length(), 8);
let ext = ExtensionField::new(0, vec![]);
assert_eq!(ext.wire_length(), 4);
}
#[test]
fn empty_extension_field() {
let ext = ExtensionField::new(0x1234, vec![]);
let bytes = ext.serialize();
assert_eq!(bytes.len(), 4);
assert_eq!(&bytes[2..4], &4u16.to_be_bytes());
let (parsed, consumed) = ExtensionField::parse(&bytes).unwrap();
assert_eq!(consumed, 4);
assert_eq!(parsed.field_type, 0x1234);
assert!(parsed.value.is_empty());
}
#[test]
fn extension_type_constants() {
assert_eq!(NTS_UNIQUE_IDENTIFIER, 0x0104);
assert_eq!(NTS_COOKIE, 0x0204);
assert_eq!(NTS_COOKIE_PLACEHOLDER, 0x0304);
assert_eq!(NTS_AUTHENTICATOR, 0x0404);
}
}