use ockam_core::Result;
use ockam_vault::AeadSecretKeyHandle;
use tracing::{trace, warn};
use tracing_attributes::instrument;
use crate::{IdentityError, Nonce};
pub(crate) struct KeyTracker {
pub(crate) current_key: AeadSecretKeyHandle,
pub(crate) previous_key: Option<AeadSecretKeyHandle>,
number_of_rekeys: u64,
max_rekeys_reached: bool,
renewal_interval: u64,
}
impl KeyTracker {
pub(crate) fn new(current_key: AeadSecretKeyHandle, renewal_interval: u64) -> Self {
KeyTracker {
current_key,
number_of_rekeys: 0,
max_rekeys_reached: false,
previous_key: None,
renewal_interval,
}
}
}
impl KeyTracker {
#[instrument(skip_all)]
pub(crate) fn get_key(&self, nonce: Nonce) -> Result<Option<&AeadSecretKeyHandle>> {
trace!(
"The current number of rekeys is {}, the rekey interval is {}",
self.number_of_rekeys,
self.renewal_interval
);
let current_interval_start = self.number_of_rekeys * self.renewal_interval;
if self.max_rekeys_reached {
warn!("The maximum number of available rekeying operation has been reached. The last interval was starting at {} and the interval size is {}",
current_interval_start, self.renewal_interval);
return Err(IdentityError::InvalidNonce)?;
};
if nonce.value() >= current_interval_start {
let nonce_age = nonce.value() - current_interval_start;
if nonce_age < self.renewal_interval {
Ok(Some(&self.current_key))
}
else if nonce_age < self.renewal_interval * 2 {
Ok(None)
}
else {
warn!("This nonce is too far in the future: {}", nonce);
Err(IdentityError::InvalidNonce)?
}
} else if current_interval_start - nonce.value() <= self.renewal_interval {
if let Some(previous) = &self.previous_key {
Ok(Some(previous))
} else {
warn!("There should be a previous key for this nonce: {}", nonce);
Err(IdentityError::InvalidNonce)?
}
} else {
warn!("This nonce is too old: {}", nonce);
Err(IdentityError::InvalidNonce)?
}
}
#[instrument(skip_all)]
pub(crate) fn update_key(
&mut self,
decryption_key: &AeadSecretKeyHandle,
) -> Result<Option<AeadSecretKeyHandle>> {
let mut key_to_delete = None;
if decryption_key != &self.current_key && Some(decryption_key) != self.previous_key.as_ref()
{
key_to_delete = self.previous_key.clone();
self.previous_key.replace(self.current_key.clone());
self.current_key = decryption_key.clone();
if u64::MAX - self.number_of_rekeys * self.renewal_interval < self.renewal_interval {
self.max_rekeys_reached = true;
} else {
self.number_of_rekeys += 1;
}
}
Ok(key_to_delete)
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::MAX_NONCE;
use ockam_vault::{Aes256GcmSecretKeyHandle, HandleToSecret};
#[test]
fn test_get_key_first_interval() {
let handle = b"handle".to_vec();
let handle = AeadSecretKeyHandle(Aes256GcmSecretKeyHandle(HandleToSecret::new(handle)));
let key_tracker = KeyTracker::new(handle.clone(), 10);
assert_eq!(key_tracker.get_key(0.into()).unwrap(), Some(&handle));
assert_eq!(key_tracker.get_key(5.into()).unwrap(), Some(&handle));
assert_eq!(key_tracker.get_key(9.into()).unwrap(), Some(&handle));
assert_eq!(
key_tracker.get_key(10.into()).unwrap(),
None,
"the next key must be created"
);
assert_eq!(
key_tracker.get_key(20.into()).ok(),
None,
"this nonce is too far in the future"
);
assert_eq!(
key_tracker.get_key(MAX_NONCE).ok(),
None,
"this nonce is too far in the future"
);
}
#[test]
fn test_get_key_middle_interval() {
let handle = b"handle".to_vec();
let handle = AeadSecretKeyHandle(Aes256GcmSecretKeyHandle(HandleToSecret::new(handle)));
let previous_handle = b"previous_handle".to_vec();
let previous_handle = AeadSecretKeyHandle(Aes256GcmSecretKeyHandle(HandleToSecret::new(
previous_handle,
)));
let key_tracker = KeyTracker {
current_key: handle.clone(),
number_of_rekeys: 5,
max_rekeys_reached: false,
previous_key: Some(previous_handle.clone()),
renewal_interval: 10,
};
assert_eq!(
key_tracker.get_key(0.into()).ok(),
None,
"this nonce is too far in the past"
);
assert_eq!(
key_tracker.get_key(30.into()).ok(),
None,
"this nonce is too far in the past"
);
assert_eq!(
key_tracker.get_key(39.into()).ok(),
None,
"this nonce is too far in the past"
);
assert_eq!(
key_tracker.get_key(40.into()).unwrap(),
Some(&previous_handle)
);
assert_eq!(
key_tracker.get_key(45.into()).unwrap(),
Some(&previous_handle)
);
assert_eq!(
key_tracker.get_key(49.into()).unwrap(),
Some(&previous_handle)
);
assert_eq!(key_tracker.get_key(50.into()).unwrap(), Some(&handle));
assert_eq!(key_tracker.get_key(59.into()).unwrap(), Some(&handle));
assert_eq!(
key_tracker.get_key(60.into()).unwrap(),
None,
"the next key must be created"
);
assert_eq!(
key_tracker.get_key(MAX_NONCE).ok(),
None,
"this nonce is too far in the future"
);
}
#[test]
fn test_get_key_last_interval() {
let handle = b"handle".to_vec();
let handle = AeadSecretKeyHandle(Aes256GcmSecretKeyHandle(HandleToSecret::new(handle)));
let previous_handle = b"previous_handle".to_vec();
let previous_handle = AeadSecretKeyHandle(Aes256GcmSecretKeyHandle(HandleToSecret::new(
previous_handle,
)));
let key_tracker = KeyTracker {
current_key: handle,
number_of_rekeys: 5,
max_rekeys_reached: true,
previous_key: Some(previous_handle),
renewal_interval: 10,
};
assert_eq!(
key_tracker.get_key(0.into()).ok(),
None,
"we reached the last interval already. The channel needs to be recreated"
);
}
#[test]
fn test_update_key() {
let handle = b"handle".to_vec();
let handle = AeadSecretKeyHandle(Aes256GcmSecretKeyHandle(HandleToSecret::new(handle)));
let previous_handle = b"previous_handle".to_vec();
let previous_handle = AeadSecretKeyHandle(Aes256GcmSecretKeyHandle(HandleToSecret::new(
previous_handle,
)));
let new_handle = b"new_handle".to_vec();
let new_handle =
AeadSecretKeyHandle(Aes256GcmSecretKeyHandle(HandleToSecret::new(new_handle)));
let mut key_tracker = KeyTracker {
current_key: handle.clone(),
number_of_rekeys: 5,
max_rekeys_reached: false,
previous_key: Some(previous_handle.clone()),
renewal_interval: 10,
};
assert_eq!(key_tracker.update_key(&handle).unwrap(), None);
assert_eq!(key_tracker.update_key(&previous_handle).unwrap(), None);
assert_eq!(
key_tracker.update_key(&new_handle).unwrap(),
Some(previous_handle),
"the previous key id must be returned in order to be deleted",
);
assert_eq!(key_tracker.current_key, new_handle);
assert_eq!(key_tracker.previous_key, Some(handle));
}
#[test]
fn test_update_key_on_last_interval() {
let handle = b"handle".to_vec();
let handle = AeadSecretKeyHandle(Aes256GcmSecretKeyHandle(HandleToSecret::new(handle)));
let previous_handle = b"previous_handle".to_vec();
let previous_handle = AeadSecretKeyHandle(Aes256GcmSecretKeyHandle(HandleToSecret::new(
previous_handle,
)));
let new_handle = b"new_handle".to_vec();
let new_handle =
AeadSecretKeyHandle(Aes256GcmSecretKeyHandle(HandleToSecret::new(new_handle)));
let mut key_tracker = KeyTracker {
current_key: handle,
number_of_rekeys: u64::MAX / 10 - 1,
max_rekeys_reached: false,
previous_key: Some(previous_handle),
renewal_interval: 10,
};
key_tracker.update_key(&new_handle).unwrap();
assert!(
!key_tracker.max_rekeys_reached,
"the maximum number of rekeys is not yet reached"
);
let new_handle2 = b"new_handle2".to_vec();
let new_handle2 =
AeadSecretKeyHandle(Aes256GcmSecretKeyHandle(HandleToSecret::new(new_handle2)));
key_tracker.update_key(&new_handle2).unwrap();
assert!(
key_tracker.max_rekeys_reached,
"the maximum number of rekeys is reached now"
);
}
}