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, [u8; 32]>,
destroyed: HashMap<String, (Timestamp, String)>,
generation: u64,
}
impl MemoryKeyRing {
#[must_use]
pub fn new() -> Self {
Self::default()
}
pub fn rotate(&self) {
self.state.lock().expect("keyring lock").generation += 1;
}
#[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) -> [u8; 32] {
*state.wrapping.entry(scope.to_owned()).or_insert_with(|| {
use rand::RngCore as _;
let mut k = [0u8; 32];
#[allow(clippy::disallowed_methods)]
rand::rng().fill_bytes(&mut k);
k
})
}
}
#[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);
let mut dek = [0u8; 32];
{
use rand::RngCore 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 mut state = self.state.lock().expect("keyring lock");
if let Some(gone) = Self::tombstone(&state, &wrapped.scope) {
return Err(gone);
}
if !state.wrapping.contains_key(&wrapped.scope) {
return Err(KeyError::Refused(format!(
"no wrapping key for scope '{}'",
wrapped.scope
)));
}
let kek = Self::kek(&mut state, &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(())
}
async fn rewrap(&self, wrapped: &WrappedKey) -> Result<WrappedKey, KeyError> {
let dek = self.open(wrapped).await?;
let mut state = self.state.lock().expect("keyring lock");
let generation = state.generation;
let kek = Self::kek(&mut state, &wrapped.scope);
Ok(WrappedKey {
scope: wrapped.scope.clone(),
wrapped_by: format!("memory-kek-{generation}"),
sealed: Self::xor(&kek, dek.expose()),
})
}
}