alien_bindings/providers/key/
aws.rs1use 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}