use std::collections::{HashMap, VecDeque};
#[cfg(not(target_arch = "wasm32"))]
use std::time::SystemTime;
#[cfg(target_arch = "wasm32")]
use web_time::SystemTime;
use crate::schedule::message_secrets::MessageSecrets;
use super::*;
impl EpochTree {
#[cfg(all(test, feature = "sqlite-provider", feature = "libcrux-provider"))]
pub(crate) fn timestamp(&self) -> Option<SystemTime> {
self.message_secrets.timestamp()
}
}
#[derive(Serialize, Deserialize)]
#[cfg_attr(any(test, feature = "test-utils"), derive(Clone, PartialEq))]
#[cfg_attr(feature = "crypto-debug", derive(Debug))]
pub(crate) struct EpochTree {
epoch: u64,
message_secrets: MessageSecrets,
leaves: Vec<Member>,
}
#[derive(Serialize, Deserialize)]
#[cfg_attr(any(test, feature = "test-utils"), derive(Clone, PartialEq))]
#[cfg_attr(feature = "crypto-debug", derive(Debug))]
pub(crate) struct MessageSecretsStore {
pub(crate) max_epochs: usize,
past_epoch_trees: VecDeque<EpochTree>,
message_secrets: MessageSecrets,
}
#[cfg(not(feature = "crypto-debug"))]
impl core::fmt::Debug for MessageSecretsStore {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("MessageSecretsStore")
.field("max_epochs", &"***")
.field("past_epoch_trees", &"***")
.field("message_secrets", &"***")
.finish()
}
}
const VECDEQUE_MAX_CAPACITY: usize = isize::MAX as usize;
fn max_epochs(policy: &PastEpochDeletionPolicy) -> usize {
let max_epochs = policy.max_epochs().unwrap_or(VECDEQUE_MAX_CAPACITY);
max_epochs.min(VECDEQUE_MAX_CAPACITY)
}
impl MessageSecretsStore {
pub(crate) fn new_with_secret(
policy: &PastEpochDeletionPolicy,
message_secrets: MessageSecrets,
) -> Self {
let max_epochs = max_epochs(policy);
Self {
max_epochs,
past_epoch_trees: VecDeque::new(),
message_secrets: message_secrets.with_timestamp(SystemTime::now()),
}
}
pub(crate) fn resize(&mut self, policy: &PastEpochDeletionPolicy) {
let max_past_epochs = max_epochs(policy);
let old_size = self.max_epochs;
self.max_epochs = max_past_epochs;
if old_size > max_past_epochs {
let num_epochs_out = old_size - max_past_epochs;
self.past_epoch_trees
.rotate_left(num_epochs_out.min(self.past_epoch_trees.len()));
self.past_epoch_trees.truncate(max_past_epochs);
}
}
pub(crate) fn replace_current_message_secrets(
&mut self,
message_secrets: MessageSecrets,
) -> MessageSecrets {
let mut message_secrets = message_secrets.with_timestamp(SystemTime::now());
std::mem::swap(&mut self.message_secrets, &mut message_secrets);
message_secrets
}
pub(crate) fn add_past_epoch_tree(
&mut self,
group_epoch: impl Into<GroupEpoch>,
message_secrets: MessageSecrets,
leaves: Vec<Member>,
) {
if self.max_epochs == 0 {
return;
}
if self.past_epoch_trees.len() >= self.max_epochs {
self.past_epoch_trees.rotate_left(1);
self.past_epoch_trees.truncate(self.max_epochs - 1);
}
self.past_epoch_trees.push_back(EpochTree {
epoch: group_epoch.into().as_u64(),
message_secrets,
leaves,
});
debug_assert!(
self.max_epochs >= self.past_epoch_trees.len(),
"Only {} past secrets must be stored but we found {}",
self.max_epochs,
self.past_epoch_trees.len()
);
}
pub(crate) fn secrets_for_epoch_mut(
&mut self,
group_epoch: impl Into<GroupEpoch>,
) -> Option<&mut MessageSecrets> {
let epoch = group_epoch.into().as_u64();
for epoch_tree in self.past_epoch_trees.iter_mut() {
if epoch_tree.epoch == epoch {
return Some(&mut epoch_tree.message_secrets);
}
}
None
}
pub(crate) fn secrets_for_epoch(
&self,
group_epoch: impl Into<GroupEpoch>,
) -> Option<&MessageSecrets> {
let epoch = group_epoch.into().as_u64();
for epoch_tree in self.past_epoch_trees.iter() {
if epoch_tree.epoch == epoch {
return Some(&epoch_tree.message_secrets);
}
}
None
}
pub(crate) fn secrets_and_leaves_for_epoch(
&self,
group_epoch: impl Into<GroupEpoch>,
) -> Option<(&MessageSecrets, &[Member])> {
let epoch = group_epoch.into().as_u64();
for epoch_tree in self.past_epoch_trees.iter() {
if epoch_tree.epoch == epoch {
return Some((&epoch_tree.message_secrets, &epoch_tree.leaves));
}
}
None
}
pub(crate) fn leaves_for_epoch(
&self,
group_epoch: impl Into<GroupEpoch>,
) -> HashMap<LeafNodeIndex, &Member> {
let epoch = group_epoch.into().as_u64();
for epoch_tree in self.past_epoch_trees.iter() {
if epoch_tree.epoch == epoch {
return epoch_tree
.leaves
.iter()
.map(|m| (m.index, m))
.collect::<HashMap<LeafNodeIndex, &Member>>();
}
}
HashMap::new()
}
pub(crate) fn epoch_has_leaf(
&self,
group_epoch: GroupEpoch,
leaf_index: LeafNodeIndex,
) -> bool {
self.past_epoch_trees.iter().any(|t| {
t.epoch == group_epoch.0
&& t.leaves
.iter()
.any(|Member { index, .. }| *index == leaf_index)
})
}
pub(crate) fn message_secrets_mut(&mut self) -> &mut MessageSecrets {
&mut self.message_secrets
}
pub(crate) fn message_secrets(&self) -> &MessageSecrets {
&self.message_secrets
}
fn delete_past_epoch_secrets_older_than_duration(&mut self, duration: std::time::Duration) {
if let Some(added_at) = self.message_secrets.timestamp() {
if let Ok(elapsed) = SystemTime::now().duration_since(added_at) {
if elapsed > duration {
self.past_epoch_trees.clear();
return;
}
}
}
let found = self
.past_epoch_trees
.iter()
.enumerate()
.rev()
.find(|(_idx, tree)| {
let Some(added_at) = tree.message_secrets.timestamp() else {
return false;
};
let Ok(elapsed) = SystemTime::now().duration_since(added_at) else {
return false;
};
elapsed > duration
})
.map(|(idx, _tree)| idx);
if let Some(found_idx) = found {
self.past_epoch_trees.drain(0..found_idx + 1);
} else {
}
}
fn delete_past_epoch_secrets_before_timestamp(&mut self, cutoff: SystemTime) {
if let Some(added_at) = self.message_secrets.timestamp() {
if added_at < cutoff {
self.past_epoch_trees.clear();
return;
}
}
let found = self
.past_epoch_trees
.iter()
.enumerate()
.rev()
.find(|(_idx, tree)| {
let Some(added_at) = tree.message_secrets.timestamp() else {
return false;
};
added_at < cutoff
})
.map(|(idx, _tree)| idx);
if let Some(found_idx) = found {
self.past_epoch_trees.drain(0..found_idx + 1);
} else {
}
}
pub(crate) fn delete_past_epoch_secrets(&mut self, policy: PastEpochDeletion) {
if let Some(config) = policy.config {
match config {
PastEpochDeletionTimeConfig::DeleteAllWithoutTimestamp => {
self.past_epoch_trees
.retain(|tree| tree.message_secrets.timestamp().is_some());
}
PastEpochDeletionTimeConfig::BeforeTimestamp(timestamp) => {
self.delete_past_epoch_secrets_before_timestamp(timestamp)
}
PastEpochDeletionTimeConfig::OlderThanDuration(duration) => {
self.delete_past_epoch_secrets_older_than_duration(duration)
}
};
if let Some(max_past_epochs) = policy.max_past_epochs {
if let Some(i) = self.past_epoch_trees.len().checked_sub(max_past_epochs) {
self.past_epoch_trees.drain(0..i);
}
}
} else {
self.past_epoch_trees.clear();
}
}
#[cfg(all(test, feature = "sqlite-provider", feature = "libcrux-provider"))]
pub(crate) fn iter_past_epoch_trees(&self) -> impl Iterator<Item = &EpochTree> {
self.past_epoch_trees.iter()
}
#[cfg(test)]
pub(crate) fn num_past_epoch_trees(&self) -> usize {
self.past_epoch_trees.len()
}
}