1use std::ops::Range;
47use std::sync::Arc;
48
49use bytes::{Bytes, BytesMut};
50
51use super::AesGcmCipher;
52use crate::error::invalid_data;
53use crate::io::{FileMetadata, FileRead, FileWrite};
54use crate::{Error, ErrorKind, Result};
55
56pub const PLAIN_BLOCK_SIZE: u32 = 1024 * 1024;
58
59pub const NONCE_LENGTH: u32 = 12;
61
62pub const GCM_TAG_LENGTH: u32 = 16;
64
65pub const CIPHER_BLOCK_SIZE: u32 = PLAIN_BLOCK_SIZE + NONCE_LENGTH + GCM_TAG_LENGTH;
67
68pub const GCM_STREAM_MAGIC: [u8; 4] = *b"AGS1";
70
71pub const GCM_STREAM_HEADER_LENGTH: u32 = 8;
73
74pub(crate) const MIN_STREAM_LENGTH: u32 = GCM_STREAM_HEADER_LENGTH + NONCE_LENGTH + GCM_TAG_LENGTH;
76
77pub(crate) fn stream_block_aad(aad_prefix: &[u8], block_index: u32) -> Vec<u8> {
83 let index_bytes = block_index.to_le_bytes();
84 if aad_prefix.is_empty() {
85 index_bytes.to_vec()
86 } else {
87 let mut aad = Vec::with_capacity(aad_prefix.len() + 4);
88 aad.extend_from_slice(aad_prefix);
89 aad.extend_from_slice(&index_bytes);
90 aad
91 }
92}
93
94pub struct AesGcmFileRead {
116 inner: Box<dyn FileRead>,
118 cipher: Arc<AesGcmCipher>,
120 aad_prefix: Box<[u8]>,
122 plain_stream_size: u64,
124 num_blocks: u64,
126 last_cipher_block_size: u32,
128}
129
130impl AesGcmFileRead {
131 pub fn new(
144 inner: Box<dyn FileRead>,
145 cipher: Arc<AesGcmCipher>,
146 aad_prefix: Box<[u8]>,
147 encrypted_file_length: u64,
148 ) -> Result<Self> {
149 if encrypted_file_length < u64::from(MIN_STREAM_LENGTH) {
150 return Err(invalid_data!(
151 "Invalid encrypted file length: {encrypted_file_length} is less than {MIN_STREAM_LENGTH}"
152 ));
153 }
154 let plain_stream_size = Self::calculate_plaintext_length(encrypted_file_length)?;
155 let stream_length = encrypted_file_length - GCM_STREAM_HEADER_LENGTH as u64;
156
157 let num_full_blocks = stream_length / CIPHER_BLOCK_SIZE as u64;
158 let cipher_bytes_in_last_block = (stream_length % CIPHER_BLOCK_SIZE as u64) as u32;
159 let full_blocks_only = cipher_bytes_in_last_block == 0;
160
161 let num_blocks = if full_blocks_only {
162 num_full_blocks
163 } else {
164 num_full_blocks + 1
165 };
166
167 if num_blocks > u32::MAX as u64 {
168 return Err(invalid_data!(
169 "AGS1 format supports at most {} blocks (~4 TiB per file), but file requires {num_blocks} blocks",
170 u32::MAX
171 ));
172 }
173
174 let last_cipher_block_size = if full_blocks_only {
175 CIPHER_BLOCK_SIZE
176 } else {
177 cipher_bytes_in_last_block
178 };
179
180 Ok(Self {
181 inner,
182 cipher,
183 aad_prefix,
184 plain_stream_size,
185 num_blocks,
186 last_cipher_block_size,
187 })
188 }
189
190 pub fn plaintext_length(&self) -> u64 {
192 self.plain_stream_size
193 }
194
195 pub fn calculate_plaintext_length(encrypted_file_length: u64) -> Result<u64> {
200 if encrypted_file_length < GCM_STREAM_HEADER_LENGTH as u64 {
201 return Err(invalid_data!(
202 "Encrypted file too short: {encrypted_file_length} bytes (minimum {GCM_STREAM_HEADER_LENGTH})"
203 ));
204 }
205
206 let stream_length = encrypted_file_length - GCM_STREAM_HEADER_LENGTH as u64;
207
208 if stream_length == 0 {
209 return Ok(0);
210 }
211
212 let num_full_blocks = stream_length / CIPHER_BLOCK_SIZE as u64;
213 let cipher_bytes_in_last_block = stream_length % CIPHER_BLOCK_SIZE as u64;
214 let full_blocks_only = cipher_bytes_in_last_block == 0;
215
216 let plain_bytes_in_last_block = if full_blocks_only {
217 0
218 } else {
219 if cipher_bytes_in_last_block < (NONCE_LENGTH + GCM_TAG_LENGTH) as u64 {
220 return Err(invalid_data!(
221 "Truncated encrypted file: last block is {} bytes (minimum {})",
222 cipher_bytes_in_last_block,
223 NONCE_LENGTH + GCM_TAG_LENGTH
224 ));
225 }
226 cipher_bytes_in_last_block - NONCE_LENGTH as u64 - GCM_TAG_LENGTH as u64
227 };
228
229 Ok(num_full_blocks * PLAIN_BLOCK_SIZE as u64 + plain_bytes_in_last_block)
230 }
231
232 fn encrypted_block_offset(block_index: u64) -> u64 {
234 block_index * CIPHER_BLOCK_SIZE as u64 + GCM_STREAM_HEADER_LENGTH as u64
235 }
236
237 fn cipher_block_size(&self, block_index: u64) -> u32 {
239 if block_index == self.num_blocks - 1 {
240 self.last_cipher_block_size
241 } else {
242 CIPHER_BLOCK_SIZE
243 }
244 }
245}
246
247#[async_trait::async_trait]
248impl FileRead for AesGcmFileRead {
249 async fn read(&self, range: Range<u64>) -> Result<Bytes> {
276 if range.start == range.end && self.plain_stream_size != 0 {
279 return Ok(Bytes::new());
280 }
281
282 if range.start > range.end {
283 return Err(invalid_data!(
284 "Invalid read range: start ({}) is greater than end ({})",
285 range.start,
286 range.end
287 ));
288 }
289
290 if range.end > self.plain_stream_size {
291 return Err(invalid_data!(
292 "Read range {}..{} exceeds plaintext size {}",
293 range.start,
294 range.end,
295 self.plain_stream_size
296 ));
297 }
298
299 let first_block = range.start / PLAIN_BLOCK_SIZE as u64;
300 let last_block = range.end.saturating_sub(1) / PLAIN_BLOCK_SIZE as u64;
301
302 let encrypted_start = Self::encrypted_block_offset(first_block);
304 let encrypted_end =
305 Self::encrypted_block_offset(last_block) + self.cipher_block_size(last_block) as u64;
306
307 let all_encrypted = self.inner.read(encrypted_start..encrypted_end).await?;
308 if all_encrypted.len() as u64 != encrypted_end - encrypted_start {
309 return Err(invalid_data!(
310 "Invalid encrypted read length: expected {} bytes, got {}",
311 encrypted_end - encrypted_start,
312 all_encrypted.len()
313 ));
314 }
315
316 let result_len = (range.end - range.start) as usize;
318 let mut result = BytesMut::with_capacity(result_len);
319 let mut encrypted_offset = 0usize;
320
321 for block_idx in first_block..=last_block {
322 let block_size = self.cipher_block_size(block_idx) as usize;
323 let cipher_block = &all_encrypted[encrypted_offset..encrypted_offset + block_size];
324 encrypted_offset += block_size;
325
326 let aad = stream_block_aad(&self.aad_prefix, block_idx as u32);
327 let decrypted = self.cipher.decrypt(cipher_block, Some(&aad))?;
328
329 let block_plain_start = block_idx * PLAIN_BLOCK_SIZE as u64;
331 let slice_start = if block_idx == first_block {
332 (range.start - block_plain_start) as usize
333 } else {
334 0
335 };
336 let slice_end = if block_idx == last_block {
337 (range.end - block_plain_start) as usize
338 } else {
339 decrypted.len()
340 };
341
342 result.extend_from_slice(&decrypted[slice_start..slice_end]);
343 }
344
345 Ok(result.freeze())
346 }
347}
348
349pub struct AesGcmFileWrite {
369 inner: Box<dyn FileWrite>,
371 cipher: Arc<AesGcmCipher>,
373 aad_prefix: Box<[u8]>,
375 buffer: Vec<u8>,
377 block_index: u32,
379 header_written: bool,
381 closed: bool,
383 poisoned: bool,
387}
388
389impl AesGcmFileWrite {
390 pub fn new(
394 inner: Box<dyn FileWrite>,
395 cipher: Arc<AesGcmCipher>,
396 aad_prefix: impl Into<Box<[u8]>>,
397 ) -> Self {
398 Self {
399 inner,
400 cipher,
401 aad_prefix: aad_prefix.into(),
402 buffer: Vec::new(),
403 block_index: 0,
404 header_written: false,
405 closed: false,
406 poisoned: false,
407 }
408 }
409
410 async fn write_header(&mut self) -> Result<()> {
412 let mut header = Vec::with_capacity(GCM_STREAM_HEADER_LENGTH as usize);
413 header.extend_from_slice(&GCM_STREAM_MAGIC);
414 header.extend_from_slice(&PLAIN_BLOCK_SIZE.to_le_bytes());
415 if let Err(e) = self.inner.write(Bytes::from(header)).await {
416 self.poisoned = true;
417 return Err(e);
418 }
419 self.header_written = true;
420 Ok(())
421 }
422
423 async fn encrypt_and_write_block(&mut self, block_data: &[u8]) -> Result<()> {
425 let aad = stream_block_aad(&self.aad_prefix, self.block_index);
426 let encrypted = self.cipher.encrypt(block_data, Some(&aad))?;
427 if let Err(e) = self.inner.write(Bytes::from(encrypted)).await {
428 self.poisoned = true;
429 return Err(e);
430 }
431 self.block_index = self.block_index.checked_add(1).ok_or_else(|| {
432 invalid_data!(
433 "AGS1 block index overflow: file exceeds the maximum supported size (~4 TiB)"
434 )
435 })?;
436 Ok(())
437 }
438
439 async fn encrypt_and_drain_block(&mut self) -> Result<()> {
442 let aad = stream_block_aad(&self.aad_prefix, self.block_index);
443 let encrypted = self
444 .cipher
445 .encrypt(&self.buffer[..PLAIN_BLOCK_SIZE as usize], Some(&aad))?;
446 if let Err(e) = self.inner.write(Bytes::from(encrypted)).await {
447 self.poisoned = true;
448 return Err(e);
449 }
450 self.block_index = self.block_index.checked_add(1).ok_or_else(|| {
451 invalid_data!(
452 "AGS1 block index overflow: file exceeds the maximum supported size (~4 TiB)"
453 )
454 })?;
455 self.buffer.drain(..PLAIN_BLOCK_SIZE as usize);
456 Ok(())
457 }
458}
459
460#[async_trait::async_trait]
461impl FileWrite for AesGcmFileWrite {
462 async fn write(&mut self, bs: Bytes) -> Result<()> {
463 if self.closed {
464 return Err(Error::new(
465 ErrorKind::Unexpected,
466 "Cannot write to a closed AesGcmFileWrite",
467 ));
468 }
469 if self.poisoned {
470 return Err(Error::new(
471 ErrorKind::Unexpected,
472 "AesGcmFileWrite is in a poisoned state due to a previous write failure",
473 ));
474 }
475
476 if !self.header_written {
477 self.write_header().await?;
478 }
479
480 self.buffer.extend_from_slice(&bs);
481
482 while self.buffer.len() >= PLAIN_BLOCK_SIZE as usize {
484 self.encrypt_and_drain_block().await?;
485 }
486
487 Ok(())
488 }
489
490 async fn close(&mut self) -> Result<FileMetadata> {
491 if self.closed {
492 return Err(Error::new(
493 ErrorKind::Unexpected,
494 "AesGcmFileWrite already closed",
495 ));
496 }
497 if self.poisoned {
498 return Err(Error::new(
499 ErrorKind::Unexpected,
500 "AesGcmFileWrite is in a poisoned state due to a previous write failure",
501 ));
502 }
503
504 if !self.header_written {
505 self.write_header().await?;
506 }
507
508 if !self.buffer.is_empty() || self.block_index == 0 {
512 let final_block = std::mem::take(&mut self.buffer);
513 self.encrypt_and_write_block(&final_block).await?;
514 }
515 self.closed = true;
516
517 self.inner.close().await
518 }
519}
520
521#[cfg(test)]
522mod tests {
523 use super::*;
524
525 fn encrypt_ags1(plaintext: &[u8], cipher: &AesGcmCipher, aad_prefix: &[u8]) -> Vec<u8> {
531 let mut result = Vec::new();
532
533 result.extend_from_slice(&GCM_STREAM_MAGIC);
535 result.extend_from_slice(&PLAIN_BLOCK_SIZE.to_le_bytes());
536
537 let mut offset = 0;
539 let mut block_index = 0u32;
540
541 loop {
542 let remaining = plaintext.len() - offset;
543 let block_size = std::cmp::min(remaining, PLAIN_BLOCK_SIZE as usize);
544
545 if block_size == 0 && block_index > 0 {
547 break;
548 }
549
550 let block_data = &plaintext[offset..offset + block_size];
551 let aad = stream_block_aad(aad_prefix, block_index);
552 let encrypted = cipher.encrypt(block_data, Some(&aad)).unwrap();
553 result.extend_from_slice(&encrypted);
554
555 offset += block_size;
556 block_index += 1;
557
558 if block_size < PLAIN_BLOCK_SIZE as usize {
560 break;
561 }
562 }
563
564 result
565 }
566
567 fn make_cipher(key: &[u8]) -> AesGcmCipher {
569 use super::super::SecureKey;
570 let secure_key = SecureKey::new(key).unwrap();
571 AesGcmCipher::new(secure_key)
572 }
573
574 fn memory_reader(data: Vec<u8>) -> Box<dyn FileRead> {
576 Box::new(MemoryFileRead(Bytes::from(data)))
577 }
578
579 struct MemoryFileRead(Bytes);
581
582 #[async_trait::async_trait]
583 impl FileRead for MemoryFileRead {
584 async fn read(&self, range: Range<u64>) -> Result<Bytes> {
585 let start = range.start as usize;
586 let end = range.end as usize;
587 if end > self.0.len() {
588 return Err(invalid_data!(
589 "Range {}..{} out of bounds for {} bytes",
590 start,
591 end,
592 self.0.len()
593 ));
594 }
595 Ok(self.0.slice(start..end))
596 }
597 }
598
599 #[tokio::test]
600 async fn test_empty_file_roundtrip() {
601 let key = b"0123456789abcdef";
602 let aad_prefix = b"test-aad-prefix!";
603 let cipher = make_cipher(key);
604
605 let encrypted = encrypt_ags1(b"", &cipher, aad_prefix);
606
607 assert_eq!(encrypted.len(), MIN_STREAM_LENGTH as usize);
609
610 let reader = AesGcmFileRead::new(
611 memory_reader(encrypted.clone()),
612 Arc::new(make_cipher(key)),
613 aad_prefix.as_slice().into(),
614 encrypted.len() as u64,
615 )
616 .unwrap();
617
618 assert_eq!(reader.plaintext_length(), 0);
619
620 let result = reader.read(0..0).await.unwrap();
622 assert!(result.is_empty());
623 }
624
625 #[tokio::test]
626 async fn test_short_ciphertext_read_is_rejected() {
627 struct ShortRead;
628
629 #[async_trait::async_trait]
630 impl FileRead for ShortRead {
631 async fn read(&self, range: Range<u64>) -> Result<Bytes> {
632 let len = (range.end - range.start).saturating_sub(1) as usize;
633 Ok(Bytes::from(vec![0; len]))
634 }
635 }
636
637 let reader = AesGcmFileRead::new(
638 Box::new(ShortRead),
639 Arc::new(make_cipher(b"0123456789abcdef")),
640 Box::default(),
641 u64::from(MIN_STREAM_LENGTH) + 10,
642 )
643 .unwrap();
644 let err = reader.read(0..10).await.unwrap_err();
645 assert_eq!(err.kind(), ErrorKind::DataInvalid);
646 assert!(err.to_string().contains("Invalid encrypted read length"));
647 }
648
649 #[tokio::test]
650 async fn test_oversized_declared_length_is_rejected() {
651 struct ClampingRead(Bytes);
653
654 #[async_trait::async_trait]
655 impl FileRead for ClampingRead {
656 async fn read(&self, range: Range<u64>) -> Result<Bytes> {
657 let start = (range.start as usize).min(self.0.len());
658 let end = (range.end as usize).min(self.0.len());
659 Ok(self.0.slice(start..end))
660 }
661 }
662
663 let key = b"0123456789abcdef";
664 let aad_prefix = b"test-aad-prefix!";
665 let plaintext = b"some bytes to measure";
666 let encrypted = write_through_ags1(plaintext, key, aad_prefix).await;
667
668 for excess in [1, u64::from(CIPHER_BLOCK_SIZE)] {
671 let reader = AesGcmFileRead::new(
672 Box::new(ClampingRead(Bytes::from(encrypted.clone()))),
673 Arc::new(make_cipher(key)),
674 aad_prefix.to_vec().into_boxed_slice(),
675 encrypted.len() as u64 + excess,
676 )
677 .unwrap();
678
679 let err = reader
680 .read(0..plaintext.len() as u64)
681 .await
682 .expect_err("an inflated declared length must not read back as plaintext");
683 assert_eq!(err.kind(), ErrorKind::DataInvalid);
684 assert!(err.to_string().contains("Invalid encrypted read length"));
685 }
686 }
687
688 #[tokio::test]
689 async fn test_small_file_roundtrip() {
690 let key = b"0123456789abcdef";
691 let aad_prefix = b"test-aad-prefix!";
692 let plaintext = b"Hello, Iceberg encryption!";
693 let cipher = make_cipher(key);
694
695 let encrypted = encrypt_ags1(plaintext, &cipher, aad_prefix);
696
697 let reader = AesGcmFileRead::new(
698 memory_reader(encrypted.clone()),
699 Arc::new(make_cipher(key)),
700 aad_prefix.as_slice().into(),
701 encrypted.len() as u64,
702 )
703 .unwrap();
704
705 assert_eq!(reader.plaintext_length(), plaintext.len() as u64);
706
707 let result = reader.read(0..plaintext.len() as u64).await.unwrap();
709 assert_eq!(&result[..], plaintext);
710 }
711
712 #[tokio::test]
713 async fn test_partial_read() {
714 let key = b"0123456789abcdef";
715 let aad_prefix = b"aad-prefix-here!";
716 let plaintext = b"ABCDEFGHIJKLMNOPQRSTUVWXYZ";
717 let cipher = make_cipher(key);
718
719 let encrypted = encrypt_ags1(plaintext, &cipher, aad_prefix);
720
721 let reader = AesGcmFileRead::new(
722 memory_reader(encrypted.clone()),
723 Arc::new(make_cipher(key)),
724 aad_prefix.as_slice().into(),
725 encrypted.len() as u64,
726 )
727 .unwrap();
728
729 let result = reader.read(10..20).await.unwrap();
731 assert_eq!(&result[..], &plaintext[10..20]);
732
733 let result = reader.read(0..1).await.unwrap();
735 assert_eq!(&result[..], &plaintext[0..1]);
736
737 let last = plaintext.len() as u64;
739 let result = reader.read(last - 1..last).await.unwrap();
740 assert_eq!(&result[..], &plaintext[plaintext.len() - 1..]);
741 }
742
743 #[tokio::test]
744 async fn test_multi_block_roundtrip() {
745 let key = b"0123456789abcdef";
746 let aad_prefix = b"multi-block-aad!";
747
748 let size = PLAIN_BLOCK_SIZE as usize + PLAIN_BLOCK_SIZE as usize / 2;
750 let plaintext: Vec<u8> = (0..size).map(|i| (i % 256) as u8).collect();
751 let cipher = make_cipher(key);
752
753 let encrypted = encrypt_ags1(&plaintext, &cipher, aad_prefix);
754
755 let reader = AesGcmFileRead::new(
756 memory_reader(encrypted.clone()),
757 Arc::new(make_cipher(key)),
758 aad_prefix.as_slice().into(),
759 encrypted.len() as u64,
760 )
761 .unwrap();
762
763 assert_eq!(reader.plaintext_length(), plaintext.len() as u64);
764
765 let result = reader.read(0..plaintext.len() as u64).await.unwrap();
767 assert_eq!(&result[..], &plaintext[..]);
768 }
769
770 #[tokio::test]
771 async fn test_cross_block_read() {
772 let key = b"0123456789abcdef";
773 let aad_prefix = b"cross-block-aad!";
774
775 let size = PLAIN_BLOCK_SIZE as usize * 2 + PLAIN_BLOCK_SIZE as usize / 2;
777 let plaintext: Vec<u8> = (0..size).map(|i| (i % 256) as u8).collect();
778 let cipher = make_cipher(key);
779
780 let encrypted = encrypt_ags1(&plaintext, &cipher, aad_prefix);
781
782 let reader = AesGcmFileRead::new(
783 memory_reader(encrypted.clone()),
784 Arc::new(make_cipher(key)),
785 aad_prefix.as_slice().into(),
786 encrypted.len() as u64,
787 )
788 .unwrap();
789
790 let boundary = PLAIN_BLOCK_SIZE as u64;
792 let result = reader.read(boundary - 100..boundary + 100).await.unwrap();
793 assert_eq!(
794 &result[..],
795 &plaintext[(boundary - 100) as usize..(boundary + 100) as usize]
796 );
797
798 let result = reader.read(boundary - 50..boundary * 2 + 50).await.unwrap();
800 assert_eq!(
801 &result[..],
802 &plaintext[(boundary - 50) as usize..(boundary * 2 + 50) as usize]
803 );
804 }
805
806 #[tokio::test]
807 async fn test_exact_block_size() {
808 let key = b"0123456789abcdef";
809 let aad_prefix = b"exact-block-aad!";
810
811 let plaintext: Vec<u8> = (0..PLAIN_BLOCK_SIZE as usize)
813 .map(|i| (i % 256) as u8)
814 .collect();
815 let cipher = make_cipher(key);
816
817 let encrypted = encrypt_ags1(&plaintext, &cipher, aad_prefix);
818
819 let reader = AesGcmFileRead::new(
820 memory_reader(encrypted.clone()),
821 Arc::new(make_cipher(key)),
822 aad_prefix.as_slice().into(),
823 encrypted.len() as u64,
824 )
825 .unwrap();
826
827 assert_eq!(reader.plaintext_length(), PLAIN_BLOCK_SIZE as u64);
828
829 let result = reader.read(0..PLAIN_BLOCK_SIZE as u64).await.unwrap();
830 assert_eq!(&result[..], &plaintext[..]);
831 }
832
833 #[tokio::test]
834 async fn test_block_size_plus_one() {
835 let key = b"0123456789abcdef";
836 let aad_prefix = b"block-plus-one!!";
837
838 let size = PLAIN_BLOCK_SIZE as usize + 1;
840 let plaintext: Vec<u8> = (0..size).map(|i| (i % 256) as u8).collect();
841 let cipher = make_cipher(key);
842
843 let encrypted = encrypt_ags1(&plaintext, &cipher, aad_prefix);
844
845 let reader = AesGcmFileRead::new(
846 memory_reader(encrypted.clone()),
847 Arc::new(make_cipher(key)),
848 aad_prefix.as_slice().into(),
849 encrypted.len() as u64,
850 )
851 .unwrap();
852
853 assert_eq!(reader.plaintext_length(), size as u64);
854
855 let result = reader.read(size as u64 - 1..size as u64).await.unwrap();
857 assert_eq!(result[0], plaintext[size - 1]);
858
859 let result = reader.read(0..size as u64).await.unwrap();
861 assert_eq!(&result[..], &plaintext[..]);
862 }
863
864 #[tokio::test]
865 async fn test_block_size_minus_one() {
866 let key = b"0123456789abcdef";
867 let aad_prefix = b"block-minus-one!";
868
869 let size = PLAIN_BLOCK_SIZE as usize - 1;
871 let plaintext: Vec<u8> = (0..size).map(|i| (i % 256) as u8).collect();
872 let cipher = make_cipher(key);
873
874 let encrypted = encrypt_ags1(&plaintext, &cipher, aad_prefix);
875
876 let reader = AesGcmFileRead::new(
877 memory_reader(encrypted.clone()),
878 Arc::new(make_cipher(key)),
879 aad_prefix.as_slice().into(),
880 encrypted.len() as u64,
881 )
882 .unwrap();
883
884 assert_eq!(reader.plaintext_length(), size as u64);
885
886 let result = reader.read(0..size as u64).await.unwrap();
887 assert_eq!(&result[..], &plaintext[..]);
888 }
889
890 #[tokio::test]
891 async fn test_wrong_aad_fails() {
892 let key = b"0123456789abcdef";
893 let aad_prefix = b"correct-aad-here";
894 let plaintext = b"sensitive data here";
895 let cipher = make_cipher(key);
896
897 let encrypted = encrypt_ags1(plaintext, &cipher, aad_prefix);
898
899 let mut bad_aad = aad_prefix.to_vec();
901 bad_aad[0] ^= 0xFF;
902
903 let reader = AesGcmFileRead::new(
904 memory_reader(encrypted.clone()),
905 Arc::new(make_cipher(key)),
906 bad_aad.as_slice().into(),
907 encrypted.len() as u64,
908 )
909 .unwrap();
910
911 let result = reader.read(0..plaintext.len() as u64).await;
912 assert!(result.is_err(), "Decryption with wrong AAD should fail");
913 }
914
915 #[tokio::test]
916 async fn test_wrong_key_fails() {
917 let key = b"0123456789abcdef";
918 let wrong_key = b"fedcba9876543210";
919 let aad_prefix = b"test-aad-prefix!";
920 let plaintext = b"sensitive data";
921 let cipher = make_cipher(key);
922
923 let encrypted = encrypt_ags1(plaintext, &cipher, aad_prefix);
924
925 let reader = AesGcmFileRead::new(
926 memory_reader(encrypted.clone()),
927 Arc::new(make_cipher(wrong_key)),
928 aad_prefix.as_slice().into(),
929 encrypted.len() as u64,
930 )
931 .unwrap();
932
933 let result = reader.read(0..plaintext.len() as u64).await;
934 assert!(result.is_err(), "Decryption with wrong key should fail");
935 }
936
937 #[tokio::test]
938 async fn test_out_of_bounds_read() {
939 let key = b"0123456789abcdef";
940 let aad_prefix = b"test-aad-prefix!";
941 let plaintext = b"short data";
942 let cipher = make_cipher(key);
943
944 let encrypted = encrypt_ags1(plaintext, &cipher, aad_prefix);
945
946 let reader = AesGcmFileRead::new(
947 memory_reader(encrypted.clone()),
948 Arc::new(make_cipher(key)),
949 aad_prefix.as_slice().into(),
950 encrypted.len() as u64,
951 )
952 .unwrap();
953
954 let result = reader.read(0..plaintext.len() as u64 + 1).await;
955 assert!(result.is_err(), "Reading past end should fail");
956 }
957
958 #[tokio::test]
959 async fn test_calculate_plaintext_length() {
960 assert_eq!(
962 AesGcmFileRead::calculate_plaintext_length(GCM_STREAM_HEADER_LENGTH as u64).unwrap(),
963 0
964 );
965
966 assert_eq!(
968 AesGcmFileRead::calculate_plaintext_length(MIN_STREAM_LENGTH as u64).unwrap(),
969 0
970 );
971
972 let one_full = GCM_STREAM_HEADER_LENGTH as u64 + CIPHER_BLOCK_SIZE as u64;
974 assert_eq!(
975 AesGcmFileRead::calculate_plaintext_length(one_full).unwrap(),
976 PLAIN_BLOCK_SIZE as u64
977 );
978
979 let one_full_plus_one = one_full + NONCE_LENGTH as u64 + 1 + GCM_TAG_LENGTH as u64;
982 assert_eq!(
983 AesGcmFileRead::calculate_plaintext_length(one_full_plus_one).unwrap(),
984 PLAIN_BLOCK_SIZE as u64 + 1
985 );
986 }
987
988 #[tokio::test]
989 async fn test_stream_block_aad() {
990 let aad = stream_block_aad(b"prefix", 0);
992 assert_eq!(&aad[..6], b"prefix");
993 assert_eq!(&aad[6..], &0u32.to_le_bytes());
994
995 let aad = stream_block_aad(b"prefix", 1);
996 assert_eq!(&aad[..6], b"prefix");
997 assert_eq!(&aad[6..], &1u32.to_le_bytes());
998
999 let aad = stream_block_aad(b"", 42);
1001 assert_eq!(&aad[..], &42u32.to_le_bytes());
1002 }
1003
1004 #[test]
1005 fn test_encrypted_file_too_short() {
1006 for length in 0..MIN_STREAM_LENGTH {
1007 let result = AesGcmFileRead::new(
1008 memory_reader(vec![0; length as usize]),
1009 Arc::new(make_cipher(b"0123456789abcdef")),
1010 [].into(),
1011 u64::from(length),
1012 );
1013 let err = result
1014 .err()
1015 .expect("a stream must contain an authenticated block");
1016 assert_eq!(err.kind(), ErrorKind::DataInvalid);
1017 assert!(err.to_string().contains("Invalid encrypted file length"));
1018 }
1019 }
1020
1021 struct SharedMemoryWrite {
1025 buffer: Arc<std::sync::Mutex<Vec<u8>>>,
1026 }
1027
1028 struct FailingFileWrite {
1030 writes_before_failure: usize,
1031 write_count: usize,
1032 }
1033
1034 #[async_trait::async_trait]
1035 impl FileWrite for FailingFileWrite {
1036 async fn write(&mut self, _bs: Bytes) -> Result<()> {
1037 if self.write_count >= self.writes_before_failure {
1038 return Err(Error::new(ErrorKind::Unexpected, "simulated write failure"));
1039 }
1040 self.write_count += 1;
1041 Ok(())
1042 }
1043
1044 async fn close(&mut self) -> Result<FileMetadata> {
1048 Err(Error::new(
1049 ErrorKind::Unexpected,
1050 "FailingFileWrite::close called unexpectedly",
1051 ))
1052 }
1053 }
1054
1055 #[async_trait::async_trait]
1056 impl FileWrite for SharedMemoryWrite {
1057 async fn write(&mut self, bs: Bytes) -> Result<()> {
1058 self.buffer.lock().unwrap().extend_from_slice(&bs);
1059 Ok(())
1060 }
1061
1062 async fn close(&mut self) -> Result<FileMetadata> {
1063 Ok(FileMetadata {
1064 size: self.buffer.lock().unwrap().len() as u64,
1065 })
1066 }
1067 }
1068
1069 async fn write_through_ags1(plaintext: &[u8], key: &[u8], aad_prefix: &[u8]) -> Vec<u8> {
1071 let buffer = Arc::new(std::sync::Mutex::new(Vec::new()));
1072 let inner: Box<dyn FileWrite> = Box::new(SharedMemoryWrite {
1073 buffer: buffer.clone(),
1074 });
1075 let cipher = Arc::new(make_cipher(key));
1076 let mut writer = AesGcmFileWrite::new(inner, cipher, aad_prefix.to_vec());
1077
1078 writer.write(Bytes::from(plaintext.to_vec())).await.unwrap();
1079 let metadata = writer.close().await.unwrap();
1080
1081 let encrypted = buffer.lock().unwrap().clone();
1082 assert_eq!(
1083 metadata.size,
1084 encrypted.len() as u64,
1085 "close() must report the full ciphertext length"
1086 );
1087 encrypted
1088 }
1089
1090 #[tokio::test]
1091 async fn test_write_empty_roundtrip() {
1092 let key = b"0123456789abcdef";
1093 let aad_prefix = b"test-aad-prefix!";
1094
1095 let encrypted = write_through_ags1(b"", key, aad_prefix).await;
1096
1097 assert_eq!(encrypted.len(), MIN_STREAM_LENGTH as usize);
1099
1100 let reader = AesGcmFileRead::new(
1101 memory_reader(encrypted.clone()),
1102 Arc::new(make_cipher(key)),
1103 aad_prefix.as_slice().into(),
1104 encrypted.len() as u64,
1105 )
1106 .unwrap();
1107
1108 assert_eq!(reader.plaintext_length(), 0);
1109 }
1110
1111 #[tokio::test]
1112 async fn test_write_small_roundtrip() {
1113 let key = b"0123456789abcdef";
1114 let aad_prefix = b"test-aad-prefix!";
1115 let plaintext = b"Hello, Iceberg encryption!";
1116
1117 let encrypted = write_through_ags1(plaintext, key, aad_prefix).await;
1118
1119 let reader = AesGcmFileRead::new(
1120 memory_reader(encrypted.clone()),
1121 Arc::new(make_cipher(key)),
1122 aad_prefix.as_slice().into(),
1123 encrypted.len() as u64,
1124 )
1125 .unwrap();
1126
1127 assert_eq!(reader.plaintext_length(), plaintext.len() as u64);
1128 let result = reader.read(0..plaintext.len() as u64).await.unwrap();
1129 assert_eq!(&result[..], plaintext);
1130 }
1131
1132 #[tokio::test]
1133 async fn test_write_multi_block_roundtrip() {
1134 let key = b"0123456789abcdef";
1135 let aad_prefix = b"multi-block-aad!";
1136
1137 let size = PLAIN_BLOCK_SIZE as usize + PLAIN_BLOCK_SIZE as usize / 2;
1139 let plaintext: Vec<u8> = (0..size).map(|i| (i % 256) as u8).collect();
1140
1141 let encrypted = write_through_ags1(&plaintext, key, aad_prefix).await;
1142
1143 let reader = AesGcmFileRead::new(
1144 memory_reader(encrypted.clone()),
1145 Arc::new(make_cipher(key)),
1146 aad_prefix.as_slice().into(),
1147 encrypted.len() as u64,
1148 )
1149 .unwrap();
1150
1151 assert_eq!(reader.plaintext_length(), plaintext.len() as u64);
1152 let result = reader.read(0..plaintext.len() as u64).await.unwrap();
1153 assert_eq!(&result[..], &plaintext[..]);
1154 }
1155
1156 #[tokio::test]
1157 async fn test_write_cross_block_accumulation() {
1158 let key = b"0123456789abcdef";
1159 let aad_prefix = b"cross-block-aad!";
1160
1161 let buffer = Arc::new(std::sync::Mutex::new(Vec::new()));
1162 let inner: Box<dyn FileWrite> = Box::new(SharedMemoryWrite {
1163 buffer: buffer.clone(),
1164 });
1165 let cipher = Arc::new(make_cipher(key));
1166 let mut writer = AesGcmFileWrite::new(inner, cipher, aad_prefix.to_vec());
1167
1168 let total_size = PLAIN_BLOCK_SIZE as usize + PLAIN_BLOCK_SIZE as usize / 2;
1170 let plaintext: Vec<u8> = (0..total_size).map(|i| (i % 256) as u8).collect();
1171 let chunk_size = 1000;
1172 for chunk in plaintext.chunks(chunk_size) {
1173 writer.write(Bytes::from(chunk.to_vec())).await.unwrap();
1174 }
1175 writer.close().await.unwrap();
1176
1177 let encrypted = buffer.lock().unwrap().clone();
1178
1179 let reader = AesGcmFileRead::new(
1180 memory_reader(encrypted.clone()),
1181 Arc::new(make_cipher(key)),
1182 aad_prefix.as_slice().into(),
1183 encrypted.len() as u64,
1184 )
1185 .unwrap();
1186
1187 assert_eq!(reader.plaintext_length(), plaintext.len() as u64);
1188 let result = reader.read(0..plaintext.len() as u64).await.unwrap();
1189 assert_eq!(&result[..], &plaintext[..]);
1190 }
1191
1192 #[tokio::test]
1193 async fn test_write_exact_block_size() {
1194 let key = b"0123456789abcdef";
1195 let aad_prefix = b"exact-block-aad!";
1196
1197 let plaintext: Vec<u8> = (0..PLAIN_BLOCK_SIZE as usize)
1199 .map(|i| (i % 256) as u8)
1200 .collect();
1201
1202 let encrypted = write_through_ags1(&plaintext, key, aad_prefix).await;
1203
1204 let reader = AesGcmFileRead::new(
1205 memory_reader(encrypted.clone()),
1206 Arc::new(make_cipher(key)),
1207 aad_prefix.as_slice().into(),
1208 encrypted.len() as u64,
1209 )
1210 .unwrap();
1211
1212 assert_eq!(reader.plaintext_length(), PLAIN_BLOCK_SIZE as u64);
1213 let result = reader.read(0..PLAIN_BLOCK_SIZE as u64).await.unwrap();
1214 assert_eq!(&result[..], &plaintext[..]);
1215 }
1216
1217 #[tokio::test]
1218 async fn test_write_block_aligned_no_spurious_empty_block() {
1219 let key = b"0123456789abcdef";
1220 let aad_prefix = b"block-align-aad!";
1221
1222 let plaintext: Vec<u8> = (0..PLAIN_BLOCK_SIZE as usize)
1225 .map(|i| (i % 256) as u8)
1226 .collect();
1227
1228 let encrypted_via_writer = write_through_ags1(&plaintext, key, aad_prefix).await;
1229 let encrypted_via_reference = encrypt_ags1(&plaintext, &make_cipher(key), aad_prefix);
1230
1231 assert_eq!(
1233 encrypted_via_writer.len(),
1234 encrypted_via_reference.len(),
1235 "Writer output should match reference encryption length (no spurious trailing block)"
1236 );
1237
1238 let reader = AesGcmFileRead::new(
1240 memory_reader(encrypted_via_writer.clone()),
1241 Arc::new(make_cipher(key)),
1242 aad_prefix.as_slice().into(),
1243 encrypted_via_writer.len() as u64,
1244 )
1245 .unwrap();
1246
1247 assert_eq!(reader.plaintext_length(), PLAIN_BLOCK_SIZE as u64);
1248 let result = reader.read(0..PLAIN_BLOCK_SIZE as u64).await.unwrap();
1249 assert_eq!(&result[..], &plaintext[..]);
1250 }
1251
1252 #[tokio::test]
1253 async fn test_write_two_blocks_aligned_no_spurious_empty_block() {
1254 let key = b"0123456789abcdef";
1255 let aad_prefix = b"2blk-align-aad!!";
1256
1257 let size = PLAIN_BLOCK_SIZE as usize * 2;
1259 let plaintext: Vec<u8> = (0..size).map(|i| (i % 256) as u8).collect();
1260
1261 let encrypted_via_writer = write_through_ags1(&plaintext, key, aad_prefix).await;
1262 let encrypted_via_reference = encrypt_ags1(&plaintext, &make_cipher(key), aad_prefix);
1263
1264 assert_eq!(
1265 encrypted_via_writer.len(),
1266 encrypted_via_reference.len(),
1267 "Writer output should match reference encryption length (no spurious trailing block)"
1268 );
1269
1270 let reader = AesGcmFileRead::new(
1271 memory_reader(encrypted_via_writer.clone()),
1272 Arc::new(make_cipher(key)),
1273 aad_prefix.as_slice().into(),
1274 encrypted_via_writer.len() as u64,
1275 )
1276 .unwrap();
1277
1278 assert_eq!(reader.plaintext_length(), size as u64);
1279 let result = reader.read(0..size as u64).await.unwrap();
1280 assert_eq!(&result[..], &plaintext[..]);
1281 }
1282
1283 #[tokio::test]
1284 async fn test_write_poisoned_after_inner_write_failure() {
1285 let cipher = Arc::new(make_cipher(b"0123456789abcdef"));
1286 let inner: Box<dyn FileWrite> = Box::new(FailingFileWrite {
1288 writes_before_failure: 1,
1289 write_count: 0,
1290 });
1291 let mut writer = AesGcmFileWrite::new(inner, cipher, b"aad-prefix-here!".to_vec());
1292
1293 let data = vec![0u8; PLAIN_BLOCK_SIZE as usize];
1295 let result = writer.write(Bytes::from(data)).await;
1296 assert!(result.is_err());
1297
1298 let result = writer.write(Bytes::from(b"more data".to_vec())).await;
1300 assert!(result.is_err());
1301 assert!(
1302 result.unwrap_err().to_string().contains("poisoned"),
1303 "expected poisoned error"
1304 );
1305
1306 let err = writer.close().await.err().expect("close should fail");
1308 assert!(
1309 err.to_string().contains("poisoned"),
1310 "expected poisoned error on close"
1311 );
1312 }
1313}