use zeroize::{Zeroize, ZeroizeOnDrop};
use crate::records::{NtsKeRecord, RecordType};
use crate::{
AEAD_AES_SIV_CMAC_256, AEAD_AES_SIV_CMAC_256_KEYLEN, NTS_NEXT_PROTOCOL_NTPV4,
NTS_TLS_EXPORTER_LABEL, NtsError,
};
#[derive(Clone, Zeroize, ZeroizeOnDrop)]
pub struct NtsKeResult {
pub c2s_key: Vec<u8>,
pub s2c_key: Vec<u8>,
pub cookies: Vec<Vec<u8>>,
pub aead_algorithm: u16,
pub server: Option<String>,
pub port: Option<u16>,
}
impl std::fmt::Debug for NtsKeResult {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("NtsKeResult")
.field("c2s_key", &"[REDACTED]")
.field("s2c_key", &"[REDACTED]")
.field("cookies", &format!("[{} cookies]", self.cookies.len()))
.field("aead_algorithm", &self.aead_algorithm)
.field("server", &self.server)
.field("port", &self.port)
.finish()
}
}
pub fn derive_keys(
tls_connection: &impl ExporterInterface,
algorithm: u16,
) -> Result<(Vec<u8>, Vec<u8>), NtsError> {
let key_length = key_length_for_algorithm(algorithm)?;
let mut context = Vec::with_capacity(4);
context.extend_from_slice(&NTS_NEXT_PROTOCOL_NTPV4.to_be_bytes());
context.extend_from_slice(&algorithm.to_be_bytes());
let export_length = key_length * 2;
let mut output = vec![0u8; export_length];
tls_connection
.export_keying_material(&mut output, NTS_TLS_EXPORTER_LABEL, Some(&context))
.map_err(|e| NtsError::Tls(format!("TLS exporter failed: {}", e)))?;
let c2s_key = output[..key_length].to_vec();
let s2c_key = output[key_length..].to_vec();
output.zeroize();
Ok((c2s_key, s2c_key))
}
pub trait ExporterInterface {
fn export_keying_material(
&self,
output: &mut [u8],
label: &str,
context: Option<&[u8]>,
) -> Result<(), ExporterError>;
}
#[derive(Debug, thiserror::Error)]
pub enum ExporterError {
#[error("TLS exporter not available (handshake not complete?)")]
NotAvailable,
#[error("export failed: {0}")]
Failed(String),
}
pub fn key_length_for_algorithm(algorithm: u16) -> Result<usize, NtsError> {
match algorithm {
AEAD_AES_SIV_CMAC_256 => Ok(AEAD_AES_SIV_CMAC_256_KEYLEN),
other => Err(NtsError::UnsupportedAlgorithm(other)),
}
}
pub fn build_client_request() -> Vec<u8> {
let records = [
NtsKeRecord::next_protocol(&[NTS_NEXT_PROTOCOL_NTPV4]),
NtsKeRecord::aead_algorithm(&[AEAD_AES_SIV_CMAC_256]),
NtsKeRecord::end_of_message(),
];
let mut data = Vec::new();
for record in &records {
data.extend_from_slice(&record.serialize());
}
data
}
pub fn build_server_response(
algorithm: u16,
cookies: &[Vec<u8>],
server: Option<&str>,
port: Option<u16>,
) -> Vec<u8> {
let mut records = Vec::new();
records.push(NtsKeRecord::next_protocol(&[NTS_NEXT_PROTOCOL_NTPV4]));
records.push(NtsKeRecord::aead_algorithm(&[algorithm]));
for cookie in cookies {
records.push(NtsKeRecord::new_cookie(cookie.clone()));
}
if let Some(srv) = server {
records.push(NtsKeRecord::server_negotiation(srv));
}
if let Some(p) = port {
records.push(NtsKeRecord::port_negotiation(p));
}
records.push(NtsKeRecord::end_of_message());
let mut data = Vec::new();
for record in &records {
data.extend_from_slice(&record.serialize());
}
data
}
pub fn parse_server_response(data: &[u8]) -> Result<ParsedKeResponse, NtsError> {
let records = NtsKeRecord::parse_all(data)?;
let mut protocol = None;
let mut algorithm = None;
let mut cookies = Vec::new();
let mut server = None;
let mut port = None;
for record in &records {
match record.record_type {
RecordType::NextProtocol => {
let ids = record.protocol_ids();
if ids.contains(&NTS_NEXT_PROTOCOL_NTPV4) {
protocol = Some(NTS_NEXT_PROTOCOL_NTPV4);
} else if let Some(&first) = ids.first() {
return Err(NtsError::UnsupportedProtocol(first));
}
}
RecordType::AeadAlgorithm => {
let algos = record.algorithm_ids();
if let Some(&alg) = algos.first() {
algorithm = Some(alg);
}
}
RecordType::NewCookieForNtpv4 => {
cookies.push(record.body.clone());
}
RecordType::NtpV4ServerNegotiation => {
server = Some(String::from_utf8(record.body.clone()).map_err(|_| {
NtsError::InvalidCookie("invalid server name UTF-8".to_string())
})?);
}
RecordType::NtpV4PortNegotiation => {
if record.body.len() >= 2 {
port = Some(u16::from_be_bytes([record.body[0], record.body[1]]));
}
}
RecordType::Error => {
if record.body.len() >= 2 {
let code = u16::from_be_bytes([record.body[0], record.body[1]]);
return Err(NtsError::KeError(code));
}
}
RecordType::Warning => {
if record.body.len() >= 2 {
let code = u16::from_be_bytes([record.body[0], record.body[1]]);
tracing::warn!(code, "NTS-KE warning received");
}
}
RecordType::EndOfMessage => break,
}
}
if protocol.is_none() {
return Err(NtsError::MissingRecord("NextProtocol"));
}
let aead_algorithm = algorithm.ok_or(NtsError::MissingRecord("AeadAlgorithm"))?;
if cookies.is_empty() {
return Err(NtsError::NoCookies);
}
Ok(ParsedKeResponse {
aead_algorithm,
cookies,
server,
port,
})
}
pub fn parse_client_request(data: &[u8]) -> Result<ParsedKeRequest, NtsError> {
let records = NtsKeRecord::parse_all(data)?;
let mut protocol = None;
let mut algorithms = Vec::new();
for record in &records {
match record.record_type {
RecordType::NextProtocol => {
let ids = record.protocol_ids();
if ids.contains(&NTS_NEXT_PROTOCOL_NTPV4) {
protocol = Some(NTS_NEXT_PROTOCOL_NTPV4);
} else if let Some(&first) = ids.first() {
return Err(NtsError::UnsupportedProtocol(first));
}
}
RecordType::AeadAlgorithm => {
algorithms.extend(record.algorithm_ids());
}
RecordType::EndOfMessage => break,
_ => {
if record.critical {
return Err(NtsError::UnknownRecordType(record.record_type as u16));
}
}
}
}
if protocol.is_none() {
return Err(NtsError::MissingRecord("NextProtocol"));
}
if algorithms.is_empty() {
return Err(NtsError::MissingRecord("AeadAlgorithm"));
}
Ok(ParsedKeRequest { algorithms })
}
#[derive(Debug, Clone)]
pub struct ParsedKeRequest {
pub algorithms: Vec<u16>,
}
#[derive(Debug, Clone)]
pub struct ParsedKeResponse {
pub aead_algorithm: u16,
pub cookies: Vec<Vec<u8>>,
pub server: Option<String>,
pub port: Option<u16>,
}
pub fn select_algorithm(offered: &[u16]) -> Option<u16> {
if offered.contains(&AEAD_AES_SIV_CMAC_256) {
Some(AEAD_AES_SIV_CMAC_256)
} else {
None
}
}
pub fn generate_cookies(
cookie_jar: &crate::cookie::CookieJar,
c2s_key: &[u8],
s2c_key: &[u8],
algorithm: u16,
count: usize,
) -> Vec<Vec<u8>> {
(0..count)
.map(|_| cookie_jar.make_cookie(c2s_key, s2c_key, algorithm))
.collect()
}
#[cfg(test)]
mod tests {
use super::*;
use crate::DEFAULT_COOKIE_COUNT;
use rand::Rng;
fn random_key() -> [u8; AEAD_AES_SIV_CMAC_256_KEYLEN] {
let mut key = [0u8; AEAD_AES_SIV_CMAC_256_KEYLEN];
rand::rng().fill_bytes(&mut key);
key
}
#[test]
fn build_and_parse_client_request() {
let data = build_client_request();
let parsed = parse_client_request(&data).unwrap();
assert_eq!(parsed.algorithms, vec![AEAD_AES_SIV_CMAC_256]);
}
#[test]
fn build_and_parse_server_response() {
let cookies = vec![vec![1, 2, 3, 4], vec![5, 6, 7, 8]];
let data = build_server_response(
AEAD_AES_SIV_CMAC_256,
&cookies,
Some("ntp.example.com"),
Some(123),
);
let parsed = parse_server_response(&data).unwrap();
assert_eq!(parsed.aead_algorithm, AEAD_AES_SIV_CMAC_256);
assert_eq!(parsed.cookies, cookies);
assert_eq!(parsed.server.as_deref(), Some("ntp.example.com"));
assert_eq!(parsed.port, Some(123));
}
#[test]
fn server_response_without_optional_fields() {
let cookies = vec![vec![1, 2, 3]];
let data = build_server_response(AEAD_AES_SIV_CMAC_256, &cookies, None, None);
let parsed = parse_server_response(&data).unwrap();
assert!(parsed.server.is_none());
assert!(parsed.port.is_none());
}
#[test]
fn server_response_no_cookies_fails() {
let data = build_server_response(AEAD_AES_SIV_CMAC_256, &[], None, None);
assert!(matches!(
parse_server_response(&data),
Err(NtsError::NoCookies)
));
}
#[test]
fn select_algorithm_supported() {
assert_eq!(
select_algorithm(&[AEAD_AES_SIV_CMAC_256]),
Some(AEAD_AES_SIV_CMAC_256)
);
assert_eq!(
select_algorithm(&[99, AEAD_AES_SIV_CMAC_256, 100]),
Some(AEAD_AES_SIV_CMAC_256)
);
}
#[test]
fn select_algorithm_unsupported() {
assert_eq!(select_algorithm(&[99, 100]), None);
assert_eq!(select_algorithm(&[]), None);
}
#[test]
fn generate_cookies_count() {
let mut jar = crate::cookie::CookieJar::new(random_key());
let c2s = random_key();
let s2c = random_key();
let cookies = generate_cookies(
&jar,
&c2s,
&s2c,
AEAD_AES_SIV_CMAC_256,
DEFAULT_COOKIE_COUNT,
);
assert_eq!(cookies.len(), DEFAULT_COOKIE_COUNT);
for cookie in &cookies {
let contents = jar.open_cookie(cookie).unwrap();
assert_eq!(contents.c2s_key, c2s);
assert_eq!(contents.s2c_key, s2c);
assert_eq!(contents.algorithm, AEAD_AES_SIV_CMAC_256);
}
}
#[test]
fn key_length_for_known_algorithm() {
assert_eq!(
key_length_for_algorithm(AEAD_AES_SIV_CMAC_256).unwrap(),
AEAD_AES_SIV_CMAC_256_KEYLEN
);
}
#[test]
fn key_length_for_unknown_algorithm() {
assert!(key_length_for_algorithm(999).is_err());
}
struct MockExporter {
material: Vec<u8>,
}
impl MockExporter {
fn new(len: usize) -> Self {
let mut material = vec![0u8; len];
rand::rng().fill_bytes(&mut material);
Self { material }
}
}
impl ExporterInterface for MockExporter {
fn export_keying_material(
&self,
output: &mut [u8],
_label: &str,
_context: Option<&[u8]>,
) -> Result<(), ExporterError> {
if output.len() > self.material.len() {
return Err(ExporterError::Failed("output too large".to_string()));
}
output.copy_from_slice(&self.material[..output.len()]);
Ok(())
}
}
#[test]
fn derive_keys_splits_correctly() {
let key_len = AEAD_AES_SIV_CMAC_256_KEYLEN;
let exporter = MockExporter::new(key_len * 2);
let (c2s, s2c) = derive_keys(&exporter, AEAD_AES_SIV_CMAC_256).unwrap();
assert_eq!(c2s.len(), key_len);
assert_eq!(s2c.len(), key_len);
assert_eq!(&c2s, &exporter.material[..key_len]);
assert_eq!(&s2c, &exporter.material[key_len..]);
assert_ne!(c2s, s2c);
}
#[test]
fn derive_keys_unsupported_algorithm() {
let exporter = MockExporter::new(64);
assert!(derive_keys(&exporter, 999).is_err());
}
#[test]
fn full_ke_exchange_simulation() {
let client_request = build_client_request();
let parsed_req = parse_client_request(&client_request).unwrap();
let algorithm =
select_algorithm(&parsed_req.algorithms).expect("should find supported algorithm");
assert_eq!(algorithm, AEAD_AES_SIV_CMAC_256);
let exporter = MockExporter::new(AEAD_AES_SIV_CMAC_256_KEYLEN * 2);
let (c2s_key, s2c_key) = derive_keys(&exporter, algorithm).unwrap();
let mut jar = crate::cookie::CookieJar::new(random_key());
let cookies = generate_cookies(&jar, &c2s_key, &s2c_key, algorithm, DEFAULT_COOKIE_COUNT);
let server_response = build_server_response(algorithm, &cookies, None, None);
let parsed_resp = parse_server_response(&server_response).unwrap();
assert_eq!(parsed_resp.aead_algorithm, algorithm);
assert_eq!(parsed_resp.cookies.len(), DEFAULT_COOKIE_COUNT);
for cookie in &parsed_resp.cookies {
let contents = jar.open_cookie(cookie).unwrap();
assert_eq!(contents.c2s_key, c2s_key);
assert_eq!(contents.s2c_key, s2c_key);
assert_eq!(contents.algorithm, algorithm);
}
}
#[test]
fn error_record_in_response() {
let mut data = Vec::new();
data.extend_from_slice(&NtsKeRecord::next_protocol(&[NTS_NEXT_PROTOCOL_NTPV4]).serialize());
data.extend_from_slice(&NtsKeRecord::error(1).serialize());
data.extend_from_slice(&NtsKeRecord::end_of_message().serialize());
assert!(matches!(
parse_server_response(&data),
Err(NtsError::KeError(1))
));
}
}