use std::collections::HashSet;
use std::collections::VecDeque;
use std::time::{SystemTime, UNIX_EPOCH};
use aes_siv::Aes128SivAead;
use aes_siv::aead::generic_array::GenericArray;
use aes_siv::aead::{Aead, KeyInit, Payload};
use rand::Rng;
use sha2::{Digest, Sha256};
use zeroize::{Zeroize, ZeroizeOnDrop};
use crate::{AEAD_AES_SIV_CMAC_256_KEYLEN, NtsError};
const KEY_ID_SIZE: usize = 4;
const COOKIE_NONCE_SIZE: usize = 16;
const DEFAULT_COOKIE_TTL_SECS: u64 = 86400;
const MAX_USED_NONCES: usize = 100_000;
const NONCE_EVICT_COUNT: usize = 10_000;
fn derive_key_id(key: &[u8; AEAD_AES_SIV_CMAC_256_KEYLEN]) -> u32 {
let hash = Sha256::digest(key);
u32::from_be_bytes([hash[0], hash[1], hash[2], hash[3]])
}
#[derive(Clone, PartialEq, Eq, Zeroize, ZeroizeOnDrop)]
pub struct CookieContents {
pub algorithm: u16,
pub c2s_key: Vec<u8>,
pub s2c_key: Vec<u8>,
}
impl std::fmt::Debug for CookieContents {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("CookieContents")
.field("algorithm", &self.algorithm)
.field("c2s_key", &"[REDACTED]")
.field("s2c_key", &"[REDACTED]")
.finish()
}
}
pub struct CookieJar {
current_key: [u8; AEAD_AES_SIV_CMAC_256_KEYLEN],
current_key_id: u32,
previous_key: Option<[u8; AEAD_AES_SIV_CMAC_256_KEYLEN]>,
previous_key_id: Option<u32>,
used_nonces: HashSet<[u8; COOKIE_NONCE_SIZE]>,
nonce_order: VecDeque<[u8; COOKIE_NONCE_SIZE]>,
cookie_ttl_secs: u64,
}
impl Drop for CookieJar {
fn drop(&mut self) {
self.current_key.zeroize();
if let Some(ref mut key) = self.previous_key {
key.zeroize();
}
}
}
impl CookieJar {
pub fn new(master_key: [u8; AEAD_AES_SIV_CMAC_256_KEYLEN]) -> Self {
let key_id = derive_key_id(&master_key);
Self {
current_key: master_key,
current_key_id: key_id,
previous_key: None,
previous_key_id: None,
used_nonces: HashSet::new(),
nonce_order: VecDeque::new(),
cookie_ttl_secs: DEFAULT_COOKIE_TTL_SECS,
}
}
pub fn with_ttl(master_key: [u8; AEAD_AES_SIV_CMAC_256_KEYLEN], ttl_secs: u64) -> Self {
let mut jar = Self::new(master_key);
jar.cookie_ttl_secs = ttl_secs;
jar
}
pub fn make_cookie(&self, c2s_key: &[u8], s2c_key: &[u8], algorithm: u16) -> Vec<u8> {
let now = SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap_or_default()
.as_secs();
let mut plaintext = Vec::with_capacity(2 + 8 + c2s_key.len() + s2c_key.len());
plaintext.extend_from_slice(&algorithm.to_be_bytes());
plaintext.extend_from_slice(&now.to_be_bytes());
plaintext.extend_from_slice(c2s_key);
plaintext.extend_from_slice(s2c_key);
let mut nonce = [0u8; COOKIE_NONCE_SIZE];
rand::rng().fill_bytes(&mut nonce);
let cipher = Aes128SivAead::new((&self.current_key).into());
let zero_nonce = GenericArray::default();
let payload = Payload {
msg: &plaintext,
aad: &nonce,
};
let ciphertext = cipher
.encrypt(&zero_nonce, payload)
.expect("cookie encryption should not fail with valid key");
plaintext.zeroize();
let mut cookie = Vec::with_capacity(KEY_ID_SIZE + COOKIE_NONCE_SIZE + ciphertext.len());
cookie.extend_from_slice(&self.current_key_id.to_be_bytes());
cookie.extend_from_slice(&nonce);
cookie.extend_from_slice(&ciphertext);
cookie
}
pub fn open_cookie(&mut self, cookie: &[u8]) -> Result<CookieContents, NtsError> {
let min_size = KEY_ID_SIZE + COOKIE_NONCE_SIZE;
if cookie.len() < min_size {
return Err(NtsError::InvalidCookie(format!(
"cookie too short: {} bytes, need at least {}",
cookie.len(),
min_size
)));
}
let key_id = u32::from_be_bytes([cookie[0], cookie[1], cookie[2], cookie[3]]);
let nonce = &cookie[KEY_ID_SIZE..KEY_ID_SIZE + COOKIE_NONCE_SIZE];
let ciphertext = &cookie[KEY_ID_SIZE + COOKIE_NONCE_SIZE..];
let nonce_array: [u8; COOKIE_NONCE_SIZE] = nonce
.try_into()
.map_err(|_| NtsError::InvalidCookie("nonce length mismatch".to_string()))?;
if self.used_nonces.contains(&nonce_array) {
return Err(NtsError::InvalidCookie(
"cookie replay detected".to_string(),
));
}
let plaintext = if key_id == self.current_key_id {
decrypt_cookie_data(&self.current_key, nonce, ciphertext)?
} else if self.previous_key_id == Some(key_id) {
if let Some(ref prev_key) = self.previous_key {
decrypt_cookie_data(prev_key, nonce, ciphertext)?
} else {
return Err(NtsError::InvalidCookie(
"key_id matches previous but no previous key available".to_string(),
));
}
} else {
return Err(NtsError::InvalidCookie(format!(
"unknown key_id: {}",
key_id
)));
};
if plaintext.len() < 10 {
return Err(NtsError::InvalidCookie(
"decrypted cookie too short for header".to_string(),
));
}
let algorithm = u16::from_be_bytes([plaintext[0], plaintext[1]]);
let created_at = u64::from_be_bytes([
plaintext[2],
plaintext[3],
plaintext[4],
plaintext[5],
plaintext[6],
plaintext[7],
plaintext[8],
plaintext[9],
]);
let now = SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap_or_default()
.as_secs();
if now.saturating_sub(created_at) >= self.cookie_ttl_secs {
return Err(NtsError::InvalidCookie("cookie has expired".to_string()));
}
let key_data = &plaintext[10..];
if key_data.len() % 2 != 0 {
return Err(NtsError::InvalidCookie(
"key data has odd length".to_string(),
));
}
let key_len = key_data.len() / 2;
let c2s_key = key_data[..key_len].to_vec();
let s2c_key = key_data[key_len..].to_vec();
if self.used_nonces.len() >= MAX_USED_NONCES {
let to_evict = NONCE_EVICT_COUNT.min(self.nonce_order.len());
for _ in 0..to_evict {
if let Some(old_nonce) = self.nonce_order.pop_front() {
self.used_nonces.remove(&old_nonce);
}
}
}
self.used_nonces.insert(nonce_array);
self.nonce_order.push_back(nonce_array);
Ok(CookieContents {
algorithm,
c2s_key,
s2c_key,
})
}
pub fn rotate_key(&mut self, new_key: [u8; AEAD_AES_SIV_CMAC_256_KEYLEN]) {
self.previous_key = Some(self.current_key);
self.previous_key_id = Some(self.current_key_id);
self.current_key = new_key;
self.current_key_id = derive_key_id(&new_key);
}
pub fn current_key(&self) -> &[u8; AEAD_AES_SIV_CMAC_256_KEYLEN] {
&self.current_key
}
}
fn decrypt_cookie_data(
key: &[u8; AEAD_AES_SIV_CMAC_256_KEYLEN],
nonce: &[u8],
ciphertext: &[u8],
) -> Result<Vec<u8>, NtsError> {
let cipher = Aes128SivAead::new(key.into());
let zero_nonce = GenericArray::default();
let payload = Payload {
msg: ciphertext,
aad: nonce,
};
cipher
.decrypt(&zero_nonce, payload)
.map_err(|_| NtsError::DecryptionFailed)
}
#[cfg(test)]
mod tests {
use super::*;
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 make_open_roundtrip() {
let master_key = random_key();
let mut jar = CookieJar::new(master_key);
let c2s_key = random_key();
let s2c_key = random_key();
let algorithm = 15u16;
let cookie = jar.make_cookie(&c2s_key, &s2c_key, algorithm);
let contents = jar.open_cookie(&cookie).unwrap();
assert_eq!(contents.algorithm, algorithm);
assert_eq!(contents.c2s_key, c2s_key);
assert_eq!(contents.s2c_key, s2c_key);
}
#[test]
fn different_cookies_are_unique() {
let master_key = random_key();
let mut jar = CookieJar::new(master_key);
let c2s_key = random_key();
let s2c_key = random_key();
let cookie1 = jar.make_cookie(&c2s_key, &s2c_key, 15);
let cookie2 = jar.make_cookie(&c2s_key, &s2c_key, 15);
assert_ne!(cookie1, cookie2);
let c1 = jar.open_cookie(&cookie1).unwrap();
let c2 = jar.open_cookie(&cookie2).unwrap();
assert_eq!(c1, c2);
}
#[test]
fn wrong_master_key_fails() {
let jar1 = CookieJar::new(random_key());
let mut jar2 = CookieJar::new(random_key());
let cookie = jar1.make_cookie(&random_key(), &random_key(), 15);
assert!(jar2.open_cookie(&cookie).is_err());
}
#[test]
fn key_rotation_current_key_works() {
let key1 = random_key();
let key2 = random_key();
let mut jar = CookieJar::new(key1);
let c2s = random_key();
let s2c = random_key();
let cookie_old = jar.make_cookie(&c2s, &s2c, 15);
jar.rotate_key(key2);
let contents = jar.open_cookie(&cookie_old).unwrap();
assert_eq!(contents.c2s_key, c2s);
assert_eq!(contents.s2c_key, s2c);
}
#[test]
fn key_rotation_new_key_works() {
let key1 = random_key();
let key2 = random_key();
let mut jar = CookieJar::new(key1);
jar.rotate_key(key2);
let c2s = random_key();
let s2c = random_key();
let cookie_new = jar.make_cookie(&c2s, &s2c, 15);
let contents = jar.open_cookie(&cookie_new).unwrap();
assert_eq!(contents.c2s_key, c2s);
assert_eq!(contents.s2c_key, s2c);
}
#[test]
fn double_rotation_drops_oldest_key() {
let key1 = random_key();
let key2 = random_key();
let key3 = random_key();
let mut jar = CookieJar::new(key1);
let cookie_k1 = jar.make_cookie(&random_key(), &random_key(), 15);
jar.rotate_key(key2);
assert!(jar.open_cookie(&cookie_k1).is_ok());
jar.rotate_key(key3);
assert!(jar.open_cookie(&cookie_k1).is_err());
}
#[test]
fn tampered_cookie_fails() {
let mut jar = CookieJar::new(random_key());
let mut cookie = jar.make_cookie(&random_key(), &random_key(), 15);
let last = cookie.len() - 1;
cookie[last] ^= 0xFF;
assert!(jar.open_cookie(&cookie).is_err());
}
#[test]
fn too_short_cookie_fails() {
let mut jar = CookieJar::new(random_key());
assert!(jar.open_cookie(&[0u8; 10]).is_err());
}
#[test]
fn cookie_preserves_algorithm() {
let mut jar = CookieJar::new(random_key());
let c2s = random_key();
let s2c = random_key();
for algo in [15u16, 16, 0, 0xFFFF] {
let cookie = jar.make_cookie(&c2s, &s2c, algo);
let contents = jar.open_cookie(&cookie).unwrap();
assert_eq!(contents.algorithm, algo);
}
}
#[test]
fn replay_detected() {
let mut jar = CookieJar::new(random_key());
let cookie = jar.make_cookie(&random_key(), &random_key(), 15);
assert!(jar.open_cookie(&cookie).is_ok());
let err = jar.open_cookie(&cookie).unwrap_err();
assert!(err.to_string().contains("replay"));
}
#[test]
fn expired_cookie_rejected() {
let mut jar = CookieJar::with_ttl(random_key(), 1);
let cookie = jar.make_cookie(&random_key(), &random_key(), 15);
std::thread::sleep(std::time::Duration::from_secs(1));
let err = jar.open_cookie(&cookie).unwrap_err();
assert!(err.to_string().contains("expired"));
}
#[test]
fn lru_eviction_preserves_recent_nonces() {
let mut jar = CookieJar::new(random_key());
let mut cookies = Vec::new();
for _ in 0..MAX_USED_NONCES {
let cookie = jar.make_cookie(&random_key(), &random_key(), 15);
jar.open_cookie(&cookie).unwrap();
cookies.push(cookie);
}
let extra = jar.make_cookie(&random_key(), &random_key(), 15);
jar.open_cookie(&extra).unwrap();
let recent = &cookies[cookies.len() - 1];
assert!(jar.open_cookie(recent).is_err());
}
#[test]
fn key_id_changes_on_rotation() {
let key1 = random_key();
let key2 = random_key();
let mut jar = CookieJar::new(key1);
let id1 = jar.current_key_id;
jar.rotate_key(key2);
let id2 = jar.current_key_id;
assert_ne!(id1, id2);
}
#[test]
fn key_id_in_cookie_matches_key() {
let key = random_key();
let jar = CookieJar::new(key);
let cookie = jar.make_cookie(&random_key(), &random_key(), 15);
let embedded_id = u32::from_be_bytes([cookie[0], cookie[1], cookie[2], cookie[3]]);
assert_eq!(embedded_id, derive_key_id(&key));
}
}