use crate::core::{CryptoError, HASH_SIZE, REPLAY_WINDOW_SIZE};
use super::{
aead::{construct_aad, decrypt, encrypt, SessionKey},
nonce::{construct_nonce, Direction},
rekey::{OldKeyRetention, RekeyState},
Role, SessionId,
};
pub struct ReplayWindow {
bitmap: [u64; REPLAY_WINDOW_SIZE / 64],
highest: u64,
initialized: bool,
}
impl ReplayWindow {
pub fn new() -> Self {
Self {
bitmap: [0; REPLAY_WINDOW_SIZE / 64],
highest: 0,
initialized: false,
}
}
pub fn is_replay(&self, nonce: u64) -> bool {
if !self.initialized {
return false;
}
if nonce > self.highest {
return false;
}
let diff = self.highest - nonce;
if diff >= REPLAY_WINDOW_SIZE as u64 {
return true; }
let bit_index = diff as usize;
let word_index = bit_index / 64;
let bit_offset = bit_index % 64;
(self.bitmap[word_index] & (1 << bit_offset)) != 0
}
pub fn check_and_update(&mut self, nonce: u64) -> Result<(), CryptoError> {
if !self.initialized {
self.highest = nonce;
self.mark_seen(nonce);
self.initialized = true;
return Ok(());
}
if nonce > self.highest {
let shift = nonce - self.highest;
self.shift_window(shift);
self.highest = nonce;
self.mark_seen(nonce);
Ok(())
} else {
let diff = self.highest - nonce;
if diff >= REPLAY_WINDOW_SIZE as u64 {
return Err(CryptoError::ReplayDetected);
}
if self.is_seen(nonce) {
return Err(CryptoError::ReplayDetected);
}
self.mark_seen(nonce);
Ok(())
}
}
fn is_seen(&self, nonce: u64) -> bool {
if nonce > self.highest {
return false;
}
let diff = self.highest - nonce;
if diff >= REPLAY_WINDOW_SIZE as u64 {
return true; }
let bit_index = diff as usize;
let word_index = bit_index / 64;
let bit_offset = bit_index % 64;
(self.bitmap[word_index] & (1 << bit_offset)) != 0
}
fn mark_seen(&mut self, nonce: u64) {
if nonce > self.highest {
return; }
let diff = self.highest - nonce;
if diff >= REPLAY_WINDOW_SIZE as u64 {
return; }
let bit_index = diff as usize;
let word_index = bit_index / 64;
let bit_offset = bit_index % 64;
self.bitmap[word_index] |= 1 << bit_offset;
}
fn shift_window(&mut self, shift: u64) {
if shift >= REPLAY_WINDOW_SIZE as u64 {
self.bitmap = [0; REPLAY_WINDOW_SIZE / 64];
return;
}
let shift_words = (shift / 64) as usize;
let shift_bits = (shift % 64) as u32;
if shift_words > 0 {
for i in (shift_words..self.bitmap.len()).rev() {
self.bitmap[i] = self.bitmap[i - shift_words];
}
for word in self.bitmap.iter_mut().take(shift_words) {
*word = 0;
}
}
if shift_bits > 0 {
let mut carry = 0u64;
for i in (0..self.bitmap.len()).rev() {
let new_carry = self.bitmap[i] >> (64 - shift_bits);
self.bitmap[i] = (self.bitmap[i] << shift_bits) | carry;
carry = new_carry;
}
}
}
pub fn reset(&mut self) {
self.bitmap = [0; REPLAY_WINDOW_SIZE / 64];
self.highest = 0;
self.initialized = false;
}
}
impl Default for ReplayWindow {
fn default() -> Self {
Self::new()
}
}
pub struct CryptoSession {
session_id: SessionId,
role: Role,
send_key: SessionKey,
recv_key: SessionKey,
rekey_state: RekeyState,
replay_window: ReplayWindow,
old_keys: OldKeyRetention,
#[allow(dead_code)]
handshake_hash: [u8; HASH_SIZE],
rekey_auth_key: [u8; HASH_SIZE],
}
impl CryptoSession {
pub fn new(
session_id: SessionId,
role: Role,
send_key: SessionKey,
recv_key: SessionKey,
handshake_hash: [u8; HASH_SIZE],
rekey_auth_key: [u8; HASH_SIZE],
) -> Self {
Self {
session_id,
role,
send_key,
recv_key,
rekey_state: RekeyState::new(),
replay_window: ReplayWindow::new(),
old_keys: OldKeyRetention::new(),
handshake_hash,
rekey_auth_key,
}
}
pub fn session_id(&self) -> &SessionId {
&self.session_id
}
pub fn role(&self) -> Role {
self.role
}
pub fn epoch(&self) -> u32 {
self.rekey_state.epoch()
}
pub fn should_rekey(&self) -> bool {
self.rekey_state.should_rekey()
}
pub fn keys_expired(&self) -> bool {
self.rekey_state.keys_expired()
}
fn send_direction(&self) -> Direction {
match self.role {
Role::Initiator => Direction::InitiatorToResponder,
Role::Responder => Direction::ResponderToInitiator,
}
}
fn recv_direction(&self) -> Direction {
self.send_direction().opposite()
}
pub fn encrypt_frame(
&mut self,
frame_type: u8,
flags: u8,
plaintext: &[u8],
) -> Result<(u64, Vec<u8>), CryptoError> {
let counter = self.rekey_state.increment_send()?;
let nonce = construct_nonce(self.rekey_state.epoch(), self.send_direction(), counter);
let aad = construct_aad(frame_type, flags, self.session_id.as_bytes(), counter);
let ciphertext = encrypt(&self.send_key, &nonce, &aad, plaintext)?;
Ok((counter, ciphertext))
}
pub fn decrypt_frame(
&mut self,
frame_type: u8,
flags: u8,
nonce_counter: u64,
ciphertext: &[u8],
) -> Result<Vec<u8>, CryptoError> {
if self.replay_window.is_replay(nonce_counter) {
return Err(CryptoError::ReplayDetected);
}
let nonce = construct_nonce(self.rekey_state.epoch(), self.recv_direction(), nonce_counter);
let aad = construct_aad(frame_type, flags, self.session_id.as_bytes(), nonce_counter);
if let Ok(plaintext) = decrypt(&self.recv_key, &nonce, &aad, ciphertext) {
let _ = self.replay_window.check_and_update(nonce_counter);
self.rekey_state.record_recv(nonce_counter);
return Ok(plaintext);
}
self.old_keys.clear_if_expired();
if let Some(old_recv_key) = self.get_old_recv_key() {
let old_epoch = self.rekey_state.epoch().saturating_sub(1);
let old_nonce = construct_nonce(old_epoch, self.recv_direction(), nonce_counter);
if let Ok(plaintext) = decrypt(old_recv_key, &old_nonce, &aad, ciphertext) {
return Ok(plaintext);
}
}
Err(CryptoError::DecryptionFailed)
}
fn get_old_recv_key(&self) -> Option<&SessionKey> {
match self.role {
Role::Initiator => self.old_keys.old_responder_key(),
Role::Responder => self.old_keys.old_initiator_key(),
}
}
pub fn rekey(&mut self, ephemeral_dh: &[u8; 32]) -> Result<(), CryptoError> {
use super::rekey::derive_rekey_keys;
self.old_keys
.retain(self.send_key.clone(), self.recv_key.clone());
self.rekey_state.advance_epoch()?;
let (new_initiator_key, new_responder_key) =
derive_rekey_keys(ephemeral_dh, &self.rekey_auth_key, self.rekey_state.epoch())?;
match self.role {
Role::Initiator => {
self.send_key = new_initiator_key;
self.recv_key = new_responder_key;
}
Role::Responder => {
self.send_key = new_responder_key;
self.recv_key = new_initiator_key;
}
}
self.replay_window.reset();
Ok(())
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_replay_window_basic() {
let mut window = ReplayWindow::new();
assert!(window.check_and_update(0).is_ok());
assert!(window.check_and_update(0).is_err());
assert!(window.check_and_update(1).is_ok());
assert!(window.check_and_update(5).is_ok());
assert!(window.check_and_update(3).is_ok());
assert!(window.check_and_update(4).is_ok());
assert!(window.check_and_update(2).is_ok());
assert!(window.check_and_update(0).is_err());
assert!(window.check_and_update(3).is_err());
assert!(window.check_and_update(5).is_err());
}
#[test]
fn test_replay_window_large_gap() {
let mut window = ReplayWindow::new();
assert!(window.check_and_update(0).is_ok());
assert!(window.check_and_update(1).is_ok());
assert!(window.check_and_update(1000).is_ok());
assert!(window.check_and_update(0).is_err());
assert!(window.check_and_update(1).is_err());
assert!(window.check_and_update(999).is_ok());
assert!(window.check_and_update(998).is_ok());
}
#[test]
fn test_replay_window_full_reset() {
let mut window = ReplayWindow::new();
for i in 0..100 {
assert!(window.check_and_update(i).is_ok());
}
assert!(window.check_and_update(100 + REPLAY_WINDOW_SIZE as u64).is_ok());
for i in 0..100 {
assert!(window.check_and_update(i).is_err());
}
}
#[test]
fn test_crypto_session_roundtrip() {
let session_id = SessionId::generate();
let send_key = SessionKey::from_bytes([0x01; 32]);
let recv_key = SessionKey::from_bytes([0x02; 32]);
let handshake_hash = [0x42; 32];
let rekey_auth_key = [0x33; 32];
let mut initiator = CryptoSession::new(
session_id,
Role::Initiator,
send_key.clone(),
recv_key.clone(),
handshake_hash,
rekey_auth_key,
);
let mut responder = CryptoSession::new(
session_id,
Role::Responder,
recv_key.clone(),
send_key.clone(),
handshake_hash,
rekey_auth_key,
);
let plaintext = b"Hello, NOMAD!";
let (counter, ciphertext) = initiator.encrypt_frame(0x03, 0x00, plaintext).unwrap();
let decrypted = responder
.decrypt_frame(0x03, 0x00, counter, &ciphertext)
.unwrap();
assert_eq!(decrypted, plaintext);
let reply = b"Hello back!";
let (reply_counter, reply_ciphertext) =
responder.encrypt_frame(0x03, 0x00, reply).unwrap();
let decrypted_reply = initiator
.decrypt_frame(0x03, 0x00, reply_counter, &reply_ciphertext)
.unwrap();
assert_eq!(decrypted_reply, reply);
}
#[test]
fn test_crypto_session_replay_detection() {
let session_id = SessionId::generate();
let send_key = SessionKey::from_bytes([0x01; 32]);
let recv_key = SessionKey::from_bytes([0x02; 32]);
let handshake_hash = [0x42; 32];
let rekey_auth_key = [0x33; 32];
let mut initiator = CryptoSession::new(
session_id,
Role::Initiator,
send_key.clone(),
recv_key.clone(),
handshake_hash,
rekey_auth_key,
);
let mut responder = CryptoSession::new(
session_id,
Role::Responder,
recv_key.clone(),
send_key.clone(),
handshake_hash,
rekey_auth_key,
);
let plaintext = b"test";
let (counter, ciphertext) = initiator.encrypt_frame(0x03, 0x00, plaintext).unwrap();
assert!(responder
.decrypt_frame(0x03, 0x00, counter, &ciphertext)
.is_ok());
assert!(responder
.decrypt_frame(0x03, 0x00, counter, &ciphertext)
.is_err());
}
#[test]
fn test_crypto_session_wrong_aad() {
let session_id = SessionId::generate();
let send_key = SessionKey::from_bytes([0x01; 32]);
let recv_key = SessionKey::from_bytes([0x02; 32]);
let handshake_hash = [0x42; 32];
let rekey_auth_key = [0x33; 32];
let mut initiator = CryptoSession::new(
session_id,
Role::Initiator,
send_key.clone(),
recv_key.clone(),
handshake_hash,
rekey_auth_key,
);
let mut responder = CryptoSession::new(
session_id,
Role::Responder,
recv_key.clone(),
send_key.clone(),
handshake_hash,
rekey_auth_key,
);
let plaintext = b"test";
let (counter, ciphertext) = initiator.encrypt_frame(0x03, 0x00, plaintext).unwrap();
assert!(responder
.decrypt_frame(0x04, 0x00, counter, &ciphertext)
.is_err());
}
}