fix(rio): preserve reader capabilities and crypto safety (#2363)

This commit is contained in:
weisd
2026-04-03 13:57:42 +08:00
committed by GitHub
parent 6a114cd2e0
commit 5d302febb7
20 changed files with 1074 additions and 699 deletions
+86 -31
View File
@@ -46,7 +46,7 @@ use rustfs_ecstore::set_disk::{MAX_PARTS_COUNT, is_valid_storage_class};
use rustfs_ecstore::store_api::{CompletePart, HTTPRangeSpec, MultipartUploadResult, ObjectIO, ObjectOptions, PutObjReader};
use rustfs_ecstore::store_api::{MultipartOperations, ObjectOperations};
use rustfs_filemeta::{ReplicationStatusType, ReplicationType};
use rustfs_rio::{CompressReader, HashReader, Reader, WarpReader};
use rustfs_rio::{CompressReader, HashReader};
use rustfs_s3_common::S3Operation;
use rustfs_targets::EventName;
use rustfs_utils::CompressionAlgorithm;
@@ -730,8 +730,6 @@ impl DefaultMultipartUsecase {
let is_compressible = rustfs_utils::http::contains_key_str(&fi.user_defined, rustfs_utils::http::SUFFIX_COMPRESSION);
let mut reader: Box<dyn Reader> = Box::new(WarpReader::new(body));
let actual_size = size;
let mut md5hex = if let Some(base64_md5) = input.content_md5 {
@@ -745,21 +743,27 @@ impl DefaultMultipartUsecase {
let mut sha256hex = get_content_sha256_with_query(&req.headers, req.uri.query());
if is_compressible {
let mut hrd = HashReader::new(reader, size, actual_size, md5hex, sha256hex, false).map_err(ApiError::from)?;
let mut reader = if is_compressible {
let mut hrd = HashReader::from_stream(body, size, actual_size, md5hex.take(), sha256hex.take(), false)
.map_err(ApiError::from)?;
if let Err(err) = hrd.add_checksum_from_s3s(&req.headers, req.trailing_headers.clone(), false) {
return Err(ApiError::from(err).into());
}
let compress_reader = CompressReader::new(hrd, CompressionAlgorithm::default());
reader = Box::new(compress_reader);
size = HashReader::SIZE_PRESERVE_LAYER;
md5hex = None;
sha256hex = None;
}
let mut reader = HashReader::new(reader, size, actual_size, md5hex, sha256hex, false).map_err(ApiError::from)?;
HashReader::from_reader(
CompressReader::new(hrd, CompressionAlgorithm::default()),
size,
actual_size,
None,
None,
false,
)
.map_err(ApiError::from)?
} else {
HashReader::from_stream(body, size, actual_size, md5hex, sha256hex, false).map_err(ApiError::from)?
};
if let Err(err) = reader.add_checksum_from_s3s(&req.headers, req.trailing_headers.clone(), size < 0) {
return Err(ApiError::from(err).into());
@@ -813,8 +817,9 @@ impl DefaultMultipartUsecase {
let requested_kms_key_id = material.kms_key_id.clone();
let encrypted_reader = material.wrap_reader(reader);
reader = HashReader::new(encrypted_reader, HashReader::SIZE_PRESERVE_LAYER, actual_size, None, None, false)
.map_err(ApiError::from)?;
reader =
HashReader::from_reader(encrypted_reader, HashReader::SIZE_PRESERVE_LAYER, actual_size, None, None, false)
.map_err(ApiError::from)?;
fi.user_defined.extend(material.metadata);
@@ -1110,8 +1115,6 @@ impl DefaultMultipartUsecase {
let is_compressible = rustfs_utils::http::contains_key_str(&mp_info.user_defined, rustfs_utils::http::SUFFIX_COMPRESSION);
let mut reader: Box<dyn Reader> = Box::new(WarpReader::new(src_stream));
let src_decryption_request = DecryptionRequest {
bucket: &src_bucket,
key: &src_key,
@@ -1123,23 +1126,74 @@ impl DefaultMultipartUsecase {
etag: src_info.etag.as_deref(),
};
if let Some(material) = sse_decryption(src_decryption_request).await? {
reader = material.wrap_single_reader(reader);
if let Some(original) = material.original_size {
src_info.actual_size = original;
}
}
let actual_size = length;
let mut size = length;
if is_compressible {
let hrd = HashReader::new(reader, size, actual_size, None, None, false).map_err(ApiError::from)?;
reader = Box::new(CompressReader::new(hrd, CompressionAlgorithm::default()));
size = HashReader::SIZE_PRESERVE_LAYER;
}
let mut reader = match sse_decryption(src_decryption_request).await? {
Some(material) => {
if let Some(original) = material.original_size {
src_info.actual_size = original;
}
let mut reader = HashReader::new(reader, size, actual_size, None, None, false).map_err(ApiError::from)?;
if material.is_multipart {
let (decrypted_stream, plaintext_size) =
material.wrap_reader(src_stream, size).await.map_err(ApiError::from)?;
size = plaintext_size;
if is_compressible {
let hrd = HashReader::from_reader(decrypted_stream, size, actual_size, None, None, false)
.map_err(ApiError::from)?;
size = HashReader::SIZE_PRESERVE_LAYER;
HashReader::from_reader(
CompressReader::new(hrd, CompressionAlgorithm::default()),
size,
actual_size,
None,
None,
false,
)
.map_err(ApiError::from)?
} else {
HashReader::from_reader(decrypted_stream, size, actual_size, None, None, false).map_err(ApiError::from)?
}
} else if is_compressible {
let hrd =
HashReader::from_stream(material.wrap_single_reader(src_stream), size, actual_size, None, None, false)
.map_err(ApiError::from)?;
size = HashReader::SIZE_PRESERVE_LAYER;
HashReader::from_reader(
CompressReader::new(hrd, CompressionAlgorithm::default()),
size,
actual_size,
None,
None,
false,
)
.map_err(ApiError::from)?
} else {
HashReader::from_stream(material.wrap_single_reader(src_stream), size, actual_size, None, None, false)
.map_err(ApiError::from)?
}
}
None => {
if is_compressible {
let hrd =
HashReader::from_stream(src_stream, size, actual_size, None, None, false).map_err(ApiError::from)?;
size = HashReader::SIZE_PRESERVE_LAYER;
HashReader::from_reader(
CompressReader::new(hrd, CompressionAlgorithm::default()),
size,
actual_size,
None,
None,
false,
)
.map_err(ApiError::from)?
} else {
HashReader::from_stream(src_stream, size, actual_size, None, None, false).map_err(ApiError::from)?
}
}
};
let server_side_encryption = mp_info
.user_defined
@@ -1180,8 +1234,9 @@ impl DefaultMultipartUsecase {
let requested_kms_key_id = material.kms_key_id.clone();
let encrypted_reader = material.wrap_reader(reader);
reader = HashReader::new(encrypted_reader, HashReader::SIZE_PRESERVE_LAYER, actual_size, None, None, false)
.map_err(ApiError::from)?;
reader =
HashReader::from_reader(encrypted_reader, HashReader::SIZE_PRESERVE_LAYER, actual_size, None, None, false)
.map_err(ApiError::from)?;
mp_info.user_defined.extend(material.metadata);
+109 -55
View File
@@ -86,7 +86,7 @@ use rustfs_filemeta::{
use rustfs_io_metrics;
use rustfs_notify::EventArgsBuilder;
use rustfs_policy::policy::action::{Action, S3Action};
use rustfs_rio::{CompressReader, EtagReader, HashReader, Reader, WarpReader};
use rustfs_rio::{CompressReader, DynReader, HashReader, wrap_reader};
use rustfs_s3_common::S3Operation;
use rustfs_s3select_api::{
object_store::bytes_stream,
@@ -183,7 +183,7 @@ struct GetObjectRequestContext {
struct GetObjectReadSetup {
info: ObjectInfo,
event_info: ObjectInfo,
final_stream: Box<dyn Reader>,
final_stream: DynReader,
rs: Option<HTTPRangeSpec>,
content_type: Option<ContentType>,
last_modified: Option<Timestamp>,
@@ -1319,14 +1319,7 @@ impl DefaultObjectUsecase {
decrypted_stream,
)
}
None => (
None,
None,
None,
None,
false,
Box::new(WarpReader::new(encrypted_stream)) as Box<dyn Reader>,
),
None => (None, None, None, None, false, wrap_reader(encrypted_stream)),
};
Ok(GetObjectReadSetup {
@@ -1824,8 +1817,6 @@ impl DefaultObjectUsecase {
}
}
let mut reader: Box<dyn Reader> = Box::new(WarpReader::new(body));
let actual_size = size;
let mut md5hex = if let Some(base64_md5) = content_md5 {
@@ -1839,12 +1830,13 @@ impl DefaultObjectUsecase {
let mut sha256hex = get_content_sha256_with_query(&req.headers, req.uri.query());
if is_compressible(&req.headers, &key) && size > MIN_COMPRESSIBLE_SIZE as i64 {
let mut reader = if is_compressible(&req.headers, &key) && size > MIN_COMPRESSIBLE_SIZE as i64 {
let algorithm = CompressionAlgorithm::default();
insert_str(&mut metadata, SUFFIX_COMPRESSION, algorithm.to_string());
insert_str(&mut metadata, SUFFIX_ACTUAL_SIZE, size.to_string());
let mut hrd = HashReader::new(reader, size as i64, size as i64, md5hex, sha256hex, false).map_err(ApiError::from)?;
let mut hrd =
HashReader::from_stream(body, size, size, md5hex.take(), sha256hex.take(), false).map_err(ApiError::from)?;
if let Err(err) = hrd.add_checksum_from_s3s(&req.headers, req.trailing_headers.clone(), false) {
return Err(ApiError::from(err).into());
@@ -1854,13 +1846,12 @@ impl DefaultObjectUsecase {
insert_str(&mut opts.user_defined, SUFFIX_COMPRESSION, algorithm.to_string());
insert_str(&mut opts.user_defined, SUFFIX_ACTUAL_SIZE, size.to_string());
reader = Box::new(CompressReader::new(hrd, algorithm));
size = HashReader::SIZE_PRESERVE_LAYER;
md5hex = None;
sha256hex = None;
}
let mut reader = HashReader::new(reader, size, actual_size, md5hex, sha256hex, false).map_err(ApiError::from)?;
HashReader::from_reader(CompressReader::new(hrd, algorithm), size, actual_size, None, None, false)
.map_err(ApiError::from)?
} else {
HashReader::from_stream(body, size, actual_size, md5hex, sha256hex, false).map_err(ApiError::from)?
};
if size >= 0 {
if let Err(err) = reader.add_checksum_from_s3s(&req.headers, req.trailing_headers.clone(), false) {
@@ -1901,7 +1892,7 @@ impl DefaultObjectUsecase {
effective_kms_key_id = material.kms_key_id.clone();
let encrypted_reader = material.wrap_reader(reader);
reader = HashReader::new(encrypted_reader, HashReader::SIZE_PRESERVE_LAYER, actual_size, None, None, false)
reader = HashReader::from_reader(encrypted_reader, HashReader::SIZE_PRESERVE_LAYER, actual_size, None, None, false)
.map_err(ApiError::from)?;
let encryption_metadata = material.metadata;
@@ -2504,7 +2495,7 @@ impl DefaultObjectUsecase {
key: &str,
info: ObjectInfo,
event_info: ObjectInfo,
final_stream: Box<dyn Reader>,
final_stream: DynReader,
rs: Option<HTTPRangeSpec>,
content_type: Option<ContentType>,
last_modified: Option<Timestamp>,
@@ -3339,8 +3330,6 @@ impl DefaultObjectUsecase {
src_info.metadata_only = true;
}
let mut reader: Box<dyn Reader> = Box::new(WarpReader::new(gr.stream));
let decryption_request = DecryptionRequest {
bucket: &src_bucket,
key: &src_key,
@@ -3352,11 +3341,12 @@ impl DefaultObjectUsecase {
etag: src_info.etag.as_deref(),
};
if let Some(material) = sse_decryption(decryption_request).await? {
reader = material.wrap_single_reader(reader);
if let Some(original) = material.original_size {
src_info.actual_size = original;
}
let decryption_material = sse_decryption(decryption_request).await?;
if let Some(material) = decryption_material.as_ref()
&& let Some(original) = material.original_size
{
src_info.actual_size = original;
}
strip_managed_encryption_metadata(&mut src_info.user_defined);
@@ -3367,16 +3357,11 @@ impl DefaultObjectUsecase {
let mut compress_metadata = HashMap::new();
if is_compressible(&req.headers, &key) && actual_size > MIN_COMPRESSIBLE_SIZE as i64 {
let should_compress = is_compressible(&req.headers, &key) && actual_size > MIN_COMPRESSIBLE_SIZE as i64;
if should_compress {
insert_str(&mut compress_metadata, SUFFIX_COMPRESSION, CompressionAlgorithm::default().to_string());
insert_str(&mut compress_metadata, SUFFIX_ACTUAL_SIZE, actual_size.to_string());
let hrd = EtagReader::new(reader, None);
// let hrd = HashReader::new(reader, length, actual_size, None, false).map_err(ApiError::from)?;
reader = Box::new(CompressReader::new(hrd, CompressionAlgorithm::default()));
length = HashReader::SIZE_PRESERVE_LAYER;
} else {
remove_str(&mut src_info.user_defined, SUFFIX_COMPRESSION);
remove_str(&mut src_info.user_defined, SUFFIX_ACTUAL_SIZE);
@@ -3408,7 +3393,68 @@ impl DefaultObjectUsecase {
src_info.user_defined.extend(object_lock_metadata);
}
let mut reader = HashReader::new(reader, length, actual_size, None, None, false).map_err(ApiError::from)?;
let mut reader = match decryption_material {
Some(material) => {
if material.is_multipart {
let (decrypted_stream, plaintext_size) =
material.wrap_reader(gr.stream, length).await.map_err(ApiError::from)?;
length = plaintext_size;
if should_compress {
let hrd = HashReader::from_reader(decrypted_stream, length, actual_size, None, None, false)
.map_err(ApiError::from)?;
length = HashReader::SIZE_PRESERVE_LAYER;
HashReader::from_reader(
CompressReader::new(hrd, CompressionAlgorithm::default()),
length,
actual_size,
None,
None,
false,
)
.map_err(ApiError::from)?
} else {
HashReader::from_reader(decrypted_stream, length, actual_size, None, None, false)
.map_err(ApiError::from)?
}
} else if should_compress {
let hrd =
HashReader::from_stream(material.wrap_single_reader(gr.stream), length, actual_size, None, None, false)
.map_err(ApiError::from)?;
length = HashReader::SIZE_PRESERVE_LAYER;
HashReader::from_reader(
CompressReader::new(hrd, CompressionAlgorithm::default()),
length,
actual_size,
None,
None,
false,
)
.map_err(ApiError::from)?
} else {
HashReader::from_stream(material.wrap_single_reader(gr.stream), length, actual_size, None, None, false)
.map_err(ApiError::from)?
}
}
None => {
if should_compress {
let hrd =
HashReader::from_stream(gr.stream, length, actual_size, None, None, false).map_err(ApiError::from)?;
length = HashReader::SIZE_PRESERVE_LAYER;
HashReader::from_reader(
CompressReader::new(hrd, CompressionAlgorithm::default()),
length,
actual_size,
None,
None,
false,
)
.map_err(ApiError::from)?
} else {
HashReader::from_stream(gr.stream, length, actual_size, None, None, false).map_err(ApiError::from)?
}
}
};
let encryption_request = EncryptionRequest {
bucket: &bucket,
@@ -3429,7 +3475,7 @@ impl DefaultObjectUsecase {
effective_kms_key_id = material.kms_key_id.clone();
let encrypted_reader = material.wrap_reader(reader);
reader = HashReader::new(encrypted_reader, HashReader::SIZE_PRESERVE_LAYER, actual_size, None, None, false)
reader = HashReader::from_reader(encrypted_reader, HashReader::SIZE_PRESERVE_LAYER, actual_size, None, None, false)
.map_err(ApiError::from)?;
src_info.user_defined.extend(material.metadata);
@@ -4816,9 +4862,8 @@ impl DefaultObjectUsecase {
let sha256hex = get_content_sha256_with_query(&req.headers, req.uri.query());
let actual_size = size;
let reader: Box<dyn Reader> = Box::new(WarpReader::new(body));
let mut archive_reader = HashReader::new(reader, size, actual_size, md5hex, sha256hex, false).map_err(ApiError::from)?;
let mut archive_reader =
HashReader::from_stream(body, size, actual_size, md5hex, sha256hex, false).map_err(ApiError::from)?;
if let Err(err) = archive_reader.add_checksum_from_s3s(&req.headers, req.trailing_headers.clone(), false) {
return Err(ApiError::from(err).into());
@@ -4935,30 +4980,39 @@ impl DefaultObjectUsecase {
debug!("Extracting file: {}, size: {} bytes", fpath, size);
let mut reader: Box<dyn Reader> = if is_dir {
if is_dir {
if extract_options.ignore_dirs {
debug!("Skipping directory entry during archive extract: {}", fpath);
continue;
}
size = 0;
Box::new(WarpReader::new(std::io::Cursor::new(Vec::new())))
} else {
Box::new(WarpReader::new(f))
};
}
let actual_size = size;
if !is_dir && is_compressible(&HeaderMap::new(), &fpath) && size > MIN_COMPRESSIBLE_SIZE as i64 {
let should_compress = !is_dir && is_compressible(&HeaderMap::new(), &fpath) && size > MIN_COMPRESSIBLE_SIZE as i64;
let mut hrd = if is_dir {
HashReader::from_stream(std::io::Cursor::new(Vec::new()), size, actual_size, None, None, false)
.map_err(ApiError::from)?
} else if should_compress {
insert_str(&mut metadata, SUFFIX_COMPRESSION, CompressionAlgorithm::default().to_string());
insert_str(&mut metadata, SUFFIX_ACTUAL_SIZE, size.to_string());
let hrd = HashReader::new(reader, size, actual_size, None, None, false).map_err(ApiError::from)?;
reader = Box::new(CompressReader::new(hrd, CompressionAlgorithm::default()));
let hrd = HashReader::from_stream(f, size, actual_size, None, None, false).map_err(ApiError::from)?;
size = HashReader::SIZE_PRESERVE_LAYER;
}
let mut hrd = HashReader::new(reader, size, actual_size, None, None, false).map_err(ApiError::from)?;
HashReader::from_reader(
CompressReader::new(hrd, CompressionAlgorithm::default()),
size,
actual_size,
None,
None,
false,
)
.map_err(ApiError::from)?
} else {
HashReader::from_stream(f, size, actual_size, None, None, false).map_err(ApiError::from)?
};
apply_put_request_object_lock_opts(
&bucket,
object_lock_legal_hold_status.clone(),
@@ -4986,7 +5040,7 @@ impl DefaultObjectUsecase {
effective_kms_key_id = material.kms_key_id.clone();
let encrypted_reader = material.wrap_reader(hrd);
hrd = HashReader::new(encrypted_reader, HashReader::SIZE_PRESERVE_LAYER, actual_size, None, None, false)
hrd = HashReader::from_reader(encrypted_reader, HashReader::SIZE_PRESERVE_LAYER, actual_size, None, None, false)
.map_err(ApiError::from)?;
let encryption_metadata = material.metadata;
-1
View File
@@ -21,7 +21,6 @@ pub(crate) mod entity;
pub(crate) mod helper;
pub mod lock_optimizer;
pub mod options;
pub(crate) mod readers;
pub mod rpc;
pub(crate) mod s3_api;
mod sse;
-55
View File
@@ -1,55 +0,0 @@
// 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.
use tokio::io::{AsyncRead, AsyncSeek};
/// Seekable in-memory async reader used by internal S3 API fast paths (e.g., GET/HEAD)
/// and by SSE flows that need a rewindable in-memory stream.
pub(crate) struct InMemoryAsyncReader {
cursor: std::io::Cursor<Vec<u8>>,
}
impl InMemoryAsyncReader {
pub(crate) fn new(data: Vec<u8>) -> Self {
Self {
cursor: std::io::Cursor::new(data),
}
}
}
impl AsyncRead for InMemoryAsyncReader {
fn poll_read(
mut self: std::pin::Pin<&mut Self>,
_cx: &mut std::task::Context<'_>,
buf: &mut tokio::io::ReadBuf<'_>,
) -> std::task::Poll<std::io::Result<()>> {
let unfilled = buf.initialize_unfilled();
let bytes_read = std::io::Read::read(&mut self.cursor, unfilled)?;
buf.advance(bytes_read);
std::task::Poll::Ready(Ok(()))
}
}
impl AsyncSeek for InMemoryAsyncReader {
fn start_seek(mut self: std::pin::Pin<&mut Self>, position: std::io::SeekFrom) -> std::io::Result<()> {
// std::io::Cursor natively supports negative SeekCurrent offsets
// It will automatically handle validation and return an error if the final position would be negative
std::io::Seek::seek(&mut self.cursor, position)?;
Ok(())
}
fn poll_complete(self: std::pin::Pin<&mut Self>, _cx: &mut std::task::Context<'_>) -> std::task::Poll<std::io::Result<u64>> {
std::task::Poll::Ready(Ok(self.cursor.position()))
}
}
+208 -88
View File
@@ -23,8 +23,8 @@
//!
//! ### Unified API
//! The module provides two core functions that automatically route to the correct encryption method:
//! - `apply_encryption()` - Unified encryption entry point
//! - `apply_decryption()` - Unified decryption entry point
//! - `sse_encryption()` - Unified encryption entry point
//! - `sse_decryption()` - Unified decryption entry point
//!
//! ### Managed SSE (SSE-S3 / SSE-KMS)
//! - Keys are managed by the server-side KMS service
@@ -52,8 +52,8 @@
//! part_number: None,
//! };
//!
//! if let Some(material) = apply_encryption(request).await? {
//! reader = material.wrap_reader(reader)?;
//! if let Some(material) = sse_encryption(request).await? {
//! reader = material.wrap_reader(reader);
//! metadata.extend(material.metadata);
//! }
//!
@@ -67,8 +67,10 @@
//! part_number: None,
//! };
//!
//! if let Some(material) = apply_decryption(request).await? {
//! reader = material.wrap_reader(reader)?;
//! if let Some(material) = sse_decryption(request).await? {
//! let (decrypted_reader, plaintext_size) = material.wrap_reader(reader, actual_size).await?;
//! reader = decrypted_reader;
//! content_size = plaintext_size;
//! }
//! ```
@@ -87,19 +89,17 @@ use rustfs_kms::{
service_manager::get_global_encryption_service,
types::{EncryptionMetadata, ObjectEncryptionContext},
};
use rustfs_rio::{DecryptReader, EncryptReader, HardLimitReader, Reader, WarpReader};
use rustfs_rio::{DecryptReader, DynReader, EncryptReader, HardLimitReader, ReadStream, boxed_reader, wrap_reader};
use rustfs_utils::get_env_opt_str;
use s3s::S3ErrorCode;
use s3s::dto::ServerSideEncryption;
use std::collections::HashMap;
use std::sync::{Arc, OnceLock};
use tokio::io::AsyncRead;
use tracing::{debug, error};
const INTERNAL_ENCRYPTION_KEY_ID_HEADER: &str = "x-rustfs-encryption-key-id";
use crate::error::ApiError;
use crate::storage::readers::InMemoryAsyncReader;
use rustfs_ecstore::bucket::metadata_sys;
use rustfs_ecstore::error::Error;
use s3s::dto::{SSECustomerAlgorithm, SSECustomerKey, SSECustomerKeyMD5, SSEKMSKeyId};
@@ -619,7 +619,7 @@ impl EncryptionMaterial {
/// Wrap a reader with encryption
pub fn wrap_reader<R>(&self, reader: R) -> Box<EncryptReader<R>>
where
R: Reader + 'static,
R: rustfs_rio::ReadStream + 'static,
{
Box::new(EncryptReader::new(reader, self.key_bytes, self.nonce))
}
@@ -630,42 +630,40 @@ impl DecryptionMaterial {
/// For multipart objects, use `wrap_multipart_stream` instead
pub fn wrap_single_reader<R>(&self, reader: R) -> Box<DecryptReader<R>>
where
R: Reader + 'static,
R: rustfs_rio::ReadStream + 'static,
{
Box::new(DecryptReader::new(reader, self.key_bytes, self.nonce))
}
/// Wrap a stream with multipart decryption
/// Returns the decrypted reader and the total plaintext size
pub async fn wrap_multipart_stream(
&self,
encrypted_stream: Box<dyn AsyncRead + Unpin + Send + Sync>,
) -> Result<(Box<dyn Reader>, i64), StorageError> {
pub async fn wrap_multipart_stream<R>(&self, encrypted_stream: R) -> Result<(DynReader, i64), StorageError>
where
R: ReadStream + 'static,
{
decrypt_multipart_managed_stream(encrypted_stream, &self.parts, self.key_bytes, self.nonce).await
}
/// Unified method to wrap stream with decryption and hard limit
/// Handles both single-part and multipart objects, applies decryption and size limiting
/// Accepts AsyncRead stream (from object storage) and returns (decrypted_reader, plaintext_size)
pub async fn wrap_reader(
self,
stream: Box<dyn AsyncRead + Unpin + Send + Sync>,
actual_size: i64,
) -> Result<(Box<dyn Reader>, i64), StorageError> {
let (mut final_stream, response_content_length): (Box<dyn Reader>, i64) = if self.is_multipart {
/// Accepts a readable stream (from object storage) and returns (decrypted_reader, plaintext_size)
pub async fn wrap_reader<R>(self, stream: R, actual_size: i64) -> Result<(DynReader, i64), StorageError>
where
R: ReadStream + 'static,
{
let (mut final_stream, response_content_length): (DynReader, i64) = if self.is_multipart {
// Multipart decryption
let (decrypted_reader, plain_size) = self.wrap_multipart_stream(stream).await?;
(decrypted_reader, plain_size)
} else {
// Single-part decryption - wrap AsyncRead into Reader first
let warp_reader = WarpReader::new(stream);
let decrypt_reader = self.wrap_single_reader(warp_reader);
// Single-part decryption keeps Reader capabilities via the generic wrapper helper.
let decrypt_reader = self.wrap_single_reader(wrap_reader(stream));
let plain_size = self.original_size.unwrap_or(actual_size);
(decrypt_reader, plain_size)
};
// Add hard limit reader to prevent over-reading
// final_stream is already Box<dyn Reader>, no need to wrap with WarpReader
// final_stream is already a DynReader, no need to wrap with WarpReader
let limit_reader = HardLimitReader::new(final_stream, response_content_length);
final_stream = Box::new(limit_reader);
@@ -711,8 +709,8 @@ impl DecryptionMaterial {
/// part_number: None,
/// };
///
/// if let Some(material) = apply_encryption(request).await? {
/// reader = material.wrap_reader(reader)?;
/// if let Some(material) = sse_encryption(request).await? {
/// reader = material.wrap_reader(reader);
/// metadata.extend(material.metadata);
/// }
/// ```
@@ -846,8 +844,10 @@ pub async fn sse_prepare_encryption(request: PrepareEncryptionRequest<'_>) -> Re
/// part_number: None,
/// };
///
/// if let Some(material) = apply_decryption(request).await? {
/// reader = material.wrap_reader(reader)?;
/// if let Some(material) = sse_decryption(request).await? {
/// let (decrypted_reader, plaintext_size) = material.wrap_reader(reader, actual_size).await?;
/// reader = decrypted_reader;
/// content_size = plaintext_size;
/// }
/// ```
pub async fn sse_decryption(request: DecryptionRequest<'_>) -> Result<Option<DecryptionMaterial>, ApiError> {
@@ -1642,49 +1642,43 @@ pub fn strip_managed_encryption_metadata(metadata: &mut HashMap<String, String>)
// Multipart Encryption Support
// ============================================================================
/// Derive a unique nonce for each part in a multipart upload
///
/// Uses the base nonce and increments the counter portion by part number.
/// This ensures each part has a unique nonce while maintaining determinism.
pub fn derive_part_nonce(base: [u8; 12], part_number: usize) -> [u8; 12] {
let mut nonce = base;
let current = u32::from_be_bytes([nonce[8], nonce[9], nonce[10], nonce[11]]);
let incremented = current.wrapping_add(part_number as u32);
nonce[8..12].copy_from_slice(&incremented.to_be_bytes());
nonce
derive_nonce_offset(base, 4, part_number)
}
pub(crate) async fn decrypt_multipart_managed_stream(
mut encrypted_stream: Box<dyn AsyncRead + Unpin + Send + Sync>,
#[cfg(test)]
fn derive_legacy_part_nonce(base: [u8; 12], part_number: usize) -> [u8; 12] {
derive_nonce_offset(base, 8, part_number)
}
fn derive_nonce_offset(mut base: [u8; 12], start: usize, offset: usize) -> [u8; 12] {
let current = u32::from_be_bytes([base[start], base[start + 1], base[start + 2], base[start + 3]]);
let incremented = current.wrapping_add(offset as u32);
base[start..start + 4].copy_from_slice(&incremented.to_be_bytes());
base
}
pub(crate) async fn decrypt_multipart_managed_stream<R>(
encrypted_stream: R,
parts: &[ObjectPartInfo],
key_bytes: [u8; 32],
base_nonce: [u8; 12],
) -> Result<(Box<dyn Reader>, i64), StorageError> {
let total_plain_capacity: usize = parts.iter().map(|part| part.actual_size.max(0) as usize).sum();
) -> Result<(DynReader, i64), StorageError>
where
R: ReadStream + 'static,
{
let total_plain_size = parts
.iter()
.map(|part| {
if part.actual_size > 0 {
part.actual_size
} else {
part.size as i64
}
})
.sum();
let mut plaintext = Vec::with_capacity(total_plain_capacity);
for part in parts {
if part.size == 0 {
continue;
}
let mut encrypted_part = vec![0u8; part.size];
tokio::io::AsyncReadExt::read_exact(&mut encrypted_stream, &mut encrypted_part)
.await
.map_err(|e| StorageError::other(format!("failed to read encrypted multipart segment {}: {}", part.number, e)))?;
let part_nonce = derive_part_nonce(base_nonce, part.number);
let cursor = std::io::Cursor::new(encrypted_part);
let mut decrypt_reader = DecryptReader::new(WarpReader::new(cursor), key_bytes, part_nonce);
tokio::io::AsyncReadExt::read_to_end(&mut decrypt_reader, &mut plaintext)
.await
.map_err(|e| StorageError::other(format!("failed to decrypt multipart segment {}: {}", part.number, e)))?;
}
let total_plain_size = plaintext.len() as i64;
let reader = Box::new(WarpReader::new(InMemoryAsyncReader::new(plaintext))) as Box<dyn Reader>;
let reader = boxed_reader(DecryptReader::new_multipart(wrap_reader(encrypted_stream), key_bytes, base_nonce));
Ok((reader, total_plain_size))
}
@@ -1951,13 +1945,139 @@ mod tests {
let part1 = derive_part_nonce(base, 1);
let part2 = derive_part_nonce(base, 2);
// First 8 bytes should be unchanged
assert_eq!(&base[..8], &part1[..8]);
assert_eq!(&base[..8], &part2[..8]);
assert_eq!(&base[..4], &part1[..4]);
assert_eq!(&base[8..], &part1[8..]);
assert_ne!(&base[4..8], &part1[4..8]);
assert_ne!(&part1[4..8], &part2[4..8]);
}
// Last 4 bytes should be incremented
assert_ne!(&base[8..], &part1[8..]);
assert_ne!(&part1[8..], &part2[8..]);
#[tokio::test]
async fn test_decrypt_multipart_managed_stream_accepts_legacy_part_nonce_layout() {
use std::io::Cursor;
use tokio::io::AsyncReadExt;
let key_bytes = [7u8; 32];
let base_nonce = [3u8; 12];
let part_one_plaintext = vec![0x11; rustfs_rio::DEFAULT_ENCRYPTION_BLOCK_SIZE + 19];
let part_two_plaintext = vec![0x22; rustfs_rio::DEFAULT_ENCRYPTION_BLOCK_SIZE + 37];
let part_one_nonce = derive_legacy_part_nonce(base_nonce, 1);
let part_two_nonce = derive_legacy_part_nonce(base_nonce, 2);
let first_part = {
let mut buf = Vec::new();
EncryptReader::new(Cursor::new(part_one_plaintext.clone()), key_bytes, part_one_nonce)
.read_to_end(&mut buf)
.await
.unwrap();
buf
};
let second_part = {
let mut buf = Vec::new();
EncryptReader::new(Cursor::new(part_two_plaintext.clone()), key_bytes, part_two_nonce)
.read_to_end(&mut buf)
.await
.unwrap();
buf
};
let mut encrypted_stream = Vec::with_capacity(first_part.len() + second_part.len());
encrypted_stream.extend_from_slice(&first_part);
encrypted_stream.extend_from_slice(&second_part);
let parts = vec![
ObjectPartInfo {
number: 1,
size: first_part.len(),
actual_size: part_one_plaintext.len() as i64,
..Default::default()
},
ObjectPartInfo {
number: 2,
size: second_part.len(),
actual_size: part_two_plaintext.len() as i64,
..Default::default()
},
];
let (mut decrypted_reader, plaintext_size) =
decrypt_multipart_managed_stream(Cursor::new(encrypted_stream), &parts, key_bytes, base_nonce)
.await
.unwrap();
let mut decrypted = Vec::new();
decrypted_reader.read_to_end(&mut decrypted).await.unwrap();
let mut expected = part_one_plaintext;
expected.extend_from_slice(&part_two_plaintext);
assert_eq!(plaintext_size, expected.len() as i64);
assert_eq!(decrypted, expected);
}
#[tokio::test]
async fn test_decrypt_multipart_managed_stream_supports_current_nonce_layout() {
use std::io::Cursor;
use tokio::io::AsyncReadExt;
let key_bytes = [9u8; 32];
let base_nonce = [5u8; 12];
let part_one_plaintext = vec![0x33; rustfs_rio::DEFAULT_ENCRYPTION_BLOCK_SIZE + 11];
let part_two_plaintext = vec![0x44; rustfs_rio::DEFAULT_ENCRYPTION_BLOCK_SIZE * 2 + 7];
let part_one_nonce = derive_part_nonce(base_nonce, 1);
let part_two_nonce = derive_part_nonce(base_nonce, 2);
let first_part = {
let mut buf = Vec::new();
EncryptReader::new(Cursor::new(part_one_plaintext.clone()), key_bytes, part_one_nonce)
.read_to_end(&mut buf)
.await
.unwrap();
buf
};
let second_part = {
let mut buf = Vec::new();
EncryptReader::new(Cursor::new(part_two_plaintext.clone()), key_bytes, part_two_nonce)
.read_to_end(&mut buf)
.await
.unwrap();
buf
};
let mut encrypted_stream = Vec::with_capacity(first_part.len() + second_part.len());
encrypted_stream.extend_from_slice(&first_part);
encrypted_stream.extend_from_slice(&second_part);
let parts = vec![
ObjectPartInfo {
number: 1,
size: first_part.len(),
actual_size: part_one_plaintext.len() as i64,
..Default::default()
},
ObjectPartInfo {
number: 2,
size: second_part.len(),
actual_size: part_two_plaintext.len() as i64,
..Default::default()
},
];
let (mut decrypted_reader, plaintext_size) =
decrypt_multipart_managed_stream(Cursor::new(encrypted_stream), &parts, key_bytes, base_nonce)
.await
.unwrap();
let mut decrypted = Vec::new();
decrypted_reader.read_to_end(&mut decrypted).await.unwrap();
let mut expected = part_one_plaintext;
expected.extend_from_slice(&part_two_plaintext);
assert_eq!(plaintext_size, expected.len() as i64);
assert_eq!(decrypted, expected);
}
#[test]
@@ -2436,8 +2556,8 @@ mod tests {
println!("Original plaintext: {:?}", String::from_utf8_lossy(plaintext));
println!("Plaintext length: {} bytes", plaintext.len());
// 4. Encrypt with EncryptReader (wrap Cursor with WarpReader)
let plaintext_reader = WarpReader::new(Cursor::new(plaintext.to_vec()));
// 4. Encrypt with EncryptReader.
let plaintext_reader = Cursor::new(plaintext.to_vec());
let mut encrypt_reader = EncryptReader::new(plaintext_reader, data_key.plaintext_key, data_key.nonce);
// Read encrypted data
@@ -2460,8 +2580,8 @@ mod tests {
"Encrypted data should be different from plaintext"
);
// 5. Decrypt with DecryptReader (wrap Cursor with WarpReader)
let encrypted_reader = WarpReader::new(Cursor::new(encrypted_data));
// 5. Decrypt with DecryptReader.
let encrypted_reader = Cursor::new(encrypted_data);
let mut decrypt_reader = DecryptReader::new(encrypted_reader, data_key.plaintext_key, data_key.nonce);
// Read decrypted data
@@ -2502,8 +2622,8 @@ mod tests {
let plaintext: Vec<u8> = (0..plaintext_size).map(|i| (i % 256) as u8).collect();
println!("Testing with {} bytes of data", plaintext.len());
// Encrypt (wrap with WarpReader)
let plaintext_reader = WarpReader::new(Cursor::new(plaintext.clone()));
// Encrypt.
let plaintext_reader = Cursor::new(plaintext.clone());
let mut encrypt_reader = EncryptReader::new(plaintext_reader, data_key.plaintext_key, data_key.nonce);
let mut encrypted_data = Vec::new();
@@ -2514,8 +2634,8 @@ mod tests {
println!("Encrypted {} bytes to {} bytes", plaintext.len(), encrypted_data.len());
// Decrypt (wrap with WarpReader)
let encrypted_reader = WarpReader::new(Cursor::new(encrypted_data));
// Decrypt.
let encrypted_reader = Cursor::new(encrypted_data);
let mut decrypt_reader = DecryptReader::new(encrypted_reader, data_key.plaintext_key, data_key.nonce);
let mut decrypted_data = Vec::new();
@@ -2560,14 +2680,14 @@ mod tests {
// Same plaintext
let plaintext = b"Same plaintext";
// Encrypt with first key (wrap with WarpReader)
let reader1 = WarpReader::new(Cursor::new(plaintext.to_vec()));
// Encrypt with first key.
let reader1 = Cursor::new(plaintext.to_vec());
let mut encrypt_reader1 = EncryptReader::new(reader1, data_key1.plaintext_key, data_key1.nonce);
let mut encrypted1 = Vec::new();
encrypt_reader1.read_to_end(&mut encrypted1).await.unwrap();
// Encrypt with second key (wrap with WarpReader)
let reader2 = WarpReader::new(Cursor::new(plaintext.to_vec()));
// Encrypt with second key.
let reader2 = Cursor::new(plaintext.to_vec());
let mut encrypt_reader2 = EncryptReader::new(reader2, data_key2.plaintext_key, data_key2.nonce);
let mut encrypted2 = Vec::new();
encrypt_reader2.read_to_end(&mut encrypted2).await.unwrap();
@@ -2620,14 +2740,14 @@ mod tests {
// 5. Use decrypted key to encrypt/decrypt data
let plaintext = b"Test data with decrypted DEK";
// Encrypt with original key (wrap with WarpReader)
let reader = WarpReader::new(Cursor::new(plaintext.to_vec()));
// Encrypt with original key.
let reader = Cursor::new(plaintext.to_vec());
let mut encrypt_reader = EncryptReader::new(reader, original_plaintext_key, original_nonce);
let mut encrypted_data = Vec::new();
encrypt_reader.read_to_end(&mut encrypted_data).await.unwrap();
// Decrypt with recovered key (simulating GET operation) (wrap with WarpReader)
let reader = WarpReader::new(Cursor::new(encrypted_data));
// Decrypt with recovered key (simulating GET operation).
let reader = Cursor::new(encrypted_data);
let mut decrypt_reader = DecryptReader::new(
reader,
decrypted_plaintext_key,
+17 -17
View File
@@ -16,7 +16,7 @@
mod tests {
use crate::storage::sse::SseDekProvider;
use crate::storage::sse::TestSseDekProvider;
use rustfs_rio::{DecryptReader, EncryptReader, WarpReader};
use rustfs_rio::{DecryptReader, EncryptReader};
use std::io::Cursor;
use tokio::io::AsyncReadExt;
@@ -51,8 +51,8 @@ mod tests {
println!("Original plaintext: {:?}", String::from_utf8_lossy(plaintext));
println!("Plaintext length: {} bytes", plaintext.len());
// Step 4: Encrypt using EncryptReader (wrap Cursor with WarpReader)
let plaintext_reader = WarpReader::new(Cursor::new(plaintext.to_vec()));
// Step 4: Encrypt using EncryptReader.
let plaintext_reader = Cursor::new(plaintext.to_vec());
let mut encrypt_reader = EncryptReader::new(plaintext_reader, data_key.plaintext_key, data_key.nonce);
// Read encrypted data
@@ -75,8 +75,8 @@ mod tests {
"Encrypted data should be different from plaintext"
);
// Step 5: Decrypt using DecryptReader (wrap Cursor with WarpReader)
let encrypted_reader = WarpReader::new(Cursor::new(encrypted_data));
// Step 5: Decrypt using DecryptReader.
let encrypted_reader = Cursor::new(encrypted_data);
let mut decrypt_reader = DecryptReader::new(encrypted_reader, data_key.plaintext_key, data_key.nonce);
// Read decrypted data
@@ -115,8 +115,8 @@ mod tests {
let plaintext: Vec<u8> = (0..plaintext_size).map(|i| (i % 256) as u8).collect();
println!("Testing with {} bytes of data", plaintext.len());
// Encrypt (wrap with WarpReader)
let plaintext_reader = WarpReader::new(Cursor::new(plaintext.clone()));
// Encrypt.
let plaintext_reader = Cursor::new(plaintext.clone());
let mut encrypt_reader = EncryptReader::new(plaintext_reader, data_key.plaintext_key, data_key.nonce);
let mut encrypted_data = Vec::new();
@@ -127,8 +127,8 @@ mod tests {
println!("Encrypted {} bytes to {} bytes", plaintext.len(), encrypted_data.len());
// Decrypt (wrap with WarpReader)
let encrypted_reader = WarpReader::new(Cursor::new(encrypted_data));
// Decrypt.
let encrypted_reader = Cursor::new(encrypted_data);
let mut decrypt_reader = DecryptReader::new(encrypted_reader, data_key.plaintext_key, data_key.nonce);
let mut decrypted_data = Vec::new();
@@ -171,14 +171,14 @@ mod tests {
// Same plaintext
let plaintext = b"Same plaintext";
// Encrypt with first key (wrap with WarpReader)
let reader1 = WarpReader::new(Cursor::new(plaintext.to_vec()));
// Encrypt with first key.
let reader1 = Cursor::new(plaintext.to_vec());
let mut encrypt_reader1 = EncryptReader::new(reader1, data_key1.plaintext_key, data_key1.nonce);
let mut encrypted1 = Vec::new();
encrypt_reader1.read_to_end(&mut encrypted1).await.unwrap();
// Encrypt with second key (wrap with WarpReader)
let reader2 = WarpReader::new(Cursor::new(plaintext.to_vec()));
// Encrypt with second key.
let reader2 = Cursor::new(plaintext.to_vec());
let mut encrypt_reader2 = EncryptReader::new(reader2, data_key2.plaintext_key, data_key2.nonce);
let mut encrypted2 = Vec::new();
encrypt_reader2.read_to_end(&mut encrypted2).await.unwrap();
@@ -226,14 +226,14 @@ mod tests {
// Step 4: Use decrypted key to encrypt/decrypt data
let plaintext = b"Test data with decrypted DEK";
// Encrypt with original key (wrap with WarpReader)
let reader = WarpReader::new(Cursor::new(plaintext.to_vec()));
// Encrypt with original key.
let reader = Cursor::new(plaintext.to_vec());
let mut encrypt_reader = EncryptReader::new(reader, original_plaintext_key, original_nonce);
let mut encrypted_data = Vec::new();
encrypt_reader.read_to_end(&mut encrypted_data).await.unwrap();
// Decrypt with recovered key (simulating GET operation) (wrap with WarpReader)
let reader = WarpReader::new(Cursor::new(encrypted_data));
// Decrypt with recovered key (simulating GET operation).
let reader = Cursor::new(encrypted_data);
let mut decrypt_reader = DecryptReader::new(
reader,
decrypted_plaintext_key,