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