1use aes_gcm::aead::{Aead, KeyInit, Payload};
9use aes_gcm::{Aes256Gcm, Key, Nonce};
10use hkdf::Hkdf;
11use sha2::Sha256;
12use x25519_dalek::{PublicKey, SharedSecret, StaticSecret};
13use zeroize::Zeroizing;
14
15pub const WRAP_HKDF_INFO: &[u8] = b"heddle-env-wrap-v1";
17
18pub const AEAD_AES256_GCM_V1: &str = "aes-256-gcm-v1";
20
21pub const PAD_BUCKETS: &[usize] = &[32, 64, 128, 256, 512, 1024, 2048, 4096];
23
24const NONCE_LEN: usize = 12;
25const X25519_LEN: usize = 32;
26const DEK_LEN: usize = 32;
27const LENGTH_PREFIX: usize = 4;
28
29#[derive(Clone)]
31pub struct Dek([u8; DEK_LEN]);
32
33impl Drop for Dek {
34 fn drop(&mut self) {
35 zeroize::Zeroize::zeroize(&mut self.0);
37 }
38}
39
40impl Dek {
41 pub fn generate() -> Result<Self, AeadError> {
42 let mut bytes = [0u8; DEK_LEN];
43 fill_random(&mut bytes)?;
44 Ok(Self(bytes))
45 }
46
47 pub fn from_bytes(bytes: [u8; DEK_LEN]) -> Self {
48 Self(bytes)
49 }
50
51 pub fn as_bytes(&self) -> &[u8; DEK_LEN] {
52 &self.0
53 }
54}
55
56#[derive(Clone)]
62pub struct SoftwareRecipientSecret(StaticSecret);
63
64impl SoftwareRecipientSecret {
65 pub fn generate() -> Result<Self, AeadError> {
66 let mut seed = [0u8; X25519_LEN];
67 fill_random(&mut seed)?;
68 Ok(Self(StaticSecret::from(seed)))
69 }
70
71 pub fn from_bytes(bytes: [u8; X25519_LEN]) -> Self {
72 Self(StaticSecret::from(bytes))
73 }
74
75 pub fn to_bytes(&self) -> [u8; X25519_LEN] {
76 self.0.to_bytes()
77 }
78
79 pub fn public_key(&self) -> [u8; X25519_LEN] {
80 PublicKey::from(&self.0).to_bytes()
81 }
82}
83
84#[derive(Clone, Debug, PartialEq, Eq)]
86pub struct AeadCiphertext {
87 pub alg: &'static str,
88 pub nonce: [u8; NONCE_LEN],
89 pub ciphertext: Vec<u8>,
90 pub pad_bucket: u32,
91}
92
93#[derive(Clone, Debug, PartialEq, Eq)]
95pub struct WrappedDek {
96 pub ephemeral_public: [u8; X25519_LEN],
97 pub nonce: [u8; NONCE_LEN],
98 pub ciphertext: Vec<u8>,
99}
100
101#[derive(Debug, thiserror::Error)]
102pub enum AeadError {
103 #[error("secure random generation failed: {0}")]
104 Random(String),
105 #[error("aead encryption failed")]
106 Encrypt,
107 #[error("aead decryption failed")]
108 Decrypt,
109 #[error("hkdf expansion failed")]
110 Hkdf,
111 #[error("wrapped dek is truncated")]
112 TruncatedWrap,
113 #[error("padded plaintext is truncated or corrupt")]
114 CorruptPadding,
115 #[error("x25519 shared secret is non-contributory (low-order point)")]
116 NonContributory,
117}
118
119pub fn pad_bucket_for(plaintext_len: usize) -> usize {
121 let needed = plaintext_len.saturating_add(LENGTH_PREFIX);
122 for &bucket in PAD_BUCKETS {
123 if needed <= bucket {
124 return bucket;
125 }
126 }
127 needed.div_ceil(4096).saturating_mul(4096)
128}
129
130fn fill_random(dest: &mut [u8]) -> Result<(), AeadError> {
131 getrandom::fill(dest).map_err(|err| AeadError::Random(err.to_string()))
132}
133
134fn hkdf_sha256(
137 salt: &[u8],
138 ikm: &[u8],
139 info: &[u8],
140) -> Result<Zeroizing<[u8; DEK_LEN]>, AeadError> {
141 let hk = Hkdf::<Sha256>::new(Some(salt), ikm);
142 let mut okm = Zeroizing::new([0u8; DEK_LEN]);
143 hk.expand(info, okm.as_mut_slice())
144 .map_err(|_| AeadError::Hkdf)?;
145 Ok(okm)
146}
147
148fn pad_plaintext(plaintext: &[u8]) -> Result<(Zeroizing<Vec<u8>>, u32), AeadError> {
149 let len = u32::try_from(plaintext.len()).map_err(|_| AeadError::CorruptPadding)?;
150 let bucket = pad_bucket_for(plaintext.len());
151 let bucket_u32 = u32::try_from(bucket).map_err(|_| AeadError::CorruptPadding)?;
152 let mut out = Zeroizing::new(vec![0u8; bucket]);
153 out[..LENGTH_PREFIX].copy_from_slice(&len.to_be_bytes());
154 let end = LENGTH_PREFIX + plaintext.len();
155 if end > bucket {
156 return Err(AeadError::CorruptPadding);
157 }
158 out[LENGTH_PREFIX..end].copy_from_slice(plaintext);
159 Ok((out, bucket_u32))
160}
161
162fn unpad_plaintext(padded: &[u8]) -> Result<Vec<u8>, AeadError> {
163 if padded.len() < LENGTH_PREFIX {
164 return Err(AeadError::CorruptPadding);
165 }
166 let mut len_bytes = [0u8; LENGTH_PREFIX];
167 len_bytes.copy_from_slice(&padded[..LENGTH_PREFIX]);
168 let len = usize::try_from(u32::from_be_bytes(len_bytes)).unwrap_or(usize::MAX);
169 let end = LENGTH_PREFIX.saturating_add(len);
170 if end > padded.len() {
171 return Err(AeadError::CorruptPadding);
172 }
173 if padded[end..].iter().any(|byte| *byte != 0) {
174 return Err(AeadError::CorruptPadding);
175 }
176 Ok(padded[LENGTH_PREFIX..end].to_vec())
177}
178
179pub fn encrypt_padded(
182 dek: &Dek,
183 plaintext: &[u8],
184 aad: &[u8],
185) -> Result<AeadCiphertext, AeadError> {
186 let (padded, pad_bucket) = pad_plaintext(plaintext)?;
187 let mut nonce = [0u8; NONCE_LEN];
188 fill_random(&mut nonce)?;
189 let cipher = Aes256Gcm::new(Key::<Aes256Gcm>::from_slice(dek.as_bytes()));
190 let ciphertext = cipher
191 .encrypt(
192 Nonce::from_slice(&nonce),
193 Payload {
194 msg: padded.as_slice(),
195 aad,
196 },
197 )
198 .map_err(|_| AeadError::Encrypt)?;
199 Ok(AeadCiphertext {
200 alg: AEAD_AES256_GCM_V1,
201 nonce,
202 ciphertext,
203 pad_bucket,
204 })
205}
206
207pub fn decrypt_padded(
208 dek: &Dek,
209 sealed: &AeadCiphertext,
210 aad: &[u8],
211) -> Result<Vec<u8>, AeadError> {
212 if sealed.alg != AEAD_AES256_GCM_V1 {
213 return Err(AeadError::Decrypt);
214 }
215 let cipher = Aes256Gcm::new(Key::<Aes256Gcm>::from_slice(dek.as_bytes()));
216 let padded = Zeroizing::new(
217 cipher
218 .decrypt(
219 Nonce::from_slice(&sealed.nonce),
220 Payload {
221 msg: &sealed.ciphertext,
222 aad,
223 },
224 )
225 .map_err(|_| AeadError::Decrypt)?,
226 );
227 if padded.len() != sealed.pad_bucket as usize {
228 return Err(AeadError::CorruptPadding);
229 }
230 unpad_plaintext(&padded)
231}
232
233fn wrap_key_from_shared(
237 shared: &SharedSecret,
238 ephemeral_public: &[u8; X25519_LEN],
239 recipient_public: &[u8; X25519_LEN],
240) -> Result<Zeroizing<[u8; DEK_LEN]>, AeadError> {
241 if !shared.was_contributory() {
242 return Err(AeadError::NonContributory);
243 }
244 let mut salt = [0u8; X25519_LEN * 2];
245 salt[..X25519_LEN].copy_from_slice(ephemeral_public);
246 salt[X25519_LEN..].copy_from_slice(recipient_public);
247 hkdf_sha256(&salt, shared.as_bytes(), WRAP_HKDF_INFO)
248}
249
250pub fn wrap_dek(
254 dek: &Dek,
255 recipient_public: &[u8; X25519_LEN],
256 aad: &[u8],
257) -> Result<WrappedDek, AeadError> {
258 let ephemeral = SoftwareRecipientSecret::generate()?;
259 let ephemeral_public = ephemeral.public_key();
260 let shared = ephemeral.0.diffie_hellman(&PublicKey::from(*recipient_public));
261 let wrap_key = wrap_key_from_shared(&shared, &ephemeral_public, recipient_public)?;
262 let mut nonce = [0u8; NONCE_LEN];
263 fill_random(&mut nonce)?;
264 let cipher = Aes256Gcm::new(Key::<Aes256Gcm>::from_slice(wrap_key.as_slice()));
265 let ciphertext = cipher
266 .encrypt(
267 Nonce::from_slice(&nonce),
268 Payload {
269 msg: dek.as_bytes().as_slice(),
270 aad,
271 },
272 )
273 .map_err(|_| AeadError::Encrypt)?;
274 Ok(WrappedDek {
275 ephemeral_public,
276 nonce,
277 ciphertext,
278 })
279}
280
281pub fn unwrap_dek(
282 wrapped: &WrappedDek,
283 recipient: &SoftwareRecipientSecret,
284 aad: &[u8],
285) -> Result<Dek, AeadError> {
286 if wrapped.ciphertext.len() < 16 {
287 return Err(AeadError::TruncatedWrap);
288 }
289 let shared = recipient
290 .0
291 .diffie_hellman(&PublicKey::from(wrapped.ephemeral_public));
292 let recipient_public = recipient.public_key();
293 let okm = wrap_key_from_shared(&shared, &wrapped.ephemeral_public, &recipient_public)?;
294 let cipher = Aes256Gcm::new(Key::<Aes256Gcm>::from_slice(okm.as_slice()));
295 let dek_bytes = Zeroizing::new(
296 cipher
297 .decrypt(
298 Nonce::from_slice(&wrapped.nonce),
299 Payload {
300 msg: wrapped.ciphertext.as_slice(),
301 aad,
302 },
303 )
304 .map_err(|_| AeadError::Decrypt)?,
305 );
306 let dek_arr: [u8; DEK_LEN] = dek_bytes
307 .as_slice()
308 .try_into()
309 .map_err(|_| AeadError::TruncatedWrap)?;
310 Ok(Dek::from_bytes(dek_arr))
311}
312
313#[cfg(test)]
314mod tests {
315 use super::*;
316
317 #[test]
318 fn padded_round_trip_hides_exact_length_inside_a_bucket() {
319 let dek = Dek::generate().expect("dek");
320 let sealed = encrypt_padded(&dek, b"secret-value", b"aad-v1").expect("encrypt");
321 assert_eq!(sealed.alg, AEAD_AES256_GCM_V1);
322 assert_eq!(sealed.pad_bucket, 32);
323 assert_ne!(&sealed.ciphertext, b"secret-value");
324 let plain = decrypt_padded(&dek, &sealed, b"aad-v1").expect("decrypt");
325 assert_eq!(plain, b"secret-value");
326 }
327
328 #[test]
329 fn wrong_aad_cannot_decrypt() {
330 let dek = Dek::generate().expect("dek");
331 let sealed = encrypt_padded(&dek, b"secret-value", b"slot-a").expect("encrypt");
332 decrypt_padded(&dek, &sealed, b"slot-b").expect_err("aad mismatch");
333 }
334
335 #[test]
336 fn wrap_round_trip_and_wrong_recipient_fails() {
337 let dek = Dek::generate().expect("dek");
338 let alice = SoftwareRecipientSecret::generate().expect("alice");
339 let bob = SoftwareRecipientSecret::generate().expect("bob");
340 let wrapped = wrap_dek(&dek, &alice.public_key(), b"wrap-aad-v1").expect("wrap");
341 let opened = unwrap_dek(&wrapped, &alice, b"wrap-aad-v1").expect("alice unwraps");
342 assert_eq!(opened.as_bytes(), dek.as_bytes());
343 assert!(
344 unwrap_dek(&wrapped, &bob, b"wrap-aad-v1").is_err(),
345 "bob cannot unwrap alice's wrap"
346 );
347 }
348
349 #[test]
350 fn wrap_aad_mismatch_cannot_unwrap() {
351 let dek = Dek::generate().expect("dek");
352 let alice = SoftwareRecipientSecret::generate().expect("alice");
353 let wrapped = wrap_dek(&dek, &alice.public_key(), b"recip|profile|SLOT|v1").expect("wrap");
354 assert!(
355 unwrap_dek(&wrapped, &alice, b"recip|profile|SLOT|v2").is_err(),
356 "a wrap must not unwrap under a different binding (transplant/rollback)"
357 );
358 }
359
360 #[test]
361 fn low_order_ephemeral_public_is_rejected() {
362 let dek = Dek::generate().expect("dek");
366 let alice = SoftwareRecipientSecret::generate().expect("alice");
367 let good = wrap_dek(&dek, &alice.public_key(), b"aad").expect("wrap");
368 let forged = WrappedDek {
370 ephemeral_public: [0u8; X25519_LEN],
371 nonce: good.nonce,
372 ciphertext: good.ciphertext,
373 };
374 assert!(
375 matches!(
376 unwrap_dek(&forged, &alice, b"aad"),
377 Err(AeadError::NonContributory)
378 ),
379 "a low-order ephemeral_public must be rejected as non-contributory"
380 );
381 }
382
383 #[test]
384 fn pad_buckets_jump_to_4k_increments_after_4k() {
385 assert_eq!(pad_bucket_for(0), 32);
386 assert_eq!(pad_bucket_for(28), 32);
387 assert_eq!(pad_bucket_for(29), 64);
388 assert_eq!(pad_bucket_for(4093), 8192);
389 }
390}