use aes_gcm::{
Aes256Gcm,
aead::{Aead, KeyInit, Payload},
};
use hkdf::Hkdf;
use sha2::Sha256;
use x25519_dalek::{PublicKey, SharedSecret, StaticSecret};
use zeroize::Zeroizing;
pub const WRAP_HKDF_INFO: &[u8] = b"heddle-env-wrap-v1";
pub const AEAD_AES256_GCM_V1: &str = "aes-256-gcm-v1";
pub const PAD_BUCKETS: &[usize] = &[32, 64, 128, 256, 512, 1024, 2048, 4096];
const NONCE_LEN: usize = 12;
const X25519_LEN: usize = 32;
const DEK_LEN: usize = 32;
const LENGTH_PREFIX: usize = 4;
#[derive(Clone)]
pub struct Dek([u8; DEK_LEN]);
impl Drop for Dek {
fn drop(&mut self) {
zeroize::Zeroize::zeroize(&mut self.0);
}
}
impl Dek {
pub fn generate() -> Result<Self, AeadError> {
let mut bytes = [0u8; DEK_LEN];
fill_random(&mut bytes)?;
Ok(Self(bytes))
}
pub fn from_bytes(bytes: [u8; DEK_LEN]) -> Self {
Self(bytes)
}
pub fn as_bytes(&self) -> &[u8; DEK_LEN] {
&self.0
}
}
#[derive(Clone)]
pub struct SoftwareRecipientSecret(StaticSecret);
impl SoftwareRecipientSecret {
pub fn generate() -> Result<Self, AeadError> {
let mut seed = [0u8; X25519_LEN];
fill_random(&mut seed)?;
Ok(Self(StaticSecret::from(seed)))
}
pub fn from_bytes(bytes: [u8; X25519_LEN]) -> Self {
Self(StaticSecret::from(bytes))
}
pub fn to_bytes(&self) -> [u8; X25519_LEN] {
self.0.to_bytes()
}
pub fn public_key(&self) -> [u8; X25519_LEN] {
PublicKey::from(&self.0).to_bytes()
}
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct AeadCiphertext {
pub alg: &'static str,
pub nonce: [u8; NONCE_LEN],
pub ciphertext: Vec<u8>,
pub pad_bucket: u32,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct WrappedDek {
pub ephemeral_public: [u8; X25519_LEN],
pub nonce: [u8; NONCE_LEN],
pub ciphertext: Vec<u8>,
}
#[derive(Debug, thiserror::Error)]
pub enum AeadError {
#[error("secure random generation failed: {0}")]
Random(String),
#[error("aead encryption failed")]
Encrypt,
#[error("aead decryption failed")]
Decrypt,
#[error("hkdf expansion failed")]
Hkdf,
#[error("wrapped dek is truncated")]
TruncatedWrap,
#[error("padded plaintext is truncated or corrupt")]
CorruptPadding,
#[error("x25519 shared secret is non-contributory (low-order point)")]
NonContributory,
}
pub fn pad_bucket_for(plaintext_len: usize) -> usize {
let needed = plaintext_len.saturating_add(LENGTH_PREFIX);
for &bucket in PAD_BUCKETS {
if needed <= bucket {
return bucket;
}
}
needed.div_ceil(4096).saturating_mul(4096)
}
fn fill_random(dest: &mut [u8]) -> Result<(), AeadError> {
getrandom::fill(dest).map_err(|err| AeadError::Random(err.to_string()))
}
fn hkdf_sha256(
salt: &[u8],
ikm: &[u8],
info: &[u8],
) -> Result<Zeroizing<[u8; DEK_LEN]>, AeadError> {
let hk = Hkdf::<Sha256>::new(Some(salt), ikm);
let mut okm = Zeroizing::new([0u8; DEK_LEN]);
hk.expand(info, okm.as_mut_slice())
.map_err(|_| AeadError::Hkdf)?;
Ok(okm)
}
fn pad_plaintext(plaintext: &[u8]) -> Result<(Zeroizing<Vec<u8>>, u32), AeadError> {
let len = u32::try_from(plaintext.len()).map_err(|_| AeadError::CorruptPadding)?;
let bucket = pad_bucket_for(plaintext.len());
let bucket_u32 = u32::try_from(bucket).map_err(|_| AeadError::CorruptPadding)?;
let mut out = Zeroizing::new(vec![0u8; bucket]);
out[..LENGTH_PREFIX].copy_from_slice(&len.to_be_bytes());
let end = LENGTH_PREFIX + plaintext.len();
if end > bucket {
return Err(AeadError::CorruptPadding);
}
out[LENGTH_PREFIX..end].copy_from_slice(plaintext);
Ok((out, bucket_u32))
}
fn unpad_plaintext(padded: &[u8]) -> Result<Vec<u8>, AeadError> {
if padded.len() < LENGTH_PREFIX {
return Err(AeadError::CorruptPadding);
}
let mut len_bytes = [0u8; LENGTH_PREFIX];
len_bytes.copy_from_slice(&padded[..LENGTH_PREFIX]);
let len = usize::try_from(u32::from_be_bytes(len_bytes)).unwrap_or(usize::MAX);
let end = LENGTH_PREFIX.saturating_add(len);
if end > padded.len() {
return Err(AeadError::CorruptPadding);
}
if padded[end..].iter().any(|byte| *byte != 0) {
return Err(AeadError::CorruptPadding);
}
Ok(padded[LENGTH_PREFIX..end].to_vec())
}
pub fn encrypt_padded(
dek: &Dek,
plaintext: &[u8],
aad: &[u8],
) -> Result<AeadCiphertext, AeadError> {
let (padded, pad_bucket) = pad_plaintext(plaintext)?;
let mut nonce = [0u8; NONCE_LEN];
fill_random(&mut nonce)?;
let cipher = Aes256Gcm::new(dek.as_bytes().into());
let ciphertext = cipher
.encrypt(
(&nonce).into(),
Payload {
msg: padded.as_slice(),
aad,
},
)
.map_err(|_| AeadError::Encrypt)?;
Ok(AeadCiphertext {
alg: AEAD_AES256_GCM_V1,
nonce,
ciphertext,
pad_bucket,
})
}
pub fn decrypt_padded(
dek: &Dek,
sealed: &AeadCiphertext,
aad: &[u8],
) -> Result<Vec<u8>, AeadError> {
if sealed.alg != AEAD_AES256_GCM_V1 {
return Err(AeadError::Decrypt);
}
let cipher = Aes256Gcm::new(dek.as_bytes().into());
let padded = Zeroizing::new(
cipher
.decrypt(
(&sealed.nonce).into(),
Payload {
msg: &sealed.ciphertext,
aad,
},
)
.map_err(|_| AeadError::Decrypt)?,
);
if padded.len() != sealed.pad_bucket as usize {
return Err(AeadError::CorruptPadding);
}
unpad_plaintext(&padded)
}
fn wrap_key_from_shared(
shared: &SharedSecret,
ephemeral_public: &[u8; X25519_LEN],
recipient_public: &[u8; X25519_LEN],
) -> Result<Zeroizing<[u8; DEK_LEN]>, AeadError> {
if !shared.was_contributory() {
return Err(AeadError::NonContributory);
}
let mut salt = [0u8; X25519_LEN * 2];
salt[..X25519_LEN].copy_from_slice(ephemeral_public);
salt[X25519_LEN..].copy_from_slice(recipient_public);
hkdf_sha256(&salt, shared.as_bytes(), WRAP_HKDF_INFO)
}
pub fn wrap_dek(
dek: &Dek,
recipient_public: &[u8; X25519_LEN],
aad: &[u8],
) -> Result<WrappedDek, AeadError> {
let ephemeral = SoftwareRecipientSecret::generate()?;
let ephemeral_public = ephemeral.public_key();
let shared = ephemeral
.0
.diffie_hellman(&PublicKey::from(*recipient_public));
let wrap_key = wrap_key_from_shared(&shared, &ephemeral_public, recipient_public)?;
let mut nonce = [0u8; NONCE_LEN];
fill_random(&mut nonce)?;
let cipher = Aes256Gcm::new((&*wrap_key).into());
let ciphertext = cipher
.encrypt(
(&nonce).into(),
Payload {
msg: dek.as_bytes().as_slice(),
aad,
},
)
.map_err(|_| AeadError::Encrypt)?;
Ok(WrappedDek {
ephemeral_public,
nonce,
ciphertext,
})
}
pub fn unwrap_dek(
wrapped: &WrappedDek,
recipient: &SoftwareRecipientSecret,
aad: &[u8],
) -> Result<Dek, AeadError> {
if wrapped.ciphertext.len() < 16 {
return Err(AeadError::TruncatedWrap);
}
let shared = recipient
.0
.diffie_hellman(&PublicKey::from(wrapped.ephemeral_public));
let recipient_public = recipient.public_key();
let okm = wrap_key_from_shared(&shared, &wrapped.ephemeral_public, &recipient_public)?;
let cipher = Aes256Gcm::new((&*okm).into());
let dek_bytes = Zeroizing::new(
cipher
.decrypt(
(&wrapped.nonce).into(),
Payload {
msg: wrapped.ciphertext.as_slice(),
aad,
},
)
.map_err(|_| AeadError::Decrypt)?,
);
let dek_arr: [u8; DEK_LEN] = dek_bytes
.as_slice()
.try_into()
.map_err(|_| AeadError::TruncatedWrap)?;
Ok(Dek::from_bytes(dek_arr))
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn padded_round_trip_hides_exact_length_inside_a_bucket() {
let dek = Dek::generate().expect("dek");
let sealed = encrypt_padded(&dek, b"secret-value", b"aad-v1").expect("encrypt");
assert_eq!(sealed.alg, AEAD_AES256_GCM_V1);
assert_eq!(sealed.pad_bucket, 32);
assert_ne!(&sealed.ciphertext, b"secret-value");
let plain = decrypt_padded(&dek, &sealed, b"aad-v1").expect("decrypt");
assert_eq!(plain, b"secret-value");
}
#[test]
fn wrong_aad_cannot_decrypt() {
let dek = Dek::generate().expect("dek");
let sealed = encrypt_padded(&dek, b"secret-value", b"slot-a").expect("encrypt");
decrypt_padded(&dek, &sealed, b"slot-b").expect_err("aad mismatch");
}
#[test]
fn wrap_round_trip_and_wrong_recipient_fails() {
let dek = Dek::generate().expect("dek");
let alice = SoftwareRecipientSecret::generate().expect("alice");
let bob = SoftwareRecipientSecret::generate().expect("bob");
let wrapped = wrap_dek(&dek, &alice.public_key(), b"wrap-aad-v1").expect("wrap");
let opened = unwrap_dek(&wrapped, &alice, b"wrap-aad-v1").expect("alice unwraps");
assert_eq!(opened.as_bytes(), dek.as_bytes());
assert!(
unwrap_dek(&wrapped, &bob, b"wrap-aad-v1").is_err(),
"bob cannot unwrap alice's wrap"
);
}
#[test]
fn wrap_aad_mismatch_cannot_unwrap() {
let dek = Dek::generate().expect("dek");
let alice = SoftwareRecipientSecret::generate().expect("alice");
let wrapped = wrap_dek(&dek, &alice.public_key(), b"recip|profile|SLOT|v1").expect("wrap");
assert!(
unwrap_dek(&wrapped, &alice, b"recip|profile|SLOT|v2").is_err(),
"a wrap must not unwrap under a different binding (transplant/rollback)"
);
}
#[test]
fn low_order_ephemeral_public_is_rejected() {
let dek = Dek::generate().expect("dek");
let alice = SoftwareRecipientSecret::generate().expect("alice");
let good = wrap_dek(&dek, &alice.public_key(), b"aad").expect("wrap");
let forged = WrappedDek {
ephemeral_public: [0u8; X25519_LEN],
nonce: good.nonce,
ciphertext: good.ciphertext,
};
assert!(
matches!(
unwrap_dek(&forged, &alice, b"aad"),
Err(AeadError::NonContributory)
),
"a low-order ephemeral_public must be rejected as non-contributory"
);
}
#[test]
fn pad_buckets_jump_to_4k_increments_after_4k() {
assert_eq!(pad_bucket_for(0), 32);
assert_eq!(pad_bucket_for(28), 32);
assert_eq!(pad_bucket_for(29), 64);
assert_eq!(pad_bucket_for(4093), 8192);
}
}