1use std::collections::HashMap;
24use std::fmt;
25use std::sync::{Arc, RwLock};
26
27use async_trait::async_trait;
28
29use super::KeyManagementClient;
30use super::factory::KmsClientFactory;
31use crate::encryption::{AesGcmCipher, AesKeySize, SecureKey, SensitiveBytes};
32use crate::error::lock_error;
33use crate::{Error, ErrorKind, Result};
34
35#[derive(Clone, Default)]
53pub struct MemoryKeyManagementClient {
54 master_keys: Arc<RwLock<HashMap<String, SensitiveBytes>>>,
55 master_key_size: AesKeySize,
56}
57
58impl fmt::Debug for MemoryKeyManagementClient {
59 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
60 f.debug_struct("MemoryKeyManagementClient")
61 .field("master_key_size", &self.master_key_size)
62 .field("key_count", &self.key_count())
63 .finish()
64 }
65}
66
67impl MemoryKeyManagementClient {
68 pub fn new() -> Self {
70 Self::default()
71 }
72
73 pub fn with_master_key_size(master_key_size: AesKeySize) -> Self {
75 Self {
76 master_keys: Arc::new(RwLock::new(HashMap::new())),
77 master_key_size,
78 }
79 }
80
81 pub fn add_master_key(&self, key_id: impl Into<String>) -> Result<()> {
83 let key = SecureKey::generate(self.master_key_size);
84 self.insert_key(key_id.into(), SensitiveBytes::new(key.as_bytes()))
85 }
86
87 pub fn add_master_key_bytes(
93 &self,
94 key_id: impl Into<String>,
95 key_bytes: SensitiveBytes,
96 ) -> Result<()> {
97 Self::check_key_length(&key_bytes)?;
98 self.insert_key(key_id.into(), key_bytes)
99 }
100
101 fn check_key_length(key_bytes: &SensitiveBytes) -> Result<()> {
103 SecureKey::new(key_bytes.as_bytes())?;
104 Ok(())
105 }
106
107 fn insert_key(&self, key_id: String, key: SensitiveBytes) -> Result<()> {
108 let mut keys = self.master_keys.write().map_err(lock_error)?;
109
110 if keys.contains_key(&key_id) {
111 return Err(Error::new(
112 ErrorKind::DataInvalid,
113 format!("Master key already exists: {key_id}"),
114 ));
115 }
116
117 keys.insert(key_id, key);
118 Ok(())
119 }
120
121 fn get_master_key(&self, key_id: &str) -> Result<SensitiveBytes> {
122 let keys = self.master_keys.read().map_err(lock_error)?;
123
124 keys.get(key_id).cloned().ok_or_else(|| {
125 Error::new(
126 ErrorKind::DataInvalid,
127 format!("Master key not found: {key_id}"),
128 )
129 })
130 }
131
132 pub fn key_count(&self) -> usize {
134 self.master_keys.read().map(|keys| keys.len()).unwrap_or(0)
135 }
136
137 pub fn has_key(&self, key_id: &str) -> bool {
139 self.master_keys
140 .read()
141 .map(|keys| keys.contains_key(key_id))
142 .unwrap_or(false)
143 }
144}
145
146#[derive(Debug, Clone, Default)]
155pub struct MemoryKmsClientFactory {
156 master_keys: Arc<RwLock<HashMap<String, SensitiveBytes>>>,
157 master_key_size: AesKeySize,
158}
159
160impl MemoryKmsClientFactory {
161 pub fn new() -> Self {
163 Self::default()
164 }
165
166 pub fn with_master_key_size(master_key_size: AesKeySize) -> Self {
168 Self {
169 master_keys: Arc::new(RwLock::new(HashMap::new())),
170 master_key_size,
171 }
172 }
173
174 fn client(&self) -> MemoryKeyManagementClient {
176 MemoryKeyManagementClient {
177 master_keys: Arc::clone(&self.master_keys),
178 master_key_size: self.master_key_size,
179 }
180 }
181
182 pub fn add_master_key(&self, key_id: impl Into<String>) -> Result<()> {
184 self.client().add_master_key(key_id)
185 }
186
187 pub fn add_master_key_bytes(
193 &self,
194 key_id: impl Into<String>,
195 key_bytes: SensitiveBytes,
196 ) -> Result<()> {
197 self.client().add_master_key_bytes(key_id, key_bytes)
198 }
199}
200
201#[async_trait]
202impl KmsClientFactory for MemoryKmsClientFactory {
203 async fn create_kms_client(
204 &self,
205 _properties: &HashMap<String, String>,
206 ) -> Result<Arc<dyn KeyManagementClient>> {
207 Ok(Arc::new(self.client()))
208 }
209}
210
211#[async_trait]
212impl KeyManagementClient for MemoryKeyManagementClient {
213 async fn wrap_key(&self, key: &[u8], wrapping_key_id: &str) -> Result<Vec<u8>> {
214 let master_key_bytes = self.get_master_key(wrapping_key_id)?;
215 let master_key = SecureKey::new(master_key_bytes.as_bytes())?;
216 let cipher = AesGcmCipher::new(master_key);
217
218 cipher.encrypt(key, None)
219 }
220
221 async fn unwrap_key(
222 &self,
223 wrapped_key: &[u8],
224 wrapping_key_id: &str,
225 ) -> Result<SensitiveBytes> {
226 let master_key_bytes = self.get_master_key(wrapping_key_id)?;
227 let master_key = SecureKey::new(master_key_bytes.as_bytes())?;
228 let cipher = AesGcmCipher::new(master_key);
229
230 Ok(SensitiveBytes::new(cipher.decrypt(wrapped_key, None)?))
231 }
232
233 fn supports_key_generation(&self) -> bool {
234 false
235 }
236
237 async fn generate_key(&self, _wrapping_key_id: &str) -> Result<super::GeneratedKey> {
238 Err(Error::new(
239 ErrorKind::FeatureUnsupported,
240 "MemoryKeyManagementClient does not support server-side key generation",
241 ))
242 }
243}
244
245#[cfg(test)]
246mod tests {
247 use super::*;
248
249 #[tokio::test]
250 async fn test_wrap_unwrap_roundtrip() {
251 let kms = MemoryKeyManagementClient::new();
252 kms.add_master_key("master-1").unwrap();
253 let dek = vec![0u8; 16];
254
255 let wrapped = kms.wrap_key(&dek, "master-1").await.unwrap();
256 let unwrapped = kms.unwrap_key(&wrapped, "master-1").await.unwrap();
257 assert_eq!(unwrapped.as_bytes(), dek.as_slice());
258 }
259
260 #[tokio::test]
261 async fn test_wrap_unknown_key_fails() {
262 let kms = MemoryKeyManagementClient::new();
263 let dek = vec![0u8; 16];
264
265 let result = kms.wrap_key(&dek, "nonexistent").await;
266 assert!(result.is_err());
267 }
268
269 #[tokio::test]
270 async fn test_wrong_master_key_fails_unwrap() {
271 let kms = MemoryKeyManagementClient::new();
272 kms.add_master_key("master-1").unwrap();
273 kms.add_master_key("master-2").unwrap();
274 let dek = vec![0u8; 16];
275
276 let wrapped = kms.wrap_key(&dek, "master-1").await.unwrap();
277
278 let result = kms.unwrap_key(&wrapped, "master-2").await;
279 assert!(result.is_err());
280 }
281
282 #[tokio::test]
283 async fn test_does_not_support_key_generation() {
284 let kms = MemoryKeyManagementClient::new();
285 assert!(!kms.supports_key_generation());
286
287 let result = kms.generate_key("master-1").await;
288 assert!(result.is_err());
289 }
290
291 #[tokio::test]
292 async fn test_multiple_master_keys() {
293 let kms = MemoryKeyManagementClient::new();
294 kms.add_master_key("master-1").unwrap();
295 kms.add_master_key("master-2").unwrap();
296 let dek1 = vec![1u8; 16];
297 let dek2 = vec![2u8; 16];
298
299 let wrapped1 = kms.wrap_key(&dek1, "master-1").await.unwrap();
300 let wrapped2 = kms.wrap_key(&dek2, "master-2").await.unwrap();
301
302 let unwrapped1 = kms.unwrap_key(&wrapped1, "master-1").await.unwrap();
303 let unwrapped2 = kms.unwrap_key(&wrapped2, "master-2").await.unwrap();
304
305 assert_eq!(unwrapped1.as_bytes(), dek1.as_slice());
306 assert_eq!(unwrapped2.as_bytes(), dek2.as_slice());
307 }
308
309 #[tokio::test]
310 async fn test_add_master_key() {
311 let kms = MemoryKeyManagementClient::new();
312
313 kms.add_master_key("my-key").unwrap();
314 assert!(kms.has_key("my-key"));
315 assert_eq!(kms.key_count(), 1);
316
317 let result = kms.add_master_key("my-key");
318 assert!(result.is_err());
319 }
320
321 #[tokio::test]
322 async fn test_add_master_key_bytes() {
323 let kms = MemoryKeyManagementClient::new();
324 let key_bytes = SensitiveBytes::new([42u8; 16]);
325
326 kms.add_master_key_bytes("my-key", key_bytes).unwrap();
327 assert!(kms.has_key("my-key"));
328
329 let dek = vec![7u8; 16];
330 let wrapped = kms.wrap_key(&dek, "my-key").await.unwrap();
331 let unwrapped = kms.unwrap_key(&wrapped, "my-key").await.unwrap();
332 assert_eq!(unwrapped.as_bytes(), dek.as_slice());
333 }
334
335 #[tokio::test]
336 async fn test_add_master_key_bytes_invalid_length() {
337 let kms = MemoryKeyManagementClient::new();
338
339 let result = kms.add_master_key_bytes("my-key", SensitiveBytes::new([0u8; 7]));
340 assert!(result.is_err());
341 }
342
343 #[tokio::test]
344 async fn test_with_master_key_size() {
345 let kms = MemoryKeyManagementClient::with_master_key_size(AesKeySize::Bits256);
346 kms.add_master_key("master-256").unwrap();
347
348 let dek = vec![0u8; 16];
349 let wrapped = kms.wrap_key(&dek, "master-256").await.unwrap();
350 let unwrapped = kms.unwrap_key(&wrapped, "master-256").await.unwrap();
351 assert_eq!(unwrapped.as_bytes(), dek.as_slice());
352 }
353
354 #[tokio::test]
355 async fn test_clone_shares_state() {
356 let kms1 = MemoryKeyManagementClient::new();
357 let kms2 = kms1.clone();
358
359 kms1.add_master_key("shared-key").unwrap();
360 assert!(kms2.has_key("shared-key"));
361 }
362
363 #[tokio::test]
364 async fn test_factory_seeds_produced_clients() {
365 let factory = MemoryKmsClientFactory::new();
368 factory
369 .add_master_key_bytes("master-1", SensitiveBytes::new([9u8; 16]))
370 .unwrap();
371
372 let client = factory.create_kms_client(&HashMap::new()).await.unwrap();
373
374 let dek = vec![3u8; 16];
375 let wrapped = client.wrap_key(&dek, "master-1").await.unwrap();
376 let unwrapped = client.unwrap_key(&wrapped, "master-1").await.unwrap();
377 assert_eq!(unwrapped.as_bytes(), dek.as_slice());
378 }
379
380 #[tokio::test]
381 async fn test_factory_seeding_after_client_creation_is_visible() {
382 let factory = MemoryKmsClientFactory::new();
385 let client = factory.create_kms_client(&HashMap::new()).await.unwrap();
386
387 factory.add_master_key("late-key").unwrap();
388
389 let dek = vec![1u8; 16];
390 assert!(client.wrap_key(&dek, "late-key").await.is_ok());
391 }
392
393 #[tokio::test]
394 async fn test_factory_with_master_key_size() {
395 let factory = MemoryKmsClientFactory::with_master_key_size(AesKeySize::Bits256);
396 factory.add_master_key("master-256").unwrap();
397
398 let client = factory.create_kms_client(&HashMap::new()).await.unwrap();
399 let dek = vec![0u8; 16];
400 let wrapped = client.wrap_key(&dek, "master-256").await.unwrap();
401 let unwrapped = client.unwrap_key(&wrapped, "master-256").await.unwrap();
402 assert_eq!(unwrapped.as_bytes(), dek.as_slice());
403 }
404}