use std::collections::HashMap;
use std::sync::Mutex;
use async_trait::async_trait;
use crate::core::{KeyId, Timestamp};
use crate::keyring::{DataKey, KeyError, KeyRing, WrappedKey};
#[derive(Debug, Default)]
pub struct MemoryKeyRing {
state: Mutex<RingState>,
}
#[derive(Debug, Default)]
struct RingState {
wrapping: HashMap<String, HashMap<u64, [u8; 32]>>,
destroyed: HashMap<String, (Timestamp, String)>,
generation: u64,
floor: u64,
}
impl MemoryKeyRing {
#[must_use]
pub fn new() -> Self {
Self::default()
}
pub fn rotate(&self) {
self.state.lock().expect("keyring lock").generation += 1;
}
pub fn retire_below(&self, generation: u64) {
self.state.lock().expect("keyring lock").floor = generation;
}
#[must_use]
pub fn current_key_id(&self) -> KeyId {
format!(
"memory-kek-{}",
self.state.lock().expect("keyring lock").generation
)
}
fn tombstone(state: &RingState, scope: &str) -> Option<KeyError> {
state
.destroyed
.get(scope)
.map(|(at, reason)| KeyError::Destroyed {
scope: scope.to_owned(),
at: *at,
reason: reason.clone(),
})
}
fn xor(kek: &[u8; 32], bytes: &[u8]) -> Vec<u8> {
bytes
.iter()
.zip(kek.iter().cycle())
.map(|(b, k)| b ^ k)
.collect()
}
fn kek(state: &mut RingState, scope: &str, generation: u64) -> [u8; 32] {
*state
.wrapping
.entry(scope.to_owned())
.or_default()
.entry(generation)
.or_insert_with(|| {
use rand::Rng as _;
let mut k = [0u8; 32];
#[allow(clippy::disallowed_methods)]
rand::rng().fill_bytes(&mut k);
k
})
}
fn generation_of(wrapped_by: &str) -> Result<u64, KeyError> {
wrapped_by
.strip_prefix("memory-kek-")
.and_then(|g| g.parse::<u64>().ok())
.ok_or_else(|| {
KeyError::Refused(format!(
"'{wrapped_by}' is not a wrapping key id this ring issues"
))
})
}
}
#[async_trait]
impl KeyRing for MemoryKeyRing {
async fn data_key(&self, scope: &str) -> Result<(DataKey, WrappedKey), KeyError> {
let mut state = self.state.lock().expect("keyring lock");
if let Some(gone) = Self::tombstone(&state, scope) {
return Err(gone);
}
let generation = state.generation;
let kek = Self::kek(&mut state, scope, generation);
let mut dek = [0u8; 32];
{
use rand::Rng as _;
#[allow(clippy::disallowed_methods)]
rand::rng().fill_bytes(&mut dek);
}
Ok((
DataKey::new(dek),
WrappedKey {
scope: scope.to_owned(),
wrapped_by: format!("memory-kek-{generation}"),
sealed: Self::xor(&kek, &dek),
},
))
}
async fn open(&self, wrapped: &WrappedKey) -> Result<DataKey, KeyError> {
let state = self.state.lock().expect("keyring lock");
if let Some(gone) = Self::tombstone(&state, &wrapped.scope) {
return Err(gone);
}
let generation = Self::generation_of(&wrapped.wrapped_by)?;
if generation < state.floor {
return Err(KeyError::Retired {
scope: wrapped.scope.clone(),
key_id: wrapped.wrapped_by.clone(),
});
}
let Some(kek) = state
.wrapping
.get(&wrapped.scope)
.and_then(|generations| generations.get(&generation))
else {
return Err(KeyError::Refused(format!(
"no wrapping key for scope '{}' at generation {generation}",
wrapped.scope
)));
};
let raw = Self::xor(kek, &wrapped.sealed);
let mut dek = [0u8; 32];
if raw.len() != dek.len() {
return Err(KeyError::Refused(
"a wrapped data key is not the right length".to_owned(),
));
}
dek.copy_from_slice(&raw);
Ok(DataKey::new(dek))
}
async fn destroy(&self, scope: &str, at: Timestamp, reason: &str) -> Result<(), KeyError> {
let mut state = self.state.lock().expect("keyring lock");
state.wrapping.remove(scope);
state
.destroyed
.entry(scope.to_owned())
.or_insert_with(|| (at, reason.to_owned()));
Ok(())
}
}