Skip to main content

alien_bindings/providers/key/
aws.rs

1use super::{encode_context, frame, unframe};
2use crate::error::{ErrorData, Result};
3use crate::traits::{Binding, Key};
4use alien_aws_clients::kms::{DecryptRequest, EncryptRequest, KmsApi};
5use alien_error::{Context, IntoAlienError};
6use async_trait::async_trait;
7use base64::{engine::general_purpose::STANDARD, Engine as _};
8use std::collections::BTreeMap;
9use std::sync::Arc;
10
11#[derive(Debug)]
12pub struct AwsKmsKey {
13    client: Arc<dyn KmsApi>,
14    key_arn: String,
15}
16
17impl AwsKmsKey {
18    pub fn new(client: Arc<dyn KmsApi>, key_arn: String) -> Self {
19        Self { client, key_arn }
20    }
21}
22
23impl Binding for AwsKmsKey {}
24
25#[async_trait]
26impl Key for AwsKmsKey {
27    async fn encrypt(
28        &self,
29        plaintext: &[u8],
30        context: Option<&BTreeMap<String, String>>,
31    ) -> Result<Vec<u8>> {
32        let canonical = encode_context(context)?;
33        let response = self
34            .client
35            .encrypt(
36                EncryptRequest::builder(
37                    self.key_arn.clone(),
38                    STANDARD.encode(frame(plaintext, &canonical)?),
39                )
40                .maybe_encryption_context(context.cloned())
41                .build(),
42            )
43            .await
44            .context(ErrorData::CloudPlatformError {
45                message: "AWS KMS encrypt failed".to_string(),
46                resource_id: None,
47            })?;
48        STANDARD
49            .decode(response.ciphertext_blob)
50            .into_alien_error()
51            .context(ErrorData::CloudPlatformError {
52                message: "AWS KMS returned invalid ciphertext encoding".to_string(),
53                resource_id: None,
54            })
55    }
56
57    async fn decrypt(
58        &self,
59        ciphertext: &[u8],
60        context: Option<&BTreeMap<String, String>>,
61    ) -> Result<Vec<u8>> {
62        let canonical = encode_context(context)?;
63        let response = self
64            .client
65            .decrypt(
66                DecryptRequest::builder(self.key_arn.clone(), STANDARD.encode(ciphertext))
67                    .maybe_encryption_context(context.cloned())
68                    .build(),
69            )
70            .await
71            .context(ErrorData::CloudPlatformError {
72                message: "AWS KMS decrypt failed".to_string(),
73                resource_id: None,
74            })?;
75        let framed = STANDARD
76            .decode(response.plaintext)
77            .into_alien_error()
78            .context(ErrorData::KeyCiphertextInvalid {
79                reason: "provider plaintext encoding is invalid".to_string(),
80            })?;
81        unframe(&framed, &canonical)
82    }
83}
84
85#[cfg(test)]
86mod tests {
87    use super::*;
88    use alien_aws_clients::kms::{DecryptResponse, EncryptResponse, MockKmsApi};
89
90    #[tokio::test]
91    async fn passes_context_to_kms_and_validates_it_after_decrypt() {
92        let context = BTreeMap::from([("tenant".to_string(), "acme".to_string())]);
93        let canonical = encode_context(Some(&context)).unwrap();
94        let framed = frame(b"root", &canonical).unwrap();
95        let expected_context = context.clone();
96        let mut client = MockKmsApi::new();
97        client
98            .expect_encrypt()
99            .withf(move |request| {
100                request.key_id == "key-arn"
101                    && request.encryption_context.as_ref() == Some(&expected_context)
102            })
103            .returning(|_| {
104                Ok(EncryptResponse {
105                    ciphertext_blob: STANDARD.encode(b"ciphertext"),
106                    key_id: "key-arn".to_string(),
107                })
108            });
109        let decrypt_context = context.clone();
110        client
111            .expect_decrypt()
112            .withf(move |request| {
113                request.key_id == "key-arn"
114                    && request.ciphertext_blob == STANDARD.encode(b"ciphertext")
115                    && request.encryption_context.as_ref() == Some(&decrypt_context)
116            })
117            .returning(move |_| {
118                Ok(DecryptResponse {
119                    plaintext: STANDARD.encode(&framed),
120                    key_id: "key-arn".to_string(),
121                })
122            });
123
124        let key = AwsKmsKey::new(Arc::new(client), "key-arn".to_string());
125        let ciphertext = key.encrypt(b"root", Some(&context)).await.unwrap();
126        assert_eq!(
127            key.decrypt(&ciphertext, Some(&context)).await.unwrap(),
128            b"root"
129        );
130    }
131}