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