Skip to main content

iceberg/encryption/kms/
memory.rs

1// Licensed to the Apache Software Foundation (ASF) under one
2// or more contributor license agreements.  See the NOTICE file
3// distributed with this work for additional information
4// regarding copyright ownership.  The ASF licenses this file
5// to you under the Apache License, Version 2.0 (the
6// "License"); you may not use this file except in compliance
7// with the License.  You may obtain a copy of the License at
8//
9//   http://www.apache.org/licenses/LICENSE-2.0
10//
11// Unless required by applicable law or agreed to in writing,
12// software distributed under the License is distributed on an
13// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
14// KIND, either express or implied.  See the License for the
15// specific language governing permissions and limitations
16// under the License.
17
18//! In-memory KMS implementation for testing and development.
19//!
20//! **WARNING**: This implementation is NOT suitable for production use.
21//! Keys are stored in memory only and will be lost when the process exits.
22
23use 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/// In-memory KMS for testing. Not suitable for production use.
36///
37/// ```
38/// use iceberg::encryption::KeyManagementClient;
39/// use iceberg::encryption::kms::MemoryKeyManagementClient;
40///
41/// # async fn example() -> iceberg::Result<()> {
42/// let kms = MemoryKeyManagementClient::new();
43/// kms.add_master_key("my-master-key")?;
44///
45/// let dek = vec![0u8; 16];
46/// let wrapped = kms.wrap_key(&dek, "my-master-key").await?;
47/// let unwrapped = kms.unwrap_key(&wrapped, "my-master-key").await?;
48/// assert_eq!(dek.as_slice(), unwrapped.as_bytes());
49/// # Ok(())
50/// # }
51/// ```
52#[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    /// Creates a new in-memory KMS with 128-bit AES keys.
69    pub fn new() -> Self {
70        Self::default()
71    }
72
73    /// Creates a new in-memory KMS with the specified master key size.
74    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    /// Adds a randomly generated master key with the given ID.
82    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    /// Adds a master key with explicit key bytes.
88    ///
89    /// Use this to seed the KMS with known key material, e.g. for
90    /// cross-language integration tests where both Java and Rust must
91    /// share the same master key bytes.
92    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    /// Check the key length is valid by constructing a SecureKey.
102    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    /// Number of master keys stored.
133    pub fn key_count(&self) -> usize {
134        self.master_keys.read().map(|keys| keys.len()).unwrap_or(0)
135    }
136
137    /// Whether a master key with the given ID exists.
138    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/// Factory for creating [`MemoryKeyManagementClient`] instances.
147///
148/// The factory owns the master-key table and its key size; seed it with
149/// [`add_master_key`](Self::add_master_key) /
150/// [`add_master_key_bytes`](Self::add_master_key_bytes), then every client it
151/// produces via [`create_kms_client`](KmsClientFactory::create_kms_client)
152/// shares that same table. Useful for testing encryption flows without a real
153/// KMS backend.
154#[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    /// Creates a new factory with 128-bit AES master keys.
162    pub fn new() -> Self {
163        Self::default()
164    }
165
166    /// Creates a new factory whose clients use the given master key size.
167    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    /// A client view over this factory's shared master-key table.
175    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    /// Adds a randomly generated master key with the given ID.
183    pub fn add_master_key(&self, key_id: impl Into<String>) -> Result<()> {
184        self.client().add_master_key(key_id)
185    }
186
187    /// Adds a master key with explicit key bytes.
188    ///
189    /// Use this to seed the factory with known key material, e.g. for
190    /// cross-language integration tests where both Java and Rust must
191    /// share the same master key bytes.
192    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        // Seed the factory, then every client it produces sees those keys and
366        // can wrap/unwrap with them.
367        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        // The factory owns the shared table, so keys added after a client is
383        // produced are still visible to that client.
384        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}