Skip to main content

alien_bindings/providers/key/
gcp.rs

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