Skip to main content

kms_aead/
kms_envelope_encryption.rs

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    // Mock provider for testing
199    #[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            // Simple mock: just clone the key and store it
246            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            // Simple mock: just return the key
263            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        // Generate a DEK
299        let (dek, encrypted_dek) = KmsAeadEnvelopeEncryption::<&str>::generate_new_dek(&encryption)
300            .await
301            .unwrap();
302
303        // Encrypt with the DEK
304        let ciphertext = encryption
305            .encrypt_value_with_dek(&aad, &plaintext, &dek)
306            .await
307            .unwrap();
308
309        // Decrypt with the same DEK
310        let decrypted = encryption
311            .decrypt_value_with_dek(&aad, &ciphertext, &dek)
312            .await
313            .unwrap();
314
315        assert_eq!(decrypted, plaintext);
316
317        // Also test decrypt with encrypted DEK
318        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        // Generate a DEK
335        let (_dek, encrypted_dek) =
336            KmsAeadEnvelopeEncryption::<&str>::generate_new_dek(&encryption)
337                .await
338                .unwrap();
339
340        // Encrypt with encrypted DEK
341        let ciphertext = encryption
342            .encrypt_value_with_encrypted_dek(&aad, &plaintext, &encrypted_dek)
343            .await
344            .unwrap();
345
346        // Decrypt with encrypted DEK
347        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]; // 100KB
391        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        // Different encryptions should produce different ciphertexts (different nonces/DEKs)
411        assert_ne!(ct1, ct2);
412
413        // But both should decrypt to the same plaintext
414        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        // Generate a single DEK
430        let (dek, _) = KmsAeadEnvelopeEncryption::<&str>::generate_new_dek(&encryption)
431            .await
432            .unwrap();
433
434        // Encrypt multiple times with same DEK
435        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        // Different ciphertexts (different nonces)
445        assert_ne!(ct1, ct2);
446
447        // Both decrypt correctly
448        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        // Encrypt successfully
487        let ciphertext = encryption.encrypt_value(&aad, &plaintext).await.unwrap();
488
489        // Create new encryption with failing provider
490        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        // DEKs should be different
598        assert_ne!(dek1, dek2);
599        assert_ne!(encrypted_dek1, encrypted_dek2);
600
601        // DEKs should have correct length for ChaCha20-Poly1305
602        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        // Corrupt the ciphertext
619        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}