mirror of
https://github.com/rustfs/rustfs.git
synced 2026-08-08 06:13:14 +00:00
feat(encryption): add managed encryption support for SSE-S3 and SSE-KMS (#583)
Signed-off-by: junxiang Mu <1948535941@qq.com>
This commit is contained in:
@@ -0,0 +1,386 @@
|
||||
// Copyright 2024 RustFS Team
|
||||
//
|
||||
// Licensed under the Apache License, Version 2.0 (the "License");
|
||||
// you may not use this file except in compliance with the License.
|
||||
// You may obtain a copy of the License at
|
||||
//
|
||||
// http://www.apache.org/licenses/LICENSE-2.0
|
||||
//
|
||||
// Unless required by applicable law or agreed to in writing, software
|
||||
// distributed under the License is distributed on an "AS IS" BASIS,
|
||||
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
// See the License for the specific language governing permissions and
|
||||
// limitations under the License.
|
||||
|
||||
//! Integration tests that focus on surface headers/metadata emitted by the
|
||||
//! managed encryption pipeline (SSE-S3/SSE-KMS).
|
||||
|
||||
use super::common::LocalKMSTestEnvironment;
|
||||
use crate::common::{TEST_BUCKET, init_logging};
|
||||
use aws_sdk_s3::primitives::ByteStream;
|
||||
use aws_sdk_s3::types::{
|
||||
CompletedMultipartUpload, CompletedPart, ServerSideEncryption, ServerSideEncryptionByDefault,
|
||||
ServerSideEncryptionConfiguration, ServerSideEncryptionRule,
|
||||
};
|
||||
use serial_test::serial;
|
||||
use std::collections::{HashMap, VecDeque};
|
||||
use tracing::info;
|
||||
|
||||
fn assert_encryption_metadata(metadata: &HashMap<String, String>, expected_size: usize) {
|
||||
for key in [
|
||||
"x-rustfs-encryption-key",
|
||||
"x-rustfs-encryption-iv",
|
||||
"x-rustfs-encryption-context",
|
||||
"x-rustfs-encryption-original-size",
|
||||
] {
|
||||
assert!(metadata.contains_key(key), "expected managed encryption metadata '{}' to be present", key);
|
||||
assert!(
|
||||
!metadata.get(key).unwrap().is_empty(),
|
||||
"managed encryption metadata '{}' should not be empty",
|
||||
key
|
||||
);
|
||||
}
|
||||
|
||||
let size_value = metadata
|
||||
.get("x-rustfs-encryption-original-size")
|
||||
.expect("managed encryption metadata should include original size");
|
||||
let parsed_size: usize = size_value
|
||||
.parse()
|
||||
.expect("x-rustfs-encryption-original-size should be numeric");
|
||||
assert_eq!(parsed_size, expected_size, "recorded original size should match uploaded payload length");
|
||||
}
|
||||
|
||||
fn assert_storage_encrypted(storage_root: &std::path::Path, bucket: &str, key: &str, plaintext: &[u8]) {
|
||||
let mut stack = VecDeque::from([storage_root.to_path_buf()]);
|
||||
let mut scanned = 0;
|
||||
let mut plaintext_path: Option<std::path::PathBuf> = None;
|
||||
|
||||
while let Some(current) = stack.pop_front() {
|
||||
let Ok(metadata) = std::fs::metadata(¤t) else { continue };
|
||||
if metadata.is_dir() {
|
||||
if let Ok(entries) = std::fs::read_dir(¤t) {
|
||||
for entry in entries.flatten() {
|
||||
stack.push_back(entry.path());
|
||||
}
|
||||
}
|
||||
continue;
|
||||
}
|
||||
|
||||
let path_str = current.to_string_lossy();
|
||||
if !(path_str.contains(bucket) || path_str.contains(key)) {
|
||||
continue;
|
||||
}
|
||||
|
||||
scanned += 1;
|
||||
let Ok(bytes) = std::fs::read(¤t) else { continue };
|
||||
if bytes.len() < plaintext.len() {
|
||||
continue;
|
||||
}
|
||||
if bytes.windows(plaintext.len()).any(|window| window == plaintext) {
|
||||
plaintext_path = Some(current);
|
||||
break;
|
||||
}
|
||||
}
|
||||
|
||||
assert!(
|
||||
scanned > 0,
|
||||
"Failed to locate stored data files for bucket '{}' and key '{}' under {:?}",
|
||||
bucket,
|
||||
key,
|
||||
storage_root
|
||||
);
|
||||
assert!(plaintext_path.is_none(), "Plaintext detected on disk at {:?}", plaintext_path.unwrap());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
#[serial]
|
||||
async fn test_head_reports_managed_metadata_for_sse_s3() -> Result<(), Box<dyn std::error::Error + Send + Sync>> {
|
||||
init_logging();
|
||||
info!("Validating SSE-S3 managed encryption metadata exposure");
|
||||
|
||||
let mut kms_env = LocalKMSTestEnvironment::new().await?;
|
||||
let _default_key = kms_env.start_rustfs_for_local_kms().await?;
|
||||
tokio::time::sleep(tokio::time::Duration::from_secs(3)).await;
|
||||
|
||||
let s3_client = kms_env.base_env.create_s3_client();
|
||||
kms_env.base_env.create_test_bucket(TEST_BUCKET).await?;
|
||||
|
||||
// Bucket level default SSE-S3 configuration.
|
||||
let encryption_config = ServerSideEncryptionConfiguration::builder()
|
||||
.rules(
|
||||
ServerSideEncryptionRule::builder()
|
||||
.apply_server_side_encryption_by_default(
|
||||
ServerSideEncryptionByDefault::builder()
|
||||
.sse_algorithm(ServerSideEncryption::Aes256)
|
||||
.build()
|
||||
.unwrap(),
|
||||
)
|
||||
.build(),
|
||||
)
|
||||
.build()
|
||||
.unwrap();
|
||||
|
||||
s3_client
|
||||
.put_bucket_encryption()
|
||||
.bucket(TEST_BUCKET)
|
||||
.server_side_encryption_configuration(encryption_config)
|
||||
.send()
|
||||
.await?;
|
||||
|
||||
let payload = b"metadata-sse-s3-payload";
|
||||
let key = "metadata-sse-s3-object";
|
||||
|
||||
s3_client
|
||||
.put_object()
|
||||
.bucket(TEST_BUCKET)
|
||||
.key(key)
|
||||
.body(payload.to_vec().into())
|
||||
.send()
|
||||
.await?;
|
||||
|
||||
let head = s3_client.head_object().bucket(TEST_BUCKET).key(key).send().await?;
|
||||
|
||||
assert_eq!(
|
||||
head.server_side_encryption(),
|
||||
Some(&ServerSideEncryption::Aes256),
|
||||
"head_object should advertise SSE-S3"
|
||||
);
|
||||
|
||||
let metadata = head
|
||||
.metadata()
|
||||
.expect("head_object should return managed encryption metadata");
|
||||
assert_encryption_metadata(metadata, payload.len());
|
||||
|
||||
assert_storage_encrypted(std::path::Path::new(&kms_env.base_env.temp_dir), TEST_BUCKET, key, payload);
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
#[serial]
|
||||
async fn test_head_reports_managed_metadata_for_sse_kms_and_copy() -> Result<(), Box<dyn std::error::Error + Send + Sync>> {
|
||||
init_logging();
|
||||
info!("Validating SSE-KMS managed encryption metadata (including copy)");
|
||||
|
||||
let mut kms_env = LocalKMSTestEnvironment::new().await?;
|
||||
let default_key_id = kms_env.start_rustfs_for_local_kms().await?;
|
||||
tokio::time::sleep(tokio::time::Duration::from_secs(3)).await;
|
||||
|
||||
let s3_client = kms_env.base_env.create_s3_client();
|
||||
kms_env.base_env.create_test_bucket(TEST_BUCKET).await?;
|
||||
|
||||
let encryption_config = ServerSideEncryptionConfiguration::builder()
|
||||
.rules(
|
||||
ServerSideEncryptionRule::builder()
|
||||
.apply_server_side_encryption_by_default(
|
||||
ServerSideEncryptionByDefault::builder()
|
||||
.sse_algorithm(ServerSideEncryption::AwsKms)
|
||||
.kms_master_key_id(&default_key_id)
|
||||
.build()
|
||||
.unwrap(),
|
||||
)
|
||||
.build(),
|
||||
)
|
||||
.build()
|
||||
.unwrap();
|
||||
|
||||
s3_client
|
||||
.put_bucket_encryption()
|
||||
.bucket(TEST_BUCKET)
|
||||
.server_side_encryption_configuration(encryption_config)
|
||||
.send()
|
||||
.await?;
|
||||
|
||||
let payload = b"metadata-sse-kms-payload";
|
||||
let source_key = "metadata-sse-kms-object";
|
||||
|
||||
s3_client
|
||||
.put_object()
|
||||
.bucket(TEST_BUCKET)
|
||||
.key(source_key)
|
||||
.body(payload.to_vec().into())
|
||||
.send()
|
||||
.await?;
|
||||
|
||||
let head_source = s3_client.head_object().bucket(TEST_BUCKET).key(source_key).send().await?;
|
||||
|
||||
assert_eq!(
|
||||
head_source.server_side_encryption(),
|
||||
Some(&ServerSideEncryption::AwsKms),
|
||||
"source object should report SSE-KMS"
|
||||
);
|
||||
assert_eq!(
|
||||
head_source.ssekms_key_id().unwrap(),
|
||||
&default_key_id,
|
||||
"source object should maintain the configured KMS key id"
|
||||
);
|
||||
let source_metadata = head_source
|
||||
.metadata()
|
||||
.expect("source object should include managed encryption metadata");
|
||||
assert_encryption_metadata(source_metadata, payload.len());
|
||||
|
||||
let dest_key = "metadata-sse-kms-object-copy";
|
||||
let copy_source = format!("{}/{}", TEST_BUCKET, source_key);
|
||||
|
||||
s3_client
|
||||
.copy_object()
|
||||
.bucket(TEST_BUCKET)
|
||||
.key(dest_key)
|
||||
.copy_source(copy_source)
|
||||
.send()
|
||||
.await?;
|
||||
|
||||
let head_dest = s3_client.head_object().bucket(TEST_BUCKET).key(dest_key).send().await?;
|
||||
|
||||
assert_eq!(
|
||||
head_dest.server_side_encryption(),
|
||||
Some(&ServerSideEncryption::AwsKms),
|
||||
"copied object should remain encrypted with SSE-KMS"
|
||||
);
|
||||
assert_eq!(
|
||||
head_dest.ssekms_key_id().unwrap(),
|
||||
&default_key_id,
|
||||
"copied object should keep the default KMS key id"
|
||||
);
|
||||
let dest_metadata = head_dest
|
||||
.metadata()
|
||||
.expect("copied object should include managed encryption metadata");
|
||||
assert_encryption_metadata(dest_metadata, payload.len());
|
||||
|
||||
let copied_body = s3_client
|
||||
.get_object()
|
||||
.bucket(TEST_BUCKET)
|
||||
.key(dest_key)
|
||||
.send()
|
||||
.await?
|
||||
.body
|
||||
.collect()
|
||||
.await?
|
||||
.into_bytes();
|
||||
assert_eq!(&copied_body[..], payload, "copied object payload should match source");
|
||||
|
||||
let storage_root = std::path::Path::new(&kms_env.base_env.temp_dir);
|
||||
assert_storage_encrypted(storage_root, TEST_BUCKET, source_key, payload);
|
||||
assert_storage_encrypted(storage_root, TEST_BUCKET, dest_key, payload);
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
#[serial]
|
||||
async fn test_multipart_upload_writes_encrypted_data() -> Result<(), Box<dyn std::error::Error + Send + Sync>> {
|
||||
init_logging();
|
||||
info!("Validating ciphertext persistence for multipart SSE-KMS uploads");
|
||||
|
||||
let mut kms_env = LocalKMSTestEnvironment::new().await?;
|
||||
let default_key_id = kms_env.start_rustfs_for_local_kms().await?;
|
||||
tokio::time::sleep(tokio::time::Duration::from_secs(3)).await;
|
||||
|
||||
let s3_client = kms_env.base_env.create_s3_client();
|
||||
kms_env.base_env.create_test_bucket(TEST_BUCKET).await?;
|
||||
|
||||
let encryption_config = ServerSideEncryptionConfiguration::builder()
|
||||
.rules(
|
||||
ServerSideEncryptionRule::builder()
|
||||
.apply_server_side_encryption_by_default(
|
||||
ServerSideEncryptionByDefault::builder()
|
||||
.sse_algorithm(ServerSideEncryption::AwsKms)
|
||||
.kms_master_key_id(&default_key_id)
|
||||
.build()
|
||||
.unwrap(),
|
||||
)
|
||||
.build(),
|
||||
)
|
||||
.build()
|
||||
.unwrap();
|
||||
|
||||
s3_client
|
||||
.put_bucket_encryption()
|
||||
.bucket(TEST_BUCKET)
|
||||
.server_side_encryption_configuration(encryption_config)
|
||||
.send()
|
||||
.await?;
|
||||
|
||||
let key = "multipart-encryption-object";
|
||||
let part_size = 5 * 1024 * 1024; // minimum part size required by S3 semantics
|
||||
let part_one = vec![0xA5; part_size];
|
||||
let part_two = vec![0x5A; part_size];
|
||||
let combined: Vec<u8> = part_one.iter().chain(part_two.iter()).copied().collect();
|
||||
|
||||
let create_output = s3_client
|
||||
.create_multipart_upload()
|
||||
.bucket(TEST_BUCKET)
|
||||
.key(key)
|
||||
.send()
|
||||
.await?;
|
||||
|
||||
let upload_id = create_output.upload_id().unwrap();
|
||||
|
||||
let part1 = s3_client
|
||||
.upload_part()
|
||||
.bucket(TEST_BUCKET)
|
||||
.key(key)
|
||||
.upload_id(upload_id)
|
||||
.part_number(1)
|
||||
.body(ByteStream::from(part_one.clone()))
|
||||
.send()
|
||||
.await?;
|
||||
|
||||
let part2 = s3_client
|
||||
.upload_part()
|
||||
.bucket(TEST_BUCKET)
|
||||
.key(key)
|
||||
.upload_id(upload_id)
|
||||
.part_number(2)
|
||||
.body(ByteStream::from(part_two.clone()))
|
||||
.send()
|
||||
.await?;
|
||||
|
||||
let completed = CompletedMultipartUpload::builder()
|
||||
.parts(CompletedPart::builder().part_number(1).e_tag(part1.e_tag().unwrap()).build())
|
||||
.parts(CompletedPart::builder().part_number(2).e_tag(part2.e_tag().unwrap()).build())
|
||||
.build();
|
||||
|
||||
s3_client
|
||||
.complete_multipart_upload()
|
||||
.bucket(TEST_BUCKET)
|
||||
.key(key)
|
||||
.upload_id(upload_id)
|
||||
.multipart_upload(completed)
|
||||
.send()
|
||||
.await?;
|
||||
|
||||
let head = s3_client.head_object().bucket(TEST_BUCKET).key(key).send().await?;
|
||||
assert_eq!(
|
||||
head.server_side_encryption(),
|
||||
Some(&ServerSideEncryption::AwsKms),
|
||||
"multipart head_object should expose SSE-KMS"
|
||||
);
|
||||
assert_eq!(
|
||||
head.ssekms_key_id().unwrap(),
|
||||
&default_key_id,
|
||||
"multipart object should retain bucket default KMS key"
|
||||
);
|
||||
|
||||
assert_encryption_metadata(
|
||||
head.metadata().expect("multipart head_object should expose managed metadata"),
|
||||
combined.len(),
|
||||
);
|
||||
|
||||
// Data returned to clients should decrypt back to original payload
|
||||
let fetched = s3_client
|
||||
.get_object()
|
||||
.bucket(TEST_BUCKET)
|
||||
.key(key)
|
||||
.send()
|
||||
.await?
|
||||
.body
|
||||
.collect()
|
||||
.await?
|
||||
.into_bytes();
|
||||
assert_eq!(&fetched[..], &combined[..]);
|
||||
|
||||
assert_storage_encrypted(std::path::Path::new(&kms_env.base_env.temp_dir), TEST_BUCKET, key, &combined);
|
||||
|
||||
Ok(())
|
||||
}
|
||||
@@ -44,3 +44,6 @@ mod test_runner;
|
||||
|
||||
#[cfg(test)]
|
||||
mod bucket_default_encryption_test;
|
||||
|
||||
#[cfg(test)]
|
||||
mod encryption_metadata_test;
|
||||
|
||||
@@ -2143,6 +2143,7 @@ impl SetDisks {
|
||||
where
|
||||
W: AsyncWrite + Send + Sync + Unpin + 'static,
|
||||
{
|
||||
tracing::debug!(bucket, object, requested_length = length, offset, "get_object_with_fileinfo start");
|
||||
let (disks, files) = Self::shuffle_disks_and_parts_metadata_by_index(disks, &files, &fi);
|
||||
|
||||
let total_size = fi.size as usize;
|
||||
@@ -2160,28 +2161,46 @@ impl SetDisks {
|
||||
|
||||
let (part_index, mut part_offset) = fi.to_part_offset(offset)?;
|
||||
|
||||
// debug!(
|
||||
// "get_object_with_fileinfo start offset:{}, part_index:{},part_offset:{}",
|
||||
// offset, part_index, part_offset
|
||||
// );
|
||||
|
||||
let mut end_offset = offset;
|
||||
if length > 0 {
|
||||
end_offset += length - 1
|
||||
}
|
||||
|
||||
let (last_part_index, _) = fi.to_part_offset(end_offset)?;
|
||||
let (last_part_index, last_part_relative_offset) = fi.to_part_offset(end_offset)?;
|
||||
|
||||
tracing::debug!(
|
||||
bucket,
|
||||
object,
|
||||
offset,
|
||||
length,
|
||||
end_offset,
|
||||
part_index,
|
||||
last_part_index,
|
||||
last_part_relative_offset,
|
||||
"Multipart read bounds"
|
||||
);
|
||||
|
||||
let erasure = erasure_coding::Erasure::new(fi.erasure.data_blocks, fi.erasure.parity_blocks, fi.erasure.block_size);
|
||||
|
||||
let part_indices: Vec<usize> = (part_index..=last_part_index).collect();
|
||||
tracing::debug!(bucket, object, ?part_indices, "Multipart part indices to stream");
|
||||
|
||||
let mut total_read = 0;
|
||||
for i in part_index..=last_part_index {
|
||||
for current_part in part_indices {
|
||||
if total_read == length {
|
||||
tracing::debug!(
|
||||
bucket,
|
||||
object,
|
||||
total_read,
|
||||
requested_length = length,
|
||||
part_index = current_part,
|
||||
"Stopping multipart stream early because accumulated bytes match request"
|
||||
);
|
||||
break;
|
||||
}
|
||||
|
||||
let part_number = fi.parts[i].number;
|
||||
let part_size = fi.parts[i].size;
|
||||
let part_number = fi.parts[current_part].number;
|
||||
let part_size = fi.parts[current_part].size;
|
||||
let mut part_length = part_size - part_offset;
|
||||
if part_length > (length - total_read) {
|
||||
part_length = length - total_read
|
||||
@@ -2191,6 +2210,21 @@ impl SetDisks {
|
||||
|
||||
let read_offset = (part_offset / erasure.block_size) * erasure.shard_size();
|
||||
|
||||
tracing::debug!(
|
||||
bucket,
|
||||
object,
|
||||
part_index = current_part,
|
||||
part_number,
|
||||
part_offset,
|
||||
part_size,
|
||||
part_length,
|
||||
read_offset,
|
||||
till_offset,
|
||||
total_read_before = total_read,
|
||||
requested_length = length,
|
||||
"Streaming multipart part"
|
||||
);
|
||||
|
||||
let mut readers = Vec::with_capacity(disks.len());
|
||||
let mut errors = Vec::with_capacity(disks.len());
|
||||
for (idx, disk_op) in disks.iter().enumerate() {
|
||||
@@ -2236,6 +2270,15 @@ impl SetDisks {
|
||||
// part_number, part_offset, part_length, part_size
|
||||
// );
|
||||
let (written, err) = erasure.decode(writer, readers, part_offset, part_length, part_size).await;
|
||||
tracing::debug!(
|
||||
bucket,
|
||||
object,
|
||||
part_index = current_part,
|
||||
part_number,
|
||||
part_length,
|
||||
bytes_written = written,
|
||||
"Finished decoding multipart part"
|
||||
);
|
||||
if let Some(e) = err {
|
||||
let de_err: DiskError = e.into();
|
||||
let mut has_err = true;
|
||||
@@ -2274,6 +2317,8 @@ impl SetDisks {
|
||||
|
||||
// debug!("read end");
|
||||
|
||||
tracing::debug!(bucket, object, total_read, expected_length = length, "Multipart read finished");
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
@@ -3462,12 +3507,13 @@ impl ObjectIO for SetDisks {
|
||||
// let _guard_to_hold = _read_lock_guard; // moved into closure below
|
||||
tokio::spawn(async move {
|
||||
// let _guard = _guard_to_hold; // keep guard alive until task ends
|
||||
let mut writer = wd;
|
||||
if let Err(e) = Self::get_object_with_fileinfo(
|
||||
&bucket,
|
||||
&object,
|
||||
offset,
|
||||
length,
|
||||
&mut Box::new(wd),
|
||||
&mut writer,
|
||||
fi,
|
||||
files,
|
||||
&disks,
|
||||
@@ -5377,6 +5423,7 @@ impl StorageAPI for SetDisks {
|
||||
}
|
||||
|
||||
let ext_part = &curr_fi.parts[i];
|
||||
tracing::info!(target:"rustfs_ecstore::set_disk", part_number = p.part_num, part_size = ext_part.size, part_actual_size = ext_part.actual_size, "Completing multipart part");
|
||||
|
||||
if p.etag != Some(ext_part.etag.clone()) {
|
||||
error!(
|
||||
@@ -5436,6 +5483,9 @@ impl StorageAPI for SetDisks {
|
||||
fi.metadata
|
||||
.insert(format!("{RESERVED_METADATA_PREFIX_LOWER}actual-size"), object_actual_size.to_string());
|
||||
|
||||
fi.metadata
|
||||
.insert("x-rustfs-encryption-original-size".to_string(), object_actual_size.to_string());
|
||||
|
||||
if fi.is_compressed() {
|
||||
fi.metadata
|
||||
.insert(format!("{RESERVED_METADATA_PREFIX_LOWER}compression-size"), object_size.to_string());
|
||||
|
||||
@@ -23,6 +23,7 @@ use rustfs_utils::{put_uvarint, put_uvarint_len};
|
||||
use std::pin::Pin;
|
||||
use std::task::{Context, Poll};
|
||||
use tokio::io::{AsyncRead, ReadBuf};
|
||||
use tracing::debug;
|
||||
|
||||
pin_project! {
|
||||
/// A reader wrapper that encrypts data on the fly using AES-256-GCM.
|
||||
@@ -119,7 +120,7 @@ where
|
||||
header[5] = ((crc >> 8) & 0xFF) as u8;
|
||||
header[6] = ((crc >> 16) & 0xFF) as u8;
|
||||
header[7] = ((crc >> 24) & 0xFF) as u8;
|
||||
println!(
|
||||
debug!(
|
||||
"encrypt block header typ=0 len={} header={:?} plaintext_len={} ciphertext_len={}",
|
||||
clen,
|
||||
header,
|
||||
@@ -184,7 +185,10 @@ pin_project! {
|
||||
#[pin]
|
||||
pub inner: R,
|
||||
key: [u8; 32], // AES-256-GCM key
|
||||
nonce: [u8; 12], // 96-bit nonce for GCM
|
||||
base_nonce: [u8; 12], // Base nonce recorded in object metadata
|
||||
current_nonce: [u8; 12], // Active nonce for the current encrypted segment
|
||||
multipart_mode: bool,
|
||||
current_part: usize,
|
||||
buffer: Vec<u8>,
|
||||
buffer_pos: usize,
|
||||
finished: bool,
|
||||
@@ -206,7 +210,35 @@ where
|
||||
Self {
|
||||
inner,
|
||||
key,
|
||||
nonce,
|
||||
base_nonce: nonce,
|
||||
current_nonce: nonce,
|
||||
multipart_mode: false,
|
||||
current_part: 0,
|
||||
buffer: Vec::new(),
|
||||
buffer_pos: 0,
|
||||
finished: false,
|
||||
header_buf: [0u8; 8],
|
||||
header_read: 0,
|
||||
header_done: false,
|
||||
ciphertext_buf: None,
|
||||
ciphertext_read: 0,
|
||||
ciphertext_len: 0,
|
||||
}
|
||||
}
|
||||
|
||||
pub fn new_multipart(inner: R, key: [u8; 32], base_nonce: [u8; 12]) -> Self {
|
||||
let first_part = 1;
|
||||
let initial_nonce = derive_part_nonce(&base_nonce, first_part);
|
||||
|
||||
debug!("decrypt_reader: initialized multipart mode");
|
||||
|
||||
Self {
|
||||
inner,
|
||||
key,
|
||||
base_nonce,
|
||||
current_nonce: initial_nonce,
|
||||
multipart_mode: true,
|
||||
current_part: first_part,
|
||||
buffer: Vec::new(),
|
||||
buffer_pos: 0,
|
||||
finished: false,
|
||||
@@ -287,7 +319,23 @@ where
|
||||
*this.header_done = false;
|
||||
|
||||
if typ == 0xFF {
|
||||
if *this.multipart_mode {
|
||||
debug!(
|
||||
next_part = *this.current_part + 1,
|
||||
"decrypt_reader: reached segment terminator, advancing to next part"
|
||||
);
|
||||
*this.current_part += 1;
|
||||
*this.current_nonce = derive_part_nonce(this.base_nonce, *this.current_part);
|
||||
this.ciphertext_buf.take();
|
||||
*this.ciphertext_read = 0;
|
||||
*this.ciphertext_len = 0;
|
||||
continue;
|
||||
}
|
||||
|
||||
*this.finished = true;
|
||||
this.ciphertext_buf.take();
|
||||
*this.ciphertext_read = 0;
|
||||
*this.ciphertext_len = 0;
|
||||
continue;
|
||||
}
|
||||
|
||||
@@ -342,11 +390,17 @@ where
|
||||
let ciphertext = &ciphertext_buf[uvarint_len as usize..];
|
||||
|
||||
let cipher = Aes256Gcm::new_from_slice(this.key).expect("key");
|
||||
let nonce = Nonce::from_slice(this.nonce);
|
||||
let nonce = Nonce::from_slice(this.current_nonce);
|
||||
let plaintext = cipher
|
||||
.decrypt(nonce, ciphertext)
|
||||
.map_err(|e| std::io::Error::other(format!("decrypt error: {e}")))?;
|
||||
|
||||
debug!(
|
||||
part = *this.current_part,
|
||||
plaintext_len = plaintext.len(),
|
||||
"decrypt_reader: decrypted chunk"
|
||||
);
|
||||
|
||||
if plaintext.len() != plaintext_len as usize {
|
||||
this.ciphertext_buf.take();
|
||||
*this.ciphertext_read = 0;
|
||||
@@ -407,6 +461,16 @@ where
|
||||
}
|
||||
}
|
||||
|
||||
fn derive_part_nonce(base: &[u8; 12], part_number: usize) -> [u8; 12] {
|
||||
let mut nonce = *base;
|
||||
let mut suffix = [0u8; 4];
|
||||
suffix.copy_from_slice(&nonce[8..12]);
|
||||
let current = u32::from_be_bytes(suffix);
|
||||
let next = current.wrapping_add(part_number as u32);
|
||||
nonce[8..12].copy_from_slice(&next.to_be_bytes());
|
||||
nonce
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use std::io::Cursor;
|
||||
@@ -495,4 +559,42 @@ mod tests {
|
||||
|
||||
assert_eq!(&decrypted, &data);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_decrypt_reader_multipart_segments() {
|
||||
let mut key = [0u8; 32];
|
||||
let mut base_nonce = [0u8; 12];
|
||||
rand::rng().fill_bytes(&mut key);
|
||||
rand::rng().fill_bytes(&mut base_nonce);
|
||||
|
||||
let part_one = vec![0xA5; 512 * 1024];
|
||||
let part_two = vec![0x5A; 256 * 1024];
|
||||
|
||||
async fn encrypt_part(data: &[u8], key: [u8; 32], base_nonce: [u8; 12], part_number: usize) -> Vec<u8> {
|
||||
let nonce = derive_part_nonce(&base_nonce, part_number);
|
||||
let reader = BufReader::new(Cursor::new(data.to_vec()));
|
||||
let mut encrypt_reader = EncryptReader::new(WarpReader::new(reader), key, nonce);
|
||||
let mut encrypted = Vec::new();
|
||||
encrypt_reader.read_to_end(&mut encrypted).await.unwrap();
|
||||
encrypted
|
||||
}
|
||||
|
||||
let encrypted_one = encrypt_part(&part_one, key, base_nonce, 1).await;
|
||||
let encrypted_two = encrypt_part(&part_two, key, base_nonce, 2).await;
|
||||
|
||||
let mut combined = Vec::with_capacity(encrypted_one.len() + encrypted_two.len());
|
||||
combined.extend_from_slice(&encrypted_one);
|
||||
combined.extend_from_slice(&encrypted_two);
|
||||
|
||||
let reader = BufReader::new(Cursor::new(combined));
|
||||
let mut decrypt_reader = DecryptReader::new_multipart(WarpReader::new(reader), key, base_nonce);
|
||||
let mut decrypted = Vec::new();
|
||||
decrypt_reader.read_to_end(&mut decrypted).await.unwrap();
|
||||
|
||||
let mut expected = Vec::with_capacity(part_one.len() + part_two.len());
|
||||
expected.extend_from_slice(&part_one);
|
||||
expected.extend_from_slice(&part_two);
|
||||
|
||||
assert_eq!(decrypted, expected);
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user