1use crate::ring_encryption::{RingAeadEncryption, RingAeadEncryptionOptions};
2use crate::*;
3use async_trait::*;
4use ring::rand::SystemRandom;
5use rsb_derive::*;
6use secret_vault_value::SecretValue;
7
8#[async_trait]
9pub trait KmsAeadRingEncryptionProvider {
10 async fn encrypt_data_encryption_key(
11 &self,
12 encryption_key: &DataEncryptionKey,
13 ) -> KmsAeadResult<EncryptedDataEncryptionKey>;
14
15 async fn decrypt_data_encryption_key(
16 &self,
17 encrypted_key: &EncryptedDataEncryptionKey,
18 ) -> KmsAeadResult<DataEncryptionKey>;
19
20 async fn generate_encryption_key(
21 &self,
22 aead_encryption: &RingAeadEncryption,
23 ) -> KmsAeadResult<DataEncryptionKey>;
24}
25
26pub struct KmsAeadRingEnvelopeEncryption<P>
27where
28 P: KmsAeadRingEncryptionProvider + Send + Sync,
29{
30 provider: P,
31 aead_encryption: RingAeadEncryption,
32}
33
34#[derive(Debug, Clone, Builder)]
35pub struct KmsAeadRingEnvelopeEncryptionOptions {
36 #[default = "RingAeadEncryptionOptions::new()"]
37 pub encryption_options: RingAeadEncryptionOptions,
38}
39
40impl<P> KmsAeadRingEnvelopeEncryption<P>
41where
42 P: KmsAeadRingEncryptionProvider + Send + Sync,
43{
44 pub async fn new(provider: P) -> KmsAeadResult<Self> {
45 Self::with_algorithm(provider, &ring::aead::CHACHA20_POLY1305).await
46 }
47
48 pub async fn with_algorithm(
49 provider: P,
50 algo: &'static ring::aead::Algorithm,
51 ) -> KmsAeadResult<Self> {
52 Self::with_algorithm_options(provider, algo, KmsAeadRingEnvelopeEncryptionOptions::new())
53 .await
54 }
55
56 pub async fn with_options(
57 provider: P,
58 options: KmsAeadRingEnvelopeEncryptionOptions,
59 ) -> KmsAeadResult<Self> {
60 Self::with_algorithm_options(provider, &ring::aead::CHACHA20_POLY1305, options).await
61 }
62
63 pub async fn with_algorithm_options(
64 provider: P,
65 algo: &'static ring::aead::Algorithm,
66 options: KmsAeadRingEnvelopeEncryptionOptions,
67 ) -> KmsAeadResult<Self> {
68 let secure_rand = SystemRandom::new();
69 let aead_encryption = RingAeadEncryption::with_algorithm_options(
70 algo,
71 secure_rand,
72 options.encryption_options,
73 )?;
74
75 Ok(Self {
76 provider,
77 aead_encryption,
78 })
79 }
80
81 async fn new_dek(&self) -> KmsAeadResult<(DataEncryptionKey, EncryptedDataEncryptionKey)> {
82 let dek = self
83 .provider
84 .generate_encryption_key(&self.aead_encryption)
85 .await?;
86
87 let encrypted_dek = self.provider.encrypt_data_encryption_key(&dek).await?;
88
89 Ok((dek, encrypted_dek))
90 }
91
92 async fn encrypt_value_with_new_dek<Aad>(
93 &self,
94 aad: &Aad,
95 plain_text: &SecretValue,
96 ) -> KmsAeadResult<(CipherText, EncryptedDataEncryptionKey)>
97 where
98 Aad: AsRef<[u8]> + Send + Sync + 'static,
99 {
100 let (new_dek, new_encrypted_dek) = self.new_dek().await?;
101
102 let cipher_text = self
103 .aead_encryption
104 .encrypt_value(aad, plain_text, &new_dek)
105 .await?;
106
107 Ok((cipher_text, new_encrypted_dek))
108 }
109}
110
111#[async_trait]
112impl<Aad, P> KmsAeadEnvelopeEncryption<Aad> for KmsAeadRingEnvelopeEncryption<P>
113where
114 Aad: AsRef<[u8]> + Send + Sync + 'static,
115 P: KmsAeadRingEncryptionProvider + Send + Sync,
116{
117 async fn encrypt_value(
118 &self,
119 aad: &Aad,
120 plain_text: &SecretValue,
121 ) -> KmsAeadResult<CipherTextWithEncryptedKey> {
122 let (cipher_text, dek) = self.encrypt_value_with_new_dek(aad, plain_text).await?;
123 Ok(CipherTextWithEncryptedKey::new(&cipher_text, &dek))
124 }
125
126 async fn decrypt_value(
127 &self,
128 aad: &Aad,
129 cipher_text: &CipherTextWithEncryptedKey,
130 ) -> KmsAeadResult<SecretValue> {
131 let (cipher_text, encrypted_dek) = cipher_text.separate()?;
132 self.decrypt_value_with_encrypted_dek(aad, &cipher_text, &encrypted_dek)
133 .await
134 }
135
136 async fn encrypt_value_with_dek(
137 &self,
138 aad: &Aad,
139 plain_text: &SecretValue,
140 dek: &DataEncryptionKey,
141 ) -> KmsAeadResult<CipherText> {
142 let cipher_text = self
143 .aead_encryption
144 .encrypt_value(aad, plain_text, dek)
145 .await?;
146
147 Ok(cipher_text)
148 }
149
150 async fn encrypt_value_with_encrypted_dek(
151 &self,
152 aad: &Aad,
153 plain_text: &SecretValue,
154 dek: &EncryptedDataEncryptionKey,
155 ) -> KmsAeadResult<CipherText> {
156 let dek = self.provider.decrypt_data_encryption_key(dek).await?;
157
158 self.encrypt_value_with_dek(aad, plain_text, &dek).await
159 }
160
161 async fn decrypt_value_with_dek(
162 &self,
163 aad: &Aad,
164 cipher_text: &CipherText,
165 data_encryption_key: &DataEncryptionKey,
166 ) -> KmsAeadResult<SecretValue> {
167 self.aead_encryption
168 .decrypt_value(aad, cipher_text, data_encryption_key)
169 .await
170 }
171 async fn decrypt_value_with_encrypted_dek(
172 &self,
173 aad: &Aad,
174 cipher_text: &CipherText,
175 encrypted_data_encryption_key: &EncryptedDataEncryptionKey,
176 ) -> KmsAeadResult<SecretValue> {
177 let dek = self
178 .provider
179 .decrypt_data_encryption_key(encrypted_data_encryption_key)
180 .await?;
181
182 self.decrypt_value_with_dek(aad, cipher_text, &dek).await
183 }
184
185 async fn generate_new_dek(
186 &self,
187 ) -> KmsAeadResult<(DataEncryptionKey, EncryptedDataEncryptionKey)> {
188 self.new_dek().await
189 }
190}
191
192#[cfg(test)]
193mod tests {
194 use super::*;
195 use rvstruct::ValueStruct;
196 use std::sync::{Arc, Mutex};
197
198 #[derive(Clone)]
200 struct MockKmsProvider {
201 encrypted_keys: Arc<Mutex<Vec<Vec<u8>>>>,
202 fail_encrypt: bool,
203 fail_decrypt: bool,
204 }
205
206 impl MockKmsProvider {
207 fn new() -> Self {
208 Self {
209 encrypted_keys: Arc::new(Mutex::new(Vec::new())),
210 fail_encrypt: false,
211 fail_decrypt: false,
212 }
213 }
214
215 fn with_fail_encrypt() -> Self {
216 Self {
217 encrypted_keys: Arc::new(Mutex::new(Vec::new())),
218 fail_encrypt: true,
219 fail_decrypt: false,
220 }
221 }
222
223 fn with_fail_decrypt() -> Self {
224 Self {
225 encrypted_keys: Arc::new(Mutex::new(Vec::new())),
226 fail_encrypt: false,
227 fail_decrypt: true,
228 }
229 }
230 }
231
232 #[async_trait]
233 impl KmsAeadRingEncryptionProvider for MockKmsProvider {
234 async fn encrypt_data_encryption_key(
235 &self,
236 encryption_key: &DataEncryptionKey,
237 ) -> KmsAeadResult<EncryptedDataEncryptionKey> {
238 if self.fail_encrypt {
239 return Err(crate::errors::KmsAeadEncryptionError::create(
240 "MOCK_ENCRYPT_FAIL",
241 "Mock provider configured to fail encryption",
242 ));
243 }
244
245 let encrypted = encryption_key.value().ref_sensitive_value().to_vec();
247 self.encrypted_keys.lock().unwrap().push(encrypted.clone());
248 Ok(EncryptedDataEncryptionKey::from(encrypted))
249 }
250
251 async fn decrypt_data_encryption_key(
252 &self,
253 encrypted_key: &EncryptedDataEncryptionKey,
254 ) -> KmsAeadResult<DataEncryptionKey> {
255 if self.fail_decrypt {
256 return Err(crate::errors::KmsAeadEncryptionError::create(
257 "MOCK_DECRYPT_FAIL",
258 "Mock provider configured to fail decryption",
259 ));
260 }
261
262 Ok(DataEncryptionKey::from(SecretValue::from(
264 encrypted_key.value().clone(),
265 )))
266 }
267
268 async fn generate_encryption_key(
269 &self,
270 aead_encryption: &RingAeadEncryption,
271 ) -> KmsAeadResult<DataEncryptionKey> {
272 aead_encryption.generate_data_encryption_key()
273 }
274 }
275
276 #[tokio::test]
277 async fn test_envelope_encryption_roundtrip() {
278 let provider = MockKmsProvider::new();
279 let encryption = KmsAeadRingEnvelopeEncryption::new(provider).await.unwrap();
280
281 let aad = "test-aad";
282 let plaintext = SecretValue::from("secret message");
283
284 let ciphertext = encryption.encrypt_value(&aad, &plaintext).await.unwrap();
285 let decrypted = encryption.decrypt_value(&aad, &ciphertext).await.unwrap();
286
287 assert_eq!(decrypted, plaintext);
288 }
289
290 #[tokio::test]
291 async fn test_envelope_encryption_with_provided_dek() {
292 let provider = MockKmsProvider::new();
293 let encryption = KmsAeadRingEnvelopeEncryption::new(provider).await.unwrap();
294
295 let aad = "test-aad";
296 let plaintext = SecretValue::from("secret message");
297
298 let (dek, encrypted_dek) = KmsAeadEnvelopeEncryption::<&str>::generate_new_dek(&encryption)
300 .await
301 .unwrap();
302
303 let ciphertext = encryption
305 .encrypt_value_with_dek(&aad, &plaintext, &dek)
306 .await
307 .unwrap();
308
309 let decrypted = encryption
311 .decrypt_value_with_dek(&aad, &ciphertext, &dek)
312 .await
313 .unwrap();
314
315 assert_eq!(decrypted, plaintext);
316
317 let decrypted2 = encryption
319 .decrypt_value_with_encrypted_dek(&aad, &ciphertext, &encrypted_dek)
320 .await
321 .unwrap();
322
323 assert_eq!(decrypted2, plaintext);
324 }
325
326 #[tokio::test]
327 async fn test_envelope_encryption_with_encrypted_dek() {
328 let provider = MockKmsProvider::new();
329 let encryption = KmsAeadRingEnvelopeEncryption::new(provider).await.unwrap();
330
331 let aad = "test-aad";
332 let plaintext = SecretValue::from("secret message");
333
334 let (_dek, encrypted_dek) =
336 KmsAeadEnvelopeEncryption::<&str>::generate_new_dek(&encryption)
337 .await
338 .unwrap();
339
340 let ciphertext = encryption
342 .encrypt_value_with_encrypted_dek(&aad, &plaintext, &encrypted_dek)
343 .await
344 .unwrap();
345
346 let decrypted = encryption
348 .decrypt_value_with_encrypted_dek(&aad, &ciphertext, &encrypted_dek)
349 .await
350 .unwrap();
351
352 assert_eq!(decrypted, plaintext);
353 }
354
355 #[tokio::test]
356 async fn test_wrong_aad_fails() {
357 let provider = MockKmsProvider::new();
358 let encryption = KmsAeadRingEnvelopeEncryption::new(provider).await.unwrap();
359
360 let aad1 = "correct-aad";
361 let aad2 = "wrong-aad";
362 let plaintext = SecretValue::from("secret message");
363
364 let ciphertext = encryption.encrypt_value(&aad1, &plaintext).await.unwrap();
365 let result = encryption.decrypt_value(&aad2, &ciphertext).await;
366
367 assert!(result.is_err(), "Decryption with wrong AAD should fail");
368 }
369
370 #[tokio::test]
371 async fn test_empty_plaintext() {
372 let provider = MockKmsProvider::new();
373 let encryption = KmsAeadRingEnvelopeEncryption::new(provider).await.unwrap();
374
375 let aad = "test-aad";
376 let plaintext = SecretValue::from(vec![]);
377
378 let ciphertext = encryption.encrypt_value(&aad, &plaintext).await.unwrap();
379 let decrypted = encryption.decrypt_value(&aad, &ciphertext).await.unwrap();
380
381 assert_eq!(decrypted.ref_sensitive_value().len(), 0);
382 }
383
384 #[tokio::test]
385 async fn test_large_plaintext() {
386 let provider = MockKmsProvider::new();
387 let encryption = KmsAeadRingEnvelopeEncryption::new(provider).await.unwrap();
388
389 let aad = "test-aad";
390 let large_data = vec![0x42; 100_000]; let plaintext = SecretValue::from(large_data.clone());
392
393 let ciphertext = encryption.encrypt_value(&aad, &plaintext).await.unwrap();
394 let decrypted = encryption.decrypt_value(&aad, &ciphertext).await.unwrap();
395
396 assert_eq!(decrypted.ref_sensitive_value(), large_data.as_slice());
397 }
398
399 #[tokio::test]
400 async fn test_multiple_encryptions_different_deks() {
401 let provider = MockKmsProvider::new();
402 let encryption = KmsAeadRingEnvelopeEncryption::new(provider).await.unwrap();
403
404 let aad = "test-aad";
405 let plaintext = SecretValue::from("secret");
406
407 let ct1 = encryption.encrypt_value(&aad, &plaintext).await.unwrap();
408 let ct2 = encryption.encrypt_value(&aad, &plaintext).await.unwrap();
409
410 assert_ne!(ct1, ct2);
412
413 let decrypted1 = encryption.decrypt_value(&aad, &ct1).await.unwrap();
415 let decrypted2 = encryption.decrypt_value(&aad, &ct2).await.unwrap();
416
417 assert_eq!(decrypted1, plaintext);
418 assert_eq!(decrypted2, plaintext);
419 }
420
421 #[tokio::test]
422 async fn test_multiple_encryptions_same_dek() {
423 let provider = MockKmsProvider::new();
424 let encryption = KmsAeadRingEnvelopeEncryption::new(provider).await.unwrap();
425
426 let aad = "test-aad";
427 let plaintext = SecretValue::from("secret");
428
429 let (dek, _) = KmsAeadEnvelopeEncryption::<&str>::generate_new_dek(&encryption)
431 .await
432 .unwrap();
433
434 let ct1 = encryption
436 .encrypt_value_with_dek(&aad, &plaintext, &dek)
437 .await
438 .unwrap();
439 let ct2 = encryption
440 .encrypt_value_with_dek(&aad, &plaintext, &dek)
441 .await
442 .unwrap();
443
444 assert_ne!(ct1, ct2);
446
447 let decrypted1 = encryption
449 .decrypt_value_with_dek(&aad, &ct1, &dek)
450 .await
451 .unwrap();
452 let decrypted2 = encryption
453 .decrypt_value_with_dek(&aad, &ct2, &dek)
454 .await
455 .unwrap();
456
457 assert_eq!(decrypted1, plaintext);
458 assert_eq!(decrypted2, plaintext);
459 }
460
461 #[tokio::test]
462 async fn test_provider_encrypt_failure() {
463 let provider = MockKmsProvider::with_fail_encrypt();
464 let encryption = KmsAeadRingEnvelopeEncryption::new(provider).await.unwrap();
465
466 let aad = "test-aad";
467 let plaintext = SecretValue::from("secret");
468
469 let result = encryption.encrypt_value(&aad, &plaintext).await;
470 assert!(
471 result.is_err(),
472 "Should fail when provider fails to encrypt DEK"
473 );
474 }
475
476 #[tokio::test]
477 async fn test_provider_decrypt_failure() {
478 let provider = MockKmsProvider::new();
479 let encryption = KmsAeadRingEnvelopeEncryption::new(provider.clone())
480 .await
481 .unwrap();
482
483 let aad = "test-aad";
484 let plaintext = SecretValue::from("secret");
485
486 let ciphertext = encryption.encrypt_value(&aad, &plaintext).await.unwrap();
488
489 let failing_provider = MockKmsProvider::with_fail_decrypt();
491 let failing_encryption = KmsAeadRingEnvelopeEncryption::new(failing_provider)
492 .await
493 .unwrap();
494
495 let result = failing_encryption.decrypt_value(&aad, &ciphertext).await;
496 assert!(
497 result.is_err(),
498 "Should fail when provider fails to decrypt DEK"
499 );
500 }
501
502 #[tokio::test]
503 async fn test_aes_256_gcm_algorithm() {
504 let provider = MockKmsProvider::new();
505 let encryption =
506 KmsAeadRingEnvelopeEncryption::with_algorithm(provider, &ring::aead::AES_256_GCM)
507 .await
508 .unwrap();
509
510 let aad = "test-aad";
511 let plaintext = SecretValue::from("secret message");
512
513 let ciphertext = encryption.encrypt_value(&aad, &plaintext).await.unwrap();
514 let decrypted = encryption.decrypt_value(&aad, &ciphertext).await.unwrap();
515
516 assert_eq!(decrypted, plaintext);
517 }
518
519 #[tokio::test]
520 async fn test_with_options() {
521 let provider = MockKmsProvider::new();
522 let options = KmsAeadRingEnvelopeEncryptionOptions::new();
523 let encryption = KmsAeadRingEnvelopeEncryption::with_options(provider, options)
524 .await
525 .unwrap();
526
527 let aad = "test-aad";
528 let plaintext = SecretValue::from("secret");
529
530 let ciphertext = encryption.encrypt_value(&aad, &plaintext).await.unwrap();
531 let decrypted = encryption.decrypt_value(&aad, &ciphertext).await.unwrap();
532
533 assert_eq!(decrypted, plaintext);
534 }
535
536 #[tokio::test]
537 async fn test_binary_aad() {
538 let provider = MockKmsProvider::new();
539 let encryption = KmsAeadRingEnvelopeEncryption::new(provider).await.unwrap();
540
541 let binary_aad: Vec<u8> = vec![0x00, 0xFF, 0xDE, 0xAD, 0xBE, 0xEF];
542 let plaintext = SecretValue::from("secret");
543
544 let ciphertext = encryption
545 .encrypt_value(&binary_aad, &plaintext)
546 .await
547 .unwrap();
548 let decrypted = encryption
549 .decrypt_value(&binary_aad, &ciphertext)
550 .await
551 .unwrap();
552
553 assert_eq!(decrypted, plaintext);
554 }
555
556 #[tokio::test]
557 async fn test_concurrent_encryptions() {
558 use tokio::task;
559
560 let provider = MockKmsProvider::new();
561 let encryption = Arc::new(KmsAeadRingEnvelopeEncryption::new(provider).await.unwrap());
562
563 let mut handles = vec![];
564 for i in 0..10 {
565 let enc = encryption.clone();
566 let handle = task::spawn(async move {
567 let aad = format!("aad-{}", i);
568 let plaintext = SecretValue::from(format!("secret-{}", i).as_bytes().to_vec());
569
570 let ciphertext = enc.encrypt_value(&aad, &plaintext).await.unwrap();
571 let decrypted = enc.decrypt_value(&aad, &ciphertext).await.unwrap();
572
573 assert_eq!(decrypted, plaintext);
574 });
575 handles.push(handle);
576 }
577
578 for handle in handles {
579 handle.await.unwrap();
580 }
581 }
582
583 #[tokio::test]
584 async fn test_dek_generation() {
585 let provider = MockKmsProvider::new();
586 let encryption = KmsAeadRingEnvelopeEncryption::new(provider).await.unwrap();
587
588 let (dek1, encrypted_dek1) =
589 KmsAeadEnvelopeEncryption::<&str>::generate_new_dek(&encryption)
590 .await
591 .unwrap();
592 let (dek2, encrypted_dek2) =
593 KmsAeadEnvelopeEncryption::<&str>::generate_new_dek(&encryption)
594 .await
595 .unwrap();
596
597 assert_ne!(dek1, dek2);
599 assert_ne!(encrypted_dek1, encrypted_dek2);
600
601 assert_eq!(
603 dek1.value().ref_sensitive_value().len(),
604 ring::aead::CHACHA20_POLY1305.key_len()
605 );
606 }
607
608 #[tokio::test]
609 async fn test_corrupted_cipher_text_with_key_fails() {
610 let provider = MockKmsProvider::new();
611 let encryption = KmsAeadRingEnvelopeEncryption::new(provider).await.unwrap();
612
613 let aad = "test-aad";
614 let plaintext = SecretValue::from("secret");
615
616 let ciphertext = encryption.encrypt_value(&aad, &plaintext).await.unwrap();
617
618 let mut corrupted = ciphertext.value().to_vec();
620 if corrupted.len() > 20 {
621 corrupted[20] ^= 0x01;
622 }
623 let corrupted_ct = CipherTextWithEncryptedKey::from(corrupted);
624
625 let result = encryption.decrypt_value(&aad, &corrupted_ct).await;
626 assert!(result.is_err(), "Decryption of corrupted data should fail");
627 }
628}