1use aes_gcm::aead::{Aead, KeyInit, Payload};
9use aes_gcm::{Aes256Gcm, Key, Nonce};
10use hkdf::Hkdf;
11use sha2::Sha256;
12use x25519_dalek::{PublicKey, StaticSecret};
13
14pub const WRAP_HKDF_INFO: &[u8] = b"heddle-runtime-wrap-v1";
16
17pub const AEAD_AES256_GCM_V1: &str = "aes-256-gcm-v1";
19
20pub const PAD_BUCKETS: &[usize] = &[32, 64, 128, 256, 512, 1024, 2048, 4096];
22
23const NONCE_LEN: usize = 12;
24const X25519_LEN: usize = 32;
25const DEK_LEN: usize = 32;
26const LENGTH_PREFIX: usize = 4;
27
28#[derive(Clone)]
30pub struct Dek([u8; DEK_LEN]);
31
32impl Drop for Dek {
33 fn drop(&mut self) {
34 self.0.fill(0);
35 }
36}
37
38impl Dek {
39 pub fn generate() -> Result<Self, AeadError> {
40 let mut bytes = [0u8; DEK_LEN];
41 fill_random(&mut bytes)?;
42 Ok(Self(bytes))
43 }
44
45 pub fn from_bytes(bytes: [u8; DEK_LEN]) -> Self {
46 Self(bytes)
47 }
48
49 pub fn as_bytes(&self) -> &[u8; DEK_LEN] {
50 &self.0
51 }
52}
53
54#[derive(Clone)]
59pub struct SoftwareRecipientSecret(StaticSecret);
60
61impl SoftwareRecipientSecret {
62 pub fn generate() -> Result<Self, AeadError> {
63 let mut seed = [0u8; X25519_LEN];
64 fill_random(&mut seed)?;
65 Ok(Self(StaticSecret::from(seed)))
66 }
67
68 pub fn from_bytes(bytes: [u8; X25519_LEN]) -> Self {
69 Self(StaticSecret::from(bytes))
70 }
71
72 pub fn to_bytes(&self) -> [u8; X25519_LEN] {
73 self.0.to_bytes()
74 }
75
76 pub fn public_key(&self) -> [u8; X25519_LEN] {
77 PublicKey::from(&self.0).to_bytes()
78 }
79}
80
81#[derive(Clone, Debug, PartialEq, Eq)]
83pub struct AeadCiphertext {
84 pub alg: &'static str,
85 pub nonce: [u8; NONCE_LEN],
86 pub ciphertext: Vec<u8>,
87 pub pad_bucket: u32,
88}
89
90#[derive(Clone, Debug, PartialEq, Eq)]
92pub struct WrappedDek {
93 pub ephemeral_public: [u8; X25519_LEN],
94 pub nonce: [u8; NONCE_LEN],
95 pub ciphertext: Vec<u8>,
96}
97
98#[derive(Debug, thiserror::Error)]
99pub enum AeadError {
100 #[error("secure random generation failed: {0}")]
101 Random(String),
102 #[error("aead encryption failed")]
103 Encrypt,
104 #[error("aead decryption failed")]
105 Decrypt,
106 #[error("hkdf expansion failed")]
107 Hkdf,
108 #[error("wrapped dek is truncated")]
109 TruncatedWrap,
110 #[error("padded plaintext is truncated or corrupt")]
111 CorruptPadding,
112}
113
114pub fn pad_bucket_for(plaintext_len: usize) -> usize {
116 let needed = plaintext_len.saturating_add(LENGTH_PREFIX);
117 for &bucket in PAD_BUCKETS {
118 if needed <= bucket {
119 return bucket;
120 }
121 }
122 needed.div_ceil(4096).saturating_mul(4096)
123}
124
125fn fill_random(dest: &mut [u8]) -> Result<(), AeadError> {
126 getrandom::fill(dest).map_err(|err| AeadError::Random(err.to_string()))
127}
128
129fn hkdf_sha256(salt: &[u8], ikm: &[u8], info: &[u8]) -> Result<[u8; DEK_LEN], AeadError> {
131 let hk = Hkdf::<Sha256>::new(Some(salt), ikm);
132 let mut okm = [0u8; DEK_LEN];
133 hk.expand(info, &mut okm).map_err(|_| AeadError::Hkdf)?;
134 Ok(okm)
135}
136
137fn pad_plaintext(plaintext: &[u8]) -> Result<(Vec<u8>, u32), AeadError> {
138 let len = u32::try_from(plaintext.len()).map_err(|_| AeadError::CorruptPadding)?;
139 let bucket = pad_bucket_for(plaintext.len());
140 let bucket_u32 = u32::try_from(bucket).map_err(|_| AeadError::CorruptPadding)?;
141 let mut out = vec![0u8; bucket];
142 out[..LENGTH_PREFIX].copy_from_slice(&len.to_be_bytes());
143 let end = LENGTH_PREFIX + plaintext.len();
144 if end > bucket {
145 return Err(AeadError::CorruptPadding);
146 }
147 out[LENGTH_PREFIX..end].copy_from_slice(plaintext);
148 Ok((out, bucket_u32))
149}
150
151fn unpad_plaintext(padded: &[u8]) -> Result<Vec<u8>, AeadError> {
152 if padded.len() < LENGTH_PREFIX {
153 return Err(AeadError::CorruptPadding);
154 }
155 let mut len_bytes = [0u8; LENGTH_PREFIX];
156 len_bytes.copy_from_slice(&padded[..LENGTH_PREFIX]);
157 let len = usize::try_from(u32::from_be_bytes(len_bytes)).unwrap_or(usize::MAX);
158 let end = LENGTH_PREFIX.saturating_add(len);
159 if end > padded.len() {
160 return Err(AeadError::CorruptPadding);
161 }
162 if padded[end..].iter().any(|byte| *byte != 0) {
163 return Err(AeadError::CorruptPadding);
164 }
165 Ok(padded[LENGTH_PREFIX..end].to_vec())
166}
167
168pub fn encrypt_padded(
171 dek: &Dek,
172 plaintext: &[u8],
173 aad: &[u8],
174) -> Result<AeadCiphertext, AeadError> {
175 let (padded, pad_bucket) = pad_plaintext(plaintext)?;
176 let mut nonce = [0u8; NONCE_LEN];
177 fill_random(&mut nonce)?;
178 let cipher = Aes256Gcm::new(Key::<Aes256Gcm>::from_slice(dek.as_bytes()));
179 let ciphertext = cipher
180 .encrypt(Nonce::from_slice(&nonce), Payload { msg: &padded, aad })
181 .map_err(|_| AeadError::Encrypt)?;
182 Ok(AeadCiphertext {
183 alg: AEAD_AES256_GCM_V1,
184 nonce,
185 ciphertext,
186 pad_bucket,
187 })
188}
189
190pub fn decrypt_padded(
191 dek: &Dek,
192 sealed: &AeadCiphertext,
193 aad: &[u8],
194) -> Result<Vec<u8>, AeadError> {
195 if sealed.alg != AEAD_AES256_GCM_V1 {
196 return Err(AeadError::Decrypt);
197 }
198 let cipher = Aes256Gcm::new(Key::<Aes256Gcm>::from_slice(dek.as_bytes()));
199 let padded = cipher
200 .decrypt(
201 Nonce::from_slice(&sealed.nonce),
202 Payload {
203 msg: &sealed.ciphertext,
204 aad,
205 },
206 )
207 .map_err(|_| AeadError::Decrypt)?;
208 if padded.len() != sealed.pad_bucket as usize {
209 return Err(AeadError::CorruptPadding);
210 }
211 unpad_plaintext(&padded)
212}
213
214fn wrap_key(
215 ephemeral_secret: &StaticSecret,
216 recipient_public: &[u8; X25519_LEN],
217) -> Result<[u8; DEK_LEN], AeadError> {
218 let shared = ephemeral_secret.diffie_hellman(&PublicKey::from(*recipient_public));
219 let ephemeral_public = PublicKey::from(ephemeral_secret).to_bytes();
220 let mut salt = [0u8; X25519_LEN * 2];
221 salt[..X25519_LEN].copy_from_slice(&ephemeral_public);
222 salt[X25519_LEN..].copy_from_slice(recipient_public);
223 hkdf_sha256(&salt, shared.as_bytes(), WRAP_HKDF_INFO)
224}
225
226pub fn wrap_dek(dek: &Dek, recipient_public: &[u8; X25519_LEN]) -> Result<WrappedDek, AeadError> {
228 let ephemeral = SoftwareRecipientSecret::generate()?;
229 let wrap_key = wrap_key(&ephemeral.0, recipient_public)?;
230 let mut nonce = [0u8; NONCE_LEN];
231 fill_random(&mut nonce)?;
232 let cipher = Aes256Gcm::new(Key::<Aes256Gcm>::from_slice(&wrap_key));
233 let ciphertext = cipher
234 .encrypt(Nonce::from_slice(&nonce), dek.as_bytes().as_slice())
235 .map_err(|_| AeadError::Encrypt)?;
236 Ok(WrappedDek {
237 ephemeral_public: ephemeral.public_key(),
238 nonce,
239 ciphertext,
240 })
241}
242
243pub fn unwrap_dek(
244 wrapped: &WrappedDek,
245 recipient: &SoftwareRecipientSecret,
246) -> Result<Dek, AeadError> {
247 if wrapped.ciphertext.len() < 16 {
248 return Err(AeadError::TruncatedWrap);
249 }
250 let ephemeral_secret_for_shared = &recipient.0;
251 let shared =
252 ephemeral_secret_for_shared.diffie_hellman(&PublicKey::from(wrapped.ephemeral_public));
253 let recipient_public = recipient.public_key();
254 let mut salt = [0u8; X25519_LEN * 2];
255 salt[..X25519_LEN].copy_from_slice(&wrapped.ephemeral_public);
256 salt[X25519_LEN..].copy_from_slice(&recipient_public);
257 let okm = hkdf_sha256(&salt, shared.as_bytes(), WRAP_HKDF_INFO)?;
258 let cipher = Aes256Gcm::new(Key::<Aes256Gcm>::from_slice(&okm));
259 let dek_bytes = cipher
260 .decrypt(
261 Nonce::from_slice(&wrapped.nonce),
262 wrapped.ciphertext.as_slice(),
263 )
264 .map_err(|_| AeadError::Decrypt)?;
265 let dek_arr: [u8; DEK_LEN] = dek_bytes.try_into().map_err(|_| AeadError::TruncatedWrap)?;
266 Ok(Dek::from_bytes(dek_arr))
267}
268
269#[cfg(test)]
270mod tests {
271 use super::*;
272
273 #[test]
274 fn padded_round_trip_hides_exact_length_inside_a_bucket() {
275 let dek = Dek::generate().expect("dek");
276 let sealed = encrypt_padded(&dek, b"secret-value", b"aad-v1").expect("encrypt");
277 assert_eq!(sealed.alg, AEAD_AES256_GCM_V1);
278 assert_eq!(sealed.pad_bucket, 32);
279 assert_ne!(&sealed.ciphertext, b"secret-value");
280 let plain = decrypt_padded(&dek, &sealed, b"aad-v1").expect("decrypt");
281 assert_eq!(plain, b"secret-value");
282 }
283
284 #[test]
285 fn wrong_aad_cannot_decrypt() {
286 let dek = Dek::generate().expect("dek");
287 let sealed = encrypt_padded(&dek, b"secret-value", b"slot-a").expect("encrypt");
288 decrypt_padded(&dek, &sealed, b"slot-b").expect_err("aad mismatch");
289 }
290
291 #[test]
292 fn wrap_round_trip_and_wrong_recipient_fails() {
293 let dek = Dek::generate().expect("dek");
294 let alice = SoftwareRecipientSecret::generate().expect("alice");
295 let bob = SoftwareRecipientSecret::generate().expect("bob");
296 let wrapped = wrap_dek(&dek, &alice.public_key()).expect("wrap");
297 let opened = unwrap_dek(&wrapped, &alice).expect("alice unwraps");
298 assert_eq!(opened.as_bytes(), dek.as_bytes());
299 assert!(
300 unwrap_dek(&wrapped, &bob).is_err(),
301 "bob cannot unwrap alice's wrap"
302 );
303 }
304
305 #[test]
306 fn pad_buckets_jump_to_4k_increments_after_4k() {
307 assert_eq!(pad_bucket_for(0), 32);
308 assert_eq!(pad_bucket_for(28), 32);
309 assert_eq!(pad_bucket_for(29), 64);
310 assert_eq!(pad_bucket_for(4093), 8192);
311 }
312}