1use std::collections::HashMap;
27use std::fmt;
28use std::sync::{Arc, RwLock};
29use std::time::Duration;
30
31use aes_gcm::aead::OsRng;
32use aes_gcm::aead::rand_core::RngCore;
33use chrono::Utc;
34use moka::future::Cache;
35use uuid::Uuid;
36
37const MILLIS_IN_DAY: i64 = 24 * 60 * 60 * 1000;
38
39use super::crypto::{AesGcmCipher, AesKeySize, SecureKey};
40use super::io::EncryptedOutputFile;
41use super::key_metadata::StandardKeyMetadata;
42use super::kms::KeyManagementClient;
43use crate::io::OutputFile;
44use crate::sensitive::SensitiveBytes;
45use crate::spec::{EncryptedKey, FormatVersion, TableMetadataRef};
46use crate::{Error, ErrorKind, Result};
47
48pub const KEK_CREATED_AT_PROPERTY: &str = "KEY_TIMESTAMP";
51
52const DEFAULT_KEK_LIFESPAN_DAYS: i64 = 730;
54
55const DEFAULT_CACHE_TTL: Duration = Duration::from_secs(3600);
57
58const AAD_PREFIX_LENGTH: usize = 16;
61
62#[derive(typed_builder::TypedBuilder)]
66#[builder(mutators(
67 pub fn add_encryption_key(&mut self, key: EncryptedKey) {
69 self.encryption_keys
70 .write()
71 .expect("encryption_keys lock poisoned")
72 .insert(key.key_id().to_string(), key);
73 }
74 pub fn encryption_keys(&mut self, keys: HashMap<String, EncryptedKey>) {
76 self.encryption_keys = RwLock::new(keys);
77 }
78))]
79pub struct EncryptionManager {
80 kms_client: Arc<dyn KeyManagementClient>,
81 #[builder(
82 default = Cache::builder().time_to_live(DEFAULT_CACHE_TTL).build(),
83 setter(skip)
84 )]
85 kek_cache: Cache<String, SensitiveBytes>,
86 #[builder(default = AesKeySize::default())]
88 key_size: AesKeySize,
89 #[builder(setter(into))]
91 table_key_id: String,
92 #[builder(default = RwLock::new(HashMap::new()), via_mutators)]
96 encryption_keys: RwLock<HashMap<String, EncryptedKey>>,
97}
98
99impl fmt::Debug for EncryptionManager {
100 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
101 f.debug_struct("EncryptionManager")
102 .field("key_size", &self.key_size)
103 .field("table_key_id", &self.table_key_id)
104 .finish_non_exhaustive()
105 }
106}
107
108impl EncryptionManager {
109 pub(crate) fn from_table_metadata(
115 kms_client: Option<&Arc<dyn KeyManagementClient>>,
116 metadata: &TableMetadataRef,
117 ) -> Result<Option<Arc<Self>>> {
118 if metadata.format_version() < FormatVersion::V3 {
119 return Ok(None);
120 }
121
122 let table_properties = metadata.table_properties()?;
123 let Some(table_key_id) = table_properties.encryption_key_id else {
124 if kms_client.is_some() {
125 tracing::warn!(
126 "KeyManagementClient provided but table does not have encryption.key-id set"
127 );
128 }
129 return Ok(None);
130 };
131
132 let kms_client = kms_client.ok_or_else(|| {
133 Error::new(
134 ErrorKind::PreconditionFailed,
135 "Table has encryption.key-id set but no KeyManagementClient was provided to TableBuilder",
136 )
137 })?;
138
139 let em = EncryptionManager::builder()
140 .kms_client(Arc::clone(kms_client))
141 .table_key_id(table_key_id)
142 .encryption_keys(metadata.encryption_keys.clone())
143 .key_size(AesKeySize::from_key_length(
144 table_properties.encryption_data_key_length,
145 )?)
146 .build();
147 Ok(Some(Arc::new(em)))
148 }
149
150 pub fn encrypt(&self, raw_output: OutputFile) -> EncryptedOutputFile {
155 let dek = SecureKey::generate(self.key_size);
156 let aad_prefix = Self::generate_aad_prefix();
157 let metadata = StandardKeyMetadata::from(dek).with_aad_prefix(&aad_prefix);
158 EncryptedOutputFile::new(raw_output, metadata)
159 }
160
161 pub async fn encrypt_manifest_list_key_metadata(
170 &self,
171 key_metadata: &StandardKeyMetadata,
172 ) -> Result<String> {
173 let kek = match self.find_active_kek()? {
174 Some(existing) => existing,
175 None => self.create_kek().await?,
176 };
177
178 let kek_bytes = self.unwrap_key_encryption_key(&kek).await?;
179
180 let aad = Self::kek_timestamp_aad(&kek)?;
182 let serialized = key_metadata.encode()?;
183 let wrapped_metadata = self.wrap_dek_with_kek(&serialized, &kek_bytes, Some(aad))?;
184
185 let wrapped_key = EncryptedKey::builder()
186 .key_id(Uuid::new_v4().to_string())
187 .encrypted_key_metadata(wrapped_metadata)
188 .encrypted_by_id(kek.key_id())
189 .build();
190
191 let wrapped_key_id = wrapped_key.key_id().to_string();
192 self.insert_encryption_key(wrapped_key);
193 Ok(wrapped_key_id)
194 }
195
196 pub async fn decrypt_manifest_list_key_metadata(
202 &self,
203 encryption_key_id: &str,
204 ) -> Result<StandardKeyMetadata> {
205 let encrypted_key = self
206 .encryption_keys
207 .read()
208 .expect("encryption_keys lock poisoned")
209 .get(encryption_key_id)
210 .cloned()
211 .ok_or_else(|| {
212 Error::new(
213 ErrorKind::DataInvalid,
214 format!("Encryption key '{encryption_key_id}' not found"),
215 )
216 })?;
217
218 let kek_key_id = encrypted_key.encrypted_by_id().ok_or_else(|| {
219 Error::new(
220 ErrorKind::DataInvalid,
221 format!(
222 "EncryptedKey '{}' has no encrypted_by_id",
223 encrypted_key.key_id()
224 ),
225 )
226 })?;
227
228 let bytes = self
229 .decrypt_dek(kek_key_id, encrypted_key.encrypted_key_metadata())
230 .await?;
231
232 StandardKeyMetadata::decode(bytes.as_bytes())
233 }
234
235 pub fn with_encryption_keys<F, R>(&self, f: F) -> R
240 where F: FnOnce(&HashMap<String, EncryptedKey>) -> R {
241 let keys = self
242 .encryption_keys
243 .read()
244 .expect("encryption_keys lock poisoned");
245 f(&keys)
246 }
247
248 fn insert_encryption_key(&self, key: EncryptedKey) {
249 self.encryption_keys
250 .write()
251 .expect("encryption_keys lock poisoned")
252 .insert(key.key_id().to_string(), key);
253 }
254
255 async fn create_kek(&self) -> Result<EncryptedKey> {
258 let (plaintext_kek, wrapped_kek) = if self.kms_client.supports_key_generation() {
259 let result = self.kms_client.generate_key(&self.table_key_id).await?;
260 (result.key().clone(), result.wrapped_key().to_vec())
261 } else {
262 let plaintext_key = SecureKey::generate(self.key_size);
263 let wrapped = self
264 .kms_client
265 .wrap_key(plaintext_key.as_bytes(), &self.table_key_id)
266 .await?;
267
268 (SensitiveBytes::new(plaintext_key.as_bytes()), wrapped)
269 };
270
271 let key_id = Uuid::new_v4().to_string();
272 let now_ms = Utc::now().timestamp_millis();
273
274 let mut properties = HashMap::new();
275 properties.insert(KEK_CREATED_AT_PROPERTY.to_string(), now_ms.to_string());
276
277 self.kek_cache.insert(key_id.clone(), plaintext_kek).await;
278
279 let kek = EncryptedKey::builder()
280 .key_id(key_id)
281 .encrypted_key_metadata(wrapped_kek)
282 .encrypted_by_id(&self.table_key_id)
283 .properties(properties)
284 .build();
285
286 self.insert_encryption_key(kek.clone());
287 Ok(kek)
288 }
289
290 fn is_kek_expired(&self, kek: &EncryptedKey) -> bool {
292 let created_at_ms = match kek
293 .properties()
294 .get(KEK_CREATED_AT_PROPERTY)
295 .and_then(|ts| ts.parse::<i64>().ok())
296 {
297 Some(ts) => ts,
298 None => return true, };
300
301 let now_ms = Utc::now().timestamp_millis();
302 let lifespan_ms = DEFAULT_KEK_LIFESPAN_DAYS * MILLIS_IN_DAY;
303 (now_ms - created_at_ms) >= lifespan_ms
304 }
305
306 fn find_active_kek(&self) -> Result<Option<EncryptedKey>> {
308 let keys = self
309 .encryption_keys
310 .read()
311 .expect("encryption_keys lock poisoned");
312 Ok(keys
313 .values()
314 .filter(|kek| {
315 kek.encrypted_by_id()
316 .map(|id| id == self.table_key_id)
317 .unwrap_or(false)
318 && !self.is_kek_expired(kek)
319 })
320 .max_by_key(|kek| {
321 kek.properties()
322 .get(KEK_CREATED_AT_PROPERTY)
323 .and_then(|ts| ts.parse::<i64>().ok())
324 .unwrap_or(0)
325 })
326 .cloned())
327 }
328
329 async fn unwrap_key_encryption_key(&self, kek: &EncryptedKey) -> Result<SensitiveBytes> {
331 let cache_key = kek.key_id().to_string();
332
333 if let Some(cached) = self.kek_cache.get(&cache_key).await {
334 return Ok(cached);
335 }
336
337 let master_key_id = kek.encrypted_by_id().ok_or_else(|| {
338 Error::new(
339 ErrorKind::DataInvalid,
340 format!("KEK '{}' has no encrypted_by_id", kek.key_id()),
341 )
342 })?;
343
344 let plaintext = self
345 .kms_client
346 .unwrap_key(kek.encrypted_key_metadata(), master_key_id)
347 .await?;
348
349 self.kek_cache.insert(cache_key, plaintext.clone()).await;
350
351 Ok(plaintext)
352 }
353
354 async fn decrypt_dek(&self, kek_key_id: &str, wrapped_dek: &[u8]) -> Result<SensitiveBytes> {
357 let kek = self
358 .encryption_keys
359 .read()
360 .expect("encryption_keys lock poisoned")
361 .get(kek_key_id)
362 .cloned()
363 .ok_or_else(|| {
364 Error::new(
365 ErrorKind::DataInvalid,
366 format!("KEK not found in encryption keys: {kek_key_id}"),
367 )
368 })?;
369
370 let aad = Self::kek_timestamp_aad(&kek)?;
372
373 let kek_bytes = self.unwrap_key_encryption_key(&kek).await?;
374 self.unwrap_dek_with_kek(wrapped_dek, &kek_bytes, Some(aad))
375 .map_err(|e| {
376 Error::new(
377 e.kind(),
378 format!("Failed to unwrap key metadata with KEK '{kek_key_id}'"),
379 )
380 .with_source(e)
381 })
382 }
383
384 fn kek_timestamp_aad(kek: &EncryptedKey) -> Result<&[u8]> {
386 kek.properties()
387 .get(KEK_CREATED_AT_PROPERTY)
388 .map(|ts| ts.as_bytes())
389 .ok_or_else(|| {
390 Error::new(
391 ErrorKind::DataInvalid,
392 format!(
393 "KEK '{}' is missing required '{}' property",
394 kek.key_id(),
395 KEK_CREATED_AT_PROPERTY
396 ),
397 )
398 })
399 }
400
401 fn generate_aad_prefix() -> Box<[u8]> {
403 let mut prefix = vec![0u8; AAD_PREFIX_LENGTH];
404 OsRng.fill_bytes(&mut prefix);
405 prefix.into_boxed_slice()
406 }
407
408 fn wrap_dek_with_kek(
410 &self,
411 dek: &[u8],
412 kek: &SensitiveBytes,
413 aad: Option<&[u8]>,
414 ) -> Result<Vec<u8>> {
415 let key = SecureKey::try_from(kek.clone())?;
416 let cipher = AesGcmCipher::new(key);
417 cipher.encrypt(dek, aad)
418 }
419
420 fn unwrap_dek_with_kek(
422 &self,
423 wrapped_dek: &[u8],
424 kek: &SensitiveBytes,
425 aad: Option<&[u8]>,
426 ) -> Result<SensitiveBytes> {
427 let key = SecureKey::try_from(kek.clone())?;
428 let cipher = AesGcmCipher::new(key);
429 cipher.decrypt(wrapped_dek, aad).map(SensitiveBytes::new)
430 }
431}
432
433#[cfg(test)]
434mod tests {
435 use super::*;
436 use crate::encryption::EncryptedInputFile;
437 use crate::encryption::kms::MemoryKeyManagementClient;
438
439 fn create_test_kms() -> Arc<dyn KeyManagementClient> {
440 let kms = MemoryKeyManagementClient::new();
441 kms.add_master_key("master-1").unwrap();
442 Arc::new(kms)
443 }
444
445 fn create_test_manager() -> EncryptionManager {
446 EncryptionManager::builder()
447 .kms_client(create_test_kms())
448 .table_key_id("master-1")
449 .build()
450 }
451
452 #[tokio::test]
453 async fn test_create_kek() {
454 let mgr = create_test_manager();
455 let kek = mgr.create_kek().await.unwrap();
456
457 assert!(!kek.key_id().is_empty());
458 assert!(!kek.encrypted_key_metadata().is_empty());
459 assert_eq!(kek.encrypted_by_id(), Some("master-1"));
460 assert!(kek.properties().contains_key(KEK_CREATED_AT_PROPERTY));
461 }
462
463 fn sample_key_metadata() -> StandardKeyMetadata {
464 StandardKeyMetadata::try_new(b"0123456789abcdef")
465 .unwrap()
466 .with_aad_prefix(b"test-aad-prefix!")
467 }
468
469 #[tokio::test]
470 async fn test_wrap_unwrap_key_metadata_roundtrip() {
471 let mgr = create_test_manager();
472 let plaintext = sample_key_metadata();
473
474 let key_id = mgr
475 .encrypt_manifest_list_key_metadata(&plaintext)
476 .await
477 .unwrap();
478
479 assert_eq!(mgr.with_encryption_keys(|k| k.len()), 2);
481
482 let decrypted = mgr
483 .decrypt_manifest_list_key_metadata(&key_id)
484 .await
485 .unwrap();
486 assert_eq!(decrypted, plaintext);
487 }
488
489 #[tokio::test]
490 async fn test_kek_reuse_when_not_expired() {
491 let mgr = create_test_manager();
492
493 let _id1 = mgr
495 .encrypt_manifest_list_key_metadata(&sample_key_metadata())
496 .await
497 .unwrap();
498 let kek_id = mgr.with_encryption_keys(|keys| {
499 assert_eq!(keys.len(), 2);
500 keys.values()
501 .find(|k| k.encrypted_by_id() == Some("master-1"))
502 .unwrap()
503 .key_id()
504 .to_string()
505 });
506
507 let id2 = mgr
509 .encrypt_manifest_list_key_metadata(&sample_key_metadata())
510 .await
511 .unwrap();
512 let entry2 = mgr.with_encryption_keys(|keys| {
513 assert_eq!(keys.len(), 3);
514 keys.get(&id2).cloned().unwrap()
515 });
516 assert_eq!(entry2.encrypted_by_id(), Some(kek_id.as_str()));
517 }
518
519 #[tokio::test]
520 async fn test_kek_rotation_when_expired() {
521 let kms = create_test_kms();
522
523 let three_years_ago_ms = Utc::now().timestamp_millis() - (3 * 365 * MILLIS_IN_DAY);
525 let mut properties = HashMap::new();
526 properties.insert(
527 KEK_CREATED_AT_PROPERTY.to_string(),
528 three_years_ago_ms.to_string(),
529 );
530
531 let kek_key = SecureKey::generate(AesKeySize::Bits128);
533 let wrapped = kms.wrap_key(kek_key.as_bytes(), "master-1").await.unwrap();
534
535 let old_kek = EncryptedKey::builder()
536 .key_id("expired-kek")
537 .encrypted_key_metadata(wrapped)
538 .encrypted_by_id("master-1")
539 .properties(properties)
540 .build();
541
542 let mgr = EncryptionManager::builder()
544 .kms_client(kms)
545 .table_key_id("master-1")
546 .add_encryption_key(old_kek.clone())
547 .build();
548
549 let new_entry_id = mgr
551 .encrypt_manifest_list_key_metadata(&sample_key_metadata())
552 .await
553 .unwrap();
554 let entry = mgr
555 .with_encryption_keys(|keys| keys.get(&new_entry_id).cloned())
556 .unwrap();
557 let used_kek_id = entry.encrypted_by_id().unwrap();
558 assert_ne!(used_kek_id, old_kek.key_id());
559 }
560
561 #[tokio::test]
562 async fn test_is_kek_expired_no_timestamp() {
563 let mgr = create_test_manager();
564
565 let kek = EncryptedKey::builder()
567 .key_id("no-ts")
568 .encrypted_key_metadata(vec![0u8; 32])
569 .build();
570
571 assert!(mgr.is_kek_expired(&kek));
572 }
573
574 #[tokio::test]
575 async fn test_decrypt_with_unknown_key_id() {
576 let mgr = create_test_manager();
577 let result = mgr.decrypt_manifest_list_key_metadata("nonexistent").await;
578 assert!(result.is_err());
579 }
580
581 #[tokio::test]
582 async fn test_kek_cache_hit() {
583 let mgr = create_test_manager();
584
585 let key_id = mgr
587 .encrypt_manifest_list_key_metadata(&sample_key_metadata())
588 .await
589 .unwrap();
590
591 let _ = mgr
593 .decrypt_manifest_list_key_metadata(&key_id)
594 .await
595 .unwrap();
596 }
597
598 #[tokio::test]
599 async fn test_unwrap_fails_when_kek_missing_timestamp() {
600 let mgr = create_test_manager();
601
602 let entry_id = mgr
604 .encrypt_manifest_list_key_metadata(&sample_key_metadata())
605 .await
606 .unwrap();
607
608 let mut keys = mgr.with_encryption_keys(|k| k.clone());
611 let kek_id = keys
612 .get(&entry_id)
613 .unwrap()
614 .encrypted_by_id()
615 .unwrap()
616 .to_string();
617 let kek = keys.remove(&kek_id).unwrap();
618 let kek_no_ts = EncryptedKey::builder()
619 .key_id(kek.key_id())
620 .encrypted_key_metadata(kek.encrypted_key_metadata())
621 .encrypted_by_id(kek.encrypted_by_id().unwrap())
622 .build();
623 keys.insert(kek_no_ts.key_id().to_string(), kek_no_ts);
624
625 let mgr = EncryptionManager::builder()
626 .kms_client(create_test_kms())
627 .table_key_id("master-1")
628 .encryption_keys(keys)
629 .build();
630
631 let result = mgr.decrypt_manifest_list_key_metadata(&entry_id).await;
632 assert!(result.is_err());
633 let err = result.unwrap_err();
634 assert_eq!(err.kind(), ErrorKind::DataInvalid);
635 assert!(
636 err.to_string().contains(KEK_CREATED_AT_PROPERTY),
637 "error should mention the missing property: {err}"
638 );
639 }
640
641 #[tokio::test]
642 async fn test_unwrap_fails_when_kek_timestamp_tampered() {
643 let mgr = create_test_manager();
644
645 let entry_id = mgr
647 .encrypt_manifest_list_key_metadata(&sample_key_metadata())
648 .await
649 .unwrap();
650
651 let mut keys = mgr.with_encryption_keys(|k| k.clone());
653 let kek_id = keys
654 .get(&entry_id)
655 .unwrap()
656 .encrypted_by_id()
657 .unwrap()
658 .to_string();
659 let kek = keys.remove(&kek_id).unwrap();
660 let mut tampered_properties = kek.properties().clone();
661 tampered_properties.insert(KEK_CREATED_AT_PROPERTY.to_string(), "9999999".to_string());
662 let tampered_kek = EncryptedKey::builder()
663 .key_id(kek.key_id())
664 .encrypted_key_metadata(kek.encrypted_key_metadata())
665 .encrypted_by_id(kek.encrypted_by_id().unwrap())
666 .properties(tampered_properties)
667 .build();
668 keys.insert(tampered_kek.key_id().to_string(), tampered_kek);
669
670 let mgr = EncryptionManager::builder()
671 .kms_client(create_test_kms())
672 .table_key_id("master-1")
673 .encryption_keys(keys)
674 .build();
675
676 let result = mgr.decrypt_manifest_list_key_metadata(&entry_id).await;
678 assert!(
679 result.is_err(),
680 "tampered timestamp should cause decryption failure"
681 );
682 }
683
684 #[tokio::test]
685 async fn test_encrypt_decrypt_roundtrip() {
686 use crate::io::FileIO;
687
688 let io = FileIO::new_with_memory();
689 let path = "memory:///test/encrypt_roundtrip.bin";
690
691 let kms = MemoryKeyManagementClient::new();
692 kms.add_master_key("master-1").unwrap();
693 let mgr = EncryptionManager::builder()
694 .kms_client(Arc::new(kms) as Arc<dyn KeyManagementClient>)
695 .table_key_id("master-1")
696 .build();
697
698 let output = io.new_output(path).unwrap();
699 let encrypted_output = mgr.encrypt(output);
700
701 let plaintext = b"Hello, encrypted Iceberg round-trip!";
702 let serialized_metadata = encrypted_output.key_metadata().encode().unwrap();
703 encrypted_output
704 .write(bytes::Bytes::from(plaintext.to_vec()))
705 .await
706 .unwrap();
707
708 let input = io.new_input(path).unwrap();
709 let parsed_metadata = StandardKeyMetadata::decode(&serialized_metadata).unwrap();
710 let decrypted_file = EncryptedInputFile::new(input, parsed_metadata);
711
712 let content = decrypted_file.read().await.unwrap();
713 assert_eq!(&content[..], plaintext);
714 }
715}