From 5d302febb772c70a661b7e9817172347b25f8ec9 Mon Sep 17 00:00:00 2001 From: weisd Date: Fri, 3 Apr 2026 13:57:42 +0800 Subject: [PATCH] fix(rio): preserve reader capabilities and crypto safety (#2363) --- crates/ecstore/src/data_movement.rs | 12 +- crates/ecstore/src/set_disk.rs | 24 +- crates/ecstore/src/store_api.rs | 2 +- crates/ecstore/src/store_api/readers.rs | 10 +- crates/protocols/src/swift/object.rs | 30 +- crates/rio/src/compress_reader.rs | 113 +++----- crates/rio/src/encrypt_reader.rs | 358 ++++++++++++++++-------- crates/rio/src/etag.rs | 32 +-- crates/rio/src/etag_reader.rs | 69 +++-- crates/rio/src/hardlimit_reader.rs | 46 +-- crates/rio/src/hash_reader.rs | 259 ++++++++++++----- crates/rio/src/lib.rs | 110 +++++++- crates/rio/src/limit_reader.rs | 37 +-- crates/rio/src/reader.rs | 4 +- rustfs/src/app/multipart_usecase.rs | 117 ++++++-- rustfs/src/app/object_usecase.rs | 164 +++++++---- rustfs/src/storage/mod.rs | 1 - rustfs/src/storage/readers.rs | 55 ---- rustfs/src/storage/sse.rs | 296 ++++++++++++++------ rustfs/src/storage/sse_test.rs | 34 +-- 20 files changed, 1074 insertions(+), 699 deletions(-) delete mode 100644 rustfs/src/storage/readers.rs diff --git a/crates/ecstore/src/data_movement.rs b/crates/ecstore/src/data_movement.rs index 5d7721911..f40840624 100644 --- a/crates/ecstore/src/data_movement.rs +++ b/crates/ecstore/src/data_movement.rs @@ -16,7 +16,7 @@ use crate::error::{Error, Result}; use crate::store::ECStore; use crate::store_api::{CompletePart, GetObjectReader, MultipartOperations, ObjectIO, ObjectInfo, ObjectOptions, PutObjReader}; use bytes::Bytes; -use rustfs_rio::{EtagResolvable, HashReader, HashReaderDetector, Index, Reader, TryGetIndex, WarpReader}; +use rustfs_rio::{EtagResolvable, HashReader, HashReaderDetector, Index, TryGetIndex}; use std::io::Cursor; use std::pin::Pin; use std::sync::{ @@ -54,8 +54,6 @@ impl TryGetIndex for IndexedDataMovementRead } } -impl Reader for IndexedDataMovementReader {} - pub fn decode_part_index(index: Option<&Bytes>) -> Option { let bytes = index?; let mut decoded = Index::new(); @@ -75,8 +73,8 @@ pub fn put_obj_reader_from_chunk(chunk: Vec, size: i64, actual_size: i64, in None }; - let reader = IndexedDataMovementReader::new(WarpReader::new(Cursor::new(chunk)), index); - let hash_reader = HashReader::new(Box::new(reader), size, actual_size, None, sha256hex, false)?; + let reader = IndexedDataMovementReader::new(Cursor::new(chunk), index); + let hash_reader = HashReader::from_stream(reader, size, actual_size, None, sha256hex, false)?; Ok(PutObjReader::new(hash_reader)) } @@ -255,8 +253,8 @@ pub(crate) async fn migrate_object( .parts .first() .and_then(|part| decode_part_index(part.index.as_ref())); - let reader = IndexedDataMovementReader::new(WarpReader::new(BufReader::new(rd.stream)), index); - let hrd = HashReader::new(Box::new(reader), object_info.size, actual_size, object_info.etag.clone(), None, false)?; + let reader = IndexedDataMovementReader::new(BufReader::new(rd.stream), index); + let hrd = HashReader::from_stream(reader, object_info.size, actual_size, object_info.etag.clone(), None, false)?; let mut data = PutObjReader::new(hrd); if let Err(err) = store diff --git a/crates/ecstore/src/set_disk.rs b/crates/ecstore/src/set_disk.rs index 5789d090d..f76fad893 100644 --- a/crates/ecstore/src/set_disk.rs +++ b/crates/ecstore/src/set_disk.rs @@ -78,7 +78,7 @@ use rustfs_lock::fast_lock::types::LockResult; use rustfs_lock::local_lock::LocalLock; use rustfs_lock::{FastLockGuard, NamespaceLock, NamespaceLockGuard, NamespaceLockWrapper, ObjectKey}; use rustfs_madmin::heal_commands::{HealDriveInfo, HealResultItem}; -use rustfs_rio::{EtagResolvable, HashReader, HashReaderMut, TryGetIndex as _, WarpReader}; +use rustfs_rio::{EtagResolvable, HashReader, HashReaderMut, TryGetIndex as _}; use rustfs_s3_common::EventName; use rustfs_utils::http::headers::AMZ_OBJECT_TAGGING; use rustfs_utils::http::headers::AMZ_STORAGE_CLASS; @@ -827,7 +827,7 @@ impl ObjectIO for SetDisks { let stream = mem::replace( &mut data.stream, - HashReader::new(Box::new(WarpReader::new(Cursor::new(Vec::new()))), 0, 0, None, None, false)?, + HashReader::from_stream(Cursor::new(Vec::new()), 0, 0, None, None, false)?, ); let (reader, w_size) = match Arc::new(erasure).encode(stream, &mut writers, write_quorum).await { @@ -1961,14 +1961,7 @@ impl ObjectOperations for SetDisks { } let gr = gr.unwrap(); let reader = BufReader::new(gr.stream); - let hash_reader = HashReader::new( - Box::new(WarpReader::new(reader)), - gr.object_info.size, - gr.object_info.size, - None, - None, - false, - )?; + let hash_reader = HashReader::from_stream(reader, gr.object_info.size, gr.object_info.size, None, None, false)?; let mut p_reader = PutObjReader::new(hash_reader); return match self_.clone().put_object(bucket, object, &mut p_reader, &ropts).await { Ok(restored_info) => { @@ -2036,14 +2029,7 @@ impl ObjectOperations for SetDisks { } }; let reader = BufReader::new(gr.stream); - let hash_reader = HashReader::new( - Box::new(WarpReader::new(reader)), - part_info.actual_size, - part_info.actual_size, - None, - None, - false, - )?; + let hash_reader = HashReader::from_stream(reader, part_info.actual_size, part_info.actual_size, None, None, false)?; let mut p_reader = PutObjReader::new(hash_reader); let p_info = self_ .clone() @@ -2349,7 +2335,7 @@ impl MultipartOperations for SetDisks { let stream = mem::replace( &mut data.stream, - HashReader::new(Box::new(WarpReader::new(Cursor::new(Vec::new()))), 0, 0, None, None, false)?, + HashReader::from_stream(Cursor::new(Vec::new()), 0, 0, None, None, false)?, ); let (reader, w_size) = Arc::new(erasure).encode(stream, &mut writers, write_quorum).await?; // TODO: delete temporary directory on error diff --git a/crates/ecstore/src/store_api.rs b/crates/ecstore/src/store_api.rs index 7ce1d2355..cdfde143b 100644 --- a/crates/ecstore/src/store_api.rs +++ b/crates/ecstore/src/store_api.rs @@ -34,7 +34,7 @@ use rustfs_filemeta::{ use rustfs_lock::NamespaceLockWrapper; use rustfs_madmin::heal_commands::HealResultItem; use rustfs_rio::Checksum; -use rustfs_rio::{DecompressReader, HashReader, LimitReader, WarpReader}; +use rustfs_rio::{DecompressReader, HashReader, LimitReader}; use rustfs_utils::CompressionAlgorithm; use rustfs_utils::http::headers::AMZ_OBJECT_TAGGING; use rustfs_utils::http::{AMZ_BUCKET_REPLICATION_STATUS, AMZ_RESTORE, AMZ_STORAGE_CLASS}; diff --git a/crates/ecstore/src/store_api/readers.rs b/crates/ecstore/src/store_api/readers.rs index dd32effb7..461e8ff7e 100644 --- a/crates/ecstore/src/store_api/readers.rs +++ b/crates/ecstore/src/store_api/readers.rs @@ -28,15 +28,7 @@ impl PutObjReader { None }; PutObjReader { - stream: HashReader::new( - Box::new(WarpReader::new(Cursor::new(data))), - content_length, - content_length, - None, - sha256hex, - false, - ) - .unwrap(), + stream: HashReader::from_stream(Cursor::new(data), content_length, content_length, None, sha256hex, false).unwrap(), } } diff --git a/crates/protocols/src/swift/object.rs b/crates/protocols/src/swift/object.rs index 7a5d0ecd4..1c59da07f 100644 --- a/crates/protocols/src/swift/object.rs +++ b/crates/protocols/src/swift/object.rs @@ -56,7 +56,7 @@ use axum::http::HeaderMap; use rustfs_credentials::Credentials; use rustfs_ecstore::new_object_layer_fn; use rustfs_ecstore::store_api::{BucketOperations, BucketOptions, ObjectIO, ObjectOperations, ObjectOptions, PutObjReader}; -use rustfs_rio::{HashReader, Reader, WarpReader}; +use rustfs_rio::HashReader; use std::collections::HashMap; use tracing::debug; use tracing::error; @@ -374,20 +374,12 @@ where ..Default::default() }; - // 13. Wrap reader in buffered reader then WarpReader (Box) + // 13. Wrap reader in buffered reader for streaming hash validation let buf_reader = tokio::io::BufReader::new(reader); - let warp_reader: Box = Box::new(WarpReader::new(buf_reader)); // 14. Create HashReader (no MD5/SHA256 validation for Swift) - let hash_reader = HashReader::new( - warp_reader, - content_length, - content_length, - None, // md5hex - None, // sha256hex - false, // disable_multipart - ) - .map_err(|e| sanitize_storage_error("Hash reader creation", e))?; + let hash_reader = HashReader::from_stream(buf_reader, content_length, content_length, None, None, false) + .map_err(|e| sanitize_storage_error("Hash reader creation", e))?; // 15. Wrap in PutObjReader as expected by storage layer let mut put_reader = PutObjReader::new(hash_reader); @@ -465,20 +457,12 @@ where // Content length (use -1 for unknown) let content_length = -1i64; - // Wrap reader in buffered reader then WarpReader + // Wrap reader in buffered reader for streaming hash validation let buf_reader = tokio::io::BufReader::new(reader); - let warp_reader: Box = Box::new(WarpReader::new(buf_reader)); // Create HashReader - let hash_reader = HashReader::new( - warp_reader, - content_length, - content_length, - None, // md5hex - None, // sha256hex - false, // disable_multipart - ) - .map_err(|e| sanitize_storage_error("Hash reader creation", e))?; + let hash_reader = HashReader::from_stream(buf_reader, content_length, content_length, None, None, false) + .map_err(|e| sanitize_storage_error("Hash reader creation", e))?; // Wrap in PutObjReader let mut put_reader = PutObjReader::new(hash_reader); diff --git a/crates/rio/src/compress_reader.rs b/crates/rio/src/compress_reader.rs index af92f8b36..418373a89 100644 --- a/crates/rio/src/compress_reader.rs +++ b/crates/rio/src/compress_reader.rs @@ -13,8 +13,6 @@ // limitations under the License. use crate::compress_index::{Index, TryGetIndex}; -use crate::{EtagResolvable, HashReaderDetector}; -use crate::{HashReaderMut, Reader}; use pin_project_lite::pin_project; use rustfs_utils::compress::{CompressionAlgorithm, compress_block, decompress_block}; use rustfs_utils::{put_uvarint, uvarint}; @@ -47,13 +45,13 @@ pin_project! { written: usize, uncomp_written: usize, temp_buffer: Vec, - temp_pos: usize, + read_buffer: Vec, } } impl CompressReader where - R: Reader, + R: AsyncRead + Unpin + Send + Sync, { pub fn new(inner: R, compression_algorithm: CompressionAlgorithm) -> Self { Self { @@ -66,8 +64,8 @@ where index: Index::new(), written: 0, uncomp_written: 0, - temp_buffer: Vec::with_capacity(DEFAULT_BLOCK_SIZE), // Pre-allocate capacity - temp_pos: 0, + temp_buffer: Vec::with_capacity(DEFAULT_BLOCK_SIZE), + read_buffer: vec![0u8; DEFAULT_BLOCK_SIZE], } } @@ -84,15 +82,12 @@ where written: 0, uncomp_written: 0, temp_buffer: Vec::with_capacity(block_size), - temp_pos: 0, + read_buffer: vec![0u8; block_size], } } } -impl TryGetIndex for CompressReader -where - R: Reader, -{ +impl TryGetIndex for CompressReader { fn try_get_index(&self) -> Option<&Index> { Some(&self.index) } @@ -121,8 +116,7 @@ where // Fill temporary buffer while this.temp_buffer.len() < *this.block_size { let remaining = *this.block_size - this.temp_buffer.len(); - let mut temp = vec![0u8; remaining]; - let mut temp_buf = ReadBuf::new(&mut temp); + let mut temp_buf = ReadBuf::new(&mut this.read_buffer[..remaining]); match this.inner.as_mut().poll_read(cx, &mut temp_buf) { Poll::Pending => { if this.temp_buffer.is_empty() { @@ -134,11 +128,12 @@ where let n = temp_buf.filled().len(); if n == 0 { if this.temp_buffer.is_empty() { + *this.done = true; return Poll::Ready(Ok(())); } break; } - this.temp_buffer.extend_from_slice(&temp[..n]); + this.temp_buffer.extend_from_slice(&temp_buf.filled()[..n]); } Poll::Ready(Err(e)) => { // error!("CompressReader poll_read: read inner error: {e}"); @@ -173,27 +168,7 @@ where } } -impl EtagResolvable for CompressReader -where - R: EtagResolvable, -{ - fn try_resolve_etag(&mut self) -> Option { - self.inner.try_resolve_etag() - } -} - -impl HashReaderDetector for CompressReader -where - R: HashReaderDetector, -{ - fn is_hash_reader(&self) -> bool { - self.inner.is_hash_reader() - } - - fn as_hash_reader_mut(&mut self) -> Option<&mut dyn HashReaderMut> { - self.inner.as_hash_reader_mut() - } -} +delegate_reader_capabilities_generic_no_index!(CompressReader, inner); pin_project! { /// A reader wrapper that decompresses data on the fly using DEFLATE algorithm. @@ -213,7 +188,7 @@ pin_project! { header_read: usize, header_done: bool, // Fields for saving compressed block read progress across polls - compressed_buf: Option>, + compressed_buf: Vec, compressed_read: usize, compressed_len: usize, compression_algorithm: CompressionAlgorithm, @@ -233,7 +208,7 @@ where header_buf: [0u8; 8], header_read: 0, header_done: false, - compressed_buf: None, + compressed_buf: Vec::new(), compressed_read: 0, compressed_len: 0, compression_algorithm, @@ -295,14 +270,22 @@ where | ((this.header_buf[7] as u32) << 24); *this.header_read = 0; *this.header_done = true; - if this.compressed_buf.is_none() { - *this.compressed_len = len; - *this.compressed_buf = Some(vec![0u8; *this.compressed_len]); + + if typ == COMPRESS_TYPE_END { *this.compressed_read = 0; + *this.compressed_len = 0; + *this.finished = true; + return Poll::Ready(Ok(())); } - let compressed_buf = this.compressed_buf.as_mut().unwrap(); + + if this.compressed_buf.len() < len { + this.compressed_buf.resize(len, 0); + } + *this.compressed_len = len; + *this.compressed_read = 0; + while *this.compressed_read < *this.compressed_len { - let mut temp_buf = ReadBuf::new(&mut compressed_buf[*this.compressed_read..]); + let mut temp_buf = ReadBuf::new(&mut this.compressed_buf[*this.compressed_read..*this.compressed_len]); match this.inner.as_mut().poll_read(cx, &mut temp_buf) { Poll::Pending => return Poll::Pending, Poll::Ready(Ok(())) => { @@ -314,13 +297,13 @@ where } Poll::Ready(Err(e)) => { // error!("DecompressReader poll_read: read compressed block error: {e}"); - this.compressed_buf.take(); *this.compressed_read = 0; *this.compressed_len = 0; return Poll::Ready(Err(e)); } } } + let compressed_buf = &this.compressed_buf[..*this.compressed_len]; let (uncompress_len, uvarint) = uvarint(&compressed_buf[0..16]); let compressed_data = &compressed_buf[uvarint as usize..]; let decompressed = if typ == COMPRESS_TYPE_COMPRESSED { @@ -328,7 +311,6 @@ where Ok(out) => out, Err(e) => { // error!("DecompressReader decompress_block error: {e}"); - this.compressed_buf.take(); *this.compressed_read = 0; *this.compressed_len = 0; return Poll::Ready(Err(e)); @@ -336,22 +318,14 @@ where } } else if typ == COMPRESS_TYPE_UNCOMPRESSED { compressed_data.to_vec() - } else if typ == COMPRESS_TYPE_END { - this.compressed_buf.take(); - *this.compressed_read = 0; - *this.compressed_len = 0; - *this.finished = true; - return Poll::Ready(Ok(())); } else { // error!("DecompressReader unknown compression type: {typ}"); - this.compressed_buf.take(); *this.compressed_read = 0; *this.compressed_len = 0; return Poll::Ready(Err(io::Error::new(io::ErrorKind::InvalidData, "Unknown compression type"))); }; if decompressed.len() != uncompress_len as usize { // error!("DecompressReader decompressed length mismatch: {} != {}", decompressed.len(), uncompress_len); - this.compressed_buf.take(); *this.compressed_read = 0; *this.compressed_len = 0; return Poll::Ready(Err(io::Error::new(io::ErrorKind::InvalidData, "Decompressed length mismatch"))); @@ -363,14 +337,12 @@ where }; if actual_crc != crc { // error!("DecompressReader CRC32 mismatch: actual {actual_crc} != expected {crc}"); - this.compressed_buf.take(); *this.compressed_read = 0; *this.compressed_len = 0; return Poll::Ready(Err(io::Error::new(io::ErrorKind::InvalidData, "CRC32 mismatch"))); } *this.buffer = decompressed; *this.buffer_pos = 0; - this.compressed_buf.take(); *this.compressed_read = 0; *this.compressed_len = 0; *this.header_done = false; @@ -385,26 +357,7 @@ where } } -impl EtagResolvable for DecompressReader -where - R: EtagResolvable, -{ - fn try_resolve_etag(&mut self) -> Option { - self.inner.try_resolve_etag() - } -} - -impl HashReaderDetector for DecompressReader -where - R: HashReaderDetector, -{ - fn is_hash_reader(&self) -> bool { - self.inner.is_hash_reader() - } - fn as_hash_reader_mut(&mut self) -> Option<&mut dyn HashReaderMut> { - self.inner.as_hash_reader_mut() - } -} +delegate_reader_capabilities_generic_no_index!(DecompressReader, inner); /// Build compressed block with header + uvarint + compressed data fn build_compressed_block(uncompressed_data: &[u8], compression_algorithm: CompressionAlgorithm) -> Vec { @@ -436,8 +389,6 @@ fn build_compressed_block(uncompressed_data: &[u8], compression_algorithm: Compr #[cfg(test)] mod tests { - use crate::WarpReader; - use super::*; use rand::RngExt; use std::io::Cursor; @@ -447,7 +398,7 @@ mod tests { async fn test_compress_reader_basic() { let data = b"hello world, hello world, hello world!"; let reader = Cursor::new(&data[..]); - let mut compress_reader = CompressReader::new(WarpReader::new(reader), CompressionAlgorithm::Gzip); + let mut compress_reader = CompressReader::new(reader, CompressionAlgorithm::Gzip); let mut compressed = Vec::new(); compress_reader.read_to_end(&mut compressed).await.unwrap(); @@ -464,7 +415,7 @@ mod tests { async fn test_compress_reader_basic_deflate() { let data = b"hello world, hello world, hello world!"; let reader = BufReader::new(&data[..]); - let mut compress_reader = CompressReader::new(WarpReader::new(reader), CompressionAlgorithm::Deflate); + let mut compress_reader = CompressReader::new(reader, CompressionAlgorithm::Deflate); let mut compressed = Vec::new(); compress_reader.read_to_end(&mut compressed).await.unwrap(); @@ -481,7 +432,7 @@ mod tests { async fn test_compress_reader_empty() { let data = b""; let reader = BufReader::new(&data[..]); - let mut compress_reader = CompressReader::new(WarpReader::new(reader), CompressionAlgorithm::Gzip); + let mut compress_reader = CompressReader::new(reader, CompressionAlgorithm::Gzip); let mut compressed = Vec::new(); compress_reader.read_to_end(&mut compressed).await.unwrap(); @@ -499,7 +450,7 @@ mod tests { let mut data = vec![0u8; 1024 * 1024 * 32]; rand::rng().fill(&mut data[..]); let reader = Cursor::new(data.clone()); - let mut compress_reader = CompressReader::new(WarpReader::new(reader), CompressionAlgorithm::Gzip); + let mut compress_reader = CompressReader::new(reader, CompressionAlgorithm::Gzip); let mut compressed = Vec::new(); compress_reader.read_to_end(&mut compressed).await.unwrap(); @@ -517,7 +468,7 @@ mod tests { let mut data = vec![0u8; 1024 * 1024 * 3 + 512]; rand::rng().fill(&mut data[..]); let reader = Cursor::new(data.clone()); - let mut compress_reader = CompressReader::new(WarpReader::new(reader), CompressionAlgorithm::default()); + let mut compress_reader = CompressReader::new(reader, CompressionAlgorithm::default()); let mut compressed = Vec::new(); compress_reader.read_to_end(&mut compressed).await.unwrap(); diff --git a/crates/rio/src/encrypt_reader.rs b/crates/rio/src/encrypt_reader.rs index 4f1f39664..4b8e275cf 100644 --- a/crates/rio/src/encrypt_reader.rs +++ b/crates/rio/src/encrypt_reader.rs @@ -12,10 +12,7 @@ // See the License for the specific language governing permissions and // limitations under the License. -use crate::HashReaderDetector; -use crate::HashReaderMut; use crate::compress_index::{Index, TryGetIndex}; -use crate::{EtagResolvable, Reader}; use aes_gcm::aead::Aead; use aes_gcm::{Aes256Gcm, KeyInit, Nonce}; use pin_project_lite::pin_project; @@ -26,32 +23,37 @@ use std::task::{Context, Poll}; use tokio::io::{AsyncRead, ReadBuf}; use tracing::debug; +const ENCRYPTION_BLOCK_SIZE: usize = 8 * 1024; + pin_project! { /// A reader wrapper that encrypts data on the fly using AES-256-GCM. /// This is a demonstration. For production, use a secure and audited crypto library. - #[derive(Debug)] pub struct EncryptReader { #[pin] pub inner: R, - key: [u8; 32], // AES-256-GCM key - nonce: [u8; 12], // 96-bit nonce for GCM + cipher: Aes256Gcm, + base_nonce: [u8; 12], // 96-bit base nonce for GCM buffer: Vec, buffer_pos: usize, + read_buffer: Vec, + block_index: usize, finished: bool, } } impl EncryptReader where - R: Reader, + R: AsyncRead + Unpin + Send + Sync, { pub fn new(inner: R, key: [u8; 32], nonce: [u8; 12]) -> Self { Self { inner, - key, - nonce, + cipher: Aes256Gcm::new_from_slice(&key).expect("key"), + base_nonce: nonce, buffer: Vec::new(), buffer_pos: 0, + read_buffer: vec![0u8; ENCRYPTION_BLOCK_SIZE], + block_index: 0, finished: false, } } @@ -77,10 +79,8 @@ where if *this.finished { return Poll::Ready(Ok(())); } - // Read a fixed block size from inner - let block_size = 8 * 1024; - let mut temp = vec![0u8; block_size]; - let mut temp_buf = ReadBuf::new(&mut temp); + // Read a fixed block size from inner. + let mut temp_buf = ReadBuf::new(&mut this.read_buffer[..]); match this.inner.as_mut().poll_read(cx, &mut temp_buf) { Poll::Pending => Poll::Pending, Poll::Ready(Ok(())) => { @@ -98,16 +98,17 @@ where Poll::Ready(Ok(())) } else { // Encrypt the chunk - let cipher = Aes256Gcm::new_from_slice(this.key).expect("key"); - let nonce = Nonce::try_from(this.nonce.as_slice()).map_err(|_| Error::other("invalid nonce length"))?; - let plaintext = &temp_buf.filled()[..n]; + let block_nonce = derive_block_nonce(this.base_nonce, *this.block_index); + let nonce = Nonce::try_from(block_nonce.as_slice()).map_err(|_| Error::other("invalid nonce length"))?; + let plaintext = &this.read_buffer[..n]; let plaintext_len = plaintext.len(); let crc = { let mut hasher = crc_fast::Digest::new(crc_fast::CrcAlgorithm::Crc32IsoHdlc); hasher.update(plaintext); hasher.finalize() as u32 }; - let ciphertext = cipher + let ciphertext = this + .cipher .encrypt(&nonce, plaintext) .map_err(|e| Error::other(format!("encrypt error: {e}")))?; let int_len = put_uvarint_len(plaintext_len as u64); @@ -134,12 +135,13 @@ where ); let mut out = Vec::with_capacity(8 + int_len + ciphertext.len()); out.extend_from_slice(&header); - let mut plaintext_len_buf = vec![0u8; int_len]; - put_uvarint(&mut plaintext_len_buf, plaintext_len as u64); - out.extend_from_slice(&plaintext_len_buf); + let mut plaintext_len_buf = [0u8; 10]; + let encoded_len = put_uvarint(&mut plaintext_len_buf, plaintext_len as u64); + out.extend_from_slice(&plaintext_len_buf[..encoded_len]); out.extend_from_slice(&ciphertext); *this.buffer = out; *this.buffer_pos = 0; + *this.block_index += 1; let to_copy = std::cmp::min(buf.remaining(), this.buffer.len()); buf.put_slice(&this.buffer[..to_copy]); *this.buffer_pos += to_copy; @@ -151,27 +153,7 @@ where } } -impl EtagResolvable for EncryptReader -where - R: EtagResolvable, -{ - fn try_resolve_etag(&mut self) -> Option { - self.inner.try_resolve_etag() - } -} - -impl HashReaderDetector for EncryptReader -where - R: EtagResolvable + HashReaderDetector, -{ - fn is_hash_reader(&self) -> bool { - self.inner.is_hash_reader() - } - - fn as_hash_reader_mut(&mut self) -> Option<&mut dyn HashReaderMut> { - self.inner.as_hash_reader_mut() - } -} +delegate_reader_capabilities_generic_no_index!(EncryptReader, inner); impl TryGetIndex for EncryptReader where @@ -185,15 +167,15 @@ where pin_project! { /// A reader wrapper that decrypts data on the fly using AES-256-GCM. /// This is a demonstration. For production, use a secure and audited crypto library. -#[derive(Debug)] pub struct DecryptReader { #[pin] pub inner: R, - key: [u8; 32], // AES-256-GCM key + cipher: Aes256Gcm, base_nonce: [u8; 12], // Base nonce recorded in object metadata - current_nonce: [u8; 12], // Active nonce for the current encrypted segment + current_nonce_base: [u8; 12], // Active base nonce for the current encrypted segment multipart_mode: bool, current_part: usize, + block_index: usize, buffer: Vec, buffer_pos: usize, finished: bool, @@ -201,7 +183,7 @@ pin_project! { header_buf: [u8; 8], header_read: usize, header_done: bool, - ciphertext_buf: Option>, + ciphertext_buf: Vec, ciphertext_read: usize, ciphertext_len: usize, } @@ -209,23 +191,24 @@ pin_project! { impl DecryptReader where - R: Reader, + R: AsyncRead + Unpin + Send + Sync, { pub fn new(inner: R, key: [u8; 32], nonce: [u8; 12]) -> Self { Self { inner, - key, + cipher: Aes256Gcm::new_from_slice(&key).expect("key"), base_nonce: nonce, - current_nonce: nonce, + current_nonce_base: nonce, multipart_mode: false, current_part: 0, + block_index: 0, buffer: Vec::new(), buffer_pos: 0, finished: false, header_buf: [0u8; 8], header_read: 0, header_done: false, - ciphertext_buf: None, + ciphertext_buf: Vec::new(), ciphertext_read: 0, ciphertext_len: 0, } @@ -239,18 +222,19 @@ where Self { inner, - key, + cipher: Aes256Gcm::new_from_slice(&key).expect("key"), base_nonce, - current_nonce: initial_nonce, + current_nonce_base: initial_nonce, multipart_mode: true, current_part: first_part, + block_index: 0, buffer: Vec::new(), buffer_pos: 0, finished: false, header_buf: [0u8; 8], header_read: 0, header_done: false, - ciphertext_buf: None, + ciphertext_buf: Vec::new(), ciphertext_read: 0, ciphertext_len: 0, } @@ -332,15 +316,14 @@ where "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.current_nonce_base = derive_part_nonce(this.base_nonce, *this.current_part); + *this.block_index = 0; *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; @@ -351,7 +334,6 @@ where if len == 0 { tracing::warn!("encountered zero-length encrypted block, treating as end of stream"); *this.finished = true; - this.ciphertext_buf.take(); *this.ciphertext_read = 0; *this.ciphertext_len = 0; continue; @@ -362,15 +344,14 @@ where return Poll::Ready(Err(Error::other("Invalid encrypted block length"))); }; - if this.ciphertext_buf.is_none() { - *this.ciphertext_buf = Some(vec![0u8; payload_len]); - *this.ciphertext_len = payload_len; - *this.ciphertext_read = 0; + if this.ciphertext_buf.len() < payload_len { + this.ciphertext_buf.resize(payload_len, 0); } + *this.ciphertext_len = payload_len; + *this.ciphertext_read = 0; - let ciphertext_buf = this.ciphertext_buf.as_mut().unwrap(); while *this.ciphertext_read < *this.ciphertext_len { - let mut temp_buf = ReadBuf::new(&mut ciphertext_buf[*this.ciphertext_read..]); + let mut temp_buf = ReadBuf::new(&mut this.ciphertext_buf[*this.ciphertext_read..*this.ciphertext_len]); match this.inner.as_mut().poll_read(cx, &mut temp_buf) { Poll::Pending => return Poll::Pending, Poll::Ready(Ok(())) => { @@ -384,7 +365,6 @@ where *this.ciphertext_read += n; } Poll::Ready(Err(e)) => { - this.ciphertext_buf.take(); *this.ciphertext_read = 0; *this.ciphertext_len = 0; return Poll::Ready(Err(e)); @@ -396,14 +376,37 @@ where return Poll::Pending; } + let ciphertext_buf = &this.ciphertext_buf[..*this.ciphertext_len]; let (plaintext_len, uvarint_len) = rustfs_utils::uvarint(&ciphertext_buf[0..16]); let ciphertext = &ciphertext_buf[uvarint_len as usize..]; + let block_nonce = derive_block_nonce(this.current_nonce_base, *this.block_index); + let nonce = Nonce::try_from(block_nonce.as_slice()).map_err(|_| Error::other("invalid nonce length"))?; + let legacy_part_nonce = if *this.multipart_mode { + derive_legacy_part_nonce(this.base_nonce, *this.current_part) + } else { + *this.base_nonce + }; + let legacy_block_nonce = derive_block_nonce(&legacy_part_nonce, *this.block_index); + let plaintext = match this.cipher.decrypt(&nonce, ciphertext) { + Ok(plaintext) => plaintext, + Err(primary_err) => { + let legacy_nonce = + Nonce::try_from(legacy_block_nonce.as_slice()).map_err(|_| Error::other("invalid nonce length"))?; - let cipher = Aes256Gcm::new_from_slice(this.key).expect("key"); - let nonce = Nonce::try_from(this.current_nonce.as_slice()).map_err(|_| Error::other("invalid nonce length"))?; - let plaintext = cipher - .decrypt(&nonce, ciphertext) - .map_err(|e| Error::other(format!("decrypt error: {e}")))?; + match this.cipher.decrypt(&legacy_nonce, ciphertext) { + Ok(plaintext) => plaintext, + Err(_) => { + // Accept previously written streams that reused the part nonce + // for every block inside a segment. + let legacy_part_nonce = Nonce::try_from(legacy_part_nonce.as_slice()) + .map_err(|_| Error::other("invalid nonce length"))?; + this.cipher + .decrypt(&legacy_part_nonce, ciphertext) + .map_err(|_| Error::other(format!("decrypt error: {primary_err}")))? + } + } + } + }; debug!( part = *this.current_part, @@ -412,7 +415,6 @@ where ); if plaintext.len() != plaintext_len as usize { - this.ciphertext_buf.take(); *this.ciphertext_read = 0; *this.ciphertext_len = 0; return Poll::Ready(Err(Error::other("Plaintext length mismatch"))); @@ -424,7 +426,6 @@ where hasher.finalize() as u32 }; if actual_crc != crc { - this.ciphertext_buf.take(); *this.ciphertext_read = 0; *this.ciphertext_len = 0; return Poll::Ready(Err(Error::other("CRC32 mismatch"))); @@ -432,7 +433,7 @@ where *this.buffer = plaintext; *this.buffer_pos = 0; - this.ciphertext_buf.take(); + *this.block_index += 1; *this.ciphertext_read = 0; *this.ciphertext_len = 0; @@ -444,27 +445,7 @@ where } } -impl EtagResolvable for DecryptReader -where - R: EtagResolvable, -{ - fn try_resolve_etag(&mut self) -> Option { - self.inner.try_resolve_etag() - } -} - -impl HashReaderDetector for DecryptReader -where - R: EtagResolvable + HashReaderDetector, -{ - fn is_hash_reader(&self) -> bool { - self.inner.is_hash_reader() - } - - fn as_hash_reader_mut(&mut self) -> Option<&mut dyn HashReaderMut> { - self.inner.as_hash_reader_mut() - } -} +delegate_reader_capabilities_generic_no_index!(DecryptReader, inner); impl TryGetIndex for DecryptReader where @@ -475,23 +456,37 @@ where } } +fn derive_block_nonce(base: &[u8; 12], block_index: usize) -> [u8; 12] { + derive_nonce_offset(base, 8, block_index) +} + fn derive_part_nonce(base: &[u8; 12], part_number: usize) -> [u8; 12] { + derive_nonce_offset(base, 4, part_number) +} + +fn derive_legacy_part_nonce(base: &[u8; 12], part_number: usize) -> [u8; 12] { + derive_nonce_offset(base, 8, part_number) +} + +fn derive_nonce_offset(base: &[u8; 12], start: usize, offset: usize) -> [u8; 12] { let mut nonce = *base; let mut suffix = [0u8; 4]; - suffix.copy_from_slice(&nonce[8..12]); + suffix.copy_from_slice(&nonce[start..start + 4]); 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()); + let next = current.wrapping_add(offset as u32); + nonce[start..start + 4].copy_from_slice(&next.to_be_bytes()); nonce } #[cfg(test)] mod tests { + use aes_gcm::aead::Aead; + use aes_gcm::{Aes256Gcm, KeyInit, Nonce}; use std::io::Cursor; use std::pin::Pin; use std::task::{Context, Poll}; - use crate::{HardLimitReader, WarpReader}; + use crate::HardLimitReader; use super::*; use futures::StreamExt; @@ -533,6 +528,73 @@ mod tests { } } + fn encrypt_with_legacy_nonce_reuse(data: &[u8], key: [u8; 32], nonce: [u8; 12]) -> Vec { + let cipher = Aes256Gcm::new_from_slice(&key).expect("valid key"); + let nonce = Nonce::try_from(nonce.as_slice()).expect("valid nonce"); + let mut encrypted = Vec::new(); + + for chunk in data.chunks(ENCRYPTION_BLOCK_SIZE) { + let crc = { + let mut hasher = crc_fast::Digest::new(crc_fast::CrcAlgorithm::Crc32IsoHdlc); + hasher.update(chunk); + hasher.finalize() as u32 + }; + let ciphertext = cipher.encrypt(&nonce, chunk).expect("legacy encrypt"); + let int_len = put_uvarint_len(chunk.len() as u64); + let clen = int_len + ciphertext.len() + 4; + let mut header = [0u8; 8]; + header[1] = (clen & 0xFF) as u8; + header[2] = ((clen >> 8) & 0xFF) as u8; + header[3] = ((clen >> 16) & 0xFF) as u8; + header[4] = (crc & 0xFF) as u8; + header[5] = ((crc >> 8) & 0xFF) as u8; + header[6] = ((crc >> 16) & 0xFF) as u8; + header[7] = ((crc >> 24) & 0xFF) as u8; + encrypted.extend_from_slice(&header); + let mut plaintext_len_buf = [0u8; 10]; + let encoded_len = put_uvarint(&mut plaintext_len_buf, chunk.len() as u64); + encrypted.extend_from_slice(&plaintext_len_buf[..encoded_len]); + encrypted.extend_from_slice(&ciphertext); + } + + encrypted.extend_from_slice(&[0xFF, 0, 0, 0, 0, 0, 0, 0]); + encrypted + } + + async fn encrypt_part_with_legacy_nonce_layout( + data: &[u8], + key: [u8; 32], + base_nonce: [u8; 12], + part_number: usize, + ) -> Vec { + let nonce = derive_legacy_part_nonce(&base_nonce, part_number); + let reader = BufReader::new(Cursor::new(data.to_vec())); + let mut encrypt_reader = EncryptReader::new(reader, key, nonce); + let mut encrypted = Vec::new(); + encrypt_reader.read_to_end(&mut encrypted).await.unwrap(); + encrypted + } + + fn extract_encrypted_payloads(encrypted: &[u8]) -> Vec> { + let mut payloads = Vec::new(); + let mut pos = 0; + + while pos + 8 <= encrypted.len() { + let header = &encrypted[pos..pos + 8]; + pos += 8; + if header[0] == 0xFF { + break; + } + + let len = (header[1] as usize) | ((header[2] as usize) << 8) | ((header[3] as usize) << 16); + let payload_len = len - 4; + payloads.push(encrypted[pos..pos + payload_len].to_vec()); + pos += payload_len; + } + + payloads + } + #[tokio::test] async fn test_encrypt_decrypt_reader_aes256gcm() { let data = b"hello sse encrypt"; @@ -542,7 +604,7 @@ mod tests { rand::rng().fill_bytes(&mut nonce); let reader = BufReader::new(&data[..]); - let encrypt_reader = EncryptReader::new(WarpReader::new(reader), key, nonce); + let encrypt_reader = EncryptReader::new(reader, key, nonce); // Encrypt let mut encrypt_reader = encrypt_reader; @@ -551,7 +613,7 @@ mod tests { // Decrypt using DecryptReader let reader = Cursor::new(encrypted.clone()); - let decrypt_reader = DecryptReader::new(WarpReader::new(reader), key, nonce); + let decrypt_reader = DecryptReader::new(reader, key, nonce); let mut decrypt_reader = decrypt_reader; let mut decrypted = Vec::new(); decrypt_reader.read_to_end(&mut decrypted).await.unwrap(); @@ -570,7 +632,7 @@ mod tests { // Encrypt let reader = BufReader::new(&data[..]); - let encrypt_reader = EncryptReader::new(WarpReader::new(reader), key, nonce); + let encrypt_reader = EncryptReader::new(reader, key, nonce); let mut encrypt_reader = encrypt_reader; let mut encrypted = Vec::new(); encrypt_reader.read_to_end(&mut encrypted).await.unwrap(); @@ -578,7 +640,7 @@ mod tests { // Now test DecryptReader let reader = Cursor::new(encrypted.clone()); - let decrypt_reader = DecryptReader::new(WarpReader::new(reader), key, nonce); + let decrypt_reader = DecryptReader::new(reader, key, nonce); let mut decrypt_reader = decrypt_reader; let mut decrypted = Vec::new(); decrypt_reader.read_to_end(&mut decrypted).await.unwrap(); @@ -598,13 +660,13 @@ mod tests { rand::rng().fill_bytes(&mut nonce); let reader = std::io::Cursor::new(data.clone()); - let encrypt_reader = EncryptReader::new(WarpReader::new(reader), key, nonce); + let encrypt_reader = EncryptReader::new(reader, key, nonce); let mut encrypt_reader = encrypt_reader; let mut encrypted = Vec::new(); encrypt_reader.read_to_end(&mut encrypted).await.unwrap(); let reader = std::io::Cursor::new(encrypted.clone()); - let decrypt_reader = DecryptReader::new(WarpReader::new(reader), key, nonce); + let decrypt_reader = DecryptReader::new(reader, key, nonce); let mut decrypt_reader = decrypt_reader; let mut decrypted = Vec::new(); decrypt_reader.read_to_end(&mut decrypted).await.unwrap(); @@ -623,12 +685,12 @@ mod tests { rand::rng().fill_bytes(&mut nonce); let reader = Cursor::new(data.clone()); - let mut encrypt_reader = EncryptReader::new(WarpReader::new(reader), key, nonce); + let mut encrypt_reader = EncryptReader::new(reader, key, nonce); let mut encrypted = Vec::new(); encrypt_reader.read_to_end(&mut encrypted).await.unwrap(); let reader = ChunkedCursor::new(encrypted, 3); - let mut decrypt_reader = DecryptReader::new(WarpReader::new(reader), key, nonce); + let mut decrypt_reader = DecryptReader::new(reader, key, nonce); let mut decrypted = Vec::new(); decrypt_reader.read_to_end(&mut decrypted).await.unwrap(); @@ -646,12 +708,12 @@ mod tests { rand::rng().fill_bytes(&mut nonce); let reader = Cursor::new(data.clone()); - let mut encrypt_reader = EncryptReader::new(WarpReader::new(reader), key, nonce); + let mut encrypt_reader = EncryptReader::new(reader, key, nonce); let mut encrypted = Vec::new(); encrypt_reader.read_to_end(&mut encrypted).await.unwrap(); let reader = ChunkedCursor::new(encrypted, 8192); - let decrypt_reader = DecryptReader::new(WarpReader::new(reader), key, nonce); + let decrypt_reader = DecryptReader::new(reader, key, nonce); let mut stream = ReaderStream::with_capacity(Box::new(decrypt_reader), 262_144); let mut decrypted = Vec::new(); @@ -674,13 +736,13 @@ mod tests { rand::rng().fill_bytes(&mut nonce); let reader = Cursor::new(data.clone()); - let mut encrypt_reader = EncryptReader::new(WarpReader::new(reader), key, nonce); + let mut encrypt_reader = EncryptReader::new(reader, key, nonce); let mut encrypted = Vec::new(); encrypt_reader.read_to_end(&mut encrypted).await.unwrap(); let reader = ChunkedCursor::new(encrypted, 8192); - let decrypt_reader = DecryptReader::new(WarpReader::new(reader), key, nonce); - let limit_reader = HardLimitReader::new(Box::new(decrypt_reader), size as i64); + let decrypt_reader = DecryptReader::new(reader, key, nonce); + let limit_reader = HardLimitReader::new(decrypt_reader, size as i64); let mut stream = ReaderStream::with_capacity(Box::new(limit_reader), 262_144); let mut decrypted = Vec::new(); @@ -705,7 +767,7 @@ mod tests { async fn encrypt_part(data: &[u8], key: [u8; 32], base_nonce: [u8; 12], part_number: usize) -> Vec { 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 encrypt_reader = EncryptReader::new(reader, key, nonce); let mut encrypted = Vec::new(); encrypt_reader.read_to_end(&mut encrypted).await.unwrap(); encrypted @@ -719,7 +781,81 @@ mod tests { 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 decrypt_reader = DecryptReader::new_multipart(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); + } + + #[tokio::test] + async fn test_encrypt_reader_uses_distinct_nonces_per_block() { + let data = vec![0xAB; ENCRYPTION_BLOCK_SIZE * 2]; + let mut key = [0u8; 32]; + let mut nonce = [0u8; 12]; + rand::rng().fill_bytes(&mut key); + rand::rng().fill_bytes(&mut nonce); + + let reader = Cursor::new(data); + let mut encrypt_reader = EncryptReader::new(reader, key, nonce); + let mut encrypted = Vec::new(); + encrypt_reader.read_to_end(&mut encrypted).await.unwrap(); + + let payloads = extract_encrypted_payloads(&encrypted); + assert!(payloads.len() >= 2); + assert_ne!(payloads[0], payloads[1]); + } + + #[test] + fn test_part_and_block_nonces_do_not_collide_across_parts() { + let base_nonce = [0u8; 12]; + let part_one_block_one = derive_block_nonce(&derive_part_nonce(&base_nonce, 1), 1); + let part_two_block_zero = derive_block_nonce(&derive_part_nonce(&base_nonce, 2), 0); + + assert_ne!(part_one_block_one, part_two_block_zero); + } + + #[tokio::test] + async fn test_decrypt_reader_accepts_legacy_single_nonce_streams() { + let mut data = vec![0u8; ENCRYPTION_BLOCK_SIZE * 3 + 17]; + rand::rng().fill(&mut data[..]); + let mut key = [0u8; 32]; + let mut nonce = [0u8; 12]; + rand::rng().fill_bytes(&mut key); + rand::rng().fill_bytes(&mut nonce); + + let encrypted = encrypt_with_legacy_nonce_reuse(&data, key, nonce); + let reader = Cursor::new(encrypted); + let mut decrypt_reader = DecryptReader::new(reader, key, nonce); + let mut decrypted = Vec::new(); + decrypt_reader.read_to_end(&mut decrypted).await.unwrap(); + + assert_eq!(decrypted, data); + } + + #[tokio::test] + async fn test_decrypt_reader_accepts_legacy_multipart_nonce_layout() { + 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![0x11; ENCRYPTION_BLOCK_SIZE + 97]; + let part_two = vec![0x22; ENCRYPTION_BLOCK_SIZE + 33]; + + let encrypted_one = encrypt_part_with_legacy_nonce_layout(&part_one, key, base_nonce, 1).await; + let encrypted_two = encrypt_part_with_legacy_nonce_layout(&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(reader, key, base_nonce); let mut decrypted = Vec::new(); decrypt_reader.read_to_end(&mut decrypted).await.unwrap(); diff --git a/crates/rio/src/etag.rs b/crates/rio/src/etag.rs index 90428a4ba..6337deef5 100644 --- a/crates/rio/src/etag.rs +++ b/crates/rio/src/etag.rs @@ -31,7 +31,6 @@ The `EtagResolvable` trait provides a clean way to handle recursive unwrapping: ```rust use rustfs_rio::{CompressReader, EtagReader, resolve_etag_generic}; -use rustfs_rio::WarpReader; use rustfs_utils::compress::CompressionAlgorithm; use tokio::io::BufReader; use std::io::Cursor; @@ -39,7 +38,6 @@ use std::io::Cursor; // Direct usage with trait-based approach let data = b"test data"; let reader = BufReader::new(Cursor::new(&data[..])); -let reader = Box::new(WarpReader::new(reader)); let etag_reader = EtagReader::new(reader, Some("test_etag".to_string())); let mut reader = CompressReader::new(etag_reader, CompressionAlgorithm::Gzip); let etag = resolve_etag_generic(&mut reader); @@ -49,8 +47,8 @@ let etag = resolve_etag_generic(&mut reader); #[cfg(test)] mod tests { + use crate::resolve_etag_generic; use crate::{CompressReader, EncryptReader, EtagReader, HashReader}; - use crate::{WarpReader, resolve_etag_generic}; use md5::Md5; use rustfs_utils::compress::CompressionAlgorithm; use std::io::Cursor; @@ -60,7 +58,6 @@ mod tests { fn test_etag_reader_resolution() { let data = b"test data"; let reader = BufReader::new(Cursor::new(&data[..])); - let reader = Box::new(WarpReader::new(reader)); let mut etag_reader = EtagReader::new(reader, Some("test_etag".to_string())); // Test direct ETag resolution @@ -71,9 +68,9 @@ mod tests { fn test_hash_reader_resolution() { let data = b"test data"; let reader = BufReader::new(Cursor::new(&data[..])); - let reader = Box::new(WarpReader::new(reader)); let mut hash_reader = - HashReader::new(reader, data.len() as i64, data.len() as i64, Some("hash_etag".to_string()), None, false).unwrap(); + HashReader::from_stream(reader, data.len() as i64, data.len() as i64, Some("hash_etag".to_string()), None, false) + .unwrap(); // Test HashReader ETag resolution assert_eq!(resolve_etag_generic(&mut hash_reader), Some("hash_etag".to_string())); @@ -83,7 +80,6 @@ mod tests { fn test_compress_reader_delegation() { let data = b"test data for compression"; let reader = BufReader::new(Cursor::new(&data[..])); - let reader = Box::new(WarpReader::new(reader)); let etag_reader = EtagReader::new(reader, Some("compress_etag".to_string())); let mut compress_reader = CompressReader::new(etag_reader, CompressionAlgorithm::Gzip); @@ -95,7 +91,6 @@ mod tests { fn test_encrypt_reader_delegation() { let data = b"test data for encryption"; let reader = BufReader::new(Cursor::new(&data[..])); - let reader = Box::new(WarpReader::new(reader)); let etag_reader = EtagReader::new(reader, Some("encrypt_etag".to_string())); let key = [0u8; 32]; @@ -118,7 +113,6 @@ mod tests { let etag_hex = hex_simd::encode_to_string(etag, hex_simd::AsciiCase::Lower); let reader = BufReader::new(Cursor::new(&data[..])); - let reader = Box::new(WarpReader::new(reader)); // Create a complex nested structure: CompressReader>>> let etag_reader = EtagReader::new(reader, Some(etag_hex.clone())); let key = [0u8; 32]; @@ -136,9 +130,8 @@ mod tests { fn test_hash_reader_in_nested_structure() { let data = b"test data for hash reader nesting"; let reader = BufReader::new(Cursor::new(&data[..])); - let reader = Box::new(WarpReader::new(reader)); // Create nested structure: CompressReader>> - let hash_reader = HashReader::new( + let hash_reader = HashReader::from_stream( reader, data.len() as i64, data.len() as i64, @@ -166,7 +159,6 @@ mod tests { let etag = hasher.finalize(); let etag_hex = hex_simd::encode_to_string(etag, hex_simd::AsciiCase::Lower); let reader1 = BufReader::new(Cursor::new(&data1[..])); - let reader1 = Box::new(WarpReader::new(reader1)); let mut etag_reader = EtagReader::new(reader1, Some(etag_hex.clone())); etag_reader.read_to_end(&mut Vec::new()).await.unwrap(); assert_eq!(resolve_etag_generic(&mut etag_reader), Some(etag_hex.clone())); @@ -178,9 +170,9 @@ mod tests { let etag = hasher.finalize(); let etag_hex = hex_simd::encode_to_string(etag, hex_simd::AsciiCase::Lower); let reader2 = BufReader::new(Cursor::new(&data2[..])); - let reader2 = Box::new(WarpReader::new(reader2)); let mut hash_reader = - HashReader::new(reader2, data2.len() as i64, data2.len() as i64, Some(etag_hex.clone()), None, false).unwrap(); + HashReader::from_stream(reader2, data2.len() as i64, data2.len() as i64, Some(etag_hex.clone()), None, false) + .unwrap(); hash_reader.read_to_end(&mut Vec::new()).await.unwrap(); assert_eq!(resolve_etag_generic(&mut hash_reader), Some(etag_hex.clone())); @@ -191,7 +183,6 @@ mod tests { let etag = hasher.finalize(); let etag_hex = hex_simd::encode_to_string(etag, hex_simd::AsciiCase::Lower); let reader3 = BufReader::new(Cursor::new(&data3[..])); - let reader3 = Box::new(WarpReader::new(reader3)); let etag_reader3 = EtagReader::new(reader3, Some(etag_hex.clone())); let mut compress_reader = CompressReader::new(etag_reader3, CompressionAlgorithm::Zstd); compress_reader.read_to_end(&mut Vec::new()).await.unwrap(); @@ -204,7 +195,6 @@ mod tests { let etag = hasher.finalize(); let etag_hex = hex_simd::encode_to_string(etag, hex_simd::AsciiCase::Lower); let reader4 = BufReader::new(Cursor::new(&data4[..])); - let reader4 = Box::new(WarpReader::new(reader4)); let etag_reader4 = EtagReader::new(reader4, Some(etag_hex.clone())); let key = [1u8; 32]; let nonce = [1u8; 12]; @@ -227,10 +217,9 @@ mod tests { let data = b"Real world test data that might be compressed and encrypted"; let base_reader = BufReader::new(Cursor::new(&data[..])); - let base_reader = Box::new(WarpReader::new(base_reader)); // Create a complex nested structure that might occur in practice: // CompressReader>>> - let hash_reader = HashReader::new( + let hash_reader = HashReader::from_stream( base_reader, data.len() as i64, data.len() as i64, @@ -253,7 +242,6 @@ mod tests { // Test another complex nesting with EtagReader at the core let data2 = b"Another real world scenario"; let base_reader2 = BufReader::new(Cursor::new(&data2[..])); - let base_reader2 = Box::new(WarpReader::new(base_reader2)); let etag_reader = EtagReader::new(base_reader2, Some("core_etag".to_string())); let key2 = [99u8; 32]; let nonce2 = [88u8; 12]; @@ -279,21 +267,19 @@ mod tests { // Test with HashReader that has no etag let data = b"no etag test"; let reader = BufReader::new(Cursor::new(&data[..])); - let reader = Box::new(WarpReader::new(reader)); - let mut hash_reader_no_etag = HashReader::new(reader, data.len() as i64, data.len() as i64, None, None, false).unwrap(); + let mut hash_reader_no_etag = + HashReader::from_stream(reader, data.len() as i64, data.len() as i64, None, None, false).unwrap(); assert_eq!(resolve_etag_generic(&mut hash_reader_no_etag), None); // Test with EtagReader that has None etag let data2 = b"no etag test 2"; let reader2 = BufReader::new(Cursor::new(&data2[..])); - let reader2 = Box::new(WarpReader::new(reader2)); let mut etag_reader_none = EtagReader::new(reader2, None); assert_eq!(resolve_etag_generic(&mut etag_reader_none), None); // Test nested structure with no ETag at the core let data3 = b"nested no etag test"; let reader3 = BufReader::new(Cursor::new(&data3[..])); - let reader3 = Box::new(WarpReader::new(reader3)); let etag_reader3 = EtagReader::new(reader3, None); let mut compress_reader3 = CompressReader::new(etag_reader3, CompressionAlgorithm::Gzip); assert_eq!(resolve_etag_generic(&mut compress_reader3), None); diff --git a/crates/rio/src/etag_reader.rs b/crates/rio/src/etag_reader.rs index 0748e013a..ba1638069 100644 --- a/crates/rio/src/etag_reader.rs +++ b/crates/rio/src/etag_reader.rs @@ -13,7 +13,7 @@ // limitations under the License. use crate::compress_index::{Index, TryGetIndex}; -use crate::{EtagResolvable, HashReaderDetector, HashReaderMut, Reader}; +use crate::{EtagResolvable, HashReaderDetector, HashReaderMut}; use md5::{Digest, Md5}; use pin_project_lite::pin_project; use std::pin::Pin; @@ -22,36 +22,51 @@ use tokio::io::{AsyncRead, ReadBuf}; use tracing::error; pin_project! { - pub struct EtagReader { + pub struct EtagReader { #[pin] - pub inner: Box, + pub inner: R, pub md5: Md5, pub finished: bool, pub checksum: Option, + resolved_etag: Option, } } -impl EtagReader { - pub fn new(inner: Box, checksum: Option) -> Self { +impl EtagReader { + pub fn new(inner: R, checksum: Option) -> Self { Self { inner, md5: Md5::new(), finished: false, checksum, + resolved_etag: None, } } /// Get the final md5 value (etag) as a hex string, only compute once. /// Can be called multiple times, always returns the same result after finished. pub fn get_etag(&mut self) -> String { + if let Some(etag) = &self.resolved_etag { + return etag.clone(); + } + let etag = self.md5.clone().finalize().to_vec(); - hex_simd::encode_to_string(etag, hex_simd::AsciiCase::Lower) + let etag = hex_simd::encode_to_string(etag, hex_simd::AsciiCase::Lower); + self.resolved_etag = Some(etag.clone()); + etag } } -impl AsyncRead for EtagReader { +impl AsyncRead for EtagReader +where + R: AsyncRead, +{ fn poll_read(self: Pin<&mut Self>, cx: &mut Context<'_>, buf: &mut ReadBuf<'_>) -> Poll> { let mut this = self.project(); + if *this.finished { + return Poll::Ready(Ok(())); + } + let orig_filled = buf.filled().len(); let poll = this.inner.as_mut().poll_read(cx, buf); if let Poll::Ready(Ok(())) = &poll { @@ -61,13 +76,20 @@ impl AsyncRead for EtagReader { } else { // EOF *this.finished = true; - if let Some(checksum) = this.checksum { + let etag = if let Some(etag) = this.resolved_etag.as_ref() { + etag.clone() + } else { let etag = this.md5.clone().finalize().to_vec(); - let etag_hex = hex_simd::encode_to_string(etag, hex_simd::AsciiCase::Lower); - if *checksum != etag_hex { - error!("Checksum mismatch, expected={:?}, actual={:?}", checksum, etag_hex); - return Poll::Ready(Err(std::io::Error::new(std::io::ErrorKind::InvalidData, "Checksum mismatch"))); - } + let etag = hex_simd::encode_to_string(etag, hex_simd::AsciiCase::Lower); + *this.resolved_etag = Some(etag.clone()); + etag + }; + + if let Some(checksum) = this.checksum + && *checksum != etag + { + error!("Checksum mismatch, expected={:?}, actual={:?}", checksum, etag); + return Poll::Ready(Err(std::io::Error::new(std::io::ErrorKind::InvalidData, "Checksum mismatch"))); } } } @@ -75,7 +97,7 @@ impl AsyncRead for EtagReader { } } -impl EtagResolvable for EtagReader { +impl EtagResolvable for EtagReader { fn is_etag_reader(&self) -> bool { true } @@ -91,7 +113,10 @@ impl EtagResolvable for EtagReader { } } -impl HashReaderDetector for EtagReader { +impl HashReaderDetector for EtagReader +where + R: HashReaderDetector, +{ fn is_hash_reader(&self) -> bool { self.inner.is_hash_reader() } @@ -101,7 +126,10 @@ impl HashReaderDetector for EtagReader { } } -impl TryGetIndex for EtagReader { +impl TryGetIndex for EtagReader +where + R: TryGetIndex, +{ fn try_get_index(&self) -> Option<&Index> { self.inner.try_get_index() } @@ -109,8 +137,6 @@ impl TryGetIndex for EtagReader { #[cfg(test)] mod tests { - use crate::WarpReader; - use super::*; use rand::RngExt; use std::io::Cursor; @@ -124,7 +150,6 @@ mod tests { let hex = faster_hex::hex_string(hasher.finalize().as_slice()); let expected = hex.to_string(); let reader = BufReader::new(&data[..]); - let reader = Box::new(WarpReader::new(reader)); let mut etag_reader = EtagReader::new(reader, None); let mut buf = Vec::new(); @@ -144,7 +169,6 @@ mod tests { let hex = faster_hex::hex_string(hasher.finalize().as_slice()); let expected = hex.to_string(); let reader = BufReader::new(&data[..]); - let reader = Box::new(WarpReader::new(reader)); let mut etag_reader = EtagReader::new(reader, None); let mut buf = Vec::new(); @@ -164,7 +188,6 @@ mod tests { let hex = faster_hex::hex_string(hasher.finalize().as_slice()); let expected = hex.to_string(); let reader = BufReader::new(&data[..]); - let reader = Box::new(WarpReader::new(reader)); let mut etag_reader = EtagReader::new(reader, None); let mut buf = Vec::new(); @@ -181,7 +204,6 @@ mod tests { async fn test_etag_reader_not_finished() { let data = b"abc123"; let reader = BufReader::new(&data[..]); - let reader = Box::new(WarpReader::new(reader)); let mut etag_reader = EtagReader::new(reader, None); // Do not read to end, etag should be None @@ -202,7 +224,6 @@ mod tests { let hex = faster_hex::hex_string(hasher.finalize().as_slice()); let expected = hex.to_string(); let reader = Cursor::new(data.clone()); - let reader = Box::new(WarpReader::new(reader)); let mut etag_reader = EtagReader::new(reader, None); let mut buf = Vec::new(); let n = etag_reader.read_to_end(&mut buf).await.unwrap(); @@ -220,7 +241,6 @@ mod tests { hasher.update(data); let expected = hex_simd::encode_to_string(hasher.finalize(), hex_simd::AsciiCase::Lower); let reader = BufReader::new(&data[..]); - let reader = Box::new(WarpReader::new(reader)); let mut etag_reader = EtagReader::new(reader, Some(expected.clone())); let mut buf = Vec::new(); @@ -236,7 +256,6 @@ mod tests { let data = b"checksum test data"; let wrong_checksum = "deadbeefdeadbeefdeadbeefdeadbeef".to_string(); let reader = BufReader::new(&data[..]); - let reader = Box::new(WarpReader::new(reader)); let mut etag_reader = EtagReader::new(reader, Some(wrong_checksum.clone())); let mut buf = Vec::new(); diff --git a/crates/rio/src/hardlimit_reader.rs b/crates/rio/src/hardlimit_reader.rs index 11c130639..e50b052f5 100644 --- a/crates/rio/src/hardlimit_reader.rs +++ b/crates/rio/src/hardlimit_reader.rs @@ -12,8 +12,6 @@ // See the License for the specific language governing permissions and // limitations under the License. -use crate::compress_index::{Index, TryGetIndex}; -use crate::{EtagResolvable, HashReaderDetector, HashReaderMut, Reader}; use pin_project_lite::pin_project; use std::io::{Error, Result}; use std::pin::Pin; @@ -21,20 +19,23 @@ use std::task::{Context, Poll}; use tokio::io::{AsyncRead, ReadBuf}; pin_project! { - pub struct HardLimitReader { + pub struct HardLimitReader { #[pin] - pub inner: Box, + pub inner: R, remaining: i64, } } -impl HardLimitReader { - pub fn new(inner: Box, limit: i64) -> Self { +impl HardLimitReader { + pub fn new(inner: R, limit: i64) -> Self { HardLimitReader { inner, remaining: limit } } } -impl AsyncRead for HardLimitReader { +impl AsyncRead for HardLimitReader +where + R: AsyncRead, +{ fn poll_read(mut self: Pin<&mut Self>, cx: &mut Context<'_>, buf: &mut ReadBuf<'_>) -> Poll> { if self.remaining < 0 { return Poll::Ready(Err(Error::other("input provided more bytes than specified"))); @@ -49,8 +50,8 @@ impl AsyncRead for HardLimitReader { if let Poll::Ready(Ok(())) = &poll { let after = buf.filled().len(); let read = (after - before) as i64; - self.remaining -= read; - if self.remaining < 0 { + *this.remaining -= read; + if *this.remaining < 0 { return Poll::Ready(Err(Error::other("input provided more bytes than specified"))); } } @@ -58,33 +59,12 @@ impl AsyncRead for HardLimitReader { } } -impl EtagResolvable for HardLimitReader { - fn try_resolve_etag(&mut self) -> Option { - self.inner.try_resolve_etag() - } -} - -impl HashReaderDetector for HardLimitReader { - fn is_hash_reader(&self) -> bool { - self.inner.is_hash_reader() - } - fn as_hash_reader_mut(&mut self) -> Option<&mut dyn HashReaderMut> { - self.inner.as_hash_reader_mut() - } -} - -impl TryGetIndex for HardLimitReader { - fn try_get_index(&self) -> Option<&Index> { - self.inner.try_get_index() - } -} +delegate_reader_capabilities_generic!(HardLimitReader, inner); #[cfg(test)] mod tests { use std::vec; - use crate::WarpReader; - use super::*; use rustfs_utils::read_full; use tokio::io::{AsyncReadExt, BufReader}; @@ -93,7 +73,6 @@ mod tests { async fn test_hardlimit_reader_normal() { let data = b"hello world"; let reader = BufReader::new(&data[..]); - let reader = Box::new(WarpReader::new(reader)); let hardlimit = HardLimitReader::new(reader, 20); let mut r = hardlimit; let mut buf = Vec::new(); @@ -106,7 +85,6 @@ mod tests { async fn test_hardlimit_reader_exact_limit() { let data = b"1234567890"; let reader = BufReader::new(&data[..]); - let reader = Box::new(WarpReader::new(reader)); let hardlimit = HardLimitReader::new(reader, 10); let mut r = hardlimit; let mut buf = Vec::new(); @@ -119,7 +97,6 @@ mod tests { async fn test_hardlimit_reader_exceed_limit() { let data = b"abcdef"; let reader = BufReader::new(&data[..]); - let reader = Box::new(WarpReader::new(reader)); let hardlimit = HardLimitReader::new(reader, 3); let mut r = hardlimit; let mut buf = vec![0u8; 10]; @@ -144,7 +121,6 @@ mod tests { async fn test_hardlimit_reader_empty() { let data = b""; let reader = BufReader::new(&data[..]); - let reader = Box::new(WarpReader::new(reader)); let hardlimit = HardLimitReader::new(reader, 5); let mut r = hardlimit; let mut buf = Vec::new(); diff --git a/crates/rio/src/hash_reader.rs b/crates/rio/src/hash_reader.rs index 0c6949a8d..aee0a50d6 100644 --- a/crates/rio/src/hash_reader.rs +++ b/crates/rio/src/hash_reader.rs @@ -12,15 +12,17 @@ // See the License for the specific language governing permissions and // limitations under the License. -//! HashReader implementation with generic support +//! HashReader implementation with stream-first construction helpers. //! -//! This module provides a generic `HashReader` that can wrap any type implementing -//! `AsyncRead + Unpin + Send + Sync + 'static + EtagResolvable`. +//! `HashReader` still stores a dynamic reader internally so it can preserve +//! capability-aware wrapping behavior. For plain async readers, prefer +//! `HashReader::from_stream(...)`. Use `HashReader::new(...)` when the input is +//! already a `DynReader` or when compatibility with existing boxed wrapping +//! logic matters. //! -//! ## Migration from the original Reader enum +//! ## Construction patterns //! -//! The original `HashReader::new` method that worked with the `Reader` enum -//! has been replaced with a generic approach. To preserve the original logic: +//! `HashReader::new(...)` keeps the original dyn-reader behavior: //! //! ### Original logic (before generics): //! ```ignore @@ -38,40 +40,23 @@ //! use rustfs_rio::{HashReader, HardLimitReader, EtagReader}; //! use tokio::io::BufReader; //! use std::io::Cursor; -//! use rustfs_rio::WarpReader; //! //! # tokio_test::block_on(async { //! let data = b"hello world"; //! let reader = BufReader::new(Cursor::new(&data[..])); -//! let reader = Box::new(WarpReader::new(reader)); //! let size = data.len() as i64; //! let actual_size = size; //! let etag = None; //! let diskable_md5 = false; //! //! // Method 1: Simple creation (recommended for most cases) -//! let hash_reader = HashReader::new(reader, size, actual_size, etag.clone(), None, diskable_md5).unwrap(); +//! let hash_reader = HashReader::from_stream(reader, size, actual_size, etag.clone(), None, diskable_md5).unwrap(); //! -//! // Method 2: With manual wrapping to recreate original logic +//! // Method 2: With a capability-aware typed wrapper //! let reader2 = BufReader::new(Cursor::new(&data[..])); -//! let reader2 = Box::new(WarpReader::new(reader2)); -//! let wrapped_reader: Box = if size > 0 { -//! if !diskable_md5 { -//! // Wrap with both HardLimitReader and EtagReader -//! let hard_limit = HardLimitReader::new(reader2, size); -//! Box::new(EtagReader::new(Box::new(hard_limit), etag.clone())) -//! } else { -//! // Only wrap with HardLimitReader -//! Box::new(HardLimitReader::new(reader2, size)) -//! } -//! } else if !diskable_md5 { -//! // Only wrap with EtagReader -//! Box::new(EtagReader::new(reader2, etag.clone())) -//! } else { -//! // No wrapping needed -//! reader2 -//! }; -//! let hash_reader2 = HashReader::new(wrapped_reader, size, actual_size, etag.clone(), None, diskable_md5).unwrap(); +//! let reader2 = HashReader::from_stream(reader2, size, actual_size, etag.clone(), None, diskable_md5).unwrap(); +//! let wrapped_reader = EtagReader::new(HardLimitReader::new(reader2, size), etag.clone()); +//! let hash_reader2 = HashReader::from_reader(wrapped_reader, size, actual_size, etag.clone(), None, diskable_md5).unwrap(); //! # }); //! ``` //! @@ -83,19 +68,18 @@ //! use rustfs_rio::{HashReader, HashReaderDetector}; //! use tokio::io::BufReader; //! use std::io::Cursor; -//! use rustfs_rio::WarpReader; //! //! # tokio_test::block_on(async { //! let data = b"test"; //! let reader = BufReader::new(Cursor::new(&data[..])); -//! let hash_reader = HashReader::new(Box::new(WarpReader::new(reader)), 4, 4, None, None,false).unwrap(); +//! let hash_reader = HashReader::from_stream(reader, 4, 4, None, None,false).unwrap(); //! //! // Check if a type is a HashReader //! assert!(hash_reader.is_hash_reader()); //! -//! // Use new for compatibility (though it's simpler to use new() directly) +//! // `from_stream` is the recommended entry point for plain readers //! let reader2 = BufReader::new(Cursor::new(&data[..])); -//! let result = HashReader::new(Box::new(WarpReader::new(reader2)), 4, 4, None, None, false); +//! let result = HashReader::from_stream(reader2, 4, 4, None, None, false); //! assert!(result.is_ok()); //! # }); //! ``` @@ -106,7 +90,7 @@ use crate::ChecksumType; use crate::Sha256Hasher; use crate::compress_index::{Index, TryGetIndex}; use crate::get_content_checksum; -use crate::{EtagReader, EtagResolvable, HardLimitReader, HashReaderDetector, Reader, WarpReader}; +use crate::{DynReader, EtagReader, EtagResolvable, HardLimitReader, HashReaderDetector, WarpReader, boxed_reader, wrap_reader}; use base64::Engine; use base64::engine::general_purpose; use http::HeaderMap; @@ -123,8 +107,8 @@ use tracing::error; /// Trait for mutable operations on HashReader pub trait HashReaderMut { - fn into_inner(self) -> Box; - fn take_inner(&mut self) -> Box; + fn into_inner(self) -> DynReader; + fn take_inner(&mut self) -> DynReader; fn bytes_read(&self) -> u64; fn checksum(&self) -> &Option; fn set_checksum(&mut self, checksum: Option); @@ -142,7 +126,7 @@ pin_project! { pub struct HashReader { #[pin] - pub inner: Box, + pub inner: DynReader, pub size: i64, checksum: Option, pub actual_size: i64, @@ -163,8 +147,89 @@ pin_project! { impl HashReader { /// Used for transformation layers (compression/encryption) pub const SIZE_PRESERVE_LAYER: i64 = -1; + + pub fn from_reader( + inner: R, + size: i64, + actual_size: i64, + md5hex: Option, + sha256hex: Option, + diskable_md5: bool, + ) -> std::io::Result + where + R: crate::Reader + 'static, + { + let inner = if size > 0 { + let hard_limit_reader = HardLimitReader::new(inner, size); + if !diskable_md5 { + boxed_reader(EtagReader::new(hard_limit_reader, md5hex.clone())) + } else { + boxed_reader(hard_limit_reader) + } + } else if size != Self::SIZE_PRESERVE_LAYER && !diskable_md5 { + boxed_reader(EtagReader::new(inner, md5hex.clone())) + } else { + boxed_reader(inner) + }; + + Ok(Self { + inner, + size, + checksum: md5hex, + actual_size, + diskable_md5, + bytes_read: 0, + content_hash: None, + content_hasher: None, + content_sha256: sha256hex.clone(), + content_sha256_hasher: sha256hex.map(|_| Sha256Hasher::new()), + checksum_on_finish: false, + trailer_s3s: None, + }) + } + + pub fn from_stream( + inner: R, + size: i64, + actual_size: i64, + md5hex: Option, + sha256hex: Option, + diskable_md5: bool, + ) -> std::io::Result + where + R: crate::ReadStream + 'static, + { + let inner = WarpReader::new(inner); + let inner = if size > 0 { + if !diskable_md5 { + boxed_reader(EtagReader::new(HardLimitReader::new(inner, size), md5hex.clone())) + } else { + boxed_reader(HardLimitReader::new(inner, size)) + } + } else if size != Self::SIZE_PRESERVE_LAYER && !diskable_md5 { + boxed_reader(EtagReader::new(inner, md5hex.clone())) + } else { + boxed_reader(inner) + }; + + Ok(Self { + inner, + size, + checksum: md5hex, + actual_size, + diskable_md5, + bytes_read: 0, + content_hash: None, + content_hasher: None, + content_sha256: sha256hex.clone(), + content_sha256_hasher: sha256hex.map(|_| Sha256Hasher::new()), + checksum_on_finish: false, + trailer_s3s: None, + }) + } + pub fn new( - mut inner: Box, + mut inner: DynReader, size: i64, actual_size: i64, md5hex: Option, @@ -262,7 +327,7 @@ impl HashReader { } } - pub fn into_inner(self) -> Box { + pub fn into_inner(self) -> DynReader { self.inner } @@ -387,13 +452,13 @@ impl HashReader { } impl HashReaderMut for HashReader { - fn into_inner(self) -> Box { + fn into_inner(self) -> DynReader { self.inner } - fn take_inner(&mut self) -> Box { + fn take_inner(&mut self) -> DynReader { // Replace inner with an empty reader to move it out safely while keeping self valid - mem::replace(&mut self.inner, Box::new(WarpReader::new(Cursor::new(Vec::new())))) + mem::replace(&mut self.inner, wrap_reader(Cursor::new(Vec::new()))) } fn bytes_read(&self) -> u64 { @@ -561,7 +626,7 @@ impl TryGetIndex for HashReader { #[cfg(test)] mod tests { use super::*; - use crate::{DecryptReader, WarpReader, encrypt_reader}; + use crate::{DecryptReader, EncryptReader, encrypt_reader, wrap_reader}; use rand::RngExt; use std::io::Cursor; use tokio::io::{AsyncReadExt, BufReader}; @@ -575,41 +640,92 @@ mod tests { // Test 1: Simple creation let reader1 = BufReader::new(Cursor::new(&data[..])); - let reader1 = Box::new(WarpReader::new(reader1)); - let hash_reader1 = HashReader::new(reader1, size, actual_size, etag.clone(), None, false).unwrap(); + let hash_reader1 = HashReader::from_stream(reader1, size, actual_size, etag.clone(), None, false).unwrap(); assert_eq!(hash_reader1.size(), size); assert_eq!(hash_reader1.actual_size(), actual_size); // Test 2: With HardLimitReader wrapping - let reader2 = BufReader::new(Cursor::new(&data[..])); - let reader2 = Box::new(WarpReader::new(reader2)); + let reader2 = + HashReader::from_stream(BufReader::new(Cursor::new(&data[..])), size, actual_size, etag.clone(), None, false) + .unwrap(); let hard_limit = HardLimitReader::new(reader2, size); - let hard_limit = Box::new(hard_limit); - let hash_reader2 = HashReader::new(hard_limit, size, actual_size, etag.clone(), None, false).unwrap(); + let hash_reader2 = HashReader::from_reader(hard_limit, size, actual_size, etag.clone(), None, false).unwrap(); assert_eq!(hash_reader2.size(), size); assert_eq!(hash_reader2.actual_size(), actual_size); // Test 3: With EtagReader wrapping - let reader3 = BufReader::new(Cursor::new(&data[..])); - let reader3 = Box::new(WarpReader::new(reader3)); + let reader3 = + HashReader::from_stream(BufReader::new(Cursor::new(&data[..])), size, actual_size, etag.clone(), None, false) + .unwrap(); let etag_reader = EtagReader::new(reader3, etag.clone()); - let etag_reader = Box::new(etag_reader); - let hash_reader3 = HashReader::new(etag_reader, size, actual_size, etag.clone(), None, false).unwrap(); + let hash_reader3 = HashReader::from_reader(etag_reader, size, actual_size, etag.clone(), None, false).unwrap(); assert_eq!(hash_reader3.size(), size); assert_eq!(hash_reader3.actual_size(), actual_size); } + #[test] + fn test_boxed_reader_capabilities_delegate() { + let data = b"boxed capabilities"; + let mut boxed_etag_reader = + Box::new(EtagReader::new(BufReader::new(Cursor::new(&data[..])), Some("boxed_etag".to_string()))); + assert_eq!(boxed_etag_reader.try_resolve_etag(), Some("boxed_etag".to_string())); + + let boxed_hash_reader = Box::new( + HashReader::from_stream( + BufReader::new(Cursor::new(&data[..])), + data.len() as i64, + data.len() as i64, + None, + None, + false, + ) + .unwrap(), + ); + assert!(boxed_hash_reader.is_hash_reader()); + } + + #[tokio::test] + async fn test_from_reader_accepts_boxed_encrypt_reader() { + let data = b"boxed encrypt reader"; + let inner = HashReader::from_stream( + BufReader::new(Cursor::new(&data[..])), + data.len() as i64, + data.len() as i64, + None, + None, + false, + ) + .unwrap(); + let boxed_encrypt_reader = Box::new(EncryptReader::new(inner, [7u8; 32], [3u8; 12])); + + assert!(boxed_encrypt_reader.is_hash_reader()); + + let mut hash_reader = HashReader::from_reader( + boxed_encrypt_reader, + HashReader::SIZE_PRESERVE_LAYER, + data.len() as i64, + None, + None, + false, + ) + .unwrap(); + let mut encrypted = Vec::new(); + hash_reader.read_to_end(&mut encrypted).await.unwrap(); + + assert!(!encrypted.is_empty()); + assert_ne!(encrypted, data); + assert_eq!(hash_reader.actual_size(), data.len() as i64); + } + #[tokio::test] async fn test_hashreader_etag_basic() { let data = b"hello hashreader"; let reader = BufReader::new(Cursor::new(&data[..])); - let reader = Box::new(WarpReader::new(reader)); - let mut hash_reader = HashReader::new(reader, data.len() as i64, data.len() as i64, None, None, false).unwrap(); + let mut hash_reader = HashReader::from_stream(reader, data.len() as i64, data.len() as i64, None, None, false).unwrap(); let mut buf = Vec::new(); let _ = hash_reader.read_to_end(&mut buf).await.unwrap(); - // Since we removed EtagReader integration, etag might be None - let _etag = hash_reader.try_resolve_etag(); - // Just check that we can call etag() without error + let etag = hash_reader.try_resolve_etag(); + assert!(etag.is_some()); assert_eq!(buf, data); } @@ -617,8 +733,7 @@ mod tests { async fn test_hashreader_diskable_md5() { let data = b"no etag"; let reader = BufReader::new(Cursor::new(&data[..])); - let reader = Box::new(WarpReader::new(reader)); - let mut hash_reader = HashReader::new(reader, data.len() as i64, data.len() as i64, None, None, true).unwrap(); + let mut hash_reader = HashReader::from_stream(reader, data.len() as i64, data.len() as i64, None, None, true).unwrap(); let mut buf = Vec::new(); let _ = hash_reader.read_to_end(&mut buf).await.unwrap(); // Etag should be None when diskable_md5 is true @@ -631,11 +746,11 @@ mod tests { async fn test_hashreader_new_logic() { let data = b"test data"; let reader = BufReader::new(Cursor::new(&data[..])); - let reader = Box::new(WarpReader::new(reader)); // Create a HashReader first let hash_reader = - HashReader::new(reader, data.len() as i64, data.len() as i64, Some("test_etag".to_string()), None, false).unwrap(); - let hash_reader = Box::new(WarpReader::new(hash_reader)); + HashReader::from_stream(reader, data.len() as i64, data.len() as i64, Some("test_etag".to_string()), None, false) + .unwrap(); + let hash_reader = wrap_reader(hash_reader); // Now try to create another HashReader from the existing one using new let result = HashReader::new( hash_reader, @@ -680,9 +795,7 @@ mod tests { let size = data.len() as i64; let actual_size = data.len() as i64; - let reader = Box::new(WarpReader::new(reader)); - // Create HashReader - let mut hr = HashReader::new(reader, size, actual_size, Some(expected.clone()), None, false).unwrap(); + let mut hr = HashReader::from_stream(reader, size, actual_size, Some(expected.clone()), None, false).unwrap(); // If compression is enabled, compress data first let compressed_data = if is_compress { @@ -710,7 +823,7 @@ mod tests { if is_encrypt { // Encrypt compressed data - let encrypt_reader = encrypt_reader::EncryptReader::new(WarpReader::new(Cursor::new(compressed_data)), key, nonce); + let encrypt_reader = encrypt_reader::EncryptReader::new(Cursor::new(compressed_data), key, nonce); let mut encrypted_data = Vec::new(); let mut encrypt_reader = encrypt_reader; encrypt_reader.read_to_end(&mut encrypted_data).await.unwrap(); @@ -718,15 +831,14 @@ mod tests { println!("Encrypted size: {}", encrypted_data.len()); // Decrypt data - let decrypt_reader = DecryptReader::new(WarpReader::new(Cursor::new(encrypted_data)), key, nonce); + let decrypt_reader = DecryptReader::new(Cursor::new(encrypted_data), key, nonce); let mut decrypt_reader = decrypt_reader; let mut decrypted_data = Vec::new(); decrypt_reader.read_to_end(&mut decrypted_data).await.unwrap(); if is_compress { // If compression was used, decompress is needed - let decompress_reader = - DecompressReader::new(WarpReader::new(Cursor::new(decrypted_data)), CompressionAlgorithm::Gzip); + let decompress_reader = DecompressReader::new(Cursor::new(decrypted_data), CompressionAlgorithm::Gzip); let mut decompress_reader = decompress_reader; let mut final_data = Vec::new(); decompress_reader.read_to_end(&mut final_data).await.unwrap(); @@ -744,8 +856,7 @@ mod tests { // When encryption is disabled, only handle compression/decompression if is_compress { - let decompress_reader = - DecompressReader::new(WarpReader::new(Cursor::new(compressed_data)), CompressionAlgorithm::Gzip); + let decompress_reader = DecompressReader::new(Cursor::new(compressed_data), CompressionAlgorithm::Gzip); let mut decompress_reader = decompress_reader; let mut decompressed = Vec::new(); decompress_reader.read_to_end(&mut decompressed).await.unwrap(); @@ -777,8 +888,7 @@ mod tests { println!("Original data size: {} bytes", data.len()); let reader = BufReader::new(Cursor::new(data.clone())); - let reader = Box::new(WarpReader::new(reader)); - let hash_reader = HashReader::new(reader, data.len() as i64, data.len() as i64, None, None, false).unwrap(); + let hash_reader = HashReader::from_stream(reader, data.len() as i64, data.len() as i64, None, None, false).unwrap(); // Test compression let compress_reader = CompressReader::new(hash_reader, CompressionAlgorithm::Gzip); @@ -823,8 +933,7 @@ mod tests { println!("\nTesting algorithm: {algorithm:?}"); let reader = BufReader::new(Cursor::new(data.clone())); - let reader = Box::new(WarpReader::new(reader)); - let hash_reader = HashReader::new(reader, data.len() as i64, data.len() as i64, None, None, false).unwrap(); + let hash_reader = HashReader::from_stream(reader, data.len() as i64, data.len() as i64, None, None, false).unwrap(); // Compress let compress_reader = CompressReader::new(hash_reader, algorithm); diff --git a/crates/rio/src/lib.rs b/crates/rio/src/lib.rs index fcfc0b3df..9663f133d 100644 --- a/crates/rio/src/lib.rs +++ b/crates/rio/src/lib.rs @@ -15,6 +15,67 @@ // Default encryption block size - aligned with system default read buffer size (1MB) pub const DEFAULT_ENCRYPTION_BLOCK_SIZE: usize = 1024 * 1024; +macro_rules! delegate_reader_capabilities_generic { + ($name:ident<$inner_ty:ident>, $inner:ident) => { + impl<$inner_ty> crate::EtagResolvable for $name<$inner_ty> + where + $inner_ty: crate::EtagResolvable, + { + fn try_resolve_etag(&mut self) -> Option { + self.$inner.try_resolve_etag() + } + } + + impl<$inner_ty> crate::HashReaderDetector for $name<$inner_ty> + where + $inner_ty: crate::HashReaderDetector, + { + fn is_hash_reader(&self) -> bool { + self.$inner.is_hash_reader() + } + + fn as_hash_reader_mut(&mut self) -> Option<&mut dyn crate::HashReaderMut> { + self.$inner.as_hash_reader_mut() + } + } + + impl<$inner_ty> crate::TryGetIndex for $name<$inner_ty> + where + $inner_ty: crate::TryGetIndex, + { + fn try_get_index(&self) -> Option<&crate::compress_index::Index> { + self.$inner.try_get_index() + } + } + }; +} + +macro_rules! delegate_reader_capabilities_generic_no_index { + ($name:ident<$inner_ty:ident>, $inner:ident) => { + impl<$inner_ty> crate::EtagResolvable for $name<$inner_ty> + where + $inner_ty: crate::EtagResolvable, + { + fn try_resolve_etag(&mut self) -> Option { + self.$inner.try_resolve_etag() + } + } + + impl<$inner_ty> crate::HashReaderDetector for $name<$inner_ty> + where + $inner_ty: crate::HashReaderDetector, + { + fn is_hash_reader(&self) -> bool { + self.$inner.is_hash_reader() + } + + fn as_hash_reader_mut(&mut self) -> Option<&mut dyn crate::HashReaderMut> { + self.$inner.as_hash_reader_mut() + } + } + }; +} + mod limit_reader; pub use limit_reader::LimitReader; @@ -53,7 +114,16 @@ pub use compress_index::{Index, TryGetIndex}; mod etag; -pub trait Reader: tokio::io::AsyncRead + Unpin + Send + Sync + EtagResolvable + HashReaderDetector + TryGetIndex {} +pub trait ReadStream: tokio::io::AsyncRead + Unpin + Send + Sync {} +impl ReadStream for T where T: tokio::io::AsyncRead + Unpin + Send + Sync {} + +pub trait ReaderCapabilities: EtagResolvable + HashReaderDetector + TryGetIndex {} +impl ReaderCapabilities for T where T: EtagResolvable + HashReaderDetector + TryGetIndex {} + +pub trait Reader: ReadStream + ReaderCapabilities {} +impl Reader for T where T: ReadStream + ReaderCapabilities {} + +pub type DynReader = Box; // Trait for types that can be recursively searched for etag capability pub trait EtagResolvable { @@ -84,20 +154,33 @@ pub trait HashReaderDetector { } } -impl Reader for crate::HashReader {} -impl Reader for crate::HardLimitReader {} -impl Reader for crate::EtagReader {} -impl Reader for crate::LimitReader where R: Reader {} -impl Reader for crate::CompressReader where R: Reader {} -impl Reader for crate::EncryptReader where R: Reader {} -impl Reader for crate::DecryptReader where R: Reader {} -impl EtagResolvable for Box { +pub fn boxed_reader(reader: R) -> DynReader +where + R: Reader + 'static, +{ + Box::new(reader) +} + +pub fn wrap_reader(reader: R) -> DynReader +where + R: ReadStream + 'static, +{ + boxed_reader(WarpReader::new(reader)) +} + +impl EtagResolvable for Box +where + T: EtagResolvable + ?Sized, +{ fn try_resolve_etag(&mut self) -> Option { self.as_mut().try_resolve_etag() } } -impl HashReaderDetector for Box { +impl HashReaderDetector for Box +where + T: HashReaderDetector + ?Sized, +{ fn is_hash_reader(&self) -> bool { self.as_ref().is_hash_reader() } @@ -107,10 +190,11 @@ impl HashReaderDetector for Box { } } -impl TryGetIndex for Box { +impl TryGetIndex for Box +where + T: TryGetIndex + ?Sized, +{ fn try_get_index(&self) -> Option<&compress_index::Index> { self.as_ref().try_get_index() } } - -impl Reader for Box {} diff --git a/crates/rio/src/limit_reader.rs b/crates/rio/src/limit_reader.rs index a4b6ebad3..7378674d6 100644 --- a/crates/rio/src/limit_reader.rs +++ b/crates/rio/src/limit_reader.rs @@ -37,8 +37,6 @@ use std::pin::Pin; use std::task::{Context, Poll}; use tokio::io::{AsyncRead, ReadBuf}; -use crate::{EtagResolvable, HashReaderDetector, HashReaderMut, TryGetIndex}; - pin_project! { #[derive(Debug)] pub struct LimitReader { @@ -46,6 +44,7 @@ pin_project! { pub inner: R, limit: usize, read: usize, + scratch: Vec, } } @@ -56,7 +55,12 @@ where { /// Create a new LimitReader wrapping `inner`, with a total read limit of `limit` bytes. pub fn new(inner: R, limit: usize) -> Self { - Self { inner, limit, read: 0 } + Self { + inner, + limit, + read: 0, + scratch: Vec::new(), + } } } @@ -84,8 +88,8 @@ where } poll } else { - let mut temp = vec![0u8; allowed]; - let mut temp_buf = ReadBuf::new(&mut temp); + this.scratch.resize(allowed, 0); + let mut temp_buf = ReadBuf::new(&mut this.scratch[..allowed]); let poll = this.inner.as_mut().poll_read(cx, &mut temp_buf); if let Poll::Ready(Ok(())) = &poll { let n = temp_buf.filled().len(); @@ -97,28 +101,7 @@ where } } -impl EtagResolvable for LimitReader -where - R: EtagResolvable, -{ - fn try_resolve_etag(&mut self) -> Option { - self.inner.try_resolve_etag() - } -} - -impl HashReaderDetector for LimitReader -where - R: HashReaderDetector, -{ - fn is_hash_reader(&self) -> bool { - self.inner.is_hash_reader() - } - fn as_hash_reader_mut(&mut self) -> Option<&mut dyn HashReaderMut> { - self.inner.as_hash_reader_mut() - } -} - -impl TryGetIndex for LimitReader where R: AsyncRead + Unpin + Send + Sync {} +delegate_reader_capabilities_generic!(LimitReader, inner); #[cfg(test)] mod tests { diff --git a/crates/rio/src/reader.rs b/crates/rio/src/reader.rs index e2a83e28e..d288abe25 100644 --- a/crates/rio/src/reader.rs +++ b/crates/rio/src/reader.rs @@ -17,7 +17,7 @@ use std::task::{Context, Poll}; use tokio::io::{AsyncRead, ReadBuf}; use crate::compress_index::TryGetIndex; -use crate::{EtagResolvable, HashReaderDetector, Reader}; +use crate::{EtagResolvable, HashReaderDetector}; pub struct WarpReader { inner: R, @@ -40,5 +40,3 @@ impl HashReaderDetector for WarpReader {} impl EtagResolvable for WarpReader {} impl TryGetIndex for WarpReader {} - -impl Reader for WarpReader {} diff --git a/rustfs/src/app/multipart_usecase.rs b/rustfs/src/app/multipart_usecase.rs index f9c9ba5a3..fc1e069d4 100644 --- a/rustfs/src/app/multipart_usecase.rs +++ b/rustfs/src/app/multipart_usecase.rs @@ -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 = 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 = 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); diff --git a/rustfs/src/app/object_usecase.rs b/rustfs/src/app/object_usecase.rs index b8c70de93..f3df0aa23 100644 --- a/rustfs/src/app/object_usecase.rs +++ b/rustfs/src/app/object_usecase.rs @@ -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, + final_stream: DynReader, rs: Option, content_type: Option, last_modified: Option, @@ -1319,14 +1319,7 @@ impl DefaultObjectUsecase { decrypted_stream, ) } - None => ( - None, - None, - None, - None, - false, - Box::new(WarpReader::new(encrypted_stream)) as Box, - ), + None => (None, None, None, None, false, wrap_reader(encrypted_stream)), }; Ok(GetObjectReadSetup { @@ -1824,8 +1817,6 @@ impl DefaultObjectUsecase { } } - let mut reader: Box = 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, + final_stream: DynReader, rs: Option, content_type: Option, last_modified: Option, @@ -3339,8 +3330,6 @@ impl DefaultObjectUsecase { src_info.metadata_only = true; } - let mut reader: Box = 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 = 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 = 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; diff --git a/rustfs/src/storage/mod.rs b/rustfs/src/storage/mod.rs index e5ba0b9b9..52de62dae 100644 --- a/rustfs/src/storage/mod.rs +++ b/rustfs/src/storage/mod.rs @@ -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; diff --git a/rustfs/src/storage/readers.rs b/rustfs/src/storage/readers.rs deleted file mode 100644 index 0d19e7609..000000000 --- a/rustfs/src/storage/readers.rs +++ /dev/null @@ -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>, -} - -impl InMemoryAsyncReader { - pub(crate) fn new(data: Vec) -> 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> { - 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::task::Poll::Ready(Ok(self.cursor.position())) - } -} diff --git a/rustfs/src/storage/sse.rs b/rustfs/src/storage/sse.rs index 2ca6e7da9..dbaacc723 100644 --- a/rustfs/src/storage/sse.rs +++ b/rustfs/src/storage/sse.rs @@ -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(&self, reader: R) -> Box> 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(&self, reader: R) -> Box> 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, - ) -> Result<(Box, i64), StorageError> { + pub async fn wrap_multipart_stream(&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, - actual_size: i64, - ) -> Result<(Box, i64), StorageError> { - let (mut final_stream, response_content_length): (Box, i64) = if self.is_multipart { + /// Accepts a readable stream (from object storage) and returns (decrypted_reader, plaintext_size) + pub async fn wrap_reader(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, 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, ApiError> { @@ -1642,49 +1642,43 @@ pub fn strip_managed_encryption_metadata(metadata: &mut HashMap) // 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, +#[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( + encrypted_stream: R, parts: &[ObjectPartInfo], key_bytes: [u8; 32], base_nonce: [u8; 12], -) -> Result<(Box, 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; + 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 = (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, diff --git a/rustfs/src/storage/sse_test.rs b/rustfs/src/storage/sse_test.rs index bcc059e5e..06b414dd4 100644 --- a/rustfs/src/storage/sse_test.rs +++ b/rustfs/src/storage/sse_test.rs @@ -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 = (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,