use std::ops::Range;
use std::sync::Arc;
use bytes::{Bytes, BytesMut};
use super::AesGcmCipher;
use crate::io::{FileRead, FileWrite};
use crate::{Error, ErrorKind, Result};
pub const PLAIN_BLOCK_SIZE: u32 = 1024 * 1024;
pub const NONCE_LENGTH: u32 = 12;
pub const GCM_TAG_LENGTH: u32 = 16;
pub const CIPHER_BLOCK_SIZE: u32 = PLAIN_BLOCK_SIZE + NONCE_LENGTH + GCM_TAG_LENGTH;
pub const GCM_STREAM_MAGIC: [u8; 4] = *b"AGS1";
pub const GCM_STREAM_HEADER_LENGTH: u32 = 8;
#[cfg(test)]
pub const MIN_STREAM_LENGTH: u32 = GCM_STREAM_HEADER_LENGTH + NONCE_LENGTH + GCM_TAG_LENGTH;
pub(crate) fn stream_block_aad(aad_prefix: &[u8], block_index: u32) -> Vec<u8> {
let index_bytes = block_index.to_le_bytes();
if aad_prefix.is_empty() {
index_bytes.to_vec()
} else {
let mut aad = Vec::with_capacity(aad_prefix.len() + 4);
aad.extend_from_slice(aad_prefix);
aad.extend_from_slice(&index_bytes);
aad
}
}
pub struct AesGcmFileRead {
inner: Box<dyn FileRead>,
cipher: Arc<AesGcmCipher>,
aad_prefix: Box<[u8]>,
plain_stream_size: u64,
num_blocks: u64,
last_cipher_block_size: u32,
}
impl AesGcmFileRead {
pub fn new(
inner: Box<dyn FileRead>,
cipher: Arc<AesGcmCipher>,
aad_prefix: Box<[u8]>,
encrypted_file_length: u64,
) -> Result<Self> {
let plain_stream_size = Self::calculate_plaintext_length(encrypted_file_length)?;
let stream_length = encrypted_file_length - GCM_STREAM_HEADER_LENGTH as u64;
if stream_length == 0 {
return Ok(Self {
inner,
cipher,
aad_prefix,
plain_stream_size: 0,
num_blocks: 0,
last_cipher_block_size: 0,
});
}
let num_full_blocks = stream_length / CIPHER_BLOCK_SIZE as u64;
let cipher_bytes_in_last_block = (stream_length % CIPHER_BLOCK_SIZE as u64) as u32;
let full_blocks_only = cipher_bytes_in_last_block == 0;
let num_blocks = if full_blocks_only {
num_full_blocks
} else {
num_full_blocks + 1
};
if num_blocks > u32::MAX as u64 {
return Err(Error::new(
ErrorKind::DataInvalid,
format!(
"AGS1 format supports at most {} blocks (~4 TiB per file), but file requires {num_blocks} blocks",
u32::MAX
),
));
}
let last_cipher_block_size = if full_blocks_only {
CIPHER_BLOCK_SIZE
} else {
cipher_bytes_in_last_block
};
Ok(Self {
inner,
cipher,
aad_prefix,
plain_stream_size,
num_blocks,
last_cipher_block_size,
})
}
pub fn plaintext_length(&self) -> u64 {
self.plain_stream_size
}
pub fn calculate_plaintext_length(encrypted_file_length: u64) -> Result<u64> {
if encrypted_file_length < GCM_STREAM_HEADER_LENGTH as u64 {
return Err(Error::new(
ErrorKind::DataInvalid,
format!(
"Encrypted file too short: {encrypted_file_length} bytes (minimum {GCM_STREAM_HEADER_LENGTH})"
),
));
}
let stream_length = encrypted_file_length - GCM_STREAM_HEADER_LENGTH as u64;
if stream_length == 0 {
return Ok(0);
}
let num_full_blocks = stream_length / CIPHER_BLOCK_SIZE as u64;
let cipher_bytes_in_last_block = stream_length % CIPHER_BLOCK_SIZE as u64;
let full_blocks_only = cipher_bytes_in_last_block == 0;
let plain_bytes_in_last_block = if full_blocks_only {
0
} else {
if cipher_bytes_in_last_block < (NONCE_LENGTH + GCM_TAG_LENGTH) as u64 {
return Err(Error::new(
ErrorKind::DataInvalid,
format!(
"Truncated encrypted file: last block is {} bytes (minimum {})",
cipher_bytes_in_last_block,
NONCE_LENGTH + GCM_TAG_LENGTH
),
));
}
cipher_bytes_in_last_block - NONCE_LENGTH as u64 - GCM_TAG_LENGTH as u64
};
Ok(num_full_blocks * PLAIN_BLOCK_SIZE as u64 + plain_bytes_in_last_block)
}
fn encrypted_block_offset(block_index: u64) -> u64 {
block_index * CIPHER_BLOCK_SIZE as u64 + GCM_STREAM_HEADER_LENGTH as u64
}
fn cipher_block_size(&self, block_index: u64) -> u32 {
if block_index == self.num_blocks - 1 {
self.last_cipher_block_size
} else {
CIPHER_BLOCK_SIZE
}
}
}
#[async_trait::async_trait]
impl FileRead for AesGcmFileRead {
async fn read(&self, range: Range<u64>) -> Result<Bytes> {
if range.start == range.end {
return Ok(Bytes::new());
}
if range.start > range.end {
return Err(Error::new(
ErrorKind::DataInvalid,
format!(
"Invalid read range: start ({}) is greater than end ({})",
range.start, range.end
),
));
}
if range.end > self.plain_stream_size {
return Err(Error::new(
ErrorKind::DataInvalid,
format!(
"Read range {}..{} exceeds plaintext size {}",
range.start, range.end, self.plain_stream_size
),
));
}
if self.num_blocks == 0 {
return Ok(Bytes::new());
}
let first_block = range.start / PLAIN_BLOCK_SIZE as u64;
let last_block = (range.end - 1) / PLAIN_BLOCK_SIZE as u64;
let encrypted_start = Self::encrypted_block_offset(first_block);
let encrypted_end =
Self::encrypted_block_offset(last_block) + self.cipher_block_size(last_block) as u64;
let all_encrypted = self.inner.read(encrypted_start..encrypted_end).await?;
let result_len = (range.end - range.start) as usize;
let mut result = BytesMut::with_capacity(result_len);
let mut encrypted_offset = 0usize;
for block_idx in first_block..=last_block {
let block_size = self.cipher_block_size(block_idx) as usize;
let cipher_block = &all_encrypted[encrypted_offset..encrypted_offset + block_size];
encrypted_offset += block_size;
let aad = stream_block_aad(&self.aad_prefix, block_idx as u32);
let decrypted = self.cipher.decrypt(cipher_block, Some(&aad))?;
let block_plain_start = block_idx * PLAIN_BLOCK_SIZE as u64;
let slice_start = if block_idx == first_block {
(range.start - block_plain_start) as usize
} else {
0
};
let slice_end = if block_idx == last_block {
(range.end - block_plain_start) as usize
} else {
decrypted.len()
};
result.extend_from_slice(&decrypted[slice_start..slice_end]);
}
Ok(result.freeze())
}
}
pub struct AesGcmFileWrite {
inner: Box<dyn FileWrite>,
cipher: Arc<AesGcmCipher>,
aad_prefix: Box<[u8]>,
buffer: Vec<u8>,
block_index: u32,
header_written: bool,
closed: bool,
poisoned: bool,
}
impl AesGcmFileWrite {
pub fn new(
inner: Box<dyn FileWrite>,
cipher: Arc<AesGcmCipher>,
aad_prefix: impl Into<Box<[u8]>>,
) -> Self {
Self {
inner,
cipher,
aad_prefix: aad_prefix.into(),
buffer: Vec::new(),
block_index: 0,
header_written: false,
closed: false,
poisoned: false,
}
}
async fn write_header(&mut self) -> Result<()> {
let mut header = Vec::with_capacity(GCM_STREAM_HEADER_LENGTH as usize);
header.extend_from_slice(&GCM_STREAM_MAGIC);
header.extend_from_slice(&PLAIN_BLOCK_SIZE.to_le_bytes());
if let Err(e) = self.inner.write(Bytes::from(header)).await {
self.poisoned = true;
return Err(e);
}
self.header_written = true;
Ok(())
}
async fn encrypt_and_write_block(&mut self, block_data: &[u8]) -> Result<()> {
let aad = stream_block_aad(&self.aad_prefix, self.block_index);
let encrypted = self.cipher.encrypt(block_data, Some(&aad))?;
if let Err(e) = self.inner.write(Bytes::from(encrypted)).await {
self.poisoned = true;
return Err(e);
}
self.block_index = self.block_index.checked_add(1).ok_or_else(|| {
Error::new(
ErrorKind::DataInvalid,
"AGS1 block index overflow: file exceeds the maximum supported size (~4 TiB)",
)
})?;
Ok(())
}
async fn encrypt_and_drain_block(&mut self) -> Result<()> {
let aad = stream_block_aad(&self.aad_prefix, self.block_index);
let encrypted = self
.cipher
.encrypt(&self.buffer[..PLAIN_BLOCK_SIZE as usize], Some(&aad))?;
if let Err(e) = self.inner.write(Bytes::from(encrypted)).await {
self.poisoned = true;
return Err(e);
}
self.block_index = self.block_index.checked_add(1).ok_or_else(|| {
Error::new(
ErrorKind::DataInvalid,
"AGS1 block index overflow: file exceeds the maximum supported size (~4 TiB)",
)
})?;
self.buffer.drain(..PLAIN_BLOCK_SIZE as usize);
Ok(())
}
}
#[async_trait::async_trait]
impl FileWrite for AesGcmFileWrite {
async fn write(&mut self, bs: Bytes) -> Result<()> {
if self.closed {
return Err(Error::new(
ErrorKind::Unexpected,
"Cannot write to a closed AesGcmFileWrite",
));
}
if self.poisoned {
return Err(Error::new(
ErrorKind::Unexpected,
"AesGcmFileWrite is in a poisoned state due to a previous write failure",
));
}
if !self.header_written {
self.write_header().await?;
}
self.buffer.extend_from_slice(&bs);
while self.buffer.len() >= PLAIN_BLOCK_SIZE as usize {
self.encrypt_and_drain_block().await?;
}
Ok(())
}
async fn close(&mut self) -> Result<()> {
if self.closed {
return Err(Error::new(
ErrorKind::Unexpected,
"AesGcmFileWrite already closed",
));
}
if self.poisoned {
return Err(Error::new(
ErrorKind::Unexpected,
"AesGcmFileWrite is in a poisoned state due to a previous write failure",
));
}
if !self.header_written {
self.write_header().await?;
}
if !self.buffer.is_empty() || self.block_index == 0 {
let final_block = std::mem::take(&mut self.buffer);
self.encrypt_and_write_block(&final_block).await?;
}
self.closed = true;
self.inner.close().await
}
}
#[cfg(test)]
mod tests {
use super::*;
fn encrypt_ags1(plaintext: &[u8], cipher: &AesGcmCipher, aad_prefix: &[u8]) -> Vec<u8> {
let mut result = Vec::new();
result.extend_from_slice(&GCM_STREAM_MAGIC);
result.extend_from_slice(&PLAIN_BLOCK_SIZE.to_le_bytes());
let mut offset = 0;
let mut block_index = 0u32;
loop {
let remaining = plaintext.len() - offset;
let block_size = std::cmp::min(remaining, PLAIN_BLOCK_SIZE as usize);
if block_size == 0 && block_index > 0 {
break;
}
let block_data = &plaintext[offset..offset + block_size];
let aad = stream_block_aad(aad_prefix, block_index);
let encrypted = cipher.encrypt(block_data, Some(&aad)).unwrap();
result.extend_from_slice(&encrypted);
offset += block_size;
block_index += 1;
if block_size < PLAIN_BLOCK_SIZE as usize {
break;
}
}
result
}
fn make_cipher(key: &[u8]) -> AesGcmCipher {
use super::super::SecureKey;
let secure_key = SecureKey::new(key).unwrap();
AesGcmCipher::new(secure_key)
}
fn memory_reader(data: Vec<u8>) -> Box<dyn FileRead> {
Box::new(MemoryFileRead(Bytes::from(data)))
}
struct MemoryFileRead(Bytes);
#[async_trait::async_trait]
impl FileRead for MemoryFileRead {
async fn read(&self, range: Range<u64>) -> Result<Bytes> {
let start = range.start as usize;
let end = range.end as usize;
if end > self.0.len() {
return Err(Error::new(
ErrorKind::DataInvalid,
format!(
"Range {}..{} out of bounds for {} bytes",
start,
end,
self.0.len()
),
));
}
Ok(self.0.slice(start..end))
}
}
#[tokio::test]
async fn test_empty_file_roundtrip() {
let key = b"0123456789abcdef";
let aad_prefix = b"test-aad-prefix!";
let cipher = make_cipher(key);
let encrypted = encrypt_ags1(b"", &cipher, aad_prefix);
assert_eq!(encrypted.len(), MIN_STREAM_LENGTH as usize);
let reader = AesGcmFileRead::new(
memory_reader(encrypted.clone()),
Arc::new(make_cipher(key)),
aad_prefix.as_slice().into(),
encrypted.len() as u64,
)
.unwrap();
assert_eq!(reader.plaintext_length(), 0);
let result = reader.read(0..0).await.unwrap();
assert!(result.is_empty());
}
#[tokio::test]
async fn test_small_file_roundtrip() {
let key = b"0123456789abcdef";
let aad_prefix = b"test-aad-prefix!";
let plaintext = b"Hello, Iceberg encryption!";
let cipher = make_cipher(key);
let encrypted = encrypt_ags1(plaintext, &cipher, aad_prefix);
let reader = AesGcmFileRead::new(
memory_reader(encrypted.clone()),
Arc::new(make_cipher(key)),
aad_prefix.as_slice().into(),
encrypted.len() as u64,
)
.unwrap();
assert_eq!(reader.plaintext_length(), plaintext.len() as u64);
let result = reader.read(0..plaintext.len() as u64).await.unwrap();
assert_eq!(&result[..], plaintext);
}
#[tokio::test]
async fn test_partial_read() {
let key = b"0123456789abcdef";
let aad_prefix = b"aad-prefix-here!";
let plaintext = b"ABCDEFGHIJKLMNOPQRSTUVWXYZ";
let cipher = make_cipher(key);
let encrypted = encrypt_ags1(plaintext, &cipher, aad_prefix);
let reader = AesGcmFileRead::new(
memory_reader(encrypted.clone()),
Arc::new(make_cipher(key)),
aad_prefix.as_slice().into(),
encrypted.len() as u64,
)
.unwrap();
let result = reader.read(10..20).await.unwrap();
assert_eq!(&result[..], &plaintext[10..20]);
let result = reader.read(0..1).await.unwrap();
assert_eq!(&result[..], &plaintext[0..1]);
let last = plaintext.len() as u64;
let result = reader.read(last - 1..last).await.unwrap();
assert_eq!(&result[..], &plaintext[plaintext.len() - 1..]);
}
#[tokio::test]
async fn test_multi_block_roundtrip() {
let key = b"0123456789abcdef";
let aad_prefix = b"multi-block-aad!";
let size = PLAIN_BLOCK_SIZE as usize + PLAIN_BLOCK_SIZE as usize / 2;
let plaintext: Vec<u8> = (0..size).map(|i| (i % 256) as u8).collect();
let cipher = make_cipher(key);
let encrypted = encrypt_ags1(&plaintext, &cipher, aad_prefix);
let reader = AesGcmFileRead::new(
memory_reader(encrypted.clone()),
Arc::new(make_cipher(key)),
aad_prefix.as_slice().into(),
encrypted.len() as u64,
)
.unwrap();
assert_eq!(reader.plaintext_length(), plaintext.len() as u64);
let result = reader.read(0..plaintext.len() as u64).await.unwrap();
assert_eq!(&result[..], &plaintext[..]);
}
#[tokio::test]
async fn test_cross_block_read() {
let key = b"0123456789abcdef";
let aad_prefix = b"cross-block-aad!";
let size = PLAIN_BLOCK_SIZE as usize * 2 + PLAIN_BLOCK_SIZE as usize / 2;
let plaintext: Vec<u8> = (0..size).map(|i| (i % 256) as u8).collect();
let cipher = make_cipher(key);
let encrypted = encrypt_ags1(&plaintext, &cipher, aad_prefix);
let reader = AesGcmFileRead::new(
memory_reader(encrypted.clone()),
Arc::new(make_cipher(key)),
aad_prefix.as_slice().into(),
encrypted.len() as u64,
)
.unwrap();
let boundary = PLAIN_BLOCK_SIZE as u64;
let result = reader.read(boundary - 100..boundary + 100).await.unwrap();
assert_eq!(
&result[..],
&plaintext[(boundary - 100) as usize..(boundary + 100) as usize]
);
let result = reader.read(boundary - 50..boundary * 2 + 50).await.unwrap();
assert_eq!(
&result[..],
&plaintext[(boundary - 50) as usize..(boundary * 2 + 50) as usize]
);
}
#[tokio::test]
async fn test_exact_block_size() {
let key = b"0123456789abcdef";
let aad_prefix = b"exact-block-aad!";
let plaintext: Vec<u8> = (0..PLAIN_BLOCK_SIZE as usize)
.map(|i| (i % 256) as u8)
.collect();
let cipher = make_cipher(key);
let encrypted = encrypt_ags1(&plaintext, &cipher, aad_prefix);
let reader = AesGcmFileRead::new(
memory_reader(encrypted.clone()),
Arc::new(make_cipher(key)),
aad_prefix.as_slice().into(),
encrypted.len() as u64,
)
.unwrap();
assert_eq!(reader.plaintext_length(), PLAIN_BLOCK_SIZE as u64);
let result = reader.read(0..PLAIN_BLOCK_SIZE as u64).await.unwrap();
assert_eq!(&result[..], &plaintext[..]);
}
#[tokio::test]
async fn test_block_size_plus_one() {
let key = b"0123456789abcdef";
let aad_prefix = b"block-plus-one!!";
let size = PLAIN_BLOCK_SIZE as usize + 1;
let plaintext: Vec<u8> = (0..size).map(|i| (i % 256) as u8).collect();
let cipher = make_cipher(key);
let encrypted = encrypt_ags1(&plaintext, &cipher, aad_prefix);
let reader = AesGcmFileRead::new(
memory_reader(encrypted.clone()),
Arc::new(make_cipher(key)),
aad_prefix.as_slice().into(),
encrypted.len() as u64,
)
.unwrap();
assert_eq!(reader.plaintext_length(), size as u64);
let result = reader.read(size as u64 - 1..size as u64).await.unwrap();
assert_eq!(result[0], plaintext[size - 1]);
let result = reader.read(0..size as u64).await.unwrap();
assert_eq!(&result[..], &plaintext[..]);
}
#[tokio::test]
async fn test_block_size_minus_one() {
let key = b"0123456789abcdef";
let aad_prefix = b"block-minus-one!";
let size = PLAIN_BLOCK_SIZE as usize - 1;
let plaintext: Vec<u8> = (0..size).map(|i| (i % 256) as u8).collect();
let cipher = make_cipher(key);
let encrypted = encrypt_ags1(&plaintext, &cipher, aad_prefix);
let reader = AesGcmFileRead::new(
memory_reader(encrypted.clone()),
Arc::new(make_cipher(key)),
aad_prefix.as_slice().into(),
encrypted.len() as u64,
)
.unwrap();
assert_eq!(reader.plaintext_length(), size as u64);
let result = reader.read(0..size as u64).await.unwrap();
assert_eq!(&result[..], &plaintext[..]);
}
#[tokio::test]
async fn test_wrong_aad_fails() {
let key = b"0123456789abcdef";
let aad_prefix = b"correct-aad-here";
let plaintext = b"sensitive data here";
let cipher = make_cipher(key);
let encrypted = encrypt_ags1(plaintext, &cipher, aad_prefix);
let mut bad_aad = aad_prefix.to_vec();
bad_aad[0] ^= 0xFF;
let reader = AesGcmFileRead::new(
memory_reader(encrypted.clone()),
Arc::new(make_cipher(key)),
bad_aad.as_slice().into(),
encrypted.len() as u64,
)
.unwrap();
let result = reader.read(0..plaintext.len() as u64).await;
assert!(result.is_err(), "Decryption with wrong AAD should fail");
}
#[tokio::test]
async fn test_wrong_key_fails() {
let key = b"0123456789abcdef";
let wrong_key = b"fedcba9876543210";
let aad_prefix = b"test-aad-prefix!";
let plaintext = b"sensitive data";
let cipher = make_cipher(key);
let encrypted = encrypt_ags1(plaintext, &cipher, aad_prefix);
let reader = AesGcmFileRead::new(
memory_reader(encrypted.clone()),
Arc::new(make_cipher(wrong_key)),
aad_prefix.as_slice().into(),
encrypted.len() as u64,
)
.unwrap();
let result = reader.read(0..plaintext.len() as u64).await;
assert!(result.is_err(), "Decryption with wrong key should fail");
}
#[tokio::test]
async fn test_out_of_bounds_read() {
let key = b"0123456789abcdef";
let aad_prefix = b"test-aad-prefix!";
let plaintext = b"short data";
let cipher = make_cipher(key);
let encrypted = encrypt_ags1(plaintext, &cipher, aad_prefix);
let reader = AesGcmFileRead::new(
memory_reader(encrypted.clone()),
Arc::new(make_cipher(key)),
aad_prefix.as_slice().into(),
encrypted.len() as u64,
)
.unwrap();
let result = reader.read(0..plaintext.len() as u64 + 1).await;
assert!(result.is_err(), "Reading past end should fail");
}
#[tokio::test]
async fn test_calculate_plaintext_length() {
assert_eq!(
AesGcmFileRead::calculate_plaintext_length(GCM_STREAM_HEADER_LENGTH as u64).unwrap(),
0
);
assert_eq!(
AesGcmFileRead::calculate_plaintext_length(MIN_STREAM_LENGTH as u64).unwrap(),
0
);
let one_full = GCM_STREAM_HEADER_LENGTH as u64 + CIPHER_BLOCK_SIZE as u64;
assert_eq!(
AesGcmFileRead::calculate_plaintext_length(one_full).unwrap(),
PLAIN_BLOCK_SIZE as u64
);
let one_full_plus_one = one_full + NONCE_LENGTH as u64 + 1 + GCM_TAG_LENGTH as u64;
assert_eq!(
AesGcmFileRead::calculate_plaintext_length(one_full_plus_one).unwrap(),
PLAIN_BLOCK_SIZE as u64 + 1
);
}
#[tokio::test]
async fn test_stream_block_aad() {
let aad = stream_block_aad(b"prefix", 0);
assert_eq!(&aad[..6], b"prefix");
assert_eq!(&aad[6..], &0u32.to_le_bytes());
let aad = stream_block_aad(b"prefix", 1);
assert_eq!(&aad[..6], b"prefix");
assert_eq!(&aad[6..], &1u32.to_le_bytes());
let aad = stream_block_aad(b"", 42);
assert_eq!(&aad[..], &42u32.to_le_bytes());
}
#[tokio::test]
async fn test_encrypted_file_too_short() {
let result = AesGcmFileRead::new(
memory_reader(vec![0; 4]),
Arc::new(make_cipher(b"0123456789abcdef")),
[].into(),
4,
);
assert!(result.is_err());
}
struct SharedMemoryWrite {
buffer: std::sync::Arc<std::sync::Mutex<Vec<u8>>>,
}
struct FailingFileWrite {
writes_before_failure: usize,
write_count: usize,
}
#[async_trait::async_trait]
impl FileWrite for FailingFileWrite {
async fn write(&mut self, _bs: Bytes) -> Result<()> {
if self.write_count >= self.writes_before_failure {
return Err(Error::new(ErrorKind::Unexpected, "simulated write failure"));
}
self.write_count += 1;
Ok(())
}
async fn close(&mut self) -> Result<()> {
Ok(())
}
}
#[async_trait::async_trait]
impl FileWrite for SharedMemoryWrite {
async fn write(&mut self, bs: Bytes) -> Result<()> {
self.buffer.lock().unwrap().extend_from_slice(&bs);
Ok(())
}
async fn close(&mut self) -> Result<()> {
Ok(())
}
}
async fn write_through_ags1(plaintext: &[u8], key: &[u8], aad_prefix: &[u8]) -> Vec<u8> {
let buffer = std::sync::Arc::new(std::sync::Mutex::new(Vec::new()));
let inner: Box<dyn FileWrite> = Box::new(SharedMemoryWrite {
buffer: buffer.clone(),
});
let cipher = Arc::new(make_cipher(key));
let mut writer = AesGcmFileWrite::new(inner, cipher, aad_prefix.to_vec());
writer.write(Bytes::from(plaintext.to_vec())).await.unwrap();
writer.close().await.unwrap();
buffer.lock().unwrap().clone()
}
#[tokio::test]
async fn test_write_empty_roundtrip() {
let key = b"0123456789abcdef";
let aad_prefix = b"test-aad-prefix!";
let encrypted = write_through_ags1(b"", key, aad_prefix).await;
assert_eq!(encrypted.len(), MIN_STREAM_LENGTH as usize);
let reader = AesGcmFileRead::new(
memory_reader(encrypted.clone()),
Arc::new(make_cipher(key)),
aad_prefix.as_slice().into(),
encrypted.len() as u64,
)
.unwrap();
assert_eq!(reader.plaintext_length(), 0);
}
#[tokio::test]
async fn test_write_small_roundtrip() {
let key = b"0123456789abcdef";
let aad_prefix = b"test-aad-prefix!";
let plaintext = b"Hello, Iceberg encryption!";
let encrypted = write_through_ags1(plaintext, key, aad_prefix).await;
let reader = AesGcmFileRead::new(
memory_reader(encrypted.clone()),
Arc::new(make_cipher(key)),
aad_prefix.as_slice().into(),
encrypted.len() as u64,
)
.unwrap();
assert_eq!(reader.plaintext_length(), plaintext.len() as u64);
let result = reader.read(0..plaintext.len() as u64).await.unwrap();
assert_eq!(&result[..], plaintext);
}
#[tokio::test]
async fn test_write_multi_block_roundtrip() {
let key = b"0123456789abcdef";
let aad_prefix = b"multi-block-aad!";
let size = PLAIN_BLOCK_SIZE as usize + PLAIN_BLOCK_SIZE as usize / 2;
let plaintext: Vec<u8> = (0..size).map(|i| (i % 256) as u8).collect();
let encrypted = write_through_ags1(&plaintext, key, aad_prefix).await;
let reader = AesGcmFileRead::new(
memory_reader(encrypted.clone()),
Arc::new(make_cipher(key)),
aad_prefix.as_slice().into(),
encrypted.len() as u64,
)
.unwrap();
assert_eq!(reader.plaintext_length(), plaintext.len() as u64);
let result = reader.read(0..plaintext.len() as u64).await.unwrap();
assert_eq!(&result[..], &plaintext[..]);
}
#[tokio::test]
async fn test_write_cross_block_accumulation() {
let key = b"0123456789abcdef";
let aad_prefix = b"cross-block-aad!";
let buffer = std::sync::Arc::new(std::sync::Mutex::new(Vec::new()));
let inner: Box<dyn FileWrite> = Box::new(SharedMemoryWrite {
buffer: buffer.clone(),
});
let cipher = Arc::new(make_cipher(key));
let mut writer = AesGcmFileWrite::new(inner, cipher, aad_prefix.to_vec());
let total_size = PLAIN_BLOCK_SIZE as usize + PLAIN_BLOCK_SIZE as usize / 2;
let plaintext: Vec<u8> = (0..total_size).map(|i| (i % 256) as u8).collect();
let chunk_size = 1000;
for chunk in plaintext.chunks(chunk_size) {
writer.write(Bytes::from(chunk.to_vec())).await.unwrap();
}
writer.close().await.unwrap();
let encrypted = buffer.lock().unwrap().clone();
let reader = AesGcmFileRead::new(
memory_reader(encrypted.clone()),
Arc::new(make_cipher(key)),
aad_prefix.as_slice().into(),
encrypted.len() as u64,
)
.unwrap();
assert_eq!(reader.plaintext_length(), plaintext.len() as u64);
let result = reader.read(0..plaintext.len() as u64).await.unwrap();
assert_eq!(&result[..], &plaintext[..]);
}
#[tokio::test]
async fn test_write_exact_block_size() {
let key = b"0123456789abcdef";
let aad_prefix = b"exact-block-aad!";
let plaintext: Vec<u8> = (0..PLAIN_BLOCK_SIZE as usize)
.map(|i| (i % 256) as u8)
.collect();
let encrypted = write_through_ags1(&plaintext, key, aad_prefix).await;
let reader = AesGcmFileRead::new(
memory_reader(encrypted.clone()),
Arc::new(make_cipher(key)),
aad_prefix.as_slice().into(),
encrypted.len() as u64,
)
.unwrap();
assert_eq!(reader.plaintext_length(), PLAIN_BLOCK_SIZE as u64);
let result = reader.read(0..PLAIN_BLOCK_SIZE as u64).await.unwrap();
assert_eq!(&result[..], &plaintext[..]);
}
#[tokio::test]
async fn test_write_block_aligned_no_spurious_empty_block() {
let key = b"0123456789abcdef";
let aad_prefix = b"block-align-aad!";
let plaintext: Vec<u8> = (0..PLAIN_BLOCK_SIZE as usize)
.map(|i| (i % 256) as u8)
.collect();
let encrypted_via_writer = write_through_ags1(&plaintext, key, aad_prefix).await;
let encrypted_via_reference = encrypt_ags1(&plaintext, &make_cipher(key), aad_prefix);
assert_eq!(
encrypted_via_writer.len(),
encrypted_via_reference.len(),
"Writer output should match reference encryption length (no spurious trailing block)"
);
let reader = AesGcmFileRead::new(
memory_reader(encrypted_via_writer.clone()),
Arc::new(make_cipher(key)),
aad_prefix.as_slice().into(),
encrypted_via_writer.len() as u64,
)
.unwrap();
assert_eq!(reader.plaintext_length(), PLAIN_BLOCK_SIZE as u64);
let result = reader.read(0..PLAIN_BLOCK_SIZE as u64).await.unwrap();
assert_eq!(&result[..], &plaintext[..]);
}
#[tokio::test]
async fn test_write_two_blocks_aligned_no_spurious_empty_block() {
let key = b"0123456789abcdef";
let aad_prefix = b"2blk-align-aad!!";
let size = PLAIN_BLOCK_SIZE as usize * 2;
let plaintext: Vec<u8> = (0..size).map(|i| (i % 256) as u8).collect();
let encrypted_via_writer = write_through_ags1(&plaintext, key, aad_prefix).await;
let encrypted_via_reference = encrypt_ags1(&plaintext, &make_cipher(key), aad_prefix);
assert_eq!(
encrypted_via_writer.len(),
encrypted_via_reference.len(),
"Writer output should match reference encryption length (no spurious trailing block)"
);
let reader = AesGcmFileRead::new(
memory_reader(encrypted_via_writer.clone()),
Arc::new(make_cipher(key)),
aad_prefix.as_slice().into(),
encrypted_via_writer.len() as u64,
)
.unwrap();
assert_eq!(reader.plaintext_length(), size as u64);
let result = reader.read(0..size as u64).await.unwrap();
assert_eq!(&result[..], &plaintext[..]);
}
#[tokio::test]
async fn test_write_poisoned_after_inner_write_failure() {
let cipher = Arc::new(make_cipher(b"0123456789abcdef"));
let inner: Box<dyn FileWrite> = Box::new(FailingFileWrite {
writes_before_failure: 1,
write_count: 0,
});
let mut writer = AesGcmFileWrite::new(inner, cipher, b"aad-prefix-here!".to_vec());
let data = vec![0u8; PLAIN_BLOCK_SIZE as usize];
let result = writer.write(Bytes::from(data)).await;
assert!(result.is_err());
let result = writer.write(Bytes::from(b"more data".to_vec())).await;
assert!(result.is_err());
assert!(
result.unwrap_err().to_string().contains("poisoned"),
"expected poisoned error"
);
let result = writer.close().await;
assert!(result.is_err());
assert!(
result.unwrap_err().to_string().contains("poisoned"),
"expected poisoned error on close"
);
}
}