asjeeves_encryption/
kek.rs1use 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 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 let mut key = [0u8; 32];
111
112 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}