Skip to main content

asjeeves_encryption/
kek.rs

1use core::fmt;
2use std::fmt::Debug;
3use std::sync::Arc;
4
5use aes_gcm::aead::{Aead, AeadCore, KeyInit};
6use aes_gcm::{Aes256Gcm, Nonce};
7use base64ct::{Base64, Encoding};
8use rand_core::{CryptoRng, RngCore};
9use serde::Deserialize;
10use tracing::instrument;
11use zeroize::Zeroizing;
12
13use crate::error::Error;
14use crate::web_key::WebKey;
15
16#[derive(Clone, Deserialize, Eq, PartialEq)]
17#[serde(try_from = "String")]
18pub struct KeyEncryptionKey(Arc<aes_gcm::Key<Aes256Gcm>>);
19
20pub struct EncryptedKey {
21    pub key: Box<[u8]>,
22    pub key_id: Box<str>,
23    pub nonce: Box<[u8]>,
24}
25
26impl KeyEncryptionKey {
27    pub fn generate<R>(rng: &mut R) -> Self
28    where
29        R: CryptoRng + RngCore,
30    {
31        let key = Aes256Gcm::generate_key(rng);
32
33        let key = Arc::new(key);
34
35        Self(key)
36    }
37
38    pub fn as_slice(&self) -> &[u8] {
39        self.0.as_slice()
40    }
41
42    #[instrument]
43    pub fn encrypt_key<R>(&self, key: &WebKey, rng: &mut R) -> Result<EncryptedKey, Error>
44    where
45        R: CryptoRng + RngCore + Debug,
46    {
47        let cipher = Aes256Gcm::new(&self.0);
48
49        let key_id: Box<str> = {
50            let key_id: &str = key.id();
51
52            key_id.into()
53        };
54
55        let nonce = Aes256Gcm::generate_nonce(rng);
56        let key_vec: Zeroizing<Vec<u8>> = key.to_bytes()?;
57
58        let key: Box<[u8]> = cipher
59            .encrypt(&nonce, key_vec.as_slice())
60            .map_err(|source| Error::AesEncryptionError { source })?
61            .as_slice()
62            .into();
63
64        Ok(EncryptedKey {
65            key_id,
66            key,
67            nonce: nonce.as_slice().into(),
68        })
69    }
70
71    #[instrument(skip(enc_key, nonce))]
72    pub fn decrypt_key(&self, enc_key: &[u8], key_id: &str, nonce: &[u8]) -> Result<WebKey, Error> {
73        let cipher = Aes256Gcm::new(&self.0);
74        let nonce = Nonce::from_slice(nonce);
75
76        // Use Zeroizing to ensure the bytes are zeroed out immediately.
77        let dec_key: Zeroizing<Vec<u8>> = cipher
78            .decrypt(nonce, enc_key)
79            .map_err(|source| Error::AesDecryptionError { source })?
80            .into();
81
82        let key_id: Arc<str> = key_id.into();
83
84        let web_key = WebKey::from_bytes(dec_key, key_id)?;
85
86        Ok(web_key)
87    }
88}
89
90impl fmt::Debug for KeyEncryptionKey {
91    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
92        f.debug_struct("KeyEncryptionKey").finish_non_exhaustive()
93    }
94}
95
96impl fmt::Display for KeyEncryptionKey {
97    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
98        let b64str: String = Base64::encode_string(&self.0);
99
100        write!(f, "{}", b64str)
101    }
102}
103
104impl TryFrom<String> for KeyEncryptionKey {
105    type Error = base64ct::Error;
106
107    #[instrument]
108    fn try_from(value: String) -> Result<Self, Self::Error> {
109        // 256-bit AES requires the key to be 32 bytes in length.
110        let mut key = [0u8; 32];
111
112        // This will throw an error if it does not decode exactly 32 bytes.
113        Base64::decode(value, &mut key)?;
114
115        let key: aes_gcm::Key<Aes256Gcm> = key.into();
116        let key = Arc::new(key);
117
118        Ok(Self(key))
119    }
120}
121
122#[cfg(test)]
123mod test {
124    use super::*;
125    use crate::seed::{Rng, Seed};
126
127    #[test]
128    fn it_should_encrypt_and_decrypt_keys() {
129        let seed = Seed::from(1);
130        let mut rng: Rng = seed.rng();
131
132        let web_key = WebKey::generate(&mut rng).unwrap();
133        let kek = KeyEncryptionKey::generate(&mut rng);
134
135        let payload = kek.encrypt_key(&web_key, &mut rng).unwrap();
136        let nonce = payload.nonce;
137        let enc_key = payload.key;
138
139        let dec_key = kek.decrypt_key(&enc_key, web_key.id(), &nonce).unwrap();
140
141        assert_eq!(dec_key, web_key)
142    }
143
144    #[test]
145    fn it_should_decode_and_encode_with_base64() {
146        let seed = Seed::from(1);
147
148        let mut rng: Rng = seed.rng();
149
150        let kek = KeyEncryptionKey::generate(&mut rng);
151
152        let kek_str: String = kek.to_string();
153
154        let kek_two = KeyEncryptionKey::try_from(kek_str).unwrap();
155
156        assert_eq!(kek, kek_two);
157        assert_eq!(kek.as_slice(), kek_two.as_slice());
158    }
159
160    #[test]
161    fn it_should_hide_details_when_debugged() {
162        let seed = Seed::from(1);
163
164        let mut rng: Rng = seed.rng();
165
166        let kek = KeyEncryptionKey::generate(&mut rng);
167
168        let dbg = format!("{:?}", kek);
169
170        assert_eq!("KeyEncryptionKey { .. }", dbg);
171    }
172}