iceberg/encryption/kms/
memory.rs1use 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::{invalid_data, 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(invalid_data!("Master key already exists: {key_id}"));
112 }
113
114 keys.insert(key_id, key);
115 Ok(())
116 }
117
118 fn get_master_key(&self, key_id: &str) -> Result<SensitiveBytes> {
119 let keys = self.master_keys.read().map_err(lock_error)?;
120
121 keys.get(key_id)
122 .cloned()
123 .ok_or_else(|| invalid_data!("Master key not found: {key_id}"))
124 }
125
126 pub fn key_count(&self) -> usize {
128 self.master_keys.read().map(|keys| keys.len()).unwrap_or(0)
129 }
130
131 pub fn has_key(&self, key_id: &str) -> bool {
133 self.master_keys
134 .read()
135 .map(|keys| keys.contains_key(key_id))
136 .unwrap_or(false)
137 }
138}
139
140#[derive(Debug, Clone, Default)]
149pub struct MemoryKmsClientFactory {
150 master_keys: Arc<RwLock<HashMap<String, SensitiveBytes>>>,
151 master_key_size: AesKeySize,
152}
153
154impl MemoryKmsClientFactory {
155 pub fn new() -> Self {
157 Self::default()
158 }
159
160 pub fn with_master_key_size(master_key_size: AesKeySize) -> Self {
162 Self {
163 master_keys: Arc::new(RwLock::new(HashMap::new())),
164 master_key_size,
165 }
166 }
167
168 fn client(&self) -> MemoryKeyManagementClient {
170 MemoryKeyManagementClient {
171 master_keys: Arc::clone(&self.master_keys),
172 master_key_size: self.master_key_size,
173 }
174 }
175
176 pub fn add_master_key(&self, key_id: impl Into<String>) -> Result<()> {
178 self.client().add_master_key(key_id)
179 }
180
181 pub fn add_master_key_bytes(
187 &self,
188 key_id: impl Into<String>,
189 key_bytes: SensitiveBytes,
190 ) -> Result<()> {
191 self.client().add_master_key_bytes(key_id, key_bytes)
192 }
193}
194
195#[async_trait]
196impl KmsClientFactory for MemoryKmsClientFactory {
197 async fn create_kms_client(
198 &self,
199 _properties: &HashMap<String, String>,
200 ) -> Result<Arc<dyn KeyManagementClient>> {
201 Ok(Arc::new(self.client()))
202 }
203}
204
205#[async_trait]
206impl KeyManagementClient for MemoryKeyManagementClient {
207 async fn wrap_key(&self, key: &[u8], wrapping_key_id: &str) -> Result<Vec<u8>> {
208 let master_key_bytes = self.get_master_key(wrapping_key_id)?;
209 let master_key = SecureKey::new(master_key_bytes.as_bytes())?;
210 let cipher = AesGcmCipher::new(master_key);
211
212 cipher.encrypt(key, None)
213 }
214
215 async fn unwrap_key(
216 &self,
217 wrapped_key: &[u8],
218 wrapping_key_id: &str,
219 ) -> Result<SensitiveBytes> {
220 let master_key_bytes = self.get_master_key(wrapping_key_id)?;
221 let master_key = SecureKey::new(master_key_bytes.as_bytes())?;
222 let cipher = AesGcmCipher::new(master_key);
223
224 Ok(SensitiveBytes::new(cipher.decrypt(wrapped_key, None)?))
225 }
226
227 fn supports_key_generation(&self) -> bool {
228 false
229 }
230
231 async fn generate_key(&self, _wrapping_key_id: &str) -> Result<super::GeneratedKey> {
232 Err(Error::new(
233 ErrorKind::FeatureUnsupported,
234 "MemoryKeyManagementClient does not support server-side key generation",
235 ))
236 }
237}
238
239#[cfg(test)]
240mod tests {
241 use super::*;
242
243 #[tokio::test]
244 async fn test_wrap_unwrap_roundtrip() {
245 let kms = MemoryKeyManagementClient::new();
246 kms.add_master_key("master-1").unwrap();
247 let dek = vec![0u8; 16];
248
249 let wrapped = kms.wrap_key(&dek, "master-1").await.unwrap();
250 let unwrapped = kms.unwrap_key(&wrapped, "master-1").await.unwrap();
251 assert_eq!(unwrapped.as_bytes(), dek.as_slice());
252 }
253
254 #[tokio::test]
255 async fn test_wrap_unknown_key_fails() {
256 let kms = MemoryKeyManagementClient::new();
257 let dek = vec![0u8; 16];
258
259 let result = kms.wrap_key(&dek, "nonexistent").await;
260 assert!(result.is_err());
261 }
262
263 #[tokio::test]
264 async fn test_wrong_master_key_fails_unwrap() {
265 let kms = MemoryKeyManagementClient::new();
266 kms.add_master_key("master-1").unwrap();
267 kms.add_master_key("master-2").unwrap();
268 let dek = vec![0u8; 16];
269
270 let wrapped = kms.wrap_key(&dek, "master-1").await.unwrap();
271
272 let result = kms.unwrap_key(&wrapped, "master-2").await;
273 assert!(result.is_err());
274 }
275
276 #[tokio::test]
277 async fn test_does_not_support_key_generation() {
278 let kms = MemoryKeyManagementClient::new();
279 assert!(!kms.supports_key_generation());
280
281 let result = kms.generate_key("master-1").await;
282 assert!(result.is_err());
283 }
284
285 #[tokio::test]
286 async fn test_multiple_master_keys() {
287 let kms = MemoryKeyManagementClient::new();
288 kms.add_master_key("master-1").unwrap();
289 kms.add_master_key("master-2").unwrap();
290 let dek1 = vec![1u8; 16];
291 let dek2 = vec![2u8; 16];
292
293 let wrapped1 = kms.wrap_key(&dek1, "master-1").await.unwrap();
294 let wrapped2 = kms.wrap_key(&dek2, "master-2").await.unwrap();
295
296 let unwrapped1 = kms.unwrap_key(&wrapped1, "master-1").await.unwrap();
297 let unwrapped2 = kms.unwrap_key(&wrapped2, "master-2").await.unwrap();
298
299 assert_eq!(unwrapped1.as_bytes(), dek1.as_slice());
300 assert_eq!(unwrapped2.as_bytes(), dek2.as_slice());
301 }
302
303 #[tokio::test]
304 async fn test_add_master_key() {
305 let kms = MemoryKeyManagementClient::new();
306
307 kms.add_master_key("my-key").unwrap();
308 assert!(kms.has_key("my-key"));
309 assert_eq!(kms.key_count(), 1);
310
311 let result = kms.add_master_key("my-key");
312 assert!(result.is_err());
313 }
314
315 #[tokio::test]
316 async fn test_add_master_key_bytes() {
317 let kms = MemoryKeyManagementClient::new();
318 let key_bytes = SensitiveBytes::new([42u8; 16]);
319
320 kms.add_master_key_bytes("my-key", key_bytes).unwrap();
321 assert!(kms.has_key("my-key"));
322
323 let dek = vec![7u8; 16];
324 let wrapped = kms.wrap_key(&dek, "my-key").await.unwrap();
325 let unwrapped = kms.unwrap_key(&wrapped, "my-key").await.unwrap();
326 assert_eq!(unwrapped.as_bytes(), dek.as_slice());
327 }
328
329 #[tokio::test]
330 async fn test_add_master_key_bytes_invalid_length() {
331 let kms = MemoryKeyManagementClient::new();
332
333 let result = kms.add_master_key_bytes("my-key", SensitiveBytes::new([0u8; 7]));
334 assert!(result.is_err());
335 }
336
337 #[tokio::test]
338 async fn test_with_master_key_size() {
339 let kms = MemoryKeyManagementClient::with_master_key_size(AesKeySize::Bits256);
340 kms.add_master_key("master-256").unwrap();
341
342 let dek = vec![0u8; 16];
343 let wrapped = kms.wrap_key(&dek, "master-256").await.unwrap();
344 let unwrapped = kms.unwrap_key(&wrapped, "master-256").await.unwrap();
345 assert_eq!(unwrapped.as_bytes(), dek.as_slice());
346 }
347
348 #[tokio::test]
349 async fn test_clone_shares_state() {
350 let kms1 = MemoryKeyManagementClient::new();
351 let kms2 = kms1.clone();
352
353 kms1.add_master_key("shared-key").unwrap();
354 assert!(kms2.has_key("shared-key"));
355 }
356
357 #[tokio::test]
358 async fn test_factory_seeds_produced_clients() {
359 let factory = MemoryKmsClientFactory::new();
362 factory
363 .add_master_key_bytes("master-1", SensitiveBytes::new([9u8; 16]))
364 .unwrap();
365
366 let client = factory.create_kms_client(&HashMap::new()).await.unwrap();
367
368 let dek = vec![3u8; 16];
369 let wrapped = client.wrap_key(&dek, "master-1").await.unwrap();
370 let unwrapped = client.unwrap_key(&wrapped, "master-1").await.unwrap();
371 assert_eq!(unwrapped.as_bytes(), dek.as_slice());
372 }
373
374 #[tokio::test]
375 async fn test_factory_seeding_after_client_creation_is_visible() {
376 let factory = MemoryKmsClientFactory::new();
379 let client = factory.create_kms_client(&HashMap::new()).await.unwrap();
380
381 factory.add_master_key("late-key").unwrap();
382
383 let dek = vec![1u8; 16];
384 assert!(client.wrap_key(&dek, "late-key").await.is_ok());
385 }
386
387 #[tokio::test]
388 async fn test_factory_with_master_key_size() {
389 let factory = MemoryKmsClientFactory::with_master_key_size(AesKeySize::Bits256);
390 factory.add_master_key("master-256").unwrap();
391
392 let client = factory.create_kms_client(&HashMap::new()).await.unwrap();
393 let dek = vec![0u8; 16];
394 let wrapped = client.wrap_key(&dek, "master-256").await.unwrap();
395 let unwrapped = client.unwrap_key(&wrapped, "master-256").await.unwrap();
396 assert_eq!(unwrapped.as_bytes(), dek.as_slice());
397 }
398}