use std::time::Instant;
use hkdf::Hkdf;
use sha2::Sha256;
use crate::core::{
CryptoError, MAX_EPOCH, OLD_KEY_RETENTION, REJECT_AFTER_MESSAGES, REJECT_AFTER_TIME,
REKEY_AFTER_MESSAGES, REKEY_AFTER_TIME,
};
use zeroize::Zeroize;
use super::{SessionKey, SESSION_KEY_SIZE};
#[derive(Debug)]
pub struct RekeyState {
epoch: u32,
epoch_start: Instant,
send_count: u64,
recv_count: u64,
}
impl RekeyState {
pub fn new() -> Self {
Self {
epoch: 0,
epoch_start: Instant::now(),
send_count: 0,
recv_count: 0,
}
}
pub fn epoch(&self) -> u32 {
self.epoch
}
pub fn send_count(&self) -> u64 {
self.send_count
}
pub fn recv_count(&self) -> u64 {
self.recv_count
}
pub fn increment_send(&mut self) -> Result<u64, CryptoError> {
if self.send_count == REJECT_AFTER_MESSAGES {
return Err(CryptoError::CounterExhaustion);
}
let counter = self.send_count;
self.send_count += 1;
Ok(counter)
}
pub fn record_recv(&mut self, counter: u64) {
if counter >= self.recv_count {
self.recv_count = counter + 1;
}
}
pub fn should_rekey(&self) -> bool {
let time_exceeded = self.epoch_start.elapsed() >= REKEY_AFTER_TIME;
let messages_exceeded = self.send_count >= REKEY_AFTER_MESSAGES;
time_exceeded || messages_exceeded
}
pub fn keys_expired(&self) -> bool {
self.epoch_start.elapsed() >= REJECT_AFTER_TIME
}
pub fn can_rekey(&self) -> bool {
self.epoch < MAX_EPOCH
}
pub fn advance_epoch(&mut self) -> Result<(), CryptoError> {
if self.epoch == MAX_EPOCH {
return Err(CryptoError::EpochExhaustion);
}
self.epoch += 1;
self.epoch_start = Instant::now();
self.send_count = 0;
self.recv_count = 0;
Ok(())
}
}
impl Default for RekeyState {
fn default() -> Self {
Self::new()
}
}
pub struct OldKeyRetention {
initiator_key: Option<SessionKey>,
responder_key: Option<SessionKey>,
retained_at: Option<Instant>,
}
impl OldKeyRetention {
pub fn new() -> Self {
Self {
initiator_key: None,
responder_key: None,
retained_at: None,
}
}
pub fn retain(&mut self, initiator_key: SessionKey, responder_key: SessionKey) {
self.initiator_key = Some(initiator_key);
self.responder_key = Some(responder_key);
self.retained_at = Some(Instant::now());
}
pub fn old_initiator_key(&self) -> Option<&SessionKey> {
if self.within_retention_window() {
self.initiator_key.as_ref()
} else {
None
}
}
pub fn old_responder_key(&self) -> Option<&SessionKey> {
if self.within_retention_window() {
self.responder_key.as_ref()
} else {
None
}
}
pub fn within_retention_window(&self) -> bool {
self.retained_at
.is_some_and(|t| t.elapsed() < OLD_KEY_RETENTION)
}
pub fn clear(&mut self) {
self.initiator_key = None;
self.responder_key = None;
self.retained_at = None;
}
pub fn should_clear(&self) -> bool {
self.retained_at
.is_some_and(|t| t.elapsed() >= OLD_KEY_RETENTION)
}
pub fn clear_if_expired(&mut self) {
if self.should_clear() {
self.clear();
}
}
}
impl Default for OldKeyRetention {
fn default() -> Self {
Self::new()
}
}
pub fn derive_rekey_keys(
ephemeral_dh: &[u8; 32],
rekey_auth_key: &[u8; 32],
epoch: u32,
) -> Result<(SessionKey, SessionKey), CryptoError> {
let mut ikm = [0u8; 64];
ikm[..32].copy_from_slice(ephemeral_dh);
ikm[32..].copy_from_slice(rekey_auth_key);
let label = b"nomad v1 rekey";
let epoch_bytes = epoch.to_le_bytes();
let mut info = Vec::with_capacity(label.len() + 4);
info.extend_from_slice(label);
info.extend_from_slice(&epoch_bytes);
let hk = Hkdf::<Sha256>::from_prk(&ikm)
.map_err(|_| CryptoError::KeyDerivationFailed)?;
let mut key_material = [0u8; 64];
hk.expand(&info, &mut key_material)
.map_err(|_| CryptoError::KeyDerivationFailed)?;
let mut initiator_key = [0u8; SESSION_KEY_SIZE];
let mut responder_key = [0u8; SESSION_KEY_SIZE];
initiator_key.copy_from_slice(&key_material[..32]);
responder_key.copy_from_slice(&key_material[32..]);
ikm.zeroize();
key_material.zeroize();
Ok((
SessionKey::from_bytes(initiator_key),
SessionKey::from_bytes(responder_key),
))
}
pub fn derive_rekey_auth_key(static_dh_secret: &[u8; 32]) -> [u8; 32] {
let info = b"nomad v1 rekey auth";
let hk = Hkdf::<Sha256>::from_prk(static_dh_secret)
.expect("32 bytes is valid PRK length for SHA-256 HKDF");
let mut rekey_auth_key = [0u8; 32];
hk.expand(info, &mut rekey_auth_key)
.expect("32 bytes is valid output length for SHA-256 HKDF");
rekey_auth_key
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_rekey_state_new() {
let state = RekeyState::new();
assert_eq!(state.epoch(), 0);
assert_eq!(state.send_count(), 0);
assert_eq!(state.recv_count(), 0);
assert!(!state.should_rekey());
assert!(!state.keys_expired());
assert!(state.can_rekey());
}
#[test]
fn test_increment_send() {
let mut state = RekeyState::new();
for i in 0..10 {
let counter = state.increment_send().unwrap();
assert_eq!(counter, i);
}
assert_eq!(state.send_count(), 10);
}
#[test]
fn test_record_recv() {
let mut state = RekeyState::new();
state.record_recv(5);
assert_eq!(state.recv_count(), 6);
state.record_recv(3); assert_eq!(state.recv_count(), 6);
state.record_recv(10);
assert_eq!(state.recv_count(), 11);
}
#[test]
fn test_advance_epoch() {
let mut state = RekeyState::new();
state.increment_send().unwrap();
state.increment_send().unwrap();
state.advance_epoch().unwrap();
assert_eq!(state.epoch(), 1);
assert_eq!(state.send_count(), 0);
assert_eq!(state.recv_count(), 0);
}
#[test]
fn test_old_key_retention() {
let mut retention = OldKeyRetention::new();
assert!(retention.old_initiator_key().is_none());
assert!(!retention.within_retention_window());
let key1 = SessionKey::from_bytes([0x01; SESSION_KEY_SIZE]);
let key2 = SessionKey::from_bytes([0x02; SESSION_KEY_SIZE]);
retention.retain(key1, key2);
assert!(retention.within_retention_window());
assert!(retention.old_initiator_key().is_some());
assert!(retention.old_responder_key().is_some());
}
#[test]
fn test_derive_rekey_keys() {
let ephemeral_dh = [0x42u8; 32];
let rekey_auth_key = [0x33u8; 32];
let (key1_epoch0, key2_epoch0) = derive_rekey_keys(&ephemeral_dh, &rekey_auth_key, 0).unwrap();
let (key1_epoch1, key2_epoch1) = derive_rekey_keys(&ephemeral_dh, &rekey_auth_key, 1).unwrap();
assert_ne!(key1_epoch0.as_bytes(), key1_epoch1.as_bytes());
assert_ne!(key2_epoch0.as_bytes(), key2_epoch1.as_bytes());
let (key1_epoch0_again, key2_epoch0_again) = derive_rekey_keys(&ephemeral_dh, &rekey_auth_key, 0).unwrap();
assert_eq!(key1_epoch0.as_bytes(), key1_epoch0_again.as_bytes());
assert_eq!(key2_epoch0.as_bytes(), key2_epoch0_again.as_bytes());
}
#[test]
fn test_derive_rekey_keys_different_ephemeral_dh() {
let ephemeral_dh1 = [0x01u8; 32];
let ephemeral_dh2 = [0x02u8; 32];
let rekey_auth_key = [0x33u8; 32];
let (key1_dh1, _) = derive_rekey_keys(&ephemeral_dh1, &rekey_auth_key, 0).unwrap();
let (key1_dh2, _) = derive_rekey_keys(&ephemeral_dh2, &rekey_auth_key, 0).unwrap();
assert_ne!(key1_dh1.as_bytes(), key1_dh2.as_bytes());
}
#[test]
fn test_derive_rekey_keys_pcs() {
let ephemeral_dh = [0x42u8; 32];
let auth_key1 = [0x01u8; 32];
let auth_key2 = [0x02u8; 32];
let (key1_auth1, _) = derive_rekey_keys(&ephemeral_dh, &auth_key1, 0).unwrap();
let (key1_auth2, _) = derive_rekey_keys(&ephemeral_dh, &auth_key2, 0).unwrap();
assert_ne!(key1_auth1.as_bytes(), key1_auth2.as_bytes());
}
#[test]
fn test_derive_rekey_auth_key() {
let static_dh1 = [0x01u8; 32];
let static_dh2 = [0x02u8; 32];
let auth_key1 = derive_rekey_auth_key(&static_dh1);
let auth_key2 = derive_rekey_auth_key(&static_dh2);
assert_ne!(auth_key1, auth_key2);
let auth_key1_again = derive_rekey_auth_key(&static_dh1);
assert_eq!(auth_key1, auth_key1_again);
}
#[test]
fn test_pcs_property() {
let ephemeral_dh = [0x42u8; 32];
let real_static_dh = [0xABu8; 32];
let attacker_guess_dh = [0xCDu8; 32];
let real_auth_key = derive_rekey_auth_key(&real_static_dh);
let attacker_auth_key = derive_rekey_auth_key(&attacker_guess_dh);
let (real_key1, _) = derive_rekey_keys(&ephemeral_dh, &real_auth_key, 1).unwrap();
let (attacker_key1, _) = derive_rekey_keys(&ephemeral_dh, &attacker_auth_key, 1).unwrap();
assert_ne!(real_key1.as_bytes(), attacker_key1.as_bytes());
}
fn hex_to_bytes(hex: &str) -> Vec<u8> {
(0..hex.len())
.step_by(2)
.map(|i| u8::from_str_radix(&hex[i..i + 2], 16).unwrap())
.collect()
}
#[test]
fn test_vector_rekey_auth_key() {
let static_dh = hex_to_bytes("57fbeea357c6ca4af3654988d78e020ccc6f4bc56db385bff4a46084b1187266");
let expected_auth_key = hex_to_bytes("48c391a58d3e6fe3e5c463cd874b4565b752da33d63b9d93f9a469549ebbbe09");
let mut static_dh_arr = [0u8; 32];
static_dh_arr.copy_from_slice(&static_dh);
let auth_key = derive_rekey_auth_key(&static_dh_arr);
assert_eq!(
auth_key.as_slice(),
expected_auth_key.as_slice(),
"rekey_auth_key derivation doesn't match test vector"
);
}
#[test]
fn test_vector_epoch_1() {
let ephemeral_dh = hex_to_bytes("813c560b94aec760c9a8d12a09bb4c2be3bfc35eb6983ceb264a13046d3aaa75");
let rekey_auth_key = hex_to_bytes("48c391a58d3e6fe3e5c463cd874b4565b752da33d63b9d93f9a469549ebbbe09");
let expected_initiator_key = hex_to_bytes("ba7ba9959a0338866994033dc46c15df92e6a08b4d5041d5e52070001187c312");
let expected_responder_key = hex_to_bytes("91f2e4123a04abe6343003d6ff5793af7aae75ede7fdc6737aaf24964d9285f8");
let mut ephemeral_dh_arr = [0u8; 32];
let mut rekey_auth_key_arr = [0u8; 32];
ephemeral_dh_arr.copy_from_slice(&ephemeral_dh);
rekey_auth_key_arr.copy_from_slice(&rekey_auth_key);
let (initiator_key, responder_key) = derive_rekey_keys(&ephemeral_dh_arr, &rekey_auth_key_arr, 1).unwrap();
assert_eq!(
initiator_key.as_bytes(),
expected_initiator_key.as_slice(),
"epoch 1 initiator key doesn't match test vector"
);
assert_eq!(
responder_key.as_bytes(),
expected_responder_key.as_slice(),
"epoch 1 responder key doesn't match test vector"
);
}
#[test]
fn test_vector_epoch_2() {
let ephemeral_dh = hex_to_bytes("7efd5673c47236ad6f9bf85e945074615c1943c528a87cc0dc9084ad278d266e");
let rekey_auth_key = hex_to_bytes("48c391a58d3e6fe3e5c463cd874b4565b752da33d63b9d93f9a469549ebbbe09");
let expected_initiator_key = hex_to_bytes("206c3c4f0838aaf5b039bad2ecd1a387d6f784afbf1d283dc0a438ad45f4db3e");
let expected_responder_key = hex_to_bytes("786554075c38e73a735b26cbfd650c9fd0f8909227e498487007fc2adfec661d");
let mut ephemeral_dh_arr = [0u8; 32];
let mut rekey_auth_key_arr = [0u8; 32];
ephemeral_dh_arr.copy_from_slice(&ephemeral_dh);
rekey_auth_key_arr.copy_from_slice(&rekey_auth_key);
let (initiator_key, responder_key) = derive_rekey_keys(&ephemeral_dh_arr, &rekey_auth_key_arr, 2).unwrap();
assert_eq!(
initiator_key.as_bytes(),
expected_initiator_key.as_slice(),
"epoch 2 initiator key doesn't match test vector"
);
assert_eq!(
responder_key.as_bytes(),
expected_responder_key.as_slice(),
"epoch 2 responder key doesn't match test vector"
);
}
#[test]
fn test_vector_epoch_100() {
let ephemeral_dh = hex_to_bytes("0038038a95c66833de6cd4a4743226d03d952d35d1885876f63b95deea271e3f");
let rekey_auth_key = hex_to_bytes("48c391a58d3e6fe3e5c463cd874b4565b752da33d63b9d93f9a469549ebbbe09");
let expected_initiator_key = hex_to_bytes("dda7dd785c4c5f75096c0ea88023b1558e26bb84f4c4eb72ba7977c6947abc1a");
let expected_responder_key = hex_to_bytes("110c7c42998204153892f1ac84634c355ed1b279174befd2f27936073567e54f");
let mut ephemeral_dh_arr = [0u8; 32];
let mut rekey_auth_key_arr = [0u8; 32];
ephemeral_dh_arr.copy_from_slice(&ephemeral_dh);
rekey_auth_key_arr.copy_from_slice(&rekey_auth_key);
let (initiator_key, responder_key) = derive_rekey_keys(&ephemeral_dh_arr, &rekey_auth_key_arr, 100).unwrap();
assert_eq!(
initiator_key.as_bytes(),
expected_initiator_key.as_slice(),
"epoch 100 initiator key doesn't match test vector"
);
assert_eq!(
responder_key.as_bytes(),
expected_responder_key.as_slice(),
"epoch 100 responder key doesn't match test vector"
);
}
}