use std::collections::{BTreeMap, BTreeSet};
use std::fmt;
use std::path::Path;
use aes_gcm::Aes256Gcm;
use aes_gcm_siv::aead::{Aead, KeyInit, Payload};
use aes_gcm_siv::{Aes256GcmSiv, Nonce};
use async_trait::async_trait;
use thiserror::Error;
use zeroize::Zeroizing;
pub const BLOB_MAGIC: &[u8; 8] = b"CORIUMB1";
pub const LOG_MAGIC: &[u8; 8] = b"CORIUML1";
const ALGORITHM_AES_256_GCM_SIV: u8 = 1;
const ALGORITHM_AES_256_GCM: u8 = 2;
const BLOB_HEADER_LEN: usize = BLOB_MAGIC.len() + 1 + size_of::<u32>() + size_of::<u64>();
const LOG_HEADER_LEN: usize = LOG_MAGIC.len() + 1 + size_of::<u32>() + size_of::<u64>() + NONCE_LEN;
const AEAD_TAG_LEN: usize = 16;
const NONCE_LEN: usize = 12;
#[derive(Clone, Eq, PartialEq)]
pub struct SecretKey(Zeroizing<[u8; 32]>);
impl SecretKey {
#[must_use]
pub fn new(bytes: [u8; 32]) -> Self {
Self(Zeroizing::new(bytes))
}
pub fn generate() -> Result<Self, CryptError> {
let mut bytes = Zeroizing::new([0_u8; 32]);
getrandom::fill(bytes.as_mut_slice()).map_err(|_| CryptError::RandomnessUnavailable)?;
Ok(Self(bytes))
}
pub fn from_slice(bytes: &[u8]) -> Result<Self, CryptError> {
let bytes = <[u8; 32]>::try_from(bytes).map_err(|_| CryptError::InvalidKeyLength)?;
Ok(Self::new(bytes))
}
fn as_bytes(&self) -> &[u8; 32] {
&self.0
}
}
impl fmt::Debug for SecretKey {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.write_str("SecretKey([REDACTED])")
}
}
#[derive(Clone, Debug, Eq, Hash, Ord, PartialEq, PartialOrd)]
pub struct KeyId(String);
impl KeyId {
pub fn new(value: impl Into<String>) -> Result<Self, KeyError> {
let value = value.into();
if value.is_empty() {
return Err(KeyError::InvalidId);
}
Ok(Self(value))
}
#[must_use]
pub fn as_str(&self) -> &str {
&self.0
}
#[must_use]
pub fn scheme(&self) -> Option<(&str, &str)> {
self.0.split_once(':')
}
}
impl fmt::Display for KeyId {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.write_str(&self.0)
}
}
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub struct BlobHeader {
pub epoch: u32,
pub plaintext_len: u64,
}
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub struct LogHeader {
pub epoch: u32,
pub t: u64,
}
#[derive(Debug, Error)]
pub enum CryptError {
#[error("secret key must be exactly 32 bytes")]
InvalidKeyLength,
#[error("cryptographic randomness is unavailable")]
RandomnessUnavailable,
#[error("invalid encrypted blob header")]
InvalidBlobHeader,
#[error("invalid encrypted log record header")]
InvalidLogHeader,
#[error("unsupported encrypted blob algorithm {0}")]
UnsupportedAlgorithm(u8),
#[error("encrypted blob length does not match its header")]
InvalidBlobLength,
#[error("encryption failed")]
EncryptionFailed,
#[error("authentication failed")]
AuthenticationFailed,
#[error("plaintext is too large to encrypt")]
PlaintextTooLarge,
}
#[derive(Debug, Error)]
pub enum KeyError {
#[error("key identity must not be empty")]
InvalidId,
#[error("key {id} has no material for epoch {epoch}")]
MissingKey {
id: KeyId,
epoch: u32,
},
#[error("key {0} has no current epoch")]
MissingCurrentEpoch(KeyId),
#[error("wrapped key did not contain a 256-bit key")]
InvalidWrappedKey,
#[error(
"key {0} names a key source this build cannot resolve; \
file: and env: are always available"
)]
UnsupportedScheme(KeyId),
#[error("key {id} could not be read: {reason}")]
Unreadable {
id: KeyId,
reason: String,
},
#[error("key {0} does not hold 32 bytes of key material (raw or 64 hex characters)")]
InvalidMaterial(KeyId),
#[error("wrapped key uses epoch {actual}, expected {expected}")]
WrappedEpochMismatch {
expected: u32,
actual: u32,
},
#[error(transparent)]
Crypt(#[from] CryptError),
}
#[async_trait]
pub trait Keyring: Send + Sync {
async fn key(&self, id: &KeyId, epoch: u32) -> Result<SecretKey, KeyError>;
async fn current_epoch(&self, id: &KeyId) -> Result<u32, KeyError>;
async fn wrap(&self, id: &KeyId, epoch: u32, dek: &SecretKey) -> Result<Vec<u8>, KeyError>;
async fn unwrap(&self, id: &KeyId, epoch: u32, wrapped: &[u8]) -> Result<SecretKey, KeyError>;
fn key_ids(&self) -> &[KeyId];
}
#[derive(Clone, Default)]
pub struct StaticKeyring {
keys: BTreeMap<(KeyId, u32), SecretKey>,
current_epochs: BTreeMap<KeyId, u32>,
key_ids: Vec<KeyId>,
}
pub const STATIC_KEY_EPOCH: u32 = 1;
impl StaticKeyring {
pub fn resolve(ids: impl IntoIterator<Item = KeyId>) -> Result<Self, KeyError> {
let mut keyring = Self::default();
for id in ids {
let key = load_key(&id)?;
keyring.insert(id, STATIC_KEY_EPOCH, key, true);
}
Ok(keyring)
}
pub fn insert(&mut self, id: KeyId, epoch: u32, key: SecretKey, current: bool) {
if current {
self.current_epochs.insert(id.clone(), epoch);
}
self.keys.insert((id, epoch), key);
self.key_ids = self
.keys
.keys()
.map(|(id, _)| id.clone())
.collect::<BTreeSet<_>>()
.into_iter()
.collect();
}
}
#[async_trait]
impl Keyring for StaticKeyring {
async fn key(&self, id: &KeyId, epoch: u32) -> Result<SecretKey, KeyError> {
self.keys
.get(&(id.clone(), epoch))
.cloned()
.ok_or_else(|| KeyError::MissingKey {
id: id.clone(),
epoch,
})
}
async fn current_epoch(&self, id: &KeyId) -> Result<u32, KeyError> {
self.current_epochs
.get(id)
.copied()
.ok_or_else(|| KeyError::MissingCurrentEpoch(id.clone()))
}
async fn wrap(&self, id: &KeyId, epoch: u32, dek: &SecretKey) -> Result<Vec<u8>, KeyError> {
let kek = self.key(id, epoch).await?;
let wrapping_key = derive_key(&kek, b"corium/key-wrap");
encrypt_blob(&wrapping_key, epoch, dek.as_bytes()).map_err(Into::into)
}
async fn unwrap(&self, id: &KeyId, epoch: u32, wrapped: &[u8]) -> Result<SecretKey, KeyError> {
let kek = self.key(id, epoch).await?;
let wrapping_key = derive_key(&kek, b"corium/key-wrap");
let header = parse_blob_header(wrapped)?;
if header.epoch != epoch {
return Err(KeyError::WrappedEpochMismatch {
expected: epoch,
actual: header.epoch,
});
}
let plaintext = Zeroizing::new(decrypt_blob(&wrapping_key, wrapped)?);
SecretKey::from_slice(plaintext.as_slice()).map_err(|_| KeyError::InvalidWrappedKey)
}
fn key_ids(&self) -> &[KeyId] {
&self.key_ids
}
}
pub fn load_key(id: &KeyId) -> Result<SecretKey, KeyError> {
let Some((scheme, rest)) = id.scheme() else {
return Err(KeyError::UnsupportedScheme(id.clone()));
};
let material = match scheme {
"file" => Zeroizing::new(std::fs::read(Path::new(rest)).map_err(|error| {
KeyError::Unreadable {
id: id.clone(),
reason: format!("cannot read {rest}: {error}"),
}
})?),
"env" => Zeroizing::new(
std::env::var(rest)
.map_err(|_| KeyError::Unreadable {
id: id.clone(),
reason: format!("environment variable {rest} is not set"),
})?
.into_bytes(),
),
_ => return Err(KeyError::UnsupportedScheme(id.clone())),
};
decode_key_material(&material).ok_or_else(|| KeyError::InvalidMaterial(id.clone()))
}
fn decode_key_material(material: &[u8]) -> Option<SecretKey> {
if let Ok(key) = SecretKey::from_slice(material) {
return Some(key);
}
let trimmed = std::str::from_utf8(material).ok()?.trim();
if trimmed.len() != 64 {
return None;
}
let mut bytes = Zeroizing::new([0_u8; 32]);
for (slot, pair) in bytes.iter_mut().zip(trimmed.as_bytes().chunks(2)) {
let pair = std::str::from_utf8(pair).ok()?;
*slot = u8::from_str_radix(pair, 16).ok()?;
}
Some(SecretKey::new(*bytes))
}
#[must_use]
pub fn derive_key(parent: &SecretKey, context: &[u8]) -> SecretKey {
let mut hasher = blake3::Hasher::new_keyed(parent.as_bytes());
hasher.update(b"corium/derived-key");
hasher.update(context);
SecretKey::new(*hasher.finalize().as_bytes())
}
pub fn parse_blob_header(object: &[u8]) -> Result<BlobHeader, CryptError> {
if object.len() < BLOB_HEADER_LEN || &object[..BLOB_MAGIC.len()] != BLOB_MAGIC {
return Err(CryptError::InvalidBlobHeader);
}
let algorithm = object[BLOB_MAGIC.len()];
if algorithm != ALGORITHM_AES_256_GCM_SIV {
return Err(CryptError::UnsupportedAlgorithm(algorithm));
}
let epoch_offset = BLOB_MAGIC.len() + 1;
let length_offset = epoch_offset + size_of::<u32>();
let epoch = u32::from_be_bytes(
object[epoch_offset..length_offset]
.try_into()
.map_err(|_| CryptError::InvalidBlobHeader)?,
);
let plaintext_len = u64::from_be_bytes(
object[length_offset..BLOB_HEADER_LEN]
.try_into()
.map_err(|_| CryptError::InvalidBlobHeader)?,
);
let plaintext_len =
usize::try_from(plaintext_len).map_err(|_| CryptError::InvalidBlobLength)?;
let expected_len = BLOB_HEADER_LEN
.checked_add(NONCE_LEN)
.and_then(|length| length.checked_add(plaintext_len))
.and_then(|length| length.checked_add(AEAD_TAG_LEN))
.ok_or(CryptError::InvalidBlobLength)?;
if object.len() != expected_len {
return Err(CryptError::InvalidBlobLength);
}
Ok(BlobHeader {
epoch,
plaintext_len: plaintext_len as u64,
})
}
pub fn encrypt_blob(key: &SecretKey, epoch: u32, plaintext: &[u8]) -> Result<Vec<u8>, CryptError> {
let plaintext_len =
u64::try_from(plaintext.len()).map_err(|_| CryptError::PlaintextTooLarge)?;
let mut header = Vec::with_capacity(BLOB_HEADER_LEN);
header.extend_from_slice(BLOB_MAGIC);
header.push(ALGORITHM_AES_256_GCM_SIV);
header.extend_from_slice(&epoch.to_be_bytes());
header.extend_from_slice(&plaintext_len.to_be_bytes());
let plaintext_digest = blake3::hash(plaintext);
let mut nonce_hasher = blake3::Hasher::new_keyed(key.as_bytes());
nonce_hasher.update(b"corium/blob-nonce");
nonce_hasher.update(&header);
nonce_hasher.update(plaintext_digest.as_bytes());
let nonce_digest = nonce_hasher.finalize();
let nonce_bytes = &nonce_digest.as_bytes()[..NONCE_LEN];
let cipher =
Aes256GcmSiv::new_from_slice(key.as_bytes()).map_err(|_| CryptError::InvalidKeyLength)?;
let ciphertext = cipher
.encrypt(
Nonce::from_slice(nonce_bytes),
Payload {
msg: plaintext,
aad: &header,
},
)
.map_err(|_| CryptError::EncryptionFailed)?;
header.extend_from_slice(nonce_bytes);
header.extend_from_slice(&ciphertext);
Ok(header)
}
pub fn decrypt_blob(key: &SecretKey, object: &[u8]) -> Result<Vec<u8>, CryptError> {
let _header = parse_blob_header(object)?;
let header = &object[..BLOB_HEADER_LEN];
let nonce_end = BLOB_HEADER_LEN + NONCE_LEN;
let nonce = Nonce::from_slice(&object[BLOB_HEADER_LEN..nonce_end]);
let ciphertext = &object[nonce_end..];
let cipher =
Aes256GcmSiv::new_from_slice(key.as_bytes()).map_err(|_| CryptError::InvalidKeyLength)?;
cipher
.decrypt(
nonce,
Payload {
msg: ciphertext,
aad: header,
},
)
.map_err(|_| CryptError::AuthenticationFailed)
}
#[must_use]
pub fn is_encrypted_log_record(payload: &[u8]) -> bool {
payload.len() >= LOG_MAGIC.len() && &payload[..LOG_MAGIC.len()] == LOG_MAGIC
}
pub fn parse_log_header(payload: &[u8]) -> Result<LogHeader, CryptError> {
if !is_encrypted_log_record(payload) || payload.len() < LOG_HEADER_LEN + AEAD_TAG_LEN {
return Err(CryptError::InvalidLogHeader);
}
let algorithm = payload[LOG_MAGIC.len()];
if algorithm != ALGORITHM_AES_256_GCM {
return Err(CryptError::UnsupportedAlgorithm(algorithm));
}
let epoch_offset = LOG_MAGIC.len() + 1;
let t_offset = epoch_offset + size_of::<u32>();
let nonce_offset = t_offset + size_of::<u64>();
let epoch = u32::from_be_bytes(
payload[epoch_offset..t_offset]
.try_into()
.map_err(|_| CryptError::InvalidLogHeader)?,
);
let t = u64::from_be_bytes(
payload[t_offset..nonce_offset]
.try_into()
.map_err(|_| CryptError::InvalidLogHeader)?,
);
Ok(LogHeader { epoch, t })
}
fn log_header_and_aad(
epoch: u32,
lineage: &[u8],
log_version: u64,
t: u64,
nonce: &[u8; NONCE_LEN],
) -> (Vec<u8>, Vec<u8>) {
let mut header = Vec::with_capacity(LOG_HEADER_LEN);
header.extend_from_slice(LOG_MAGIC);
header.push(ALGORITHM_AES_256_GCM);
header.extend_from_slice(&epoch.to_be_bytes());
header.extend_from_slice(&t.to_be_bytes());
header.extend_from_slice(nonce);
let mut aad = Vec::with_capacity(b"corium/log-v1".len() + 16 + lineage.len() + header.len());
aad.extend_from_slice(b"corium/log-v1");
aad.extend_from_slice(&(lineage.len() as u64).to_be_bytes());
aad.extend_from_slice(lineage);
aad.extend_from_slice(&log_version.to_be_bytes());
aad.extend_from_slice(&header);
(header, aad)
}
pub fn encrypt_log_record(
key: &SecretKey,
epoch: u32,
lineage: &[u8],
log_version: u64,
t: u64,
plaintext: &[u8],
) -> Result<Vec<u8>, CryptError> {
let mut nonce_bytes = [0_u8; NONCE_LEN];
getrandom::fill(&mut nonce_bytes).map_err(|_| CryptError::RandomnessUnavailable)?;
let (mut header, aad) = log_header_and_aad(epoch, lineage, log_version, t, &nonce_bytes);
let cipher =
Aes256Gcm::new_from_slice(key.as_bytes()).map_err(|_| CryptError::InvalidKeyLength)?;
let ciphertext = cipher
.encrypt(
Nonce::from_slice(&nonce_bytes),
Payload {
msg: plaintext,
aad: &aad,
},
)
.map_err(|_| CryptError::EncryptionFailed)?;
header.extend_from_slice(&ciphertext);
Ok(header)
}
pub fn decrypt_log_record(
key: &SecretKey,
lineage: &[u8],
log_version: u64,
payload: &[u8],
) -> Result<Vec<u8>, CryptError> {
let LogHeader { epoch, t } = parse_log_header(payload)?;
let header = &payload[..LOG_HEADER_LEN];
let nonce_offset = LOG_HEADER_LEN - NONCE_LEN;
let nonce_bytes = <[u8; NONCE_LEN]>::try_from(&header[nonce_offset..])
.map_err(|_| CryptError::InvalidLogHeader)?;
let (_, aad) = log_header_and_aad(epoch, lineage, log_version, t, &nonce_bytes);
let cipher =
Aes256Gcm::new_from_slice(key.as_bytes()).map_err(|_| CryptError::InvalidKeyLength)?;
cipher
.decrypt(
Nonce::from_slice(&nonce_bytes),
Payload {
msg: &payload[LOG_HEADER_LEN..],
aad: &aad,
},
)
.map_err(|_| CryptError::AuthenticationFailed)
}
#[cfg(test)]
mod tests {
use super::*;
use proptest::prelude::*;
fn key(byte: u8) -> SecretKey {
SecretKey::new([byte; 32])
}
proptest! {
#[test]
fn blobs_are_deterministic_and_round_trip(
plaintext in prop::collection::vec(any::<u8>(), 0..4096)
) {
let encrypted = encrypt_blob(&key(7), 3, &plaintext).expect("encrypt");
let repeated = encrypt_blob(&key(7), 3, &plaintext).expect("repeat");
prop_assert_eq!(&encrypted, &repeated);
prop_assert_ne!(
encrypt_blob(&key(7), 4, &plaintext).expect("different epoch"),
encrypted.clone()
);
prop_assert_eq!(decrypt_blob(&key(7), &encrypted).expect("decrypt"), plaintext);
}
}
#[test]
fn header_and_ciphertext_are_authenticated() {
let encrypted = encrypt_blob(&key(1), 9, b"sentinel").expect("encrypt");
assert_eq!(
parse_blob_header(&encrypted).expect("header"),
BlobHeader {
epoch: 9,
plaintext_len: 8,
}
);
assert!(!encrypted.windows(8).any(|window| window == b"sentinel"));
assert!(decrypt_blob(&key(2), &encrypted).is_err());
let mut tampered = encrypted;
*tampered.last_mut().expect("ciphertext") ^= 1;
assert!(decrypt_blob(&key(1), &tampered).is_err());
let mut tampered_nonce = encrypt_blob(&key(1), 9, b"sentinel").expect("encrypt");
tampered_nonce[BLOB_HEADER_LEN] ^= 1;
assert!(decrypt_blob(&key(1), &tampered_nonce).is_err());
}
proptest! {
#[test]
fn log_records_round_trip(payload in prop::collection::vec(any::<u8>(), 0..4096)) {
let encrypted = encrypt_log_record(&key(5), 2, b"people", 7, 42, &payload)
.expect("encrypt");
prop_assert_eq!(
parse_log_header(&encrypted).expect("header"),
LogHeader { epoch: 2, t: 42 }
);
prop_assert_eq!(
decrypt_log_record(&key(5), b"people", 7, &encrypted).expect("decrypt"),
payload
);
}
}
#[test]
fn log_records_are_bound_to_their_position() {
let plaintext = b"sentinel payload".as_slice();
let encrypted = encrypt_log_record(&key(5), 2, b"people", 7, 42, plaintext).expect("seal");
assert!(is_encrypted_log_record(&encrypted));
assert!(
!encrypted
.windows(plaintext.len())
.any(|window| window == plaintext)
);
assert!(decrypt_log_record(&key(5), b"other", 7, &encrypted).is_err());
assert!(decrypt_log_record(&key(5), b"people", 8, &encrypted).is_err());
assert!(decrypt_log_record(&key(6), b"people", 7, &encrypted).is_err());
let mut moved = encrypted.clone();
let t_offset = LOG_MAGIC.len() + 1 + size_of::<u32>();
moved[t_offset..t_offset + size_of::<u64>()].copy_from_slice(&43_u64.to_be_bytes());
assert_eq!(parse_log_header(&moved).expect("header").t, 43);
assert!(decrypt_log_record(&key(5), b"people", 7, &moved).is_err());
let mut retagged = encrypted;
retagged[LOG_MAGIC.len() + 1] ^= 1;
assert!(decrypt_log_record(&key(5), b"people", 7, &retagged).is_err());
}
#[test]
fn log_records_do_not_reuse_a_nonce() {
let first = encrypt_log_record(&key(5), 2, b"people", 7, 42, b"first").expect("first");
let second = encrypt_log_record(&key(5), 2, b"people", 7, 42, b"first").expect("second");
assert_ne!(first, second);
assert_eq!(
decrypt_log_record(&key(5), b"people", 7, &second).expect("decrypt"),
b"first"
);
}
#[test]
fn plaintext_records_are_not_mistaken_for_encrypted_ones() {
let mut plaintext = 42_u64.to_be_bytes().to_vec();
plaintext.extend_from_slice(&[0; 32]);
assert!(!is_encrypted_log_record(&plaintext));
assert!(matches!(
parse_log_header(&plaintext),
Err(CryptError::InvalidLogHeader)
));
assert!(!is_encrypted_log_record(&[]));
assert!(!is_encrypted_log_record(LOG_MAGIC.as_slice().split_at(4).0));
}
#[test]
fn generated_keys_are_distinct() {
let first = SecretKey::generate().expect("generate");
assert_ne!(first, SecretKey::generate().expect("generate"));
assert_ne!(first, SecretKey::new([0; 32]));
}
#[test]
fn debug_never_reveals_key_material() {
let rendered = format!("{:?}", key(0xA5));
assert_eq!(rendered, "SecretKey([REDACTED])");
assert!(!rendered.contains("165"));
}
#[tokio::test]
async fn local_key_sources_accept_raw_and_hex_material() {
let dir = std::env::temp_dir().join(format!("corium-crypt-{}", std::process::id()));
std::fs::create_dir_all(&dir).expect("temp dir");
let raw_path = dir.join("raw.key");
std::fs::write(&raw_path, [7_u8; 32]).expect("write raw");
let hex_path = dir.join("hex.key");
std::fs::write(&hex_path, format!("{}\n", "07".repeat(32))).expect("write hex");
let raw = KeyId::new(format!("file:{}", raw_path.display())).expect("id");
let hex = KeyId::new(format!("file:{}", hex_path.display())).expect("id");
assert_eq!(load_key(&raw).expect("raw"), key(7));
assert_eq!(load_key(&hex).expect("hex"), key(7));
assert!(matches!(
load_key(&KeyId::new("env:CORIUM_KEY_THAT_IS_NOT_SET").expect("id")),
Err(KeyError::Unreadable { .. })
));
let keyring = StaticKeyring::resolve([raw.clone(), hex]).expect("resolve");
assert_eq!(keyring.key_ids().len(), 2);
assert_eq!(
keyring.current_epoch(&raw).await.expect("epoch"),
STATIC_KEY_EPOCH
);
std::fs::remove_dir_all(&dir).expect("clean up");
}
#[test]
fn unresolvable_key_identities_name_what_failed() {
let missing = KeyId::new("file:/nonexistent/corium/storage.key").expect("id");
assert!(matches!(
load_key(&missing),
Err(KeyError::Unreadable { .. })
));
let kms = KeyId::new("awskms:arn:aws:kms:us-west-2:1:key/2f1c").expect("id");
assert!(matches!(
load_key(&kms),
Err(KeyError::UnsupportedScheme(_))
));
assert!(matches!(
load_key(&KeyId::new("storage.key").expect("id")),
Err(KeyError::UnsupportedScheme(_))
));
assert!(decode_key_material(b"too short").is_none());
assert!(decode_key_material(&[0; 31]).is_none());
assert!(decode_key_material("zz".repeat(32).as_bytes()).is_none());
}
#[tokio::test]
async fn static_keyring_resolves_and_wraps_keys() {
let id = KeyId::new("file:test-kek").expect("key id");
let mut keyring = StaticKeyring::default();
keyring.insert(id.clone(), 4, key(4), true);
assert_eq!(keyring.current_epoch(&id).await.expect("epoch"), 4);
assert_eq!(keyring.key_ids(), std::slice::from_ref(&id));
let wrapped = keyring.wrap(&id, 4, &key(8)).await.expect("wrap");
assert_eq!(
keyring.unwrap(&id, 4, &wrapped).await.expect("unwrap"),
key(8)
);
}
}