Skip to main content

alien_bindings/providers/key/
azure.rs

1use super::{encode_context, frame, unframe};
2use crate::error::{ErrorData, Result};
3use crate::traits::{Binding, Key};
4use alien_azure_clients::keyvault::{KeyOperationRequest, KeyVaultKeysApi};
5use alien_error::{Context, IntoAlienError};
6use async_trait::async_trait;
7use base64::{engine::general_purpose::URL_SAFE_NO_PAD, Engine as _};
8use std::collections::BTreeMap;
9use std::sync::Arc;
10
11#[derive(Debug)]
12pub struct AzureKeyVaultKey {
13    client: Arc<dyn KeyVaultKeysApi>,
14    key_id: String,
15}
16
17impl AzureKeyVaultKey {
18    pub fn new(client: Arc<dyn KeyVaultKeysApi>, key_id: String) -> Self {
19        Self { client, key_id }
20    }
21}
22
23impl Binding for AzureKeyVaultKey {}
24
25#[async_trait]
26impl Key for AzureKeyVaultKey {
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                &self.key_id,
37                KeyOperationRequest {
38                    alg: "RSA-OAEP-256".to_string(),
39                    value: URL_SAFE_NO_PAD.encode(frame(plaintext, &canonical)?),
40                },
41            )
42            .await
43            .context(ErrorData::CloudPlatformError {
44                message: "Azure Key Vault encrypt failed".to_string(),
45                resource_id: None,
46            })?;
47        URL_SAFE_NO_PAD
48            .decode(response.value)
49            .into_alien_error()
50            .context(ErrorData::CloudPlatformError {
51                message: "Azure Key Vault returned invalid ciphertext encoding".to_string(),
52                resource_id: None,
53            })
54    }
55
56    async fn decrypt(
57        &self,
58        ciphertext: &[u8],
59        context: Option<&BTreeMap<String, String>>,
60    ) -> Result<Vec<u8>> {
61        let canonical = encode_context(context)?;
62        let response = self
63            .client
64            .decrypt(
65                &self.key_id,
66                KeyOperationRequest {
67                    alg: "RSA-OAEP-256".to_string(),
68                    value: URL_SAFE_NO_PAD.encode(ciphertext),
69                },
70            )
71            .await
72            .context(ErrorData::CloudPlatformError {
73                message: "Azure Key Vault decrypt failed".to_string(),
74                resource_id: None,
75            })?;
76        let framed = URL_SAFE_NO_PAD
77            .decode(response.value)
78            .into_alien_error()
79            .context(ErrorData::KeyCiphertextInvalid {
80                reason: "provider plaintext encoding is invalid".to_string(),
81            })?;
82        unframe(&framed, &canonical)
83    }
84}
85
86#[cfg(test)]
87mod tests {
88    use super::*;
89    use alien_azure_clients::keyvault::{KeyOperationResponse, MockKeyVaultKeysApi};
90
91    #[tokio::test]
92    async fn binds_context_inside_the_portable_frame() {
93        let context = BTreeMap::from([("tenant".to_string(), "acme".to_string())]);
94        let canonical = encode_context(Some(&context)).unwrap();
95        let framed = frame(b"root", &canonical).unwrap();
96        let mut client = MockKeyVaultKeysApi::new();
97        client.expect_encrypt().returning(|key_id, request| {
98            assert_eq!(key_id, "versioned-key-id");
99            assert_eq!(request.alg, "RSA-OAEP-256");
100            Ok(KeyOperationResponse {
101                kid: key_id.to_string(),
102                value: URL_SAFE_NO_PAD.encode(b"ciphertext"),
103            })
104        });
105        client
106            .expect_decrypt()
107            .times(2)
108            .returning(move |key_id, request| {
109                assert_eq!(key_id, "versioned-key-id");
110                assert_eq!(request.alg, "RSA-OAEP-256");
111                Ok(KeyOperationResponse {
112                    kid: key_id.to_string(),
113                    value: URL_SAFE_NO_PAD.encode(&framed),
114                })
115            });
116
117        let key = AzureKeyVaultKey::new(Arc::new(client), "versioned-key-id".to_string());
118        let ciphertext = key.encrypt(b"root", Some(&context)).await.unwrap();
119        assert_eq!(
120            key.decrypt(&ciphertext, Some(&context)).await.unwrap(),
121            b"root"
122        );
123        assert!(key.decrypt(&ciphertext, None).await.is_err());
124    }
125}