From 6696703343fd18df4884f5c6a61346673e613d20 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E5=AE=89=E6=AD=A3=E8=B6=85?= Date: Fri, 3 Apr 2026 11:30:56 +0800 Subject: [PATCH 01/22] test: cover delete-group percent decoding (#2373) --- rustfs/src/admin/handlers/group.rs | 109 +++++++++++++++++++++++------ 1 file changed, 88 insertions(+), 21 deletions(-) diff --git a/rustfs/src/admin/handlers/group.rs b/rustfs/src/admin/handlers/group.rs index 5db394645..b99977e6e 100644 --- a/rustfs/src/admin/handlers/group.rs +++ b/rustfs/src/admin/handlers/group.rs @@ -207,30 +207,11 @@ impl Operation for DeleteGroup { ) .await?; - let group_raw = params - .get("group") - .ok_or_else(|| s3_error!(InvalidArgument, "missing group name in request"))? - .trim(); - - // Path segments stay percent-encoded in `req.uri.path()` / matchit; IAM uses decoded names (same as GET query). - let group_decoded = percent_decode_str(group_raw) - .decode_utf8() - .map_err(|_| s3_error!(InvalidArgument, "invalid group name encoding"))?; - let group = group_decoded.trim(); - - // Validate the group name format - if group.is_empty() || group.len() > 256 { - return Err(s3_error!(InvalidArgument, "invalid group name")); - } - - // Sanity check the group name - if group.contains(['/', '\\', '\0']) { - return Err(s3_error!(InvalidArgument, "group name contains invalid characters")); - } + let group = decode_delete_group_name(¶ms)?; let Ok(iam_store) = rustfs_iam::get() else { return Err(s3_error!(InternalError, "iam not init")) }; - let updated_at = iam_store.remove_users_from_group(group, vec![]).await.map_err(|e| { + let updated_at = iam_store.remove_users_from_group(&group, vec![]).await.map_err(|e| { warn!("delete group failed, e: {:?}", e); match e { rustfs_iam::error::Error::GroupNotEmpty => { @@ -276,6 +257,33 @@ impl Operation for DeleteGroup { } } +fn decode_delete_group_name<'a>(params: &'a Params<'_, '_>) -> S3Result> { + let group_raw = params + .get("group") + .ok_or_else(|| s3_error!(InvalidArgument, "missing group name in request"))? + .trim(); + + // Path segments stay percent-encoded in `req.uri.path()` / matchit; IAM uses decoded names (same as GET query). + let decoded = percent_decode_str(group_raw) + .decode_utf8() + .map_err(|_| s3_error!(InvalidArgument, "invalid group name encoding"))?; + let group = decoded.trim(); + + if group.is_empty() || group.len() > 256 { + return Err(s3_error!(InvalidArgument, "invalid group name")); + } + + if group.contains(['/', '\\', '\0']) { + return Err(s3_error!(InvalidArgument, "group name contains invalid characters")); + } + + if group.len() == decoded.len() { + Ok(decoded) + } else { + Ok(std::borrow::Cow::Owned(group.to_string())) + } +} + pub struct SetGroupStatus {} #[async_trait::async_trait] impl Operation for SetGroupStatus { @@ -484,3 +492,62 @@ impl Operation for UpdateGroupMembers { Ok(S3Response::with_headers((StatusCode::OK, Body::empty()), header)) } } + +#[cfg(test)] +mod tests { + use super::*; + use matchit::Router; + + fn with_delete_group_params(path: &str, f: impl FnOnce(&Params<'_, '_>) -> T) -> T { + let mut router = Router::new(); + router + .insert("/rustfs/admin/v3/group/{group}", ()) + .expect("route should insert"); + + let matched = router.at(path).expect("route should match"); + f(&matched.params) + } + + #[test] + fn decode_delete_group_name_percent_decodes_path_segment() { + let group = with_delete_group_params("/rustfs/admin/v3/group/dev%2Bops%20team", |params| { + decode_delete_group_name(params).map(|group| group.into_owned()) + }) + .expect("encoded group name should decode"); + + assert_eq!(group, "dev+ops team"); + } + + #[test] + fn decode_delete_group_name_rejects_invalid_utf8() { + let err = with_delete_group_params("/rustfs/admin/v3/group/%FF", |params| { + decode_delete_group_name(params).map(|group| group.into_owned()) + }) + .expect_err("invalid utf-8 should fail"); + + assert_eq!(err.code(), &S3ErrorCode::InvalidArgument); + assert_eq!(err.message(), Some("invalid group name encoding")); + } + + #[test] + fn decode_delete_group_name_rejects_blank_name_after_decoding() { + let err = with_delete_group_params("/rustfs/admin/v3/group/%20", |params| { + decode_delete_group_name(params).map(|group| group.into_owned()) + }) + .expect_err("blank group should fail"); + + assert_eq!(err.code(), &S3ErrorCode::InvalidArgument); + assert_eq!(err.message(), Some("invalid group name")); + } + + #[test] + fn decode_delete_group_name_rejects_path_separator_after_decoding() { + let err = with_delete_group_params("/rustfs/admin/v3/group/team%2Fops", |params| { + decode_delete_group_name(params).map(|group| group.into_owned()) + }) + .expect_err("decoded slash should fail"); + + assert_eq!(err.code(), &S3ErrorCode::InvalidArgument); + assert_eq!(err.message(), Some("group name contains invalid characters")); + } +} From 6a114cd2e06faaba1eb8cf35227c3c1129e6863b Mon Sep 17 00:00:00 2001 From: weisd Date: Fri, 3 Apr 2026 13:27:56 +0800 Subject: [PATCH 02/22] fix: bump s3s for presigned checksum handling (#2379) --- Cargo.lock | 113 ++++++++++++++++++++++++++--------------------------- Cargo.toml | 2 +- 2 files changed, 57 insertions(+), 58 deletions(-) diff --git a/Cargo.lock b/Cargo.lock index 33842deb3..e51a91424 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -1222,7 +1222,7 @@ version = "0.11.0-rc.5" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "d52965399b470437fc7f4d4b51134668dbc96573fea6f1b83318a420e4605745" dependencies = [ - "digest 0.11.1", + "digest 0.11.2", ] [[package]] @@ -3108,9 +3108,9 @@ dependencies = [ [[package]] name = "digest" -version = "0.11.1" +version = "0.11.2" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "285743a676ccb6b3e116bc14cc69319b957867930ae9c4822f8e0f54509d7243" +checksum = "4850db49bf08e663084f7fb5c87d202ef91a3907271aff24a94eb97ff039153c" dependencies = [ "block-buffer 0.12.0", "const-oid 0.10.2", @@ -3204,7 +3204,7 @@ dependencies = [ "serde", "serde_json", "serial_test", - "sha2 0.11.0-rc.5", + "sha2 0.11.0", "suppaftp", "time", "tokio", @@ -3412,7 +3412,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "39cab71617ae0d63f51a36d69f866391735b51691dbda63cf6f96d042b63efeb" dependencies = [ "libc", - "windows-sys 0.52.0", + "windows-sys 0.59.0", ] [[package]] @@ -4205,11 +4205,11 @@ dependencies = [ [[package]] name = "hmac" -version = "0.13.0-rc.5" +version = "0.13.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "ef451d73f36d8a3f93ad32c332ea01146c9650e1ec821a9b0e46c01277d544f8" +checksum = "6303bc9732ae41b04cb554b844a762b4115a61bfaa81e3e83050991eeb56863f" dependencies = [ - "digest 0.11.1", + "digest 0.11.2", ] [[package]] @@ -4322,9 +4322,9 @@ dependencies = [ [[package]] name = "hyper" -version = "1.8.1" +version = "1.9.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "2ab2d4f250c3d7b1c9fcdff1cece94ea4e2dfbec68614f7b87cb205f24ca9d11" +checksum = "6299f016b246a94207e63da54dbe807655bf9e00044f73ded42c3ac5305fbcca" dependencies = [ "atomic-waker", "bytes", @@ -4337,7 +4337,6 @@ dependencies = [ "httpdate", "itoa", "pin-project-lite", - "pin-utils", "smallvec", "tokio", "want", @@ -4667,7 +4666,7 @@ checksum = "3640c1c38b8e4e43584d8df18be5fc6b0aa314ce6ebf51b53313d4306cca8e46" dependencies = [ "hermit-abi", "libc", - "windows-sys 0.52.0", + "windows-sys 0.59.0", ] [[package]] @@ -4744,7 +4743,7 @@ dependencies = [ "portable-atomic", "portable-atomic-util", "serde_core", - "windows-sys 0.52.0", + "windows-sys 0.59.0", ] [[package]] @@ -5233,12 +5232,12 @@ dependencies = [ [[package]] name = "md-5" -version = "0.11.0-rc.5" +version = "0.11.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "59e715bb6f273068fc89403d6c4f5eeb83708c62b74c8d43e3e8772ca73a6288" +checksum = "69b6441f590336821bb897fb28fc622898ccceb1d6cea3fde5ea86b090c4de98" dependencies = [ "cfg-if", - "digest 0.11.1", + "digest 0.11.2", ] [[package]] @@ -6257,8 +6256,8 @@ version = "0.13.0-rc.9" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "c8dfa4e14084d963d35bfb4cdb38712cde78dcf83054c0e8b9b8e899150f374e" dependencies = [ - "digest 0.11.1", - "hmac 0.13.0-rc.5", + "digest 0.11.2", + "hmac 0.13.0", ] [[package]] @@ -7474,7 +7473,7 @@ dependencies = [ "const-oid 0.10.2", "crypto-bigint 0.7.1", "crypto-primes", - "digest 0.11.1", + "digest 0.11.2", "pkcs1 0.8.0-rc.4", "pkcs8 0.11.0-rc.11", "rand_core 0.10.0", @@ -7648,7 +7647,7 @@ dependencies = [ "serde_json", "serde_urlencoded", "serial_test", - "sha2 0.11.0-rc.5", + "sha2 0.11.0", "shadow-rs", "socket2", "starshard", @@ -7718,10 +7717,10 @@ dependencies = [ "bytes", "crc-fast", "http 1.4.0", - "md-5 0.11.0-rc.5", + "md-5 0.11.0", "pretty_assertions", - "sha1 0.11.0-rc.5", - "sha2 0.11.0-rc.5", + "sha1 0.11.0", + "sha2 0.11.0", ] [[package]] @@ -7785,7 +7784,7 @@ dependencies = [ "pbkdf2 0.13.0-rc.9", "rand 0.10.0", "serde_json", - "sha2 0.11.0-rc.5", + "sha2 0.11.0", "test-case", "thiserror 2.0.18", "time", @@ -7819,7 +7818,7 @@ dependencies = [ "google-cloud-auth", "google-cloud-storage", "hex-simd", - "hmac 0.13.0-rc.5", + "hmac 0.13.0", "http 1.4.0", "http-body 1.0.1", "http-body-util", @@ -7827,7 +7826,7 @@ dependencies = [ "hyper-rustls", "hyper-util", "lazy_static", - "md-5 0.11.0-rc.5", + "md-5 0.11.0", "memmap2 0.9.10", "metrics", "num_cpus", @@ -7864,8 +7863,8 @@ dependencies = [ "serde_json", "serde_urlencoded", "serial_test", - "sha1 0.11.0-rc.5", - "sha2 0.11.0-rc.5", + "sha1 0.11.0", + "sha2 0.11.0", "shadow-rs", "smallvec", "temp-env", @@ -8028,7 +8027,7 @@ dependencies = [ "rustfs-utils", "serde", "serde_json", - "sha2 0.11.0-rc.5", + "sha2 0.11.0", "temp-env", "tempfile", "thiserror 2.0.18", @@ -8215,7 +8214,7 @@ dependencies = [ "futures", "futures-util", "hex", - "hmac 0.13.0-rc.5", + "hmac 0.13.0", "http 1.4.0", "http-body-util", "hyper", @@ -8236,8 +8235,8 @@ dependencies = [ "s3s", "serde", "serde_json", - "sha1 0.11.0-rc.5", - "sha2 0.11.0-rc.5", + "sha1 0.11.0", + "sha2 0.11.0", "thiserror 2.0.18", "time", "tokio", @@ -8277,7 +8276,7 @@ dependencies = [ "hex-simd", "http 1.4.0", "http-body-util", - "md-5 0.11.0-rc.5", + "md-5 0.11.0", "pin-project-lite", "rand 0.10.0", "reqwest 0.13.2", @@ -8287,8 +8286,8 @@ dependencies = [ "s3s", "serde", "serde_json", - "sha1 0.11.0-rc.5", - "sha2 0.11.0-rc.5", + "sha1 0.11.0", + "sha2 0.11.0", "thiserror 2.0.18", "tokio", "tokio-test", @@ -8452,13 +8451,13 @@ dependencies = [ "hashbrown 0.16.1", "hex-simd", "highway", - "hmac 0.13.0-rc.5", + "hmac 0.13.0", "http 1.4.0", "hyper", "libc", "local-ip-address", "lz4", - "md-5 0.11.0-rc.5", + "md-5 0.11.0", "netif", "rand 0.10.0", "regex", @@ -8468,8 +8467,8 @@ dependencies = [ "rustls-pki-types", "s3s", "serde", - "sha1 0.11.0-rc.5", - "sha2 0.11.0-rc.5", + "sha1 0.11.0", + "sha2 0.11.0", "siphasher", "snap", "sysinfo", @@ -8568,7 +8567,7 @@ dependencies = [ "errno", "libc", "linux-raw-sys 0.12.1", - "windows-sys 0.52.0", + "windows-sys 0.59.0", ] [[package]] @@ -8636,7 +8635,7 @@ dependencies = [ "security-framework", "security-framework-sys", "webpki-root-certs", - "windows-sys 0.52.0", + "windows-sys 0.59.0", ] [[package]] @@ -8683,7 +8682,7 @@ checksum = "9774ba4a74de5f7b1c1451ed6cd5285a32eddb5cccb8cc655a4e50009e06477f" [[package]] name = "s3s" version = "0.14.0-dev" -source = "git+https://github.com/rustfs/s3s?rev=f1815ced732e180f71935feee6ae5ef44fe39b22#f1815ced732e180f71935feee6ae5ef44fe39b22" +source = "git+https://github.com/rustfs/s3s?rev=738f85792c92781bd8af862a074d7379d9fbfabc#738f85792c92781bd8af862a074d7379d9fbfabc" dependencies = [ "arc-swap", "arrayvec", @@ -8697,14 +8696,14 @@ dependencies = [ "crc-fast", "futures", "hex-simd", - "hmac 0.13.0-rc.5", + "hmac 0.13.0", "http 1.4.0", "http-body 1.0.1", "http-body-util", "httparse", "hyper", "itoa", - "md-5 0.11.0-rc.5", + "md-5 0.11.0", "memchr", "mime", "nom 8.0.0", @@ -8714,8 +8713,8 @@ dependencies = [ "serde", "serde_json", "serde_urlencoded", - "sha1 0.11.0-rc.5", - "sha2 0.11.0-rc.5", + "sha1 0.11.0", + "sha2 0.11.0", "smallvec", "std-next", "subtle", @@ -9051,13 +9050,13 @@ dependencies = [ [[package]] name = "sha1" -version = "0.11.0-rc.5" +version = "0.11.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "3b167252f3c126be0d8926639c4c4706950f01445900c4b3db0fd7e89fcb750a" +checksum = "aacc4cc499359472b4abe1bf11d0b12e688af9a805fa5e3016f9a386dc2d0214" dependencies = [ "cfg-if", - "cpufeatures 0.2.17", - "digest 0.11.1", + "cpufeatures 0.3.0", + "digest 0.11.2", ] [[package]] @@ -9073,13 +9072,13 @@ dependencies = [ [[package]] name = "sha2" -version = "0.11.0-rc.5" +version = "0.11.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "7c5f3b1e2dc8aad28310d8410bd4d7e180eca65fca176c52ab00d364475d0024" +checksum = "446ba717509524cb3f22f17ecc096f10f4822d76ab5c0b9822c5f9c284e825f4" dependencies = [ "cfg-if", - "cpufeatures 0.2.17", - "digest 0.11.1", + "cpufeatures 0.3.0", + "digest 0.11.2", ] [[package]] @@ -9156,7 +9155,7 @@ version = "3.0.0-rc.10" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "7f1880df446116126965eeec169136b2e0251dba37c6223bcc819569550edea3" dependencies = [ - "digest 0.11.1", + "digest 0.11.2", "rand_core 0.10.0", ] @@ -9641,7 +9640,7 @@ dependencies = [ "getrandom 0.4.2", "once_cell", "rustix 1.1.4", - "windows-sys 0.52.0", + "windows-sys 0.59.0", ] [[package]] @@ -10600,7 +10599,7 @@ version = "0.1.11" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "c2a7b1c03c876122aa43f3020e6c3c3ee5c05081c9a00739faf7503aeba10d22" dependencies = [ - "windows-sys 0.52.0", + "windows-sys 0.59.0", ] [[package]] diff --git a/Cargo.toml b/Cargo.toml index eb8491d5d..143f3e762 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -250,7 +250,7 @@ rumqttc = { version = "0.25.1" } rustix = { version = "1.1.4", features = ["fs"] } rust-embed = { version = "8.11.0" } rustc-hash = { version = "2.1.2" } -s3s = { git = "https://github.com/rustfs/s3s", rev = "f1815ced732e180f71935feee6ae5ef44fe39b22", features = ["minio"] } +s3s = { git = "https://github.com/rustfs/s3s", rev = "738f85792c92781bd8af862a074d7379d9fbfabc", features = ["minio"] } serial_test = "3.4.0" shadow-rs = { version = "1.7.1", default-features = false } siphasher = "1.0.2" From 5d302febb772c70a661b7e9817172347b25f8ec9 Mon Sep 17 00:00:00 2001 From: weisd Date: Fri, 3 Apr 2026 13:57:42 +0800 Subject: [PATCH 03/22] 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, From 1fe036cb70a22005391e70ba5baa000887610b63 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E5=AE=89=E6=AD=A3=E8=B6=85?= Date: Fri, 3 Apr 2026 20:49:09 +0800 Subject: [PATCH 04/22] ci: update CLA workflow for corrected comments (#2384) --- .github/workflows/cla.yml | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/.github/workflows/cla.yml b/.github/workflows/cla.yml index 94bb97a03..8f5b1f44c 100644 --- a/.github/workflows/cla.yml +++ b/.github/workflows/cla.yml @@ -18,7 +18,7 @@ on: pull_request_target: types: [opened, synchronize, reopened] issue_comment: - types: [created] + types: [created, edited] permissions: contents: write @@ -42,7 +42,7 @@ jobs: permission-contents: write - name: Run CLA Bot - uses: overtrue/cla-bot@v0.0.6 + uses: overtrue/cla-bot@v0.0.8 with: github-token: ${{ github.token }} registry-token: ${{ steps.registry-token.outputs.token }} From c4efb46827090db8e1f4a4c161251526ebdd73ea Mon Sep 17 00:00:00 2001 From: Andy Brown Date: Fri, 3 Apr 2026 14:09:05 +0100 Subject: [PATCH 05/22] fix(notify): emit delete webhooks for prefix deletes and align replication headers (#2383) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Co-authored-by: 安正超 --- crates/notify/src/event.rs | 51 ++++++++++++++++++++++++++++++-- rustfs/src/app/object_usecase.rs | 20 +++++++++++-- rustfs/src/storage/helper.rs | 2 +- 3 files changed, 67 insertions(+), 6 deletions(-) diff --git a/crates/notify/src/event.rs b/crates/notify/src/event.rs index 0cf4662ee..ee803bd87 100644 --- a/crates/notify/src/event.rs +++ b/crates/notify/src/event.rs @@ -336,9 +336,21 @@ pub struct EventArgs { } impl EventArgs { - // Helper function to check if it is a copy request + /// True when the RustFS replication header is explicitly enabled (`true` or `1`). + /// + /// Only `x-rustfs-source-replication-request` is considered here. Many clients (including the + /// console) send `x-minio-source-replication-request` for MinIO compatibility; treating that + /// as replication would suppress webhooks on normal browser deletes. Storage still honors both + /// prefixes when parsing the typed HTTP headers for `ObjectOptions`. pub fn is_replication_request(&self) -> bool { - self.req_params.contains_key("x-rustfs-source-replication-request") + self.replication_header_value_true("x-rustfs-source-replication-request") + } + + fn replication_header_value_true(&self, key: &str) -> bool { + self.req_params + .get(key) + .map(|v| v.eq_ignore_ascii_case("true") || v == "1") + .unwrap_or(false) } } @@ -524,3 +536,38 @@ mod tests { assert_eq!(glacier.restore_event_data.lifecycle_restore_storage_class, "GLACIER"); } } + +#[cfg(test)] +mod event_args_tests { + use super::EventArgs; + use hashbrown::HashMap; + use rustfs_ecstore::store_api::ObjectInfo; + use rustfs_s3_common::EventName; + + fn args_with_headers(pairs: &[(&str, &str)]) -> EventArgs { + let mut req_params = HashMap::new(); + for (k, v) in pairs { + req_params.insert((*k).to_string(), (*v).to_string()); + } + EventArgs { + event_name: EventName::ObjectRemovedDelete, + bucket_name: "b".to_string(), + object: ObjectInfo::default(), + req_params, + resp_elements: HashMap::new(), + version_id: String::new(), + host: String::new(), + port: 0, + user_agent: String::new(), + } + } + + #[test] + fn replication_request_requires_true_value() { + assert!(!args_with_headers(&[("x-rustfs-source-replication-request", "")]).is_replication_request()); + assert!(!args_with_headers(&[("x-rustfs-source-replication-request", "false")]).is_replication_request()); + assert!(args_with_headers(&[("x-rustfs-source-replication-request", "true")]).is_replication_request()); + assert!(args_with_headers(&[("x-rustfs-source-replication-request", "True")]).is_replication_request()); + assert!(!args_with_headers(&[("x-minio-source-replication-request", "true")]).is_replication_request()); + } +} diff --git a/rustfs/src/app/object_usecase.rs b/rustfs/src/app/object_usecase.rs index f3df0aa23..f97b3519c 100644 --- a/rustfs/src/app/object_usecase.rs +++ b/rustfs/src/app/object_usecase.rs @@ -24,7 +24,7 @@ use crate::storage::concurrency::{ }; use crate::storage::ecfs::*; use crate::storage::head_prefix::{head_prefix_not_found_message, probe_prefix_has_children}; -use crate::storage::helper::OperationHelper; +use crate::storage::helper::{OperationHelper, spawn_background}; use crate::storage::options::{ copy_dst_opts, copy_src_opts, del_opts, extract_metadata, extract_metadata_from_mime_with_object_name, filter_object_metadata, get_content_sha256_with_query, get_opts, normalize_content_encoding_for_storage, put_opts, @@ -3824,7 +3824,7 @@ impl DefaultObjectUsecase { .as_ref() .map(|context| context.notify()) .unwrap_or_else(default_notify_interface); - tokio::spawn(async move { + spawn_background(async move { for res in delete_results { if let Some(dobj) = res.delete_object { let event_name = if dobj.delete_marker { @@ -3996,7 +3996,21 @@ impl DefaultObjectUsecase { }) .await; } - return Ok(S3Response::with_status(DeleteObjectOutput::default(), StatusCode::NO_CONTENT)); + // Prefix/force-delete returns empty ObjectInfo; still emit bucket notification so webhooks match S3 DELETE. + helper = helper + .event_name(EventName::ObjectRemovedDelete) + .object(ObjectInfo { + name: key.clone(), + bucket: bucket.clone(), + ..Default::default() + }) + .version_id(String::new()); + let result = Ok(S3Response::with_status(DeleteObjectOutput::default(), StatusCode::NO_CONTENT)); + // Match non-empty delete path: capacity manager write-op telemetry. + let manager = get_capacity_manager(); + manager.record_write_operation().await; + let _ = helper.complete(&result); + return result; } if obj_info.replication_status == ReplicationStatusType::Replica diff --git a/rustfs/src/storage/helper.rs b/rustfs/src/storage/helper.rs index b4466f538..4a2bd7352 100644 --- a/rustfs/src/storage/helper.rs +++ b/rustfs/src/storage/helper.rs @@ -31,7 +31,7 @@ use tokio::runtime::{Builder, Handle}; /// Schedules an asynchronous task on the current runtime; /// if there is no runtime, creates a minimal runtime execution on a new thread. -fn spawn_background(fut: F) +pub(crate) fn spawn_background(fut: F) where F: Future + Send + 'static, { From b3f31ad694c3d6150c3621b7ba059e64c0cb3a7a Mon Sep 17 00:00:00 2001 From: xxkeming Date: Fri, 3 Apr 2026 21:09:30 +0800 Subject: [PATCH 06/22] fix(build): enable tokio_unstable in build script (#2376) Co-authored-by: houseme --- build-rustfs.sh | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/build-rustfs.sh b/build-rustfs.sh index 0ed3352d6..82100ff68 100755 --- a/build-rustfs.sh +++ b/build-rustfs.sh @@ -213,7 +213,7 @@ setup_rust_environment() { # Set up environment variables for musl targets if [[ "$PLATFORM" == *"musl"* ]]; then print_message $YELLOW "Setting up environment for musl target..." - export RUSTFLAGS="-C target-feature=-crt-static" + export RUSTFLAGS="'--cfg tokio_unstable -C target-feature=-crt-static'" # For cargo-zigbuild, set up additional environment variables if command -v cargo-zigbuild &> /dev/null; then @@ -430,7 +430,7 @@ build_binary() { fi else # Native compilation - build_cmd="RUSTFLAGS=-Clink-arg=-lm cargo build" + build_cmd="RUSTFLAGS='--cfg tokio_unstable -Clink-arg=-lm' cargo build" fi if [ "$BUILD_TYPE" = "release" ]; then From c2449433136230275ef481e87fa55004dcea740d Mon Sep 17 00:00:00 2001 From: GatewayJ <835269233@qq.com> Date: Fri, 3 Apr 2026 21:10:27 +0800 Subject: [PATCH 07/22] feat(iam): retry OIDC discovery with issuer URL slash variants (#2360) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Co-authored-by: GatewayJ <8352692332qq.com> Co-authored-by: 安正超 --- crates/iam/src/oidc.rs | 236 +++++++++++++++++++++++++++++++++++++++-- 1 file changed, 227 insertions(+), 9 deletions(-) diff --git a/crates/iam/src/oidc.rs b/crates/iam/src/oidc.rs index b3d68016a..f28f3b9ff 100644 --- a/crates/iam/src/oidc.rs +++ b/crates/iam/src/oidc.rs @@ -833,18 +833,38 @@ impl OidcSys { async fn discover_provider(config: &OidcProviderConfig, http_client: &ReqwestHttpClient) -> Result { // The openidconnect crate expects the issuer URL (base), not the // .well-known/openid-configuration URL. - let issuer_str = normalize_config_url(&config.config_url)?; + let base_issuer = normalize_config_url(&config.config_url)?; + let candidates = issuer_candidates(&base_issuer); + let mut last_errors = Vec::new(); - let issuer_url = IssuerUrl::new(issuer_str).map_err(|e| format!("invalid issuer URL: {e}"))?; + for candidate_issuer in candidates.iter() { + let issuer_url = IssuerUrl::new(candidate_issuer.clone()).map_err(|e| format!("invalid issuer URL: {e}"))?; - let metadata = CoreProviderMetadata::discover_async(issuer_url, http_client) - .await - .map_err(|e| format!("discovery failed: {e}"))?; + match CoreProviderMetadata::discover_async(issuer_url, http_client) + .await + .map_err(|e| format!("discovery failed: {e}")) + { + Ok(metadata) => { + return Ok(ProviderState { + metadata, + discovered_at: Instant::now(), + }); + } + Err(error) => { + last_errors.push(format!("issuer '{candidate_issuer}': {error}")); + warn!( + "OIDC provider '{}' discovery attempt failed for issuer '{}': {}", + config.id, candidate_issuer, error + ); + } + } + } - Ok(ProviderState { - metadata, - discovered_at: Instant::now(), - }) + Err(format!( + "discovery failed for all issuer variants {:?}: {}", + candidates, + last_errors.join("; ") + )) } } @@ -959,6 +979,21 @@ fn normalize_config_url(config_url: &str) -> Result { Ok(issuer) } +fn issuer_candidates(base: &str) -> Vec { + let original = base.trim(); + let mut variants = Vec::with_capacity(2); + variants.push(original.to_string()); + + let toggled = if original.ends_with('/') { + original.trim_end_matches('/').to_string() + } else { + format!("{original}/") + }; + variants.push(toggled); + + variants +} + /// Decode the payload section of a JWT without validation (token must already be verified). pub(crate) fn decode_jwt_payload(token: &str) -> HashMap { let parts: Vec<&str> = token.split('.').collect(); @@ -1116,6 +1151,189 @@ mod tests { assert!(normalize_config_url("not-a-url").is_err()); } + #[test] + fn test_issuer_candidates() { + assert_eq!( + issuer_candidates("https://idp.example.com/realm"), + vec![ + "https://idp.example.com/realm".to_string(), + "https://idp.example.com/realm/".to_string() + ] + ); + assert_eq!( + issuer_candidates("https://idp.example.com/realm/"), + vec![ + "https://idp.example.com/realm/".to_string(), + "https://idp.example.com/realm".to_string() + ] + ); + assert_eq!( + issuer_candidates("https://idp.example.com"), + vec!["https://idp.example.com".to_string(), "https://idp.example.com/".to_string()] + ); + } + + fn build_mocked_oidc_provider_config(id: &str, config_url: &str) -> OidcProviderConfig { + OidcProviderConfig { + id: id.to_string(), + enabled: true, + config_url: config_url.to_string(), + client_id: "rustfs-oidc-test".to_string(), + client_secret: None, + scopes: vec!["openid".to_string()], + redirect_uri: None, + redirect_uri_dynamic: false, + claim_name: "sub".to_string(), + claim_prefix: "oidc".to_string(), + role_policy: String::new(), + display_name: "mock-oidc".to_string(), + groups_claim: "groups".to_string(), + email_claim: "email".to_string(), + username_claim: "username".to_string(), + } + } + + fn start_mock_oidc_discovery_server( + build_discovery_issuer: F, + max_requests: usize, + ) -> (String, std::thread::JoinHandle<()>) + where + F: Fn(&str) -> String + Send + 'static, + { + use std::io::Read; + use std::io::Write; + use std::net::{Shutdown, TcpListener}; + use std::time::{Duration, Instant}; + + // After the last completed response, exit if no new connection arrives within this window. + const IDLE_SHUTDOWN: Duration = Duration::from_millis(100); + const ABSOLUTE_CAP: Duration = Duration::from_millis(500); + + let listener = TcpListener::bind("127.0.0.1:0").unwrap(); + let base = format!("http://{}", listener.local_addr().unwrap()); + let discovery_issuer = build_discovery_issuer(&base); + let discovery_body = serde_json::json!({ + "issuer": discovery_issuer, + "authorization_endpoint": format!("{base}/authorize"), + "token_endpoint": format!("{base}/token"), + "jwks_uri": format!("{base}/jwks"), + "response_types_supported": ["code"], + "response_modes_supported": ["query"], + "subject_types_supported": ["public"], + "id_token_signing_alg_values_supported": ["RS256"], + }) + .to_string(); + let jwks_body = r#"{"keys":[]}"#; + + let handle = std::thread::spawn(move || { + listener + .set_nonblocking(true) + .expect("failed to set discovery mock listener non-blocking"); + + let mut seen = 0usize; + let start = Instant::now(); + let mut last_completed = Instant::now(); + + loop { + if seen > 0 && last_completed.elapsed() >= IDLE_SHUTDOWN { + break; + } + if start.elapsed() >= ABSOLUTE_CAP { + break; + } + + let mut stream = match listener.accept() { + Ok((stream, _)) => stream, + Err(e) if e.kind() == std::io::ErrorKind::WouldBlock => { + std::thread::sleep(Duration::from_millis(5)); + continue; + } + Err(_) => break, + }; + + seen += 1; + + let mut request_bytes = Vec::new(); + let mut buffer = [0u8; 4096]; + loop { + let n = stream.read(&mut buffer).unwrap_or_default(); + if n == 0 { + break; + } + request_bytes.extend_from_slice(&buffer[..n]); + if request_bytes.windows(4).any(|w| w == b"\r\n\r\n") { + break; + } + if request_bytes.len() >= 8192 { + break; + } + } + let request = String::from_utf8_lossy(&request_bytes); + let path = request.lines().next().unwrap_or("").split_whitespace().nth(1).unwrap_or(""); + + let (status, body) = if path.contains("/.well-known/openid-configuration") { + (200, discovery_body.as_str()) + } else if path.contains("/jwks") { + (200, jwks_body) + } else { + (404, r#"{"error":"not found"}"#) + }; + + let response = format!( + "HTTP/1.1 {status} {}\r\nContent-Type: application/json\r\nContent-Length: {}\r\nConnection: close\r\n\r\n{body}", + if status == 200 { "OK" } else { "Not Found" }, + body.len() + ); + + let _ = stream.write_all(response.as_bytes()); + let _ = stream.flush(); + let _ = stream.shutdown(Shutdown::Both); + last_completed = Instant::now(); + + if seen >= max_requests { + break; + } + } + }); + + (base, handle) + } + + fn discovery_error_contains_all_variants(err: &str, base: &str) -> bool { + err.contains(base) && err.contains(&format!("{base}/")) && err.contains("discovery failed for all issuer variants") + } + + #[tokio::test] + async fn test_validate_oidc_provider_config_retries_with_issuer_candidates() { + // Discovery document must advertise the canonical issuer path. The first candidate has no + // trailing slash; openidconnect rejects issuer mismatch, then the second variant succeeds. + let (base, handle) = start_mock_oidc_discovery_server(|base| format!("{base}/application/o/rustfs/"), 8); + let config_url = format!("{base}/application/o/rustfs"); + let config = build_mocked_oidc_provider_config("default", &config_url); + + let result = validate_oidc_provider_config(&config).await; + + let validation_result = result.expect("OIDC provider validation should succeed"); + assert_eq!(validation_result.issuer, format!("{base}/application/o/rustfs/")); + assert!(handle.join().is_ok()); + } + + #[tokio::test] + async fn test_validate_oidc_provider_config_returns_detailed_errors() { + let (base, handle) = start_mock_oidc_discovery_server(|base| format!("{base}/application/o/other"), 8); + let config_url = format!("{base}/application/o/rustfs"); + let config = build_mocked_oidc_provider_config("default", &config_url); + + let err = validate_oidc_provider_config(&config) + .await + .expect_err("OIDC provider validation should fail"); + assert!(discovery_error_contains_all_variants(&err, &base)); + assert!(err.contains("issuer '")); + assert!(err.contains(&format!("issuer '{base}/application/o/rustfs'"))); + assert!(err.contains(&format!("issuer '{base}/application/o/rustfs/'"))); + assert!(handle.join().is_ok()); + } + #[test] fn test_decode_jwt_payload_invalid() { assert!(decode_jwt_payload("not-a-jwt").is_empty()); From 2d91e2f580eb600962dbb7202f0f31279d0f201c Mon Sep 17 00:00:00 2001 From: Logan Ye Date: Fri, 3 Apr 2026 21:45:56 +0800 Subject: [PATCH 08/22] fix(oidc): support case-insensitive claim name matching (#2362) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Co-authored-by: loverustfs Co-authored-by: 安正超 --- crates/iam/src/oidc.rs | 58 +++++++++++++++++++++++++++--- crates/policy/src/policy/policy.rs | 53 ++++++++++++++++++++++++++- 2 files changed, 105 insertions(+), 6 deletions(-) diff --git a/crates/iam/src/oidc.rs b/crates/iam/src/oidc.rs index f28f3b9ff..cac4c7bb6 100644 --- a/crates/iam/src/oidc.rs +++ b/crates/iam/src/oidc.rs @@ -1007,14 +1007,27 @@ pub(crate) fn decode_jwt_payload(token: &str) -> HashMap, key: &str) -> String { - claims.get(key).and_then(|v| v.as_str()).unwrap_or_default().to_string() +/// Get a claim value from raw claims with case-insensitive fallback. +/// First tries exact match, then falls back to case-insensitive match if not found. +fn get_claim_case_insensitive<'a>(claims: &'a HashMap, key: &str) -> Option<&'a serde_json::Value> { + if let Some(v) = claims.get(key) { + return Some(v); + } + let key_lower = key.to_lowercase(); + claims.iter().find(|(k, _)| k.to_lowercase() == key_lower).map(|(_, v)| v) } -/// Extract a groups/array claim from raw claims. Handles both string arrays and single strings. +/// Extract a string claim from raw claims with case-insensitive fallback. +fn extract_string_claim(claims: &HashMap, key: &str) -> String { + get_claim_case_insensitive(claims, key) + .and_then(serde_json::Value::as_str) + .unwrap_or_default() + .to_string() +} + +/// Extract a groups/array claim from raw claims with case-insensitive fallback. Handles both string arrays and single strings. fn extract_groups_claim(claims: &HashMap, key: &str) -> Vec { - match claims.get(key) { + match get_claim_case_insensitive(claims, key) { Some(serde_json::Value::Array(arr)) => arr.iter().filter_map(|v| v.as_str().map(String::from)).collect(), Some(serde_json::Value::String(s)) => s.split(',').map(|s| s.trim().to_string()).collect(), _ => vec![], @@ -1069,6 +1082,41 @@ mod tests { assert!(groups.is_empty()); } + #[test] + fn test_extract_string_claim_case_insensitive() { + let mut claims = HashMap::new(); + claims.insert("policyminio".to_string(), serde_json::json!("consoleAdmin")); + + assert_eq!(extract_string_claim(&claims, "policyMinio"), "consoleAdmin"); + assert_eq!(extract_string_claim(&claims, "POLICYMINIO"), "consoleAdmin"); + assert_eq!(extract_string_claim(&claims, "policyminio"), "consoleAdmin"); + } + + #[test] + fn test_extract_groups_claim_case_insensitive() { + let mut claims = HashMap::new(); + claims.insert("policyminio".to_string(), serde_json::json!(["consoleAdmin", "readwrite"])); + + let groups = extract_groups_claim(&claims, "policyMinio"); + assert_eq!(groups, vec!["consoleAdmin", "readwrite"]); + + let groups = extract_groups_claim(&claims, "POLICYMINIO"); + assert_eq!(groups, vec!["consoleAdmin", "readwrite"]); + + let groups = extract_groups_claim(&claims, "policyminio"); + assert_eq!(groups, vec!["consoleAdmin", "readwrite"]); + } + + #[test] + fn test_extract_groups_claim_exact_match_preferred() { + let mut claims = HashMap::new(); + claims.insert("Policy".to_string(), serde_json::json!(["exact_match"])); + claims.insert("policy".to_string(), serde_json::json!(["lowercase"])); + + let groups = extract_groups_claim(&claims, "Policy"); + assert_eq!(groups, vec!["exact_match"]); + } + #[test] fn test_decode_jwt_payload() { let payload = r#"{"sub":"user123","email":"user@example.com"}"#; diff --git a/crates/policy/src/policy/policy.rs b/crates/policy/src/policy/policy.rs index 6464fcfd3..88201e66b 100644 --- a/crates/policy/src/policy/policy.rs +++ b/crates/policy/src/policy/policy.rs @@ -239,9 +239,20 @@ impl Validator for BucketPolicy { } } +fn get_claim_case_insensitive<'a>(claims: &'a HashMap, claim_name: &str) -> Option<&'a Value> { + if let Some(v) = claims.get(claim_name) { + return Some(v); + } + let claim_name_lower = claim_name.to_lowercase(); + claims + .iter() + .find(|(k, _)| k.to_lowercase() == claim_name_lower) + .map(|(_, v)| v) +} + fn get_values_from_claims(claims: &HashMap, claim_name: &str) -> (HashSet, bool) { let mut s = HashSet::new(); - if let Some(pname) = claims.get(claim_name) { + if let Some(pname) = get_claim_case_insensitive(claims, claim_name) { if let Some(pnames) = pname.as_array() { for pname in pnames { if let Some(pname_str) = pname.as_str() { @@ -1693,4 +1704,44 @@ mod test { "principal and resource match should keep ExistingObjectTag fetch hint" ); } + + #[test] + fn test_get_values_from_claims_case_insensitive() { + let mut claims = HashMap::new(); + claims.insert("policyminio".to_string(), Value::Array(vec![Value::String("consoleAdmin".to_string())])); + + let (policies, found) = get_values_from_claims(&claims, "policyMinio"); + assert!(found); + assert!(policies.contains("consoleAdmin")); + + let (policies, found) = get_values_from_claims(&claims, "POLICYMINIO"); + assert!(found); + assert!(policies.contains("consoleAdmin")); + + let (policies, found) = get_values_from_claims(&claims, "policyminio"); + assert!(found); + assert!(policies.contains("consoleAdmin")); + } + + #[test] + fn test_get_values_from_claims_exact_match_preferred() { + let mut claims = HashMap::new(); + claims.insert("Policy".to_string(), Value::Array(vec![Value::String("exact_match".to_string())])); + claims.insert("policy".to_string(), Value::Array(vec![Value::String("lowercase".to_string())])); + + let (policies, _) = get_values_from_claims(&claims, "Policy"); + assert!(policies.contains("exact_match")); + assert!(!policies.contains("lowercase")); + } + + #[test] + fn test_get_policies_from_claims_case_insensitive_string() { + let mut claims = HashMap::new(); + claims.insert("policyminio".to_string(), Value::String("consoleAdmin,readwrite".to_string())); + + let (policies, found) = get_policies_from_claims(&claims, "policyMinio"); + assert!(found); + assert!(policies.contains("consoleAdmin")); + assert!(policies.contains("readwrite")); + } } From 25512e2635ad61fcb6c87ea00f8c75b9d50f5916 Mon Sep 17 00:00:00 2001 From: weisd Date: Fri, 3 Apr 2026 21:46:51 +0800 Subject: [PATCH 09/22] perf(ecstore): batch delete object lock acquisition (#2374) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Co-authored-by: 安正超 --- .../e2e_test/src/reliant/grpc_lock_client.rs | 148 ++++++++--- .../e2e_test/src/reliant/grpc_lock_server.rs | 105 +++++++- crates/e2e_test/src/reliant/lock.rs | 56 +++- crates/ecstore/src/rpc/remote_locker.rs | 148 ++++++++--- crates/ecstore/src/set_disk.rs | 243 ++++++++++++++++-- crates/ecstore/src/sets.rs | 125 ++++++++- crates/lock/src/client/mod.rs | 17 ++ crates/lock/src/distributed_lock.rs | 25 +- crates/lock/src/fast_lock/manager.rs | 77 ++++-- crates/lock/src/namespace/tests.rs | 70 +++++ .../src/generated/proto_gen/node_service.rs | 114 ++++++++ crates/protos/src/node.proto | 16 ++ rustfs/src/storage/rpc/lock.rs | 104 ++++++++ rustfs/src/storage/rpc/node_service.rs | 14 + 14 files changed, 1120 insertions(+), 142 deletions(-) diff --git a/crates/e2e_test/src/reliant/grpc_lock_client.rs b/crates/e2e_test/src/reliant/grpc_lock_client.rs index fc772d520..1045dbf5e 100644 --- a/crates/e2e_test/src/reliant/grpc_lock_client.rs +++ b/crates/e2e_test/src/reliant/grpc_lock_client.rs @@ -21,7 +21,7 @@ use rustfs_lock::{ LockClient, LockError, LockId, LockInfo, LockRequest, LockResponse, LockStats, LockStatus, LockType, Result, types::{LockMetadata, LockPriority}, }; -use rustfs_protos::proto_gen::node_service::{GenerallyLockRequest, PingRequest}; +use rustfs_protos::proto_gen::node_service::{BatchGenerallyLockRequest, GenerallyLockRequest, PingRequest}; use tonic::Request; use tracing::{info, warn}; @@ -64,6 +64,44 @@ impl GrpcLockClient { suppress_contention_logs: false, } } + + fn build_lock_info(request: &LockRequest, lock_info_json: Option) -> LockInfo { + if let Some(lock_info_json) = lock_info_json { + match serde_json::from_str::(&lock_info_json) { + Ok(info) => info, + Err(e) => { + warn!("Failed to deserialize lock_info from response: {}, using request data", e); + LockInfo { + id: request.lock_id.clone(), + resource: request.resource.clone(), + lock_type: request.lock_type, + status: LockStatus::Acquired, + owner: request.owner.clone(), + acquired_at: std::time::SystemTime::now(), + expires_at: std::time::SystemTime::now() + request.ttl, + last_refreshed: std::time::SystemTime::now(), + metadata: request.metadata.clone(), + priority: request.priority, + wait_start_time: None, + } + } + } + } else { + LockInfo { + id: request.lock_id.clone(), + resource: request.resource.clone(), + lock_type: request.lock_type, + status: LockStatus::Acquired, + owner: request.owner.clone(), + acquired_at: std::time::SystemTime::now(), + expires_at: std::time::SystemTime::now() + request.ttl, + last_refreshed: std::time::SystemTime::now(), + metadata: request.metadata.clone(), + priority: request.priority, + wait_start_time: None, + } + } + } } #[async_trait] @@ -89,46 +127,10 @@ impl LockClient for GrpcLockClient { // Check if the lock acquisition was successful if resp.success { - // Try to deserialize lock_info from response - let lock_info = if let Some(lock_info_json) = resp.lock_info { - match serde_json::from_str::(&lock_info_json) { - Ok(info) => info, - Err(e) => { - // If deserialization fails, fall back to constructing from request - warn!("Failed to deserialize lock_info from response: {}, using request data", e); - LockInfo { - id: request.lock_id.clone(), - resource: request.resource.clone(), - lock_type: request.lock_type, - status: LockStatus::Acquired, - owner: request.owner.clone(), - acquired_at: std::time::SystemTime::now(), - expires_at: std::time::SystemTime::now() + request.ttl, - last_refreshed: std::time::SystemTime::now(), - metadata: request.metadata.clone(), - priority: request.priority, - wait_start_time: None, - } - } - } - } else { - // If lock_info is not provided, construct from request - LockInfo { - id: request.lock_id.clone(), - resource: request.resource.clone(), - lock_type: request.lock_type, - status: LockStatus::Acquired, - owner: request.owner.clone(), - acquired_at: std::time::SystemTime::now(), - expires_at: std::time::SystemTime::now() + request.ttl, - last_refreshed: std::time::SystemTime::now(), - metadata: request.metadata.clone(), - priority: request.priority, - wait_start_time: None, - } - }; - - Ok(LockResponse::success(lock_info, std::time::Duration::ZERO)) + Ok(LockResponse::success( + Self::build_lock_info(request, resp.lock_info), + std::time::Duration::ZERO, + )) } else { // Lock acquisition failed Ok(LockResponse::failure( @@ -138,6 +140,45 @@ impl LockClient for GrpcLockClient { } } + async fn acquire_locks_batch(&self, requests: &[LockRequest]) -> Result> { + let mut client = self.get_client().await?; + let req = Request::new(BatchGenerallyLockRequest { + args: requests + .iter() + .map(|request| { + serde_json::to_string(request).map_err(|e| LockError::internal(format!("Failed to serialize request: {e}"))) + }) + .collect::>>()?, + }); + + let resp = client + .lock_batch(req) + .await + .map_err(|e| LockError::internal(e.to_string()))? + .into_inner(); + + Ok(requests + .iter() + .enumerate() + .map(|(idx, request)| match resp.results.get(idx) { + Some(result) if result.success => { + LockResponse::success(Self::build_lock_info(request, result.lock_info.clone()), std::time::Duration::ZERO) + } + Some(result) => LockResponse::failure( + result + .error_info + .clone() + .unwrap_or_else(|| "Lock acquisition failed on remote server".to_string()), + std::time::Duration::ZERO, + ), + None => LockResponse::failure( + format!("Lock batch response missing entry for request index {idx}"), + std::time::Duration::ZERO, + ), + }) + .collect()) + } + async fn release(&self, lock_id: &LockId) -> Result { info!("grpc release for {}", lock_id); @@ -161,6 +202,31 @@ impl LockClient for GrpcLockClient { Ok(resp.success) } + async fn release_locks_batch(&self, lock_ids: &[LockId]) -> Result> { + let mut client = self.get_client().await?; + let req = Request::new(BatchGenerallyLockRequest { + args: lock_ids + .iter() + .map(|lock_id| { + serde_json::to_string(&Self::create_unlock_request(lock_id)) + .map_err(|e| LockError::internal(format!("Failed to serialize request: {e}"))) + }) + .collect::>>()?, + }); + + let resp = client + .un_lock_batch(req) + .await + .map_err(|e| LockError::internal(e.to_string()))? + .into_inner(); + + Ok(lock_ids + .iter() + .enumerate() + .map(|(idx, _)| resp.results.get(idx).map(|result| result.success).unwrap_or(false)) + .collect()) + } + async fn refresh(&self, lock_id: &LockId) -> Result { info!("grpc refresh for {}", lock_id); let refresh_request = Self::create_unlock_request(lock_id); diff --git a/crates/e2e_test/src/reliant/grpc_lock_server.rs b/crates/e2e_test/src/reliant/grpc_lock_server.rs index c1a927124..628152e5b 100644 --- a/crates/e2e_test/src/reliant/grpc_lock_server.rs +++ b/crates/e2e_test/src/reliant/grpc_lock_server.rs @@ -21,7 +21,8 @@ use rustfs_lock::{LockClient, LockRequest}; use rustfs_protos::{ models::PingBodyBuilder, proto_gen::node_service::{ - GenerallyLockRequest, GenerallyLockResponse, PingRequest, PingResponse, node_service_server::NodeService, + BatchGenerallyLockRequest, BatchGenerallyLockResponse, GenerallyLockRequest, GenerallyLockResponse, GenerallyLockResult, + PingRequest, PingResponse, node_service_server::NodeService, }, }; use std::pin::Pin; @@ -33,6 +34,22 @@ use tracing::debug; type ResponseStream = Pin> + Send>>; +fn lock_result_from_response(response: rustfs_lock::LockResponse) -> GenerallyLockResult { + GenerallyLockResult { + success: response.success, + error_info: response.error, + lock_info: response.lock_info.and_then(|info| serde_json::to_string(&info).ok()), + } +} + +fn lock_result_from_error(error: impl Into) -> GenerallyLockResult { + GenerallyLockResult { + success: false, + error_info: Some(error.into()), + lock_info: None, + } +} + /// Minimal NodeService implementation that only supports Lock RPCs /// Used for testing distributed lock scenarios with real gRPC #[derive(Debug)] @@ -187,6 +204,92 @@ impl NodeService for MinimalLockNodeService { } } + async fn lock_batch( + &self, + request: Request, + ) -> Result, Status> { + let request = request.into_inner(); + let mut results = vec![lock_result_from_error("request was not processed"); request.args.len()]; + let mut valid_requests = Vec::with_capacity(request.args.len()); + let mut valid_indices = Vec::with_capacity(request.args.len()); + + for (idx, arg) in request.args.iter().enumerate() { + match serde_json::from_str::(arg) { + Ok(args) => { + valid_requests.push(args); + valid_indices.push(idx); + } + Err(err) => { + results[idx] = lock_result_from_error(format!("can not decode args, err: {err}")); + } + } + } + + if !valid_requests.is_empty() { + match self.lock_client.acquire_locks_batch(&valid_requests).await { + Ok(batch_results) => { + for (result_idx, response) in batch_results.into_iter().enumerate() { + if let Some(request_idx) = valid_indices.get(result_idx) { + results[*request_idx] = lock_result_from_response(response); + } + } + } + Err(err) => { + for request_idx in valid_indices { + results[request_idx] = lock_result_from_error(format!("can not batch lock, err: {err}")); + } + } + } + } + + Ok(Response::new(BatchGenerallyLockResponse { results })) + } + + async fn un_lock_batch( + &self, + request: Request, + ) -> Result, Status> { + let request = request.into_inner(); + let mut results = vec![lock_result_from_error("request was not processed"); request.args.len()]; + let mut lock_ids = Vec::with_capacity(request.args.len()); + let mut valid_indices = Vec::with_capacity(request.args.len()); + + for (idx, arg) in request.args.iter().enumerate() { + match serde_json::from_str::(arg) { + Ok(args) => { + lock_ids.push(args.lock_id); + valid_indices.push(idx); + } + Err(err) => { + results[idx] = lock_result_from_error(format!("can not decode args, err: {err}")); + } + } + } + + if !lock_ids.is_empty() { + match self.lock_client.release_locks_batch(&lock_ids).await { + Ok(batch_results) => { + for (result_idx, success) in batch_results.into_iter().enumerate() { + if let Some(request_idx) = valid_indices.get(result_idx) { + results[*request_idx] = GenerallyLockResult { + success, + error_info: None, + lock_info: None, + }; + } + } + } + Err(err) => { + for request_idx in valid_indices { + results[request_idx] = lock_result_from_error(format!("can not batch unlock, err: {err}")); + } + } + } + } + + Ok(Response::new(BatchGenerallyLockResponse { results })) + } + // All other methods return unimplemented async fn heal_bucket( &self, diff --git a/crates/e2e_test/src/reliant/lock.rs b/crates/e2e_test/src/reliant/lock.rs index 7cae2c390..d7dfc536c 100644 --- a/crates/e2e_test/src/reliant/lock.rs +++ b/crates/e2e_test/src/reliant/lock.rs @@ -14,8 +14,10 @@ // limitations under the License. use super::{grpc_lock_client::GrpcLockClient, grpc_lock_server::spawn_lock_server}; -use rustfs_lock::client::local::LocalClient; -use rustfs_lock::{GlobalLockManager, LockError, LockInfo, LockResponse, LockStats, NamespaceLock, ObjectKey}; +use rustfs_lock::client::{LockClient, local::LocalClient}; +use rustfs_lock::{ + GlobalLockManager, LockError, LockInfo, LockRequest, LockResponse, LockStats, LockType, NamespaceLock, ObjectKey, +}; use std::sync::Arc; use std::time::Duration; @@ -223,6 +225,56 @@ async fn test_distributed_lock_2_nodes_grpc_read_survives_failed_node() { failing_handle.abort(); } +#[tokio::test] +async fn test_grpc_lock_client_batch_acquire_and_release() { + let manager = Arc::new(GlobalLockManager::new()); + let local_client: Arc = Arc::new(LocalClient::with_manager(manager)); + + let (addr, handle) = spawn_lock_server(local_client).await.expect("Failed to spawn server"); + tokio::time::sleep(Duration::from_millis(100)).await; + + let grpc_client = GrpcLockClient::new(addr); + let requests = vec![ + LockRequest::new(test_resource(), LockType::Exclusive, "owner-a").with_acquire_timeout(Duration::from_secs(2)), + LockRequest::new( + ObjectKey { + bucket: Arc::from("test-bucket"), + object: Arc::from("test-object-2"), + version: None, + }, + LockType::Exclusive, + "owner-a", + ) + .with_acquire_timeout(Duration::from_secs(2)), + ]; + + let responses = grpc_client + .acquire_locks_batch(&requests) + .await + .expect("batch acquire should succeed"); + assert_eq!(responses.len(), requests.len()); + assert!(responses.iter().all(|response| response.success)); + + let lock_ids = responses + .iter() + .map(|response| { + response + .lock_info + .as_ref() + .expect("batch response should include lock info") + .id + .clone() + }) + .collect::>(); + let released = grpc_client + .release_locks_batch(&lock_ids) + .await + .expect("batch release should succeed"); + assert_eq!(released, vec![true, true]); + + handle.abort(); +} + #[tokio::test] async fn test_distributed_lock_4_nodes_grpc_read_write_quorum_split_with_two_failed_nodes() { let manager1 = Arc::new(GlobalLockManager::new()); diff --git a/crates/ecstore/src/rpc/remote_locker.rs b/crates/ecstore/src/rpc/remote_locker.rs index e3f34ed4d..4e786a4d2 100644 --- a/crates/ecstore/src/rpc/remote_locker.rs +++ b/crates/ecstore/src/rpc/remote_locker.rs @@ -19,7 +19,7 @@ use rustfs_lock::{ types::{LockId, LockMetadata, LockPriority}, }; use rustfs_protos::proto_gen::node_service::node_service_client::NodeServiceClient; -use rustfs_protos::proto_gen::node_service::{GenerallyLockRequest, PingRequest}; +use rustfs_protos::proto_gen::node_service::{BatchGenerallyLockRequest, GenerallyLockRequest, PingRequest}; use tonic::Request; use tonic::service::interceptor::InterceptedService; use tonic::transport::Channel; @@ -61,6 +61,44 @@ impl RemoteClient { .await .map_err(|err| LockError::internal(format!("can not get client, err: {err}"))) } + + fn build_lock_info(request: &LockRequest, lock_info_json: Option) -> LockInfo { + if let Some(lock_info_json) = lock_info_json { + match serde_json::from_str::(&lock_info_json) { + Ok(info) => info, + Err(e) => { + warn!("Failed to deserialize lock_info from response: {}, using request data", e); + LockInfo { + id: request.lock_id.clone(), + resource: request.resource.clone(), + lock_type: request.lock_type, + status: LockStatus::Acquired, + owner: request.owner.clone(), + acquired_at: std::time::SystemTime::now(), + expires_at: std::time::SystemTime::now() + request.ttl, + last_refreshed: std::time::SystemTime::now(), + metadata: request.metadata.clone(), + priority: request.priority, + wait_start_time: None, + } + } + } + } else { + LockInfo { + id: request.lock_id.clone(), + resource: request.resource.clone(), + lock_type: request.lock_type, + status: LockStatus::Acquired, + owner: request.owner.clone(), + acquired_at: std::time::SystemTime::now(), + expires_at: std::time::SystemTime::now() + request.ttl, + last_refreshed: std::time::SystemTime::now(), + metadata: request.metadata.clone(), + priority: request.priority, + wait_start_time: None, + } + } + } } #[async_trait] @@ -86,46 +124,10 @@ impl LockClient for RemoteClient { // Check if the lock acquisition was successful if resp.success { - // Try to deserialize lock_info from response - let lock_info = if let Some(lock_info_json) = resp.lock_info { - match serde_json::from_str::(&lock_info_json) { - Ok(info) => info, - Err(e) => { - // If deserialization fails, fall back to constructing from request - warn!("Failed to deserialize lock_info from response: {}, using request data", e); - LockInfo { - id: request.lock_id.clone(), - resource: request.resource.clone(), - lock_type: request.lock_type, - status: LockStatus::Acquired, - owner: request.owner.clone(), - acquired_at: std::time::SystemTime::now(), - expires_at: std::time::SystemTime::now() + request.ttl, - last_refreshed: std::time::SystemTime::now(), - metadata: request.metadata.clone(), - priority: request.priority, - wait_start_time: None, - } - } - } - } else { - // If lock_info is not provided, construct from request - LockInfo { - id: request.lock_id.clone(), - resource: request.resource.clone(), - lock_type: request.lock_type, - status: LockStatus::Acquired, - owner: request.owner.clone(), - acquired_at: std::time::SystemTime::now(), - expires_at: std::time::SystemTime::now() + request.ttl, - last_refreshed: std::time::SystemTime::now(), - metadata: request.metadata.clone(), - priority: request.priority, - wait_start_time: None, - } - }; - - Ok(LockResponse::success(lock_info, std::time::Duration::ZERO)) + Ok(LockResponse::success( + Self::build_lock_info(request, resp.lock_info), + std::time::Duration::ZERO, + )) } else { // Lock acquisition failed Ok(LockResponse::failure( @@ -135,6 +137,45 @@ impl LockClient for RemoteClient { } } + async fn acquire_locks_batch(&self, requests: &[LockRequest]) -> Result> { + let mut client = self.get_client().await?; + let req = Request::new(BatchGenerallyLockRequest { + args: requests + .iter() + .map(|request| { + serde_json::to_string(request).map_err(|e| LockError::internal(format!("Failed to serialize request: {e}"))) + }) + .collect::>>()?, + }); + + let resp = client + .lock_batch(req) + .await + .map_err(|e| LockError::internal(e.to_string()))? + .into_inner(); + + Ok(requests + .iter() + .enumerate() + .map(|(idx, request)| match resp.results.get(idx) { + Some(result) if result.success => { + LockResponse::success(Self::build_lock_info(request, result.lock_info.clone()), std::time::Duration::ZERO) + } + Some(result) => LockResponse::failure( + result + .error_info + .clone() + .unwrap_or_else(|| "Lock acquisition failed on remote server".to_string()), + std::time::Duration::ZERO, + ), + None => LockResponse::failure( + format!("Lock batch response missing entry for request index {idx}"), + std::time::Duration::ZERO, + ), + }) + .collect()) + } + async fn release(&self, lock_id: &LockId) -> Result { info!("remote release for {}", lock_id); @@ -154,6 +195,31 @@ impl LockClient for RemoteClient { Ok(resp.success) } + async fn release_locks_batch(&self, lock_ids: &[LockId]) -> Result> { + let mut client = self.get_client().await?; + let req = Request::new(BatchGenerallyLockRequest { + args: lock_ids + .iter() + .map(|lock_id| { + serde_json::to_string(&Self::create_unlock_request(lock_id)) + .map_err(|e| LockError::internal(format!("Failed to serialize request: {e}"))) + }) + .collect::>>()?, + }); + + let resp = client + .un_lock_batch(req) + .await + .map_err(|e| LockError::internal(e.to_string()))? + .into_inner(); + + Ok(lock_ids + .iter() + .enumerate() + .map(|(idx, _)| resp.results.get(idx).map(|result| result.success).unwrap_or(false)) + .collect()) + } + async fn refresh(&self, lock_id: &LockId) -> Result { info!("remote refresh for {}", lock_id); let refresh_request = Self::create_unlock_request(lock_id); diff --git a/crates/ecstore/src/set_disk.rs b/crates/ecstore/src/set_disk.rs index f76fad893..7a150245b 100644 --- a/crates/ecstore/src/set_disk.rs +++ b/crates/ecstore/src/set_disk.rs @@ -76,7 +76,7 @@ use rustfs_filemeta::{ use rustfs_lock::LockClient; use rustfs_lock::fast_lock::types::LockResult; use rustfs_lock::local_lock::LocalLock; -use rustfs_lock::{FastLockGuard, NamespaceLock, NamespaceLockGuard, NamespaceLockWrapper, ObjectKey}; +use rustfs_lock::{FastLockGuard, LockManager, NamespaceLock, NamespaceLockGuard, NamespaceLockWrapper, ObjectKey}; use rustfs_madmin::heal_commands::{HealDriveInfo, HealResultItem}; use rustfs_rio::{EtagResolvable, HashReader, HashReaderMut, TryGetIndex as _}; use rustfs_s3_common::EventName; @@ -964,6 +964,132 @@ impl ObjectIO for SetDisks { } } +impl SetDisks { + async fn acquire_dist_delete_object_locks_batch( + &self, + batch: &rustfs_lock::BatchLockRequest, + ) -> (HashMap<(String, String), String>, HashSet, Vec>) { + let requests: Vec = batch + .requests + .iter() + .map(|req| { + rustfs_lock::LockRequest::new(req.key.clone(), rustfs_lock::LockType::Exclusive, self.locker_owner.clone()) + .with_acquire_timeout(get_lock_acquire_timeout()) + .with_ttl(rustfs_lock::fast_lock::DEFAULT_LOCK_TIMEOUT) + }) + .collect(); + + let write_quorum = if self.lockers.len() > 1 { + (self.lockers.len() / 2) + 1 + } else { + 1 + }; + + let client_results = join_all(self.lockers.iter().cloned().enumerate().map(|(client_idx, client)| { + let requests = requests.clone(); + async move { (client_idx, client.acquire_locks_batch(&requests).await) } + })) + .await; + + let mut lock_ids_by_object: Vec> = vec![Vec::new(); requests.len()]; + let mut errors_by_object: Vec> = vec![None; requests.len()]; + + for (client_idx, result) in client_results { + match result { + Ok(responses) => { + for (req_idx, request) in requests.iter().enumerate() { + match responses.get(req_idx) { + Some(response) if response.success => { + if let Some(lock_info) = response.lock_info.as_ref() { + lock_ids_by_object[req_idx].push((client_idx, lock_info.id.clone())); + } else if errors_by_object[req_idx].is_none() { + errors_by_object[req_idx] = Some(format!( + "missing distributed lock id for {}/{}", + request.resource.bucket, request.resource.object + )); + } + } + Some(response) => { + if errors_by_object[req_idx].is_none() { + errors_by_object[req_idx] = Some( + response + .error + .clone() + .unwrap_or_else(|| "distributed lock acquisition failed".to_string()), + ); + } + } + None => { + if errors_by_object[req_idx].is_none() { + errors_by_object[req_idx] = + Some(format!("client {client_idx} returned incomplete batch lock response")); + } + } + } + } + } + Err(err) => { + for error in errors_by_object.iter_mut().take(requests.len()) { + if error.is_none() { + *error = Some(format!("client {client_idx} batch lock request failed: {err}")); + } + } + } + } + } + + let mut failed_map = HashMap::new(); + let mut locked_objects = HashSet::new(); + let mut held_lock_ids_by_client = vec![Vec::new(); self.lockers.len()]; + let mut rollback_lock_ids_by_client = vec![Vec::new(); self.lockers.len()]; + + for (req_idx, req) in batch.requests.iter().enumerate() { + let success_count = lock_ids_by_object[req_idx].len(); + if success_count >= write_quorum { + for (client_idx, lock_id) in lock_ids_by_object[req_idx].drain(..) { + held_lock_ids_by_client[client_idx].push(lock_id); + } + locked_objects.insert(req.key.object.as_ref().to_string()); + } else { + for (client_idx, lock_id) in lock_ids_by_object[req_idx].drain(..) { + rollback_lock_ids_by_client[client_idx].push(lock_id); + } + failed_map.insert( + (req.key.bucket.as_ref().to_string(), req.key.object.as_ref().to_string()), + errors_by_object[req_idx].clone().unwrap_or_else(|| { + format!("failed to acquire distributed delete lock quorum: {success_count}/{write_quorum}") + }), + ); + } + } + + self.release_dist_delete_object_locks_batch(rollback_lock_ids_by_client).await; + + (failed_map, locked_objects, held_lock_ids_by_client) + } + + async fn release_dist_delete_object_locks_batch(&self, lock_ids_by_client: Vec>) { + join_all(self.lockers.iter().cloned().enumerate().filter_map(|(client_idx, client)| { + let lock_ids = lock_ids_by_client.get(client_idx).cloned().unwrap_or_default(); + if lock_ids.is_empty() { + None + } else { + Some(async move { + if let Err(err) = client.release_locks_batch(&lock_ids).await { + tracing::warn!( + client_idx, + lock_count = lock_ids.len(), + "failed to release distributed delete locks in batch: {}", + err + ); + } + }) + } + })) + .await; + } +} + #[async_trait::async_trait] impl StorageAPI for SetDisks { #[tracing::instrument(skip(self))] @@ -1242,27 +1368,25 @@ impl ObjectOperations for SetDisks { } let mut failed_map = HashMap::new(); - let mut batch_guards = Vec::with_capacity(batch.requests.len()); - + let mut _local_batch_guards: Vec = Vec::with_capacity(batch.requests.len()); let mut locked_objects = HashSet::new(); - for req in batch.requests.iter() { - let ns_lock = match self.new_ns_lock(req.key.bucket.as_ref(), req.key.object.as_ref()).await { - Ok(ns_lock) => ns_lock, - Err(e) => { - failed_map.insert((req.key.bucket.as_ref().to_string(), req.key.object.as_ref().to_string()), e.to_string()); - continue; - } - }; - let _lock_guard = match ns_lock.get_write_lock(get_lock_acquire_timeout()).await { - Ok(lock_guard) => lock_guard, - Err(e) => { - failed_map.insert((req.key.bucket.as_ref().to_string(), req.key.object.as_ref().to_string()), e.to_string()); - continue; - } - }; - batch_guards.push(_lock_guard); - locked_objects.insert(req.key.object.as_ref().to_string()); + let dist_erasure = is_dist_erasure().await; + let mut dist_batch_lock_ids = vec![Vec::new(); self.lockers.len()]; + + if dist_erasure { + (failed_map, locked_objects, dist_batch_lock_ids) = self.acquire_dist_delete_object_locks_batch(&batch).await; + } else { + let batch_result = self.local_lock_manager.acquire_locks_batch(batch).await; + _local_batch_guards = batch_result.guards; + + for key in batch_result.successful_locks { + locked_objects.insert(key.object.as_ref().to_string()); + } + + for (key, err) in batch_result.failed_locks { + failed_map.insert((key.bucket.as_ref().to_string(), key.object.as_ref().to_string()), format!("{err:?}")); + } } // Mark failures for objects that could not be locked @@ -1428,6 +1552,10 @@ impl ObjectOperations for SetDisks { // TODO: add_partial + if dist_erasure { + self.release_dist_delete_object_locks_batch(dist_batch_lock_ids).await; + } + (del_objects, del_errs) } @@ -4267,6 +4395,81 @@ mod tests { ); } + #[tokio::test(flavor = "multi_thread")] + #[serial] + async fn test_acquire_dist_delete_object_locks_batch_succeeds_with_two_healthy_lockers() { + let _setup_type_guard = SetupTypeGuard::switch_to(SetupType::DistErasure).await; + + let manager1 = Arc::new(rustfs_lock::GlobalLockManager::new()); + let manager2 = Arc::new(rustfs_lock::GlobalLockManager::new()); + let client1: Arc = Arc::new(LocalClient::with_manager(manager1.clone())); + let client2: Arc = Arc::new(LocalClient::with_manager(manager2.clone())); + let set_disks = make_test_set_disks(vec![client1, client2]).await; + + let batch = rustfs_lock::BatchLockRequest::new(set_disks.locker_owner.as_str()) + .with_all_or_nothing(false) + .add_write_lock(ObjectKey::new("bucket", "object-a")) + .add_write_lock(ObjectKey::new("bucket", "object-b")); + + let (failed_map, locked_objects, held_lock_ids_by_client) = + set_disks.acquire_dist_delete_object_locks_batch(&batch).await; + + assert!(failed_map.is_empty()); + assert_eq!(locked_objects.len(), 2); + assert!(locked_objects.contains("object-a")); + assert!(locked_objects.contains("object-b")); + assert_eq!(held_lock_ids_by_client.iter().map(Vec::len).sum::(), batch.requests.len() * 2); + + set_disks + .release_dist_delete_object_locks_batch(held_lock_ids_by_client) + .await; + + let local_lock_1 = NamespaceLock::with_local_manager("node-1".to_string(), manager1); + let local_lock_2 = NamespaceLock::with_local_manager("node-2".to_string(), manager2); + + let guard_1 = local_lock_1 + .get_write_lock(ObjectKey::new("bucket", "object-a"), "owner-b", Duration::from_millis(100)) + .await + .expect("released batch lock should free node 1"); + let guard_2 = local_lock_2 + .get_write_lock(ObjectKey::new("bucket", "object-b"), "owner-b", Duration::from_millis(100)) + .await + .expect("released batch lock should free node 2"); + + drop(guard_1); + drop(guard_2); + } + + #[tokio::test(flavor = "multi_thread")] + #[serial] + async fn test_acquire_dist_delete_object_locks_batch_rolls_back_when_quorum_not_reached() { + let _setup_type_guard = SetupTypeGuard::switch_to(SetupType::DistErasure).await; + + let manager = Arc::new(rustfs_lock::GlobalLockManager::new()); + let healthy_client: Arc = Arc::new(LocalClient::with_manager(manager.clone())); + let failing_client: Arc = Arc::new(FailingClient); + let set_disks = make_test_set_disks(vec![healthy_client, failing_client]).await; + + let batch = rustfs_lock::BatchLockRequest::new(set_disks.locker_owner.as_str()) + .with_all_or_nothing(false) + .add_write_lock(ObjectKey::new("bucket", "object-a")); + + let (failed_map, locked_objects, held_lock_ids_by_client) = + set_disks.acquire_dist_delete_object_locks_batch(&batch).await; + + assert!(locked_objects.is_empty()); + assert!(failed_map.contains_key(&("bucket".to_string(), "object-a".to_string()))); + assert_eq!(held_lock_ids_by_client.iter().map(Vec::len).sum::(), 0); + + let local_lock = NamespaceLock::with_local_manager("node-1".to_string(), manager); + let guard = local_lock + .get_write_lock(ObjectKey::new("bucket", "object-a"), "owner-b", Duration::from_millis(100)) + .await + .expect("quorum rollback should release the healthy node lock"); + + drop(guard); + } + #[test] fn test_common_parity() { // Test common parity calculation diff --git a/crates/ecstore/src/sets.rs b/crates/ecstore/src/sets.rs index d623d6cd4..48502a4b6 100644 --- a/crates/ecstore/src/sets.rs +++ b/crates/ecstore/src/sets.rs @@ -35,7 +35,10 @@ use crate::{ }, store_init::{check_format_erasure_values, get_format_erasure_in_quorum, load_format_erasure_all, save_format_file}, }; -use futures::future::join_all; +use futures::{ + future::join_all, + stream::{FuturesUnordered, StreamExt}, +}; use http::HeaderMap; use rustfs_common::heal_channel::HealOpts; use rustfs_common::{ @@ -336,6 +339,26 @@ struct DelObj { obj: ObjectToDelete, } +fn apply_delete_objects_results( + del_objects: &mut [DeletedObject], + del_errs: &mut [Option], + set_objects: &[DelObj], + dobjects: &[DeletedObject], + errs: Vec>, +) { + for (i, err) in errs.into_iter().enumerate() { + let obj = set_objects + .get(i) + .expect("delete_objects should return errors aligned with input objects"); + + del_errs[obj.orig_idx] = err; + del_objects[obj.orig_idx] = dobjects + .get(i) + .expect("delete_objects should return objects aligned with input objects") + .clone(); + } +} + #[async_trait::async_trait] impl ObjectIO for Sets { #[tracing::instrument(level = "debug", skip(self, object, h, opts))] @@ -508,19 +531,30 @@ impl ObjectOperations for Sets { } } - // TODO: concurrency + let max_concurrent = set_obj_map.len().min(num_cpus::get()).max(1); + let semaphore = Arc::new(tokio::sync::Semaphore::new(max_concurrent)); + let mut futures = FuturesUnordered::new(); + let bucket = bucket.to_string(); + for (k, v) in set_obj_map { let disks = self.get_disks(k); let objs: Vec = v.iter().map(|v| v.obj.clone()).collect(); - let (dobjects, errs) = disks.delete_objects(bucket, objs, opts.clone()).await; + let bucket = bucket.clone(); + let opts = opts.clone(); + let semaphore = semaphore.clone(); - for (i, err) in errs.into_iter().enumerate() { - let obj = v.get(i).unwrap(); + futures.push(async move { + let _permit = semaphore + .acquire_owned() + .await + .expect("delete_objects semaphore should remain open"); + let (dobjects, errs) = disks.delete_objects(&bucket, objs, opts).await; + (v, dobjects, errs) + }); + } - del_errs[obj.orig_idx] = err; - - del_objects[obj.orig_idx] = dobjects.get(i).unwrap().clone(); - } + while let Some((v, dobjects, errs)) = futures.next().await { + apply_delete_objects_results(&mut del_objects, &mut del_errs, &v, &dobjects, errs); } (del_objects, del_errs) @@ -1015,3 +1049,76 @@ fn new_heal_format_sets( (new_formats, current_disks_info) } + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn test_apply_delete_objects_results_preserves_original_order_for_out_of_order_batches() { + let mut del_objects = vec![DeletedObject::default(); 3]; + let mut del_errs = vec![None, None, None]; + + let early_batch = vec![DelObj { + orig_idx: 1, + obj: ObjectToDelete { + object_name: "second".to_string(), + ..Default::default() + }, + }]; + let early_objects = vec![DeletedObject { + object_name: "second".to_string(), + found: true, + ..Default::default() + }]; + + let late_batch = vec![ + DelObj { + orig_idx: 2, + obj: ObjectToDelete { + object_name: "third".to_string(), + ..Default::default() + }, + }, + DelObj { + orig_idx: 0, + obj: ObjectToDelete { + object_name: "first".to_string(), + ..Default::default() + }, + }, + ]; + let late_objects = vec![ + DeletedObject { + object_name: "third".to_string(), + found: true, + ..Default::default() + }, + DeletedObject { + object_name: "first".to_string(), + found: true, + ..Default::default() + }, + ]; + + apply_delete_objects_results(&mut del_objects, &mut del_errs, &early_batch, &early_objects, vec![None]); + apply_delete_objects_results( + &mut del_objects, + &mut del_errs, + &late_batch, + &late_objects, + vec![Some(Error::other("third failed")), None], + ); + + assert_eq!(del_objects[0].object_name, "first"); + assert_eq!(del_objects[1].object_name, "second"); + assert_eq!(del_objects[2].object_name, "third"); + + assert!(del_errs[0].is_none()); + assert!(del_errs[1].is_none()); + assert_eq!( + del_errs[2].as_ref().map(ToString::to_string), + Some(Error::other("third failed").to_string()) + ); + } +} diff --git a/crates/lock/src/client/mod.rs b/crates/lock/src/client/mod.rs index ee3f4786d..d16db4bb6 100644 --- a/crates/lock/src/client/mod.rs +++ b/crates/lock/src/client/mod.rs @@ -17,6 +17,7 @@ pub mod local; use crate::{LockId, LockInfo, LockRequest, LockResponse, LockStats, Result}; use async_trait::async_trait; +use futures::future::join_all; use std::sync::Arc; /// Lock client trait @@ -25,9 +26,25 @@ pub trait LockClient: Send + Sync + std::fmt::Debug { /// Acquire lock (generic method) async fn acquire_lock(&self, request: &LockRequest) -> Result; + /// Acquire multiple locks. Default implementation fans out to single-lock requests. + async fn acquire_locks_batch(&self, requests: &[LockRequest]) -> Result> { + Ok(join_all(requests.iter().map(|request| self.acquire_lock(request))) + .await + .into_iter() + .collect::>>()?) + } + /// Release lock async fn release(&self, lock_id: &LockId) -> Result; + /// Release multiple locks. Default implementation fans out to single-lock releases. + async fn release_locks_batch(&self, lock_ids: &[LockId]) -> Result> { + Ok(join_all(lock_ids.iter().map(|lock_id| self.release(lock_id))) + .await + .into_iter() + .collect::>>()?) + } + /// Refresh lock async fn refresh(&self, lock_id: &LockId) -> Result; diff --git a/crates/lock/src/distributed_lock.rs b/crates/lock/src/distributed_lock.rs index f461fffff..9928daffe 100644 --- a/crates/lock/src/distributed_lock.rs +++ b/crates/lock/src/distributed_lock.rs @@ -18,6 +18,7 @@ use crate::{ error::{LockError, Result}, types::{LockId, LockInfo, LockRequest, LockResponse, LockStatus, LockType}, }; +use futures::future::join_all; use std::sync::{Arc, LazyLock}; use std::time::Duration; use tokio::sync::mpsc; @@ -52,12 +53,13 @@ static UNLOCK_RUNTIME: LazyLock = LazyLock::new(|| { tokio::spawn(async move { while let Some(job) = rx.recv().await { // Best-effort release across all (LockId, client) entries. - let mut any_ok = false; - for (lock_id, client) in job.entries.into_iter() { - if client.release(&lock_id).await.unwrap_or(false) { - any_ok = true; - } - } + let results = join_all( + job.entries + .into_iter() + .map(|(lock_id, client)| async move { client.release(&lock_id).await.unwrap_or(false) }), + ) + .await; + let any_ok = results.into_iter().any(|released| released); if !any_ok { tracing::warn!("DistributedLockGuard background release failed for one or more entries"); @@ -142,7 +144,7 @@ impl DistributedLockGuard { let futures_iter = entries .into_iter() .map(|(lock_id, client)| async move { client.release(&lock_id).await.unwrap_or(false) }); - let _ = futures::future::join_all(futures_iter).await; + let _ = join_all(futures_iter).await; }); // Explicitly drop the JoinHandle to acknowledge detaching the task. drop(handle); @@ -411,8 +413,13 @@ impl DistributedLock { } else { // Rollback: release all locks that were successfully acquired let rollback_count = individual_locks.len(); - for (individual_lock_id, client) in &individual_locks { - if let Err(e) = client.release(individual_lock_id).await { + let rollback_results = join_all(individual_locks.iter().map(|(individual_lock_id, client)| async move { + (individual_lock_id, client.release(individual_lock_id).await) + })) + .await; + + for (individual_lock_id, result) in rollback_results { + if let Err(e) = result { tracing::warn!("Failed to rollback lock {} on client: {}", individual_lock_id, e); } } diff --git a/crates/lock/src/fast_lock/manager.rs b/crates/lock/src/fast_lock/manager.rs index d8c5c9295..4eca21210 100644 --- a/crates/lock/src/fast_lock/manager.rs +++ b/crates/lock/src/fast_lock/manager.rs @@ -138,7 +138,7 @@ impl FastObjectLockManager { shard_a.cmp(&shard_b).then_with(|| a.key.cmp(&b.key)) }); - // Try to use stack-allocated vectors for small batches, fallback to heap if needed + // Preserve shard order so every concurrent batch acquires locks in the same global order. let shard_groups = self.group_requests_by_shard(sorted_requests); // Choose strategy based on request type @@ -150,31 +150,28 @@ impl FastObjectLockManager { } /// Group requests by shard with proper fallback handling - fn group_requests_by_shard( - &self, - requests: Vec, - ) -> std::collections::HashMap> { - let mut shard_groups = std::collections::HashMap::new(); + fn group_requests_by_shard(&self, requests: Vec) -> Vec<(usize, Vec)> { + let mut shard_groups: Vec<(usize, Vec)> = Vec::new(); for request in requests { let shard_id = request.key.shard_index(self.shard_mask); - shard_groups.entry(shard_id).or_insert_with(Vec::new).push(request); + match shard_groups.last_mut() { + Some((last_shard_id, grouped_requests)) if *last_shard_id == shard_id => grouped_requests.push(request), + _ => shard_groups.push((shard_id, vec![request])), + } } shard_groups } /// Best effort acquisition (allows partial success) - async fn acquire_locks_best_effort( - &self, - shard_groups: &std::collections::HashMap>, - ) -> BatchLockResult { + async fn acquire_locks_best_effort(&self, shard_groups: &[(usize, Vec)]) -> BatchLockResult { let mut all_successful = Vec::new(); let mut all_failed = Vec::new(); let mut guards = Vec::new(); - for (&shard_id, requests) in shard_groups { - let shard = self.shards[shard_id].clone(); + for (shard_id, requests) in shard_groups { + let shard = self.shards[*shard_id].clone(); for request in requests { let key = request.key.clone(); @@ -212,16 +209,13 @@ impl FastObjectLockManager { } /// Two-phase commit for atomic acquisition - async fn acquire_locks_two_phase_commit( - &self, - shard_groups: &std::collections::HashMap>, - ) -> BatchLockResult { + async fn acquire_locks_two_phase_commit(&self, shard_groups: &[(usize, Vec)]) -> BatchLockResult { // Phase 1: Try to acquire all locks let mut acquired_guards = Vec::new(); let mut failed_locks = Vec::new(); - 'outer: for (&shard_id, requests) in shard_groups { - let shard = self.shards[shard_id].clone(); + 'outer: for (shard_id, requests) in shard_groups { + let shard = self.shards[*shard_id].clone(); for request in requests { match shard.acquire_lock(request).await { @@ -438,3 +432,48 @@ impl LockManager for FastObjectLockManager { false } } + +#[cfg(test)] +mod tests { + use super::*; + + fn make_request(manager: &FastObjectLockManager, shard_id: usize, suffix: usize) -> ObjectLockRequest { + let mut candidate = 0usize; + loop { + let object = format!("object-{shard_id}-{suffix}-{candidate}"); + let key = ObjectKey::new("bucket", object); + if key.shard_index(manager.shard_mask) == shard_id { + return ObjectLockRequest::new_write(key, "owner"); + } + candidate += 1; + } + } + + #[tokio::test] + async fn test_group_requests_by_shard_preserves_sorted_shard_order() { + let manager = FastObjectLockManager::new(); + let mut requests = vec![ + make_request(&manager, 3, 0), + make_request(&manager, 1, 0), + make_request(&manager, 2, 0), + make_request(&manager, 1, 1), + make_request(&manager, 3, 1), + ]; + + requests.sort_unstable_by(|a, b| { + let shard_a = a.key.shard_index(manager.shard_mask); + let shard_b = b.key.shard_index(manager.shard_mask); + shard_a.cmp(&shard_b).then_with(|| a.key.cmp(&b.key)) + }); + + let shard_groups = manager.group_requests_by_shard(requests); + let shard_ids: Vec<_> = shard_groups.iter().map(|(shard_id, _)| *shard_id).collect(); + + assert_eq!(shard_ids, vec![1, 2, 3]); + assert_eq!(shard_groups[0].1.len(), 2); + assert_eq!(shard_groups[1].1.len(), 1); + assert_eq!(shard_groups[2].1.len(), 2); + + manager.shutdown().await; + } +} diff --git a/crates/lock/src/namespace/tests.rs b/crates/lock/src/namespace/tests.rs index f8b48c882..1b11caa79 100644 --- a/crates/lock/src/namespace/tests.rs +++ b/crates/lock/src/namespace/tests.rs @@ -97,6 +97,37 @@ async fn test_namespace_lock_with_clients() { assert_eq!(lock.namespace(), "multi-client"); } +#[tokio::test] +async fn test_lock_client_default_batch_acquire_and_release() { + let manager = Arc::new(GlobalLockManager::new()); + let client = LocalClient::with_manager(manager); + let requests = vec![ + LockRequest::new(create_test_object_key("bucket", "object-a"), LockType::Exclusive, "owner-a") + .with_acquire_timeout(Duration::from_secs(1)), + LockRequest::new(create_test_object_key("bucket", "object-b"), LockType::Exclusive, "owner-a") + .with_acquire_timeout(Duration::from_secs(1)), + ]; + + let responses = client.acquire_locks_batch(&requests).await.unwrap(); + assert_eq!(responses.len(), requests.len()); + assert!(responses.iter().all(|response| response.success)); + + let lock_ids = responses + .iter() + .map(|response| { + response + .lock_info + .as_ref() + .expect("successful batch acquire should return lock info") + .id + .clone() + }) + .collect::>(); + let released = client.release_locks_batch(&lock_ids).await.unwrap(); + + assert_eq!(released, vec![true, true]); +} + #[tokio::test] async fn test_namespace_lock_get_resource_key() { let client = ClientFactory::create_local(); @@ -452,6 +483,45 @@ async fn test_namespace_lock_distributed_write_lock_fails_with_two_nodes_one_off ); } +#[tokio::test] +async fn test_namespace_lock_distributed_quorum_failure_rolls_back_successful_nodes() { + let manager1 = Arc::new(GlobalLockManager::new()); + let manager2 = Arc::new(GlobalLockManager::new()); + + let client1: Arc = Arc::new(LocalClient::with_manager(manager1.clone())); + let client2: Arc = Arc::new(LocalClient::with_manager(manager2.clone())); + let client3: Arc = Arc::new(FailingClient); + + let resource = create_test_object_key("bucket", "object"); + + let distributed_lock = NamespaceLock::with_clients_and_quorum("three-node".to_string(), vec![client1, client2, client3], 3); + let err = distributed_lock + .get_write_lock(resource.clone(), "owner-a", Duration::from_millis(100)) + .await + .expect_err("write lock should fail when quorum requires all three nodes"); + + let err_str = err.to_string().to_lowercase(); + assert!( + err_str.contains("quorum") || err_str.contains("not reached"), + "expected quorum error, got: {err}" + ); + + let local_lock_1 = NamespaceLock::with_local_manager("node-1".to_string(), manager1); + let local_lock_2 = NamespaceLock::with_local_manager("node-2".to_string(), manager2); + + let guard1 = local_lock_1 + .get_write_lock(resource.clone(), "owner-b", Duration::from_millis(100)) + .await + .expect("quorum rollback should release node 1"); + let guard2 = local_lock_2 + .get_write_lock(resource, "owner-b", Duration::from_millis(100)) + .await + .expect("quorum rollback should release node 2"); + + drop(guard1); + drop(guard2); +} + #[tokio::test] async fn test_namespace_lock_distributed_even_node_read_write_quorum_split() { let manager1 = Arc::new(GlobalLockManager::new()); diff --git a/crates/protos/src/generated/proto_gen/node_service.rs b/crates/protos/src/generated/proto_gen/node_service.rs index 0efecb167..b71cc1393 100644 --- a/crates/protos/src/generated/proto_gen/node_service.rs +++ b/crates/protos/src/generated/proto_gen/node_service.rs @@ -658,6 +658,26 @@ pub struct GenerallyLockResponse { #[prost(string, optional, tag = "3")] pub lock_info: ::core::option::Option<::prost::alloc::string::String>, } +#[derive(Clone, PartialEq, Eq, Hash, ::prost::Message)] +pub struct BatchGenerallyLockRequest { + #[prost(string, repeated, tag = "1")] + pub args: ::prost::alloc::vec::Vec<::prost::alloc::string::String>, +} +#[derive(Clone, PartialEq, Eq, Hash, ::prost::Message)] +pub struct GenerallyLockResult { + #[prost(bool, tag = "1")] + pub success: bool, + #[prost(string, optional, tag = "2")] + pub error_info: ::core::option::Option<::prost::alloc::string::String>, + /// JSON serialized LockInfo + #[prost(string, optional, tag = "3")] + pub lock_info: ::core::option::Option<::prost::alloc::string::String>, +} +#[derive(Clone, PartialEq, Eq, Hash, ::prost::Message)] +pub struct BatchGenerallyLockResponse { + #[prost(message, repeated, tag = "1")] + pub results: ::prost::alloc::vec::Vec, +} #[derive(Clone, PartialEq, ::prost::Message)] pub struct Mss { #[prost(map = "string, string", tag = "1")] @@ -1776,6 +1796,36 @@ pub mod node_service_client { .insert(GrpcMethod::new("node_service.NodeService", "Refresh")); self.inner.unary(req, path, codec).await } + pub async fn lock_batch( + &mut self, + request: impl tonic::IntoRequest, + ) -> std::result::Result, tonic::Status> { + self.inner + .ready() + .await + .map_err(|e| tonic::Status::unknown(format!("Service was not ready: {}", e.into())))?; + let codec = tonic_prost::ProstCodec::default(); + let path = http::uri::PathAndQuery::from_static("/node_service.NodeService/LockBatch"); + let mut req = request.into_request(); + req.extensions_mut() + .insert(GrpcMethod::new("node_service.NodeService", "LockBatch")); + self.inner.unary(req, path, codec).await + } + pub async fn un_lock_batch( + &mut self, + request: impl tonic::IntoRequest, + ) -> std::result::Result, tonic::Status> { + self.inner + .ready() + .await + .map_err(|e| tonic::Status::unknown(format!("Service was not ready: {}", e.into())))?; + let codec = tonic_prost::ProstCodec::default(); + let path = http::uri::PathAndQuery::from_static("/node_service.NodeService/UnLockBatch"); + let mut req = request.into_request(); + req.extensions_mut() + .insert(GrpcMethod::new("node_service.NodeService", "UnLockBatch")); + self.inner.unary(req, path, codec).await + } pub async fn local_storage_info( &mut self, request: impl tonic::IntoRequest, @@ -2512,6 +2562,14 @@ pub mod node_service_server { &self, request: tonic::Request, ) -> std::result::Result, tonic::Status>; + async fn lock_batch( + &self, + request: tonic::Request, + ) -> std::result::Result, tonic::Status>; + async fn un_lock_batch( + &self, + request: tonic::Request, + ) -> std::result::Result, tonic::Status>; async fn local_storage_info( &self, request: tonic::Request, @@ -3828,6 +3886,62 @@ pub mod node_service_server { }; Box::pin(fut) } + "/node_service.NodeService/LockBatch" => { + #[allow(non_camel_case_types)] + struct LockBatchSvc(pub Arc); + impl tonic::server::UnaryService for LockBatchSvc { + type Response = super::BatchGenerallyLockResponse; + type Future = BoxFuture, tonic::Status>; + fn call(&mut self, request: tonic::Request) -> Self::Future { + let inner = Arc::clone(&self.0); + let fut = async move { ::lock_batch(&inner, request).await }; + Box::pin(fut) + } + } + let accept_compression_encodings = self.accept_compression_encodings; + let send_compression_encodings = self.send_compression_encodings; + let max_decoding_message_size = self.max_decoding_message_size; + let max_encoding_message_size = self.max_encoding_message_size; + let inner = self.inner.clone(); + let fut = async move { + let method = LockBatchSvc(inner); + let codec = tonic_prost::ProstCodec::default(); + let mut grpc = tonic::server::Grpc::new(codec) + .apply_compression_config(accept_compression_encodings, send_compression_encodings) + .apply_max_message_size_config(max_decoding_message_size, max_encoding_message_size); + let res = grpc.unary(method, req).await; + Ok(res) + }; + Box::pin(fut) + } + "/node_service.NodeService/UnLockBatch" => { + #[allow(non_camel_case_types)] + struct UnLockBatchSvc(pub Arc); + impl tonic::server::UnaryService for UnLockBatchSvc { + type Response = super::BatchGenerallyLockResponse; + type Future = BoxFuture, tonic::Status>; + fn call(&mut self, request: tonic::Request) -> Self::Future { + let inner = Arc::clone(&self.0); + let fut = async move { ::un_lock_batch(&inner, request).await }; + Box::pin(fut) + } + } + let accept_compression_encodings = self.accept_compression_encodings; + let send_compression_encodings = self.send_compression_encodings; + let max_decoding_message_size = self.max_decoding_message_size; + let max_encoding_message_size = self.max_encoding_message_size; + let inner = self.inner.clone(); + let fut = async move { + let method = UnLockBatchSvc(inner); + let codec = tonic_prost::ProstCodec::default(); + let mut grpc = tonic::server::Grpc::new(codec) + .apply_compression_config(accept_compression_encodings, send_compression_encodings) + .apply_max_message_size_config(max_decoding_message_size, max_encoding_message_size); + let res = grpc.unary(method, req).await; + Ok(res) + }; + Box::pin(fut) + } "/node_service.NodeService/LocalStorageInfo" => { #[allow(non_camel_case_types)] struct LocalStorageInfoSvc(pub Arc); diff --git a/crates/protos/src/node.proto b/crates/protos/src/node.proto index c3a20fd62..1c7111ceb 100644 --- a/crates/protos/src/node.proto +++ b/crates/protos/src/node.proto @@ -456,6 +456,20 @@ message GenerallyLockResponse { optional string lock_info = 3; // JSON serialized LockInfo } +message BatchGenerallyLockRequest { + repeated string args = 1; +} + +message GenerallyLockResult { + bool success = 1; + optional string error_info = 2; + optional string lock_info = 3; // JSON serialized LockInfo +} + +message BatchGenerallyLockResponse { + repeated GenerallyLockResult results = 1; +} + message Mss { map value = 1; } @@ -837,6 +851,8 @@ service NodeService { rpc UnLock(GenerallyLockRequest) returns (GenerallyLockResponse) {}; rpc ForceUnLock(GenerallyLockRequest) returns (GenerallyLockResponse) {}; rpc Refresh(GenerallyLockRequest) returns (GenerallyLockResponse) {}; + rpc LockBatch(BatchGenerallyLockRequest) returns (BatchGenerallyLockResponse) {}; + rpc UnLockBatch(BatchGenerallyLockRequest) returns (BatchGenerallyLockResponse) {}; /* -------------------------------peer rest service-------------------------- */ diff --git a/rustfs/src/storage/rpc/lock.rs b/rustfs/src/storage/rpc/lock.rs index 419457476..01e547190 100644 --- a/rustfs/src/storage/rpc/lock.rs +++ b/rustfs/src/storage/rpc/lock.rs @@ -14,6 +14,22 @@ use super::*; +fn lock_result_from_response(response: rustfs_lock::LockResponse) -> GenerallyLockResult { + GenerallyLockResult { + success: response.success, + error_info: response.error, + lock_info: response.lock_info.and_then(|info| serde_json::to_string(&info).ok()), + } +} + +fn lock_result_from_error(error: impl Into) -> GenerallyLockResult { + GenerallyLockResult { + success: false, + error_info: Some(error.into()), + lock_info: None, + } +} + impl NodeService { pub(super) async fn handle_refresh( &self, @@ -144,4 +160,92 @@ impl NodeService { })), } } + + pub(super) async fn handle_lock_batch( + &self, + request: Request, + ) -> Result, Status> { + let request = request.into_inner(); + let mut results = vec![lock_result_from_error("request was not processed"); request.args.len()]; + let mut valid_requests = Vec::with_capacity(request.args.len()); + let mut valid_indices = Vec::with_capacity(request.args.len()); + + for (idx, arg) in request.args.iter().enumerate() { + match serde_json::from_str::(arg) { + Ok(args) => { + valid_requests.push(args); + valid_indices.push(idx); + } + Err(err) => { + results[idx] = lock_result_from_error(format!("can not decode args, err: {err}")); + } + } + } + + if !valid_requests.is_empty() { + let lock_client = self.get_lock_client()?; + match lock_client.acquire_locks_batch(&valid_requests).await { + Ok(batch_results) => { + for (result_idx, response) in batch_results.into_iter().enumerate() { + if let Some(request_idx) = valid_indices.get(result_idx) { + results[*request_idx] = lock_result_from_response(response); + } + } + } + Err(err) => { + for request_idx in valid_indices { + results[request_idx] = lock_result_from_error(format!("can not batch lock, err: {err}")); + } + } + } + } + + Ok(Response::new(BatchGenerallyLockResponse { results })) + } + + pub(super) async fn handle_un_lock_batch( + &self, + request: Request, + ) -> Result, Status> { + let request = request.into_inner(); + let mut results = vec![lock_result_from_error("request was not processed"); request.args.len()]; + let mut lock_ids = Vec::with_capacity(request.args.len()); + let mut valid_indices = Vec::with_capacity(request.args.len()); + + for (idx, arg) in request.args.iter().enumerate() { + match serde_json::from_str::(arg) { + Ok(args) => { + lock_ids.push(args.lock_id); + valid_indices.push(idx); + } + Err(err) => { + results[idx] = lock_result_from_error(format!("can not decode args, err: {err}")); + } + } + } + + if !lock_ids.is_empty() { + let lock_client = self.get_lock_client()?; + match lock_client.release_locks_batch(&lock_ids).await { + Ok(batch_results) => { + for (result_idx, success) in batch_results.into_iter().enumerate() { + if let Some(request_idx) = valid_indices.get(result_idx) { + results[*request_idx] = GenerallyLockResult { + success, + error_info: None, + lock_info: None, + }; + } + } + } + Err(err) => { + for request_idx in valid_indices { + results[request_idx] = lock_result_from_error(format!("can not batch unlock, err: {err}")); + } + } + } + } + + Ok(Response::new(BatchGenerallyLockResponse { results })) + } } diff --git a/rustfs/src/storage/rpc/node_service.rs b/rustfs/src/storage/rpc/node_service.rs index c0989a9e1..a4bc7a814 100644 --- a/rustfs/src/storage/rpc/node_service.rs +++ b/rustfs/src/storage/rpc/node_service.rs @@ -386,6 +386,20 @@ impl Node for NodeService { self.handle_refresh(request).await } + async fn lock_batch( + &self, + request: Request, + ) -> Result, Status> { + self.handle_lock_batch(request).await + } + + async fn un_lock_batch( + &self, + request: Request, + ) -> Result, Status> { + self.handle_un_lock_batch(request).await + } + async fn local_storage_info( &self, _request: Request, From 0ac3b5b992456f55dc887b71c408d0daf0f0d1da Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E5=AE=89=E6=AD=A3=E8=B6=85?= Date: Fri, 3 Apr 2026 22:51:29 +0800 Subject: [PATCH 10/22] docs: remind agents to clean build artifacts (#2387) --- AGENTS.md | 1 + 1 file changed, 1 insertion(+) diff --git a/AGENTS.md b/AGENTS.md index eb0310381..7b52de3e1 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -52,6 +52,7 @@ make pre-commit If `make` is unavailable, run the equivalent checks defined under `.config/make/`. Documentation-only or instruction-only changes are exempt from the verification commands above (including the `.config/make/` equivalents), though any installed git pre-commit hooks (for example, from `make setup-hooks`) may still run on commit unless explicitly skipped. +After build-based verification completes, clean generated build artifacts before wrapping up to avoid unnecessary disk usage. Do not open a PR with code changes when the required checks fail. ## Git and PR Baseline From 372f004a8da5ad060fcec9c323ddca4adf840bc8 Mon Sep 17 00:00:00 2001 From: houseme Date: Sat, 4 Apr 2026 06:37:24 +0800 Subject: [PATCH 11/22] refactor(server): unify TLS loading, optimize HTTP transport, add hot reload (#2388) --- crates/config/src/constants/tls.rs | 66 ++++ rustfs/Cargo.toml | 1 + rustfs/src/main.rs | 15 +- rustfs/src/server/cert.rs | 296 ----------------- rustfs/src/server/compress.rs | 196 +++++++++++ rustfs/src/server/http.rs | 399 ++++++++++++---------- rustfs/src/server/mod.rs | 3 +- rustfs/src/server/tls_material.rs | 510 +++++++++++++++++++++++++++++ 8 files changed, 1015 insertions(+), 471 deletions(-) delete mode 100644 rustfs/src/server/cert.rs create mode 100644 rustfs/src/server/tls_material.rs diff --git a/crates/config/src/constants/tls.rs b/crates/config/src/constants/tls.rs index 7772abff5..f6d012e2f 100644 --- a/crates/config/src/constants/tls.rs +++ b/crates/config/src/constants/tls.rs @@ -84,3 +84,69 @@ pub const ENV_SERVER_MTLS_ENABLE: &str = "RUSTFS_SERVER_MTLS_ENABLE"; /// By default, RustFS server mTLS is disabled. /// To change this behavior, set the environment variable RUSTFS_SERVER_MTLS_ENABLE=1 pub const DEFAULT_SERVER_MTLS_ENABLE: bool = false; + +// ── HTTP Transport Tuning Parameters ── + +/// Environment variable for HTTP/2 initial stream window size (bytes) +/// Default: 4194304 (4 MB) +pub const ENV_H2_INITIAL_STREAM_WINDOW_SIZE: &str = "RUSTFS_H2_INITIAL_STREAM_WINDOW_SIZE"; +pub const DEFAULT_H2_INITIAL_STREAM_WINDOW_SIZE: u32 = 4 * 1024 * 1024; // 4 MB + +/// Environment variable for HTTP/2 initial connection window size (bytes) +/// Default: 8388608 (8 MB) +pub const ENV_H2_INITIAL_CONN_WINDOW_SIZE: &str = "RUSTFS_H2_INITIAL_CONN_WINDOW_SIZE"; +pub const DEFAULT_H2_INITIAL_CONN_WINDOW_SIZE: u32 = 8 * 1024 * 1024; // 8 MB + +/// Environment variable for HTTP/2 max frame size (bytes) +/// Range: 16384 (16 KB) to 16777216 (16 MB) per RFC 7540 +/// Default: 524288 (512 KB) +pub const ENV_H2_MAX_FRAME_SIZE: &str = "RUSTFS_H2_MAX_FRAME_SIZE"; +pub const DEFAULT_H2_MAX_FRAME_SIZE: u32 = 512 * 1024; // 512 KB + +/// Environment variable for HTTP/2 max header list size (bytes) +/// Default: 65536 (64 KB) +pub const ENV_H2_MAX_HEADER_LIST_SIZE: &str = "RUSTFS_H2_MAX_HEADER_LIST_SIZE"; +pub const DEFAULT_H2_MAX_HEADER_LIST_SIZE: u32 = 64 * 1024; // 64 KB + +/// Environment variable for HTTP/2 max concurrent streams +/// Default: 2048 +pub const ENV_H2_MAX_CONCURRENT_STREAMS: &str = "RUSTFS_H2_MAX_CONCURRENT_STREAMS"; +pub const DEFAULT_H2_MAX_CONCURRENT_STREAMS: u32 = 2048; + +/// Environment variable for HTTP/2 keep-alive interval (seconds) +/// Default: 20 +pub const ENV_H2_KEEP_ALIVE_INTERVAL: &str = "RUSTFS_H2_KEEP_ALIVE_INTERVAL"; +pub const DEFAULT_H2_KEEP_ALIVE_INTERVAL: u64 = 20; + +/// Environment variable for HTTP/2 keep-alive timeout (seconds) +/// Default: 10 +pub const ENV_H2_KEEP_ALIVE_TIMEOUT: &str = "RUSTFS_H2_KEEP_ALIVE_TIMEOUT"; +pub const DEFAULT_H2_KEEP_ALIVE_TIMEOUT: u64 = 10; + +/// Environment variable for HTTP/1.1 header read timeout (seconds) +/// Default: 5 +pub const ENV_HTTP1_HEADER_READ_TIMEOUT: &str = "RUSTFS_HTTP1_HEADER_READ_TIMEOUT"; +pub const DEFAULT_HTTP1_HEADER_READ_TIMEOUT: u64 = 5; + +/// Environment variable for HTTP/1.1 max buffer size (bytes) +/// Default: 65536 (64 KB) +pub const ENV_HTTP1_MAX_BUF_SIZE: &str = "RUSTFS_HTTP1_MAX_BUF_SIZE"; +pub const DEFAULT_HTTP1_MAX_BUF_SIZE: usize = 64 * 1024; // 64 KB + +// ── TLS Hot Reload Parameters ── + +/// Environment variable to enable TLS certificate hot reload +/// Default: false +/// To enable, set the environment variable RUSTFS_TLS_RELOAD_ENABLE=1 +pub const ENV_TLS_RELOAD_ENABLE: &str = "RUSTFS_TLS_RELOAD_ENABLE"; + +/// Default value for TLS certificate hot reload +/// By default, RustFS does not reload TLS certificates automatically. +pub const DEFAULT_TLS_RELOAD_ENABLE: bool = false; + +/// Environment variable for TLS certificate reload interval (seconds) +/// Default: 30 seconds. Minimum: 5 seconds. +pub const ENV_TLS_RELOAD_INTERVAL: &str = "RUSTFS_TLS_RELOAD_INTERVAL"; + +/// Default interval for TLS certificate reload check +pub const DEFAULT_TLS_RELOAD_INTERVAL: u64 = 30; diff --git a/rustfs/Cargo.toml b/rustfs/Cargo.toml index 30210f5e8..16a27178c 100644 --- a/rustfs/Cargo.toml +++ b/rustfs/Cargo.toml @@ -47,6 +47,7 @@ webdav = ["rustfs-protocols/webdav"] license = [] direct-io = [] # Aligned direct I/O reader support (uses aligned pread, does not set O_DIRECT) io-scheduler-debug = [] # Enable debug information in I/O scheduler +tracing-chunk-debug = [] # Enable per-chunk tracing in data plane (high noise, for debugging only) full = ["metrics-gpu", "ftps", "swift", "webdav", "direct-io"] manual-test-runners = [] diff --git a/rustfs/src/main.rs b/rustfs/src/main.rs index a7db64faf..48f714643 100644 --- a/rustfs/src/main.rs +++ b/rustfs/src/main.rs @@ -43,7 +43,7 @@ use crate::init::init_webdav_system; use crate::capacity::capacity_integration::init_capacity_management; use crate::server::{ - SHUTDOWN_TIMEOUT, ServiceState, ServiceStateManager, ShutdownSignal, init_cert, init_event_notifier, shutdown_event_notifier, + SHUTDOWN_TIMEOUT, ServiceState, ServiceStateManager, ShutdownSignal, init_event_notifier, shutdown_event_notifier, start_audit_system, start_http_server, stop_audit_system, wait_for_shutdown, }; use license::{current_license, init_license, license_status}; @@ -229,15 +229,18 @@ async fn async_main() -> Result<()> { // A crypto provider is already installed (e.g. by the host process); this is fine. debug!("rustls crypto provider already installed, skipping aws-lc-rs default install"); } - // Initialize TLS if a certificate path is provided + // Initialize TLS outbound material (root CAs, mTLS identity) if configured. + // Server-side TLS acceptor is built separately inside start_http_server() + // using the same TlsMaterialSnapshot loading logic. if let Some(tls_path) = &config.tls_path { - match init_cert(tls_path).await { - Ok(_) => { - info!(target: "rustfs::main", "TLS initialized successfully with certs from {}", tls_path); + match crate::server::tls_material::TlsMaterialSnapshot::load(tls_path).await { + Ok(snapshot) => { + snapshot.apply_outbound().await; + info!(target: "rustfs::main", "TLS outbound material initialized from {}", tls_path); } Err(e) => { error!("Failed to initialize TLS from {}: {}", tls_path, e); - return Err(Error::other(e)); + return Err(Error::other(e.to_string())); } } } diff --git a/rustfs/src/server/cert.rs b/rustfs/src/server/cert.rs deleted file mode 100644 index 21c4954d9..000000000 --- a/rustfs/src/server/cert.rs +++ /dev/null @@ -1,296 +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 rustfs_common::{MtlsIdentityPem, set_global_mtls_identity, set_global_root_cert}; -use rustfs_config::{RUSTFS_CA_CERT, RUSTFS_PUBLIC_CERT, RUSTFS_TLS_CERT}; -use rustls::pki_types::{CertificateDer, PrivateKeyDer, pem::PemObject}; -use std::path::{Path, PathBuf}; -use tracing::{debug, info}; - -#[derive(Debug)] -pub enum RustFSError { - Cert(String), -} - -impl std::fmt::Display for RustFSError { - fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { - match self { - RustFSError::Cert(msg) => write!(f, "Certificate error: {msg}"), - } - } -} - -impl std::error::Error for RustFSError {} - -/// Parse PEM-encoded certificates into DER format. -/// Returns a vector of DER-encoded certificates. -/// -/// # Arguments -/// * `pem` - A byte slice containing the PEM-encoded certificates. -/// -/// # Returns -/// A vector of `CertificateDer` containing the DER-encoded certificates. -/// -/// # Errors -/// Returns `RustFSError` if parsing fails. -fn parse_pem_certs(pem: &[u8]) -> Result>, RustFSError> { - let mut out = Vec::new(); - let mut reader = std::io::Cursor::new(pem); - for item in CertificateDer::pem_reader_iter(&mut reader) { - let c = item.map_err(|e| RustFSError::Cert(format!("parse cert pem: {e}")))?; - out.push(c); - } - Ok(out) -} - -/// Parse a PEM-encoded private key into DER format. -/// Supports PKCS#8 and RSA private keys. -/// -/// # Arguments -/// * `pem` - A byte slice containing the PEM-encoded private key. -/// -/// # Returns -/// A `PrivateKeyDer` containing the DER-encoded private key. -/// -/// # Errors -/// Returns `RustFSError` if parsing fails or no key is found. -fn parse_pem_private_key(pem: &[u8]) -> Result, RustFSError> { - let mut reader = std::io::Cursor::new(pem); - PrivateKeyDer::from_pem_reader(&mut reader).map_err(|e| RustFSError::Cert(format!("parse private key pem: {e}"))) -} - -/// Helper function to read a file and return its contents. -/// Returns the file contents as a vector of bytes. -/// # Errors -/// Returns `RustFSError` if reading fails. -async fn read_file(path: &PathBuf, desc: &str) -> Result, RustFSError> { - tokio::fs::read(path) - .await - .map_err(|e| RustFSError::Cert(format!("read {desc} {path:?}: {e}"))) -} - -/// Initialize TLS material for both server and outbound client connections. -/// -/// Loads roots from: -/// - `${RUSTFS_TLS_PATH}/ca.crt` (or `tls/ca.crt`) -/// - `${RUSTFS_TLS_PATH}/public.crt` (optional additional root bundle) -/// - system roots if `RUSTFS_TRUST_SYSTEM_CA=true` (default: false) -/// - if `RUSTFS_TRUST_LEAF_CERT_AS_CA=true`, also loads leaf cert(s) from -/// `${RUSTFS_TLS_PATH}/rustfs_cert.pem` into the root store. -/// -/// Loads mTLS client identity (optional) from: -/// - `${RUSTFS_TLS_PATH}/client_cert.pem` -/// - `${RUSTFS_TLS_PATH}/client_key.pem` -/// -/// Environment overrides: -/// - RUSTFS_TLS_PATH -/// - RUSTFS_MTLS_CLIENT_CERT -/// - RUSTFS_MTLS_CLIENT_KEY -pub(crate) async fn init_cert(tls_path: &str) -> Result<(), RustFSError> { - if tls_path.is_empty() { - info!("No TLS path configured; skipping certificate initialization"); - return Ok(()); - } - - let tls_dir = PathBuf::from(tls_path); - - // Load root certificates - load_root_certs(&tls_dir).await?; - - // Load optional mTLS identity - load_mtls_identity(&tls_dir).await?; - - Ok(()) -} - -/// Load root certificates from various sources. -async fn load_root_certs(tls_dir: &Path) -> Result<(), RustFSError> { - let mut cert_data = Vec::new(); - - let trust_leaf_as_ca = - rustfs_utils::get_env_bool(rustfs_config::ENV_TRUST_LEAF_CERT_AS_CA, rustfs_config::DEFAULT_TRUST_LEAF_CERT_AS_CA); - if trust_leaf_as_ca { - walk_dir(tls_dir.to_path_buf(), RUSTFS_TLS_CERT, &mut cert_data).await; - info!("Loaded leaf certificate(s) as root CA as per RUSTFS_TRUST_LEAF_CERT_AS_CA"); - } - - // Try public.crt and ca.crt - let public_cert_path = tls_dir.join(RUSTFS_PUBLIC_CERT); - load_cert_file(public_cert_path.to_str().unwrap_or_default(), &mut cert_data, "CA certificate").await; - - let ca_cert_path = tls_dir.join(RUSTFS_CA_CERT); - load_cert_file(ca_cert_path.to_str().unwrap_or_default(), &mut cert_data, "CA certificate").await; - - // Load system root certificates if enabled - let trust_system_ca = rustfs_utils::get_env_bool(rustfs_config::ENV_TRUST_SYSTEM_CA, rustfs_config::DEFAULT_TRUST_SYSTEM_CA); - if trust_system_ca { - let system_ca_paths = [ - "/etc/ssl/certs/ca-certificates.crt", // Debian/Ubuntu/Alpine - "/etc/pki/tls/certs/ca-bundle.crt", // Fedora/RHEL/CentOS - "/etc/ssl/ca-bundle.pem", // OpenSUSE - "/etc/pki/tls/cacert.pem", // OpenELEC - "/etc/ssl/cert.pem", // macOS/FreeBSD - "/usr/local/etc/openssl/cert.pem", // macOS/Homebrew OpenSSL - "/usr/local/share/certs/ca-root-nss.crt", // FreeBSD - "/etc/pki/ca-trust/extracted/pem/tls-ca-bundle.pem", // RHEL - "/usr/share/pki/ca-trust-legacy/ca-bundle.legacy.crt", // RHEL legacy - ]; - - let mut system_cert_loaded = false; - for path in system_ca_paths { - if load_cert_file(path, &mut cert_data, "system root certificates").await { - system_cert_loaded = true; - info!("Loaded system root certificates from {}", path); - break; - } - } - - if !system_cert_loaded { - debug!("Could not find system root certificates in common locations."); - } - } else { - info!("Loading system root certificates disabled via RUSTFS_TRUST_SYSTEM_CA"); - } - - if !cert_data.is_empty() { - set_global_root_cert(cert_data).await; - info!("Configured custom root certificates for inter-node communication"); - } - - Ok(()) -} - -/// Load optional mTLS identity. -async fn load_mtls_identity(tls_dir: &Path) -> Result<(), RustFSError> { - let client_cert_path = match rustfs_utils::get_env_opt_str(rustfs_config::ENV_MTLS_CLIENT_CERT) { - Some(p) => PathBuf::from(p), - None => tls_dir.join(rustfs_config::RUSTFS_CLIENT_CERT_FILENAME), - }; - - let client_key_path = match rustfs_utils::get_env_opt_str(rustfs_config::ENV_MTLS_CLIENT_KEY) { - Some(p) => PathBuf::from(p), - None => tls_dir.join(rustfs_config::RUSTFS_CLIENT_KEY_FILENAME), - }; - - if client_cert_path.exists() && client_key_path.exists() { - let cert_bytes = read_file(&client_cert_path, "client cert").await?; - let key_bytes = read_file(&client_key_path, "client key").await?; - - // Validate parse-ability early; store as PEM bytes for tonic. - parse_pem_certs(&cert_bytes)?; - parse_pem_private_key(&key_bytes)?; - - let identity_pem = MtlsIdentityPem { - cert_pem: cert_bytes, - key_pem: key_bytes, - }; - - set_global_mtls_identity(Some(identity_pem)).await; - info!("Loaded mTLS client identity cert={:?} key={:?}", client_cert_path, client_key_path); - } else { - set_global_mtls_identity(None).await; - info!( - "mTLS client identity not configured (missing {:?} and/or {:?}); proceeding with server-only TLS", - client_cert_path, client_key_path - ); - } - - Ok(()) -} - -/// Helper function to load a certificate file and append to cert_data. -/// Returns true if the file was successfully loaded. -async fn load_cert_file(path: &str, cert_data: &mut Vec, desc: &str) -> bool { - if tokio::fs::metadata(path).await.is_ok() { - if let Ok(data) = tokio::fs::read(path).await { - cert_data.extend(data); - cert_data.push(b'\n'); - info!("Loaded {} from {}", desc, path); - true - } else { - debug!("Failed to read {} from {}", desc, path); - false - } - } else { - debug!("{} file not found at {}", desc, path); - false - } -} - -/// Load the certificate file if its name matches `cert_name`. -/// If it matches, the certificate data is appended to `cert_data`. -/// -/// # Parameters -/// - `entry`: The directory entry to check. -/// - `cert_name`: The name of the certificate file to match. -/// - `cert_data`: A mutable vector to append loaded certificate data. -async fn load_if_matches(entry: &tokio::fs::DirEntry, cert_name: &str, cert_data: &mut Vec) { - let fname = entry.file_name().to_string_lossy().to_string(); - if fname == cert_name { - let p = entry.path(); - load_cert_file(&p.to_string_lossy(), cert_data, "certificate").await; - } -} - -/// Search the directory at `path` and one level of subdirectories to find and load -/// certificates matching `cert_name`. Loaded certificate data is appended to -/// `cert_data`. -/// # Parameters -/// - `path`: The starting directory path to search for certificates. -/// - `cert_name`: The name of the certificate file to look for. -/// - `cert_data`: A mutable vector to append loaded certificate data. -async fn walk_dir(path: PathBuf, cert_name: &str, cert_data: &mut Vec) { - if let Ok(mut rd) = tokio::fs::read_dir(&path).await { - while let Ok(Some(entry)) = rd.next_entry().await { - if let Ok(ft) = entry.file_type().await { - if ft.is_file() { - load_if_matches(&entry, cert_name, cert_data).await; - } else if ft.is_dir() { - // Only check direct subdirectories, no deeper recursion - if let Ok(mut sub_rd) = tokio::fs::read_dir(&entry.path()).await { - while let Ok(Some(sub_entry)) = sub_rd.next_entry().await { - if let Ok(sub_ft) = sub_entry.file_type().await - && sub_ft.is_file() - { - load_if_matches(&sub_entry, cert_name, cert_data).await; - } - // Ignore subdirectories and symlinks in subdirs to limit to one level - } - } - } else if ft.is_symlink() { - // Follow symlink and treat target as file or directory, but limit to one level - if let Ok(meta) = tokio::fs::metadata(&entry.path()).await { - if meta.is_file() { - load_if_matches(&entry, cert_name, cert_data).await; - } else if meta.is_dir() { - // Treat as directory but only check its direct contents - if let Ok(mut sub_rd) = tokio::fs::read_dir(&entry.path()).await { - while let Ok(Some(sub_entry)) = sub_rd.next_entry().await { - if let Ok(sub_ft) = sub_entry.file_type().await - && sub_ft.is_file() - { - load_if_matches(&sub_entry, cert_name, cert_data).await; - } - // Ignore deeper levels - } - } - } - } - } - } - } - } else { - debug!("Certificate directory not found: {}", path.display()); - } -} diff --git a/rustfs/src/server/compress.rs b/rustfs/src/server/compress.rs index 24c4f39e3..09844ece3 100644 --- a/rustfs/src/server/compress.rs +++ b/rustfs/src/server/compress.rs @@ -50,6 +50,11 @@ use std::str::FromStr; use tower_http::compression::predicate::Predicate; use tracing::debug; +/// Response extension key for storing the request path category. +/// Set by `PathCategoryInjectionLayer` before the compression predicate evaluates. +#[derive(Debug, Clone, Copy)] +pub(crate) struct RequestPathCategory(pub(crate) PathCategory); + /// Configuration for HTTP response compression. /// /// This structure holds the whitelist-based compression settings: @@ -319,6 +324,156 @@ impl Predicate for CompressionPredicate { } } +// ── Path-Aware Compression ── + +/// Classifies request paths to determine if compression should apply. +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub(crate) enum PathCategory { + /// S3 data plane (bucket/key operations) — compression applies via whitelist + S3DataPlane, + /// Admin API paths — skip compression (small JSON responses) + AdminApi, + /// Console paths — skip compression (static assets, already optimized) + Console, + /// Internode RPC paths — skip compression (binary protocol data) + InternodeRpc, + /// Health/probe paths — skip compression (tiny responses) + Probe, +} + +impl PathCategory { + /// Classify a request URI path into a category. + pub(crate) fn classify(path: &str) -> Self { + if path.starts_with("/rustfs/rpc/") || path.starts_with("/rustfs/peer/") { + PathCategory::InternodeRpc + } else if path.starts_with("/rustfs/admin/") || path.starts_with("/minio/admin/") { + PathCategory::AdminApi + } else if path.starts_with("/rustfs/console") { + PathCategory::Console + } else if path.starts_with("/minio/health/") { + PathCategory::Probe + } else { + PathCategory::S3DataPlane + } + } + + /// Returns true if compression should be considered for this path category. + /// Only S3 data plane paths go through the full compression predicate. + #[inline] + pub(crate) fn should_evaluate_compression(self) -> bool { + matches!(self, PathCategory::S3DataPlane) + } +} + +/// A compression predicate that first checks the request path category +/// before evaluating the full compression rules. +/// +/// This avoids running MIME type / extension matching for admin, RPC, console, +/// and health probe paths where compression is never beneficial. +#[derive(Clone, Debug)] +pub(crate) struct PathAwareCompressionPredicate { + inner: CompressionPredicate, +} + +impl PathAwareCompressionPredicate { + pub(crate) fn new(config: CompressionConfig) -> Self { + Self { + inner: CompressionPredicate::new(config), + } + } +} + +impl Predicate for PathAwareCompressionPredicate { + fn should_compress(&self, response: &Response) -> bool + where + B: http_body::Body, + { + // Fast path: skip full predicate evaluation for non-S3 paths + if let Some(RequestPathCategory(category)) = response.extensions().get::() + && !category.should_evaluate_compression() + { + return false; + } + self.inner.should_compress(response) + } +} + +use http::Request; +use http_body::Body; +use std::pin::Pin; +use std::task::{Context, Poll}; +use tower::{Layer, Service}; + +/// Tower layer that injects `RequestPathCategory` into each response's extensions +/// based on the incoming request URI path. Must be placed before `CompressionLayer`. +#[derive(Clone, Copy, Debug)] +pub(crate) struct PathCategoryInjectionLayer; + +impl Layer for PathCategoryInjectionLayer { + type Service = PathCategoryInjectionService; + + fn layer(&self, inner: S) -> Self::Service { + PathCategoryInjectionService { inner } + } +} + +/// Service wrapper that adds `RequestPathCategory` to response extensions. +#[derive(Clone)] +pub(crate) struct PathCategoryInjectionService { + inner: S, +} + +pin_project_lite::pin_project! { + /// Future for `PathCategoryInjectionService` that injects path category into response. + #[project = InjectCategoryFutProj] + pub(crate) struct InjectCategoryFut { + #[pin] + inner: F, + category: PathCategory, + } +} + +impl std::future::Future for InjectCategoryFut +where + F: std::future::Future, E>>, +{ + type Output = Result, E>; + + fn poll(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll { + let this = self.project(); + match this.inner.poll(cx) { + Poll::Ready(Ok(mut resp)) => { + resp.extensions_mut().insert(RequestPathCategory(*this.category)); + Poll::Ready(Ok(resp)) + } + Poll::Ready(Err(e)) => Poll::Ready(Err(e)), + Poll::Pending => Poll::Pending, + } + } +} + +impl Service> for PathCategoryInjectionService +where + S: Service, Response = Response>, + ResBody: Body, +{ + type Response = Response; + type Error = S::Error; + type Future = InjectCategoryFut; + + fn poll_ready(&mut self, cx: &mut Context<'_>) -> Poll> { + self.inner.poll_ready(cx) + } + + fn call(&mut self, req: Request) -> Self::Future { + let category = PathCategory::classify(req.uri().path()); + InjectCategoryFut { + inner: self.inner.call(req), + category, + } + } +} + #[cfg(test)] mod tests { use super::*; @@ -471,4 +626,45 @@ mod tests { assert_eq!(predicate.config.mime_patterns.len(), 2); assert_eq!(predicate.config.min_size, 1000); } + + #[test] + fn test_path_category_classify_s3() { + assert_eq!(PathCategory::classify("/"), PathCategory::S3DataPlane); + assert_eq!(PathCategory::classify("/mybucket"), PathCategory::S3DataPlane); + assert_eq!(PathCategory::classify("/mybucket/mykey"), PathCategory::S3DataPlane); + assert_eq!(PathCategory::classify("/bucket?list-type=2"), PathCategory::S3DataPlane); + } + + #[test] + fn test_path_category_classify_admin() { + assert_eq!(PathCategory::classify("/rustfs/admin/v3/service"), PathCategory::AdminApi); + assert_eq!(PathCategory::classify("/minio/admin/v3/info"), PathCategory::AdminApi); + } + + #[test] + fn test_path_category_classify_console() { + assert_eq!(PathCategory::classify("/rustfs/console/index.html"), PathCategory::Console); + assert_eq!(PathCategory::classify("/rustfs/console"), PathCategory::Console); + } + + #[test] + fn test_path_category_classify_rpc() { + assert_eq!(PathCategory::classify("/rustfs/rpc/read_file_stream"), PathCategory::InternodeRpc); + assert_eq!(PathCategory::classify("/rustfs/peer/health"), PathCategory::InternodeRpc); + } + + #[test] + fn test_path_category_classify_probe() { + assert_eq!(PathCategory::classify("/minio/health/live"), PathCategory::Probe); + assert_eq!(PathCategory::classify("/minio/health/ready"), PathCategory::Probe); + } + + #[test] + fn test_path_category_should_evaluate() { + assert!(PathCategory::S3DataPlane.should_evaluate_compression()); + assert!(!PathCategory::AdminApi.should_evaluate_compression()); + assert!(!PathCategory::Console.should_evaluate_compression()); + assert!(!PathCategory::InternodeRpc.should_evaluate_compression()); + assert!(!PathCategory::Probe.should_evaluate_compression()); + } } diff --git a/rustfs/src/server/http.rs b/rustfs/src/server/http.rs index 85dc8e932..0fdcf3d1d 100644 --- a/rustfs/src/server/http.rs +++ b/rustfs/src/server/http.rs @@ -19,9 +19,10 @@ use crate::auth_keystone; use crate::config; use crate::server::{ ReadinessGateLayer, RemoteAddr, ServiceState, ServiceStateManager, - compress::{CompressionConfig, CompressionPredicate}, + compress::{CompressionConfig, PathAwareCompressionPredicate, PathCategoryInjectionLayer}, hybrid::hybrid, layer::{AdminChunkedContentLengthCompatLayer, ConditionalCorsLayer, ObjectAttributesEtagFixLayer, RedirectLayer}, + tls_material::{TlsAcceptorHolder, TlsHandshakeFailureKind, TlsMaterialSnapshot, spawn_reload_loop}, }; use crate::storage; use crate::storage::rpc::InternodeRpcService; @@ -38,7 +39,6 @@ use metrics::{counter, histogram}; use opentelemetry::global; use opentelemetry::trace::TraceContextExt; use rustfs_common::GlobalReadiness; -use rustfs_config::{RUSTFS_TLS_CERT, RUSTFS_TLS_KEY}; use rustfs_ecstore::rpc::{TONIC_RPC_PREFIX, verify_rpc_signature}; use rustfs_keystone::KeystoneAuthLayer; #[cfg(feature = "swift")] @@ -46,7 +46,6 @@ use rustfs_protocols::SwiftService; use rustfs_protos::proto_gen::node_service::node_service_server::NodeServiceServer; use rustfs_trusted_proxies::ClientInfo; use rustfs_utils::net::parse_and_resolve_address; -use rustls::ServerConfig; use s3s::{host::MultiDomain, service::S3Service, service::S3ServiceBuilder}; use socket2::{SockRef, TcpKeepalive}; use std::io::{Error, Result}; @@ -54,7 +53,6 @@ use std::net::SocketAddr; use std::sync::Arc; use std::time::Duration; use tokio::net::{TcpListener, TcpStream}; -use tokio_rustls::TlsAcceptor; use tonic::{Request, Status}; use tower::ServiceBuilder; use tower_http::add_extension::AddExtensionLayer; @@ -156,9 +154,24 @@ pub async fn start_http_server( TcpListener::from_std(socket.into())? }; - let tls_acceptor = setup_tls_acceptor(config.tls_path.as_deref().unwrap_or_default()).await?; + let tls_path = config.tls_path.as_deref().unwrap_or_default(); + // Load TLS materials and build server acceptor. + // Note: outbound material (root CAs, mTLS identity) is already applied in main.rs. + let tls_snapshot = TlsMaterialSnapshot::load(tls_path) + .await + .map_err(|e| std::io::Error::other(e.to_string()))?; + + let tls_acceptor = tls_snapshot + .build_tls_acceptor(tls_path) + .await + .map_err(|e| std::io::Error::other(e.to_string()))?; let tls_enabled = tls_acceptor.is_some(); let protocol = if tls_enabled { "https" } else { "http" }; + + // Spawn background TLS certificate hot-reload loop (if enabled). + if let Some(holder) = &tls_acceptor { + spawn_reload_loop(tls_path.to_string(), holder.clone()); + } // Obtain the listener address let local_addr: SocketAddr = listener.local_addr()?; let local_ip = match rustfs_utils::get_local_ip() { @@ -273,11 +286,53 @@ pub async fn start_http_server( (sigterm_inner, sigint_inner) }; - // RustFS Transport Layer Configuration Constants - Optimized for S3 Workloads - const H2_INITIAL_STREAM_WINDOW_SIZE: u32 = 1024 * 1024 * 4; // 4MB: Optimize large file throughput - const H2_INITIAL_CONN_WINDOW_SIZE: u32 = 1024 * 1024 * 8; // 8MB: Link-level flow control - const H2_MAX_FRAME_SIZE: u32 = 512 * 1024; // 512KB: Reduce framing overhead for large objects - const H2_MAX_HEADER_LIST_SIZE: u32 = 64 * 1024; // 64KB: Conservative header limit to mitigate DoS risk + // ── HTTP Transport Tuning (configurable via env vars) ── + // Read all transport parameters from environment, falling back to defaults. + // H2 frame size is clamped to RFC 7540 range: 2^14 (16KB) to 2^24 (16MB). + + let h2_stream_window = rustfs_utils::get_env_u32( + rustfs_config::ENV_H2_INITIAL_STREAM_WINDOW_SIZE, + rustfs_config::DEFAULT_H2_INITIAL_STREAM_WINDOW_SIZE, + ); + let h2_conn_window = rustfs_utils::get_env_u32( + rustfs_config::ENV_H2_INITIAL_CONN_WINDOW_SIZE, + rustfs_config::DEFAULT_H2_INITIAL_CONN_WINDOW_SIZE, + ); + let h2_max_frame_size = + rustfs_utils::get_env_u32(rustfs_config::ENV_H2_MAX_FRAME_SIZE, rustfs_config::DEFAULT_H2_MAX_FRAME_SIZE) + .clamp(16_384, 16_777_216); // RFC 7540 + let h2_max_header_list_size = + rustfs_utils::get_env_u32(rustfs_config::ENV_H2_MAX_HEADER_LIST_SIZE, rustfs_config::DEFAULT_H2_MAX_HEADER_LIST_SIZE); + let h2_max_concurrent_streams = rustfs_utils::get_env_u32( + rustfs_config::ENV_H2_MAX_CONCURRENT_STREAMS, + rustfs_config::DEFAULT_H2_MAX_CONCURRENT_STREAMS, + ) + .max(1); + let h2_keep_alive_interval = + rustfs_utils::get_env_u64(rustfs_config::ENV_H2_KEEP_ALIVE_INTERVAL, rustfs_config::DEFAULT_H2_KEEP_ALIVE_INTERVAL); + let h2_keep_alive_timeout = + rustfs_utils::get_env_u64(rustfs_config::ENV_H2_KEEP_ALIVE_TIMEOUT, rustfs_config::DEFAULT_H2_KEEP_ALIVE_TIMEOUT); + let http1_header_read_timeout = rustfs_utils::get_env_u64( + rustfs_config::ENV_HTTP1_HEADER_READ_TIMEOUT, + rustfs_config::DEFAULT_HTTP1_HEADER_READ_TIMEOUT, + ); + let http1_max_buf_size = + rustfs_utils::get_env_usize(rustfs_config::ENV_HTTP1_MAX_BUF_SIZE, rustfs_config::DEFAULT_HTTP1_MAX_BUF_SIZE); + + info!( + "HTTP transport parameters: h2_stream_window={}, h2_conn_window={}, h2_max_frame={}, \ + h2_max_header_list={}, h2_max_concurrent_streams={}, h2_keepalive_interval={}s, \ + h2_keepalive_timeout={}s, http1_header_timeout={}s, http1_max_buf={}", + h2_stream_window, + h2_conn_window, + h2_max_frame_size, + h2_max_header_list_size, + h2_max_concurrent_streams, + h2_keep_alive_interval, + h2_keep_alive_timeout, + http1_header_read_timeout, + http1_max_buf_size, + ); let mut conn_builder = ConnBuilder::new(TokioExecutor::new()); @@ -286,8 +341,8 @@ pub async fn start_http_server( .http1() .timer(TokioTimer::new()) .keep_alive(true) - .header_read_timeout(Duration::from_secs(5)) - .max_buf_size(64 * 1024) + .header_read_timeout(Duration::from_secs(http1_header_read_timeout)) + .max_buf_size(http1_max_buf_size) .writev(true); // Optimize for HTTP/2 (AI/Data Lake high concurrency synchronization) @@ -295,13 +350,13 @@ pub async fn start_http_server( .http2() .timer(TokioTimer::new()) .adaptive_window(true) - .initial_stream_window_size(H2_INITIAL_STREAM_WINDOW_SIZE) - .initial_connection_window_size(H2_INITIAL_CONN_WINDOW_SIZE) - .max_frame_size(H2_MAX_FRAME_SIZE) - .max_concurrent_streams(Some(2048)) - .max_header_list_size(H2_MAX_HEADER_LIST_SIZE) - .keep_alive_interval(Some(Duration::from_secs(20))) - .keep_alive_timeout(Duration::from_secs(10)); + .initial_stream_window_size(h2_stream_window) + .initial_connection_window_size(h2_conn_window) + .max_frame_size(h2_max_frame_size) + .max_concurrent_streams(Some(h2_max_concurrent_streams)) + .max_header_list_size(h2_max_header_list_size) + .keep_alive_interval(Some(Duration::from_secs(h2_keep_alive_interval))) + .keep_alive_timeout(Duration::from_secs(h2_keep_alive_timeout)); let http_server = Arc::new(conn_builder); let mut ctrl_c = std::pin::pin!(tokio::signal::ctrl_c()); @@ -310,10 +365,7 @@ pub async fn start_http_server( // service ready worker_state_manager.update(ServiceState::Ready); - let tls_acceptor = tls_acceptor.map(Arc::new); - - // Initialize keepalive configuration once to avoid recreation in the loop - let keepalive_conf = get_default_tcp_keepalive(); + // tls_acceptor is already Option>, clone for the loop loop { debug!("Waiting for new connection..."); @@ -374,28 +426,29 @@ pub async fn start_http_server( let socket_ref = SockRef::from(&socket); - // Enable TCP Keepalive to detect dead clients (e.g. power loss) - if let Err(err) = socket_ref.set_tcp_keepalive(&keepalive_conf) { - warn!(?err, "Failed to set TCP_KEEPALIVE"); - } + // ── POST-ACCEPT SOCKET SYSCALLS ── + // The listening socket already sets TCP_NODELAY, TCP_KEEPALIVE, + // SO_RCVBUF, and SO_SNDBUF. On Linux/BSD, these are inherited by + // accepted sockets, so we skip redundant re-application here. + // + // Only TCP_QUICKACK (Linux) is kept — it is inherently per-connection + // and NOT inherited from the listening socket. + // + // T03 optimized: syscall count reduced from 5 → 1 (Linux) / 0 (other) - // Disable Nagle algorithm: Critical for 4KB Payload, achieving ultra-low latency - if let Err(err) = socket_ref.set_tcp_nodelay(true) { - warn!(?err, "Failed to set TCP_NODELAY"); - } - - // Enable TCP QuickAck to reduce latency for small requests + // Enable TCP QuickAck to reduce latency for small requests (Linux only) #[cfg(target_os = "linux")] if let Err(err) = socket_ref.set_tcp_quickack(true) { debug!(?err, "Failed to set TCP_QUICKACK"); } - // Increase receive/send buffer to support BDP at GB-level throughput - if let Err(err) = socket_ref.set_recv_buffer_size(4 * rustfs_config::MI_B) { - warn!(?err, "Failed to set set_recv_buffer_size"); - } - if let Err(err) = socket_ref.set_send_buffer_size(4 * rustfs_config::MI_B) { - warn!(?err, "Failed to set set_send_buffer_size"); + // Debug-only: verify listening socket options were inherited + #[cfg(debug_assertions)] + { + debug!( + nodelay = socket_ref.tcp_nodelay().unwrap_or(false), + "TCP_NODELAY inherited from listening socket" + ); } let connection_ctx = ConnectionContext { @@ -404,6 +457,8 @@ pub async fn start_http_server( compression_config: compression_config.clone(), is_console, readiness: readiness.clone(), + keystone_auth: auth_keystone::get_keystone_auth(), + trusted_proxy_layer: rustfs_trusted_proxies::is_enabled().then(|| rustfs_trusted_proxies::layer().clone()), }; process_connection(socket, tls_acceptor.clone(), connection_ctx, graceful.clone()); @@ -433,88 +488,6 @@ pub async fn start_http_server( Ok(shutdown_tx) } -/// Sets up the TLS acceptor if certificates are available. -#[instrument(skip(tls_path))] -async fn setup_tls_acceptor(tls_path: &str) -> Result> { - if tls_path.is_empty() || tokio::fs::metadata(tls_path).await.is_err() { - debug!("TLS path is not provided or does not exist, starting with HTTP"); - return Ok(None); - } - debug!("Found TLS directory, checking for certificates"); - - let mtls_verifier = rustfs_utils::build_webpki_client_verifier(tls_path)?; - // 1. Attempt to load all certificates in the directory (multi-certificate support, for SNI) - if let Ok(cert_key_pairs) = rustfs_utils::load_all_certs_from_directory(tls_path) - && !cert_key_pairs.is_empty() - { - debug!("Found {} certificates, creating SNI-aware multi-cert resolver", cert_key_pairs.len()); - - // Create an SNI-enabled certificate resolver - let resolver = rustfs_utils::create_multi_cert_resolver(cert_key_pairs)?; - - // Configure the server to enable SNI support - let mut server_config = if let Some(verifier) = mtls_verifier.clone() { - ServerConfig::builder() - .with_client_cert_verifier(verifier) - .with_cert_resolver(Arc::new(resolver)) - } else { - ServerConfig::builder() - .with_no_client_auth() - .with_cert_resolver(Arc::new(resolver)) - }; - - // Configure ALPN protocol priority - server_config.alpn_protocols = vec![b"h2".to_vec(), b"http/1.1".to_vec(), b"http/1.0".to_vec()]; - - // Enable session resumption to reduce handshake overhead for returning clients - server_config.session_storage = rustls::server::ServerSessionMemoryCache::new(10000); - - // Log SNI requests - if rustfs_utils::tls_key_log() { - server_config.key_log = Arc::new(rustls::KeyLogFile::new()); - } - - return Ok(Some(TlsAcceptor::from(Arc::new(server_config)))); - } - - // 2. Revert to the traditional single-certificate mode - let key_path = format!("{tls_path}/{RUSTFS_TLS_KEY}"); - let cert_path = format!("{tls_path}/{RUSTFS_TLS_CERT}"); - if tokio::try_join!(tokio::fs::metadata(&key_path), tokio::fs::metadata(&cert_path)).is_ok() { - debug!("Found legacy single TLS certificate, starting with HTTPS"); - let certs = rustfs_utils::load_certs(&cert_path).map_err(|e| rustfs_utils::certs_error(e.to_string()))?; - let key = rustfs_utils::load_private_key(&key_path).map_err(|e| rustfs_utils::certs_error(e.to_string()))?; - - let mut server_config = if let Some(verifier) = mtls_verifier { - ServerConfig::builder() - .with_client_cert_verifier(verifier) - .with_single_cert(certs, key) - .map_err(|e| rustfs_utils::certs_error(e.to_string()))? - } else { - ServerConfig::builder() - .with_no_client_auth() - .with_single_cert(certs, key) - .map_err(|e| rustfs_utils::certs_error(e.to_string()))? - }; - - // Configure ALPN protocol priority - server_config.alpn_protocols = vec![b"h2".to_vec(), b"http/1.1".to_vec(), b"http/1.0".to_vec()]; - - // Enable session resumption to reduce handshake overhead for returning clients - server_config.session_storage = rustls::server::ServerSessionMemoryCache::new(10000); - - // Log SNI requests - if rustfs_utils::tls_key_log() { - server_config.key_log = Arc::new(rustls::KeyLogFile::new()); - } - - return Ok(Some(TlsAcceptor::from(Arc::new(server_config)))); - } - - debug!("No valid TLS certificates found in the directory, starting with HTTP"); - Ok(None) -} - #[derive(Clone)] struct ConnectionContext { http_server: Arc>, @@ -522,6 +495,10 @@ struct ConnectionContext { compression_config: CompressionConfig, is_console: bool, readiness: Arc, + /// Pre-computed Keystone auth provider (avoids per-connection OnceLock read). + keystone_auth: Option>, + /// Pre-computed trusted proxy layer (avoids per-connection is_enabled() check). + trusted_proxy_layer: Option, } /// Adapter that implements the OpenTelemetry [`Extractor`] trait for Hyper's @@ -569,7 +546,7 @@ impl<'a> opentelemetry::propagation::Extractor for HeaderMapCarrier<'a> { ))] fn process_connection( socket: TcpStream, - tls_acceptor: Option>, + tls_acceptor: Option>, context: ConnectionContext, graceful: Arc, ) { @@ -580,15 +557,16 @@ fn process_connection( compression_config, is_console, readiness, + keystone_auth, + trusted_proxy_layer, } = context; - // Build services inside each connected task to avoid passing complex service types across tasks, - // It also ensures that each connection has an independent service instance. + // Build the hybrid service per-connection. + // Note: NodeService is not Clone (holds LocalPeerS3Client), and the SwiftService + // type is feature-gated, so we cannot pre-build the full hybrid service. + // The construction cost is negligible (struct wrapping only, no I/O). let rpc_service = NodeServiceServer::with_interceptor(make_server(), check_auth); - // Wrap S3 service with Swift service to handle Swift API requests - // Swift API is only available when compiled with the 'swift' feature - // When enabled, Swift routes are handled at /v1/AUTH_* paths by default #[cfg(feature = "swift")] let http_service = SwiftService::new(true, None, s3_service); #[cfg(not(feature = "swift"))] @@ -607,6 +585,26 @@ fn process_connection( None } }; + // ── Canonical Middleware Stack Order (outermost → innermost) ── + // This order MUST be preserved across refactorings. + // Only AddExtensionLayer (layers 1-2) are per-connection; layers 3-15 are stateless. + // + // 1. AddExtensionLayer — per-connection peer address + // 2. AddExtensionLayer — per-connection raw socket addr (TrustedProxy) + // 3. TrustedProxyLayer — conditional, parses X-Forwarded-For + // 4. SetRequestIdLayer — generates X-Request-ID + // 5. AdminChunkedContentLengthCompatLayer — admin API compat + // 6. CatchPanicLayer — panic → 500 + // 7. ReadinessGateLayer — blocks until ready + // 8. KeystoneAuthLayer — X-Auth-Token validation + // 9. TraceLayer — request/response tracing + metrics + // 10. PropagateRequestIdLayer — X-Request-ID → response + // 11. PathCategoryInjectionLayer — injects path category for compression + // 12. CompressionLayer — response compression (whitelist, path-aware) + // 13. ObjectAttributesEtagFixLayer — ETag fix for GetObjectAttributes + // 14. ConditionalCorsLayer — S3 API CORS + // 15. RedirectLayer — console redirect (conditional) + // ───────────────────────────────────────────────────────────── let hybrid_service = ServiceBuilder::new() // NOTE: Both extension types are intentionally inserted to maintain compatibility: // 1. `Option` - Used by existing admin/storage handlers throughout the codebase @@ -618,11 +616,8 @@ fn process_connection( .option_layer(remote_addr.map(|ra| AddExtensionLayer::new(ra.0))) // Add TrustedProxyLayer to handle X-Forwarded-For and other proxy headers // This should be placed before TraceLayer so that logs reflect the real client IP - .option_layer(if rustfs_trusted_proxies::is_enabled() { - Some(rustfs_trusted_proxies::layer().clone()) - } else { - None - }) + // Pre-computed in ConnectionContext to avoid per-connection is_enabled() check. + .option_layer(trusted_proxy_layer) .layer(SetRequestIdLayer::x_request_id(MakeRequestUuid)) .layer(AdminChunkedContentLengthCompatLayer) .layer(CatchPanicLayer::new()) @@ -632,10 +627,8 @@ fn process_connection( // Add Keystone authentication middleware // This validates X-Auth-Token headers and stores credentials in task-local storage // Must be placed AFTER ReadinessGateLayer but BEFORE business logic - .layer({ - let keystone_auth = auth_keystone::get_keystone_auth(); - KeystoneAuthLayer::new(keystone_auth) - }) + // Pre-computed in ConnectionContext to avoid per-connection OnceLock read. + .layer(KeystoneAuthLayer::new(keystone_auth)) .layer( TraceLayer::new_for_http() .make_span_with(|request: &HttpRequest<_>| { @@ -702,13 +695,27 @@ fn process_connection( debug!("http response generated in {:?}", latency) }) .on_body_chunk(|chunk: &Bytes, latency: Duration, span: &Span| { - let _enter = span.enter(); - histogram!("rustfs.request.body.len").record(chunk.len() as f64); - debug!("http body sending {} bytes in {:?}", chunk.len(), latency); + #[cfg(feature = "tracing-chunk-debug")] + { + let _enter = span.enter(); + histogram!("rustfs.request.body.len").record(chunk.len() as f64); + debug!("http body sending {} bytes in {:?}", chunk.len(), latency); + } + #[cfg(not(feature = "tracing-chunk-debug"))] + { + let _ = (chunk, latency, span); + } }) .on_eos(|_trailers: Option<&HeaderMap>, stream_duration: Duration, span: &Span| { - let _enter = span.enter(); - debug!("http stream closed after {:?}", stream_duration) + #[cfg(feature = "tracing-chunk-debug")] + { + let _enter = span.enter(); + debug!("http stream closed after {:?}", stream_duration); + } + #[cfg(not(feature = "tracing-chunk-debug"))] + { + let _ = (_trailers, stream_duration, span); + } }) .on_failure(|_error, latency: Duration, span: &Span| { let _enter = span.enter(); @@ -719,7 +726,8 @@ fn process_connection( .layer(PropagateRequestIdLayer::x_request_id()) // Compress responses based on whitelist configuration // Only compresses when enabled and matches configured extensions/MIME types - .layer(CompressionLayer::new().compress_when(CompressionPredicate::new(compression_config))) + .layer(PathCategoryInjectionLayer) + .layer(CompressionLayer::new().compress_when(PathAwareCompressionPredicate::new(compression_config))) .layer(ObjectAttributesEtagFixLayer) // Conditional CORS layer: only applies to S3 API requests (not Admin, not Console) // Admin has its own CORS handling in router.rs @@ -733,12 +741,13 @@ fn process_connection( let hybrid_service = TowerToHyperService::new(hybrid_service); // Decide whether to handle HTTPS or HTTP connections based on the existence of TLS Acceptor - if let Some(acceptor) = tls_acceptor { + if let Some(holder) = tls_acceptor { debug!("TLS handshake start"); let peer_addr = socket .peer_addr() .ok() .map_or_else(|| "unknown".to_string(), |addr| addr.to_string()); + let acceptor = holder.get(); match acceptor.accept(socket).await { Ok(tls_socket) => { debug!("TLS handshake successful"); @@ -749,32 +758,26 @@ fn process_connection( } } Err(err) => { - // Detailed analysis of the reasons why the TLS handshake fails let err_str = err.to_string(); - let mut key_failure_type_str: &str = "UNKNOWN"; - if err_str.contains("unexpected EOF") || err_str.contains("handshake eof") { - warn!(peer_addr = %peer_addr, "TLS handshake failed. If this client needs HTTP, it should connect to the HTTP port instead"); - key_failure_type_str = "UNEXPECTED_EOF"; - } else if err_str.contains("protocol version") { - error!( - peer_addr = %peer_addr, - "TLS handshake failed due to protocol version mismatch: {}", err - ); - key_failure_type_str = "PROTOCOL_VERSION"; - } else if err_str.contains("certificate") { - error!( - peer_addr = %peer_addr, - "TLS handshake failed due to certificate issues: {}", err - ); - key_failure_type_str = "CERTIFICATE"; - } else { - error!( - peer_addr = %peer_addr, - "TLS handshake failed: {}", err - ); + let kind = TlsHandshakeFailureKind::classify(&err_str); + match kind { + TlsHandshakeFailureKind::UnexpectedEof => { + warn!(peer_addr = %peer_addr, "TLS handshake failed (unexpected EOF). If this client needs HTTP, it should connect to the HTTP port instead"); + } + TlsHandshakeFailureKind::ProtocolVersion => { + error!(peer_addr = %peer_addr, "TLS handshake failed (protocol version mismatch): {}", err); + } + TlsHandshakeFailureKind::Certificate => { + error!(peer_addr = %peer_addr, "TLS handshake failed (certificate issue): {}", err); + } + TlsHandshakeFailureKind::Alert => { + error!(peer_addr = %peer_addr, "TLS handshake failed (alert): {}", err); + } + TlsHandshakeFailureKind::Unknown => { + error!(peer_addr = %peer_addr, "TLS handshake failed: {}", err); + } } - counter!("rustfs_tls_handshake_failures", &[("key_failure_type", key_failure_type_str)]).increment(1); - // Record detailed diagnostic information + counter!("rustfs_tls_handshake_failures", &[("failure_type", kind.as_str())]).increment(1); debug!( peer_addr = %peer_addr, error_type = %std::any::type_name_of_val(&err), @@ -914,6 +917,68 @@ mod tests { use http::HeaderMap; use opentelemetry::propagation::Extractor; + /// Baseline constants — reference the authoritative config defaults. + /// If a config default changes, tests automatically follow. + mod baseline { + use rustfs_config::{ + DEFAULT_H2_INITIAL_CONN_WINDOW_SIZE, DEFAULT_H2_INITIAL_STREAM_WINDOW_SIZE, DEFAULT_H2_MAX_FRAME_SIZE, + DEFAULT_H2_MAX_HEADER_LIST_SIZE, DEFAULT_HTTP1_HEADER_READ_TIMEOUT, DEFAULT_HTTP1_MAX_BUF_SIZE, + }; + + /// Number of middleware layers in the canonical stack order (see http.rs). + /// Layers 1-2 are per-connection (AddExtension), 3-15 are stateless. + pub const MIDDLEWARE_LAYER_COUNT: usize = 15; + + /// Current HTTP/2 defaults (from rustfs_config). + pub const H2_INITIAL_STREAM_WINDOW_SIZE: u32 = DEFAULT_H2_INITIAL_STREAM_WINDOW_SIZE; + pub const H2_INITIAL_CONN_WINDOW_SIZE: u32 = DEFAULT_H2_INITIAL_CONN_WINDOW_SIZE; + pub const H2_MAX_FRAME_SIZE: u32 = DEFAULT_H2_MAX_FRAME_SIZE; + pub const H2_MAX_HEADER_LIST_SIZE: u32 = DEFAULT_H2_MAX_HEADER_LIST_SIZE; + + /// Current HTTP/1.1 defaults (from rustfs_config). + pub const HTTP1_HEADER_READ_TIMEOUT_SECS: u64 = DEFAULT_HTTP1_HEADER_READ_TIMEOUT; + pub const HTTP1_MAX_BUF_SIZE: usize = DEFAULT_HTTP1_MAX_BUF_SIZE; + + /// Post-accept socket syscalls after T03 optimization. + /// Linux: 1 (TCP_QUICKACK only). Other platforms: 0. + #[cfg(target_os = "linux")] + pub const POST_ACCEPT_SYSCALL_COUNT_LINUX: usize = 1; + #[cfg(not(target_os = "linux"))] + pub const POST_ACCEPT_SYSCALL_COUNT_OTHER: usize = 0; + } + + #[test] + fn test_baseline_h2_constants() { + use rustfs_config::{ + DEFAULT_H2_INITIAL_CONN_WINDOW_SIZE, DEFAULT_H2_INITIAL_STREAM_WINDOW_SIZE, DEFAULT_H2_MAX_FRAME_SIZE, + DEFAULT_H2_MAX_HEADER_LIST_SIZE, + }; + assert_eq!(baseline::H2_INITIAL_STREAM_WINDOW_SIZE, DEFAULT_H2_INITIAL_STREAM_WINDOW_SIZE); + assert_eq!(baseline::H2_INITIAL_CONN_WINDOW_SIZE, DEFAULT_H2_INITIAL_CONN_WINDOW_SIZE); + assert_eq!(baseline::H2_MAX_FRAME_SIZE, DEFAULT_H2_MAX_FRAME_SIZE); + assert_eq!(baseline::H2_MAX_HEADER_LIST_SIZE, DEFAULT_H2_MAX_HEADER_LIST_SIZE); + } + + #[test] + fn test_baseline_http1_constants() { + use rustfs_config::{DEFAULT_HTTP1_HEADER_READ_TIMEOUT, DEFAULT_HTTP1_MAX_BUF_SIZE}; + assert_eq!(baseline::HTTP1_HEADER_READ_TIMEOUT_SECS, DEFAULT_HTTP1_HEADER_READ_TIMEOUT); + assert_eq!(baseline::HTTP1_MAX_BUF_SIZE, DEFAULT_HTTP1_MAX_BUF_SIZE); + } + + #[test] + fn test_baseline_middleware_count() { + assert_eq!(baseline::MIDDLEWARE_LAYER_COUNT, 15); + } + + #[test] + fn test_baseline_post_accept_syscall_count() { + #[cfg(target_os = "linux")] + assert_eq!(baseline::POST_ACCEPT_SYSCALL_COUNT_LINUX, 1); + #[cfg(not(target_os = "linux"))] + assert_eq!(baseline::POST_ACCEPT_SYSCALL_COUNT_OTHER, 0); + } + #[test] fn test_headermap_carrier_new() { let headers = HeaderMap::new(); diff --git a/rustfs/src/server/mod.rs b/rustfs/src/server/mod.rs index 3da8c71c1..c460626ca 100644 --- a/rustfs/src/server/mod.rs +++ b/rustfs/src/server/mod.rs @@ -13,7 +13,6 @@ // limitations under the License. mod audit; -mod cert; mod compress; pub mod cors; mod event; @@ -24,9 +23,9 @@ mod prefix; mod readiness; mod runtime; mod service_state; +pub(crate) mod tls_material; pub(crate) use audit::{start_audit_system, stop_audit_system}; -pub(crate) use cert::init_cert; pub(crate) use event::{init_event_notifier, shutdown_event_notifier}; pub(crate) use http::start_http_server; pub(crate) use prefix::*; diff --git a/rustfs/src/server/tls_material.rs b/rustfs/src/server/tls_material.rs new file mode 100644 index 000000000..f6caf9ada --- /dev/null +++ b/rustfs/src/server/tls_material.rs @@ -0,0 +1,510 @@ +// 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. + +//! Unified TLS Material Snapshot +//! +//! Provides a single loading point for all TLS materials, eliminating duplicate +//! directory scanning and PEM parsing between outbound and inbound paths. +//! +//! Usage: +//! 1. Call `TlsMaterialSnapshot::load(tls_path)` once at startup. +//! 2. Call `snapshot.apply_outbound()` to set global root CAs and mTLS identity. +//! 3. Call `snapshot.build_tls_acceptor(tls_path)` to build the server TLS acceptor. + +use rustfs_common::{MtlsIdentityPem, set_global_mtls_identity, set_global_root_cert}; +use rustfs_config::{ + DEFAULT_TLS_RELOAD_ENABLE, DEFAULT_TLS_RELOAD_INTERVAL, DEFAULT_TRUST_LEAF_CERT_AS_CA, DEFAULT_TRUST_SYSTEM_CA, + ENV_MTLS_CLIENT_CERT, ENV_MTLS_CLIENT_KEY, ENV_TLS_RELOAD_ENABLE, ENV_TLS_RELOAD_INTERVAL, ENV_TRUST_LEAF_CERT_AS_CA, + ENV_TRUST_SYSTEM_CA, RUSTFS_CA_CERT, RUSTFS_CLIENT_CERT_FILENAME, RUSTFS_CLIENT_KEY_FILENAME, RUSTFS_PUBLIC_CERT, + RUSTFS_TLS_CERT, RUSTFS_TLS_KEY, +}; +use rustfs_utils::{get_env_bool, get_env_opt_str}; +use rustls::pki_types::{CertificateDer, PrivateKeyDer, pem::PemObject}; +use std::path::{Path, PathBuf}; +use std::sync::Arc; +use std::sync::RwLock; +use std::time::Duration; +use tokio_rustls::TlsAcceptor; +use tracing::{debug, info, warn}; + +/// System CA certificate search paths (platform-specific). +const SYSTEM_CA_PATHS: &[&str] = &[ + "/etc/ssl/certs/ca-certificates.crt", // Debian/Ubuntu/Alpine + "/etc/pki/tls/certs/ca-bundle.crt", // Fedora/RHEL/CentOS + "/etc/ssl/ca-bundle.pem", // OpenSUSE + "/etc/pki/tls/cacert.pem", // OpenELEC + "/etc/ssl/cert.pem", // macOS/FreeBSD + "/usr/local/etc/openssl/cert.pem", // macOS/Homebrew OpenSSL + "/usr/local/share/certs/ca-root-nss.crt", // FreeBSD + "/etc/pki/ca-trust/extracted/pem/tls-ca-bundle.pem", // RHEL + "/usr/share/pki/ca-trust-legacy/ca-bundle.legacy.crt", // RHEL legacy +]; + +/// Outbound TLS material for client connections (inter-node RPC). +#[derive(Debug, Clone)] +pub struct OutboundTlsMaterial { + /// Concatenated PEM-encoded root CA certificates. + pub root_ca_pem: Vec, + /// Optional mTLS client identity. + pub mtls_identity: Option, +} + +/// Complete TLS material snapshot loaded once at startup. +#[derive(Debug)] +pub struct TlsMaterialSnapshot { + /// Material for outbound client connections. + pub outbound: OutboundTlsMaterial, + /// Whether any server certificates were found. + pub has_server_certs: bool, +} + +impl TlsMaterialSnapshot { + /// Load all TLS materials from the given directory. + /// + /// This is the single entry point that replaces both the old + /// `cert.rs::init_cert()` and `http.rs::setup_tls_acceptor()` loading logic. + pub async fn load(tls_path: &str) -> Result { + if tls_path.is_empty() { + info!("No TLS path configured; skipping TLS material loading"); + return Ok(Self::empty()); + } + + let tls_dir = PathBuf::from(tls_path); + + // Load outbound material (root CAs + mTLS identity) + let outbound = load_outbound_material(&tls_dir).await?; + + // Check if server certs exist (actual loading happens in build_tls_acceptor) + let has_server_certs = has_server_certificates(tls_path).await; + + Ok(Self { + outbound, + has_server_certs, + }) + } + + /// Apply outbound material to global state (root CAs, mTLS identity). + pub async fn apply_outbound(&self) { + if !self.outbound.root_ca_pem.is_empty() { + set_global_root_cert(self.outbound.root_ca_pem.clone()).await; + info!("Configured custom root certificates for inter-node communication"); + } + set_global_mtls_identity(self.outbound.mtls_identity.clone()).await; + } + + /// Build a `TlsAcceptorHolder` from the loaded snapshot. + /// + /// This is the single place that constructs the server `ServerConfig`, + /// handling both multi-cert (SNI resolver) and single-cert fallback. + /// Returns `None` if no TLS certificates are available. + pub async fn build_tls_acceptor(&self, tls_path: &str) -> Result>, TlsMaterialError> { + if tls_path.is_empty() || !self.has_server_certs { + return Ok(None); + } + + let mtls_verifier = rustfs_utils::build_webpki_client_verifier(tls_path) + .map_err(|e| TlsMaterialError::Io(format!("build mTLS verifier: {e}")))?; + + // Try multi-cert (SNI) first + match rustfs_utils::load_all_certs_from_directory(tls_path) { + Ok(cert_key_pairs) if !cert_key_pairs.is_empty() => match rustfs_utils::create_multi_cert_resolver(cert_key_pairs) { + Ok(resolver) => { + let config = build_server_config(ServerCertSource::Resolver(Arc::new(resolver)), mtls_verifier)?; + info!("Created TLS acceptor with SNI resolver"); + let acceptor = Arc::new(TlsAcceptor::from(Arc::new(config))); + return Ok(Some(Arc::new(TlsAcceptorHolder::new(acceptor)))); + } + Err(e) => warn!("Failed to build multi-cert resolver: {}, falling back to single-cert", e), + }, + Ok(_) => debug!("No valid multi-cert directory structure found"), + Err(_) => debug!("load_all_certs_from_directory failed, trying single-cert fallback"), + } + + // Fallback: single cert + let key_path = format!("{tls_path}/{RUSTFS_TLS_KEY}"); + let cert_path = format!("{tls_path}/{RUSTFS_TLS_CERT}"); + if tokio::try_join!(tokio::fs::metadata(&key_path), tokio::fs::metadata(&cert_path)).is_ok() { + let certs = rustfs_utils::load_certs(&cert_path).map_err(|e| TlsMaterialError::Io(format!("load certs: {e}")))?; + let key = rustfs_utils::load_private_key(&key_path).map_err(|e| TlsMaterialError::Io(format!("load key: {e}")))?; + + let config = build_server_config(ServerCertSource::SingleCert { certs, key }, mtls_verifier)?; + info!("Created TLS acceptor with single certificate"); + let acceptor = Arc::new(TlsAcceptor::from(Arc::new(config))); + return Ok(Some(Arc::new(TlsAcceptorHolder::new(acceptor)))); + } + + debug!("No valid TLS certificates found, starting with HTTP"); + Ok(None) + } + + fn empty() -> Self { + Self { + outbound: OutboundTlsMaterial { + root_ca_pem: Vec::new(), + mtls_identity: None, + }, + has_server_certs: false, + } + } +} + +// ── Server Config Construction ── + +/// Certificate source for building a `ServerConfig`. +enum ServerCertSource { + /// Pre-built SNI resolver from multi-cert directory. + Resolver(Arc), + /// Single certificate/key pair. + SingleCert { + certs: Vec>, + key: PrivateKeyDer<'static>, + }, +} + +/// Build a `ServerConfig` with standardized ALPN, session cache, and key log settings. +/// +/// This is the single place for `ServerConfig` construction, used by both +/// initial startup and hot-reload. +fn build_server_config( + cert_source: ServerCertSource, + mtls_verifier: Option>, +) -> Result { + let mut config = match cert_source { + ServerCertSource::Resolver(resolver) => { + if let Some(verifier) = mtls_verifier { + rustls::ServerConfig::builder() + .with_client_cert_verifier(verifier) + .with_cert_resolver(resolver) + } else { + rustls::ServerConfig::builder() + .with_no_client_auth() + .with_cert_resolver(resolver) + } + } + ServerCertSource::SingleCert { certs, key } => { + if let Some(verifier) = mtls_verifier { + rustls::ServerConfig::builder() + .with_client_cert_verifier(verifier) + .with_single_cert(certs, key) + .map_err(|e| TlsMaterialError::Io(format!("configure single cert with mTLS: {e}")))? + } else { + rustls::ServerConfig::builder() + .with_no_client_auth() + .with_single_cert(certs, key) + .map_err(|e| TlsMaterialError::Io(format!("configure single cert: {e}")))? + } + } + }; + + config.alpn_protocols = vec![b"h2".to_vec(), b"http/1.1".to_vec(), b"http/1.0".to_vec()]; + config.session_storage = rustls::server::ServerSessionMemoryCache::new(10000); + + if rustfs_utils::tls_key_log() { + config.key_log = Arc::new(rustls::KeyLogFile::new()); + } + + Ok(config) +} + +// ── Outbound Material Loading ── + +/// Load root CA certificates and mTLS identity for outbound connections. +async fn load_outbound_material(tls_dir: &Path) -> Result { + let mut root_ca_pem = Vec::new(); + + // 1. Optional: load leaf certs as root CAs + if get_env_bool(ENV_TRUST_LEAF_CERT_AS_CA, DEFAULT_TRUST_LEAF_CERT_AS_CA) + && load_cert_file_by_name(tls_dir, RUSTFS_TLS_CERT, &mut root_ca_pem).await + { + info!("Loaded leaf certificate(s) as root CA as per RUSTFS_TRUST_LEAF_CERT_AS_CA"); + } + + // 2. Load public.crt and ca.crt + load_cert_file(&tls_dir.join(RUSTFS_PUBLIC_CERT), &mut root_ca_pem, "CA certificate").await; + load_cert_file(&tls_dir.join(RUSTFS_CA_CERT), &mut root_ca_pem, "CA certificate").await; + + // 3. Optional: load system root CAs + if get_env_bool(ENV_TRUST_SYSTEM_CA, DEFAULT_TRUST_SYSTEM_CA) { + let mut system_loaded = false; + for path in SYSTEM_CA_PATHS { + if load_cert_file(Path::new(path), &mut root_ca_pem, "system root certificates").await { + system_loaded = true; + info!("Loaded system root certificates from {}", path); + break; + } + } + if !system_loaded { + debug!("Could not find system root certificates in common locations."); + } + } else { + info!("Loading system root certificates disabled via RUSTFS_TRUST_SYSTEM_CA"); + } + + // 4. Load optional mTLS identity + let mtls_identity = load_mtls_identity(tls_dir).await?; + + Ok(OutboundTlsMaterial { + root_ca_pem, + mtls_identity, + }) +} + +/// Quick check whether server certificate files exist in the TLS directory. +async fn has_server_certificates(tls_path: &str) -> bool { + if tokio::fs::metadata(tls_path).await.is_err() { + return false; + } + // Check for multi-cert directory structure OR single cert files + if rustfs_utils::load_all_certs_from_directory(tls_path).is_ok_and(|p| !p.is_empty()) { + return true; + } + let key_path = format!("{tls_path}/{RUSTFS_TLS_KEY}"); + let cert_path = format!("{tls_path}/{RUSTFS_TLS_CERT}"); + tokio::try_join!(tokio::fs::metadata(&key_path), tokio::fs::metadata(&cert_path)).is_ok() +} + +/// Load mTLS client identity from the TLS directory. +async fn load_mtls_identity(tls_dir: &Path) -> Result, TlsMaterialError> { + let client_cert_path = match get_env_opt_str(ENV_MTLS_CLIENT_CERT) { + Some(p) => PathBuf::from(p), + None => tls_dir.join(RUSTFS_CLIENT_CERT_FILENAME), + }; + + let client_key_path = match get_env_opt_str(ENV_MTLS_CLIENT_KEY) { + Some(p) => PathBuf::from(p), + None => tls_dir.join(RUSTFS_CLIENT_KEY_FILENAME), + }; + + if !client_cert_path.exists() || !client_key_path.exists() { + info!( + "mTLS client identity not configured (missing {:?} and/or {:?}); proceeding with server-only TLS", + client_cert_path, client_key_path + ); + return Ok(None); + } + + let cert_pem = tokio::fs::read(&client_cert_path) + .await + .map_err(|e| TlsMaterialError::Io(format!("read client cert {client_cert_path:?}: {e}")))?; + let key_pem = tokio::fs::read(&client_key_path) + .await + .map_err(|e| TlsMaterialError::Io(format!("read client key {client_key_path:?}: {e}")))?; + + // Validate parse-ability + let mut reader = std::io::Cursor::new(&cert_pem); + if CertificateDer::pem_reader_iter(&mut reader).next().is_none() { + return Err(TlsMaterialError::Parse("no valid certificate in client cert PEM".into())); + } + let mut reader = std::io::Cursor::new(&key_pem); + PrivateKeyDer::from_pem_reader(&mut reader).map_err(|e| TlsMaterialError::Parse(format!("invalid client key PEM: {e}")))?; + + info!("Loaded mTLS client identity cert={:?} key={:?}", client_cert_path, client_key_path); + Ok(Some(MtlsIdentityPem { cert_pem, key_pem })) +} + +/// Load a single certificate file and append PEM data. +/// Returns true if the file was successfully loaded. +async fn load_cert_file(path: &Path, pem_data: &mut Vec, desc: &str) -> bool { + if tokio::fs::metadata(path).await.is_err() { + debug!("{} file not found at {:?}", desc, path); + return false; + } + match tokio::fs::read(path).await { + Ok(data) => { + pem_data.extend_from_slice(&data); + pem_data.push(b'\n'); + info!("Loaded {} from {:?}", desc, path); + true + } + Err(e) => { + debug!("Failed to read {} from {:?}: {}", desc, path, e); + false + } + } +} + +/// Search for and load certificate files matching `cert_name` in the directory +/// and one level of subdirectories. +/// Returns `true` if at least one matching file was loaded. +async fn load_cert_file_by_name(dir: &Path, cert_name: &str, pem_data: &mut Vec) -> bool { + let Ok(mut rd) = tokio::fs::read_dir(dir).await else { + debug!("Certificate directory not found: {}", dir.display()); + return false; + }; + + let mut loaded = false; + while let Ok(Some(entry)) = rd.next_entry().await { + let Ok(ft) = entry.file_type().await else { continue }; + + if ft.is_file() { + let fname = entry.file_name().to_string_lossy().to_string(); + if fname == cert_name && load_cert_file(&entry.path(), pem_data, "certificate").await { + loaded = true; + } + } else if ft.is_dir() { + // Only check direct subdirectories (one level deep) + if let Ok(mut sub_rd) = tokio::fs::read_dir(&entry.path()).await { + while let Ok(Some(sub_entry)) = sub_rd.next_entry().await { + if let Ok(sub_ft) = sub_entry.file_type().await + && sub_ft.is_file() + { + let fname = sub_entry.file_name().to_string_lossy().to_string(); + if fname == cert_name && load_cert_file(&sub_entry.path(), pem_data, "certificate").await { + loaded = true; + } + } + } + } + } + } + loaded +} + +/// Errors that can occur during TLS material loading. +#[derive(Debug)] +pub enum TlsMaterialError { + /// I/O error (file read, directory access). + Io(String), + /// PEM parsing error. + Parse(String), +} + +impl std::fmt::Display for TlsMaterialError { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + match self { + TlsMaterialError::Io(msg) => write!(f, "TLS material I/O error: {msg}"), + TlsMaterialError::Parse(msg) => write!(f, "TLS material parse error: {msg}"), + } + } +} + +impl std::error::Error for TlsMaterialError {} + +// ── TLS Handshake Error Classification ── + +/// Structured classification of TLS handshake failures. +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub(crate) enum TlsHandshakeFailureKind { + UnexpectedEof, + ProtocolVersion, + Certificate, + Alert, + Unknown, +} + +impl TlsHandshakeFailureKind { + /// Classify a TLS accept error into a structured failure kind. + pub(crate) fn classify(err_msg: &str) -> Self { + if err_msg.contains("unexpected EOF") || err_msg.contains("handshake eof") { + Self::UnexpectedEof + } else if err_msg.contains("protocol version") { + Self::ProtocolVersion + } else if err_msg.contains("certificate") || err_msg.contains("invalid peer certificate") { + Self::Certificate + } else if err_msg.contains("alert") { + Self::Alert + } else { + Self::Unknown + } + } + + /// Metric label string for Prometheus. + pub(crate) fn as_str(self) -> &'static str { + match self { + Self::UnexpectedEof => "UNEXPECTED_EOF", + Self::ProtocolVersion => "PROTOCOL_VERSION", + Self::Certificate => "CERTIFICATE", + Self::Alert => "ALERT", + Self::Unknown => "UNKNOWN", + } + } +} + +// ── TLS Acceptor Holder (for hot reload) ── + +/// Holds the current TLS acceptor and supports atomic swap for certificate rotation. +/// +/// Uses `RwLock` so that multiple readers (per-connection `get()` calls) +/// do not block each other. The write lock is held only briefly during swap. +pub(crate) struct TlsAcceptorHolder { + current: RwLock>, +} + +impl TlsAcceptorHolder { + pub(crate) fn new(acceptor: Arc) -> Self { + Self { + current: RwLock::new(acceptor), + } + } + + /// Get the current TLS acceptor for handling a new connection. + #[inline] + pub(crate) fn get(&self) -> Arc { + match self.current.read() { + Ok(guard) => guard.clone(), + Err(poisoned) => poisoned.into_inner().clone(), + } + } + + /// Atomically replace the TLS acceptor with a new one. + fn swap(&self, new_holder: &TlsAcceptorHolder) { + let new_acceptor = new_holder.get(); + match self.current.write() { + Ok(mut guard) => *guard = new_acceptor, + Err(poisoned) => { + let mut guard = poisoned.into_inner(); + *guard = new_acceptor; + } + } + } +} + +/// Spawn a background task that periodically checks for TLS certificate changes. +pub(crate) fn spawn_reload_loop(tls_path: String, holder: Arc) { + let enabled = get_env_bool(ENV_TLS_RELOAD_ENABLE, DEFAULT_TLS_RELOAD_ENABLE); + if !enabled { + debug!("TLS certificate hot reload is disabled (set {}=1 to enable)", ENV_TLS_RELOAD_ENABLE); + return; + } + + let interval_secs = rustfs_utils::get_env_u64(ENV_TLS_RELOAD_INTERVAL, DEFAULT_TLS_RELOAD_INTERVAL).max(5); + + info!("TLS certificate hot reload enabled, checking every {}s", interval_secs); + + tokio::spawn(async move { + let mut interval = tokio::time::interval(Duration::from_secs(interval_secs)); + loop { + interval.tick().await; + + match TlsMaterialSnapshot::load(&tls_path).await { + Ok(snapshot) => { + // Always refresh outbound material (root CAs, mTLS identity) on reload. + snapshot.apply_outbound().await; + + match snapshot.build_tls_acceptor(&tls_path).await { + Ok(Some(new_holder)) => { + info!("TLS certificates reloaded successfully"); + holder.swap(&new_holder); + } + Ok(None) => debug!("TLS reload: no server certificates found in directory, skipping"), + Err(e) => warn!("TLS certificate reload failed (will retry): {}", e), + } + } + Err(e) => { + warn!("TLS material reload failed (will retry): {}", e); + } + } + } + }); +} From a9be9af094f6221d23c394e790782f2c13704033 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E5=AE=89=E6=AD=A3=E8=B6=85?= Date: Sat, 4 Apr 2026 08:35:57 +0800 Subject: [PATCH 12/22] ci: bump cla-bot to v0.0.9 (#2389) --- .github/workflows/cla.yml | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/.github/workflows/cla.yml b/.github/workflows/cla.yml index 8f5b1f44c..3c0c23ee9 100644 --- a/.github/workflows/cla.yml +++ b/.github/workflows/cla.yml @@ -42,7 +42,7 @@ jobs: permission-contents: write - name: Run CLA Bot - uses: overtrue/cla-bot@v0.0.8 + uses: overtrue/cla-bot@v0.0.9 with: github-token: ${{ github.token }} registry-token: ${{ steps.registry-token.outputs.token }} From 67863630b2367fef30b9b3bc9acf263f958fb740 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E5=AE=89=E6=AD=A3=E8=B6=85?= Date: Sat, 4 Apr 2026 08:36:14 +0800 Subject: [PATCH 13/22] fix(auth): reject ambiguous case-insensitive claim matches (#2386) --- crates/iam/src/oidc.rs | 42 ++++++++++------- crates/policy/src/policy.rs | 1 + crates/policy/src/policy/policy.rs | 73 ++++++++++++++++++------------ crates/policy/src/policy/utils.rs | 71 ++++++++++++++++++++++++++++- 4 files changed, 142 insertions(+), 45 deletions(-) diff --git a/crates/iam/src/oidc.rs b/crates/iam/src/oidc.rs index cac4c7bb6..923e99c99 100644 --- a/crates/iam/src/oidc.rs +++ b/crates/iam/src/oidc.rs @@ -27,6 +27,7 @@ use openidconnect::{ use rustfs_config::oidc::*; use rustfs_config::{DEFAULT_DELIMITER, ENABLE_KEY, EnableState}; use rustfs_ecstore::config::{Config as ServerConfig, KVS, get_global_server_config}; +use rustfs_policy::policy::{ClaimLookup, get_claim_case_insensitive}; use serde::{Deserialize, Serialize}; use std::borrow::Cow; use std::collections::HashMap; @@ -1007,29 +1008,19 @@ pub(crate) fn decode_jwt_payload(token: &str) -> HashMap(claims: &'a HashMap, key: &str) -> Option<&'a serde_json::Value> { - if let Some(v) = claims.get(key) { - return Some(v); - } - let key_lower = key.to_lowercase(); - claims.iter().find(|(k, _)| k.to_lowercase() == key_lower).map(|(_, v)| v) -} - /// Extract a string claim from raw claims with case-insensitive fallback. fn extract_string_claim(claims: &HashMap, key: &str) -> String { - get_claim_case_insensitive(claims, key) - .and_then(serde_json::Value::as_str) - .unwrap_or_default() - .to_string() + match get_claim_case_insensitive(claims, key) { + ClaimLookup::Found(value) => value.as_str().unwrap_or_default().to_string(), + ClaimLookup::Missing | ClaimLookup::Ambiguous => String::new(), + } } /// Extract a groups/array claim from raw claims with case-insensitive fallback. Handles both string arrays and single strings. fn extract_groups_claim(claims: &HashMap, key: &str) -> Vec { match get_claim_case_insensitive(claims, key) { - Some(serde_json::Value::Array(arr)) => arr.iter().filter_map(|v| v.as_str().map(String::from)).collect(), - Some(serde_json::Value::String(s)) => s.split(',').map(|s| s.trim().to_string()).collect(), + ClaimLookup::Found(serde_json::Value::Array(arr)) => arr.iter().filter_map(|v| v.as_str().map(String::from)).collect(), + ClaimLookup::Found(serde_json::Value::String(s)) => s.split(',').map(|s| s.trim().to_string()).collect(), _ => vec![], } } @@ -1117,6 +1108,25 @@ mod tests { assert_eq!(groups, vec!["exact_match"]); } + #[test] + fn test_extract_string_claim_ambiguous_case_insensitive_match_returns_empty() { + let mut claims = HashMap::new(); + claims.insert("Policy".to_string(), serde_json::json!("exact_match")); + claims.insert("policy".to_string(), serde_json::json!("lowercase")); + + assert_eq!(extract_string_claim(&claims, "POLICY"), ""); + } + + #[test] + fn test_extract_groups_claim_ambiguous_case_insensitive_match_returns_empty() { + let mut claims = HashMap::new(); + claims.insert("Policy".to_string(), serde_json::json!(["exact_match"])); + claims.insert("policy".to_string(), serde_json::json!(["lowercase"])); + + let groups = extract_groups_claim(&claims, "POLICY"); + assert!(groups.is_empty()); + } + #[test] fn test_decode_jwt_payload() { let payload = r#"{"sub":"user123","email":"user@example.com"}"#; diff --git a/crates/policy/src/policy.rs b/crates/policy/src/policy.rs index 92329983f..756f44a2e 100644 --- a/crates/policy/src/policy.rs +++ b/crates/policy/src/policy.rs @@ -35,6 +35,7 @@ pub use policy::*; pub use principal::Principal; pub use resource::ResourceSet; pub use statement::Statement; +pub use utils::{ClaimLookup, get_claim_case_insensitive}; #[derive(thiserror::Error, Debug)] #[cfg_attr(test, derive(Eq, PartialEq))] diff --git a/crates/policy/src/policy/policy.rs b/crates/policy/src/policy/policy.rs index 88201e66b..1176b862b 100644 --- a/crates/policy/src/policy/policy.rs +++ b/crates/policy/src/policy/policy.rs @@ -13,8 +13,8 @@ // limitations under the License. use super::{ - Effect, Error as IamError, Functions, ID, Statement, action::Action, statement::BPStatement, - statement::variable_resolver_for_policy_args, + ClaimLookup, Effect, Error as IamError, Functions, ID, Statement, action::Action, get_claim_case_insensitive, + statement::BPStatement, statement::variable_resolver_for_policy_args, }; use crate::error::{Error, Result}; use serde::{Deserialize, Serialize}; @@ -239,42 +239,37 @@ impl Validator for BucketPolicy { } } -fn get_claim_case_insensitive<'a>(claims: &'a HashMap, claim_name: &str) -> Option<&'a Value> { - if let Some(v) = claims.get(claim_name) { - return Some(v); - } - let claim_name_lower = claim_name.to_lowercase(); - claims - .iter() - .find(|(k, _)| k.to_lowercase() == claim_name_lower) - .map(|(_, v)| v) -} - fn get_values_from_claims(claims: &HashMap, claim_name: &str) -> (HashSet, bool) { let mut s = HashSet::new(); - if let Some(pname) = get_claim_case_insensitive(claims, claim_name) { - if let Some(pnames) = pname.as_array() { - for pname in pnames { - if let Some(pname_str) = pname.as_str() { - for pname in pname_str.split(',') { - let pname = pname.trim(); - if !pname.is_empty() { - s.insert(pname.to_string()); + match get_claim_case_insensitive(claims, claim_name) { + ClaimLookup::Found(pname) => { + if let Some(pnames) = pname.as_array() { + for pname in pnames { + if let Some(pname_str) = pname.as_str() { + for pname in pname_str.split(',') { + let pname = pname.trim(); + if !pname.is_empty() { + s.insert(pname.to_string()); + } } } } + return (s, true); } - return (s, true); - } else if let Some(pname_str) = pname.as_str() { - for pname in pname_str.split(',') { - let pname = pname.trim(); - if !pname.is_empty() { - s.insert(pname.to_string()); + + if let Some(pname_str) = pname.as_str() { + for pname in pname_str.split(',') { + let pname = pname.trim(); + if !pname.is_empty() { + s.insert(pname.to_string()); + } } + return (s, true); } - return (s, true); } + ClaimLookup::Missing | ClaimLookup::Ambiguous => {} } + (s, false) } @@ -1744,4 +1739,26 @@ mod test { assert!(policies.contains("consoleAdmin")); assert!(policies.contains("readwrite")); } + + #[test] + fn test_get_values_from_claims_ambiguous_case_insensitive_match_returns_missing() { + let mut claims = HashMap::new(); + claims.insert("Policy".to_string(), Value::Array(vec![Value::String("exact_match".to_string())])); + claims.insert("policy".to_string(), Value::Array(vec![Value::String("lowercase".to_string())])); + + let (policies, found) = get_values_from_claims(&claims, "POLICY"); + assert!(!found); + assert!(policies.is_empty()); + } + + #[test] + fn test_get_policies_from_claims_ambiguous_case_insensitive_match_returns_missing() { + let mut claims = HashMap::new(); + claims.insert("Policy".to_string(), Value::String("consoleAdmin".to_string())); + claims.insert("policy".to_string(), Value::String("readwrite".to_string())); + + let (policies, found) = get_policies_from_claims(&claims, "POLICY"); + assert!(!found); + assert!(policies.is_empty()); + } } diff --git a/crates/policy/src/policy/utils.rs b/crates/policy/src/policy/utils.rs index 832f6f5a6..32e95c858 100644 --- a/crates/policy/src/policy/utils.rs +++ b/crates/policy/src/policy/utils.rs @@ -19,6 +19,41 @@ use serde_json::Value; pub mod path; pub mod wildcard; +#[derive(Clone, Copy, Debug, Eq, PartialEq)] +pub enum ClaimLookup<'a> { + Missing, + Found(&'a Value), + Ambiguous, +} + +fn case_insensitive_eq(left: &str, right: &str) -> bool { + left.chars() + .flat_map(char::to_lowercase) + .eq(right.chars().flat_map(char::to_lowercase)) +} + +pub fn get_claim_case_insensitive<'a>(claims: &'a HashMap, claim_name: &str) -> ClaimLookup<'a> { + if let Some(value) = claims.get(claim_name) { + return ClaimLookup::Found(value); + } + + let mut matched = None; + + for (candidate, value) in claims { + if case_insensitive_eq(candidate, claim_name) { + if matched.is_some() { + return ClaimLookup::Ambiguous; + } + matched = Some(value); + } + } + + match matched { + Some(value) => ClaimLookup::Found(value), + None => ClaimLookup::Missing, + } +} + pub fn _get_values_from_claims(claim: &HashMap, chaim_name: &str) -> (Vec, bool) { let mut result = vec![]; let Some(pname) = claim.get(chaim_name) else { @@ -77,7 +112,9 @@ pub fn _split_path(path: &str, second_index: bool) -> (&str, &str) { #[cfg(test)] mod tests { - use super::_split_path; + use super::{_split_path, ClaimLookup, get_claim_case_insensitive}; + use serde_json::{Value, json}; + use std::collections::HashMap; #[test_case::test_case("format.json", false => ("format.json", ""))] #[test_case::test_case("users/tester.json", false => ("users/", "tester.json"))] @@ -98,4 +135,36 @@ mod tests { fn test_split_path(path: &str, second_index: bool) -> (&str, &str) { _split_path(path, second_index) } + + #[test] + fn test_get_claim_case_insensitive_prefers_exact_match() { + let mut claims = HashMap::new(); + claims.insert("Policy".to_string(), json!("exact_match")); + claims.insert("policy".to_string(), json!("lowercase")); + + assert_eq!( + get_claim_case_insensitive(&claims, "Policy"), + ClaimLookup::Found(&Value::String("exact_match".to_string())) + ); + } + + #[test] + fn test_get_claim_case_insensitive_returns_ambiguous_for_multiple_folded_matches() { + let mut claims = HashMap::new(); + claims.insert("Policy".to_string(), json!("exact_match")); + claims.insert("policy".to_string(), json!("lowercase")); + + assert_eq!(get_claim_case_insensitive(&claims, "POLICY"), ClaimLookup::Ambiguous); + } + + #[test] + fn test_get_claim_case_insensitive_matches_unicode_without_allocation() { + let mut claims = HashMap::new(); + claims.insert("Straße".to_string(), json!("value")); + + assert_eq!( + get_claim_case_insensitive(&claims, "straße"), + ClaimLookup::Found(&Value::String("value".to_string())) + ); + } } From d2901fd78ce481dfebbf5364043bb54ad79989a8 Mon Sep 17 00:00:00 2001 From: houseme Date: Sat, 4 Apr 2026 09:07:22 +0800 Subject: [PATCH 14/22] feat(admin): add audit target APIs and harden target source handling (#2350) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Signed-off-by: houseme Co-authored-by: 安正超 Co-authored-by: copilot-swe-agent[bot] <198982749+Copilot@users.noreply.github.com> Co-authored-by: houseme <4829346+houseme@users.noreply.github.com> --- Cargo.lock | 3 + crates/audit/Cargo.toml | 2 + crates/audit/src/registry.rs | 96 +- crates/audit/src/system.rs | 267 +++--- crates/audit/tests/integration_test.rs | 49 +- crates/config/src/constants/targets.rs | 7 + crates/e2e_test/src/object_lambda_test.rs | 139 ++- crates/ecstore/src/config/audit.rs | 19 +- crates/ecstore/src/config/com.rs | 788 ++++++++++++++++- crates/ecstore/src/config/notify.rs | 13 +- crates/filemeta/Cargo.toml | 1 + crates/filemeta/src/filemeta.rs | 11 +- crates/notify/src/integration.rs | 202 +++-- crates/notify/src/notifier.rs | 6 +- crates/notify/src/registry.rs | 98 +-- crates/notify/src/stream.rs | 117 ++- crates/targets/Cargo.toml | 7 + .../targets/benches/queue_store_benchmark.rs | 94 ++ crates/targets/src/store.rs | 472 +++++++--- crates/targets/src/target/mod.rs | 179 +++- crates/targets/src/target/mqtt.rs | 126 ++- crates/targets/src/target/webhook.rs | 199 +++-- rustfs/src/admin/handlers/audit.rs | 818 ++++++++++++++++++ rustfs/src/admin/handlers/event.rs | 632 ++++++++++++-- rustfs/src/admin/handlers/mod.rs | 2 + rustfs/src/admin/mod.rs | 5 +- rustfs/src/admin/route_registration_test.rs | 10 +- rustfs/src/server/audit.rs | 26 +- scripts/run.sh | 2 +- 29 files changed, 3534 insertions(+), 856 deletions(-) create mode 100644 crates/targets/benches/queue_store_benchmark.rs create mode 100644 rustfs/src/admin/handlers/audit.rs diff --git a/Cargo.lock b/Cargo.lock index e51a91424..a74a30c29 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -7703,6 +7703,7 @@ dependencies = [ "rustfs-targets", "serde", "serde_json", + "temp-env", "thiserror 2.0.18", "tokio", "tracing", @@ -7897,6 +7898,7 @@ dependencies = [ "rustfs-utils", "s3s", "serde", + "tempfile", "thiserror 2.0.18", "time", "tokio", @@ -8396,6 +8398,7 @@ name = "rustfs-targets" version = "0.0.5" dependencies = [ "async-trait", + "criterion", "reqwest 0.13.2", "rumqttc", "rustfs-config", diff --git a/crates/audit/Cargo.toml b/crates/audit/Cargo.toml index ec904e1d6..0e0b03c09 100644 --- a/crates/audit/Cargo.toml +++ b/crates/audit/Cargo.toml @@ -44,6 +44,8 @@ tracing = { workspace = true, features = ["std", "attributes"] } url = { workspace = true } rumqttc = { workspace = true } +[dev-dependencies] +temp-env = { workspace = true } [lints] workspace = true diff --git a/crates/audit/src/registry.rs b/crates/audit/src/registry.rs index c6eccae57..f15bdba87 100644 --- a/crates/audit/src/registry.rs +++ b/crates/audit/src/registry.rs @@ -109,9 +109,6 @@ impl AuditRegistry { let all_env: Vec<(String, String)> = std::env::vars().filter(|(key, _)| key.starts_with(ENV_PREFIX)).collect(); // A collection of asynchronous tasks for concurrently executing target creation let mut tasks = FuturesUnordered::new(); - // let final_config = config.clone(); // Clone a configuration for aggregating the final result - // Record the defaults for each segment so that the segment can eventually be rebuilt - let mut section_defaults: HashMap = HashMap::new(); // 1. Traverse all registered plants and process them by target type for (target_type, factory) in &self.factories { tracing::Span::current().record("target_type", target_type.as_str()); @@ -125,9 +122,6 @@ impl AuditRegistry { let default_cfg = file_configs.get(DEFAULT_DELIMITER).cloned().unwrap_or_default(); debug!(?default_cfg, "Get the default configuration"); - // Save defaults for eventual write back - section_defaults.insert(section_name.clone(), default_cfg.clone()); - // *** Optimization point 1: Get all legitimate fields of the current target type *** let valid_fields = factory.get_valid_fields(); debug!(?valid_fields, "Get the legitimate configuration fields"); @@ -235,103 +229,28 @@ impl AuditRegistry { if enabled { info!(instance_id = %id, "Target is enabled, ready to create a task"); // 5.3. Create asynchronous tasks for enabled instances - let target_type_clone = target_type.clone(); let tid = id.clone(); let merged_config_arc = Arc::new(merged_config); tasks.push(async move { let result = factory.create_target(tid.clone(), &merged_config_arc).await; - (target_type_clone, tid, result, Arc::clone(&merged_config_arc)) + (tid, result) }); } else { - info!(instance_id = %id, "Skip the disabled target and will be removed from the final configuration"); - // Remove disabled target from final configuration - // final_config.0.entry(section_name.clone()).or_default().remove(&id); + info!(instance_id = %id, "Skip disabled target"); } } } // 6. Concurrently execute all creation tasks and collect results let mut successful_targets = Vec::new(); - let mut successful_configs = Vec::new(); - while let Some((target_type, id, result, final_config)) = tasks.next().await { + while let Some((id, result)) = tasks.next().await { match result { Ok(target) => { - info!(target_type = %target_type, instance_id = %id, "Create a target successfully"); + info!(target_type = %target.id().name, instance_id = %id, "Create a target successfully"); successful_targets.push(target); - successful_configs.push((target_type, id, final_config)); } Err(e) => { - error!(target_type = %target_type, instance_id = %id, error = %e, "Failed to create a target"); - } - } - } - - // 7. Aggregate new configuration and write back to system configuration - if !successful_configs.is_empty() || !section_defaults.is_empty() { - info!( - "Prepare to update {} successfully created target configurations to the system configuration...", - successful_configs.len() - ); - - let mut successes_by_section: HashMap> = HashMap::new(); - - for (target_type, id, kvs) in successful_configs { - let section_name = format!("{AUDIT_ROUTE_PREFIX}{target_type}").to_lowercase(); - successes_by_section - .entry(section_name) - .or_default() - .insert(id.to_lowercase(), (*kvs).clone()); - } - - let mut new_config = config.clone(); - // Collection of segments that need to be processed: Collect all segments where default items exist or where successful instances exist - let mut sections: HashSet = HashSet::new(); - sections.extend(section_defaults.keys().cloned()); - sections.extend(successes_by_section.keys().cloned()); - - for section in sections { - let mut section_map: std::collections::HashMap = std::collections::HashMap::new(); - // Add default item - if let Some(default_kvs) = section_defaults.get(§ion) - && !default_kvs.is_empty() - { - section_map.insert(DEFAULT_DELIMITER.to_string(), default_kvs.clone()); - } - - // Add successful instance item - if let Some(instances) = successes_by_section.get(§ion) { - for (id, kvs) in instances { - section_map.insert(id.clone(), kvs.clone()); - } - } - - // Empty breaks are removed and non-empty breaks are replaced entirely. - if section_map.is_empty() { - new_config.0.remove(§ion); - } else { - new_config.0.insert(section, section_map); - } - } - - if &new_config == config { - info!("Audit target configuration unchanged, skip persisting server config"); - info!(count = successful_targets.len(), "All target processing completed"); - return Ok(successful_targets); - } - - let Some(store) = rustfs_ecstore::global::new_object_layer_fn() else { - return Err(AuditError::StorageNotAvailable( - "Failed to save target configuration: server storage not initialized".to_string(), - )); - }; - - match rustfs_ecstore::config::com::save_server_config(store, &new_config).await { - Ok(_) => { - info!("The new configuration was saved to the system successfully.") - } - Err(e) => { - error!("Failed to save the new configuration: {}", e); - return Err(AuditError::SaveConfig(Box::new(e))); + error!(instance_id = %id, error = %e, "Failed to create a target"); } } } @@ -371,6 +290,11 @@ impl AuditRegistry { self.targets.get(id).map(|t| t.as_ref()) } + /// Lists cloned target values for runtime inspection without exposing mutable registry access. + pub fn list_target_values(&self) -> Vec + Send + Sync>> { + self.targets.values().map(|target| target.clone_dyn()).collect() + } + /// Lists all target IDs /// /// # Returns diff --git a/crates/audit/src/system.rs b/crates/audit/src/system.rs index d9116b659..d47729722 100644 --- a/crates/audit/src/system.rs +++ b/crates/audit/src/system.rs @@ -13,14 +13,15 @@ // limitations under the License. use crate::{AuditEntry, AuditError, AuditRegistry, AuditResult, observability}; +use hashbrown::HashMap; use rustfs_ecstore::config::Config; use rustfs_targets::{ StoreError, Target, TargetError, store::{Key, Store}, - target::EntityTarget, + target::{EntityTarget, QueuedPayload}, }; use std::sync::Arc; -use tokio::sync::{Mutex, RwLock}; +use tokio::sync::{Mutex, RwLock, mpsc}; use tracing::{error, info, warn}; /// State of the audit system @@ -39,6 +40,8 @@ pub struct AuditSystem { registry: Arc>, state: Arc>, config: Arc>>, + /// Cancellation senders for active audit stream tasks (target_id -> cancel tx) + stream_cancellers: Arc>>>, } impl Default for AuditSystem { @@ -54,6 +57,7 @@ impl AuditSystem { registry: Arc::new(Mutex::new(AuditRegistry::new())), state: Arc::new(RwLock::new(AuditSystemState::Stopped)), config: Arc::new(RwLock::new(None)), + stream_cancellers: Arc::new(RwLock::new(HashMap::new())), } } @@ -110,27 +114,7 @@ impl AuditSystem { // Initialize all targets for target in targets { - let target_id = target.id().to_string(); - if let Err(e) = target.init().await { - error!(target_id = %target_id, error = %e, "Failed to initialize audit target"); - } else { - // After successful initialization, if enabled and there is a store, start the send from storage task - if target.is_enabled() { - if let Some(store) = target.store() { - info!(target_id = %target_id, "Start audit stream processing for target"); - let store_clone: Box, Error = StoreError, Key = Key> + Send> = - store.boxed_clone(); - let target_arc: Arc + Send + Sync> = Arc::from(target.clone_dyn()); - self.start_audit_stream_with_batching(store_clone, target_arc); - info!(target_id = %target_id, "Audit stream processing started"); - } else { - info!(target_id = %target_id, "No store configured, skip audit stream processing"); - } - } else { - info!(target_id = %target_id, "Target disabled, skip audit stream processing"); - } - registry.add_target(target_id, target); - } + self.init_and_register_target(target, &mut registry).await; } // Update state to running @@ -214,6 +198,9 @@ impl AuditSystem { info!("Stopping audit system"); + // Stop all stream tasks first + self.stop_all_streams().await; + // Close all targets let mut registry = self.registry.lock().await; if let Err(e) = registry.close_all().await { @@ -258,52 +245,49 @@ impl AuditSystem { let state = self.state.read().await; match *state { - AuditSystemState::Running => { - // Continue with dispatch - info!("Dispatching audit log entry"); - } + AuditSystemState::Running => {} AuditSystemState::Paused => { - // Skip dispatch when paused return Ok(()); } _ => { - // Don't dispatch when not running return Err(AuditError::NotInitialized("Audit system is not running".to_string())); } } drop(state); - let registry = self.registry.lock().await; - let target_keys = registry.list_targets(); + // Collect cloned targets under lock, then dispatch without holding it + let targets: Vec<(String, Box + Send + Sync>)> = { + let registry = self.registry.lock().await; + let target_keys = registry.list_targets(); - if target_keys.is_empty() { - warn!("No audit targets configured for dispatch"); - return Ok(()); - } + if target_keys.is_empty() { + warn!("No audit targets configured for dispatch"); + return Ok(()); + } - // Dispatch to all targets concurrently + target_keys + .into_iter() + .filter_map(|key| registry.get_target(&key).map(|t| (key, t.clone_dyn()))) + .collect() + }; + + // Dispatch to all targets concurrently (no lock held) let mut tasks = Vec::new(); - for target_key in target_keys { - if let Some(target) = registry.get_target(&target_key) { - let entry_clone = Arc::clone(&entry); - let target_key_clone = target_key.clone(); + for (target_key, target) in targets { + let entity_target = EntityTarget { + object_name: entry.api.name.clone().unwrap_or_default(), + bucket_name: entry.api.bucket.clone().unwrap_or_default(), + event_name: entry.event, + data: (*entry).clone(), + }; - // Create EntityTarget for the audit log entry - let entity_target = EntityTarget { - object_name: entry.api.name.clone().unwrap_or_default(), - bucket_name: entry.api.bucket.clone().unwrap_or_default(), - event_name: entry.event, // Default, should be derived from entry - data: (*entry_clone).clone(), - }; + let task = async move { + let result = target.save(Arc::new(entity_target)).await; + (target_key, result) + }; - let task = async move { - let result = target.save(Arc::new(entity_target)).await; - (target_key_clone, result) - }; - - tasks.push(task); - } + tasks.push(task); } // Execute all dispatch tasks @@ -359,39 +343,45 @@ impl AuditSystem { } drop(state); - let registry = self.registry.lock().await; - let target_keys = registry.list_targets(); + // Collect targets under lock, then dispatch without holding it + let targets: Vec<(String, Box + Send + Sync>)> = { + let registry = self.registry.lock().await; + let target_keys = registry.list_targets(); - if target_keys.is_empty() { - warn!("No audit targets configured for batch dispatch"); - return Ok(()); - } + if target_keys.is_empty() { + warn!("No audit targets configured for batch dispatch"); + return Ok(()); + } + + target_keys + .into_iter() + .filter_map(|key| registry.get_target(&key).map(|t| (key, t.clone_dyn()))) + .collect() + }; let mut tasks = Vec::new(); - for target_key in target_keys { - if let Some(target) = registry.get_target(&target_key) { - let entries_clone: Vec<_> = entries.iter().map(Arc::clone).collect(); - let target_key_clone = target_key.clone(); + for (target_key, target) in targets { + let entries_clone: Vec<_> = entries.iter().map(Arc::clone).collect(); + let target_key_clone = target_key.clone(); - let task = async move { - let mut success_count = 0; - let mut errors = Vec::new(); - for entry in entries_clone { - let entity_target = EntityTarget { - object_name: entry.api.name.clone().unwrap_or_default(), - bucket_name: entry.api.bucket.clone().unwrap_or_default(), - event_name: entry.event, - data: (*entry).clone(), - }; - match target.save(Arc::new(entity_target)).await { - Ok(_) => success_count += 1, - Err(e) => errors.push(e), - } + let task = async move { + let mut success_count = 0; + let mut errors = Vec::new(); + for entry in entries_clone { + let entity_target = EntityTarget { + object_name: entry.api.name.clone().unwrap_or_default(), + bucket_name: entry.api.bucket.clone().unwrap_or_default(), + event_name: entry.event, + data: (*entry).clone(), + }; + match target.save(Arc::new(entity_target)).await { + Ok(_) => success_count += 1, + Err(e) => errors.push(e), } - (target_key_clone, success_count, errors) - }; - tasks.push(task); - } + } + (target_key_clone, success_count, errors) + }; + tasks.push(task); } let results = futures::future::join_all(tasks).await; @@ -417,6 +407,60 @@ impl AuditSystem { Ok(()) } + /// Stops all active audit stream tasks by sending cancellation signals. + async fn stop_all_streams(&self) { + let mut cancellers = self.stream_cancellers.write().await; + for (target_id, cancel_tx) in cancellers.drain() { + info!(target_id = %target_id, "Stopping audit stream"); + let _ = cancel_tx.send(()).await; + } + } + + /// Initializes a single target: runs init(), starts stream if store is present, + /// and adds it to the registry. For store-backed targets, registration and stream + /// startup proceed even if init() fails so queued entries can be drained later. + async fn init_and_register_target( + &self, + target: Box + Send + Sync>, + registry: &mut AuditRegistry, + ) -> Option { + let target_id = target.id().to_string(); + let has_store = target.store().is_some(); + + if let Err(e) = target.init().await { + error!(target_id = %target_id, error = %e, "Failed to initialize audit target"); + // Non-store targets: init failure is fatal. + if !has_store { + return None; + } + // Store-backed targets: still register and start the stream so queued + // entries can be drained when connectivity recovers. + warn!( + target_id = %target_id, + "Proceeding with store-backed audit target despite init failure" + ); + } + + if target.is_enabled() { + if let Some(store) = target.store() { + info!(target_id = %target_id, "Start audit stream processing for target"); + let store_clone: Box + Send> = store.boxed_clone(); + let target_arc: Arc + Send + Sync> = Arc::from(target.clone_dyn()); + let cancel_tx = self.start_audit_stream_with_batching(store_clone, target_arc); + + self.stream_cancellers.write().await.insert(target_id.clone(), cancel_tx); + info!(target_id = %target_id, "Audit stream processing started"); + } else { + info!(target_id = %target_id, "No store configured, skip audit stream processing"); + } + } else { + info!(target_id = %target_id, "Target disabled, skip audit stream processing"); + } + + registry.add_target(target_id.clone(), target); + Some(target_id) + } + /// Starts the audit stream processing for a target with batching and retry logic /// /// # Arguments @@ -427,9 +471,10 @@ impl AuditSystem { /// and attempts to send them to the specified target. It implements retry logic with exponential backoff fn start_audit_stream_with_batching( &self, - store: Box, Error = StoreError, Key = Key> + Send>, + store: Box + Send>, target: Arc + Send + Sync>, - ) { + ) -> mpsc::Sender<()> { + let (cancel_tx, mut cancel_rx) = mpsc::channel(1); let state = self.state.clone(); tokio::spawn(async move { @@ -442,6 +487,12 @@ impl AuditSystem { const BASE_RETRY_DELAY: Duration = Duration::from_secs(2); loop { + // Check for cancellation signal + if cancel_rx.try_recv().is_ok() { + info!("Audit stream cancelled for target: {}", target.id()); + break; + } + match *state.read().await { AuditSystemState::Running | AuditSystemState::Paused | AuditSystemState::Starting => {} _ => { @@ -452,11 +503,22 @@ impl AuditSystem { let keys: Vec = store.list(); if keys.is_empty() { - sleep(Duration::from_millis(500)).await; + tokio::select! { + _ = sleep(Duration::from_millis(500)) => {}, + _ = cancel_rx.recv() => { + info!("Audit stream cancelled during idle for target: {}", target.id()); + return; + } + } continue; } for key in keys { + if cancel_rx.try_recv().is_ok() { + info!("Audit stream cancelled during processing for target: {}", target.id()); + return; + } + let mut retries = 0usize; let mut success = false; @@ -497,6 +559,8 @@ impl AuditSystem { sleep(Duration::from_millis(100)).await; } }); + + cancel_tx } /// Enables a specific target @@ -594,6 +658,12 @@ impl AuditSystem { registry.list_targets() } + /// Returns cloned target values for read-only runtime inspection. + pub async fn get_target_values(&self) -> Vec + Send + Sync>> { + let registry = self.registry.lock().await; + registry.list_target_values() + } + /// Gets information about a specific target /// /// # Arguments @@ -616,9 +686,11 @@ impl AuditSystem { pub async fn reload_config(&self, new_config: Config) -> AuditResult<()> { info!("Reloading audit system configuration"); - // Record config reload observability::record_config_reload(); + // Stop all existing stream tasks first + self.stop_all_streams().await; + // Store new configuration { let mut config_guard = self.config.write().await; @@ -636,28 +708,9 @@ impl AuditSystem { Ok(targets) => { info!(target_count = targets.len(), "Reloaded audit targets successfully"); - // Initialize all new targets for target in targets { - let target_id = target.id().to_string(); - if let Err(e) = target.init().await { - error!(target_id = %target_id, error = %e, "Failed to initialize reloaded audit target"); - } else { - // Same starts the storage stream after a heavy load - if target.is_enabled() { - if let Some(store) = target.store() { - info!(target_id = %target_id, "Start audit stream processing for target (reload)"); - let store_clone: Box, Error = StoreError, Key = Key> + Send> = - store.boxed_clone(); - let target_arc: Arc + Send + Sync> = Arc::from(target.clone_dyn()); - self.start_audit_stream_with_batching(store_clone, target_arc); - info!(target_id = %target_id, "Audit stream processing started (reload)"); - } else { - info!(target_id = %target_id, "No store configured, skip audit stream processing (reload)"); - } - } else { - info!(target_id = %target_id, "Target disabled, skip audit stream processing (reload)"); - } - registry.add_target(target.id().to_string(), target); + if let Some(target_id) = self.init_and_register_target(target, &mut registry).await { + info!(target_id = %target_id, "Target initialized (reload)"); } } diff --git a/crates/audit/tests/integration_test.rs b/crates/audit/tests/integration_test.rs index f2ef342e1..08f5f51e2 100644 --- a/crates/audit/tests/integration_test.rs +++ b/crates/audit/tests/integration_test.rs @@ -15,6 +15,7 @@ use rustfs_audit::*; use rustfs_ecstore::config::{Config, KVS}; use std::collections::HashMap; +use temp_env::with_vars; #[tokio::test] async fn test_audit_system_creation() { @@ -35,34 +36,42 @@ async fn test_config_parsing_webhook() { let mut config = Config(HashMap::new()); let mut audit_webhook_section = HashMap::new(); - // Create default configuration let mut default_kvs = KVS::new(); - default_kvs.insert("enable".to_string(), "on".to_string()); - default_kvs.insert("endpoint".to_string(), "http://localhost:3020/webhook".to_string()); - + default_kvs.insert("enable".to_string(), "off".to_string()); + default_kvs.insert("endpoint".to_string(), "".to_string()); audit_webhook_section.insert("_".to_string(), default_kvs); + let mut instance_kvs = KVS::new(); + instance_kvs.insert("enable".to_string(), "on".to_string()); + instance_kvs.insert("endpoint".to_string(), "http://localhost:3020/webhook".to_string()); + audit_webhook_section.insert("primary".to_string(), instance_kvs); config.0.insert("audit_webhook".to_string(), audit_webhook_section); let registry = AuditRegistry::new(); - // This should not fail even if server storage is not initialized - // as it's an integration test let result = registry.create_audit_targets_from_config(&config).await; + assert!(result.is_ok(), "audit target creation should not require server storage"); +} - // We expect this to fail due to server storage not being initialized - // but the parsing should work correctly - match result { - Err(AuditError::StorageNotAvailable(_)) => { - // This is expected in test environment - } - Err(e) => { - // Other errors might indicate parsing issues - println!("Unexpected error: {e}"); - } - Ok(_) => { - // Unexpected success in test environment without server storage - } - } +#[test] +fn test_env_only_audit_target_does_not_require_server_storage() { + with_vars( + [ + ("RUSTFS_AUDIT_WEBHOOK_ENABLE_PRIMARY", Some("on")), + ("RUSTFS_AUDIT_WEBHOOK_ENDPOINT_PRIMARY", Some("http://localhost:3020/webhook")), + ], + || { + let runtime = tokio::runtime::Builder::new_current_thread() + .enable_all() + .build() + .expect("failed to create tokio runtime"); + runtime.block_on(async { + let config = Config(HashMap::new()); + let registry = AuditRegistry::new(); + let result = registry.create_audit_targets_from_config(&config).await; + assert!(result.is_ok(), "env-only audit target creation should not require server storage"); + }); + }, + ) } #[test] diff --git a/crates/config/src/constants/targets.rs b/crates/config/src/constants/targets.rs index d6d3bb50c..ce341ae44 100644 --- a/crates/config/src/constants/targets.rs +++ b/crates/config/src/constants/targets.rs @@ -34,3 +34,10 @@ pub const MQTT_RECONNECT_INTERVAL: &str = "reconnect_interval"; pub const MQTT_KEEP_ALIVE_INTERVAL: &str = "keep_alive_interval"; pub const MQTT_QUEUE_DIR: &str = "queue_dir"; pub const MQTT_QUEUE_LIMIT: &str = "queue_limit"; + +/// Environment variable controlling whether target queue files are Snappy-compressed. +/// Applies to both notify and audit target queue stores. +pub const ENV_TARGET_STORE_COMPRESS: &str = "RUSTFS_TARGET_STORE_COMPRESS"; + +/// Queue-store compression is enabled by default to reduce disk footprint. +pub const DEFAULT_TARGET_STORE_COMPRESS: bool = true; diff --git a/crates/e2e_test/src/object_lambda_test.rs b/crates/e2e_test/src/object_lambda_test.rs index f4ed795a9..658b799de 100644 --- a/crates/e2e_test/src/object_lambda_test.rs +++ b/crates/e2e_test/src/object_lambda_test.rs @@ -313,6 +313,31 @@ async fn list_target_arns(env: &RustFSTestEnvironment) -> Result, Bo Ok(serde_json::from_slice(&body)?) } +async fn delete_webhook_target(env: &RustFSTestEnvironment, target_name: &str) -> Result<(), Box> { + let url = format!("{}/rustfs/admin/v3/target/notify_webhook/{target_name}/reset", env.url); + let response = signed_request(http::Method::DELETE, &url, &env.access_key, &env.secret_key, None, None).await?; + let status = response.status(); + let body = response.text().await.unwrap_or_default(); + if status != StatusCode::OK { + return Err(format!("failed to delete webhook target {target_name}: {status} {body}").into()); + } + + Ok(()) +} + +fn notification_target_is_listed(targets: &serde_json::Value, target_name: &str) -> bool { + targets["notification_endpoints"] + .as_array() + .into_iter() + .flatten() + .any(|entry| { + entry["account_id"].as_str() == Some(target_name) + && entry["service"] + .as_str() + .is_some_and(|service| service == "webhook" || service.starts_with("webhook-")) + }) +} + async fn wait_for_target_visibility( env: &RustFSTestEnvironment, target_name: &str, @@ -324,18 +349,7 @@ async fn wait_for_target_visibility( last_targets = list_notification_targets(env).await?; last_arns = list_target_arns(env).await?; - let listed = last_targets["notification_endpoints"] - .as_array() - .into_iter() - .flatten() - .any(|entry| { - entry["account_id"].as_str() == Some(target_name) - && entry["service"] - .as_str() - .is_some_and(|service| service == "webhook" || service.starts_with("webhook-")) - }); - - if listed { + if notification_target_is_listed(&last_targets, target_name) { return Ok((last_targets, last_arns)); } @@ -345,10 +359,51 @@ async fn wait_for_target_visibility( Err(format!("target {target_name} did not become visible in admin APIs; targets={last_targets}, arns={last_arns:?}").into()) } +async fn wait_for_target_absence( + env: &RustFSTestEnvironment, + target_name: &str, +) -> Result<(serde_json::Value, Vec), Box> { + let mut last_targets = serde_json::Value::Null; + let mut last_arns = Vec::new(); + + for _ in 0..20 { + last_targets = list_notification_targets(env).await?; + last_arns = list_target_arns(env).await?; + + let listed = notification_target_is_listed(&last_targets, target_name); + let arn_listed = last_arns.iter().any(|arn| arn.ends_with(&format!(":{target_name}:webhook"))); + if !listed && !arn_listed { + return Ok((last_targets, last_arns)); + } + + tokio::time::sleep(Duration::from_millis(250)).await; + } + + Err(format!("target {target_name} remained visible in admin APIs; targets={last_targets}, arns={last_arns:?}").into()) +} + +async fn restart_rustfs_server(env: &mut RustFSTestEnvironment) -> Result<(), Box> { + env.stop_server(); + env.start_rustfs_server_without_cleanup(vec![]).await +} + async fn read_persisted_server_config(env: &RustFSTestEnvironment) -> String { let path = format!("{}/.rustfs.sys/config/config.json", env.temp_dir); match tokio::fs::read_to_string(&path).await { Ok(content) => content, + Err(err) if err.kind() == std::io::ErrorKind::IsADirectory => { + let mut entries = Vec::new(); + match tokio::fs::read_dir(&path).await { + Ok(mut dir) => { + while let Ok(Some(entry)) = dir.next_entry().await { + entries.push(entry.file_name().to_string_lossy().to_string()); + } + entries.sort(); + format!("persisted config stored as object directory at {path}; entries={entries:?}") + } + Err(dir_err) => format!("persisted config directory exists at {path} but could not be listed: {dir_err}"), + } + } Err(err) => format!("failed to read persisted config at {path}: {err}"), } } @@ -400,6 +455,66 @@ async fn read_listen_notification_event( } } +#[tokio::test] +#[serial] +async fn test_notification_target_persists_across_restart_and_delete() -> Result<(), Box> { + init_logging(); + + let (webhook_url, _request_rx, webhook_handle) = spawn_object_lambda_webhook_server().await?; + + let mut env = RustFSTestEnvironment::new().await?; + env.start_rustfs_server(vec![]).await?; + + let target_name = "restart-target"; + configure_webhook_target(&env, target_name, &webhook_url, "secret-token").await?; + + let (visible_targets, visible_arns) = wait_for_target_visibility(&env, target_name).await?; + assert!(notification_target_is_listed(&visible_targets, target_name)); + assert!( + visible_arns + .iter() + .any(|arn| arn.ends_with(&format!(":{target_name}:webhook"))), + "target ARN missing after initial configure: {visible_arns:?}" + ); + + restart_rustfs_server(&mut env).await?; + + let (targets_after_restart, arns_after_restart) = wait_for_target_visibility(&env, target_name).await?; + assert!(notification_target_is_listed(&targets_after_restart, target_name)); + assert!( + arns_after_restart + .iter() + .any(|arn| arn.ends_with(&format!(":{target_name}:webhook"))), + "target ARN missing after restart: {arns_after_restart:?}" + ); + + delete_webhook_target(&env, target_name).await?; + let (targets_after_delete, arns_after_delete) = wait_for_target_absence(&env, target_name).await?; + assert!(!notification_target_is_listed(&targets_after_delete, target_name)); + assert!( + !arns_after_delete + .iter() + .any(|arn| arn.ends_with(&format!(":{target_name}:webhook"))), + "target ARN still visible after delete: {arns_after_delete:?}" + ); + + restart_rustfs_server(&mut env).await?; + + let (targets_after_delete_restart, arns_after_delete_restart) = wait_for_target_absence(&env, target_name).await?; + assert!(!notification_target_is_listed(&targets_after_delete_restart, target_name)); + assert!( + !arns_after_delete_restart + .iter() + .any(|arn| arn.ends_with(&format!(":{target_name}:webhook"))), + "target ARN still visible after delete + restart: {arns_after_delete_restart:?}" + ); + + webhook_handle.abort(); + let _ = webhook_handle.await; + + Ok(()) +} + #[tokio::test] #[serial] async fn test_get_object_lambda_accepts_presigned_requests() -> Result<(), Box> { diff --git a/crates/ecstore/src/config/audit.rs b/crates/ecstore/src/config/audit.rs index f0c864030..5574daeb4 100644 --- a/crates/ecstore/src/config/audit.rs +++ b/crates/ecstore/src/config/audit.rs @@ -16,8 +16,8 @@ use crate::config::{KV, KVS}; use rustfs_config::{ COMMENT_KEY, DEFAULT_LIMIT, ENABLE_KEY, EVENT_DEFAULT_DIR, EnableState, MQTT_BROKER, MQTT_KEEP_ALIVE_INTERVAL, MQTT_PASSWORD, MQTT_QOS, MQTT_QUEUE_DIR, MQTT_QUEUE_LIMIT, MQTT_RECONNECT_INTERVAL, MQTT_TOPIC, MQTT_USERNAME, WEBHOOK_AUTH_TOKEN, - WEBHOOK_BATCH_SIZE, WEBHOOK_CLIENT_CERT, WEBHOOK_CLIENT_KEY, WEBHOOK_ENDPOINT, WEBHOOK_HTTP_TIMEOUT, WEBHOOK_MAX_RETRY, - WEBHOOK_QUEUE_DIR, WEBHOOK_QUEUE_LIMIT, WEBHOOK_RETRY_INTERVAL, + WEBHOOK_BATCH_SIZE, WEBHOOK_CLIENT_CA, WEBHOOK_CLIENT_CERT, WEBHOOK_CLIENT_KEY, WEBHOOK_ENDPOINT, WEBHOOK_HTTP_TIMEOUT, + WEBHOOK_MAX_RETRY, WEBHOOK_QUEUE_DIR, WEBHOOK_QUEUE_LIMIT, WEBHOOK_RETRY_INTERVAL, WEBHOOK_SKIP_TLS_VERIFY, }; use std::sync::LazyLock; @@ -51,6 +51,16 @@ pub static DEFAULT_AUDIT_WEBHOOK_KVS: LazyLock = LazyLock::new(|| { value: "".to_owned(), hidden_if_empty: false, }, + KV { + key: WEBHOOK_CLIENT_CA.to_owned(), + value: "".to_owned(), + hidden_if_empty: false, + }, + KV { + key: WEBHOOK_SKIP_TLS_VERIFY.to_owned(), + value: EnableState::Off.to_string(), + hidden_if_empty: false, + }, KV { key: WEBHOOK_BATCH_SIZE.to_owned(), value: "1".to_owned(), @@ -81,6 +91,11 @@ pub static DEFAULT_AUDIT_WEBHOOK_KVS: LazyLock = LazyLock::new(|| { value: "5s".to_owned(), hidden_if_empty: false, }, + KV { + key: COMMENT_KEY.to_owned(), + value: "".to_owned(), + hidden_if_empty: false, + }, ]) }); diff --git a/crates/ecstore/src/config/com.rs b/crates/ecstore/src/config/com.rs index 5885969b0..12a534cbd 100644 --- a/crates/ecstore/src/config/com.rs +++ b/crates/ecstore/src/config/com.rs @@ -12,12 +12,14 @@ // See the License for the specific language governing permissions and // limitations under the License. -use crate::config::{Config, GLOBAL_STORAGE_CLASS, KVS, oidc, storageclass}; +use crate::config::{Config, GLOBAL_STORAGE_CLASS, KVS, audit, notify, oidc, storageclass}; use crate::disk::{MIGRATING_META_BUCKET, RUSTFS_META_BUCKET}; use crate::error::{Error, Result}; use crate::global::is_first_cluster_node_local; use crate::store_api::{ObjectInfo, ObjectOptions, PutObjReader, StorageAPI}; use http::HeaderMap; +use rustfs_config::audit::{AUDIT_MQTT_KEYS, AUDIT_MQTT_SUB_SYS, AUDIT_WEBHOOK_KEYS, AUDIT_WEBHOOK_SUB_SYS}; +use rustfs_config::notify::{NOTIFY_MQTT_KEYS, NOTIFY_MQTT_SUB_SYS, NOTIFY_WEBHOOK_KEYS, NOTIFY_WEBHOOK_SUB_SYS}; use rustfs_config::oidc::{IDENTITY_OPENID_KEYS, IDENTITY_OPENID_SUB_SYS, OIDC_REDIRECT_URI_DYNAMIC}; use rustfs_config::{COMMENT_KEY, DEFAULT_DELIMITER, ENABLE_KEY, EnableState, RUSTFS_REGION}; use rustfs_utils::path::SLASH_SEPARATOR; @@ -261,6 +263,159 @@ fn apply_external_oidc_map(cfg: &mut Config, root: &Map) -> bool applied } +fn parse_notify_scalar_value(key: &str, value: &Value) -> Option { + match value { + Value::String(v) => Some(v.trim().to_string()), + Value::Bool(v) if key == ENABLE_KEY || key == rustfs_config::WEBHOOK_SKIP_TLS_VERIFY => Some(if *v { + EnableState::On.to_string() + } else { + EnableState::Off.to_string() + }), + Value::Bool(v) => Some(v.to_string()), + Value::Number(v) => Some(v.to_string()), + Value::Null => None, + _ => None, + } +} + +fn decode_notify_instance_object(instance: &Map, valid_keys: &[&str]) -> KVS { + let mut kvs = KVS::new(); + + for (key, value) in instance { + if !valid_keys.contains(&key.as_str()) || key == COMMENT_KEY { + continue; + } + + if let Some(parsed) = parse_notify_scalar_value(key, value) { + kvs.insert(key.clone(), parsed); + } + } + + kvs +} + +fn decode_notify_instance_value(value: &Value, valid_keys: &[&str]) -> Option { + match value { + Value::Object(instance) => Some(decode_notify_instance_object(instance, valid_keys)), + Value::Array(_) => serde_json::from_value::(value.clone()).ok(), + _ => None, + } +} + +fn is_notify_instance_shorthand(section: &Map, valid_keys: &[&str]) -> bool { + section + .iter() + .any(|(key, value)| valid_keys.contains(&key.as_str()) && parse_notify_scalar_value(key, value).is_some()) +} + +fn apply_external_notify_section( + cfg: &mut Config, + notify_obj: &Map, + external_key: &str, + subsystem_key: &str, + default_kvs: &KVS, + valid_keys: &[&str], +) -> bool { + let Some(Value::Object(section_obj)) = notify_obj.get(external_key).or_else(|| notify_obj.get(subsystem_key)) else { + return false; + }; + + if section_obj.is_empty() { + return false; + } + + let subsystem = cfg.0.entry(subsystem_key.to_string()).or_default(); + let mut applied = false; + + if is_notify_instance_shorthand(section_obj, valid_keys) { + let kvs = decode_notify_instance_object(section_obj, valid_keys); + if !kvs.is_empty() { + let mut merged = default_kvs.clone(); + merged.extend(kvs); + subsystem.insert(DEFAULT_DELIMITER.to_string(), merged); + applied = true; + } + return applied; + } + + for (raw_instance, value) in section_obj { + let Some(mut kvs) = decode_notify_instance_value(value, valid_keys) else { + continue; + }; + if kvs.is_empty() { + continue; + } + + let instance_key = if raw_instance == "default" { + DEFAULT_DELIMITER.to_string() + } else { + raw_instance.to_string() + }; + + if instance_key == DEFAULT_DELIMITER { + let mut merged = default_kvs.clone(); + merged.extend(kvs); + kvs = merged; + } + + subsystem.insert(instance_key, kvs); + applied = true; + } + + applied +} + +fn apply_external_notify_map(cfg: &mut Config, root: &Map) -> bool { + let Some(Value::Object(notify_obj)) = root.get("notify") else { + return false; + }; + + let mut applied = false; + applied |= apply_external_notify_section( + cfg, + notify_obj, + "webhook", + NOTIFY_WEBHOOK_SUB_SYS, + ¬ify::DEFAULT_NOTIFY_WEBHOOK_KVS, + NOTIFY_WEBHOOK_KEYS, + ); + applied |= apply_external_notify_section( + cfg, + notify_obj, + "mqtt", + NOTIFY_MQTT_SUB_SYS, + ¬ify::DEFAULT_NOTIFY_MQTT_KVS, + NOTIFY_MQTT_KEYS, + ); + applied +} + +fn apply_external_audit_map(cfg: &mut Config, root: &Map) -> bool { + let audit_root = root.get("audit").or_else(|| root.get("logger")).and_then(Value::as_object); + let Some(audit_obj) = audit_root else { + return false; + }; + + let mut applied = false; + applied |= apply_external_notify_section( + cfg, + audit_obj, + "webhook", + AUDIT_WEBHOOK_SUB_SYS, + &audit::DEFAULT_AUDIT_WEBHOOK_KVS, + AUDIT_WEBHOOK_KEYS, + ); + applied |= apply_external_notify_section( + cfg, + audit_obj, + "mqtt", + AUDIT_MQTT_SUB_SYS, + &audit::DEFAULT_AUDIT_MQTT_KVS, + AUDIT_MQTT_KEYS, + ); + applied +} + fn apply_external_storage_class_map(cfg: &mut Config, root: &Map) -> bool { let sc = root.get("storageclass").or_else(|| root.get("storage_class")); let Some(Value::Object(sc_obj)) = sc else { @@ -305,8 +460,10 @@ fn decode_server_config_blob(data: &[u8]) -> Result { let mut cfg = Config::new(); let has_storage = apply_external_storage_class_map(&mut cfg, &root); let has_oidc = apply_external_oidc_map(&mut cfg, &root); + let has_notify = apply_external_notify_map(&mut cfg, &root); + let has_audit = apply_external_audit_map(&mut cfg, &root); let has_header = root.contains_key("version") || root.contains_key("region") || root.contains_key("credential"); - if !has_storage && !has_oidc && !has_header { + if !has_storage && !has_oidc && !has_notify && !has_audit && !has_header { return Err(Error::other("unrecognized external server config shape")); } Ok(cfg) @@ -449,6 +606,139 @@ fn build_semantic_oidc_object(cfg: &Config) -> Map { oidc_obj } +fn is_notify_bool_key(key: &str) -> bool { + key == ENABLE_KEY || key == rustfs_config::WEBHOOK_SKIP_TLS_VERIFY +} + +fn encode_notify_scalar_value(key: &str, value: &str) -> Value { + if is_notify_bool_key(key) { + if let Ok(state) = value.parse::() { + return Value::Bool(state.is_enabled()); + } + if let Ok(boolean) = value.parse::() { + return Value::Bool(boolean); + } + } + + Value::String(value.to_string()) +} + +fn is_hidden_if_empty(default_kvs: &KVS, key: &str) -> bool { + default_kvs + .0 + .iter() + .find(|kv| kv.key == key) + .map(|kv| kv.hidden_if_empty) + .unwrap_or(false) +} + +fn build_notify_instance_diff_object(kvs: &KVS, baseline: &KVS, valid_keys: &[&str], default_kvs: &KVS) -> Map { + let mut instance = Map::new(); + + for key in valid_keys { + if *key == COMMENT_KEY { + continue; + } + + let baseline_value = baseline.lookup(key).unwrap_or_default(); + let effective_value = kvs.lookup(key).unwrap_or_else(|| baseline_value.clone()); + + if effective_value == baseline_value { + continue; + } + + if effective_value.trim().is_empty() && baseline_value.trim().is_empty() { + continue; + } + + if is_hidden_if_empty(default_kvs, key) && effective_value.trim().is_empty() && baseline_value.trim().is_empty() { + continue; + } + + instance.insert((*key).to_string(), encode_notify_scalar_value(key, &effective_value)); + } + + instance +} + +fn merged_notify_default_kvs(subsystem: &HashMap, default_kvs: &KVS) -> KVS { + let mut merged = default_kvs.clone(); + if let Some(kvs) = subsystem.get(DEFAULT_DELIMITER) { + merged.extend(kvs.clone()); + } + merged +} + +fn build_notify_subsystem_object( + cfg: &Config, + subsystem_key: &str, + default_kvs: &KVS, + valid_keys: &[&str], +) -> Map { + let Some(subsystem) = cfg.0.get(subsystem_key) else { + return Map::new(); + }; + + let effective_default = merged_notify_default_kvs(subsystem, default_kvs); + let mut subsystem_obj = Map::new(); + + if let Some(default_instance) = subsystem.get(DEFAULT_DELIMITER) { + let default_obj = build_notify_instance_diff_object(default_instance, default_kvs, valid_keys, default_kvs); + if !default_obj.is_empty() { + subsystem_obj.insert("default".to_string(), Value::Object(default_obj)); + } + } + + let mut instances = subsystem + .iter() + .filter(|(instance_key, _)| instance_key.as_str() != DEFAULT_DELIMITER) + .collect::>(); + instances.sort_by(|(lhs, _), (rhs, _)| lhs.cmp(rhs)); + + for (instance_key, kvs) in instances { + let instance_obj = build_notify_instance_diff_object(kvs, &effective_default, valid_keys, default_kvs); + if !instance_obj.is_empty() { + subsystem_obj.insert(instance_key.clone(), Value::Object(instance_obj)); + } + } + + subsystem_obj +} + +fn build_notify_object(cfg: &Config) -> Map { + let mut notify_obj = Map::new(); + + let webhook_obj = + build_notify_subsystem_object(cfg, NOTIFY_WEBHOOK_SUB_SYS, ¬ify::DEFAULT_NOTIFY_WEBHOOK_KVS, NOTIFY_WEBHOOK_KEYS); + if !webhook_obj.is_empty() { + notify_obj.insert("webhook".to_string(), Value::Object(webhook_obj)); + } + + let mqtt_obj = build_notify_subsystem_object(cfg, NOTIFY_MQTT_SUB_SYS, ¬ify::DEFAULT_NOTIFY_MQTT_KVS, NOTIFY_MQTT_KEYS); + if !mqtt_obj.is_empty() { + notify_obj.insert("mqtt".to_string(), Value::Object(mqtt_obj)); + } + + notify_obj +} + +fn build_audit_object(cfg: &Config) -> Map { + let mut audit_obj = Map::new(); + + let webhook_obj = + build_notify_subsystem_object(cfg, AUDIT_WEBHOOK_SUB_SYS, &audit::DEFAULT_AUDIT_WEBHOOK_KVS, AUDIT_WEBHOOK_KEYS); + if !webhook_obj.is_empty() { + audit_obj.insert("webhook".to_string(), Value::Object(webhook_obj)); + } + + let mqtt_obj = build_notify_subsystem_object(cfg, AUDIT_MQTT_SUB_SYS, &audit::DEFAULT_AUDIT_MQTT_KVS, AUDIT_MQTT_KEYS); + if !mqtt_obj.is_empty() { + audit_obj.insert("mqtt".to_string(), Value::Object(mqtt_obj)); + } + + audit_obj +} + fn encode_server_config_blob(cfg: &Config, seed: Option<&[u8]>) -> Result> { let mut root = seed.and_then(parse_object_seed).unwrap_or_default(); @@ -478,6 +768,73 @@ fn encode_server_config_blob(cfg: &Config, seed: Option<&[u8]>) -> Result v, + _ => Map::new(), + }; + let rendered_notify = build_notify_object(cfg); + match rendered_notify.get("webhook") { + Some(Value::Object(v)) => { + notify_obj.insert("webhook".to_string(), Value::Object(v.clone())); + notify_obj.remove(NOTIFY_WEBHOOK_SUB_SYS); + } + _ => { + notify_obj.remove("webhook"); + notify_obj.remove(NOTIFY_WEBHOOK_SUB_SYS); + } + } + match rendered_notify.get("mqtt") { + Some(Value::Object(v)) => { + notify_obj.insert("mqtt".to_string(), Value::Object(v.clone())); + notify_obj.remove(NOTIFY_MQTT_SUB_SYS); + } + _ => { + notify_obj.remove("mqtt"); + notify_obj.remove(NOTIFY_MQTT_SUB_SYS); + } + } + if notify_obj.is_empty() { + root.remove("notify"); + } else { + root.insert("notify".to_string(), Value::Object(notify_obj)); + } + root.remove(NOTIFY_WEBHOOK_SUB_SYS); + root.remove(NOTIFY_MQTT_SUB_SYS); + + let mut logger_obj = match root.remove("logger") { + Some(Value::Object(v)) => v, + _ => Map::new(), + }; + let rendered_audit = build_audit_object(cfg); + match rendered_audit.get("webhook") { + Some(Value::Object(v)) => { + logger_obj.insert("webhook".to_string(), Value::Object(v.clone())); + logger_obj.remove(AUDIT_WEBHOOK_SUB_SYS); + } + _ => { + logger_obj.remove("webhook"); + logger_obj.remove(AUDIT_WEBHOOK_SUB_SYS); + } + } + match rendered_audit.get("mqtt") { + Some(Value::Object(v)) => { + logger_obj.insert("mqtt".to_string(), Value::Object(v.clone())); + logger_obj.remove(AUDIT_MQTT_SUB_SYS); + } + _ => { + logger_obj.remove("mqtt"); + logger_obj.remove(AUDIT_MQTT_SUB_SYS); + } + } + if logger_obj.is_empty() { + root.remove("logger"); + } else { + root.insert("logger".to_string(), Value::Object(logger_obj)); + } + root.remove("audit"); + root.remove(AUDIT_WEBHOOK_SUB_SYS); + root.remove(AUDIT_MQTT_SUB_SYS); + Ok(serde_json::to_vec(&Value::Object(root))?) } @@ -496,6 +853,8 @@ fn is_standard_object_server_config(data: &[u8]) -> bool { fn configs_semantically_equal(lhs: &Config, rhs: &Config) -> bool { build_storageclass_object(lhs) == build_storageclass_object(rhs) && build_semantic_oidc_object(lhs) == build_semantic_oidc_object(rhs) + && build_notify_object(lhs) == build_notify_object(rhs) + && build_audit_object(lhs) == build_audit_object(rhs) } fn is_object_not_found(err: &Error) -> bool { @@ -712,7 +1071,7 @@ mod tests { configs_semantically_equal, decode_server_config_blob, encode_server_config_blob, is_standard_object_server_config, read_config_with_metadata, storage_class_kvs_mut, }; - use crate::config::{Config, oidc}; + use crate::config::{Config, audit, notify, oidc}; use crate::disk::endpoint::Endpoint; use crate::endpoints::SetupType; use crate::error::{Error, Result}; @@ -725,6 +1084,8 @@ mod tests { ObjectOptions, ObjectToDelete, PartInfo, PutObjReader, StorageAPI, WalkOptions, }; use http::HeaderMap; + use rustfs_config::audit::{AUDIT_MQTT_SUB_SYS, AUDIT_WEBHOOK_SUB_SYS}; + use rustfs_config::notify::{NOTIFY_MQTT_SUB_SYS, NOTIFY_WEBHOOK_SUB_SYS}; use rustfs_config::oidc::IDENTITY_OPENID_SUB_SYS; use rustfs_config::{DEFAULT_DELIMITER, ENABLE_KEY, EnableState}; use rustfs_filemeta::FileInfo; @@ -1332,6 +1693,152 @@ mod tests { ); } + #[test] + fn test_decode_server_config_reads_notify_targets() { + let input = r#"{ + "version":"33", + "storageclass":{"standard":"EC:2","rrs":"EC:1"}, + "notify":{ + "webhook":{ + "primary":{ + "enable":true, + "endpoint":"https://example.com/hook", + "queue_dir":"/tmp/webhook-queue" + } + }, + "mqtt":{ + "default":{ + "enable":true, + "topic":"events" + }, + "analytics":{ + "enable":true, + "broker":"tcp://127.0.0.1:1883", + "topic":"events", + "queue_dir":"" + } + } + } + }"#; + + let cfg = decode_server_config_blob(input.as_bytes()).expect("decode should succeed"); + + let webhook = cfg + .get_value(NOTIFY_WEBHOOK_SUB_SYS, "primary") + .expect("webhook target should be decoded"); + assert_eq!(webhook.get(ENABLE_KEY), EnableState::On.to_string()); + assert_eq!(webhook.get(rustfs_config::WEBHOOK_ENDPOINT), "https://example.com/hook"); + assert_eq!(webhook.get(rustfs_config::WEBHOOK_QUEUE_DIR), "/tmp/webhook-queue"); + + let mqtt_default = cfg + .get_value(NOTIFY_MQTT_SUB_SYS, DEFAULT_DELIMITER) + .expect("mqtt default should be decoded"); + assert_eq!(mqtt_default.get(ENABLE_KEY), EnableState::On.to_string()); + assert_eq!(mqtt_default.get(rustfs_config::MQTT_TOPIC), "events"); + assert_eq!( + mqtt_default.get(rustfs_config::MQTT_QUEUE_DIR), + notify::DEFAULT_NOTIFY_MQTT_KVS.get(rustfs_config::MQTT_QUEUE_DIR) + ); + + let mqtt = cfg + .get_value(NOTIFY_MQTT_SUB_SYS, "analytics") + .expect("mqtt target should be decoded"); + assert_eq!(mqtt.get(rustfs_config::MQTT_BROKER), "tcp://127.0.0.1:1883"); + assert_eq!(mqtt.get(rustfs_config::MQTT_QUEUE_DIR), ""); + } + + #[test] + fn test_decode_server_config_reads_notify_shorthand_default() { + let input = r#"{ + "version":"33", + "storageclass":{"standard":"EC:2","rrs":"EC:1"}, + "notify":{ + "webhook":{ + "enable":true, + "endpoint":"https://example.com/shorthand" + } + } + }"#; + + let cfg = decode_server_config_blob(input.as_bytes()).expect("decode should succeed"); + let webhook_default = cfg + .get_value(NOTIFY_WEBHOOK_SUB_SYS, DEFAULT_DELIMITER) + .expect("default webhook config should be decoded"); + assert_eq!(webhook_default.get(ENABLE_KEY), EnableState::On.to_string()); + assert_eq!(webhook_default.get(rustfs_config::WEBHOOK_ENDPOINT), "https://example.com/shorthand"); + } + + #[test] + fn test_decode_server_config_keeps_instance_named_like_field() { + let input = r#"{ + "version":"33", + "storageclass":{"standard":"EC:2","rrs":"EC:1"}, + "notify":{ + "webhook":{ + "enable":{ + "enable":true, + "endpoint":"https://example.com/instance-enable" + } + } + } + }"#; + + let cfg = decode_server_config_blob(input.as_bytes()).expect("decode should succeed"); + let named = cfg + .get_value(NOTIFY_WEBHOOK_SUB_SYS, "enable") + .expect("instance named 'enable' should be decoded"); + assert_eq!(named.get(ENABLE_KEY), EnableState::On.to_string()); + assert_eq!(named.get(rustfs_config::WEBHOOK_ENDPOINT), "https://example.com/instance-enable"); + } + + #[test] + fn test_decode_server_config_reads_audit_targets() { + let input = r#"{ + "version":"33", + "storageclass":{"standard":"EC:2","rrs":"EC:1"}, + "logger":{ + "webhook":{ + "primary":{ + "enable":true, + "endpoint":"https://example.com/audit-hook", + "queue_dir":"/tmp/audit-queue" + } + }, + "mqtt":{ + "default":{ + "enable":true, + "topic":"audit-events" + }, + "analytics":{ + "enable":true, + "broker":"tcp://127.0.0.1:1883", + "topic":"audit-events" + } + } + } + }"#; + + let cfg = decode_server_config_blob(input.as_bytes()).expect("decode should succeed"); + + let webhook = cfg + .get_value(AUDIT_WEBHOOK_SUB_SYS, "primary") + .expect("audit webhook target should be decoded"); + assert_eq!(webhook.get(ENABLE_KEY), EnableState::On.to_string()); + assert_eq!(webhook.get(rustfs_config::WEBHOOK_ENDPOINT), "https://example.com/audit-hook"); + assert_eq!(webhook.get(rustfs_config::WEBHOOK_QUEUE_DIR), "/tmp/audit-queue"); + + let mqtt_default = cfg + .get_value(AUDIT_MQTT_SUB_SYS, DEFAULT_DELIMITER) + .expect("audit mqtt default should be decoded"); + assert_eq!(mqtt_default.get(ENABLE_KEY), EnableState::On.to_string()); + assert_eq!(mqtt_default.get(rustfs_config::MQTT_TOPIC), "audit-events"); + + let mqtt = cfg + .get_value(AUDIT_MQTT_SUB_SYS, "analytics") + .expect("audit mqtt target should be decoded"); + assert_eq!(mqtt.get(rustfs_config::MQTT_BROKER), "tcp://127.0.0.1:1883"); + } + #[test] fn test_encode_server_config_writes_external_object_shape() { let mut cfg = Config::new(); @@ -1388,6 +1895,174 @@ mod tests { assert_eq!(default_provider.get(ENABLE_KEY).and_then(Value::as_bool), Some(true)); } + #[test] + fn test_encode_server_config_writes_notify_object_shape() { + let mut cfg = Config::new(); + let mut webhook_section = std::collections::HashMap::new(); + webhook_section.insert(DEFAULT_DELIMITER.to_string(), notify::DEFAULT_NOTIFY_WEBHOOK_KVS.clone()); + webhook_section.insert( + "primary".to_string(), + crate::config::KVS(vec![ + crate::config::KV { + key: ENABLE_KEY.to_string(), + value: EnableState::On.to_string(), + hidden_if_empty: false, + }, + crate::config::KV { + key: rustfs_config::WEBHOOK_ENDPOINT.to_string(), + value: "https://example.com/hook".to_string(), + hidden_if_empty: false, + }, + crate::config::KV { + key: rustfs_config::WEBHOOK_QUEUE_DIR.to_string(), + value: "/tmp/webhook-queue".to_string(), + hidden_if_empty: false, + }, + ]), + ); + cfg.0.insert(NOTIFY_WEBHOOK_SUB_SYS.to_string(), webhook_section); + + let mut mqtt_default = notify::DEFAULT_NOTIFY_MQTT_KVS.clone(); + mqtt_default.insert(ENABLE_KEY.to_string(), EnableState::On.to_string()); + mqtt_default.insert(rustfs_config::MQTT_TOPIC.to_string(), "events".to_string()); + let mut mqtt_section = std::collections::HashMap::new(); + mqtt_section.insert(DEFAULT_DELIMITER.to_string(), mqtt_default); + mqtt_section.insert( + "analytics".to_string(), + crate::config::KVS(vec![ + crate::config::KV { + key: ENABLE_KEY.to_string(), + value: EnableState::On.to_string(), + hidden_if_empty: false, + }, + crate::config::KV { + key: rustfs_config::MQTT_BROKER.to_string(), + value: "tcp://127.0.0.1:1883".to_string(), + hidden_if_empty: false, + }, + crate::config::KV { + key: rustfs_config::MQTT_QUEUE_DIR.to_string(), + value: "".to_string(), + hidden_if_empty: false, + }, + ]), + ); + cfg.0.insert(NOTIFY_MQTT_SUB_SYS.to_string(), mqtt_section); + + let out = encode_server_config_blob(&cfg, None).expect("encode should succeed"); + let v: Value = serde_json::from_slice(&out).expect("output should be json"); + let notify = v + .get("notify") + .and_then(Value::as_object) + .expect("notify object should be present"); + let webhook = notify + .get("webhook") + .and_then(Value::as_object) + .and_then(|targets| targets.get("primary")) + .and_then(Value::as_object) + .expect("webhook target should be encoded"); + assert_eq!( + webhook.get(rustfs_config::WEBHOOK_ENDPOINT).and_then(Value::as_str), + Some("https://example.com/hook") + ); + assert_eq!(webhook.get(ENABLE_KEY).and_then(Value::as_bool), Some(true)); + + let mqtt_default = notify + .get("mqtt") + .and_then(Value::as_object) + .and_then(|targets| targets.get("default")) + .and_then(Value::as_object) + .expect("mqtt default should be encoded"); + assert_eq!(mqtt_default.get(ENABLE_KEY).and_then(Value::as_bool), Some(true)); + assert_eq!(mqtt_default.get(rustfs_config::MQTT_TOPIC).and_then(Value::as_str), Some("events")); + + let mqtt = notify + .get("mqtt") + .and_then(Value::as_object) + .and_then(|targets| targets.get("analytics")) + .and_then(Value::as_object) + .expect("mqtt target should be encoded"); + assert_eq!(mqtt.get(rustfs_config::MQTT_BROKER).and_then(Value::as_str), Some("tcp://127.0.0.1:1883")); + assert_eq!(mqtt.get(rustfs_config::MQTT_QUEUE_DIR).and_then(Value::as_str), Some("")); + } + + #[test] + fn test_encode_server_config_writes_audit_object_shape() { + let mut cfg = Config::new(); + let mut webhook_section = std::collections::HashMap::new(); + webhook_section.insert(DEFAULT_DELIMITER.to_string(), audit::DEFAULT_AUDIT_WEBHOOK_KVS.clone()); + webhook_section.insert( + "primary".to_string(), + crate::config::KVS(vec![ + crate::config::KV { + key: ENABLE_KEY.to_string(), + value: EnableState::On.to_string(), + hidden_if_empty: false, + }, + crate::config::KV { + key: rustfs_config::WEBHOOK_ENDPOINT.to_string(), + value: "https://example.com/audit-hook".to_string(), + hidden_if_empty: false, + }, + crate::config::KV { + key: rustfs_config::WEBHOOK_QUEUE_DIR.to_string(), + value: "/tmp/audit-queue".to_string(), + hidden_if_empty: false, + }, + ]), + ); + cfg.0.insert(AUDIT_WEBHOOK_SUB_SYS.to_string(), webhook_section); + + let mut mqtt_default = audit::DEFAULT_AUDIT_MQTT_KVS.clone(); + mqtt_default.insert(ENABLE_KEY.to_string(), EnableState::On.to_string()); + mqtt_default.insert(rustfs_config::MQTT_TOPIC.to_string(), "audit-events".to_string()); + let mut mqtt_section = std::collections::HashMap::new(); + mqtt_section.insert(DEFAULT_DELIMITER.to_string(), mqtt_default); + mqtt_section.insert( + "analytics".to_string(), + crate::config::KVS(vec![ + crate::config::KV { + key: ENABLE_KEY.to_string(), + value: EnableState::On.to_string(), + hidden_if_empty: false, + }, + crate::config::KV { + key: rustfs_config::MQTT_BROKER.to_string(), + value: "tcp://127.0.0.1:1883".to_string(), + hidden_if_empty: false, + }, + ]), + ); + cfg.0.insert(AUDIT_MQTT_SUB_SYS.to_string(), mqtt_section); + + let out = encode_server_config_blob(&cfg, None).expect("encode should succeed"); + let v: Value = serde_json::from_slice(&out).expect("output should be json"); + let logger = v + .get("logger") + .and_then(Value::as_object) + .expect("logger object should be present"); + let webhook = logger + .get("webhook") + .and_then(Value::as_object) + .and_then(|targets| targets.get("primary")) + .and_then(Value::as_object) + .expect("audit webhook target should be encoded"); + assert_eq!( + webhook.get(rustfs_config::WEBHOOK_ENDPOINT).and_then(Value::as_str), + Some("https://example.com/audit-hook") + ); + assert_eq!(webhook.get(ENABLE_KEY).and_then(Value::as_bool), Some(true)); + + let mqtt_default = logger + .get("mqtt") + .and_then(Value::as_object) + .and_then(|targets| targets.get("default")) + .and_then(Value::as_object) + .expect("audit mqtt default should be encoded"); + assert_eq!(mqtt_default.get(ENABLE_KEY).and_then(Value::as_bool), Some(true)); + assert_eq!(mqtt_default.get(rustfs_config::MQTT_TOPIC).and_then(Value::as_str), Some("audit-events")); + } + #[test] fn test_is_standard_object_server_config_detection() { let external = br#"{"version":"33","storageclass":{"standard":"EC:2","rrs":"EC:1"}}"#; @@ -1441,6 +2116,113 @@ mod tests { assert!(configs_semantically_equal(&lhs, &rhs)); } + #[test] + fn test_configs_semantically_equal_accounts_for_notify() { + let external = br#"{ + "version":"33", + "storageclass":{"standard":"EC:2","rrs":"EC:1","optimize":"availability"}, + "notify":{ + "webhook":{ + "primary":{ + "enable":true, + "endpoint":"https://example.com/hook" + } + } + } + }"#; + let legacy = br#"{ + "storage_class":{"_":[ + {"key":"standard","value":"EC:2"}, + {"key":"rrs","value":"EC:1"}, + {"key":"optimize","value":"availability"} + ]}, + "notify_webhook":{ + "_":[ + {"key":"enable","value":"off"}, + {"key":"endpoint","value":""}, + {"key":"queue_limit","value":"100000"}, + {"key":"queue_dir","value":"/opt/rustfs/events"}, + {"key":"client_cert","value":""}, + {"key":"client_key","value":""}, + {"key":"comment","value":""}, + {"key":"client_ca","value":""}, + {"key":"skip_tls_verify","value":"off"} + ], + "primary":[ + {"key":"enable","value":"on"}, + {"key":"endpoint","value":"https://example.com/hook"} + ] + } + }"#; + + let lhs = decode_server_config_blob(external).expect("decode external"); + let rhs = decode_server_config_blob(legacy).expect("decode legacy"); + assert!(configs_semantically_equal(&lhs, &rhs)); + } + + #[test] + fn test_configs_semantically_equal_detects_notify_changes() { + let lhs = decode_server_config_blob( + br#"{"version":"33","storageclass":{"standard":"EC:2","rrs":"EC:1"},"notify":{"webhook":{"primary":{"enable":true,"endpoint":"https://example.com/a"}}}}"#, + ) + .expect("decode lhs"); + let rhs = decode_server_config_blob( + br#"{"version":"33","storageclass":{"standard":"EC:2","rrs":"EC:1"},"notify":{"webhook":{"primary":{"enable":true,"endpoint":"https://example.com/b"}}}}"#, + ) + .expect("decode rhs"); + + assert!(!configs_semantically_equal(&lhs, &rhs)); + } + + #[test] + fn test_configs_semantically_equal_accounts_for_audit() { + let external = br#"{ + "version":"33", + "storageclass":{"standard":"EC:2","rrs":"EC:1","optimize":"availability"}, + "logger":{ + "webhook":{ + "primary":{ + "enable":true, + "endpoint":"https://example.com/audit-hook" + } + } + } + }"#; + let legacy = br#"{ + "storage_class":{"_":[ + {"key":"standard","value":"EC:2"}, + {"key":"rrs","value":"EC:1"}, + {"key":"optimize","value":"availability"} + ]}, + "audit_webhook":{ + "_":[ + {"key":"enable","value":"off"}, + {"key":"endpoint","value":""}, + {"key":"auth_token","value":""}, + {"key":"client_cert","value":""}, + {"key":"client_key","value":""}, + {"key":"client_ca","value":""}, + {"key":"skip_tls_verify","value":"off"}, + {"key":"batch_size","value":"1"}, + {"key":"queue_limit","value":"100000"}, + {"key":"queue_dir","value":"/opt/rustfs/events"}, + {"key":"max_retry","value":"0"}, + {"key":"retry_interval","value":"3s"}, + {"key":"http_timeout","value":"5s"}, + {"key":"comment","value":""} + ], + "primary":[ + {"key":"enable","value":"on"}, + {"key":"endpoint","value":"https://example.com/audit-hook"} + ] + } + }"#; + + let lhs = decode_server_config_blob(external).expect("decode external"); + let rhs = decode_server_config_blob(legacy).expect("decode legacy"); + assert!(configs_semantically_equal(&lhs, &rhs)); + } + #[tokio::test(flavor = "multi_thread")] #[serial] async fn test_read_config_with_metadata_succeeds_with_one_healthy_locker_in_two_node_dist_setup() { diff --git a/crates/ecstore/src/config/notify.rs b/crates/ecstore/src/config/notify.rs index c9ebf3ba6..5f3bb9442 100644 --- a/crates/ecstore/src/config/notify.rs +++ b/crates/ecstore/src/config/notify.rs @@ -16,7 +16,8 @@ use crate::config::{KV, KVS}; use rustfs_config::{ COMMENT_KEY, DEFAULT_LIMIT, ENABLE_KEY, EVENT_DEFAULT_DIR, EnableState, MQTT_BROKER, MQTT_KEEP_ALIVE_INTERVAL, MQTT_PASSWORD, MQTT_QOS, MQTT_QUEUE_DIR, MQTT_QUEUE_LIMIT, MQTT_RECONNECT_INTERVAL, MQTT_TOPIC, MQTT_USERNAME, WEBHOOK_AUTH_TOKEN, - WEBHOOK_CLIENT_CERT, WEBHOOK_CLIENT_KEY, WEBHOOK_ENDPOINT, WEBHOOK_QUEUE_DIR, WEBHOOK_QUEUE_LIMIT, + WEBHOOK_CLIENT_CA, WEBHOOK_CLIENT_CERT, WEBHOOK_CLIENT_KEY, WEBHOOK_ENDPOINT, WEBHOOK_QUEUE_DIR, WEBHOOK_QUEUE_LIMIT, + WEBHOOK_SKIP_TLS_VERIFY, }; use std::sync::LazyLock; @@ -60,6 +61,16 @@ pub static DEFAULT_NOTIFY_WEBHOOK_KVS: LazyLock = LazyLock::new(|| { value: "".to_owned(), hidden_if_empty: false, }, + KV { + key: WEBHOOK_CLIENT_CA.to_owned(), + value: "".to_owned(), + hidden_if_empty: false, + }, + KV { + key: WEBHOOK_SKIP_TLS_VERIFY.to_owned(), + value: EnableState::Off.to_string(), + hidden_if_empty: false, + }, KV { key: COMMENT_KEY.to_owned(), value: "".to_owned(), diff --git a/crates/filemeta/Cargo.toml b/crates/filemeta/Cargo.toml index b2d9feff5..3ac6c96d6 100644 --- a/crates/filemeta/Cargo.toml +++ b/crates/filemeta/Cargo.toml @@ -44,6 +44,7 @@ regex.workspace = true [dev-dependencies] criterion = { workspace = true } +tempfile = { workspace = true } [[bench]] name = "xl_meta_bench" diff --git a/crates/filemeta/src/filemeta.rs b/crates/filemeta/src/filemeta.rs index 2ed9d8e1f..33eee875a 100644 --- a/crates/filemeta/src/filemeta.rs +++ b/crates/filemeta/src/filemeta.rs @@ -1907,7 +1907,6 @@ mod test { #[tokio::test] async fn test_read_xl_meta_no_data() { - use tokio::fs; use tokio::fs::File; use tokio::io::AsyncWriteExt; @@ -1926,13 +1925,15 @@ async fn test_read_xl_meta_no_data() { buff.resize(buff.len() + 100, 0); - let filepath = "./test_xl.meta"; + // Use tempfile to avoid conflicts with parallel tests or previous runs + let dir = tempfile::tempdir().unwrap(); + let filepath = dir.path().join("test_xl.meta"); - let mut file = File::create(filepath).await.unwrap(); + let mut file = File::create(&filepath).await.unwrap(); // Write string data file.write_all(&buff).await.unwrap(); - let mut f = File::open(filepath).await.unwrap(); + let mut f = File::open(&filepath).await.unwrap(); let stat = f.metadata().await.unwrap(); @@ -1941,7 +1942,5 @@ async fn test_read_xl_meta_no_data() { let mut newfm = FileMeta::default(); newfm.unmarshal_msg(&data).unwrap(); - fs::remove_file(filepath).await.unwrap(); - assert_eq!(fm, newfm) } diff --git a/crates/notify/src/integration.rs b/crates/notify/src/integration.rs index 89657f4c4..4ac5087e7 100644 --- a/crates/notify/src/integration.rs +++ b/crates/notify/src/integration.rs @@ -18,12 +18,14 @@ use crate::{ Event, error::NotificationError, notifier::EventNotifier, registry::TargetRegistry, rules::BucketNotificationConfig, stream, }; use hashbrown::HashMap; -use rustfs_config::notify::{DEFAULT_NOTIFY_TARGET_STREAM_CONCURRENCY, ENV_NOTIFY_TARGET_STREAM_CONCURRENCY}; +use rustfs_config::notify::{ + DEFAULT_NOTIFY_TARGET_STREAM_CONCURRENCY, ENV_NOTIFY_TARGET_STREAM_CONCURRENCY, NOTIFY_MQTT_SUB_SYS, NOTIFY_WEBHOOK_SUB_SYS, +}; use rustfs_ecstore::config::{Config, KVS}; use rustfs_s3_common::EventName; use rustfs_targets::arn::TargetID; use rustfs_targets::store::{Key, Store}; -use rustfs_targets::target::EntityTarget; +use rustfs_targets::target::QueuedPayload; use rustfs_targets::{StoreError, Target}; use std::collections::VecDeque; use std::sync::Arc; @@ -34,6 +36,21 @@ use tracing::{debug, error, info, warn}; const MAX_RECENT_LIVE_EVENTS: usize = 1024; +fn subsystem_target_type(target_type: &str) -> &str { + match target_type { + NOTIFY_WEBHOOK_SUB_SYS => "webhook", + NOTIFY_MQTT_SUB_SYS => "mqtt", + _ => target_type, + } +} + +fn runtime_target_id_for_subsystem(target_type: &str, target_name: &str) -> TargetID { + TargetID { + id: target_name.to_lowercase(), + name: subsystem_target_type(target_type).to_string(), + } +} + #[derive(Clone)] pub struct LiveEventBatch { pub events: Vec>, @@ -183,58 +200,84 @@ impl NotificationSystem { } } + /// Initializes targets and starts event streams for those with stores. + /// Returns a map of (target_id -> cancel_sender) for streams that were started. + async fn init_targets_and_start_streams( + &self, + targets: &[Box + Send + Sync>], + ) -> HashMap> { + let mut cancellers = HashMap::new(); + for target in targets { + let target_id = target.id(); + info!("Initializing target: {}", target_id); + + let has_store = target.store().is_some(); + + if let Err(e) = target.init().await { + warn!("Target {} Initialization failed: {}", target_id, e); + // For targets without a store, init failure is fatal — skip. + // For store-backed targets, still start the stream so queued events + // can be drained when connectivity recovers (send_from_store retries). + if !has_store { + continue; + } + warn!( + "Target {} has a store, starting stream despite init failure — \ + connectivity will be retried by send_from_store", + target_id + ); + } else { + debug!("Target {} initialized successfully, enabled: {}", target_id, target.is_enabled()); + } + + if !target.is_enabled() { + info!("Target {} is not enabled, event stream processing is skipped", target_id); + continue; + } + + if let Some(store) = target.store() { + info!("Start event stream processing for target {}", target_id); + + let store_clone = store.boxed_clone(); + let target_arc = Arc::from(target.clone_dyn()); + + let cancel_tx = self.enhanced_start_event_stream( + store_clone, + target_arc, + self.metrics.clone(), + self.concurrency_limiter.clone(), + ); + + let target_id_clone = target_id.clone(); + cancellers.insert(target_id, cancel_tx); + info!("Event stream processing for target {} is started successfully", target_id_clone); + } else { + info!("Target {} No storage is configured, event stream processing is skipped", target_id); + } + } + cancellers + } + /// Initializes the notification system pub async fn init(&self) -> Result<(), NotificationError> { info!("Initialize notification system..."); - let config = self.config.read().await; - debug!("Initializing notification system with config: {:?}", *config); + let config = { + let guard = self.config.read().await; + debug!("Initializing notification system with config: {:?}", *guard); + guard.clone() + }; + let targets: Vec + Send + Sync>> = self.registry.create_targets_from_config(&config).await?; info!("{} notification targets were created", targets.len()); - // Initiate event stream processing for each storage enabled target - let mut cancellers = HashMap::new(); - for target in &targets { - let target_id = target.id(); - info!("Initializing target: {}", target.id()); - // Initialize the target - if let Err(e) = target.init().await { - warn!("Target {} Initialization failed:{}", target.id(), e); - continue; - } - debug!("Target {} initialized successfully,enabled:{}", target_id, target.is_enabled()); - // Check if the target is enabled and has storage - if target.is_enabled() { - if let Some(store) = target.store() { - info!("Start event stream processing for target {}", target.id()); + // Initialize targets and start event streams + let cancellers = self.init_targets_and_start_streams(&targets).await; - // The storage of the cloned target and the target itself - let store_clone = store.boxed_clone(); - let target_box = target.clone_dyn(); - let target_arc = Arc::from(target_box); - - // Add a reference to the monitoring metrics - let metrics = self.metrics.clone(); - let semaphore = self.concurrency_limiter.clone(); - - // Encapsulated enhanced version of start_event_stream - let cancel_tx = self.enhanced_start_event_stream(store_clone, target_arc, metrics, semaphore); - - // Start event stream processing and save cancel sender - let target_id_clone = target_id.clone(); - cancellers.insert(target_id, cancel_tx); - info!("Event stream processing for target {} is started successfully", target_id_clone); - } else { - info!("Target {} No storage is configured, event stream processing is skipped", target_id); - } - } else { - info!("Target {} is not enabled, event stream processing is skipped", target_id); - } - } - - // Update canceler collection + // Update canceller collection *self.stream_cancellers.write().await = cancellers; + // Initialize the bucket target self.notifier.init_bucket_targets(targets).await?; info!("Notification system initialized"); @@ -333,7 +376,7 @@ impl NotificationSystem { info!("Attempting to remove target: {}", target_id); let ttype = target_type.to_lowercase(); - let tname = target_id.name.to_lowercase(); + let tname = target_id.id.to_lowercase(); self.update_config_and_reload(|config| { let mut changed = false; @@ -405,11 +448,7 @@ impl NotificationSystem { let ttype = target_type.to_lowercase(); let tname = target_name.to_lowercase(); - - let target_id = TargetID { - id: tname.clone(), - name: ttype.clone(), - }; + let target_id = runtime_target_id_for_subsystem(&ttype, &tname); // Deletion is prohibited if bucket rules refer to it if self.notifier.is_target_bound_to_any_bucket(&target_id).await { @@ -451,7 +490,7 @@ impl NotificationSystem { /// Enhanced event stream startup function, including monitoring and concurrency control fn enhanced_start_event_stream( &self, - store: Box, Error = StoreError, Key = Key> + Send>, + store: Box + Send>, target: Arc + Send + Sync>, metrics: Arc, semaphore: Arc, @@ -476,14 +515,13 @@ impl NotificationSystem { let _ = cancel_tx.send(()).await; } - // Clear the target_list and ensure that reload is a replacement reconstruction (solve the target_list len unchanged/residual problem) + // Clear the target_list and ensure that reload is a replacement reconstruction self.notifier.remove_all_bucket_targets().await; // Update the config self.update_config(new_config.clone()).await; - // Create a new target from configuration - // This function will now be responsible for merging env, creating and persisting the final configuration. + // Create new targets from configuration let targets: Vec + Send + Sync>> = self .registry .create_targets_from_config(&new_config) @@ -492,46 +530,8 @@ impl NotificationSystem { info!("{} notification targets were created from the new configuration", targets.len()); - // Start new event stream processing for each storage enabled target - let mut new_cancellers = HashMap::new(); - for target in &targets { - let target_id = target.id(); - - // Initialize the target - if let Err(e) = target.init().await { - error!("Target {} Initialization failed:{}", target_id, e); - continue; - } - // Check if the target is enabled and has storage - if target.is_enabled() { - if let Some(store) = target.store() { - info!("Start new event stream processing for target {}", target_id); - - // The storage of the cloned target and the target itself - let store_clone = store.boxed_clone(); - // let target_box = target.clone_dyn(); - let target_arc = Arc::from(target.clone_dyn()); - - // Encapsulated enhanced version of start_event_stream - let cancel_tx = self.enhanced_start_event_stream( - store_clone, - target_arc, - self.metrics.clone(), - self.concurrency_limiter.clone(), - ); - - // Start event stream processing and save cancel sender - // let cancel_tx = start_event_stream(store_clone, target_clone); - let target_id_clone = target_id.clone(); - new_cancellers.insert(target_id, cancel_tx); - info!("Event stream processing of target {} is restarted successfully", target_id_clone); - } else { - info!("Target {} No storage is configured, event stream processing is skipped", target_id); - } - } else { - info!("Target {} disabled, event stream processing is skipped", target_id); - } - } + // Initialize targets and start event streams using shared helper + let new_cancellers = self.init_targets_and_start_streams(&targets).await; // Update canceler collection *cancellers = new_cancellers; @@ -665,4 +665,18 @@ mod tests { assert_eq!(batch.events.len(), 1); assert_eq!(batch.events[0].s3.object.key, "one"); } + + #[test] + fn runtime_target_id_for_subsystem_maps_notify_webhook_to_runtime_type() { + let target_id = runtime_target_id_for_subsystem(NOTIFY_WEBHOOK_SUB_SYS, "Primary"); + assert_eq!(target_id.id, "primary"); + assert_eq!(target_id.name, "webhook"); + } + + #[test] + fn runtime_target_id_for_subsystem_maps_notify_mqtt_to_runtime_type() { + let target_id = runtime_target_id_for_subsystem(NOTIFY_MQTT_SUB_SYS, "Analytics"); + assert_eq!(target_id.id, "analytics"); + assert_eq!(target_id.name, "mqtt"); + } } diff --git a/crates/notify/src/notifier.rs b/crates/notify/src/notifier.rs index 57c1febb2..4d6fe40c3 100644 --- a/crates/notify/src/notifier.rs +++ b/crates/notify/src/notifier.rs @@ -399,7 +399,7 @@ mod tests { use rustfs_targets::{ TargetError, store::{Key, Store}, - target::EntityTarget, + target::{EntityTarget, QueuedPayload, QueuedPayloadMeta}, }; use serde::{Serialize, de::DeserializeOwned}; use std::sync::{ @@ -442,7 +442,7 @@ mod tests { Ok(()) } - async fn send_from_store(&self, _key: Key) -> Result<(), TargetError> { + async fn send_raw_from_store(&self, _key: Key, _body: Vec, _meta: QueuedPayloadMeta) -> Result<(), TargetError> { Ok(()) } @@ -450,7 +450,7 @@ mod tests { Ok(()) } - fn store(&self) -> Option<&(dyn Store, Error = StoreError, Key = Key> + Send + Sync)> { + fn store(&self) -> Option<&(dyn Store + Send + Sync)> { None } diff --git a/crates/notify/src/registry.rs b/crates/notify/src/registry.rs index 74ccbdfb9..6b5e3e515 100644 --- a/crates/notify/src/registry.rs +++ b/crates/notify/src/registry.rs @@ -89,9 +89,6 @@ impl TargetRegistry { let all_env: Vec<(String, String)> = std::env::vars().filter(|(key, _)| key.starts_with(ENV_PREFIX)).collect(); // A collection of asynchronous tasks for concurrently executing target creation let mut tasks = FuturesUnordered::new(); - // let final_config = config.clone(); // Clone a configuration for aggregating the final result - // Record the defaults for each segment so that the segment can eventually be rebuilt - let mut section_defaults: HashMap = HashMap::new(); // 1. Traverse all registered plants and process them by target type for (target_type, factory) in &self.factories { tracing::Span::current().record("target_type", target_type.as_str()); @@ -105,9 +102,6 @@ impl TargetRegistry { let default_cfg = file_configs.get(DEFAULT_DELIMITER).cloned().unwrap_or_default(); debug!(?default_cfg, "Get the default configuration"); - // Save defaults for eventual write back - section_defaults.insert(section_name.clone(), default_cfg.clone()); - // *** Optimization point 1: Get all legitimate fields of the current target type *** let valid_fields = factory.get_valid_fields(); debug!(?valid_fields, "Get the legitimate configuration fields"); @@ -215,110 +209,28 @@ impl TargetRegistry { if enabled { info!(instance_id = %id, "Target is enabled, ready to create a task"); // 5.3. Create asynchronous tasks for enabled instances - let target_type_clone = target_type.clone(); let tid = id.clone(); let merged_config_arc = Arc::new(merged_config); tasks.push(async move { let result = factory.create_target(tid.clone(), &merged_config_arc).await; - (target_type_clone, tid, result, Arc::clone(&merged_config_arc)) + (tid, result) }); } else { - info!(instance_id = %id, "Skip the disabled target and will be removed from the final configuration"); - // Remove disabled target from final configuration - // final_config.0.entry(section_name.clone()).or_default().remove(&id); + info!(instance_id = %id, "Skip disabled target"); } } } // 6. Concurrently execute all creation tasks and collect results let mut successful_targets = Vec::new(); - let mut successful_configs = Vec::new(); - while let Some((target_type, id, result, final_config)) = tasks.next().await { + while let Some((id, result)) = tasks.next().await { match result { Ok(target) => { - info!(target_type = %target_type, instance_id = %id, "Create a target successfully"); + info!(instance_id = %id, "Create target successfully"); successful_targets.push(target); - successful_configs.push((target_type, id, final_config)); } Err(e) => { - error!(target_type = %target_type, instance_id = %id, error = %e, "Failed to create a target"); - } - } - } - - // 7. Aggregate new configuration and write back to system configuration - if !successful_configs.is_empty() || !section_defaults.is_empty() { - info!( - "Prepare to update {} successfully created target configurations to the system configuration...", - successful_configs.len() - ); - - let mut successes_by_section: HashMap> = HashMap::new(); - - for (target_type, id, kvs) in successful_configs { - let section_name = format!("{NOTIFY_ROUTE_PREFIX}{target_type}").to_lowercase(); - successes_by_section - .entry(section_name) - .or_default() - .insert(id.to_lowercase(), (*kvs).clone()); - } - - let mut new_config = config.clone(); - // Collection of segments that need to be processed: Collect all segments where default items exist or where successful instances exist - let mut sections: HashSet = HashSet::new(); - sections.extend(section_defaults.keys().cloned()); - sections.extend(successes_by_section.keys().cloned()); - - for section in sections { - let mut section_map: std::collections::HashMap = std::collections::HashMap::new(); - // Add default item - if let Some(default_kvs) = section_defaults.get(§ion) - && !default_kvs.is_empty() - { - section_map.insert(DEFAULT_DELIMITER.to_string(), default_kvs.clone()); - } - - // Add successful instance item - if let Some(instances) = successes_by_section.get(§ion) { - for (id, kvs) in instances { - section_map.insert(id.clone(), kvs.clone()); - } - } - - // Empty breaks are removed and non-empty breaks are replaced entirely. - if section_map.is_empty() { - new_config.0.remove(§ion); - } else { - new_config.0.insert(section, section_map); - } - } - - if &new_config == config { - info!("Notification target configuration unchanged, skip persisting server config"); - info!(count = successful_targets.len(), "All target processing completed"); - return Ok(successful_targets); - } - - let store = match rustfs_ecstore::global::new_object_layer_fn() { - Some(s) => s, - None => { - warn!( - "Object store not available at notification init; skipping config persistence. \ - {} target(s) active in memory.", - successful_targets.len() - ); - info!(count = successful_targets.len(), "All target processing completed"); - return Ok(successful_targets); - } - }; - - match rustfs_ecstore::config::com::save_server_config(store, &new_config).await { - Ok(_) => { - info!("The new configuration was saved to the system successfully.") - } - Err(e) => { - error!("Failed to save the new configuration: {}", e); - return Err(TargetError::SaveConfig(e.to_string())); + error!(instance_id = %id, error = %e, "Failed to create target"); } } } diff --git a/crates/notify/src/stream.rs b/crates/notify/src/stream.rs index bbb784a87..597d824e1 100644 --- a/crates/notify/src/stream.rs +++ b/crates/notify/src/stream.rs @@ -15,8 +15,8 @@ use crate::{Event, integration::NotificationMetrics}; use rustfs_targets::{ StoreError, Target, TargetError, - store::{Key, Store}, - target::EntityTarget, + store::{Key, Store, ensure_store_entry_raw_readable}, + target::QueuedPayload, }; use rustfs_utils::get_env_usize; use std::sync::Arc; @@ -32,7 +32,7 @@ use tracing::{debug, error, info, warn}; /// - `target`: The target to send events to /// - `cancel_rx`: Receiver to listen for cancellation signals pub async fn stream_events( - store: &mut (dyn Store + Send), + store: &mut (dyn Store + Send), target: &dyn Target, mut cancel_rx: mpsc::Receiver<()>, ) { @@ -119,7 +119,7 @@ pub async fn stream_events( /// # Returns /// A sender to signal cancellation of the event stream pub fn start_event_stream( - mut store: Box + Send>, + mut store: Box + Send>, target: Arc + Send + Sync>, ) -> mpsc::Sender<()> { let (cancel_tx, cancel_rx) = mpsc::channel(1); @@ -143,7 +143,7 @@ pub fn start_event_stream( /// # Returns /// A sender to signal cancellation of the event stream pub fn start_event_stream_with_batching( - mut store: Box, Error = StoreError, Key = Key> + Send>, + mut store: Box + Send>, target: Arc + Send + Sync>, metrics: Arc, semaphore: Arc, @@ -170,7 +170,7 @@ pub fn start_event_stream_with_batching( /// # Notes /// This function processes events in batches to improve efficiency. pub async fn stream_events_with_batching( - store: &mut (dyn Store, Error = StoreError, Key = Key> + Send), + store: &mut (dyn Store + Send), target: &dyn Target, mut cancel_rx: mpsc::Receiver<()>, metrics: Arc, @@ -185,7 +185,6 @@ pub async fn stream_events_with_batching( const MAX_RETRIES: usize = 5; const BASE_RETRY_DELAY: Duration = Duration::from_secs(2); - let mut batch: Vec> = Vec::with_capacity(batch_size); let mut batch_keys = Vec::with_capacity(batch_size); let mut last_flush = Instant::now(); @@ -201,8 +200,8 @@ pub async fn stream_events_with_batching( debug!("Found {} keys in store for target: {}", keys.len(), target.name()); if keys.is_empty() { // If there is data in the batch and timeout, refresh the batch - if !batch.is_empty() && last_flush.elapsed() >= BATCH_TIMEOUT { - process_batch(&mut batch, &mut batch_keys, target, MAX_RETRIES, BASE_RETRY_DELAY, &metrics, &semaphore).await; + if !batch_keys.is_empty() && last_flush.elapsed() >= BATCH_TIMEOUT { + process_batch(&mut batch_keys, target, MAX_RETRIES, BASE_RETRY_DELAY, &metrics, &semaphore).await; last_flush = Instant::now(); } @@ -218,41 +217,31 @@ pub async fn stream_events_with_batching( info!("Cancellation received during processing for target: {}", target.name()); // Processing collected batches before exiting - if !batch.is_empty() { - process_batch(&mut batch, &mut batch_keys, target, MAX_RETRIES, BASE_RETRY_DELAY, &metrics, &semaphore).await; + if !batch_keys.is_empty() { + process_batch(&mut batch_keys, target, MAX_RETRIES, BASE_RETRY_DELAY, &metrics, &semaphore).await; } return; } - // Try to get events from storage - match store.get(&key) { - Ok(event) => { - // Add to batch - batch.push(event); - batch_keys.push(key); - metrics.increment_processing(); - - // If the batch is full or enough time has passed since the last refresh, the batch will be processed - if batch.len() >= batch_size || last_flush.elapsed() >= BATCH_TIMEOUT { - process_batch(&mut batch, &mut batch_keys, target, MAX_RETRIES, BASE_RETRY_DELAY, &metrics, &semaphore) - .await; - last_flush = Instant::now(); - } + // Skip unreadable entries so a single corrupt file cannot stall the stream. + // ensure_store_entry_raw_readable attempts get_raw; on I/O error it calls del() to + // remove the corrupt entry before returning Err, so no cleanup is needed here. + match ensure_store_entry_raw_readable(&*store, &key) { + Ok(true) => {} // entry is readable, proceed + Ok(false) => continue, // entry not found (already removed), skip + Err(err) => { + warn!("Skipping unreadable store entry {} for target {}: {}", key, target.name(), err); + continue; // corrupt entry was already deleted by ensure_store_entry_raw_readable } - Err(e) => { - error!("Failed to target: {}, get event {} from store: {}", target.name(), key.to_string(), e); - // Consider deleting unreadable events to prevent infinite loops from trying to read - match store.del(&key) { - Ok(_) => { - info!("Deleted corrupted event {} from store", key.to_string()); - } - Err(del_err) => { - error!("Failed to delete corrupted event {}: {}", key.to_string(), del_err); - } - } + } - metrics.increment_failed(); - } + batch_keys.push(key); + metrics.increment_processing(); + + // If the batch is full or enough time has passed since the last refresh, the batch will be processed + if batch_keys.len() >= batch_size || last_flush.elapsed() >= BATCH_TIMEOUT { + process_batch(&mut batch_keys, target, MAX_RETRIES, BASE_RETRY_DELAY, &metrics, &semaphore).await; + last_flush = Instant::now(); } } @@ -273,7 +262,6 @@ pub async fn stream_events_with_batching( /// # Notes /// This function processes a batch of events, sending each event to the target with retry async fn process_batch( - batch: &mut Vec>, batch_keys: &mut Vec, target: &dyn Target, max_retries: usize, @@ -281,8 +269,8 @@ async fn process_batch( metrics: &Arc, semaphore: &Arc, ) { - debug!("Processing batch of {} events for target: {}", batch.len(), target.name()); - if batch.is_empty() { + debug!("Processing batch of {} events for target: {}", batch_keys.len(), target.name()); + if batch_keys.is_empty() { return; } @@ -296,44 +284,42 @@ async fn process_batch( }; // Handle every event in the batch - for (_event, key) in batch.iter().zip(batch_keys.iter()) { + for key in batch_keys.iter() { let mut retry_count = 0; let mut success = false; // Retry logic while retry_count < max_retries && !success { - // After sending successfully, the event in the storage is deleted synchronously. match target.send_from_store(key.clone()).await { Ok(_) => { - info!("Successfully sent event for target: {}, Key: {}", target.name(), key.to_string()); + debug!("Successfully sent event for target: {}, Key: {}", target.name(), key.to_string()); success = true; metrics.increment_processed(); } - Err(e) => { - // Different retry strategies are adopted according to the error type - match &e { - TargetError::NotConnected => { - warn!("Target {} not connected, retrying...", target.name()); - retry_count += 1; - tokio::time::sleep(base_delay * (1 << retry_count)).await; // Exponential backoff - } - TargetError::Timeout(_) => { - warn!("Timeout for target {}, retrying...", target.name()); - retry_count += 1; - tokio::time::sleep(base_delay * (1 << retry_count)).await; - } - _ => { - // Permanent error, skip this event - error!("Permanent error for target {}: {}", target.name(), e); - metrics.increment_failed(); - break; - } + Err(e) => match &e { + TargetError::NotConnected => { + warn!("Target {} not connected, retrying...", target.name()); + retry_count += 1; + let jitter = Duration::from_millis(key.to_string().len() as u64 % 500); + let backoff = 1u32 << retry_count as u32; + tokio::time::sleep(base_delay * backoff + jitter).await; } - } + TargetError::Timeout(_) => { + warn!("Timeout for target {}, retrying...", target.name()); + retry_count += 1; + let jitter = Duration::from_millis(key.to_string().len() as u64 % 500); + let backoff = 1u32 << retry_count as u32; + tokio::time::sleep(base_delay * backoff + jitter).await; + } + _ => { + error!("Permanent error for target {}: {}", target.name(), e); + metrics.increment_failed(); + break; + } + }, } } - // Handle the situation where the maximum number of retry exhaustion is exhausted if retry_count >= max_retries && !success { warn!("Max retries exceeded for event {}, target: {}, skipping", key.to_string(), target.name()); metrics.increment_failed(); @@ -341,7 +327,6 @@ async fn process_batch( } // Clear processed batches - batch.clear(); batch_keys.clear(); // Release semaphore permission (via drop) diff --git a/crates/targets/Cargo.toml b/crates/targets/Cargo.toml index eb8d157dc..65fe2079c 100644 --- a/crates/targets/Cargo.toml +++ b/crates/targets/Cargo.toml @@ -28,5 +28,12 @@ url = { workspace = true } urlencoding = { workspace = true } uuid = { workspace = true, features = ["v4", "serde"] } +[dev-dependencies] +criterion = { workspace = true } + +[[bench]] +name = "queue_store_benchmark" +harness = false + [lints] workspace = true diff --git a/crates/targets/benches/queue_store_benchmark.rs b/crates/targets/benches/queue_store_benchmark.rs new file mode 100644 index 000000000..205c3bc4a --- /dev/null +++ b/crates/targets/benches/queue_store_benchmark.rs @@ -0,0 +1,94 @@ +use criterion::{BenchmarkId, Criterion, Throughput, criterion_group, criterion_main}; +use rustfs_targets::store::{QueueStore, Store}; +use serde::{Deserialize, Serialize}; +use std::hint::black_box; +use std::sync::Arc; +use uuid::Uuid; + +#[derive(Clone, Serialize, Deserialize)] +struct BenchEvent { + bucket: String, + key: String, + metadata: Vec<(String, String)>, + payload: String, +} + +fn bench_dir(prefix: &str) -> std::path::PathBuf { + std::env::temp_dir().join(format!("rustfs-targets-bench-{prefix}-{}", Uuid::new_v4())) +} + +fn build_payload(payload_len: usize) -> Vec { + let event = BenchEvent { + bucket: "bench-bucket".to_string(), + key: format!("objects/{payload_len}/file.json"), + metadata: (0..8) + .map(|idx| (format!("x-amz-meta-{idx}"), "benchmark-value".repeat(4))) + .collect(), + payload: "abcdefghijklmnopqrstuvwxyz0123456789".repeat(payload_len / 36 + 1)[..payload_len].to_string(), + }; + serde_json::to_vec(&event).unwrap() +} + +fn queue_store_write_benchmark(c: &mut Criterion) { + let mut group = c.benchmark_group("queue_store_put_raw"); + + for payload_size in [512usize, 8 * 1024, 64 * 1024] { + let payload = build_payload(payload_size); + group.throughput(Throughput::Bytes(payload.len() as u64)); + + for compress in [false, true] { + let dir = bench_dir(if compress { "put-compress" } else { "put-plain" }); + let store = QueueStore::::new_with_compression(&dir, 100_000, ".bench", compress); + store.open().unwrap(); + + group.bench_with_input( + BenchmarkId::new(if compress { "snap_on" } else { "snap_off" }, payload_size), + &payload, + |b, payload| { + b.iter(|| { + let key = store.put_raw(payload).unwrap(); + store.del(&key).unwrap(); + }); + }, + ); + + let _ = store.delete(); + } + } + + group.finish(); +} + +fn queue_store_read_benchmark(c: &mut Criterion) { + let mut group = c.benchmark_group("queue_store_get_raw"); + + for payload_size in [512usize, 8 * 1024, 64 * 1024] { + let payload = build_payload(payload_size); + group.throughput(Throughput::Bytes(payload.len() as u64)); + + for compress in [false, true] { + let dir = bench_dir(if compress { "get-compress" } else { "get-plain" }); + let store = Arc::new(QueueStore::::new_with_compression(&dir, 100_000, ".bench", compress)); + store.open().unwrap(); + let key = store.put_raw(&payload).unwrap(); + + group.bench_with_input( + BenchmarkId::new(if compress { "snap_on" } else { "snap_off" }, payload_size), + &(Arc::clone(&store), key), + |b, (store, key)| { + b.iter(|| { + let raw = store.get_raw(key).unwrap(); + black_box(raw); + }); + }, + ); + + let _ = store.delete(); + } + } + + group.finish(); +} + +criterion_group!(benches, queue_store_write_benchmark, queue_store_read_benchmark); +criterion_main!(benches); diff --git a/crates/targets/src/store.rs b/crates/targets/src/store.rs index ee3a616f1..b568f9ba8 100644 --- a/crates/targets/src/store.rs +++ b/crates/targets/src/store.rs @@ -13,20 +13,34 @@ // limitations under the License. use crate::error::StoreError; -use rustfs_config::DEFAULT_LIMIT; use rustfs_config::notify::{COMPRESS_EXT, DEFAULT_EXT}; +use rustfs_config::{DEFAULT_LIMIT, DEFAULT_TARGET_STORE_COMPRESS, ENV_TARGET_STORE_COMPRESS, EnableState}; use serde::{Serialize, de::DeserializeOwned}; use snap::raw::{Decoder, Encoder}; -use std::sync::{Arc, RwLock}; use std::{ collections::HashMap, marker::PhantomData, path::PathBuf, + sync::{ + Arc, RwLock, + atomic::{AtomicU64, Ordering}, + }, time::{SystemTime, UNIX_EPOCH}, }; use tracing::{debug, warn}; use uuid::Uuid; +fn resolve_queue_store_compression_from_env_value(value: Option<&str>) -> bool { + value + .and_then(|value| value.parse::().ok().map(|state| state.is_enabled())) + .unwrap_or(DEFAULT_TARGET_STORE_COMPRESS) +} + +fn queue_store_compression_enabled() -> bool { + let value = std::env::var(ENV_TARGET_STORE_COMPRESS).ok(); + resolve_queue_store_compression_from_env_value(value.as_deref()) +} + /// Represents a key for an entry in the store #[derive(Debug, Clone)] pub struct Key { @@ -63,21 +77,7 @@ impl Key { impl std::fmt::Display for Key { fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { - let name_part = if self.item_count > 1 { - format!("{}:{}", self.item_count, self.name) - } else { - self.name.clone() - }; - - let mut file_name = name_part; - if !self.extension.is_empty() { - file_name.push_str(&self.extension); - } - - if self.compress { - file_name.push_str(COMPRESS_EXT); - } - write!(f, "{file_name}") + f.write_str(&self.to_key_string()) } } @@ -123,6 +123,28 @@ pub fn parse_key(s: &str) -> Key { } } +pub fn ensure_store_entry_raw_readable( + store: &(dyn Store + Send), + key: &Key, +) -> Result +where + T: Send + Sync + 'static + Clone + Serialize, +{ + match store.get_raw(key) { + Ok(_) => Ok(true), + Err(StoreError::NotFound) => Ok(false), + Err(err) => { + match store.del(key) { + Ok(()) | Err(StoreError::NotFound) => {} + Err(del_err) => { + return Err(StoreError::Internal(format!("Failed to remove unreadable store entry {key}: {del_err}"))); + } + } + Err(err) + } + } +} + /// Trait for a store that can store and retrieve items of type T pub trait Store: Send + Sync where @@ -142,15 +164,24 @@ where /// Stores multiple items in a single batch fn put_multiple(&self, items: Vec) -> Result; + /// Stores raw bytes in a single entry. + fn put_raw(&self, data: &[u8]) -> Result; + /// Retrieves a single item by key fn get(&self, key: &Self::Key) -> Result; /// Retrieves multiple items by key fn get_multiple(&self, key: &Self::Key) -> Result, Self::Error>; + /// Retrieves the raw bytes stored for a key. + fn get_raw(&self, key: &Self::Key) -> Result, Self::Error>; + /// Deletes an item by key fn del(&self, key: &Self::Key) -> Result<(), Self::Error>; + /// Deletes the underlying store directory and clears all in-memory state. + fn delete(&self) -> Result<(), Self::Error>; + /// Lists all keys in the store fn list(&self) -> Vec; @@ -169,7 +200,10 @@ pub struct QueueStore { entry_limit: u64, directory: PathBuf, file_ext: String, + compress: bool, entries: Arc>>, // key -> modtime as unix nano + pending_entries: Arc, + fs_guard: Arc>, _phantom: PhantomData, } @@ -179,35 +213,70 @@ impl Clone for QueueStore { entry_limit: self.entry_limit, directory: self.directory.clone(), file_ext: self.file_ext.clone(), + compress: self.compress, entries: Arc::clone(&self.entries), + pending_entries: Arc::clone(&self.pending_entries), + fs_guard: Arc::clone(&self.fs_guard), _phantom: PhantomData, } } } +struct EntryReservation<'a> { + pending_entries: &'a AtomicU64, +} + +impl Drop for EntryReservation<'_> { + fn drop(&mut self) { + self.pending_entries.fetch_sub(1, Ordering::SeqCst); + } +} + impl QueueStore { /// Creates a new QueueStore pub fn new(directory: impl Into, limit: u64, ext: &str) -> Self { + Self::new_with_compression(directory, limit, ext, queue_store_compression_enabled()) + } + + /// Creates a new QueueStore with an explicit compression setting. + pub fn new_with_compression(directory: impl Into, limit: u64, ext: &str, compress: bool) -> Self { let file_ext = if ext.is_empty() { DEFAULT_EXT } else { ext }; + let entry_limit = if limit == 0 { DEFAULT_LIMIT } else { limit }; QueueStore { directory: directory.into(), - entry_limit: if limit == 0 { DEFAULT_LIMIT } else { limit }, + entry_limit, file_ext: file_ext.to_string(), - entries: Arc::new(RwLock::new(HashMap::with_capacity(limit as usize))), + compress, + entries: Arc::new(RwLock::new(HashMap::with_capacity(entry_limit as usize))), + pending_entries: Arc::new(AtomicU64::new(0)), + fs_guard: Arc::new(RwLock::new(())), _phantom: PhantomData, } } /// Returns the full path for a key fn file_path(&self, key: &Key) -> PathBuf { - self.directory.join(key.to_string()) + self.directory.join(key.to_key_string()) + } + + fn build_key(&self, item_count: usize) -> Key { + Key { + name: Uuid::new_v4().to_string(), + extension: self.file_ext.clone(), + item_count, + compress: self.compress, + } } /// Reads a file for the given key fn read_file(&self, key: &Key) -> Result, StoreError> { + let _fs_guard = self + .fs_guard + .read() + .map_err(|_| StoreError::Internal("Failed to acquire read lock on store filesystem".to_string()))?; let path = self.file_path(key); - debug!("Reading file for key: {},path: {}", key.to_string(), path.display()); + debug!("Reading file for key: {},path: {}", key, path.display()); let data = std::fs::read(&path).map_err(|e| { if e.kind() == std::io::ErrorKind::NotFound { StoreError::NotFound @@ -220,41 +289,89 @@ impl QueueStore { return Err(StoreError::NotFound); } - if key.compress { - let mut decoder = Decoder::new(); - decoder - .decompress_vec(&data) - .map_err(|e| StoreError::Compression(e.to_string())) - } else { - Ok(data) + if !key.compress { + return Ok(data); + } + + let mut decoder = Decoder::new(); + decoder + .decompress_vec(&data) + .map_err(|e| StoreError::Compression(e.to_string())) + } + + fn reserve_entry_slot(&self) -> Result, StoreError> { + loop { + let entries = self + .entries + .read() + .map_err(|_| StoreError::Internal("Failed to acquire read lock on entries".to_string()))?; + let entries_len = entries.len() as u64; + let pending = self.pending_entries.load(Ordering::SeqCst); + + if entries_len + pending >= self.entry_limit { + return Err(StoreError::LimitExceeded); + } + + if self + .pending_entries + .compare_exchange(pending, pending + 1, Ordering::SeqCst, Ordering::SeqCst) + .is_ok() + { + return Ok(EntryReservation { + pending_entries: self.pending_entries.as_ref(), + }); + } } } - /// Writes data to a file for the given key - fn write_file(&self, key: &Key, data: &[u8]) -> Result<(), StoreError> { + /// Writes data to a file for the given key. + fn write_file(&self, key: &Key, data: &[u8]) -> Result { let path = self.file_path(key); // Create directory if it doesn't exist if let Some(parent) = path.parent() { std::fs::create_dir_all(parent).map_err(StoreError::Io)?; } - let data = if key.compress { + if key.compress { let mut encoder = Encoder::new(); - encoder + let compressed = encoder .compress_vec(data) - .map_err(|e| StoreError::Compression(e.to_string()))? + .map_err(|e| StoreError::Compression(e.to_string()))?; + std::fs::write(&path, &compressed).map_err(StoreError::Io)?; } else { - data.to_vec() - }; - - std::fs::write(&path, &data).map_err(StoreError::Io)?; + std::fs::write(&path, data).map_err(StoreError::Io)?; + } let modified = SystemTime::now().duration_since(UNIX_EPOCH).unwrap_or_default().as_nanos() as i64; + debug!("Wrote event to store: {}", key); + Ok(modified) + } + + fn insert_entry(&self, key: &Key, modified: i64) -> Result<(), StoreError> { let mut entries = self .entries .write() .map_err(|_| StoreError::Internal("Failed to acquire write lock on entries".to_string()))?; - entries.insert(key.to_string(), modified); - debug!("Wrote event to store: {}", key.to_string()); + entries.insert(key.to_key_string(), modified); + Ok(()) + } + + fn remove_file_if_present(&self, key: &Key) -> Result<(), StoreError> { + let path = self.file_path(key); + match std::fs::remove_file(&path) { + Ok(()) => Ok(()), + Err(err) if err.kind() == std::io::ErrorKind::NotFound => Ok(()), + Err(err) => Err(StoreError::Io(err)), + } + } + + fn write_and_index(&self, key: &Key, data: &[u8]) -> Result<(), StoreError> { + let modified = self.write_file(key, data)?; + if let Err(err) = self.insert_entry(key, modified) { + self.remove_file_if_present(key).map_err(|cleanup_err| { + StoreError::Internal(format!("Failed to index store entry {key}: {err}; cleanup failed: {cleanup_err}")) + })?; + return Err(err); + } Ok(()) } } @@ -267,15 +384,20 @@ where type Key = Key; fn open(&self) -> Result<(), Self::Error> { + let _fs_guard = self + .fs_guard + .write() + .map_err(|_| StoreError::Internal("Failed to acquire write lock on store filesystem".to_string()))?; std::fs::create_dir_all(&self.directory).map_err(StoreError::Io)?; - let entries = std::fs::read_dir(&self.directory).map_err(StoreError::Io)?; - // Get the write lock to update the internal state + let dir_entries = std::fs::read_dir(&self.directory).map_err(StoreError::Io)?; let mut entries_map = self .entries .write() .map_err(|_| StoreError::Internal("Failed to acquire write lock on entries".to_string()))?; - for entry in entries { + self.pending_entries.store(0, Ordering::SeqCst); + entries_map.clear(); + for entry in dir_entries { let entry = entry.map_err(StoreError::Io)?; let metadata = entry.metadata().map_err(StoreError::Io)?; if metadata.is_file() { @@ -292,71 +414,47 @@ where } fn put(&self, item: Arc) -> Result { - // Check storage limits - { - let entries = self - .entries - .read() - .map_err(|_| StoreError::Internal("Failed to acquire read lock on entries".to_string()))?; - - if entries.len() as u64 >= self.entry_limit { - return Err(StoreError::LimitExceeded); - } - } - - let uuid = Uuid::new_v4(); - let key = Key { - name: uuid.to_string(), - extension: self.file_ext.clone(), - item_count: 1, - compress: true, - }; - + let _fs_guard = self + .fs_guard + .read() + .map_err(|_| StoreError::Internal("Failed to acquire read lock on store filesystem".to_string()))?; + let _reservation = self.reserve_entry_slot()?; + let key = self.build_key(1); let data = serde_json::to_vec(&*item).map_err(|e| StoreError::Serialization(e.to_string()))?; - self.write_file(&key, &data)?; + self.write_and_index(&key, &data)?; Ok(key) } fn put_multiple(&self, items: Vec) -> Result { - // Check storage limits - { - let entries = self - .entries - .read() - .map_err(|_| StoreError::Internal("Failed to acquire read lock on entries".to_string()))?; - - if entries.len() as u64 >= self.entry_limit { - return Err(StoreError::LimitExceeded); - } - } if items.is_empty() { - // Or return an error, or a special key? return Err(StoreError::Internal("Cannot put_multiple with empty items list".to_string())); } - let uuid = Uuid::new_v4(); - let key = Key { - name: uuid.to_string(), - extension: self.file_ext.clone(), - item_count: items.len(), - compress: true, - }; + let _fs_guard = self + .fs_guard + .read() + .map_err(|_| StoreError::Internal("Failed to acquire read lock on store filesystem".to_string()))?; + let _reservation = self.reserve_entry_slot()?; + let key = self.build_key(items.len()); - // Serialize all items into a single Vec - // This current approach for get_multiple/put_multiple assumes items are concatenated JSON objects. - // This might be problematic for deserialization if not handled carefully. - // A better approach for multiple items might be to store them as a JSON array `Vec`. - // For now, sticking to current logic of concatenating. let mut buffer = Vec::new(); for item in items { - // If items are Vec, and Event is large, this could be inefficient. - // The current get_multiple deserializes one by one. - let item_data = serde_json::to_vec(&item).map_err(|e| StoreError::Serialization(e.to_string()))?; - buffer.extend_from_slice(&item_data); - // If using JSON array: buffer = serde_json::to_vec(&items)? + serde_json::to_writer(&mut buffer, &item).map_err(|e| StoreError::Serialization(e.to_string()))?; } - self.write_file(&key, &buffer)?; + self.write_and_index(&key, &buffer)?; + + Ok(key) + } + + fn put_raw(&self, data: &[u8]) -> Result { + let _fs_guard = self + .fs_guard + .read() + .map_err(|_| StoreError::Internal("Failed to acquire read lock on store filesystem".to_string()))?; + let _reservation = self.reserve_entry_slot()?; + let key = self.build_key(1); + self.write_and_index(&key, data)?; Ok(key) } @@ -373,8 +471,8 @@ where } fn get_multiple(&self, key: &Self::Key) -> Result, Self::Error> { - debug!("Reading items from store for key: {}", key.to_string()); - let data = self.read_file(key)?; + debug!("Reading items from store for key: {}", key); + let data = self.get_raw(key)?; if data.is_empty() { return Err(StoreError::Deserialization("Cannot deserialize empty data".to_string())); } @@ -404,7 +502,7 @@ where warn!( "Expected {} items for key {}, but only found {}. Possible data corruption or incorrect item_count.", key.item_count, - key.to_string(), + key, items.len() ); // Depending on strictness, this could be an error. @@ -426,20 +524,24 @@ where Ok(items) } + fn get_raw(&self, key: &Self::Key) -> Result, Self::Error> { + self.read_file(key) + } + fn del(&self, key: &Self::Key) -> Result<(), Self::Error> { + let _fs_guard = self + .fs_guard + .read() + .map_err(|_| StoreError::Internal("Failed to acquire read lock on store filesystem".to_string()))?; let path = self.file_path(key); - std::fs::remove_file(&path).map_err(|e| { - if e.kind() == std::io::ErrorKind::NotFound { - // If file not found, still try to remove from entries map in case of inconsistency - warn!( - "File not found for key {} during del, but proceeding to remove from entries map.", - key.to_string() - ); - StoreError::NotFound - } else { - StoreError::Io(e) + match std::fs::remove_file(&path) { + Ok(()) => {} + Err(e) if e.kind() == std::io::ErrorKind::NotFound => { + // File already gone — still clean up the entries map to avoid stale keys. + warn!("File not found for key {} during del, cleaning up entries map.", key); } - })?; + Err(e) => return Err(StoreError::Io(e)), + } // Get the write lock to update the internal state let mut entries = self @@ -447,15 +549,32 @@ where .write() .map_err(|_| StoreError::Internal("Failed to acquire write lock on entries".to_string()))?; - if entries.remove(&key.to_string()).is_none() { - // Key was not in the map, could be an inconsistency or already deleted. - // This is not necessarily an error if the file deletion succeeded or was NotFound. + if entries.remove(&key.to_key_string()).is_none() { debug!("Key {} not found in entries map during del, might have been already removed.", key); } debug!("Deleted event from store: {}", key.to_string()); Ok(()) } + fn delete(&self) -> Result<(), Self::Error> { + let _fs_guard = self + .fs_guard + .write() + .map_err(|_| StoreError::Internal("Failed to acquire write lock on store filesystem".to_string()))?; + let mut entries = self + .entries + .write() + .map_err(|_| StoreError::Internal("Failed to acquire write lock on entries".to_string()))?; + entries.clear(); + self.pending_entries.store(0, Ordering::SeqCst); + + match std::fs::remove_dir_all(&self.directory) { + Ok(()) => Ok(()), + Err(err) if err.kind() == std::io::ErrorKind::NotFound => Ok(()), + Err(err) => Err(StoreError::Io(err)), + } + } + fn list(&self) -> Vec { // Get the read lock to read the internal state let entries = match self.entries.read() { @@ -492,3 +611,136 @@ where Box::new(self.clone()) as Box + Send + Sync> } } + +#[cfg(test)] +mod tests { + use super::*; + use std::{ + sync::{Arc, Barrier}, + thread, + }; + + fn temp_store_dir(name: &str) -> PathBuf { + std::env::temp_dir().join(format!("rustfs-targets-{name}-{}", Uuid::new_v4())) + } + + #[test] + fn resolve_queue_store_compression_defaults_to_true() { + assert!(resolve_queue_store_compression_from_env_value(None)); + } + + #[test] + fn resolve_queue_store_compression_respects_disabled_env_value() { + assert!(!resolve_queue_store_compression_from_env_value(Some("off"))); + assert!(!resolve_queue_store_compression_from_env_value(Some("false"))); + } + + #[test] + fn put_uses_store_compression_setting_in_key() { + let dir = temp_store_dir("put-key"); + let store = QueueStore::::new_with_compression(&dir, 8, ".test", false); + store.open().unwrap(); + + let key = store.put(Arc::new("payload".to_string())).unwrap(); + + assert!(!key.compress); + assert!(store.file_path(&key).exists()); + + let _ = std::fs::remove_dir_all(dir); + } + + #[test] + fn parse_key_round_trips_batch_and_compression_suffixes() { + let key = Key { + name: "event-id".to_string(), + extension: ".json".to_string(), + item_count: 3, + compress: true, + }; + + let parsed = parse_key(&key.to_key_string()); + + assert_eq!(parsed.name, key.name); + assert_eq!(parsed.extension, key.extension); + assert_eq!(parsed.item_count, key.item_count); + assert_eq!(parsed.compress, key.compress); + } + + #[test] + fn put_raw_and_get_raw_round_trip_bytes() { + let dir = temp_store_dir("raw-roundtrip"); + let store = QueueStore::::new_with_compression(&dir, 8, ".test", true); + store.open().unwrap(); + + let payload = br#"{"kind":"notify","bucket":"demo","key":"alpha.txt"}"#; + let key = store.put_raw(payload).unwrap(); + let raw = store.get_raw(&key).unwrap(); + + assert_eq!(raw, payload); + + let _ = store.delete(); + } + + #[test] + fn delete_removes_directory_and_clears_entries() { + let dir = temp_store_dir("delete-store"); + let store = QueueStore::::new_with_compression(&dir, 8, ".test", false); + store.open().unwrap(); + let _ = store.put(Arc::new("payload".to_string())).unwrap(); + + store.delete().unwrap(); + + assert!(store.list().is_empty()); + assert!(!dir.exists()); + } + + #[test] + fn put_enforces_entry_limit() { + let dir = temp_store_dir("limit"); + let store = QueueStore::::new_with_compression(&dir, 1, ".test", false); + store.open().unwrap(); + + let _ = store.put(Arc::new("first".to_string())).unwrap(); + let err = store.put(Arc::new("second".to_string())).unwrap_err(); + + assert!(matches!(err, StoreError::LimitExceeded)); + + let _ = store.delete(); + } + + #[test] + fn concurrent_put_raw_respects_entry_limit() { + let dir = temp_store_dir("concurrent-limit"); + let store = Arc::new(QueueStore::::new_with_compression(&dir, 1, ".test", true)); + store.open().unwrap(); + + let start = Arc::new(Barrier::new(4)); + let mut handles = Vec::new(); + + for idx in 0..4 { + let store = Arc::clone(&store); + let start = Arc::clone(&start); + handles.push(thread::spawn(move || { + let payload = vec![b'x'; 32 * 1024 + idx]; + start.wait(); + store.put_raw(&payload) + })); + } + + let mut successes = 0; + let mut limit_errors = 0; + for handle in handles { + match handle.join().unwrap() { + Ok(_) => successes += 1, + Err(StoreError::LimitExceeded) => limit_errors += 1, + Err(err) => panic!("unexpected error: {err}"), + } + } + + assert_eq!(successes, 1); + assert_eq!(limit_errors, 3); + assert_eq!(store.len(), 1); + + let _ = store.delete(); + } +} diff --git a/crates/targets/src/target/mod.rs b/crates/targets/src/target/mod.rs index 1a097a054..dc16dcae2 100644 --- a/crates/targets/src/target/mod.rs +++ b/crates/targets/src/target/mod.rs @@ -21,6 +21,8 @@ use serde::de::DeserializeOwned; use serde::{Deserialize, Serialize}; use std::fmt::Formatter; use std::sync::Arc; +use std::time::{SystemTime, UNIX_EPOCH}; +use tracing::warn; pub mod mqtt; pub mod webhook; @@ -45,14 +47,43 @@ where /// Saves an event (either sends it immediately or stores it for later) async fn save(&self, event: Arc>) -> Result<(), TargetError>; - /// Sends an event from the store - async fn send_from_store(&self, key: Key) -> Result<(), TargetError>; + /// Sends an event from the store using the queued raw body and metadata. + async fn send_raw_from_store(&self, key: Key, body: Vec, meta: QueuedPayloadMeta) -> Result<(), TargetError>; + + /// Sends an event from the store. + async fn send_from_store(&self, key: Key) -> Result<(), TargetError> { + let store = self + .store() + .ok_or_else(|| TargetError::Configuration("No store configured".to_string()))?; + + let raw = match store.get_raw(&key) { + Ok(raw) => raw, + Err(StoreError::NotFound) => return Ok(()), + Err(err) => return Err(TargetError::Storage(format!("Failed to read queued payload from store: {err}"))), + }; + + let queued = match QueuedPayload::decode(&raw) { + Ok(queued) => queued, + Err(err) => { + delete_stored_payload(store, &key).map_err(|delete_err| { + TargetError::Storage(format!( + "Failed to delete invalid queued payload {key} after decode error '{err}': {delete_err}" + )) + })?; + warn!("Dropped invalid queued payload {key}: {err}"); + return Ok(()); + } + }; + + self.send_raw_from_store(key.clone(), queued.body, queued.meta).await?; + delete_stored_payload(store, &key) + } /// Closes the target and releases resources async fn close(&self) -> Result<(), TargetError>; /// Returns the store associated with the target (if any) - fn store(&self) -> Option<&(dyn Store, Error = StoreError, Key = Key> + Send + Sync)>; + fn store(&self) -> Option<&(dyn Store + Send + Sync)>; /// Returns the type of the target fn clone_dyn(&self) -> Box + Send + Sync>; @@ -78,6 +109,106 @@ where pub data: E, } +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct QueuedPayloadMeta { + pub event_name: EventName, + pub bucket_name: String, + pub object_name: String, + pub content_type: String, + pub queued_at_unix_ms: u64, + pub payload_len: usize, +} + +impl QueuedPayloadMeta { + pub fn new( + event_name: EventName, + bucket_name: String, + object_name: String, + content_type: impl Into, + payload_len: usize, + ) -> Self { + Self { + event_name, + bucket_name, + object_name, + content_type: content_type.into(), + queued_at_unix_ms: SystemTime::now().duration_since(UNIX_EPOCH).unwrap_or_default().as_millis() as u64, + payload_len, + } + } + + pub fn best_effort_preview(&self, body: &[u8], limit: usize) -> String { + if limit == 0 || body.is_empty() { + return String::new(); + } + + let slice = &body[..body.len().min(limit)]; + match std::str::from_utf8(slice) { + Ok(text) => { + if body.len() > limit { + format!("{text}...") + } else { + text.to_string() + } + } + Err(_) => format!("<{} bytes binary>", body.len()), + } + } +} + +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct QueuedPayload { + pub meta: QueuedPayloadMeta, + pub body: Vec, +} + +impl QueuedPayload { + const MAGIC: [u8; 4] = *b"RQP1"; + + pub fn new(meta: QueuedPayloadMeta, body: Vec) -> Self { + Self { meta, body } + } + + pub fn encode(&self) -> Result, TargetError> { + let meta = serde_json::to_vec(&self.meta) + .map_err(|err| TargetError::Serialization(format!("Failed to serialize queued payload metadata: {err}")))?; + let meta_len = u32::try_from(meta.len()) + .map_err(|_| TargetError::Serialization("Queued payload metadata is too large".to_string()))?; + + let mut out = Vec::with_capacity(Self::MAGIC.len() + 4 + meta.len() + self.body.len()); + out.extend_from_slice(&Self::MAGIC); + out.extend_from_slice(&meta_len.to_le_bytes()); + out.extend_from_slice(&meta); + out.extend_from_slice(&self.body); + Ok(out) + } + + pub fn decode(raw: &[u8]) -> Result { + if raw.len() < Self::MAGIC.len() + 4 { + return Err(TargetError::Serialization("Queued payload is too short".to_string())); + } + if raw[..Self::MAGIC.len()] != Self::MAGIC { + return Err(TargetError::Serialization("Queued payload magic mismatch".to_string())); + } + + let mut meta_len_bytes = [0u8; 4]; + meta_len_bytes.copy_from_slice(&raw[Self::MAGIC.len()..Self::MAGIC.len() + 4]); + let meta_len = u32::from_le_bytes(meta_len_bytes) as usize; + let meta_start = Self::MAGIC.len() + 4; + let meta_end = meta_start + meta_len; + + if meta_end > raw.len() { + return Err(TargetError::Serialization("Queued payload metadata length exceeds input".to_string())); + } + + let meta = serde_json::from_slice(&raw[meta_start..meta_end]) + .map_err(|err| TargetError::Serialization(format!("Failed to deserialize queued payload metadata: {err}")))?; + let body = raw[meta_end..].to_vec(); + + Ok(Self { meta, body }) + } +} + /// The `ChannelTargetType` enum represents the different types of channel Target /// used in the notification system. /// @@ -187,3 +318,45 @@ pub fn decode_object_name(encoded: &str) -> Result { .map(|s| s.into_owned()) .map_err(|e| TargetError::Encoding(format!("Failed to decode object key: {e}"))) } + +pub(crate) fn delete_stored_payload( + store: &(dyn Store + Send + Sync), + key: &Key, +) -> Result<(), TargetError> { + match store.del(key) { + Ok(()) | Err(StoreError::NotFound) => Ok(()), + Err(err) => Err(TargetError::Storage(format!("Failed to delete event from store: {err}"))), + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn queued_payload_round_trips_meta_and_body() { + let meta = QueuedPayloadMeta::new( + EventName::ObjectCreatedPut, + "bucket-a".to_string(), + "folder/object.txt".to_string(), + "application/json", + 12, + ); + let payload = QueuedPayload::new(meta.clone(), br#"{"ok":true}"#.to_vec()); + + let encoded = payload.encode().unwrap(); + let decoded = QueuedPayload::decode(&encoded).unwrap(); + + assert_eq!(decoded.meta.event_name, meta.event_name); + assert_eq!(decoded.meta.bucket_name, meta.bucket_name); + assert_eq!(decoded.meta.object_name, meta.object_name); + assert_eq!(decoded.meta.content_type, meta.content_type); + assert_eq!(decoded.body, br#"{"ok":true}"#); + } + + #[test] + fn queued_payload_decode_rejects_invalid_magic() { + let err = QueuedPayload::decode(b"bad-payload").unwrap_err(); + assert!(err.to_string().contains("magic") || err.to_string().contains("short")); + } +} diff --git a/crates/targets/src/target/mqtt.rs b/crates/targets/src/target/mqtt.rs index eb876d9c1..df32b2fff 100644 --- a/crates/targets/src/target/mqtt.rs +++ b/crates/targets/src/target/mqtt.rs @@ -17,7 +17,7 @@ use crate::{ arn::TargetID, error::TargetError, store::{Key, QueueStore, Store}, - target::{ChannelTargetType, EntityTarget, TargetType}, + target::{ChannelTargetType, EntityTarget, QueuedPayload, QueuedPayloadMeta, TargetType}, }; use async_trait::async_trait; use rumqttc::{AsyncClient, ConnectionError, EventLoop, MqttOptions, Outgoing, Packet, QoS, mqttbytes::Error as MqttBytesError}; @@ -25,6 +25,7 @@ use serde::Serialize; use serde::de::DeserializeOwned; use std::sync::Arc; use std::{ + marker::PhantomData, path::PathBuf, sync::atomic::{AtomicBool, Ordering}, time::Duration, @@ -110,9 +111,10 @@ where id: TargetID, args: MQTTArgs, client: Arc>>, - store: Option, Error = StoreError, Key = Key> + Send + Sync>>, + store: Option + Send + Sync>>, connected: Arc, bg_task_manager: Arc, + _phantom: PhantomData, } impl MQTTTarget @@ -135,7 +137,7 @@ where TargetType::NotifyEvent => rustfs_config::notify::NOTIFY_STORE_EXTENSION, }; - let store = QueueStore::>::new(specific_queue_path, args.queue_limit, extension); + let store = QueueStore::::new(specific_queue_path, args.queue_limit, extension); if let Err(e) = store.open() { error!( target_id = %target_id, @@ -144,7 +146,7 @@ where ); return Err(TargetError::Storage(format!("{e}"))); } - Some(Box::new(store) as Box, Error = StoreError, Key = Key> + Send + Sync>) + Some(Box::new(store) as Box + Send + Sync>) } else { None }; @@ -157,13 +159,14 @@ where }); info!(target_id = %target_id, "MQTT target created"); - Ok(MQTTTarget { + Ok(MQTTTarget:: { id: target_id, args, client: Arc::new(Mutex::new(None)), store: queue_store, connected: Arc::new(AtomicBool::new(false)), bg_task_manager, + _phantom: PhantomData, }) } @@ -251,14 +254,7 @@ where } } - #[instrument(skip(self, event), fields(target_id = %self.id))] - async fn send(&self, event: &EntityTarget) -> Result<(), TargetError> { - let client_guard = self.client.lock().await; - let client = client_guard - .as_ref() - .ok_or_else(|| TargetError::Configuration("MQTT client not initialized".to_string()))?; - - // Decode form-urlencoded object name + fn build_queued_payload(&self, event: &EntityTarget) -> Result { let object_name = crate::target::decode_object_name(&event.object_name)?; let key = format!("{}/{}", event.bucket_name, object_name); @@ -269,14 +265,35 @@ where records: vec![event.clone()], }; - let data = serde_json::to_vec(&log).map_err(|e| TargetError::Serialization(format!("Failed to serialize event: {e}")))?; + let body = serde_json::to_vec(&log).map_err(|e| TargetError::Serialization(format!("Failed to serialize event: {e}")))?; + let meta = QueuedPayloadMeta::new( + event.event_name, + event.bucket_name.clone(), + event.object_name.clone(), + "application/json", + body.len(), + ); + Ok(QueuedPayload::new(meta, body)) + } - let data_string = String::from_utf8(data.clone()) - .map_err(|e| TargetError::Encoding(format!("Failed to convert event data to UTF-8: {e}")))?; - debug!("Sending event to mqtt target: {}, event log: {}", self.id, data_string); + #[instrument(skip(self, body, meta), fields(target_id = %self.id))] + async fn send_body(&self, body: Vec, meta: &QueuedPayloadMeta) -> Result<(), TargetError> { + let client_guard = self.client.lock().await; + let client = client_guard + .as_ref() + .ok_or_else(|| TargetError::Configuration("MQTT client not initialized".to_string()))?; + + debug!( + target = %self.id, + bucket = %meta.bucket_name, + object = %meta.object_name, + event = %meta.event_name, + preview = %meta.best_effort_preview(&body, 256), + "Sending MQTT payload" + ); client - .publish(&self.args.topic, self.args.qos, false, data) + .publish(&self.args.topic, self.args.qos, false, body) .await .map_err(|e| { if e.to_string().contains("Connection") || e.to_string().contains("Timeout") { @@ -293,13 +310,14 @@ where } pub fn clone_target(&self) -> Box + Send + Sync> { - Box::new(MQTTTarget { + Box::new(MQTTTarget:: { id: self.id.clone(), args: self.args.clone(), client: self.client.clone(), store: self.store.as_ref().map(|s| s.boxed_clone()), connected: self.connected.clone(), bg_task_manager: self.bg_task_manager.clone(), + _phantom: PhantomData, }) } } @@ -494,11 +512,15 @@ where #[instrument(skip(self, event), fields(target_id = %self.id))] async fn save(&self, event: Arc>) -> Result<(), TargetError> { + let queued = self.build_queued_payload(&event)?; + if let Some(store) = &self.store { debug!(target_id = %self.id, "Event saved to store start"); - // If store is configured, ONLY put the event into the store. - // Do NOT send it directly here. - match store.put(event.clone()) { + match store.put_raw( + &queued + .encode() + .map_err(|e| TargetError::Storage(format!("Failed to encode queued payload: {e}")))?, + ) { Ok(_) => { debug!(target_id = %self.id, "Event saved to store for MQTT target successfully."); Ok(()) @@ -516,7 +538,7 @@ where if !self.connected.load(Ordering::SeqCst) { warn!(target_id = %self.id, "Attempting to send directly but not connected; trying to init."); // Call the struct's init method, not the trait's default - match MQTTTarget::init(self).await { + match MQTTTarget::::init(self).await { Ok(_) => debug!(target_id = %self.id, "MQTT target initialized successfully."), Err(e) => { error!(target_id = %self.id, error = %e, "Failed to initialize MQTT target."); @@ -528,13 +550,13 @@ where return Err(TargetError::NotConnected); } } - self.send(&event).await + self.send_body(queued.body, &queued.meta).await } } - #[instrument(skip(self), fields(target_id = %self.id))] - async fn send_from_store(&self, key: Key) -> Result<(), TargetError> { - debug!(target_id = %self.id, ?key, "Attempting to send event from store with key."); + #[instrument(skip(self, body, meta), fields(target_id = %self.id))] + async fn send_raw_from_store(&self, key: Key, body: Vec, meta: QueuedPayloadMeta) -> Result<(), TargetError> { + debug!(target_id = %self.id, ?key, "Attempting to send queued payload from store."); if !self.is_enabled() { return Err(TargetError::Disabled); @@ -542,7 +564,7 @@ where if !self.connected.load(Ordering::SeqCst) { warn!(target_id = %self.id, "Not connected; trying to init before sending from store."); - match MQTTTarget::init(self).await { + match MQTTTarget::::init(self).await { Ok(_) => debug!(target_id = %self.id, "MQTT target initialized successfully."), Err(e) => { error!(target_id = %self.id, error = %e, "Failed to initialize MQTT target."); @@ -555,33 +577,8 @@ where } } - let store = self - .store - .as_ref() - .ok_or_else(|| TargetError::Configuration("No store configured".to_string()))?; - - let event = match store.get(&key) { - Ok(event) => { - debug!(target_id = %self.id, ?key, "Retrieved event from store for sending."); - event - } - Err(StoreError::NotFound) => { - // Assuming NotFound takes the key - debug!(target_id = %self.id, ?key, "Event not found in store for sending."); - return Ok(()); - } - Err(e) => { - error!( - target_id = %self.id, - error = %e, - "Failed to get event from store" - ); - return Err(TargetError::Storage(format!("Failed to get event from store: {e}"))); - } - }; - debug!(target_id = %self.id, ?key, "Sending event from store."); - if let Err(e) = self.send(&event).await { + if let Err(e) = self.send_body(body, &meta).await { if matches!(e, TargetError::NotConnected) { warn!(target_id = %self.id, "Failed to send event from store: Not connected. Event remains in store."); return Err(TargetError::NotConnected); @@ -589,22 +586,7 @@ where error!(target_id = %self.id, error = %e, "Failed to send event from store with an unexpected error."); return Err(e); } - debug!(target_id = %self.id, ?key, "Event sent from store successfully. deleting from store. "); - - match store.del(&key) { - Ok(_) => { - debug!(target_id = %self.id, ?key, "Event deleted from store after successful send.") - } - Err(StoreError::NotFound) => { - debug!(target_id = %self.id, ?key, "Event already deleted from store."); - } - Err(e) => { - error!(target_id = %self.id, error = %e, "Failed to delete event from store after send."); - return Err(TargetError::Storage(format!("Failed to delete event from store: {e}"))); - } - } - - debug!(target_id = %self.id, ?key, "Event deleted from store."); + debug!(target_id = %self.id, ?key, "Event sent from store successfully."); Ok(()) } @@ -637,7 +619,7 @@ where Ok(()) } - fn store(&self) -> Option<&(dyn Store, Error = StoreError, Key = Key> + Send + Sync)> { + fn store(&self) -> Option<&(dyn Store + Send + Sync)> { self.store.as_deref() } @@ -651,7 +633,7 @@ where return Ok(()); } // Call the internal init logic - MQTTTarget::init(self).await + MQTTTarget::::init(self).await } fn is_enabled(&self) -> bool { diff --git a/crates/targets/src/target/webhook.rs b/crates/targets/src/target/webhook.rs index 525cf5b9f..a34b16e10 100644 --- a/crates/targets/src/target/webhook.rs +++ b/crates/targets/src/target/webhook.rs @@ -17,7 +17,7 @@ use crate::{ arn::TargetID, error::TargetError, store::{Key, QueueStore, Store}, - target::{ChannelTargetType, EntityTarget, TargetType}, + target::{ChannelTargetType, EntityTarget, QueuedPayload, QueuedPayloadMeta, TargetType}, }; use async_trait::async_trait; use reqwest::{Client, StatusCode, Url}; @@ -26,6 +26,7 @@ use rustfs_config::notify::NOTIFY_STORE_EXTENSION; use serde::Serialize; use serde::de::DeserializeOwned; use std::{ + marker::PhantomData, path::PathBuf, sync::{ Arc, @@ -33,7 +34,6 @@ use std::{ }, time::Duration, }; -use tokio::net::lookup_host; use tokio::sync::mpsc; use tracing::{debug, error, info, instrument, warn}; @@ -105,10 +105,10 @@ where args: WebhookArgs, http_client: Arc, // Add Send + Sync constraints to ensure thread safety - store: Option, Error = StoreError, Key = Key> + Send + Sync>>, + store: Option + Send + Sync>>, initialized: AtomicBool, - addr: String, cancel_sender: mpsc::Sender<()>, + _phantom: PhantomData, } impl WebhookTarget @@ -117,14 +117,14 @@ where { /// Clones the WebhookTarget, creating a new instance with the same configuration pub fn clone_box(&self) -> Box + Send + Sync> { - Box::new(WebhookTarget { + Box::new(WebhookTarget:: { id: self.id.clone(), args: self.args.clone(), http_client: Arc::clone(&self.http_client), store: self.store.as_ref().map(|s| s.boxed_clone()), initialized: AtomicBool::new(self.initialized.load(Ordering::SeqCst)), - addr: self.addr.clone(), cancel_sender: self.cancel_sender.clone(), + _phantom: PhantomData, }) } @@ -149,7 +149,7 @@ where TargetType::NotifyEvent => NOTIFY_STORE_EXTENSION, }; - let store = QueueStore::>::new(queue_dir, args.queue_limit, extension); + let store = QueueStore::::new(queue_dir, args.queue_limit, extension); if let Err(e) = store.open() { error!("Failed to open store for Webhook target {}: {}", target_id.id, e); @@ -157,32 +157,22 @@ where } // Make sure that the Store trait implemented by QueueStore matches the expected error type - Some(Box::new(store) as Box, Error = StoreError, Key = Key> + Send + Sync>) + Some(Box::new(store) as Box + Send + Sync>) } else { None }; - // resolved address - let addr = { - let host = args.endpoint.host_str().unwrap_or("localhost"); - let port = args - .endpoint - .port() - .unwrap_or_else(|| if args.endpoint.scheme() == "https" { 443 } else { 80 }); - format!("{host}:{port}") - }; - // Create a cancel channel let (cancel_sender, _) = mpsc::channel(1); info!(target_id = %target_id.id, "Webhook target created"); - Ok(WebhookTarget { + Ok(WebhookTarget:: { id: target_id, args, http_client, store: queue_store, initialized: AtomicBool::new(false), - addr, cancel_sender, + _phantom: PhantomData, }) } @@ -226,53 +216,80 @@ where .map_err(|e| TargetError::Configuration(format!("Failed to build HTTP client: {e}"))) } - async fn init(&self) -> Result<(), TargetError> { - // Use CAS operations to ensure thread-safe initialization - if !self.initialized.load(Ordering::SeqCst) { - // Check the connection - match self.is_active().await { - Ok(true) => { - info!("Webhook target {} is active", self.id); - } - Ok(false) => { - return Err(TargetError::NotConnected); - } - Err(e) => { - error!("Failed to check if Webhook target {} is active: {}", self.id, e); - return Err(e); + async fn init_inner(&self) -> Result<(), TargetError> { + if self.initialized.load(Ordering::SeqCst) { + return Ok(()); + } + + // HTTP HEAD probe: verifies the full request path (proxy, TLS, firewall) + // unlike TCP connect which can't detect proxy issues. + let probe_timeout = Duration::from_secs(5); + match tokio::time::timeout(probe_timeout, self.http_client.head(self.args.endpoint.as_str()).send()).await { + Ok(Ok(resp)) => { + let status = resp.status(); + if status.is_success() || status == StatusCode::NOT_FOUND { + // NOT_FOUND is acceptable for HEAD probes — the endpoint may not + // exist as a HEAD route, but the server is reachable. + debug!("Webhook target {} HEAD probe returned {}", self.id, status); + } else if status == StatusCode::METHOD_NOT_ALLOWED { + // Server is reachable but doesn't support HEAD — still valid. + debug!("Webhook target {} HEAD probe: METHOD_NOT_ALLOWED (reachable)", self.id); + } else { + warn!("Webhook target {} HEAD probe returned {}", self.id, status); } } - self.initialized.store(true, Ordering::SeqCst); - info!("Webhook target {} initialized", self.id); + Ok(Err(e)) => { + // Connection-level error (DNS, TLS, refused, timeout) + return Err(if e.is_timeout() || e.is_connect() { + TargetError::NotConnected + } else { + TargetError::Network(format!("Webhook HEAD probe failed: {e}")) + }); + } + Err(_) => { + return Err(TargetError::Timeout("Webhook HEAD probe timed out".to_string())); + } } + + self.initialized.store(true, Ordering::SeqCst); + info!("Webhook target {} initialized", self.id); Ok(()) } - async fn send(&self, event: &EntityTarget) -> Result<(), TargetError> { - info!("Webhook Sending event to webhook target: {}", self.id); - // Decode form-urlencoded object name + fn build_queued_payload(&self, event: &EntityTarget) -> Result { let object_name = crate::target::decode_object_name(&event.object_name)?; - let key = format!("{}/{}", event.bucket_name, object_name); - let log = TargetLog { event_name: event.event_name, key, records: vec![event.data.clone()], }; + let body = serde_json::to_vec(&log).map_err(|e| TargetError::Serialization(format!("Failed to serialize event: {e}")))?; + let meta = QueuedPayloadMeta::new( + event.event_name, + event.bucket_name.clone(), + event.object_name.clone(), + "application/json", + body.len(), + ); + Ok(QueuedPayload::new(meta, body)) + } - let data = serde_json::to_vec(&log).map_err(|e| TargetError::Serialization(format!("Failed to serialize event: {e}")))?; + async fn send_body(&self, body: Vec, meta: &QueuedPayloadMeta) -> Result<(), TargetError> { + info!("Webhook sending queued payload to target: {}", self.id); + debug!( + target = %self.id, + bucket = %meta.bucket_name, + object = %meta.object_name, + event = %meta.event_name, + preview = %meta.best_effort_preview(&body, 256), + "Sending webhook payload" + ); - // Vec Convert to String - let data_string = String::from_utf8(data.clone()) - .map_err(|e| TargetError::Encoding(format!("Failed to convert event data to UTF-8: {e}")))?; - debug!("Sending event to webhook target: {}, event log: {}", self.id, data_string); - - // build request let mut req_builder = self .http_client .post(self.args.endpoint.as_str()) - .header("Content-Type", "application/json"); + .header("Content-Type", meta.content_type.as_str()); if !self.args.auth_token.is_empty() { // Split auth_token string to check if the authentication type is included @@ -293,7 +310,7 @@ where } // Send a request - let resp = req_builder.body(data).send().await.map_err(|e| { + let resp = req_builder.body(body).send().await.map_err(|e| { if e.is_timeout() || e.is_connect() { TargetError::NotConnected } else { @@ -329,34 +346,39 @@ where } async fn is_active(&self) -> Result { - let socket_addr = lookup_host(&self.addr) - .await - .map_err(|e| TargetError::Network(format!("Failed to resolve host: {e}")))? - .next() - .ok_or_else(|| TargetError::Network("No address found".to_string()))?; - debug!("is_active socket addr: {},target id:{}", socket_addr, self.id.id); - match tokio::time::timeout(Duration::from_secs(5), tokio::net::TcpStream::connect(socket_addr)).await { - Ok(Ok(_)) => { - debug!("Connection to {} is active", self.addr); - Ok(true) - } - Ok(Err(e)) => { - debug!("Connection to {} failed: {}", self.addr, e); - if e.kind() == std::io::ErrorKind::ConnectionRefused { - Err(TargetError::NotConnected) + match tokio::time::timeout(Duration::from_secs(5), self.http_client.head(self.args.endpoint.as_str()).send()).await { + Ok(Ok(resp)) => { + let status = resp.status(); + if status.is_server_error() { + debug!("Webhook {} server error: {}", self.id, status); + Ok(false) } else { - Err(TargetError::Network(format!("Connection failed: {e}"))) + debug!("Webhook {} is reachable (status: {})", self.id, status); + Ok(true) } } - Err(_) => Err(TargetError::Timeout("Connection timed out".to_string())), + Ok(Err(e)) => { + debug!("Webhook {} request failed: {}", self.id, e); + if e.is_timeout() || e.is_connect() { + Err(TargetError::NotConnected) + } else { + Err(TargetError::Network(format!("Webhook health check failed: {e}"))) + } + } + Err(_) => Err(TargetError::Timeout("Webhook health check timed out".to_string())), } } async fn save(&self, event: Arc>) -> Result<(), TargetError> { + let queued = self.build_queued_payload(&event)?; + if let Some(store) = &self.store { - // Call the store method directly, no longer need to acquire the lock store - .put(event) + .put_raw( + &queued + .encode() + .map_err(|e| TargetError::Storage(format!("Failed to encode queued payload: {e}")))?, + ) .map_err(|e| TargetError::Storage(format!("Failed to save event to store: {e}")))?; debug!("Event saved to store for target: {}", self.id); Ok(()) @@ -368,12 +390,12 @@ where return Err(TargetError::NotConnected); } } - self.send(&event).await + self.send_body(queued.body, &queued.meta).await } } - async fn send_from_store(&self, key: Key) -> Result<(), TargetError> { - debug!("Sending event from store for target: {}", self.id); + async fn send_raw_from_store(&self, key: Key, body: Vec, meta: QueuedPayloadMeta) -> Result<(), TargetError> { + debug!("Sending queued payload from store for target: {}, key: {}", self.id, key); match self.init().await { Ok(_) => { debug!("Event sent to store for target: {}", self.name()); @@ -384,37 +406,13 @@ where } } - let store = self - .store - .as_ref() - .ok_or_else(|| TargetError::Configuration("No store configured".to_string()))?; - - // Get events directly from the store, no longer need to acquire locks - let event = match store.get(&key) { - Ok(event) => event, - Err(StoreError::NotFound) => return Ok(()), - Err(e) => { - return Err(TargetError::Storage(format!("Failed to get event from store: {e}"))); - } - }; - - if let Err(e) = self.send(&event).await { + if let Err(e) = self.send_body(body, &meta).await { if let TargetError::NotConnected = e { return Err(TargetError::NotConnected); } return Err(e); } - // Use the immutable reference of the store to delete the event content corresponding to the key - debug!("Deleting event from store for target: {}, key:{}, start", self.id, key.to_string()); - match store.del(&key) { - Ok(_) => debug!("Event deleted from store for target: {}, key:{}, end", self.id, key.to_string()), - Err(e) => { - error!("Failed to delete event from store: {}", e); - return Err(TargetError::Storage(format!("Failed to delete event from store: {e}"))); - } - } - debug!("Event sent from store and deleted for target: {}", self.id); Ok(()) } @@ -426,7 +424,7 @@ where Ok(()) } - fn store(&self) -> Option<&(dyn Store, Error = StoreError, Key = Key> + Send + Sync)> { + fn store(&self) -> Option<&(dyn Store + Send + Sync)> { // Returns the reference to the internal store self.store.as_deref() } @@ -436,14 +434,11 @@ where } async fn init(&self) -> Result<(), TargetError> { - // If the target is disabled, return to success directly if !self.is_enabled() { debug!("Webhook target {} is disabled, skipping initialization", self.id); return Ok(()); } - - // Use existing initialization logic - WebhookTarget::init(self).await + self.init_inner().await } fn is_enabled(&self) -> bool { diff --git a/rustfs/src/admin/handlers/audit.rs b/rustfs/src/admin/handlers/audit.rs new file mode 100644 index 000000000..c8b91fed8 --- /dev/null +++ b/rustfs/src/admin/handlers/audit.rs @@ -0,0 +1,818 @@ +// 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 crate::admin::router::{AdminOperation, Operation, S3Router}; +use crate::auth::{check_key_valid, get_session_token}; +use crate::server::ADMIN_PREFIX; +use futures::stream::{FuturesUnordered, StreamExt}; +use hashbrown::HashSet as HbHashSet; +use http::{HeaderMap, StatusCode}; +use hyper::Method; +use matchit::Params; +use rustfs_audit::{audit_system, start_audit_system as start_global_audit_system, system::AuditSystemState}; +use rustfs_config::audit::{AUDIT_MQTT_KEYS, AUDIT_MQTT_SUB_SYS, AUDIT_ROUTE_PREFIX, AUDIT_WEBHOOK_KEYS, AUDIT_WEBHOOK_SUB_SYS}; +use rustfs_config::{DEFAULT_DELIMITER, ENABLE_KEY, ENV_PREFIX, EnableState, MAX_ADMIN_REQUEST_BODY_SIZE}; +use rustfs_ecstore::config::Config; +use rustfs_targets::check_mqtt_broker_available; +use s3s::{Body, S3Request, S3Response, S3Result, header::CONTENT_TYPE, s3_error}; +use serde::{Deserialize, Serialize}; +use std::collections::{HashMap, HashSet}; +use std::future::Future; +use std::io::{Error, ErrorKind}; +use std::path::Path; +use std::sync::Arc; +use tokio::sync::Semaphore; +use tokio::time::{Duration, sleep, timeout}; +use tracing::{Span, warn}; +use url::Url; + +pub fn register_audit_target_route(r: &mut S3Router) -> std::io::Result<()> { + r.insert( + Method::GET, + format!("{}{}", ADMIN_PREFIX, "/v3/audit/target/list").as_str(), + AdminOperation(&ListAuditTargets {}), + )?; + + r.insert( + Method::PUT, + format!("{}{}", ADMIN_PREFIX, "/v3/audit/target/{target_type}/{target_name}").as_str(), + AdminOperation(&AuditTargetConfig {}), + )?; + + r.insert( + Method::DELETE, + format!("{}{}", ADMIN_PREFIX, "/v3/audit/target/{target_type}/{target_name}/reset").as_str(), + AdminOperation(&RemoveAuditTarget {}), + )?; + + Ok(()) +} + +#[derive(Debug, Deserialize)] +pub struct KeyValue { + pub key: String, + pub value: String, +} + +#[derive(Debug, Deserialize)] +pub struct AuditTargetBody { + pub key_values: Vec, +} + +#[derive(Serialize, Debug)] +struct AuditEndpoint { + account_id: String, + service: String, + status: String, + source: AuditEndpointSource, +} + +#[derive(Serialize, Debug)] +struct AuditEndpointsResponse { + audit_endpoints: Vec, +} + +type EndpointKey = (String, String); + +#[derive(Clone, Copy, Debug, PartialEq, Eq, Serialize)] +#[serde(rename_all = "lowercase")] +enum AuditEndpointSource { + Config, + Env, + Mixed, + Runtime, +} + +fn normalized_endpoint_key(account_id: &str, service: &str) -> EndpointKey { + (account_id.to_lowercase(), service.to_string()) +} + +async fn check_permissions(req: &S3Request) -> S3Result<()> { + let Some(input_cred) = &req.credentials else { + return Err(s3_error!(InvalidRequest, "credentials not found")); + }; + check_key_valid(get_session_token(&req.uri, &req.headers).unwrap_or_default(), &input_cred.access_key).await?; + Ok(()) +} + +fn build_response(status: StatusCode, body: Body, request_id: Option<&http::HeaderValue>) -> S3Response<(StatusCode, Body)> { + let mut header = HeaderMap::new(); + header.insert(CONTENT_TYPE, "application/json".parse().unwrap()); + if let Some(v) = request_id { + header.insert("x-request-id", v.clone()); + } + S3Response::with_headers((status, body), header) +} + +async fn retry_with_backoff(mut operation: F, max_attempts: usize, base_delay: Duration) -> Result +where + F: FnMut() -> Fut, + Fut: Future>, +{ + let mut attempts = 0; + let mut delay = base_delay; + let mut last_err = None; + + while attempts < max_attempts { + match operation().await { + Ok(result) => return Ok(result), + Err(e) => { + last_err = Some(e); + attempts += 1; + if attempts < max_attempts { + sleep(delay).await; + delay = delay.saturating_mul(2); + } + } + } + } + Err(last_err.unwrap_or_else(|| Error::other("retry_with_backoff: unknown error"))) +} + +async fn validate_queue_dir(queue_dir: &str) -> S3Result<()> { + if !queue_dir.is_empty() { + if !Path::new(queue_dir).is_absolute() { + return Err(s3_error!(InvalidArgument, "queue_dir must be absolute path")); + } + retry_with_backoff( + || async { tokio::fs::metadata(queue_dir).await.map(|_| ()) }, + 3, + Duration::from_millis(100), + ) + .await + .map_err(|e| match e.kind() { + ErrorKind::NotFound => s3_error!(InvalidArgument, "queue_dir does not exist"), + ErrorKind::PermissionDenied => s3_error!(InvalidArgument, "queue_dir exists but permission denied"), + _ => s3_error!(InvalidArgument, "failed to access queue_dir: {}", e), + })?; + } + Ok(()) +} + +fn config_enable_is_on(value: &str) -> bool { + matches!(value.trim().to_ascii_lowercase().as_str(), "on" | "true" | "yes" | "1") +} + +fn has_any_audit_targets(config: &Config) -> bool { + for subsystem in [AUDIT_WEBHOOK_SUB_SYS, AUDIT_MQTT_SUB_SYS] { + let Some(targets) = config.0.get(subsystem) else { + continue; + }; + if targets.keys().any(|key| key != DEFAULT_DELIMITER) { + return true; + } + } + false +} + +fn collect_configured_audit_endpoint_keys(config: &Config) -> Vec { + let mut endpoints = Vec::new(); + for (subsystem, service) in [(AUDIT_WEBHOOK_SUB_SYS, "webhook"), (AUDIT_MQTT_SUB_SYS, "mqtt")] { + let Some(targets) = config.0.get(subsystem) else { + continue; + }; + + for (target_name, kvs) in targets { + if target_name == DEFAULT_DELIMITER { + continue; + } + let enabled = kvs.lookup(ENABLE_KEY).as_deref().map(config_enable_is_on).unwrap_or(false); + if enabled { + endpoints.push((target_name.clone(), service.to_string())); + } + } + } + endpoints +} + +fn collect_config_entry_keys(config: &Config) -> HbHashSet { + let mut endpoints = HbHashSet::new(); + for (subsystem, service) in [(AUDIT_WEBHOOK_SUB_SYS, "webhook"), (AUDIT_MQTT_SUB_SYS, "mqtt")] { + let Some(targets) = config.0.get(subsystem) else { + continue; + }; + + for target_name in targets.keys() { + if target_name == DEFAULT_DELIMITER { + continue; + } + endpoints.insert(normalized_endpoint_key(target_name, service)); + } + } + endpoints +} + +fn collect_env_endpoint_keys() -> HbHashSet { + let mut endpoints = HbHashSet::new(); + + for (service, valid_keys) in [("webhook", AUDIT_WEBHOOK_KEYS), ("mqtt", AUDIT_MQTT_KEYS)] { + let env_prefix = format!("{ENV_PREFIX}{AUDIT_ROUTE_PREFIX}{service}{DEFAULT_DELIMITER}").to_uppercase(); + + for (key, _value) in std::env::vars() { + let Some(rest) = key.strip_prefix(&env_prefix) else { + continue; + }; + + let mut parts = rest.rsplitn(2, DEFAULT_DELIMITER); + let instance_id_part = parts.next().unwrap_or(DEFAULT_DELIMITER); + let field_name_part = parts.next(); + + let (field_name, instance_id) = match field_name_part { + Some(field) => (field.to_lowercase(), instance_id_part.to_lowercase()), + None => (instance_id_part.to_lowercase(), DEFAULT_DELIMITER.to_string()), + }; + + if instance_id == DEFAULT_DELIMITER || instance_id.is_empty() { + continue; + } + + if valid_keys.contains(&field_name.as_str()) { + endpoints.insert(normalized_endpoint_key(&instance_id, service)); + } + } + } + + endpoints +} + +fn classify_audit_endpoint_source( + config_targets: &HbHashSet, + env_targets: &HbHashSet, + key: &EndpointKey, +) -> AuditEndpointSource { + match (config_targets.contains(key), env_targets.contains(key)) { + (true, true) => AuditEndpointSource::Mixed, + (true, false) => AuditEndpointSource::Config, + (false, true) => AuditEndpointSource::Env, + (false, false) => AuditEndpointSource::Runtime, + } +} + +fn audit_endpoint_source(config: &Config, target_type: &str, target_name: &str) -> AuditEndpointSource { + let config_targets = collect_config_entry_keys(config); + let env_targets = collect_env_endpoint_keys(); + let service = match target_type { + AUDIT_WEBHOOK_SUB_SYS => "webhook", + AUDIT_MQTT_SUB_SYS => "mqtt", + _ => "", + }; + + let key = normalized_endpoint_key(target_name, service); + classify_audit_endpoint_source(&config_targets, &env_targets, &key) +} + +fn audit_target_mutation_block_reason(config: &Config, target_type: &str, target_name: &str) -> Option { + match audit_endpoint_source(config, target_type, target_name) { + AuditEndpointSource::Env => Some(format!( + "audit target '{}' is managed by environment variables and cannot be modified from the console", + target_name + )), + AuditEndpointSource::Mixed => Some(format!( + "audit target '{}' is configured by both persisted config and environment variables; remove the environment variables first", + target_name + )), + AuditEndpointSource::Config | AuditEndpointSource::Runtime => None, + } +} + +fn merge_audit_endpoints(config: &Config, runtime_statuses: HashMap) -> Vec { + let mut audit_endpoints = Vec::new(); + let mut seen = HashSet::new(); + let configured_keys = collect_configured_audit_endpoint_keys(config); + let config_targets = collect_config_entry_keys(config); + let env_targets = collect_env_endpoint_keys(); + let mut normalized_runtime_statuses: HashMap = HashMap::new(); + for ((account_id, service), status) in runtime_statuses { + let normalized = normalized_endpoint_key(&account_id, &service); + normalized_runtime_statuses + .entry(normalized) + .or_insert((account_id, service, status)); + } + + for key in configured_keys { + let normalized = normalized_endpoint_key(&key.0, &key.1); + if !seen.insert(normalized.clone()) { + continue; + } + let status = normalized_runtime_statuses + .remove(&normalized) + .map(|(_, _, status)| status) + .unwrap_or_else(|| "offline".to_string()); + let source = classify_audit_endpoint_source(&config_targets, &env_targets, &normalized); + audit_endpoints.push(AuditEndpoint { + account_id: key.0, + service: key.1, + status, + source, + }); + } + + for (normalized, (account_id, service, status)) in normalized_runtime_statuses { + if seen.insert(normalized.clone()) { + audit_endpoints.push(AuditEndpoint { + account_id, + service, + status, + source: classify_audit_endpoint_source(&config_targets, &env_targets, &normalized), + }); + } + } + + for key in &env_targets { + if !seen.insert(key.clone()) { + continue; + } + + audit_endpoints.push(AuditEndpoint { + account_id: key.0.clone(), + service: key.1.clone(), + status: "offline".to_string(), + source: classify_audit_endpoint_source(&config_targets, &env_targets, key), + }); + } + + audit_endpoints.sort_by(|a, b| a.service.cmp(&b.service).then_with(|| a.account_id.cmp(&b.account_id))); + audit_endpoints +} + +fn collect_validated_key_values( + key_values: &[KeyValue], + allowed_keys: &HashSet<&str>, + target_type: &str, +) -> S3Result> { + let mut kv_map = HashMap::new(); + let mut seen = HashSet::new(); + + for kv in key_values { + if !allowed_keys.contains(kv.key.as_str()) { + return Err(s3_error!( + InvalidArgument, + "key '{}' not allowed for audit target type '{}'", + kv.key, + target_type + )); + } + + if !seen.insert(kv.key.as_str()) { + return Err(s3_error!(InvalidArgument, "duplicate key '{}' in request body", kv.key)); + } + + kv_map.insert(kv.key.clone(), kv.value.clone()); + } + + Ok(kv_map) +} + +fn extract_target_params<'a>(params: &'a Params<'_, '_>) -> S3Result<(&'a str, &'a str)> { + let target_type = params + .get("target_type") + .ok_or_else(|| s3_error!(InvalidArgument, "missing required parameter: 'target_type'"))?; + if target_type != AUDIT_WEBHOOK_SUB_SYS && target_type != AUDIT_MQTT_SUB_SYS { + return Err(s3_error!(InvalidArgument, "unsupported audit target type: '{}'", target_type)); + } + let target_name = params + .get("target_name") + .ok_or_else(|| s3_error!(InvalidArgument, "missing required parameter: 'target_name'"))?; + Ok((target_type, target_name)) +} + +async fn load_server_config_from_store() -> S3Result { + let Some(store) = rustfs_ecstore::global::new_object_layer_fn() else { + return Ok(Config::new()); + }; + + rustfs_ecstore::config::com::read_config_without_migrate(store) + .await + .map_err(|e| s3_error!(InternalError, "failed to read server config: {}", e)) +} + +async fn apply_audit_runtime_config(config: Config) -> S3Result<()> { + let has_targets = has_any_audit_targets(&config); + + if let Some(system) = audit_system() { + match system.get_state().await { + AuditSystemState::Running | AuditSystemState::Paused | AuditSystemState::Starting => { + if has_targets { + system + .reload_config(config) + .await + .map_err(|e| s3_error!(InternalError, "failed to reload audit config: {}", e))?; + } else { + system + .close() + .await + .map_err(|e| s3_error!(InternalError, "failed to stop audit system: {}", e))?; + } + } + AuditSystemState::Stopped | AuditSystemState::Stopping => { + if has_targets { + system + .start(config) + .await + .map_err(|e| s3_error!(InternalError, "failed to start audit system: {}", e))?; + } + } + } + } else if has_targets { + start_global_audit_system(config) + .await + .map_err(|e| s3_error!(InternalError, "failed to start audit system: {}", e))?; + } + + Ok(()) +} + +async fn update_audit_config_and_reload(mut modifier: F) -> S3Result<()> +where + F: FnMut(&mut Config) -> bool, +{ + let Some(store) = rustfs_ecstore::global::new_object_layer_fn() else { + return Err(s3_error!(InternalError, "server storage not initialized")); + }; + + let mut config = rustfs_ecstore::config::com::read_config_without_migrate(store.clone()) + .await + .map_err(|e| s3_error!(InternalError, "failed to read server config: {}", e))?; + + if !modifier(&mut config) { + return Ok(()); + } + + rustfs_ecstore::config::com::save_server_config(store, &config) + .await + .map_err(|e| s3_error!(InternalError, "failed to save audit config: {}", e))?; + + apply_audit_runtime_config(config).await +} + +pub struct AuditTargetConfig {} + +#[async_trait::async_trait] +impl Operation for AuditTargetConfig { + async fn call(&self, req: S3Request, params: Params<'_, '_>) -> S3Result> { + let span = Span::current(); + let _enter = span.enter(); + let (target_type, target_name) = extract_target_params(¶ms)?; + + check_permissions(&req).await?; + let config_snapshot = load_server_config_from_store().await?; + if let Some(reason) = audit_target_mutation_block_reason(&config_snapshot, target_type, target_name) { + return Err(s3_error!(InvalidRequest, "{reason}")); + } + + let mut input = req.input; + let body_bytes = input.store_all_limited(MAX_ADMIN_REQUEST_BODY_SIZE).await.map_err(|e| { + warn!("failed to read request body: {:?}", e); + s3_error!(InvalidRequest, "failed to read request body") + })?; + + let audit_body: AuditTargetBody = serde_json::from_slice(&body_bytes) + .map_err(|e| s3_error!(InvalidArgument, "invalid json body for audit target config: {}", e))?; + + let allowed_keys: HashSet<&str> = match target_type { + AUDIT_WEBHOOK_SUB_SYS => AUDIT_WEBHOOK_KEYS.iter().cloned().collect(), + AUDIT_MQTT_SUB_SYS => AUDIT_MQTT_KEYS.iter().cloned().collect(), + _ => unreachable!(), + }; + + let kv_map = collect_validated_key_values(&audit_body.key_values, &allowed_keys, target_type)?; + + if target_type == AUDIT_WEBHOOK_SUB_SYS { + let endpoint = kv_map + .get("endpoint") + .map(String::as_str) + .ok_or_else(|| s3_error!(InvalidArgument, "endpoint is required"))?; + let parsed_endpoint = Url::parse(endpoint).map_err(|e| s3_error!(InvalidArgument, "invalid endpoint url: {}", e))?; + match parsed_endpoint.scheme() { + "http" | "https" => {} + other => { + return Err(s3_error!( + InvalidArgument, + "unsupported endpoint scheme: {} (only http and https are allowed)", + other + )); + } + } + if let Some(queue_dir) = kv_map.get("queue_dir") { + validate_queue_dir(queue_dir.as_str()).await?; + } + if kv_map.contains_key("client_cert") != kv_map.contains_key("client_key") { + return Err(s3_error!(InvalidArgument, "client_cert and client_key must be specified as a pair")); + } + } else if target_type == AUDIT_MQTT_SUB_SYS { + let endpoint = kv_map + .get(rustfs_config::MQTT_BROKER) + .map(String::as_str) + .ok_or_else(|| s3_error!(InvalidArgument, "broker endpoint is required"))?; + let topic = kv_map + .get(rustfs_config::MQTT_TOPIC) + .map(String::as_str) + .ok_or_else(|| s3_error!(InvalidArgument, "topic is required"))?; + let username = kv_map.get(rustfs_config::MQTT_USERNAME).map(String::as_str); + let password = kv_map.get(rustfs_config::MQTT_PASSWORD).map(String::as_str); + check_mqtt_broker_available(endpoint, topic, username, password) + .await + .map_err(|e| s3_error!(InvalidArgument, "MQTT Broker unavailable: {}", e))?; + + if let Some(queue_dir) = kv_map.get("queue_dir") { + validate_queue_dir(queue_dir.as_str()).await?; + if let Some(qos) = kv_map.get("qos") { + match qos.parse::() { + Ok(1) | Ok(2) => {} + Ok(0) => return Err(s3_error!(InvalidArgument, "qos should be 1 or 2 if queue_dir is set")), + _ => return Err(s3_error!(InvalidArgument, "qos must be an integer 0, 1, or 2")), + } + } + } + } + + let mut kvs = rustfs_ecstore::config::KVS::new(); + for (key, value) in kv_map { + kvs.insert(key, value); + } + kvs.insert(ENABLE_KEY.to_string(), EnableState::On.to_string()); + + update_audit_config_and_reload(|config| { + config + .0 + .entry(target_type.to_lowercase()) + .or_default() + .insert(target_name.to_lowercase(), kvs.clone()); + true + }) + .await?; + + Ok(build_response(StatusCode::OK, Body::empty(), req.headers.get("x-request-id"))) + } +} + +pub struct ListAuditTargets {} + +#[async_trait::async_trait] +impl Operation for ListAuditTargets { + async fn call(&self, req: S3Request, _params: Params<'_, '_>) -> S3Result> { + let span = Span::current(); + let _enter = span.enter(); + check_permissions(&req).await?; + + let mut runtime_statuses = HashMap::new(); + if let Some(system) = audit_system() { + let targets = system.get_target_values().await; + let semaphore = Arc::new(Semaphore::new(10)); + let mut futures = FuturesUnordered::new(); + + for target in targets { + let sem = Arc::clone(&semaphore); + futures.push(async move { + let _permit = sem.acquire().await; + let status = match timeout(Duration::from_secs(3), target.is_active()).await { + Ok(Ok(true)) => "online", + _ => "offline", + }; + ((target.id().id.clone(), target.id().name.to_string()), status.to_string()) + }); + } + + while let Some((key, status)) = futures.next().await { + runtime_statuses.insert(key, status); + } + } + + let config = load_server_config_from_store().await?; + let audit_endpoints = merge_audit_endpoints(&config, runtime_statuses); + let data = serde_json::to_vec(&AuditEndpointsResponse { audit_endpoints }) + .map_err(|e| s3_error!(InternalError, "failed to serialize audit targets: {}", e))?; + + Ok(build_response(StatusCode::OK, Body::from(data), req.headers.get("x-request-id"))) + } +} + +pub struct RemoveAuditTarget {} + +#[async_trait::async_trait] +impl Operation for RemoveAuditTarget { + async fn call(&self, req: S3Request, params: Params<'_, '_>) -> S3Result> { + let span = Span::current(); + let _enter = span.enter(); + let (target_type, target_name) = extract_target_params(¶ms)?; + + check_permissions(&req).await?; + let config_snapshot = load_server_config_from_store().await?; + if let Some(reason) = audit_target_mutation_block_reason(&config_snapshot, target_type, target_name) { + return Err(s3_error!(InvalidRequest, "{reason}")); + } + + update_audit_config_and_reload(|config| { + let mut changed = false; + if let Some(targets) = config.0.get_mut(&target_type.to_lowercase()) { + if targets.remove(&target_name.to_lowercase()).is_some() { + changed = true; + } + if targets.is_empty() { + config.0.remove(&target_type.to_lowercase()); + } + } + changed + }) + .await?; + + Ok(build_response(StatusCode::OK, Body::empty(), req.headers.get("x-request-id"))) + } +} + +#[cfg(test)] +mod tests { + use super::*; + use rustfs_ecstore::config::{KV, KVS}; + use std::collections::{HashMap, HashSet}; + use temp_env::{with_var, with_vars}; + + fn enabled_kvs(value: &str) -> KVS { + KVS(vec![KV { + key: ENABLE_KEY.to_string(), + value: value.to_string(), + hidden_if_empty: false, + }]) + } + + #[test] + fn merge_audit_endpoints_marks_config_env_and_mixed_sources() { + let config = Config(HashMap::from([( + AUDIT_WEBHOOK_SUB_SYS.to_string(), + HashMap::from([ + ("mixed-target".to_string(), enabled_kvs("on")), + ("config-target".to_string(), enabled_kvs("on")), + ]), + )])); + + with_vars( + [ + ("RUSTFS_AUDIT_WEBHOOK_ENDPOINT_MIXED-TARGET", Some("https://example.com/hook")), + ("RUSTFS_AUDIT_WEBHOOK_ENABLE_ENV-ONLY", Some("on")), + ("RUSTFS_AUDIT_WEBHOOK_ENDPOINT_ENV-ONLY", Some("https://example.com/env")), + ], + || { + let runtime = HashMap::from([ + (("mixed-target".to_string(), "webhook".to_string()), "online".to_string()), + (("env-only".to_string(), "webhook".to_string()), "online".to_string()), + ]); + let merged = merge_audit_endpoints(&config, runtime); + + let mixed = merged + .iter() + .find(|entry| entry.account_id == "mixed-target") + .expect("mixed target should be present"); + assert_eq!(mixed.source, AuditEndpointSource::Mixed); + + let env_only = merged + .iter() + .find(|entry| entry.account_id == "env-only") + .expect("env-only target should be present"); + assert_eq!(env_only.source, AuditEndpointSource::Env); + + let config_only = merged + .iter() + .find(|entry| entry.account_id == "config-target") + .expect("config target should be present"); + assert_eq!(config_only.source, AuditEndpointSource::Config); + }, + ); + } + + #[test] + fn audit_target_mutation_block_reason_rejects_env_managed_target() { + with_vars( + [ + ("RUSTFS_AUDIT_WEBHOOK_ENABLE_PRIMARY", Some("on")), + ("RUSTFS_AUDIT_WEBHOOK_ENDPOINT_PRIMARY", Some("https://example.com/hook")), + ], + || { + let config = Config(HashMap::new()); + let reason = audit_target_mutation_block_reason(&config, AUDIT_WEBHOOK_SUB_SYS, "primary"); + assert!(reason.is_some()); + assert!(reason.unwrap().contains("managed by environment variables")); + }, + ); + } + + #[test] + fn audit_target_mutation_block_reason_rejects_mixed_target() { + with_var("RUSTFS_AUDIT_WEBHOOK_ENDPOINT_PRIMARY", Some("https://example.com/hook"), || { + let config = Config(HashMap::from([( + AUDIT_WEBHOOK_SUB_SYS.to_string(), + HashMap::from([("primary".to_string(), enabled_kvs("on"))]), + )])); + let reason = audit_target_mutation_block_reason(&config, AUDIT_WEBHOOK_SUB_SYS, "primary"); + assert!(reason.is_some()); + assert!(reason.unwrap().contains("both persisted config and environment variables")); + }); + } + + #[test] + fn merge_audit_endpoints_marks_disabled_config_with_env_override_as_mixed() { + let config = Config(HashMap::from([( + AUDIT_WEBHOOK_SUB_SYS.to_string(), + HashMap::from([("mixed-disabled".to_string(), enabled_kvs("off"))]), + )])); + + with_vars( + [ + ("RUSTFS_AUDIT_WEBHOOK_ENABLE_MIXED-DISABLED", Some("on")), + ("RUSTFS_AUDIT_WEBHOOK_ENDPOINT_MIXED-DISABLED", Some("https://example.com/hook")), + ], + || { + let merged = merge_audit_endpoints(&config, HashMap::new()); + let mixed = merged + .iter() + .find(|entry| entry.account_id == "mixed-disabled") + .expect("mixed target should be present"); + assert_eq!(mixed.source, AuditEndpointSource::Mixed); + assert_eq!(mixed.status, "offline"); + }, + ); + } + + #[test] + fn merge_audit_endpoints_includes_env_only_target_without_runtime_status() { + let config = Config(HashMap::new()); + + with_vars( + [ + ("RUSTFS_AUDIT_WEBHOOK_ENABLE_ENV-ONLY", Some("on")), + ("RUSTFS_AUDIT_WEBHOOK_ENDPOINT_ENV-ONLY", Some("https://example.com/env")), + ], + || { + let merged = merge_audit_endpoints(&config, HashMap::new()); + let env_only = merged + .iter() + .find(|entry| entry.account_id == "env-only") + .expect("env-only target should be present"); + assert_eq!(env_only.source, AuditEndpointSource::Env); + assert_eq!(env_only.status, "offline"); + }, + ); + } + + #[test] + fn collect_validated_key_values_rejects_duplicate_keys() { + let allowed_keys: HashSet<&str> = ["endpoint", "auth_token"].into_iter().collect(); + let key_values = vec![ + KeyValue { + key: "endpoint".to_string(), + value: "https://example.com/one".to_string(), + }, + KeyValue { + key: "endpoint".to_string(), + value: "https://example.com/two".to_string(), + }, + ]; + + let err = collect_validated_key_values(&key_values, &allowed_keys, AUDIT_WEBHOOK_SUB_SYS).unwrap_err(); + assert!(err.to_string().contains("duplicate key")); + } + + #[test] + fn merge_audit_endpoints_marks_mixed_with_case_insensitive_instance_id() { + let config = Config(HashMap::from([( + AUDIT_WEBHOOK_SUB_SYS.to_string(), + HashMap::from([("PrimaryCase".to_string(), enabled_kvs("on"))]), + )])); + + with_vars( + [ + ("RUSTFS_AUDIT_WEBHOOK_ENABLE_PRIMARYCASE", Some("on")), + ("RUSTFS_AUDIT_WEBHOOK_ENDPOINT_PRIMARYCASE", Some("https://example.com/hook")), + ], + || { + let runtime = HashMap::from([(("PrimaryCase".to_string(), "webhook".to_string()), "online".to_string())]); + let merged = merge_audit_endpoints(&config, runtime); + let mixed = merged + .iter() + .find(|entry| entry.account_id == "PrimaryCase" && entry.service == "webhook") + .expect("mixed target should be present"); + assert_eq!(mixed.source, AuditEndpointSource::Mixed); + }, + ); + } + + #[test] + fn audit_target_mutation_block_reason_allows_case_insensitive_config_target_lookup() { + let config = Config(HashMap::from([( + AUDIT_WEBHOOK_SUB_SYS.to_string(), + HashMap::from([("PrimaryCase".to_string(), enabled_kvs("on"))]), + )])); + + assert!(audit_target_mutation_block_reason(&config, AUDIT_WEBHOOK_SUB_SYS, "primarycase").is_none()); + } +} diff --git a/rustfs/src/admin/handlers/event.rs b/rustfs/src/admin/handlers/event.rs index 4d0c879ee..07b6f73b0 100644 --- a/rustfs/src/admin/handlers/event.rs +++ b/rustfs/src/admin/handlers/event.rs @@ -16,21 +16,23 @@ use crate::admin::router::{AdminOperation, Operation, S3Router}; use crate::auth::{check_key_valid, get_session_token}; use crate::server::ADMIN_PREFIX; use futures::stream::{FuturesUnordered, StreamExt}; +use hashbrown::HashSet as HbHashSet; use http::{HeaderMap, StatusCode}; use hyper::Method; use matchit::Params; -use rustfs_config::notify::{NOTIFY_MQTT_SUB_SYS, NOTIFY_WEBHOOK_SUB_SYS}; -use rustfs_config::{ENABLE_KEY, EnableState, MAX_ADMIN_REQUEST_BODY_SIZE}; +use rustfs_config::notify::{ + NOTIFY_MQTT_KEYS, NOTIFY_MQTT_SUB_SYS, NOTIFY_ROUTE_PREFIX, NOTIFY_WEBHOOK_KEYS, NOTIFY_WEBHOOK_SUB_SYS, +}; +use rustfs_config::{DEFAULT_DELIMITER, ENABLE_KEY, ENV_PREFIX, EnableState, MAX_ADMIN_REQUEST_BODY_SIZE}; +use rustfs_ecstore::config::Config; use rustfs_targets::check_mqtt_broker_available; use s3s::{Body, S3Request, S3Response, S3Result, header::CONTENT_TYPE, s3_error}; use serde::{Deserialize, Serialize}; use std::collections::{HashMap, HashSet}; use std::future::Future; use std::io::{Error, ErrorKind}; -use std::net::SocketAddr; use std::path::Path; use std::sync::Arc; -use tokio::net::lookup_host; use tokio::sync::Semaphore; use tokio::time::{Duration, sleep, timeout}; use tracing::{Span, info, warn}; @@ -80,6 +82,7 @@ struct NotificationEndpoint { account_id: String, service: String, status: String, + source: NotificationEndpointSource, } #[derive(Serialize, Debug)] @@ -87,6 +90,21 @@ struct NotificationEndpointsResponse { notification_endpoints: Vec, } +type EndpointKey = (String, String); + +#[derive(Clone, Copy, Debug, PartialEq, Eq, Serialize)] +#[serde(rename_all = "lowercase")] +enum NotificationEndpointSource { + Config, + Env, + Mixed, + Runtime, +} + +fn normalized_endpoint_key(account_id: &str, service: &str) -> EndpointKey { + (account_id.to_lowercase(), service.to_string()) +} + // --- Helper Functions --- async fn check_permissions(req: &S3Request) -> S3Result<()> { @@ -155,6 +173,215 @@ async fn validate_queue_dir(queue_dir: &str) -> S3Result<()> { Ok(()) } +fn config_enable_is_on(value: &str) -> bool { + matches!(value.trim().to_ascii_lowercase().as_str(), "on" | "true" | "yes" | "1") +} + +fn collect_configured_endpoint_keys(config: &Config) -> Vec { + let mut endpoints = Vec::new(); + for (subsystem, service) in [(NOTIFY_WEBHOOK_SUB_SYS, "webhook"), (NOTIFY_MQTT_SUB_SYS, "mqtt")] { + let Some(targets) = config.0.get(subsystem) else { + continue; + }; + + for (target_name, kvs) in targets { + if target_name == DEFAULT_DELIMITER { + continue; + } + let enabled = kvs.lookup(ENABLE_KEY).as_deref().map(config_enable_is_on).unwrap_or(false); + if enabled { + endpoints.push((target_name.clone(), service.to_string())); + } + } + } + endpoints +} + +fn collect_config_entry_keys(config: &Config) -> HbHashSet { + let mut endpoints = HbHashSet::new(); + for (subsystem, service) in [(NOTIFY_WEBHOOK_SUB_SYS, "webhook"), (NOTIFY_MQTT_SUB_SYS, "mqtt")] { + let Some(targets) = config.0.get(subsystem) else { + continue; + }; + + for target_name in targets.keys() { + if target_name == DEFAULT_DELIMITER { + continue; + } + endpoints.insert(normalized_endpoint_key(target_name, service)); + } + } + endpoints +} + +fn collect_env_endpoint_keys() -> HbHashSet { + let mut endpoints = HbHashSet::new(); + + for (service, valid_keys) in [("webhook", NOTIFY_WEBHOOK_KEYS), ("mqtt", NOTIFY_MQTT_KEYS)] { + let env_prefix = format!("{ENV_PREFIX}{NOTIFY_ROUTE_PREFIX}{service}{DEFAULT_DELIMITER}").to_uppercase(); + + for (key, _value) in std::env::vars() { + let Some(rest) = key.strip_prefix(&env_prefix) else { + continue; + }; + + let mut parts = rest.rsplitn(2, DEFAULT_DELIMITER); + let instance_id_part = parts.next().unwrap_or(DEFAULT_DELIMITER); + let field_name_part = parts.next(); + + let (field_name, instance_id) = match field_name_part { + Some(field) => (field.to_lowercase(), instance_id_part.to_lowercase()), + None => (instance_id_part.to_lowercase(), DEFAULT_DELIMITER.to_string()), + }; + + if instance_id == DEFAULT_DELIMITER || instance_id.is_empty() { + continue; + } + + if valid_keys.contains(&field_name.as_str()) { + endpoints.insert(normalized_endpoint_key(&instance_id, service)); + } + } + } + + endpoints +} + +fn classify_notification_endpoint_source( + config_targets: &HbHashSet, + env_targets: &HbHashSet, + key: &EndpointKey, +) -> NotificationEndpointSource { + match (config_targets.contains(key), env_targets.contains(key)) { + (true, true) => NotificationEndpointSource::Mixed, + (true, false) => NotificationEndpointSource::Config, + (false, true) => NotificationEndpointSource::Env, + (false, false) => NotificationEndpointSource::Runtime, + } +} + +fn notification_endpoint_source(config: &Config, target_type: &str, target_name: &str) -> NotificationEndpointSource { + let config_targets = collect_config_entry_keys(config); + let env_targets = collect_env_endpoint_keys(); + let service = match target_type { + NOTIFY_WEBHOOK_SUB_SYS => "webhook", + NOTIFY_MQTT_SUB_SYS => "mqtt", + _ => "", + }; + + let key = normalized_endpoint_key(target_name, service); + classify_notification_endpoint_source(&config_targets, &env_targets, &key) +} + +fn target_mutation_block_reason(config: &Config, target_type: &str, target_name: &str) -> Option { + match notification_endpoint_source(config, target_type, target_name) { + NotificationEndpointSource::Env => Some(format!( + "target '{}' is managed by environment variables and cannot be modified from the console", + target_name + )), + NotificationEndpointSource::Mixed => Some(format!( + "target '{}' is configured by both persisted config and environment variables; remove the environment variables first", + target_name + )), + NotificationEndpointSource::Config | NotificationEndpointSource::Runtime => None, + } +} + +fn merge_notification_endpoints(config: &Config, runtime_statuses: HashMap) -> Vec { + let mut notification_endpoints = Vec::new(); + let mut seen = HashSet::new(); + let configured_keys = collect_configured_endpoint_keys(config); + let config_targets = collect_config_entry_keys(config); + let env_targets = collect_env_endpoint_keys(); + let mut normalized_runtime_statuses: HashMap = HashMap::new(); + for ((account_id, service), status) in runtime_statuses { + let normalized = normalized_endpoint_key(&account_id, &service); + normalized_runtime_statuses + .entry(normalized) + .or_insert((account_id, service, status)); + } + + for key in configured_keys { + let normalized = normalized_endpoint_key(&key.0, &key.1); + if !seen.insert(normalized.clone()) { + continue; + } + let status = normalized_runtime_statuses + .remove(&normalized) + .map(|(_, _, status)| status) + .unwrap_or_else(|| "offline".to_string()); + let source = classify_notification_endpoint_source(&config_targets, &env_targets, &normalized); + notification_endpoints.push(NotificationEndpoint { + account_id: key.0, + service: key.1, + status, + source, + }); + } + + for (normalized, (account_id, service, status)) in normalized_runtime_statuses { + if seen.insert(normalized.clone()) { + notification_endpoints.push(NotificationEndpoint { + account_id, + service, + status, + source: classify_notification_endpoint_source(&config_targets, &env_targets, &normalized), + }); + } + } + + for key in &env_targets { + if !seen.insert(key.clone()) { + continue; + } + + notification_endpoints.push(NotificationEndpoint { + account_id: key.0.clone(), + service: key.1.clone(), + status: "offline".to_string(), + source: classify_notification_endpoint_source(&config_targets, &env_targets, key), + }); + } + + notification_endpoints.sort_by(|a, b| a.service.cmp(&b.service).then_with(|| a.account_id.cmp(&b.account_id))); + notification_endpoints +} + +fn collect_online_target_arns(region: &str, target_statuses: Vec<(rustfs_targets::arn::TargetID, String)>) -> Vec { + target_statuses + .into_iter() + .filter_map(|(target_id, status)| (status == "online").then(|| target_id.to_arn(region).to_string())) + .collect() +} + +fn collect_validated_key_values( + key_values: &[KeyValue], + allowed_keys: &HashSet<&str>, + target_type: &str, +) -> S3Result> { + let mut kv_map = HashMap::new(); + let mut seen = HashSet::new(); + + for kv in key_values { + if !allowed_keys.contains(kv.key.as_str()) { + return Err(s3_error!( + InvalidArgument, + "key '{}' not allowed for target type '{}'", + kv.key, + target_type + )); + } + + if !seen.insert(kv.key.as_str()) { + return Err(s3_error!(InvalidArgument, "duplicate key '{}' in request body", kv.key)); + } + + kv_map.insert(kv.key.clone(), kv.value.clone()); + } + + Ok(kv_map) +} + // --- Operations --- pub struct NotificationTarget {} @@ -167,6 +394,10 @@ impl Operation for NotificationTarget { check_permissions(&req).await?; let ns = get_notification_system()?; + let config_snapshot = ns.config.read().await.clone(); + if let Some(reason) = target_mutation_block_reason(&config_snapshot, target_type, target_name) { + return Err(s3_error!(InvalidRequest, "{reason}")); + } let mut input = req.input; let body_bytes = input.store_all_limited(MAX_ADMIN_REQUEST_BODY_SIZE).await.map_err(|e| { @@ -183,37 +414,27 @@ impl Operation for NotificationTarget { _ => unreachable!(), }; - let kv_map: HashMap<&str, &str> = notification_body - .key_values - .iter() - .map(|kv| (kv.key.as_str(), kv.value.as_str())) - .collect(); - - // Validate keys - for key in kv_map.keys() { - if !allowed_keys.contains(key) { - return Err(s3_error!(InvalidArgument, "key '{}' not allowed for target type '{}'", key, target_type)); - } - } + let kv_map = collect_validated_key_values(¬ification_body.key_values, &allowed_keys, target_type)?; // Type-specific validation if target_type == NOTIFY_WEBHOOK_SUB_SYS { let endpoint = kv_map .get("endpoint") + .map(String::as_str) .ok_or_else(|| s3_error!(InvalidArgument, "endpoint is required"))?; - let url = Url::parse(endpoint).map_err(|e| s3_error!(InvalidArgument, "invalid endpoint url: {}", e))?; - let host = url - .host_str() - .ok_or_else(|| s3_error!(InvalidArgument, "endpoint missing host"))?; - let port = url - .port_or_known_default() - .ok_or_else(|| s3_error!(InvalidArgument, "endpoint missing port"))?; - let addr = format!("{host}:{port}"); - if addr.parse::().is_err() && lookup_host(&addr).await.is_err() { - return Err(s3_error!(InvalidArgument, "invalid or unresolvable endpoint address")); + let parsed_endpoint = Url::parse(endpoint).map_err(|e| s3_error!(InvalidArgument, "invalid endpoint url: {}", e))?; + match parsed_endpoint.scheme() { + "http" | "https" => {} + other => { + return Err(s3_error!( + InvalidArgument, + "unsupported endpoint scheme: {} (only http and https are allowed)", + other + )); + } } if let Some(queue_dir) = kv_map.get("queue_dir") { - validate_queue_dir(queue_dir).await?; + validate_queue_dir(queue_dir.as_str()).await?; } if kv_map.contains_key("client_cert") != kv_map.contains_key("client_key") { return Err(s3_error!(InvalidArgument, "client_cert and client_key must be specified as a pair")); @@ -221,18 +442,20 @@ impl Operation for NotificationTarget { } else if target_type == NOTIFY_MQTT_SUB_SYS { let endpoint = kv_map .get(rustfs_config::MQTT_BROKER) + .map(String::as_str) .ok_or_else(|| s3_error!(InvalidArgument, "broker endpoint is required"))?; let topic = kv_map .get(rustfs_config::MQTT_TOPIC) + .map(String::as_str) .ok_or_else(|| s3_error!(InvalidArgument, "topic is required"))?; - let username = kv_map.get(rustfs_config::MQTT_USERNAME).copied(); - let password = kv_map.get(rustfs_config::MQTT_PASSWORD).copied(); + let username = kv_map.get(rustfs_config::MQTT_USERNAME).map(String::as_str); + let password = kv_map.get(rustfs_config::MQTT_PASSWORD).map(String::as_str); check_mqtt_broker_available(endpoint, topic, username, password) .await .map_err(|e| s3_error!(InvalidArgument, "MQTT Broker unavailable: {}", e))?; if let Some(queue_dir) = kv_map.get("queue_dir") { - validate_queue_dir(queue_dir).await?; + validate_queue_dir(queue_dir.as_str()).await?; if let Some(qos) = kv_map.get("qos") { match qos.parse::() { Ok(1) | Ok(2) => {} @@ -243,24 +466,14 @@ impl Operation for NotificationTarget { } } - let mut kvs_vec: Vec<_> = notification_body - .key_values - .into_iter() - .map(|kv| rustfs_ecstore::config::KV { - key: kv.key, - value: kv.value, - hidden_if_empty: false, - }) - .collect(); - - kvs_vec.push(rustfs_ecstore::config::KV { - key: ENABLE_KEY.to_string(), - value: EnableState::On.to_string(), - hidden_if_empty: false, - }); + let mut kvs = rustfs_ecstore::config::KVS::new(); + for (key, value) in kv_map { + kvs.insert(key, value); + } + kvs.insert(ENABLE_KEY.to_string(), EnableState::On.to_string()); info!("Setting target config for type '{}', name '{}'", target_type, target_name); - ns.set_target_config(target_type, target_name, rustfs_ecstore::config::KVS(kvs_vec)) + ns.set_target_config(target_type, target_name, kvs) .await .map_err(|e| s3_error!(InternalError, "failed to set target config: {}", e))?; @@ -278,8 +491,6 @@ impl Operation for ListNotificationTargets { let ns = get_notification_system()?; let targets = ns.get_target_values().await; - let target_count = targets.len(); - let semaphore = Arc::new(Semaphore::new(10)); let mut futures = FuturesUnordered::new(); @@ -291,18 +502,16 @@ impl Operation for ListNotificationTargets { Ok(Ok(true)) => "online", _ => "offline", }; - NotificationEndpoint { - account_id: target.id().id.clone(), - service: target.id().name.to_string(), - status: status.to_string(), - } + ((target.id().id.clone(), target.id().name.to_string()), status.to_string()) }); } - let mut notification_endpoints = Vec::with_capacity(target_count); - while let Some(endpoint) = futures.next().await { - notification_endpoints.push(endpoint); + let mut runtime_statuses = HashMap::new(); + while let Some((key, status)) = futures.next().await { + runtime_statuses.insert(key, status); } + let config = ns.config.read().await.clone(); + let notification_endpoints = merge_notification_endpoints(&config, runtime_statuses); let data = serde_json::to_vec(&NotificationEndpointsResponse { notification_endpoints }) .map_err(|e| s3_error!(InternalError, "failed to serialize targets: {}", e))?; @@ -320,16 +529,32 @@ impl Operation for ListTargetsArns { check_permissions(&req).await?; let ns = get_notification_system()?; - let active_targets = ns.get_active_targets().await; + let targets = ns.get_target_values().await; let region = req .region .clone() .ok_or_else(|| s3_error!(InvalidRequest, "region not found"))?; + let semaphore = Arc::new(Semaphore::new(10)); + let mut futures = FuturesUnordered::new(); - let data_target_arn_list: Vec<_> = active_targets - .iter() - .map(|id| id.to_arn(region.as_str()).to_string()) - .collect(); + for target in targets { + let sem = Arc::clone(&semaphore); + futures.push(async move { + let _permit = sem.acquire().await; + let status = match timeout(Duration::from_secs(3), target.is_active()).await { + Ok(Ok(true)) => "online", + _ => "offline", + }; + (target.id(), status.to_string()) + }); + } + + let mut target_statuses = Vec::new(); + while let Some(target_status) = futures.next().await { + target_statuses.push(target_status); + } + + let data_target_arn_list = collect_online_target_arns(region.as_str(), target_statuses); let data = serde_json::to_vec(&data_target_arn_list) .map_err(|e| s3_error!(InternalError, "failed to serialize targets: {}", e))?; @@ -348,6 +573,10 @@ impl Operation for RemoveNotificationTarget { check_permissions(&req).await?; let ns = get_notification_system()?; + let config_snapshot = ns.config.read().await.clone(); + if let Some(reason) = target_mutation_block_reason(&config_snapshot, target_type, target_name) { + return Err(s3_error!(InvalidRequest, "{reason}")); + } info!("Removing target config for type '{}', name '{}'", target_type, target_name); ns.remove_target_config(target_type, target_name) @@ -372,3 +601,284 @@ fn extract_target_params<'a>(params: &'a Params<'_, '_>) -> S3Result<(&'a str, & let target_name = extract_param(params, "target_name")?; Ok((target_type, target_name)) } + +#[cfg(test)] +mod tests { + use super::*; + use rustfs_ecstore::config::{KV, KVS}; + use rustfs_targets::arn::TargetID; + use std::collections::{HashMap, HashSet}; + use temp_env::{with_var, with_vars}; + + fn enabled_kvs(value: &str) -> KVS { + KVS(vec![KV { + key: ENABLE_KEY.to_string(), + value: value.to_string(), + hidden_if_empty: false, + }]) + } + + #[test] + fn merge_notification_endpoints_keeps_configured_targets_after_runtime_loss() { + let mut cfg_map = HashMap::new(); + cfg_map.insert( + NOTIFY_WEBHOOK_SUB_SYS.to_string(), + HashMap::from([("webhook-a".to_string(), enabled_kvs("on"))]), + ); + cfg_map.insert( + NOTIFY_MQTT_SUB_SYS.to_string(), + HashMap::from([("mqtt-a".to_string(), enabled_kvs("on"))]), + ); + let config = Config(cfg_map); + + let runtime = HashMap::from([(("webhook-a".to_string(), "webhook".to_string()), "online".to_string())]); + let merged = merge_notification_endpoints(&config, runtime); + + let mqtt = merged + .iter() + .find(|entry| entry.account_id == "mqtt-a" && entry.service == "mqtt") + .expect("mqtt-a should be present"); + assert_eq!(mqtt.status, "offline"); + assert_eq!(mqtt.source, NotificationEndpointSource::Config); + + let webhook = merged + .iter() + .find(|entry| entry.account_id == "webhook-a" && entry.service == "webhook") + .expect("webhook-a should be present"); + assert_eq!(webhook.status, "online"); + assert_eq!(webhook.source, NotificationEndpointSource::Config); + } + + #[test] + fn merge_notification_endpoints_skips_disabled_and_default_entries() { + let mut webhook_targets = HashMap::new(); + webhook_targets.insert(DEFAULT_DELIMITER.to_string(), enabled_kvs("on")); + webhook_targets.insert("webhook-disabled".to_string(), enabled_kvs("off")); + webhook_targets.insert("webhook-enabled".to_string(), enabled_kvs("on")); + let config = Config(HashMap::from([(NOTIFY_WEBHOOK_SUB_SYS.to_string(), webhook_targets)])); + + let runtime = HashMap::from([ + (("webhook-enabled".to_string(), "webhook".to_string()), "online".to_string()), + (("env-only".to_string(), "mqtt".to_string()), "offline".to_string()), + ]); + let merged = merge_notification_endpoints(&config, runtime); + + let env_only = merged + .iter() + .find(|entry| entry.account_id == "env-only" && entry.service == "mqtt") + .expect("env-only should be present"); + assert_eq!(env_only.status, "offline"); + assert_eq!(env_only.source, NotificationEndpointSource::Runtime); + + let enabled = merged + .iter() + .find(|entry| entry.account_id == "webhook-enabled" && entry.service == "webhook") + .expect("webhook-enabled should be present"); + assert_eq!(enabled.status, "online"); + assert_eq!(enabled.source, NotificationEndpointSource::Config); + } + + #[test] + fn merge_notification_endpoints_marks_env_and_mixed_sources() { + let config = Config(HashMap::from([ + ( + NOTIFY_WEBHOOK_SUB_SYS.to_string(), + HashMap::from([("mixed-target".to_string(), enabled_kvs("on"))]), + ), + ( + NOTIFY_MQTT_SUB_SYS.to_string(), + HashMap::from([("config-target".to_string(), enabled_kvs("on"))]), + ), + ])); + + with_vars( + [ + ("RUSTFS_NOTIFY_WEBHOOK_ENDPOINT_MIXED-TARGET", Some("https://example.com/hook")), + ("RUSTFS_NOTIFY_WEBHOOK_ENABLE_ENV-ONLY", Some("on")), + ("RUSTFS_NOTIFY_WEBHOOK_ENDPOINT_ENV-ONLY", Some("https://example.com/env")), + ], + || { + let runtime = HashMap::from([ + (("mixed-target".to_string(), "webhook".to_string()), "online".to_string()), + (("env-only".to_string(), "webhook".to_string()), "online".to_string()), + ]); + let merged = merge_notification_endpoints(&config, runtime); + + let mixed = merged + .iter() + .find(|entry| entry.account_id == "mixed-target") + .expect("mixed target should be present"); + assert_eq!(mixed.source, NotificationEndpointSource::Mixed); + + let env_only = merged + .iter() + .find(|entry| entry.account_id == "env-only") + .expect("env-only target should be present"); + assert_eq!(env_only.source, NotificationEndpointSource::Env); + + let config_only = merged + .iter() + .find(|entry| entry.account_id == "config-target") + .expect("config target should be present"); + assert_eq!(config_only.source, NotificationEndpointSource::Config); + }, + ); + } + + #[test] + fn target_mutation_block_reason_rejects_env_managed_target() { + with_vars( + [ + ("RUSTFS_NOTIFY_WEBHOOK_ENABLE_PRIMARY", Some("on")), + ("RUSTFS_NOTIFY_WEBHOOK_ENDPOINT_PRIMARY", Some("https://example.com/hook")), + ], + || { + let config = Config(HashMap::new()); + let reason = target_mutation_block_reason(&config, NOTIFY_WEBHOOK_SUB_SYS, "primary"); + assert!(reason.is_some()); + assert!(reason.unwrap().contains("managed by environment variables")); + }, + ); + } + + #[test] + fn target_mutation_block_reason_rejects_mixed_target() { + with_var("RUSTFS_NOTIFY_WEBHOOK_ENDPOINT_PRIMARY", Some("https://example.com/hook"), || { + let config = Config(HashMap::from([( + NOTIFY_WEBHOOK_SUB_SYS.to_string(), + HashMap::from([("primary".to_string(), enabled_kvs("on"))]), + )])); + let reason = target_mutation_block_reason(&config, NOTIFY_WEBHOOK_SUB_SYS, "primary"); + assert!(reason.is_some()); + assert!(reason.unwrap().contains("both persisted config and environment variables")); + }); + } + + #[test] + fn target_mutation_block_reason_allows_config_only_target() { + let target_name = "config-only-target"; + let config = Config(HashMap::from([( + NOTIFY_WEBHOOK_SUB_SYS.to_string(), + HashMap::from([(target_name.to_string(), enabled_kvs("on"))]), + )])); + assert!(target_mutation_block_reason(&config, NOTIFY_WEBHOOK_SUB_SYS, target_name).is_none()); + } + + #[test] + fn merge_notification_endpoints_marks_disabled_config_with_env_override_as_mixed() { + let config = Config(HashMap::from([( + NOTIFY_WEBHOOK_SUB_SYS.to_string(), + HashMap::from([("mixed-disabled".to_string(), enabled_kvs("off"))]), + )])); + + with_vars( + [ + ("RUSTFS_NOTIFY_WEBHOOK_ENABLE_MIXED-DISABLED", Some("on")), + ("RUSTFS_NOTIFY_WEBHOOK_ENDPOINT_MIXED-DISABLED", Some("https://example.com/hook")), + ], + || { + let merged = merge_notification_endpoints(&config, HashMap::new()); + let mixed = merged + .iter() + .find(|entry| entry.account_id == "mixed-disabled") + .expect("mixed target should be present"); + assert_eq!(mixed.source, NotificationEndpointSource::Mixed); + assert_eq!(mixed.status, "offline"); + }, + ); + } + + #[test] + fn merge_notification_endpoints_includes_env_only_target_without_runtime_status() { + let config = Config(HashMap::new()); + + with_vars( + [ + ("RUSTFS_NOTIFY_WEBHOOK_ENABLE_ENV-ONLY", Some("on")), + ("RUSTFS_NOTIFY_WEBHOOK_ENDPOINT_ENV-ONLY", Some("https://example.com/env")), + ], + || { + let merged = merge_notification_endpoints(&config, HashMap::new()); + let env_only = merged + .iter() + .find(|entry| entry.account_id == "env-only") + .expect("env-only target should be present"); + assert_eq!(env_only.source, NotificationEndpointSource::Env); + assert_eq!(env_only.status, "offline"); + }, + ); + } + + #[test] + fn collect_validated_key_values_rejects_duplicate_keys() { + let allowed_keys: HashSet<&str> = ["endpoint", "auth_token"].into_iter().collect(); + let key_values = vec![ + KeyValue { + key: "endpoint".to_string(), + value: "https://example.com/one".to_string(), + }, + KeyValue { + key: "endpoint".to_string(), + value: "https://example.com/two".to_string(), + }, + ]; + + let err = collect_validated_key_values(&key_values, &allowed_keys, NOTIFY_WEBHOOK_SUB_SYS).unwrap_err(); + assert!(err.to_string().contains("duplicate key")); + } + + #[test] + fn merge_notification_endpoints_marks_mixed_with_case_insensitive_instance_id() { + let config = Config(HashMap::from([( + NOTIFY_WEBHOOK_SUB_SYS.to_string(), + HashMap::from([("PrimaryCase".to_string(), enabled_kvs("on"))]), + )])); + + with_vars( + [ + ("RUSTFS_NOTIFY_WEBHOOK_ENABLE_PRIMARYCASE", Some("on")), + ("RUSTFS_NOTIFY_WEBHOOK_ENDPOINT_PRIMARYCASE", Some("https://example.com/hook")), + ], + || { + let runtime = HashMap::from([(("PrimaryCase".to_string(), "webhook".to_string()), "online".to_string())]); + let merged = merge_notification_endpoints(&config, runtime); + let mixed = merged + .iter() + .find(|entry| entry.account_id == "PrimaryCase" && entry.service == "webhook") + .expect("mixed target should be present"); + assert_eq!(mixed.source, NotificationEndpointSource::Mixed); + }, + ); + } + + #[test] + fn collect_online_target_arns_filters_offline_targets() { + let arns = collect_online_target_arns( + "us-east-1", + vec![ + (TargetID::new("webhook-a".to_string(), "webhook".to_string()), "online".to_string()), + (TargetID::new("mqtt-a".to_string(), "mqtt".to_string()), "offline".to_string()), + ], + ); + + assert_eq!(arns, vec!["arn:rustfs:sqs:us-east-1:webhook-a:webhook".to_string()]); + } + + #[test] + fn target_mutation_block_reason_allows_case_insensitive_config_target_lookup() { + let config = Config(HashMap::from([( + NOTIFY_WEBHOOK_SUB_SYS.to_string(), + HashMap::from([("PrimaryCase".to_string(), enabled_kvs("on"))]), + )])); + + with_vars( + [ + ("RUSTFS_NOTIFY_WEBHOOK_ENABLE_PRIMARYCASE", None::<&str>), + ("RUSTFS_NOTIFY_WEBHOOK_ENDPOINT_PRIMARYCASE", None::<&str>), + ], + || { + assert!(target_mutation_block_reason(&config, NOTIFY_WEBHOOK_SUB_SYS, "primarycase").is_none()); + }, + ); + } +} diff --git a/rustfs/src/admin/handlers/mod.rs b/rustfs/src/admin/handlers/mod.rs index 522cb0056..f26f3ea89 100644 --- a/rustfs/src/admin/handlers/mod.rs +++ b/rustfs/src/admin/handlers/mod.rs @@ -13,6 +13,7 @@ // limitations under the License. pub mod account_info; +pub mod audit; pub mod bucket_meta; pub mod event; pub mod group; @@ -51,6 +52,7 @@ mod tests { fn test_handler_struct_creation() { // Test that handler structs can be created let _account_handler = account_info::AccountInfoHandler {}; + let _list_audit_targets = audit::ListAuditTargets {}; let _service_handler = system::ServiceHandle {}; let _server_info_handler = system::ServerInfoHandler {}; let _inspect_data_handler = system::InspectDataHandler {}; diff --git a/rustfs/src/admin/mod.rs b/rustfs/src/admin/mod.rs index bf54583ab..588781ecc 100644 --- a/rustfs/src/admin/mod.rs +++ b/rustfs/src/admin/mod.rs @@ -25,8 +25,8 @@ mod console_test; mod route_registration_test; use handlers::{ - bucket_meta, heal, health, kms, oidc, pools, profile_admin, quota, rebalance, replication, site_replication, sts, system, - tier, user, + audit, bucket_meta, heal, health, kms, oidc, pools, profile_admin, quota, rebalance, replication, site_replication, sts, + system, tier, user, }; use router::{AdminOperation, S3Router}; use s3s::route::S3Route; @@ -55,6 +55,7 @@ pub fn make_admin_route(console_enabled: bool) -> std::io::Result quota::register_quota_route(&mut r)?; bucket_meta::register_bucket_meta_route(&mut r)?; + audit::register_audit_target_route(&mut r)?; replication::register_replication_route(&mut r)?; site_replication::register_site_replication_route(&mut r)?; diff --git a/rustfs/src/admin/route_registration_test.rs b/rustfs/src/admin/route_registration_test.rs index d58556e8e..d490d4d9d 100644 --- a/rustfs/src/admin/route_registration_test.rs +++ b/rustfs/src/admin/route_registration_test.rs @@ -14,8 +14,8 @@ use crate::admin::{ handlers::{ - bucket_meta, heal, health, kms, oidc, pools, profile_admin, quota, rebalance, replication, site_replication, sts, system, - tier, user, + audit, bucket_meta, heal, health, kms, oidc, pools, profile_admin, quota, rebalance, replication, site_replication, sts, + system, tier, user, }, router::{AdminOperation, S3Router}, }; @@ -50,6 +50,7 @@ fn register_admin_routes(router: &mut S3Router) { tier::register_tier_route(router).expect("register tier route"); quota::register_quota_route(router).expect("register quota route"); bucket_meta::register_bucket_meta_route(router).expect("register bucket meta route"); + audit::register_audit_target_route(router).expect("register audit target route"); replication::register_replication_route(router).expect("register replication route"); site_replication::register_site_replication_route(router).expect("register site replication route"); profile_admin::register_profiling_route(router).expect("register profile route"); @@ -60,7 +61,6 @@ fn register_admin_routes(router: &mut S3Router) { #[test] fn test_register_routes_cover_representative_admin_paths() { let mut router: S3Router = S3Router::new(false); - register_admin_routes(&mut router); assert_route(&router, Method::GET, HEALTH_PREFIX); assert_route(&router, Method::HEAD, HEALTH_PREFIX); @@ -91,6 +91,9 @@ fn test_register_routes_cover_representative_admin_paths() { assert_route(&router, Method::POST, &admin_path("/v3/idp/builtin/policy/detach")); assert_route(&router, Method::GET, &admin_path("/v3/idp/builtin/policy-entities")); assert_route(&router, Method::GET, &admin_path("/v3/target/list")); + assert_route(&router, Method::GET, &admin_path("/v3/audit/target/list")); + assert_route(&router, Method::PUT, &admin_path("/v3/audit/target/audit_webhook/test-audit")); + assert_route(&router, Method::DELETE, &admin_path("/v3/audit/target/audit_webhook/test-audit/reset")); assert_route(&router, Method::GET, &admin_path("/v3/accountinfo")); assert_route(&router, Method::POST, &admin_path("/v3/service")); @@ -165,7 +168,6 @@ fn test_register_routes_cover_representative_admin_paths() { #[test] fn test_admin_alias_paths_match_existing_admin_routes() { let mut router: S3Router = S3Router::new(false); - register_admin_routes(&mut router); for (method, path) in [ diff --git a/rustfs/src/server/audit.rs b/rustfs/src/server/audit.rs index 98105be0a..7809684cd 100644 --- a/rustfs/src/server/audit.rs +++ b/rustfs/src/server/audit.rs @@ -14,13 +14,27 @@ use crate::app::context::resolve_server_config; use rustfs_audit::{AuditError, AuditResult, audit_system, init_audit_system, system::AuditSystemState}; -use rustfs_config::DEFAULT_DELIMITER; use tracing::{info, warn}; fn server_config_from_context() -> Option { resolve_server_config() } +fn has_any_audit_targets(config: &rustfs_ecstore::config::Config) -> bool { + for subsystem in [ + rustfs_config::audit::AUDIT_MQTT_SUB_SYS, + rustfs_config::audit::AUDIT_WEBHOOK_SUB_SYS, + ] { + let Some(targets) = config.0.get(subsystem) else { + continue; + }; + if targets.keys().any(|key| key != rustfs_config::DEFAULT_DELIMITER) { + return true; + } + } + false +} + /// Start the audit system. /// This function checks if the audit subsystem is configured in the global server configuration. /// If configured, it initializes and starts the audit system. @@ -55,10 +69,8 @@ pub(crate) async fn start_audit_system() -> AuditResult<()> { "The global server configuration is loaded" ); // 2. Check if the notify subsystem exists in the configuration, and skip initialization if it doesn't - let mqtt_config = server_config.get_value(rustfs_config::audit::AUDIT_MQTT_SUB_SYS, DEFAULT_DELIMITER); - let webhook_config = server_config.get_value(rustfs_config::audit::AUDIT_WEBHOOK_SUB_SYS, DEFAULT_DELIMITER); - - if mqtt_config.is_none() && webhook_config.is_none() { + let has_targets = has_any_audit_targets(&server_config); + if !has_targets { info!( target: "rustfs::main::start_audit_system", "Audit subsystem (MQTT/Webhook) is not configured, and audit system initialization is skipped." @@ -68,9 +80,7 @@ pub(crate) async fn start_audit_system() -> AuditResult<()> { info!( target: "rustfs::main::start_audit_system", - "Audit subsystem configuration detected (MQTT: {}, Webhook: {}) and started initializing the audit system.", - mqtt_config.is_some(), - webhook_config.is_some() + "Audit subsystem configuration detected and started initializing the audit system." ); // 3. Initialize and start the audit system let system = init_audit_system(); diff --git a/scripts/run.sh b/scripts/run.sh index 5aa862da5..36c4f788c 100755 --- a/scripts/run.sh +++ b/scripts/run.sh @@ -71,7 +71,7 @@ export RUSTFS_OBS_SERVICE_NAME=rustfs # Service name export RUSTFS_OBS_SERVICE_VERSION=0.1.0 # Service version export RUSTFS_OBS_ENVIRONMENT=production # Environment name development, staging, production export RUSTFS_OBS_LOGGER_LEVEL=info # Log level, supports trace, debug, info, warn, error -#export RUSTFS_OBS_LOG_STDOUT_ENABLED=true # Whether to enable local stdout logging +export RUSTFS_OBS_LOG_STDOUT_ENABLED=true # Whether to enable local stdout logging export RUSTFS_OBS_LOG_DIRECTORY="$current_dir/deploy/logs" # Log directory export RUSTFS_OBS_LOG_ROTATION_TIME="minutely" # Log rotation time unit, can be "minutely", "hourly", "daily" export RUSTFS_OBS_LOG_KEEP_FILES=10 # Number of log files to keep From eabbea46d384b3ca172409c122818d08e7f5c362 Mon Sep 17 00:00:00 2001 From: houseme Date: Sat, 4 Apr 2026 12:20:59 +0800 Subject: [PATCH 15/22] refactor(tracing): unify request-context propagation and fix tracing chain breaks (#2394) Co-authored-by: heihutu --- rustfs/src/admin/router.rs | 2 + rustfs/src/app/bucket_usecase.rs | 8 +- rustfs/src/app/multipart_usecase.rs | 3 +- rustfs/src/app/object_usecase.rs | 38 +++--- rustfs/src/protocols/client.rs | 1 + rustfs/src/server/http.rs | 39 ++++-- rustfs/src/server/layer.rs | 115 +++++++++++++++++ rustfs/src/storage/access.rs | 14 ++ rustfs/src/storage/helper.rs | 135 ++++++++++++++++++-- rustfs/src/storage/mod.rs | 1 + rustfs/src/storage/request_context.rs | 176 ++++++++++++++++++++++++++ rustfs/src/storage/timeout_wrapper.rs | 17 ++- 12 files changed, 499 insertions(+), 50 deletions(-) create mode 100644 rustfs/src/storage/request_context.rs diff --git a/rustfs/src/admin/router.rs b/rustfs/src/admin/router.rs index 3551c4365..f25441187 100644 --- a/rustfs/src/admin/router.rs +++ b/rustfs/src/admin/router.rs @@ -1442,6 +1442,7 @@ async fn authorize_replication_extension_request(req: &mut S3Request, ext_ object: None, version_id: None, region: get_global_region(), + ..Default::default() }); license_check().map_err(|er| match er.kind() { @@ -2163,6 +2164,7 @@ async fn authorize_misc_extension_request(req: &mut S3Request, route: &Mis object, version_id: None, region: get_global_region(), + ..Default::default() }); license_check().map_err(|er| match er.kind() { diff --git a/rustfs/src/app/bucket_usecase.rs b/rustfs/src/app/bucket_usecase.rs index 64e963f82..e23e5ef91 100644 --- a/rustfs/src/app/bucket_usecase.rs +++ b/rustfs/src/app/bucket_usecase.rs @@ -22,7 +22,7 @@ use crate::auth::get_condition_values; use crate::error::ApiError; use crate::server::RemoteAddr; use crate::storage::access::{ReqInfo, authorize_request, req_info_ref}; -use crate::storage::helper::OperationHelper; +use crate::storage::helper::{OperationHelper, spawn_background_with_context}; use crate::storage::s3_api::bucket::{build_list_buckets_output, build_list_objects_v2_output}; use crate::storage::s3_api::common::rustfs_owner; use crate::storage::s3_api::{acl, encryption, replication, tagging}; @@ -1494,7 +1494,11 @@ impl DefaultBucketUsecase { && let Some(store) = new_object_layer_fn() { let bucket_name = bucket.clone(); - tokio::spawn(async move { + let request_context = req + .extensions + .get::() + .cloned(); + spawn_background_with_context(request_context, async move { if let Err(err) = enqueue_transition_for_existing_objects(store, &bucket_name).await { warn!(bucket = %bucket_name, error = ?err, "failed to enqueue transition for existing objects"); } diff --git a/rustfs/src/app/multipart_usecase.rs b/rustfs/src/app/multipart_usecase.rs index fc1e069d4..0f97f1fe1 100644 --- a/rustfs/src/app/multipart_usecase.rs +++ b/rustfs/src/app/multipart_usecase.rs @@ -25,6 +25,7 @@ use crate::storage::options::{ copy_src_opts, extract_metadata, get_complete_multipart_upload_opts, get_content_sha256_with_query, get_opts, parse_copy_source_range, put_opts, }; +use crate::storage::request_context::spawn_traced; use crate::storage::s3_api::multipart::build_list_parts_output; use crate::storage::*; use bytes::Bytes; @@ -406,7 +407,7 @@ impl DefaultMultipartUsecase { }; let mpu_version_clone = mpu_version.clone(); let mpu_version_for_event = mpu_version.clone(); - tokio::spawn(async move { + spawn_traced(async move { manager .invalidate_cache_versioned(&mpu_bucket, &mpu_key, mpu_version_clone.as_deref()) .await; diff --git a/rustfs/src/app/object_usecase.rs b/rustfs/src/app/object_usecase.rs index f97b3519c..efebe94e4 100644 --- a/rustfs/src/app/object_usecase.rs +++ b/rustfs/src/app/object_usecase.rs @@ -24,11 +24,12 @@ use crate::storage::concurrency::{ }; use crate::storage::ecfs::*; use crate::storage::head_prefix::{head_prefix_not_found_message, probe_prefix_has_children}; -use crate::storage::helper::{OperationHelper, spawn_background}; +use crate::storage::helper::{OperationHelper, spawn_background, spawn_background_with_context}; use crate::storage::options::{ copy_dst_opts, copy_src_opts, del_opts, extract_metadata, extract_metadata_from_mime_with_object_name, filter_object_metadata, get_content_sha256_with_query, get_opts, normalize_content_encoding_for_storage, put_opts, }; +use crate::storage::request_context::spawn_traced; use crate::storage::s3_api::multipart::parse_list_parts_params; use crate::storage::s3_api::{acl, restore, select}; use crate::storage::timeout_wrapper::{RequestTimeoutWrapper, TimeoutConfig}; @@ -928,7 +929,7 @@ impl DefaultObjectUsecase { fn spawn_cache_invalidation(bucket: String, key: String, version_id: Option) { let manager = get_concurrency_manager(); - tokio::spawn(async move { + spawn_traced(async move { manager.invalidate_cache_versioned(&bucket, &key, version_id.as_deref()).await; }); } @@ -1015,9 +1016,9 @@ impl DefaultObjectUsecase { ))) } - fn init_get_object_bootstrap(bucket: &str, key: &str) -> S3Result { + fn init_get_object_bootstrap(bucket: &str, key: &str, request_id: &str) -> S3Result { let timeout_config = TimeoutConfig::from_env(); - let wrapper = RequestTimeoutWrapper::with_request_id(timeout_config.clone(), format!("get-{bucket}-{key}")); + let wrapper = RequestTimeoutWrapper::with_request_id(timeout_config.clone(), request_id.to_string()); let request_start = std::time::Instant::now(); let request_guard = ConcurrencyManager::track_request(); let concurrent_requests = GetObjectGuard::concurrent_requests(); @@ -1535,7 +1536,7 @@ impl DefaultObjectUsecase { .with_last_modified(last_modified_str.unwrap_or_default()); let cache_key_clone = cache_key.to_string(); - tokio::spawn(async move { + spawn_traced(async move { let manager = get_concurrency_manager(); manager.put_cached_object(cache_key_clone.clone(), cached_response).await; debug!("Object cached successfully with metadata: {}", cache_key_clone); @@ -2369,7 +2370,7 @@ impl DefaultObjectUsecase { let cache_key = ConcurrencyManager::make_cache_key(&bucket, &object, version_id.clone().as_deref()); let cache_bucket = bucket.clone(); let cache_object = object.clone(); - tokio::spawn(async move { + spawn_traced(async move { manager .invalidate_cache_versioned(&cache_bucket, &cache_object, version_id.as_deref()) .await; @@ -2599,7 +2600,12 @@ impl DefaultObjectUsecase { let _ = context.object_store(); } - let bootstrap = Self::init_get_object_bootstrap(&req.input.bucket, &req.input.key)?; + let request_id = req + .extensions + .get::() + .map(|ctx| ctx.request_id.clone()) + .unwrap_or_else(|| crate::storage::request_context::RequestContext::fallback().request_id); + let bootstrap = Self::init_get_object_bootstrap(&req.input.bucket, &req.input.key, &request_id)?; let timeout_config = bootstrap.timeout_config; let wrapper = bootstrap.wrapper; let request_start = bootstrap.request_start; @@ -3711,7 +3717,7 @@ impl DefaultObjectUsecase { let manager = get_concurrency_manager(); let bucket_clone = bucket.clone(); let deleted_objects = dobjs.clone(); - tokio::spawn(async move { + spawn_traced(async move { for dobj in deleted_objects { manager .invalidate_cache_versioned( @@ -4114,7 +4120,7 @@ impl DefaultObjectUsecase { let version_id_clone = version_id.clone(); let cache_bucket = bucket.clone(); let cache_object = object.clone(); - tokio::spawn(async move { + spawn_traced(async move { manager .invalidate_cache_versioned(&cache_bucket, &cache_object, version_id_clone.as_deref()) .await; @@ -4626,7 +4632,7 @@ impl DefaultObjectUsecase { let rreq_clone = rreq.clone(); let version_id_clone = version_id.clone(); - tokio::spawn(async move { + spawn_traced(async move { let opts = ObjectOptions { transition: TransitionOptions { restore_request: rreq_clone, @@ -4647,8 +4653,6 @@ impl DefaultObjectUsecase { object_clone, err.to_string() ); - // Note: Errors from background tasks cannot be returned to client - // Consider adding to monitoring/metrics system } else { info!("successfully restored transitioned object: {}/{}", bucket_clone, object_clone); } @@ -4721,7 +4725,7 @@ impl DefaultObjectUsecase { let (tx, rx) = mpsc::channel::>(2); let stream = ReceiverStream::new(rx); - tokio::spawn(async move { + spawn_traced(async move { let _ = tx .send(Ok(SelectObjectContentEvent::Cont(ContinuationEvent::default()))) .await; @@ -5078,7 +5082,7 @@ impl DefaultObjectUsecase { let manager = get_concurrency_manager(); let fpath_clone = fpath.clone(); let bucket_clone = bucket.clone(); - tokio::spawn(async move { + spawn_traced(async move { manager.invalidate_cache_versioned(&bucket_clone, &fpath_clone, None).await; }); @@ -5102,7 +5106,11 @@ impl DefaultObjectUsecase { }; let notify = notify.clone(); - tokio::spawn(async move { + let request_context = req + .extensions + .get::() + .cloned(); + spawn_background_with_context(request_context, async move { notify.notify(event_args).await; }); } diff --git a/rustfs/src/protocols/client.rs b/rustfs/src/protocols/client.rs index 65347d84f..766d3c147 100644 --- a/rustfs/src/protocols/client.rs +++ b/rustfs/src/protocols/client.rs @@ -81,6 +81,7 @@ impl ProtocolStorageClient { object: params.object, version_id: None, region: None, + request_context: Some(crate::storage::request_context::RequestContext::fallback()), }); let req = S3Request { diff --git a/rustfs/src/server/http.rs b/rustfs/src/server/http.rs index 0fdcf3d1d..58a5b7af6 100644 --- a/rustfs/src/server/http.rs +++ b/rustfs/src/server/http.rs @@ -21,7 +21,10 @@ use crate::server::{ ReadinessGateLayer, RemoteAddr, ServiceState, ServiceStateManager, compress::{CompressionConfig, PathAwareCompressionPredicate, PathCategoryInjectionLayer}, hybrid::hybrid, - layer::{AdminChunkedContentLengthCompatLayer, ConditionalCorsLayer, ObjectAttributesEtagFixLayer, RedirectLayer}, + layer::{ + AdminChunkedContentLengthCompatLayer, ConditionalCorsLayer, ObjectAttributesEtagFixLayer, RedirectLayer, + RequestContextLayer, + }, tls_material::{TlsAcceptorHolder, TlsHandshakeFailureKind, TlsMaterialSnapshot, spawn_reload_loop}, }; use crate::storage; @@ -593,17 +596,18 @@ fn process_connection( // 2. AddExtensionLayer — per-connection raw socket addr (TrustedProxy) // 3. TrustedProxyLayer — conditional, parses X-Forwarded-For // 4. SetRequestIdLayer — generates X-Request-ID - // 5. AdminChunkedContentLengthCompatLayer — admin API compat - // 6. CatchPanicLayer — panic → 500 - // 7. ReadinessGateLayer — blocks until ready - // 8. KeystoneAuthLayer — X-Auth-Token validation - // 9. TraceLayer — request/response tracing + metrics - // 10. PropagateRequestIdLayer — X-Request-ID → response - // 11. PathCategoryInjectionLayer — injects path category for compression - // 12. CompressionLayer — response compression (whitelist, path-aware) - // 13. ObjectAttributesEtagFixLayer — ETag fix for GetObjectAttributes - // 14. ConditionalCorsLayer — S3 API CORS - // 15. RedirectLayer — console redirect (conditional) + // 5. RequestContextLayer — creates RequestContext in extensions + // 6. AdminChunkedContentLengthCompatLayer — admin API compat + // 7. CatchPanicLayer — panic → 500 + // 8. ReadinessGateLayer — blocks until ready + // 9. KeystoneAuthLayer — X-Auth-Token validation + // 10. TraceLayer — request/response tracing + metrics + // 11. PropagateRequestIdLayer — X-Request-ID → response + // 12. PathCategoryInjectionLayer — injects path category for compression + // 13. CompressionLayer — response compression (whitelist, path-aware) + // 14. ObjectAttributesEtagFixLayer — ETag fix for GetObjectAttributes + // 15. ConditionalCorsLayer — S3 API CORS + // 16. RedirectLayer — console redirect (conditional) // ───────────────────────────────────────────────────────────── let hybrid_service = ServiceBuilder::new() // NOTE: Both extension types are intentionally inserted to maintain compatibility: @@ -619,6 +623,7 @@ fn process_connection( // Pre-computed in ConnectionContext to avoid per-connection is_enabled() check. .option_layer(trusted_proxy_layer) .layer(SetRequestIdLayer::x_request_id(MakeRequestUuid)) + .layer(RequestContextLayer) .layer(AdminChunkedContentLengthCompatLayer) .layer(CatchPanicLayer::new()) // CRITICAL: Insert ReadinessGateLayer before business logic @@ -687,6 +692,12 @@ fn process_connection( debug!("http started method: {}, url path: {}", request.method(), request.uri().path()); let labels = [("key_request_method", request.method().to_string())]; counter!("rustfs.api.requests.total", &labels).increment(1); + // Aggregate request body size for throughput monitoring (lightweight) + if let Some(cl) = request.headers().get("content-length") + && let Some(len) = cl.to_str().ok().and_then(|s| s.parse::().ok()) + { + counter!("rustfs.request.body.bytes_total", "direction" => "request").increment(len); + } }) .on_response(|response: &Response<_>, latency: Duration, span: &Span| { span.record("status_code", tracing::field::display(response.status())); @@ -695,6 +706,8 @@ fn process_connection( debug!("http response generated in {:?}", latency) }) .on_body_chunk(|chunk: &Bytes, latency: Duration, span: &Span| { + // Always track aggregate body bytes (lightweight counter, no debug logging) + counter!("rustfs.request.body.bytes_total", "direction" => "response").increment(chunk.len() as u64); #[cfg(feature = "tracing-chunk-debug")] { let _enter = span.enter(); @@ -703,7 +716,7 @@ fn process_connection( } #[cfg(not(feature = "tracing-chunk-debug"))] { - let _ = (chunk, latency, span); + let _ = (latency, span); } }) .on_eos(|_trailers: Option<&HeaderMap>, stream_duration: Duration, span: &Span| { diff --git a/rustfs/src/server/layer.rs b/rustfs/src/server/layer.rs index da2c95046..13a198141 100644 --- a/rustfs/src/server/layer.rs +++ b/rustfs/src/server/layer.rs @@ -17,19 +17,134 @@ use crate::server::cors; use crate::server::hybrid::HybridBody; use crate::server::{ADMIN_PREFIX, CONSOLE_PREFIX, MINIO_ADMIN_PREFIX, MINIO_ADMIN_V3_PREFIX, RPC_PREFIX, RUSTFS_ADMIN_PREFIX}; use crate::storage::apply_cors_headers; +use crate::storage::request_context::{RequestContext, extract_request_id_from_headers}; use bytes::Bytes; use http::{HeaderMap, HeaderValue, Method, Request as HttpRequest, Response, StatusCode}; use http_body::Body; use http_body_util::BodyExt; use hyper::body::Incoming; +use opentelemetry::global; +use opentelemetry::trace::TraceContextExt; use rustfs_utils::get_env_opt_str; +use rustfs_utils::http::headers::AMZ_REQUEST_ID; use std::future::Future; use std::pin::Pin; use std::sync::Arc; use std::task::{Context, Poll}; +use std::time::Instant; use tower::{Layer, Service}; use tracing::debug; +/// A carrier that adapts [`HeaderMap`] for OpenTelemetry trace context propagation. +struct HeaderMapCarrier<'a>(&'a HeaderMap); + +impl<'a> opentelemetry::propagation::Extractor for HeaderMapCarrier<'a> { + fn get(&self, key: &str) -> Option<&str> { + self.0.get(key).and_then(|v| v.to_str().ok()) + } + + fn keys(&self) -> Vec<&str> { + self.0.keys().map(|k| k.as_str()).collect() + } + + fn get_all(&self, key: &str) -> Option> { + let headers = self + .0 + .get_all(key) + .iter() + .filter_map(|value| value.to_str().ok()) + .collect::>(); + + if headers.is_empty() { None } else { Some(headers) } + } +} + +/// Tower middleware layer that creates a canonical [`RequestContext`] from HTTP headers +/// and injects it into `request.extensions()`. +/// +/// This layer must be placed after `SetRequestIdLayer` in the middleware stack, +/// as it reads the `x-request-id` header that `SetRequestIdLayer` generates. +/// +/// Additionally, it sets the `x-amz-request-id` request header for S3 compatibility +/// if not already present. +#[derive(Clone, Default)] +pub struct RequestContextLayer; + +impl Layer for RequestContextLayer { + type Service = RequestContextService; + + fn layer(&self, inner: S) -> Self::Service { + RequestContextService { inner } + } +} + +/// Service that injects [`RequestContext`] into every request. +#[derive(Clone)] +pub struct RequestContextService { + inner: S, +} + +impl Service> for RequestContextService +where + S: Service>, +{ + type Response = S::Response; + type Error = S::Error; + type Future = S::Future; + + fn poll_ready(&mut self, cx: &mut Context<'_>) -> Poll> { + self.inner.poll_ready(cx) + } + + fn call(&mut self, mut req: HttpRequest) -> Self::Future { + let request_id = extract_request_id_from_headers(req.headers()); + + // Extract OpenTelemetry trace/span context from incoming headers + let parent_cx = global::get_text_map_propagator(|propagator| propagator.extract(&HeaderMapCarrier(req.headers()))); + let span_ref = parent_cx.span(); + let span_context = span_ref.span_context(); + let trace_id = if span_context.is_valid() { + Some(span_context.trace_id().to_string()) + } else { + None + }; + let span_id = if span_context.is_valid() { + Some(span_context.span_id().to_string()) + } else { + None + }; + + // Preserve the upstream x-amz-request-id if present (S3 client forwarding), + // otherwise fall back to the canonical request_id. + let x_amz_request_id = req + .headers() + .get(AMZ_REQUEST_ID) + .and_then(|v| v.to_str().ok()) + .map(String::from) + .unwrap_or_else(|| request_id.clone()); + + let ctx = RequestContext { + request_id: request_id.clone(), + x_amz_request_id, + trace_id, + span_id, + start_time: Instant::now(), + }; + + req.extensions_mut().insert(ctx); + + // Set x-amz-request-id for S3 compatibility downstream + if !req.headers().contains_key(AMZ_REQUEST_ID) + && let Ok(val) = HeaderValue::from_str(&request_id) + { + req.headers_mut() + .insert(http::header::HeaderName::from_static(AMZ_REQUEST_ID), val); + } + + self.inner.call(req) + } +} + /// Redirect layer that redirects browser requests to the console #[derive(Clone)] pub struct RedirectLayer; diff --git a/rustfs/src/storage/access.rs b/rustfs/src/storage/access.rs index a681dda4e..86937d650 100644 --- a/rustfs/src/storage/access.rs +++ b/rustfs/src/storage/access.rs @@ -17,6 +17,7 @@ use crate::auth::{check_key_valid, get_condition_values_with_query, get_session_ use crate::error::ApiError; use crate::license::license_check; use crate::server::RemoteAddr; +use crate::storage::request_context::RequestContext; use metrics::counter; use rustfs_ecstore::bucket::metadata_sys; use rustfs_ecstore::bucket::policy_sys::PolicySys; @@ -45,6 +46,7 @@ pub(crate) struct ReqInfo { pub version_id: Option, #[allow(dead_code)] pub region: Option, + pub request_context: Option, } #[derive(Clone, Debug)] @@ -67,6 +69,15 @@ fn ext_req_info_mut(ext: &mut http::Extensions) -> S3Result<&mut ReqInfo> { .ok_or_else(|| s3_error!(InternalError, "ReqInfo not found in request extensions")) } +/// Extract the canonical `RequestContext` from a request, checking both +/// the request extensions directly and the `ReqInfo.request_context` field. +pub(crate) fn request_context_from_req(req: &S3Request) -> Option { + req.extensions + .get::() + .cloned() + .or_else(|| req.extensions.get::().and_then(|ri| ri.request_context.clone())) +} + #[derive(Clone, Debug)] pub(crate) struct ObjectTagConditions { bucket: String, @@ -731,10 +742,13 @@ impl S3Access for FS { (None, false) }; + let request_context = cx.extensions_mut().get::().cloned(); + let req_info = ReqInfo { cred, is_owner, region: rustfs_ecstore::global::get_global_region(), + request_context, ..Default::default() }; diff --git a/rustfs/src/storage/helper.rs b/rustfs/src/storage/helper.rs index 4a2bd7352..968dfb363 100644 --- a/rustfs/src/storage/helper.rs +++ b/rustfs/src/storage/helper.rs @@ -12,7 +12,9 @@ // See the License for the specific language governing permissions and // limitations under the License. -use crate::storage::access::ReqInfo; +use crate::storage::access::{ReqInfo, request_context_from_req}; +use crate::storage::request_context::{RequestContext, extract_request_id_from_headers}; +use hashbrown::HashMap; use http::StatusCode; use rustfs_audit::{ entity::{ApiDetails, ApiDetailsBuilder, AuditEntryBuilder}, @@ -24,10 +26,13 @@ use rustfs_s3_common::record_s3_op; use rustfs_s3_common::{EventName, S3Operation}; use rustfs_utils::{ extract_params_header, extract_req_params, extract_resp_elements, get_request_host, get_request_port, get_request_user_agent, + http::headers::AMZ_REQUEST_ID, }; use s3s::{S3Request, S3Response, S3Result}; +use serde_json::Value; use std::future::Future; use tokio::runtime::{Builder, Handle}; +use tracing::{Instrument, info_span}; /// Schedules an asynchronous task on the current runtime; /// if there is no runtime, creates a minimal runtime execution on a new thread. @@ -46,12 +51,30 @@ where } } +/// Spawn a background task with request context correlation. +/// Creates a child span with the request_id for tracing continuity, +/// ensuring audit/notify tasks can be traced back to the original request. +pub(crate) fn spawn_background_with_context(request_context: Option, fut: F) +where + F: Future + Send + 'static, +{ + match request_context { + Some(ctx) => { + let request_id = ctx.request_id.clone(); + let span = info_span!("background-task", request_id = %request_id); + spawn_background(Instrument::instrument(fut, span)); + } + None => spawn_background(fut), + } +} + /// A unified helper structure for building and distributing audit logs and event notifications via RAII mode at the end of an S3 operation scope. pub struct OperationHelper { audit_builder: Option, api_builder: ApiDetailsBuilder, event_builder: Option, start_time: std::time::Instant, + request_context: Option, } impl OperationHelper { @@ -95,18 +118,21 @@ impl OperationHelper { api_builder = api_builder.object(&object_key); } // Audit builder - let mut audit_builder = AuditEntryBuilder::new("1.0", event, trigger, ApiDetails::default()) + // Resolve canonical request context and request_id in a single pass: + // RequestContext.request_id > extract_request_id_from_headers() > "unknown" + let request_context = request_context_from_req(req); + let request_id = request_context + .as_ref() + .map(|ctx| ctx.request_id.clone()) + .unwrap_or_else(|| extract_request_id_from_headers(&req.headers)); + + let audit_builder = AuditEntryBuilder::new("1.0", event, trigger, ApiDetails::default()) .remote_host(remote_host) .user_agent(get_request_user_agent(&req.headers)) .req_host(get_request_host(&req.headers)) .req_path(req.uri.path().to_string()) - .req_query(extract_req_params(req)); - - if let Some(req_id) = req.headers.get("x-amz-request-id") - && let Ok(id_str) = req_id.to_str() - { - audit_builder = audit_builder.request_id(id_str); - } + .req_query(extract_req_params(req)) + .request_id(&request_id); let event_object = ObjectInfo { bucket: bucket.clone(), @@ -115,6 +141,12 @@ impl OperationHelper { }; let mut req_params = extract_params_header(&req.headers); + // Inject x-amz-request-id from RequestContext into req_params for event correlation + if let Some(ref ctx) = request_context { + req_params + .entry(AMZ_REQUEST_ID.to_string()) + .or_insert_with(|| ctx.x_amz_request_id.clone()); + } if let Some(principal_id) = req_info .and_then(|info| info.cred.as_ref()) .map(|cred| cred.access_key.clone()) @@ -141,7 +173,11 @@ impl OperationHelper { audit_builder: Some(audit_builder), api_builder, event_builder: Some(event_builder), - start_time: std::time::Instant::now(), + start_time: request_context + .as_ref() + .map(|ctx| ctx.start_time) + .unwrap_or_else(std::time::Instant::now), + request_context, } } @@ -211,6 +247,20 @@ impl OperationHelper { final_builder = final_builder.access_key(&sk); } + // Inject OpenTelemetry trace context into audit tags for distributed tracing correlation + if let Some(ref ctx) = self.request_context + && (ctx.trace_id.is_some() || ctx.span_id.is_some()) + { + let mut tags = HashMap::new(); + if let Some(ref tid) = ctx.trace_id { + tags.insert("traceId".to_string(), Value::String(tid.clone())); + } + if let Some(ref sid) = ctx.span_id { + tags.insert("spanId".to_string(), Value::String(sid.clone())); + } + final_builder = final_builder.tags(tags); + } + self.audit_builder = Some(final_builder); self.api_builder = ApiDetailsBuilder(api_details); // Store final details for Drop use } @@ -234,7 +284,8 @@ impl Drop for OperationHelper { fn drop(&mut self) { // Distribute audit logs if let Some(builder) = self.audit_builder.take() { - spawn_background(async move { + let ctx = self.request_context.clone(); + spawn_background_with_context(ctx, async move { AuditLogger::log(builder.build()).await; }); } @@ -246,7 +297,8 @@ impl Drop for OperationHelper { let event_args = builder.build(); // Avoid generating notifications for copy requests if !event_args.is_replication_request() { - spawn_background(async move { + let ctx = self.request_context.clone(); + spawn_background_with_context(ctx, async move { notifier_global::notify(event_args).await; }); } @@ -305,4 +357,63 @@ mod tests { assert_eq!(event_args.version_id, "version-123"); assert_eq!(event_args.req_params.get("principalId").map(String::as_str), Some("notifyTag")); } + + #[test] + fn operation_helper_prioritizes_request_context_for_request_id() { + let input = DeleteObjectTaggingInput::builder() + .bucket("test-bucket".to_string()) + .key("test-key".to_string()) + .build() + .unwrap(); + let mut req = build_request(input, Method::DELETE, Uri::from_static("/test-bucket/test-key")); + req.headers.insert("host", HeaderValue::from_static("example.com")); + req.headers.insert("user-agent", HeaderValue::from_static("rustfs-test")); + + // Insert RequestContext (set by ingress layer) with a specific request_id + req.extensions.insert(RequestContext { + request_id: "ingress-canonical-uuid".to_string(), + x_amz_request_id: "ingress-canonical-uuid".to_string(), + trace_id: None, + span_id: None, + start_time: std::time::Instant::now(), + }); + + req.extensions.insert(ReqInfo { + bucket: Some("test-bucket".to_string()), + object: Some("test-key".to_string()), + ..Default::default() + }); + + let helper = OperationHelper::new(&req, EventName::ObjectAccessedGet, S3Operation::GetObject); + + // Verify the helper stored the RequestContext + assert!(helper.request_context.is_some()); + assert_eq!(helper.request_context.as_ref().unwrap().request_id, "ingress-canonical-uuid"); + } + + #[test] + fn operation_helper_no_request_context_when_absent() { + let input = DeleteObjectTaggingInput::builder() + .bucket("test-bucket".to_string()) + .key("test-key".to_string()) + .build() + .unwrap(); + let mut req = build_request(input, Method::DELETE, Uri::from_static("/test-bucket/test-key")); + req.headers.insert("host", HeaderValue::from_static("example.com")); + req.headers.insert("user-agent", HeaderValue::from_static("rustfs-test")); + req.headers + .insert("x-amz-request-id", HeaderValue::from_static("amz-header-uuid")); + + // No RequestContext inserted + req.extensions.insert(ReqInfo { + bucket: Some("test-bucket".to_string()), + object: Some("test-key".to_string()), + ..Default::default() + }); + + let helper = OperationHelper::new(&req, EventName::ObjectAccessedGet, S3Operation::GetObject); + + // Verify the helper has no RequestContext + assert!(helper.request_context.is_none()); + } } diff --git a/rustfs/src/storage/mod.rs b/rustfs/src/storage/mod.rs index 52de62dae..8b514ce33 100644 --- a/rustfs/src/storage/mod.rs +++ b/rustfs/src/storage/mod.rs @@ -21,6 +21,7 @@ pub(crate) mod entity; pub(crate) mod helper; pub mod lock_optimizer; pub mod options; +pub mod request_context; pub mod rpc; pub(crate) mod s3_api; mod sse; diff --git a/rustfs/src/storage/request_context.rs b/rustfs/src/storage/request_context.rs new file mode 100644 index 000000000..9e4514282 --- /dev/null +++ b/rustfs/src/storage/request_context.rs @@ -0,0 +1,176 @@ +// 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. + +//! Canonical request context carried through the entire request lifecycle. +//! +//! # Architecture +//! +//! ```text +//! HTTP Ingress (SetRequestIdLayer) +//! → generates x-request-id UUID +//! → RequestContextLayer creates RequestContext +//! → stores in request.extensions() +//! → sets x-amz-request-id header +//! Auth (FS::check) +//! → copies RequestContext into ReqInfo.request_context +//! Storage (FS methods) +//! → reads ReqInfo for bucket/object/version +//! → reads RequestContext for request_id/trace_id/span_id +//! Timeout Wrapper +//! → receives canonical request_id from caller +//! → passes to deadlock_detector.register_request() +//! OperationHelper +//! → reads RequestContext.request_id for audit log +//! → spawn_background_with_context() for audit/notify +//! tokio::spawn (request-internal) +//! → spawn_traced() = tokio::spawn + .instrument(Span::current()) +//! ``` +//! +//! # Frozen Rules (T00 Guardrails) +//! +//! ## request-id +//! - Canonical source: HTTP ingress `x-request-id` header (set by `SetRequestIdLayer`) +//! - `x-amz-request_id` is an alias for S3 compatibility, always equal to `request_id` +//! - Internal modules MUST NOT generate a second request-id under the name `request_id` +//! - Internal identifiers for sub-operations should use `operation_id` or `subtask_id` +//! +//! ## tokio::spawn usage +//! - **Request-internal tasks** (cache invalidation, metrics, read/write subtasks): +//! Use `spawn_traced()` which wraps `tokio::spawn` with `.instrument(Span::current())` +//! - **Post-request side effects** (audit flush, notify, replication enqueue): +//! Use `spawn_background_with_context()` which creates a correlated child span +//! with explicit `request_id` +//! - **Infrastructure tasks** (server loop, TLS reload, deadlock detection): +//! Plain `tokio::spawn` is acceptable; these are not request-scoped +//! - NEVER use bare `tokio::spawn` in request-handling code paths + +use http::HeaderMap; +use rustfs_utils::http::headers::AMZ_REQUEST_ID; +use std::time::Instant; + +/// Canonical request context carried through the entire request lifecycle. +/// +/// Created exactly once at HTTP ingress. Cloned by value; never mutated after creation. +#[derive(Clone, Debug)] +pub struct RequestContext { + /// Canonical request ID (from `x-request-id` header, set by `SetRequestIdLayer`). + pub request_id: String, + /// S3-compatible request ID alias (preserves upstream `x-amz-request-id` if present, + /// otherwise equals `request_id`). + pub x_amz_request_id: String, + /// OpenTelemetry trace ID (if present from upstream propagation). + pub trace_id: Option, + /// OpenTelemetry span ID (if present from upstream propagation). + pub span_id: Option, + /// Request ingress timestamp. + pub start_time: Instant, +} + +impl RequestContext { + /// Create a fallback `RequestContext` for paths that bypass HTTP ingress. + /// Generates a `req-{uuid}` format request-id. + pub fn fallback() -> Self { + let id = format!("req-{}", &uuid::Uuid::new_v4().to_string()[..8]); + Self { + request_id: id.clone(), + x_amz_request_id: id, + trace_id: None, + span_id: None, + start_time: Instant::now(), + } + } +} + +/// Extract the canonical request ID from HTTP headers. +/// +/// Priority: +/// 1. `x-request-id` (primary, set by `SetRequestIdLayer`) +/// 2. `x-amz-request-id` (fallback, from S3 client forwarding) +/// 3. `"unknown"` (no header present) +pub fn extract_request_id_from_headers(headers: &HeaderMap) -> String { + headers + .get("x-request-id") + .and_then(|v| v.to_str().ok()) + .map(String::from) + .or_else(|| headers.get(AMZ_REQUEST_ID).and_then(|v| v.to_str().ok()).map(String::from)) + .unwrap_or_else(|| "unknown".to_string()) +} + +/// Spawn a request-internal task that inherits the current tracing span. +/// +/// Use this for tasks that are part of the request processing pipeline +/// (e.g., cache invalidation, metrics recording, read/write subtasks). +/// +/// # Rules +/// - Do NOT use this for post-request side effects (audit, notify). +/// Use `crate::storage::helper::spawn_background_with_context` instead. +/// - Do NOT use bare `tokio::spawn` in request-handling code paths. +pub fn spawn_traced(fut: F) +where + F: std::future::Future + Send + 'static, +{ + tokio::spawn(tracing::Instrument::instrument(fut, tracing::Span::current())); +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn test_request_context_clone_send_sync() { + fn assert_clone_send_sync() {} + assert_clone_send_sync::(); + } + + #[test] + fn test_request_context_fallback_generates_id() { + let ctx = RequestContext::fallback(); + assert!(ctx.request_id.starts_with("req-")); + assert_eq!(ctx.request_id, ctx.x_amz_request_id); + assert!(ctx.trace_id.is_none()); + assert!(ctx.span_id.is_none()); + } + + #[test] + fn test_extract_request_id_from_x_request_id() { + let mut headers = HeaderMap::new(); + headers.insert("x-request-id", "test-uuid-123".parse().unwrap()); + let id = extract_request_id_from_headers(&headers); + assert_eq!(id, "test-uuid-123"); + } + + #[test] + fn test_extract_request_id_fallback_to_amz() { + let mut headers = HeaderMap::new(); + headers.insert("x-amz-request-id", "amz-uuid-456".parse().unwrap()); + let id = extract_request_id_from_headers(&headers); + assert_eq!(id, "amz-uuid-456"); + } + + #[test] + fn test_extract_request_id_priority() { + let mut headers = HeaderMap::new(); + headers.insert("x-request-id", "x-req-789".parse().unwrap()); + headers.insert("x-amz-request-id", "amz-req-000".parse().unwrap()); + let id = extract_request_id_from_headers(&headers); + assert_eq!(id, "x-req-789"); + } + + #[test] + fn test_extract_request_id_no_headers() { + let headers = HeaderMap::new(); + let id = extract_request_id_from_headers(&headers); + assert_eq!(id, "unknown"); + } +} diff --git a/rustfs/src/storage/timeout_wrapper.rs b/rustfs/src/storage/timeout_wrapper.rs index 47d711854..017239096 100644 --- a/rustfs/src/storage/timeout_wrapper.rs +++ b/rustfs/src/storage/timeout_wrapper.rs @@ -234,12 +234,15 @@ pub struct RequestTimeoutWrapper { impl RequestTimeoutWrapper { /// Create a new timeout wrapper with the given configuration. + /// + /// Note: This uses a sentinel request_id. Prefer `with_request_id()` to pass + /// the canonical request-id from `RequestContext`. pub fn new(config: TimeoutConfig) -> Self { Self { config, start_time: Instant::now(), cancel_token: CancellationToken::new(), - request_id: format!("req-{}", &uuid::Uuid::new_v4().to_string()[..8]), + request_id: "no-request-id".to_string(), } } @@ -253,17 +256,17 @@ impl RequestTimeoutWrapper { } } - /// Create a new timeout wrapper with operation size for dynamic timeout calculation + /// Create a new timeout wrapper with operation size for dynamic timeout calculation. + /// + /// Note: This uses a sentinel request_id. Prefer `with_request_id()` to pass + /// the canonical request-id from `RequestContext`. pub fn with_operation_size(config: TimeoutConfig, operation_size: Option) -> Self { - // Store operation size in config for later use - // Note: Currently we don't store the size in the wrapper itself, - // but the config can be used to calculate appropriate timeout - let _ = operation_size; // Suppress unused warning for now + let _ = operation_size; Self { config, start_time: Instant::now(), cancel_token: CancellationToken::new(), - request_id: format!("req-{}", &uuid::Uuid::new_v4().to_string()[..8]), + request_id: "no-request-id".to_string(), } } From 77229fe4268621d851f680ea469e8e802621a56c Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E5=AE=89=E6=AD=A3=E8=B6=85?= Date: Sat, 4 Apr 2026 14:41:35 +0800 Subject: [PATCH 16/22] test(admin): cover audit target validation gaps (#2390) --- rustfs/src/admin/handlers/audit.rs | 81 ++++++++++++++++++++- rustfs/src/admin/route_registration_test.rs | 3 + 2 files changed, 82 insertions(+), 2 deletions(-) diff --git a/rustfs/src/admin/handlers/audit.rs b/rustfs/src/admin/handlers/audit.rs index c8b91fed8..1325afa2d 100644 --- a/rustfs/src/admin/handlers/audit.rs +++ b/rustfs/src/admin/handlers/audit.rs @@ -634,9 +634,10 @@ impl Operation for RemoveAuditTarget { #[cfg(test)] mod tests { use super::*; + use matchit::Router; use rustfs_ecstore::config::{KV, KVS}; use std::collections::{HashMap, HashSet}; - use temp_env::{with_var, with_vars}; + use temp_env::{with_var, with_vars, with_vars_unset}; fn enabled_kvs(value: &str) -> KVS { KVS(vec![KV { @@ -646,6 +647,31 @@ mod tests { }]) } + fn with_audit_webhook_target_env_cleared(target_name: &str, f: F) + where + F: FnOnce(), + { + let target_name = target_name.to_ascii_uppercase(); + let mut env_keys = vec![format!( + "{ENV_PREFIX}{}{DEFAULT_DELIMITER}{}{DEFAULT_DELIMITER}{target_name}", + AUDIT_WEBHOOK_SUB_SYS.to_ascii_uppercase(), + ENABLE_KEY.to_ascii_uppercase(), + )]; + + for key in AUDIT_WEBHOOK_KEYS { + let env_key = format!( + "{ENV_PREFIX}{}{DEFAULT_DELIMITER}{}{DEFAULT_DELIMITER}{target_name}", + AUDIT_WEBHOOK_SUB_SYS.to_ascii_uppercase(), + key.to_ascii_uppercase(), + ); + if !env_keys.contains(&env_key) { + env_keys.push(env_key); + } + } + + with_vars_unset(env_keys, f); + } + #[test] fn merge_audit_endpoints_marks_config_env_and_mixed_sources() { let config = Config(HashMap::from([( @@ -782,6 +808,47 @@ mod tests { assert!(err.to_string().contains("duplicate key")); } + #[test] + fn collect_validated_key_values_rejects_unsupported_key() { + let allowed_keys: HashSet<&str> = AUDIT_WEBHOOK_KEYS.iter().copied().collect(); + let key_values = vec![KeyValue { + key: "not_a_real_key".to_string(), + value: "/tmp/rustfs-audit".to_string(), + }]; + + let err = collect_validated_key_values(&key_values, &allowed_keys, AUDIT_WEBHOOK_SUB_SYS).unwrap_err(); + assert!(err.to_string().contains("not allowed for audit target type")); + } + + #[test] + fn extract_target_params_rejects_missing_or_unsupported_values() { + let mut root_router = Router::new(); + root_router.insert("/", ()).expect("route should insert"); + let missing_type_params = root_router.at("/").expect("route should match"); + let missing_type = extract_target_params(&missing_type_params.params).unwrap_err(); + assert!(missing_type.to_string().contains("missing required parameter: 'target_type'")); + + let mut full_router = Router::new(); + full_router + .insert("/v3/audit/target/{target_type}/{target_name}", ()) + .expect("route should insert"); + let unsupported_type_params = full_router + .at("/v3/audit/target/audit_kafka/primary") + .expect("route should match"); + let unsupported_type = extract_target_params(&unsupported_type_params.params).unwrap_err(); + assert!(unsupported_type.to_string().contains("unsupported audit target type")); + + let mut partial_router = Router::new(); + partial_router + .insert("/v3/audit/target/{target_type}", ()) + .expect("route should insert"); + let missing_name_params = partial_router + .at("/v3/audit/target/audit_webhook") + .expect("route should match"); + let missing_name = extract_target_params(&missing_name_params.params).unwrap_err(); + assert!(missing_name.to_string().contains("missing required parameter: 'target_name'")); + } + #[test] fn merge_audit_endpoints_marks_mixed_with_case_insensitive_instance_id() { let config = Config(HashMap::from([( @@ -813,6 +880,16 @@ mod tests { HashMap::from([("PrimaryCase".to_string(), enabled_kvs("on"))]), )])); - assert!(audit_target_mutation_block_reason(&config, AUDIT_WEBHOOK_SUB_SYS, "primarycase").is_none()); + with_audit_webhook_target_env_cleared("primarycase", || { + assert!(audit_target_mutation_block_reason(&config, AUDIT_WEBHOOK_SUB_SYS, "primarycase").is_none()); + }); + } + + #[test] + fn audit_target_mutation_block_reason_allows_runtime_only_target() { + with_audit_webhook_target_env_cleared("primary", || { + let config = Config(HashMap::new()); + assert!(audit_target_mutation_block_reason(&config, AUDIT_WEBHOOK_SUB_SYS, "primary").is_none()); + }); } } diff --git a/rustfs/src/admin/route_registration_test.rs b/rustfs/src/admin/route_registration_test.rs index d490d4d9d..e25238556 100644 --- a/rustfs/src/admin/route_registration_test.rs +++ b/rustfs/src/admin/route_registration_test.rs @@ -182,6 +182,9 @@ fn test_admin_alias_paths_match_existing_admin_routes() { (Method::PUT, compat_admin_alias_path("/v3/set-policy")), (Method::PUT, compat_admin_alias_path("/v3/set-bucket-quota")), (Method::GET, compat_admin_alias_path("/v3/get-bucket-quota")), + (Method::GET, compat_admin_alias_path("/v3/audit/target/list")), + (Method::PUT, compat_admin_alias_path("/v3/audit/target/audit_webhook/test-audit")), + (Method::DELETE, compat_admin_alias_path("/v3/audit/target/audit_webhook/test-audit/reset")), (Method::POST, compat_admin_alias_path("/v3/heal/")), (Method::POST, compat_admin_alias_path("/v3/heal/test-bucket")), (Method::POST, compat_admin_alias_path("/v3/heal/test-bucket/prefix")), From 0e5fc4bec129b6390642e9c114404fd006b4fd9b Mon Sep 17 00:00:00 2001 From: ankohuu Date: Sun, 5 Apr 2026 21:39:10 +0800 Subject: [PATCH 17/22] fix(metrics): use Prometheus-compatible metric names (#2312) (#2317) Signed-off-by: Shunchao Hu --- .../grafana/dashboards/rustfs.json | 26 ++++++------- crates/metrics/src/collectors/system_cpu.rs | 2 +- .../metrics/src/collectors/system_memory.rs | 2 +- .../src/metrics_type/entry/descriptor.rs | 38 +++++++++++++++++-- crates/metrics/src/metrics_type/entry/mod.rs | 4 +- .../src/metrics_type/entry/subsystem.rs | 4 +- 6 files changed, 54 insertions(+), 22 deletions(-) diff --git a/.docker/observability/grafana/dashboards/rustfs.json b/.docker/observability/grafana/dashboards/rustfs.json index 2f28ac4d0..d71eaf3e6 100644 --- a/.docker/observability/grafana/dashboards/rustfs.json +++ b/.docker/observability/grafana/dashboards/rustfs.json @@ -91,7 +91,7 @@ "uid": "${datasource}" }, "editorMode": "code", - "expr": "gauge_rustfs_process_uptime_seconds{job=~\"$job\"}", + "expr": "rustfs_system_process_uptime_seconds{job=~\"$job\"}", "legendFormat": "__auto", "range": true, "refId": "A" @@ -223,7 +223,7 @@ "uid": "${datasource}" }, "editorMode": "code", - "expr": "sum(gauge_rustfs_cluster_buckets_total{job=~\"$job\"})", + "expr": "sum(rustfs_cluster_buckets_total{job=~\"$job\"})", "legendFormat": "__auto", "range": true, "refId": "A" @@ -289,7 +289,7 @@ "uid": "${datasource}" }, "editorMode": "code", - "expr": "sum(gauge_rustfs_cluster_objects_total{job=~\"$job\"})", + "expr": "sum(rustfs_cluster_objects_total{job=~\"$job\"})", "legendFormat": "__auto", "range": true, "refId": "A" @@ -427,7 +427,7 @@ "uid": "${datasource}" }, "editorMode": "code", - "expr": "sum(gauge_rustfs_cluster_capacity_used_bytes{job=~\"$job\"})", + "expr": "sum(rustfs_cluster_capacity_used_bytes{job=~\"$job\"})", "legendFormat": "Used", "range": true, "refId": "A" @@ -438,7 +438,7 @@ "uid": "${datasource}" }, "editorMode": "code", - "expr": "sum(gauge_rustfs_cluster_capacity_raw_total_bytes{job=~\"$job\"})", + "expr": "sum(rustfs_cluster_capacity_raw_total_bytes{job=~\"$job\"})", "hide": false, "legendFormat": "Total", "range": true, @@ -450,7 +450,7 @@ "uid": "${datasource}" }, "editorMode": "code", - "expr": "sum(gauge_rustfs_cluster_capacity_used_bytes{job=~\"$job\"}) / sum(gauge_rustfs_cluster_capacity_raw_total_bytes{job=~\"$job\"})", + "expr": "sum(rustfs_cluster_capacity_used_bytes{job=~\"$job\"}) / sum(rustfs_cluster_capacity_raw_total_bytes{job=~\"$job\"})", "hide": false, "instant": false, "legendFormat": "Percent", @@ -1971,7 +1971,7 @@ "uid": "${datasource}" }, "editorMode": "code", - "expr": "sum by (drive) (gauge_rustfs_node_disk_used_bytes{job=~\"$job\", drive=~\"$drive\"})", + "expr": "sum by (drive) (rustfs_node_disk_used_bytes{job=~\"$job\", drive=~\"$drive\"})", "legendFormat": "{{drive}} (bytes)", "range": true, "refId": "A" @@ -1982,7 +1982,7 @@ "uid": "${datasource}" }, "editorMode": "code", - "expr": "sum by (drive) (gauge_rustfs_node_disk_used_bytes{job=~\"$job\", drive=~\"$drive\"}) / sum by (drive)(gauge_rustfs_node_disk_total_bytes{job=~\"$job\", drive=~\"$drive\"})", + "expr": "sum by (drive) (rustfs_node_disk_used_bytes{job=~\"$job\", drive=~\"$drive\"}) / sum by (drive)(rustfs_node_disk_total_bytes{job=~\"$job\", drive=~\"$drive\"})", "hide": false, "instant": false, "legendFormat": "{{drive}} (percent)", @@ -2473,7 +2473,7 @@ "uid": "${datasource}" }, "editorMode": "code", - "expr": "sum by (job) (gauge_rustfs_process_cpu_usage{job=~\"$job\"})", + "expr": "sum by (job) (rustfs_system_process_cpu_usage{job=~\"$job\"})", "legendFormat": "{{job}}", "range": true, "refId": "A" @@ -2572,7 +2572,7 @@ "uid": "${datasource}" }, "editorMode": "code", - "expr": "sum by (job) (gauge_rustfs_process_resident_memory_bytes{job=~\"$job\"})", + "expr": "sum by (job) (rustfs_system_process_resident_memory_bytes{job=~\"$job\"})", "legendFormat": "{{job}}", "range": true, "refId": "B" @@ -2670,7 +2670,7 @@ "uid": "${datasource}" }, "editorMode": "code", - "expr": "sum by (job) (rate(gauge_rustfs_process_network_io{job=~\"$job\", direction=\"received\"}[5m]))", + "expr": "sum by (job) (rate(rustfs_system_process_network_io{job=~\"$job\", direction=\"received\"}[5m]))", "legendFormat": "RX - {{job}}", "range": true, "refId": "C" @@ -2681,7 +2681,7 @@ "uid": "${datasource}" }, "editorMode": "code", - "expr": "sum by (job) (rate(gauge_rustfs_process_network_io{job=~\"$job\", direction=\"transmitted\"}[5m]))", + "expr": "sum by (job) (rate(rustfs_system_process_network_io{job=~\"$job\", direction=\"transmitted\"}[5m]))", "legendFormat": "TX - {{job}}", "range": true, "refId": "D" @@ -4922,4 +4922,4 @@ "title": "RustFS", "uid": "rustfs-s3", "version": 12 -} \ No newline at end of file +} diff --git a/crates/metrics/src/collectors/system_cpu.rs b/crates/metrics/src/collectors/system_cpu.rs index 71de1f224..07ea9cba9 100644 --- a/crates/metrics/src/collectors/system_cpu.rs +++ b/crates/metrics/src/collectors/system_cpu.rs @@ -123,7 +123,7 @@ mod tests { assert_eq!(metrics.len(), 8); // Verify that metric names are properly generated from descriptors - assert!(metrics.iter().all(|m| m.name.starts_with("gauge.rustfs_system_cpu_"))); + assert!(metrics.iter().all(|m| m.name.starts_with("rustfs_system_cpu_"))); } #[test] diff --git a/crates/metrics/src/collectors/system_memory.rs b/crates/metrics/src/collectors/system_memory.rs index ff8855ffe..54697cca5 100644 --- a/crates/metrics/src/collectors/system_memory.rs +++ b/crates/metrics/src/collectors/system_memory.rs @@ -122,7 +122,7 @@ mod tests { report_metrics(&metrics); assert_eq!(metrics.len(), 8); - assert!(metrics.iter().all(|m| m.name.starts_with("gauge.rustfs_system_memory_"))); + assert!(metrics.iter().all(|m| m.name.starts_with("rustfs_system_memory_"))); } #[test] diff --git a/crates/metrics/src/metrics_type/entry/descriptor.rs b/crates/metrics/src/metrics_type/entry/descriptor.rs index e8ddb699f..c4612e1f8 100644 --- a/crates/metrics/src/metrics_type/entry/descriptor.rs +++ b/crates/metrics/src/metrics_type/entry/descriptor.rs @@ -51,14 +51,13 @@ impl MetricDescriptor { } } - /// Get the full metric name, including the prefix and formatting path + /// Get the full metric name in Prometheus style: __ #[allow(dead_code)] pub fn get_full_metric_name(&self) -> String { - let prefix = self.metric_type.as_prom(); let namespace = self.namespace.as_str(); let formatted_subsystem = self.subsystem.as_str(); - format!("{}{}_{}_{}", prefix, namespace, formatted_subsystem, self.name.as_str()) + format!("{}_{}_{}", namespace, formatted_subsystem, self.name.as_str()) } /// check whether the label is in the label set @@ -79,3 +78,36 @@ impl MetricDescriptor { self.label_set.as_ref().unwrap() } } + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn full_metric_name_uses_prometheus_convention_without_type_prefix() { + let descriptor = MetricDescriptor::new( + MetricName::ApiRequestsTotal, + MetricType::Counter, + "test help".to_string(), + vec![], + MetricNamespace::RustFS, + MetricSubsystem::ApiRequests, + ); + + assert_eq!(descriptor.get_full_metric_name(), "rustfs_api_requests_total"); + } + + #[test] + fn full_metric_name_formats_custom_subsystems_without_type_prefix() { + let descriptor = MetricDescriptor::new( + MetricName::Custom("latency_seconds".to_string()), + MetricType::Histogram, + "test help".to_string(), + vec![], + MetricNamespace::RustFS, + MetricSubsystem::new("/custom/path-metrics"), + ); + + assert_eq!(descriptor.get_full_metric_name(), "rustfs_custom_path_metrics_latency_seconds"); + } +} diff --git a/crates/metrics/src/metrics_type/entry/mod.rs b/crates/metrics/src/metrics_type/entry/mod.rs index 9d7881e3b..87215d15c 100644 --- a/crates/metrics/src/metrics_type/entry/mod.rs +++ b/crates/metrics/src/metrics_type/entry/mod.rs @@ -110,7 +110,7 @@ mod tests { assert_eq!(histogram_md.subsystem, MetricSubsystem::ApiRequests); // Verify that the full metric name generated is formatted correctly - assert_eq!(histogram_md.get_full_metric_name(), "histogram.rustfs_api_requests_seconds_distribution"); + assert_eq!(histogram_md.get_full_metric_name(), "rustfs_api_requests_seconds_distribution"); // Tests use custom subsystems let custom_histogram_md = new_histogram_md( @@ -123,7 +123,7 @@ mod tests { // Verify the custom name and subsystem assert_eq!( custom_histogram_md.get_full_metric_name(), - "histogram.rustfs_custom_path_metrics_custom_latency_distribution" + "rustfs_custom_path_metrics_custom_latency_distribution" ); } } diff --git a/crates/metrics/src/metrics_type/entry/subsystem.rs b/crates/metrics/src/metrics_type/entry/subsystem.rs index e6f0b83c1..2567560b9 100644 --- a/crates/metrics/src/metrics_type/entry/subsystem.rs +++ b/crates/metrics/src/metrics_type/entry/subsystem.rs @@ -233,7 +233,7 @@ mod tests { MetricSubsystem::ApiRequests, ); - assert_eq!(md.get_full_metric_name(), "counter.rustfs_api_requests_total"); + assert_eq!(md.get_full_metric_name(), "rustfs_api_requests_total"); let custom_md = MetricDescriptor::new( MetricName::Custom("test_metric".to_string()), @@ -244,6 +244,6 @@ mod tests { MetricSubsystem::new("/custom/path-with-dash"), ); - assert_eq!(custom_md.get_full_metric_name(), "gauge.rustfs_custom_path_with_dash_test_metric"); + assert_eq!(custom_md.get_full_metric_name(), "rustfs_custom_path_with_dash_test_metric"); } } From 97d3fb15fbd0eaf813e55845de6c60d0a925b758 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E5=AE=89=E6=AD=A3=E8=B6=85?= Date: Mon, 6 Apr 2026 10:16:53 +0800 Subject: [PATCH 18/22] test(metrics): cover Prometheus descriptor names (#2405) --- crates/metrics/src/format.rs | 36 ++++++++++++++++++++++++++++++++++++ 1 file changed, 36 insertions(+) diff --git a/crates/metrics/src/format.rs b/crates/metrics/src/format.rs index 976cc3ce8..e87138314 100644 --- a/crates/metrics/src/format.rs +++ b/crates/metrics/src/format.rs @@ -182,3 +182,39 @@ impl PrometheusMetric { self } } + +#[cfg(test)] +mod tests { + use super::*; + use crate::{MetricDescriptor, MetricName, MetricNamespace, MetricSubsystem}; + + #[test] + fn from_descriptor_uses_prometheus_metric_names_for_all_types() { + let cases = [ + (MetricType::Counter, "rustfs_api_requests_total"), + (MetricType::Gauge, "rustfs_system_memory_used_bytes"), + (MetricType::Histogram, "rustfs_custom_path_latency_seconds"), + ]; + + for (metric_type, expected_name) in cases { + let subsystem = match metric_type { + MetricType::Counter => MetricSubsystem::ApiRequests, + MetricType::Gauge => MetricSubsystem::SystemMemory, + MetricType::Histogram => MetricSubsystem::new("/custom/path"), + }; + let name = match metric_type { + MetricType::Counter => MetricName::ApiRequestsTotal, + MetricType::Gauge => MetricName::Custom("used_bytes".to_string()), + MetricType::Histogram => MetricName::Custom("latency_seconds".to_string()), + }; + + let metric = PrometheusMetric::from_descriptor( + &MetricDescriptor::new(name, metric_type, "test help".to_string(), vec![], MetricNamespace::RustFS, subsystem), + 1.0, + ); + + assert_eq!(metric.name, expected_name); + assert_eq!(metric.metric_type, metric_type); + } + } +} From e69d1dc3e68975a6515cee6cf95355f792dae9ee Mon Sep 17 00:00:00 2001 From: houseme Date: Mon, 6 Apr 2026 17:55:38 +0800 Subject: [PATCH 19/22] build(deps): bump the dependencies group with 14 updates (#2407) Co-authored-by: heihutu --- Cargo.lock | 66 +++++++++++++++++++++++++++--------------------------- Cargo.toml | 28 +++++++++++------------ 2 files changed, 47 insertions(+), 47 deletions(-) diff --git a/Cargo.lock b/Cargo.lock index a74a30c29..7199861d6 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -234,9 +234,9 @@ dependencies = [ [[package]] name = "arc-swap" -version = "1.9.0" +version = "1.9.1" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "a07d1f37ff60921c83bdfc7407723bdefe89b44b98a9b772f225c8f9d67141a6" +checksum = "6a3a1fd6f75306b68087b831f025c712524bcb19aad54e557b1129cfa0a2b207" dependencies = [ "rustversion", ] @@ -727,9 +727,9 @@ dependencies = [ [[package]] name = "aws-sdk-s3" -version = "1.127.0" +version = "1.128.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "151783f64e0dcddeb4965d08e36c276b4400a46caa88805a2e36d497deaf031a" +checksum = "99304b64672e0d81a3c100a589b93d9ef5e9c0ce12e21c848fd39e50f493c2a1" dependencies = [ "aws-credential-types", "aws-runtime", @@ -3136,7 +3136,7 @@ dependencies = [ "libc", "option-ext", "redox_users", - "windows-sys 0.59.0", + "windows-sys 0.61.2", ] [[package]] @@ -3412,7 +3412,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "39cab71617ae0d63f51a36d69f866391735b51691dbda63cf6f96d042b63efeb" dependencies = [ "libc", - "windows-sys 0.59.0", + "windows-sys 0.61.2", ] [[package]] @@ -4666,7 +4666,7 @@ checksum = "3640c1c38b8e4e43584d8df18be5fc6b0aa314ce6ebf51b53313d4306cca8e46" dependencies = [ "hermit-abi", "libc", - "windows-sys 0.59.0", + "windows-sys 0.61.2", ] [[package]] @@ -4743,7 +4743,7 @@ dependencies = [ "portable-atomic", "portable-atomic-util", "serde_core", - "windows-sys 0.59.0", + "windows-sys 0.61.2", ] [[package]] @@ -4934,9 +4934,9 @@ checksum = "2c4a545a15244c7d945065b5d392b2d2d7f21526fba56ce51467b06ed445e8f7" [[package]] name = "libc" -version = "0.2.183" +version = "0.2.184" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "b5b646652bf6661599e1da8901b3b9522896f01e736bad5f723fe7a3a27f899d" +checksum = "48f5d2a454e16a5ea0f4ced81bd44e4cfc7bd3a507b61887c99fd3538b28e4af" [[package]] name = "libflate" @@ -5090,9 +5090,9 @@ checksum = "6373607a59f0be73a39b6fe456b8192fcc3585f602af20751600e974dd455e77" [[package]] name = "local-ip-address" -version = "0.6.10" +version = "0.6.11" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "79ef8c257c92ade496781a32a581d43e3d512cf8ce714ecf04ea80f93ed0ff4a" +checksum = "d4a59a0cb1c7f84471ad5cd38d768c2a29390d17f1ff2827cdf49bc53e8ac70b" dependencies = [ "libc", "neli", @@ -5457,9 +5457,9 @@ dependencies = [ [[package]] name = "mio" -version = "1.1.1" +version = "1.2.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "a69bcab0ad47271a0234d9422b131806bf3968021e5dc9328caf2d4cd58557fc" +checksum = "50b7e5b27aa02a74bac8c3f23f448f8d87ff11f92d3aac1a6ed369ee08cc56c1" dependencies = [ "libc", "wasi", @@ -5619,7 +5619,7 @@ version = "0.50.3" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "7957b9740744892f114936ab4a57b3f487491bbeafaf8083688b16841a4240e5" dependencies = [ - "windows-sys 0.59.0", + "windows-sys 0.61.2", ] [[package]] @@ -5755,9 +5755,9 @@ checksum = "a3c00a0c9600379bd32f8972de90676a7672cba3bf4886986bc05902afc1e093" [[package]] name = "nvml-wrapper" -version = "0.12.0" +version = "0.12.1" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "7d9e6eebc1fe424d24c864e40092072618169bd0130f103919aaf615f153e4d0" +checksum = "f049ae562349fefb8e837eb15443da1e7c6dcbd8a11f52a228f92220c2e5c85e" dependencies = [ "bitflags 2.11.0", "libloading", @@ -5769,9 +5769,9 @@ dependencies = [ [[package]] name = "nvml-wrapper-sys" -version = "0.9.0" +version = "0.9.1" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "dd23dbe2eb8d8335d2bce0299e0a07d6a63c089243d626ca75b770a962ff49e6" +checksum = "6b4d594420fcda43b1c2c4bd44d48974aa3c7a9ab2cbf10dc18e35265767bf0b" dependencies = [ "libloading", ] @@ -6252,9 +6252,9 @@ dependencies = [ [[package]] name = "pbkdf2" -version = "0.13.0-rc.9" +version = "0.13.0-rc.10" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "c8dfa4e14084d963d35bfb4cdb38712cde78dcf83054c0e8b9b8e899150f374e" +checksum = "1f24f3eb2f4471b1730d59e4b730b747939960a8c7eb0c33c5a9076f2d3dddea" dependencies = [ "digest 0.11.2", "hmac 0.13.0", @@ -7782,7 +7782,7 @@ dependencies = [ "cfg-if", "chacha20poly1305", "jsonwebtoken", - "pbkdf2 0.13.0-rc.9", + "pbkdf2 0.13.0-rc.10", "rand 0.10.0", "serde_json", "sha2 0.11.0", @@ -8570,7 +8570,7 @@ dependencies = [ "errno", "libc", "linux-raw-sys 0.12.1", - "windows-sys 0.59.0", + "windows-sys 0.61.2", ] [[package]] @@ -8638,7 +8638,7 @@ dependencies = [ "security-framework", "security-framework-sys", "webpki-root-certs", - "windows-sys 0.59.0", + "windows-sys 0.61.2", ] [[package]] @@ -8685,7 +8685,7 @@ checksum = "9774ba4a74de5f7b1c1451ed6cd5285a32eddb5cccb8cc655a4e50009e06477f" [[package]] name = "s3s" version = "0.14.0-dev" -source = "git+https://github.com/rustfs/s3s?rev=738f85792c92781bd8af862a074d7379d9fbfabc#738f85792c92781bd8af862a074d7379d9fbfabc" +source = "git+https://github.com/rustfs/s3s?rev=79bb9abe6b353025ef4572a25db8eb64f1409faf#79bb9abe6b353025ef4572a25db8eb64f1409faf" dependencies = [ "arc-swap", "arrayvec", @@ -9643,7 +9643,7 @@ dependencies = [ "getrandom 0.4.2", "once_cell", "rustix 1.1.4", - "windows-sys 0.59.0", + "windows-sys 0.61.2", ] [[package]] @@ -9850,9 +9850,9 @@ checksum = "1f3ccbac311fea05f86f61904b462b55fb3df8837a366dfc601a0161d0532f20" [[package]] name = "tokio" -version = "1.50.0" +version = "1.51.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "27ad5e34374e03cfffefc301becb44e9dc3c17584f414349ebe29ed26661822d" +checksum = "2bd1c4c0fc4a7ab90fc15ef6daaa3ec3b893f004f915f2392557ed23237820cd" dependencies = [ "bytes", "libc", @@ -9867,9 +9867,9 @@ dependencies = [ [[package]] name = "tokio-macros" -version = "2.6.1" +version = "2.7.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "5c55a2eff8b69ce66c84f85e1da1c233edc36ceb85a2058d11b0d6a3c7e7569c" +checksum = "385a6cb71ab9ab790c5fe8d67f1645e6c450a7ce006a33de03daa956cf70a496" dependencies = [ "proc-macro2", "quote", @@ -10602,7 +10602,7 @@ version = "0.1.11" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "c2a7b1c03c876122aa43f3020e6c3c3ee5c05081c9a00739faf7503aeba10d22" dependencies = [ - "windows-sys 0.59.0", + "windows-sys 0.61.2", ] [[package]] @@ -11193,9 +11193,9 @@ dependencies = [ [[package]] name = "zip" -version = "8.4.0" +version = "8.5.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "7756d0206d058333667493c4014f545f4b9603c4330ccd6d9b3f86dcab59f7d9" +checksum = "2726508a48f38dceb22b35ecbbd2430efe34ff05c62bd3285f965d7911b33464" dependencies = [ "aes 0.8.4", "bzip2", diff --git a/Cargo.toml b/Cargo.toml index 143f3e762..1bc8933b2 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -123,7 +123,7 @@ futures = "0.3.32" futures-core = "0.3.32" futures-util = "0.3.32" pollster = "0.4.0" -hyper = { version = "1.8.1", features = ["http2", "http1", "server"] } +hyper = { version = "1.9.0", features = ["http2", "http1", "server"] } hyper-rustls = { version = "0.27.7", default-features = false, features = ["native-tokio", "http1", "tls12", "logging", "http2", "aws-lc-rs", "webpki-roots"] } hyper-util = { version = "0.1.20", features = ["tokio", "server-auto", "server-graceful", "tracing"] } http = "1.4.0" @@ -131,7 +131,7 @@ http-body = "1.0.1" http-body-util = "0.1.3" reqwest = { version = "0.13.2", default-features = false, features = ["rustls", "charset", "http2", "system-proxy", "stream", "json", "blocking", "query", "form"] } socket2 = { version = "0.6.3", features = ["all"] } -tokio = { version = "1.50.0", features = ["fs", "rt-multi-thread"] } +tokio = { version = "1.51.0", features = ["fs", "rt-multi-thread"] } tokio-rustls = { version = "0.26.4", default-features = false, features = ["logging", "tls12", "aws-lc-rs"] } tokio-stream = { version = "0.1.18" } tokio-test = "0.4.5" @@ -164,15 +164,15 @@ argon2 = { version = "0.6.0-rc.8" } blake2 = "0.11.0-rc.5" chacha20poly1305 = { version = "0.11.0-rc.3" } crc-fast = "1.9.0" -hmac = { version = "0.13.0-rc.5" } +hmac = { version = "0.13.0" } jsonwebtoken = { version = "10.3.0", features = ["aws_lc_rs"] } openidconnect = { version = "4.0", default-features = false } -pbkdf2 = "0.13.0-rc.9" +pbkdf2 = "0.13.0-rc.10" rsa = { version = "0.10.0-rc.17" } rustls = { version = "0.23.37", default-features = false, features = ["aws-lc-rs", "logging", "tls12", "prefer-post-quantum", "std"] } rustls-pki-types = "1.14.0" -sha1 = "0.11.0-rc.5" -sha2 = "0.11.0-rc.5" +sha1 = "0.11.0" +sha2 = "0.11.0" subtle = "2.6" zeroize = { version = "1.8.2", features = ["derive"] } @@ -184,13 +184,13 @@ time = { version = "0.3.47", features = ["std", "parsing", "formatting", "macros # Utilities and Tools anyhow = "1.0.102" -arc-swap = "1.9.0" +arc-swap = "1.9.1" astral-tokio-tar = "0.6.0" atoi = "2.0.0" atomic_enum = "0.3.0" aws-config = { version = "1.8.15" } aws-credential-types = { version = "1.2.14" } -aws-sdk-s3 = { version = "1.127.0", default-features = false, features = ["sigv4a", "default-https-client", "rt-tokio"] } +aws-sdk-s3 = { version = "1.128.0", default-features = false, features = ["sigv4a", "default-https-client", "rt-tokio"] } aws-smithy-http-client = { version = "1.1.12", default-features = false, features = ["default-client", "rustls-aws-lc"] } aws-smithy-types = { version = "1.4.7" } backtrace = "0.3.76" @@ -220,19 +220,19 @@ hex-simd = "0.8.0" highway = { version = "1.3.0" } ipnetwork = { version = "0.21.1", features = ["serde"] } lazy_static = "1.5.0" -libc = "0.2.183" +libc = "0.2.184" libsystemd = "0.7.2" -local-ip-address = "0.6.10" +local-ip-address = "0.6.11" memmap2 = "0.9.10" lz4 = "1.28.1" matchit = "0.9.1" -md-5 = "0.11.0-rc.5" +md-5 = "0.11.0" md5 = "0.8.0" mime_guess = "2.0.5" moka = { version = "0.12.15", features = ["future"] } netif = "0.1.6" num_cpus = { version = "1.17.0" } -nvml-wrapper = "0.12.0" +nvml-wrapper = "0.12.1" object_store = "0.13.2" parking_lot = "0.12.5" path-absolutize = "3.1.1" @@ -250,7 +250,7 @@ rumqttc = { version = "0.25.1" } rustix = { version = "1.1.4", features = ["fs"] } rust-embed = { version = "8.11.0" } rustc-hash = { version = "2.1.2" } -s3s = { git = "https://github.com/rustfs/s3s", rev = "738f85792c92781bd8af862a074d7379d9fbfabc", features = ["minio"] } +s3s = { git = "https://github.com/rustfs/s3s", rev = "79bb9abe6b353025ef4572a25db8eb64f1409faf", features = ["minio"] } serial_test = "3.4.0" shadow-rs = { version = "1.7.1", default-features = false } siphasher = "1.0.2" @@ -279,7 +279,7 @@ walkdir = "2.5.0" wildmatch = { version = "2.6.1", features = ["serde"] } windows = { version = "0.62.2" } xxhash-rust = { version = "0.8.15", features = ["xxh64", "xxh3"] } -zip = "8.4.0" +zip = "8.5.0" zstd = "0.13.3" # Observability and Metrics From dd68a419e3539ec8b6779ddc8708e0f03fbd7984 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E5=AE=89=E6=AD=A3=E8=B6=85?= Date: Mon, 6 Apr 2026 20:35:37 +0800 Subject: [PATCH 20/22] test(server): cover request context layer propagation (#2398) Co-authored-by: loverustfs --- rustfs/src/server/layer.rs | 62 ++++++++++++++++++++++++++++++++++++++ 1 file changed, 62 insertions(+) diff --git a/rustfs/src/server/layer.rs b/rustfs/src/server/layer.rs index 13a198141..09f6cf10a 100644 --- a/rustfs/src/server/layer.rs +++ b/rustfs/src/server/layer.rs @@ -678,11 +678,30 @@ where #[cfg(test)] mod tests { use super::*; + use futures::future::{Ready, ready}; use http::Request; use http_body_util::BodyExt; use http_body_util::Full; + use std::convert::Infallible; use temp_env::with_var; + #[derive(Clone, Debug)] + struct CaptureService; + + impl Service> for CaptureService { + type Response = Request; + type Error = Infallible; + type Future = Ready>; + + fn poll_ready(&mut self, _cx: &mut Context<'_>) -> Poll> { + Poll::Ready(Ok(())) + } + + fn call(&mut self, req: Request) -> Self::Future { + ready(Ok(req)) + } + } + #[test] fn admin_chunked_put_without_content_length_is_normalized() { let request = Request::builder() @@ -816,6 +835,49 @@ mod tests { }); } + #[test] + fn request_context_layer_populates_context_and_s3_request_id_from_x_request_id() { + let mut service = RequestContextLayer.layer(CaptureService); + let request = Request::builder() + .uri("/bucket/object") + .header("x-request-id", "req-123") + .body(()) + .expect("request"); + + let request = service.call(request).into_inner().expect("service call should succeed"); + let context = request + .extensions() + .get::() + .expect("request context should be present"); + + assert_eq!(context.request_id, "req-123"); + assert_eq!(context.x_amz_request_id, "req-123"); + assert!(context.trace_id.is_none()); + assert!(context.span_id.is_none()); + assert_eq!(request.headers().get(AMZ_REQUEST_ID).unwrap(), "req-123"); + } + + #[test] + fn request_context_layer_preserves_upstream_s3_request_id() { + let mut service = RequestContextLayer.layer(CaptureService); + let request = Request::builder() + .uri("/bucket/object") + .header("x-request-id", "req-123") + .header(AMZ_REQUEST_ID, "amz-456") + .body(()) + .expect("request"); + + let request = service.call(request).into_inner().expect("service call should succeed"); + let context = request + .extensions() + .get::() + .expect("request context should be present"); + + assert_eq!(context.request_id, "req-123"); + assert_eq!(context.x_amz_request_id, "amz-456"); + assert_eq!(request.headers().get(AMZ_REQUEST_ID).unwrap(), "amz-456"); + } + #[tokio::test] async fn test_resolve_s3_options_cors_headers_no_headers_without_match() { let mut req_headers = HeaderMap::new(); From 8d27170ce44ed8a6320e91f2c8d98073340f3478 Mon Sep 17 00:00:00 2001 From: ankohuu Date: Mon, 6 Apr 2026 23:05:30 +0800 Subject: [PATCH 21/22] fix: limit pyroscope profiling to supported Unix targets (#2399) Signed-off-by: houseme Co-authored-by: loverustfs Co-authored-by: houseme Co-authored-by: Copilot <175728472+Copilot@users.noreply.github.com> --- crates/obs/Cargo.toml | 2 +- crates/obs/src/telemetry/guard.rs | 6 +++--- crates/obs/src/telemetry/local.rs | 4 ++-- crates/obs/src/telemetry/otel.rs | 13 +++++++------ 4 files changed, 13 insertions(+), 12 deletions(-) diff --git a/crates/obs/Cargo.toml b/crates/obs/Cargo.toml index 03aa44b81..7a00bc879 100644 --- a/crates/obs/Cargo.toml +++ b/crates/obs/Cargo.toml @@ -56,7 +56,7 @@ dial9-tokio-telemetry = { workspace = true } thiserror = { workspace = true } zstd = { workspace = true, features = ["zstdmt"] } -[target.'cfg(unix)'.dependencies] +[target.'cfg(any(target_os = "linux", target_os = "macos"))'.dependencies] pyroscope = { workspace = true, features = ["backend-pprof-rs"] } diff --git a/crates/obs/src/telemetry/guard.rs b/crates/obs/src/telemetry/guard.rs index f91d25ce7..26c868d50 100644 --- a/crates/obs/src/telemetry/guard.rs +++ b/crates/obs/src/telemetry/guard.rs @@ -41,7 +41,7 @@ pub struct OtelGuard { pub(crate) meter_provider: Option, /// Optional logger provider for OTLP log export. pub(crate) logger_provider: Option, - #[cfg(unix)] + #[cfg(any(target_os = "linux", target_os = "macos"))] pub(crate) profiling_agent: Option>, /// Handle to the background log-cleanup task; aborted on drop. pub(crate) cleanup_handle: Option>, @@ -58,7 +58,7 @@ impl std::fmt::Debug for OtelGuard { s.field("tracer_provider", &self.tracer_provider.is_some()) .field("meter_provider", &self.meter_provider.is_some()) .field("logger_provider", &self.logger_provider.is_some()); - #[cfg(unix)] + #[cfg(any(target_os = "linux", target_os = "macos"))] s.field("profiling_agent", &self.profiling_agent.is_some()); s.field("cleanup_handle", &self.cleanup_handle.is_some()) .field("tracing_guard", &self.tracing_guard.is_some()) @@ -91,7 +91,7 @@ impl Drop for OtelGuard { eprintln!("Logger shutdown error: {err:?}"); } - #[cfg(unix)] + #[cfg(any(target_os = "linux", target_os = "macos"))] if let Some(agent) = self.profiling_agent.take() { match agent.stop() { Err(err) => eprintln!("Profiling agent stop error: {err:?}"), diff --git a/crates/obs/src/telemetry/local.rs b/crates/obs/src/telemetry/local.rs index c7d8cedf6..b73d3ac23 100644 --- a/crates/obs/src/telemetry/local.rs +++ b/crates/obs/src/telemetry/local.rs @@ -155,7 +155,7 @@ fn init_stdout_only(_config: &OtelConfig, logger_level: &str, is_production: boo tracer_provider: None, meter_provider: None, logger_provider: None, - #[cfg(unix)] + #[cfg(any(target_os = "linux", target_os = "macos"))] profiling_agent: None, tracing_guard: Some(guard), stdout_guard: None, @@ -289,7 +289,7 @@ fn init_file_logging_internal( tracer_provider: None, meter_provider: None, logger_provider: None, - #[cfg(unix)] + #[cfg(any(target_os = "linux", target_os = "macos"))] profiling_agent: None, tracing_guard: Some(guard), stdout_guard, diff --git a/crates/obs/src/telemetry/otel.rs b/crates/obs/src/telemetry/otel.rs index 3d4320356..736dc0dcd 100644 --- a/crates/obs/src/telemetry/otel.rs +++ b/crates/obs/src/telemetry/otel.rs @@ -162,7 +162,7 @@ pub(super) fn init_observability_http( // ── Meter provider (HTTP) ───────────────────────────────────────────────── let meter_provider = build_meter_provider(&metric_ep, config, res.clone(), &service_name, use_stdout)?; - #[cfg(unix)] + #[cfg(any(target_os = "linux", target_os = "macos"))] let profiling_agent = init_profiler(config); // ── Logger Logic ────────────────────────────────────────────────────────── @@ -205,7 +205,7 @@ pub(super) fn init_observability_http( let file_logging_result = (|| -> Result<_, TelemetryError> { fs::create_dir_all(log_directory).map_err(|e| TelemetryError::Io(e.to_string()))?; - #[cfg(unix)] + #[cfg(any(target_os = "linux", target_os = "macos"))] crate::telemetry::local::ensure_dir_permissions(log_directory)?; let rotation_str = config @@ -312,7 +312,7 @@ pub(super) fn init_observability_http( tracer_provider, meter_provider, logger_provider, - #[cfg(unix)] + #[cfg(any(target_os = "linux", target_os = "macos"))] profiling_agent, tracing_guard, stdout_guard, @@ -462,9 +462,10 @@ fn build_logger_provider( /// Start the Pyroscope continuous profiling agent when profiling is enabled. /// -/// Returns `None` on non-Unix platforms, when the feature is disabled, or when -/// no usable profiling endpoint is configured. -#[cfg(unix)] +/// Returns `None` when profiling export is disabled, when no usable +/// profiling endpoint is configured, or when building or starting the agent +/// fails. +#[cfg(any(target_os = "linux", target_os = "macos"))] fn init_profiler(config: &OtelConfig) -> Option> { use pyroscope::backend::{BackendConfig, PprofConfig, pprof_backend}; use pyroscope::pyroscope::PyroscopeAgentBuilder; From 32bf8f5bf30c21323c155dbcb6b2fe69ee0d62aa Mon Sep 17 00:00:00 2001 From: houseme Date: Tue, 7 Apr 2026 08:33:46 +0800 Subject: [PATCH 22/22] feat(storage): add direct chunk GET fast path (#2351) Signed-off-by: houseme Co-authored-by: heihutu Co-authored-by: cxymds --- Cargo.lock | 467 +-- Cargo.toml | 4 +- _typos.toml | 1 + crates/config/src/constants/zero_copy.rs | 28 + crates/e2e_test/Cargo.toml | 8 + crates/e2e_test/src/bin/small_put_bench.rs | 441 +++ crates/e2e_test/src/checksum_upload_test.rs | 276 +- crates/e2e_test/src/kms/kms_local_test.rs | 96 + crates/e2e_test/src/lib.rs | 4 + crates/e2e_test/src/range_request_test.rs | 80 + crates/ecstore/Cargo.toml | 14 + crates/ecstore/README.md | 37 + .../ecstore/benches/bitrot_chunk_benchmark.rs | 106 + .../ecstore/benches/direct_chunk_benchmark.rs | 390 +++ .../benches/reconstructed_chunk_benchmark.rs | 393 +++ crates/ecstore/run_benchmarks.sh | 19 +- crates/ecstore/src/bitrot.rs | 790 ++++- crates/ecstore/src/bucket/migration.rs | 14 +- crates/ecstore/src/config/com.rs | 20 +- crates/ecstore/src/config/storageclass.rs | 51 +- crates/ecstore/src/data_movement.rs | 23 +- crates/ecstore/src/disk/disk_store.rs | 9 + crates/ecstore/src/disk/local.rs | 858 +++++- crates/ecstore/src/disk/mod.rs | 12 + crates/ecstore/src/erasure_coding/bitrot.rs | 68 + crates/ecstore/src/erasure_coding/decode.rs | 245 +- crates/ecstore/src/erasure_coding/encode.rs | 365 ++- crates/ecstore/src/erasure_coding/erasure.rs | 282 +- crates/ecstore/src/erasure_coding/heal.rs | 2 +- crates/ecstore/src/rpc/remote_disk.rs | 31 +- crates/ecstore/src/set_disk.rs | 222 +- crates/ecstore/src/set_disk/read.rs | 1009 ++++++- crates/ecstore/src/set_disk/write.rs | 132 + crates/ecstore/src/sets.rs | 33 +- crates/ecstore/src/store.rs | 31 +- crates/ecstore/src/store/multipart.rs | 2 +- crates/ecstore/src/store/object.rs | 31 +- crates/ecstore/src/store_api.rs | 1 + crates/ecstore/src/store_api/readers.rs | 227 +- crates/ecstore/src/store_api/traits.rs | 10 +- crates/ecstore/src/store_api/types.rs | 16 +- crates/ecstore/src/tier/tier.rs | 13 +- crates/filemeta/src/filemeta.rs | 6 +- crates/heal/src/heal/storage.rs | 2 +- crates/heal/tests/heal_integration_test.rs | 4 +- crates/io-core/Cargo.toml | 2 + crates/io-core/src/adapter.rs | 124 + crates/io-core/src/chunk.rs | 276 ++ crates/io-core/src/lib.rs | 4 + crates/io-core/src/pool.rs | 42 +- crates/io-metrics/src/lib.rs | 598 ++-- crates/io-metrics/src/metric_names.rs | 55 +- crates/object-io/Cargo.toml | 37 + crates/object-io/src/get.rs | 1703 +++++++++++ crates/object-io/src/lib.rs | 16 + crates/object-io/src/put.rs | 1564 ++++++++++ crates/protocols/src/swift/object.rs | 10 +- crates/rio/src/checksum.rs | 40 +- crates/rio/src/compress_reader.rs | 94 + crates/rio/src/encrypt_reader.rs | 182 +- crates/rio/src/etag_reader.rs | 52 +- crates/rio/src/hardlimit_reader.rs | 57 + crates/rio/src/hash_reader.rs | 251 +- crates/rio/src/http_reader.rs | 104 +- crates/rio/src/lib.rs | 140 +- .../tests/lifecycle_integration_test.rs | 14 +- rustfs/Cargo.toml | 1 + .../src/app/lifecycle_transition_api_test.rs | 8 +- rustfs/src/app/multipart_usecase.rs | 173 +- rustfs/src/app/object_usecase.rs | 2623 +---------------- rustfs/src/app/object_usecase/app_adapters.rs | 616 ++++ .../src/app/object_usecase/get_object_flow.rs | 280 ++ .../object_usecase/get_object_zero_copy.rs | 338 +++ .../app/object_usecase/put_object_extract.rs | 499 ++++ .../src/app/object_usecase/put_object_flow.rs | 868 ++++++ rustfs/src/app/object_usecase/types.rs | 52 + .../src/app/object_usecase/zero_copy_tests.rs | 1222 ++++++++ rustfs/src/error.rs | 14 + rustfs/src/server/http.rs | 2 +- .../src/storage/concurrency/object_cache.rs | 94 + rustfs/src/storage/ecfs.rs | 39 - scripts/bench-small-put-local.sh | 152 + scripts/bench-small-put-mc-local.sh | 300 ++ scripts/run.sh | 5 +- 84 files changed, 15932 insertions(+), 3592 deletions(-) create mode 100644 crates/e2e_test/src/bin/small_put_bench.rs create mode 100644 crates/e2e_test/src/range_request_test.rs create mode 100644 crates/ecstore/benches/bitrot_chunk_benchmark.rs create mode 100644 crates/ecstore/benches/direct_chunk_benchmark.rs create mode 100644 crates/ecstore/benches/reconstructed_chunk_benchmark.rs create mode 100644 crates/io-core/src/adapter.rs create mode 100644 crates/io-core/src/chunk.rs create mode 100644 crates/object-io/Cargo.toml create mode 100644 crates/object-io/src/get.rs create mode 100644 crates/object-io/src/lib.rs create mode 100644 crates/object-io/src/put.rs create mode 100644 rustfs/src/app/object_usecase/app_adapters.rs create mode 100644 rustfs/src/app/object_usecase/get_object_flow.rs create mode 100644 rustfs/src/app/object_usecase/get_object_zero_copy.rs create mode 100644 rustfs/src/app/object_usecase/put_object_extract.rs create mode 100644 rustfs/src/app/object_usecase/put_object_flow.rs create mode 100644 rustfs/src/app/object_usecase/types.rs create mode 100644 rustfs/src/app/object_usecase/zero_copy_tests.rs create mode 100755 scripts/bench-small-put-local.sh create mode 100755 scripts/bench-small-put-mc-local.sh diff --git a/Cargo.lock b/Cargo.lock index 7199861d6..456df56fb 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -410,7 +410,7 @@ dependencies = [ "arrow-schema", "chrono", "half", - "indexmap 2.13.0", + "indexmap 2.13.1", "itoa", "lexical-core", "memchr", @@ -687,9 +687,9 @@ dependencies = [ [[package]] name = "aws-lc-sys" -version = "0.39.0" +version = "0.39.1" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "1fa7e52a4c5c547c741610a2c6f123f3881e409b714cd27e6798ef020c514f0a" +checksum = "83a25cf98105baa966497416dbd42565ce3a8cf8dbfd59803ec9ad46f3126399" dependencies = [ "cc", "cmake", @@ -1227,16 +1227,16 @@ dependencies = [ [[package]] name = "blake3" -version = "1.8.3" +version = "1.8.4" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "2468ef7d57b3fb7e16b576e8377cdbde2320c60e1491e961d11da40fc4f02a2d" +checksum = "4d2d5991425dfd0785aed03aedcf0b321d61975c9b5b3689c774a2610ae0b51e" dependencies = [ "arrayref", "arrayvec", "cc", "cfg-if", "constant_time_eq", - "cpufeatures 0.2.17", + "cpufeatures 0.3.0", ] [[package]] @@ -1405,9 +1405,9 @@ checksum = "37b2a672a2cb129a2e41c10b1224bb368f9f37a2b16b612598138befd7b37eb5" [[package]] name = "cc" -version = "1.2.57" +version = "1.2.59" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "7a0dd1ca384932ff3641c8718a02769f1698e7563dc6974ffd03346116310423" +checksum = "b7a4d3ec6524d28a329fc53654bbadc9bdd7b0431f5d65f1a56ffb28a1ee5283" dependencies = [ "find-msvc-tools", "jobserver", @@ -1514,7 +1514,7 @@ version = "0.4.4" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "773f3b9af64447d2ce9850330c473515014aa235e6a783b02db81ff39e4a3dad" dependencies = [ - "crypto-common 0.1.7", + "crypto-common 0.1.6", "inout 0.1.4", ] @@ -1582,18 +1582,18 @@ dependencies = [ [[package]] name = "cmake" -version = "0.1.57" +version = "0.1.58" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "75443c44cd6b379beb8c5b45d85d0773baf31cce901fe7bb252f4eff3008ef7d" +checksum = "c0f78a02292a74a88ac736019ab962ece0bc380e3f977bf72e376c5d78ff0678" dependencies = [ "cc", ] [[package]] name = "cmov" -version = "0.5.2" +version = "0.5.3" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "de0758edba32d61d1fd9f4d69491b47604b91ee2f7e6b33de7e54ca4ebe55dc3" +checksum = "3f88a43d011fc4a6876cb7344703e297c71dda42494fee094d5f7c76bf13f746" [[package]] name = "colorchoice" @@ -1971,9 +1971,9 @@ dependencies = [ [[package]] name = "crypto-bigint" -version = "0.7.1" +version = "0.7.3" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "9fde2467e74147f492aebb834985186b2c74761927b8b9b3bd303bcb2e72199d" +checksum = "42a0d26b245348befa0c121944541476763dcc46ede886c88f9d12e1697d27c3" dependencies = [ "cpubits", "ctutils", @@ -1985,9 +1985,9 @@ dependencies = [ [[package]] name = "crypto-common" -version = "0.1.7" +version = "0.1.6" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "78c8292055d1c1df0cce5d180393dc8cce0abec0a7102adb6c7b1eef6016d60a" +checksum = "1bfb12502f3fc46cca1bb51ac28df9d618d813cdc3d2f25b9fe775a34af26bb3" dependencies = [ "generic-array", "typenum", @@ -2010,7 +2010,7 @@ version = "0.7.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "21f41f23de7d24cdbda7f0c4d9c0351f99a4ceb258ef30e5c1927af8987ffe5a" dependencies = [ - "crypto-bigint 0.7.1", + "crypto-bigint 0.7.3", "libm", "rand_core 0.10.0", ] @@ -2047,9 +2047,9 @@ dependencies = [ [[package]] name = "ctutils" -version = "0.4.0" +version = "0.4.2" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "1005a6d4446f5120ef475ad3d2af2b30c49c2c9c6904258e3bb30219bebed5e4" +checksum = "7d5515a3834141de9eafb9717ad39eea8247b5674e6066c404e8c4b365d2a29e" dependencies = [ "cmov", ] @@ -2325,7 +2325,7 @@ dependencies = [ "chrono", "half", "hashbrown 0.16.1", - "indexmap 2.13.0", + "indexmap 2.13.1", "itertools 0.14.0", "libc", "log", @@ -2529,7 +2529,7 @@ dependencies = [ "datafusion-functions-aggregate-common", "datafusion-functions-window-common", "datafusion-physical-expr-common", - "indexmap 2.13.0", + "indexmap 2.13.1", "itertools 0.14.0", "paste", "recursive", @@ -2545,7 +2545,7 @@ checksum = "ab05fdd00e05d5a6ee362882546d29d6d3df43a6c55355164a7fbee12d163bc9" dependencies = [ "arrow", "datafusion-common", - "indexmap 2.13.0", + "indexmap 2.13.1", "itertools 0.14.0", "paste", ] @@ -2709,7 +2709,7 @@ dependencies = [ "datafusion-expr", "datafusion-expr-common", "datafusion-physical-expr", - "indexmap 2.13.0", + "indexmap 2.13.1", "itertools 0.14.0", "log", "recursive", @@ -2732,7 +2732,7 @@ dependencies = [ "datafusion-physical-expr-common", "half", "hashbrown 0.16.1", - "indexmap 2.13.0", + "indexmap 2.13.1", "itertools 0.14.0", "parking_lot 0.12.5", "paste", @@ -2768,7 +2768,7 @@ dependencies = [ "datafusion-common", "datafusion-expr-common", "hashbrown 0.16.1", - "indexmap 2.13.0", + "indexmap 2.13.1", "itertools 0.14.0", "parking_lot 0.12.5", ] @@ -2815,7 +2815,7 @@ dependencies = [ "futures", "half", "hashbrown 0.16.1", - "indexmap 2.13.0", + "indexmap 2.13.1", "itertools 0.14.0", "log", "num-traits", @@ -2867,7 +2867,7 @@ dependencies = [ "datafusion-common", "datafusion-expr", "datafusion-functions-nested", - "indexmap 2.13.0", + "indexmap 2.13.1", "log", "recursive", "regex", @@ -2917,9 +2917,9 @@ dependencies = [ [[package]] name = "deflate64" -version = "0.1.11" +version = "0.1.12" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "807800ff3288b621186fe0a8f3392c4652068257302709c24efd918c3dffcdc2" +checksum = "ac6b926516df9c60bfa16e107b21086399f8285a44ca9711344b9e553c5146e2" [[package]] name = "der" @@ -3045,8 +3045,7 @@ dependencies = [ [[package]] name = "dial9-tokio-telemetry" version = "0.2.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "0fab5b5b736126e4a4a3ed06e15389ac199c2ac4f72395197addb305e6ba1759" +source = "git+https://github.com/dial9-rs/dial9-tokio-telemetry.git?rev=60502082601b647c4a51962595721f631b7bbce1#60502082601b647c4a51962595721f631b7bbce1" dependencies = [ "arc-swap", "bon", @@ -3070,8 +3069,7 @@ dependencies = [ [[package]] name = "dial9-trace-format" version = "0.2.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "80e0ee560b05f09bf817602d57644947e31e83c521d4e0277f723a6e64d44f92" +source = "git+https://github.com/dial9-rs/dial9-tokio-telemetry.git?rev=60502082601b647c4a51962595721f631b7bbce1#60502082601b647c4a51962595721f631b7bbce1" dependencies = [ "dial9-trace-format-derive", "serde", @@ -3080,8 +3078,7 @@ dependencies = [ [[package]] name = "dial9-trace-format-derive" version = "0.2.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "9dbbd8126d4d6613931317cfe2a7275c1cd487e41c961e42456ab5f956570030" +source = "git+https://github.com/dial9-rs/dial9-tokio-telemetry.git?rev=60502082601b647c4a51962595721f631b7bbce1#60502082601b647c4a51962595721f631b7bbce1" dependencies = [ "proc-macro2", "quote", @@ -3102,7 +3099,7 @@ checksum = "9ed9a281f7bc9b7576e61468ba615a66a5c8cfdff42420a70aa82701a3b1e292" dependencies = [ "block-buffer 0.10.4", "const-oid 0.9.6", - "crypto-common 0.1.7", + "crypto-common 0.1.6", "subtle", ] @@ -3182,10 +3179,13 @@ dependencies = [ "base64 0.22.1", "bytes", "chrono", + "clap", "flatbuffers", "flate2", "futures", "http 1.4.0", + "http-body 1.0.1", + "http-body-util", "md5", "rand 0.10.0", "rcgen", @@ -3353,18 +3353,18 @@ dependencies = [ [[package]] name = "env_filter" -version = "1.0.0" +version = "1.0.1" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "7a1c3cc8e57274ec99de65301228b537f1e4eedc1b8e0f9411c6caac8ae7308f" +checksum = "32e90c2accc4b07a8456ea0debdc2e7587bdd890680d71173a15d4ae604f6eef" dependencies = [ "log", ] [[package]] name = "env_logger" -version = "0.11.9" +version = "0.11.10" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "b2daee4ea451f429a58296525ddf28b45a3b64f1acf6587e2067437bb11e218d" +checksum = "0621c04f2196ac3f488dd583365b9c09be011a4ab8b9f37248ffcc8f6198b56a" dependencies = [ "env_filter", "log", @@ -3454,9 +3454,9 @@ dependencies = [ [[package]] name = "fastrand" -version = "2.3.0" +version = "2.4.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "37909eebbb50d72f9059c3b6d82c0463f2ff062c9e95845c43a6c9c0355411be" +checksum = "a043dc74da1e37d6afe657061213aa6f425f855399a11d3463c6ecccc4dfda1f" [[package]] name = "ff" @@ -3701,9 +3701,9 @@ dependencies = [ [[package]] name = "generic-array" -version = "0.14.7" +version = "0.14.9" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "85649ca51fd72272d7821adaf274ad91c288277713d9c18820d8499a7ff69e9a" +checksum = "4bb6743198531e02858aeaea5398fcc883e71851fcbcb5a2f773e2fb6cb1edf2" dependencies = [ "typenum", "version_check", @@ -4051,7 +4051,7 @@ dependencies = [ "futures-core", "futures-sink", "http 1.4.0", - "indexmap 2.13.0", + "indexmap 2.13.1", "slab", "tokio", "tokio-util", @@ -4313,9 +4313,9 @@ checksum = "135b12329e5e3ce057a9f972339ea52bc954fe1e9358ef27f95e89716fbc5424" [[package]] name = "hybrid-array" -version = "0.4.8" +version = "0.4.10" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "8655f91cd07f2b9d0c24137bd650fe69617773435ee5ec83022377777ce65ef1" +checksum = "3944cf8cf766b40e2a1a333ee5e9b563f854d5fa49d6a8ca2764e97c6eddb214" dependencies = [ "typenum", ] @@ -4425,12 +4425,13 @@ dependencies = [ [[package]] name = "icu_collections" -version = "2.1.1" +version = "2.2.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "4c6b649701667bbe825c3b7e6388cb521c23d88644678e83c0c4d0a621a34b43" +checksum = "2984d1cd16c883d7935b9e07e44071dca8d917fd52ecc02c04d5fa0b5a3f191c" dependencies = [ "displaydoc", "potential_utf", + "utf8_iter", "yoke", "zerofrom", "zerovec", @@ -4438,9 +4439,9 @@ dependencies = [ [[package]] name = "icu_locale_core" -version = "2.1.1" +version = "2.2.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "edba7861004dd3714265b4db54a3c390e880ab658fec5f7db895fae2046b5bb6" +checksum = "92219b62b3e2b4d88ac5119f8904c10f8f61bf7e95b640d25ba3075e6cac2c29" dependencies = [ "displaydoc", "litemap", @@ -4451,9 +4452,9 @@ dependencies = [ [[package]] name = "icu_normalizer" -version = "2.1.1" +version = "2.2.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "5f6c8828b67bf8908d82127b2054ea1b4427ff0230ee9141c54251934ab1b599" +checksum = "c56e5ee99d6e3d33bd91c5d85458b6005a22140021cc324cea84dd0e72cff3b4" dependencies = [ "icu_collections", "icu_normalizer_data", @@ -4465,15 +4466,15 @@ dependencies = [ [[package]] name = "icu_normalizer_data" -version = "2.1.1" +version = "2.2.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "7aedcccd01fc5fe81e6b489c15b247b8b0690feb23304303a9e560f37efc560a" +checksum = "da3be0ae77ea334f4da67c12f149704f19f81d1adf7c51cf482943e84a2bad38" [[package]] name = "icu_properties" -version = "2.1.2" +version = "2.2.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "020bfc02fe870ec3a66d93e677ccca0562506e5872c650f893269e08615d74ec" +checksum = "bee3b67d0ea5c2cca5003417989af8996f8604e34fb9ddf96208a033901e70de" dependencies = [ "icu_collections", "icu_locale_core", @@ -4485,15 +4486,15 @@ dependencies = [ [[package]] name = "icu_properties_data" -version = "2.1.2" +version = "2.2.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "616c294cf8d725c6afcd8f55abc17c56464ef6211f9ed59cccffe534129c77af" +checksum = "8e2bbb201e0c04f7b4b3e14382af113e17ba4f63e2c9d2ee626b720cbce54a14" [[package]] name = "icu_provider" -version = "2.1.1" +version = "2.2.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "85962cf0ce02e1e0a629cc34e7ca3e373ce20dda4c4d7294bbd0bf1fdb59e614" +checksum = "139c4cf31c8b5f33d7e199446eff9c1e02decfc2f0eec2c8d71f65befa45b421" dependencies = [ "displaydoc", "icu_locale_core", @@ -4550,9 +4551,9 @@ dependencies = [ [[package]] name = "indexmap" -version = "2.13.0" +version = "2.13.1" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "7714e70437a7dc3ac8eb7e6f8df75fd8eb422675fc7678aff7364301092b1017" +checksum = "45a8a2b9cb3e0b0c1803dbb0758ffac5de2f425b23c28f518faabd9d805342ff" dependencies = [ "equivalent", "hashbrown 0.16.1", @@ -4567,7 +4568,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "232929e1d75fe899576a3d5c7416ad0d88dbfbb3c3d6aa00873a7408a50ddb88" dependencies = [ "ahash 0.8.12", - "indexmap 2.13.0", + "indexmap 2.13.1", "is-terminal", "itoa", "log", @@ -4590,7 +4591,7 @@ dependencies = [ "crossbeam-utils", "dashmap", "env_logger", - "indexmap 2.13.0", + "indexmap 2.13.1", "itoa", "log", "num-format", @@ -4650,9 +4651,9 @@ dependencies = [ [[package]] name = "iri-string" -version = "0.7.10" +version = "0.7.12" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "c91338f0783edbd6195decb37bae672fd3b165faffb89bf7b9e6942f8b1a731a" +checksum = "25e659a4bb38e810ebc252e53b5814ff908a8c58c2a9ce2fae1bbec24cbf4e20" dependencies = [ "memchr", "serde", @@ -4781,7 +4782,7 @@ dependencies = [ "cesu8", "cfg-if", "combine", - "jni-sys", + "jni-sys 0.3.1", "log", "thiserror 1.0.69", "walkdir", @@ -4790,9 +4791,31 @@ dependencies = [ [[package]] name = "jni-sys" -version = "0.3.0" +version = "0.3.1" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "8eaf4bc02d17cbdd7ff4c7438cafcdf7fb9a4613313ad11b4f8fefe7d3fa0130" +checksum = "41a652e1f9b6e0275df1f15b32661cf0d4b78d4d87ddec5e0c3c20f097433258" +dependencies = [ + "jni-sys 0.4.1", +] + +[[package]] +name = "jni-sys" +version = "0.4.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c6377a88cb3910bee9b0fa88d4f42e1d2da8e79915598f65fb0c7ee14c878af2" +dependencies = [ + "jni-sys-macros", +] + +[[package]] +name = "jni-sys-macros" +version = "0.4.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "38c0b942f458fe50cdac086d2f946512305e5631e720728f2a61aabcd47a6264" +dependencies = [ + "quote", + "syn 2.0.117", +] [[package]] name = "jobserver" @@ -4806,10 +4829,12 @@ dependencies = [ [[package]] name = "js-sys" -version = "0.3.91" +version = "0.3.94" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "b49715b7073f385ba4bc528e5747d02e66cb39c6146efb66b781f131f0fb399c" +checksum = "2e04e2ef80ce82e13552136fabeef8a5ed1f985a96805761cbb9a2c34e7664d9" dependencies = [ + "cfg-if", + "futures-util", "once_cell", "wasm-bindgen", ] @@ -4983,9 +5008,9 @@ dependencies = [ [[package]] name = "liblzma-sys" -version = "0.4.5" +version = "0.4.6" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "9f2db66f3268487b5033077f266da6777d057949b8f93c8ad82e441df25e6186" +checksum = "1a60851d15cd8c5346eca4ab8babff585be2ae4bc8097c067291d3ffe2add3b6" dependencies = [ "cc", "libc", @@ -5010,9 +5035,9 @@ dependencies = [ [[package]] name = "libredox" -version = "0.1.14" +version = "0.1.15" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "1744e39d1d6a9948f4f388969627434e31128196de472883b39f148769bfe30a" +checksum = "7ddbf48fd451246b1f8c2610bd3b4ac0cc6e149d89832867093ab69a17194f08" dependencies = [ "bitflags 2.11.0", "libc", @@ -5084,9 +5109,9 @@ checksum = "32a66949e030da00e8c7d4434b251670a91556f4144941d37452769c25d58a53" [[package]] name = "litemap" -version = "0.8.1" +version = "0.8.2" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "6373607a59f0be73a39b6fe456b8192fcc3585f602af20751600e974dd455e77" +checksum = "92daf443525c4cce67b150400bc2316076100ce0b3686209eb8cf3c31612e6f0" [[package]] name = "local-ip-address" @@ -5299,9 +5324,9 @@ dependencies = [ "crossbeam-epoch", "crossbeam-utils", "hashbrown 0.16.1", - "indexmap 2.13.0", + "indexmap 2.13.1", "metrics", - "ordered-float 5.1.0", + "ordered-float 5.3.0", "quanta", "radix_trie", "rand 0.9.2", @@ -5311,9 +5336,9 @@ dependencies = [ [[package]] name = "metrique" -version = "0.1.22" +version = "0.1.23" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "4f3e5ecbbefec32dafed0fd98ef23768aaade6de35b8434fc3e44f6346b73cd6" +checksum = "c41212ded0b2ba808836b36db163067eb6da41447e7f224267a7331b670009b9" dependencies = [ "itoa", "jiff", @@ -5331,9 +5356,9 @@ dependencies = [ [[package]] name = "metrique-core" -version = "0.1.17" +version = "0.1.18" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "ad6478374c256ffbb0d2de67b7d93e43ac94e35a083f40bd5f72a9770f6110bb" +checksum = "a5f67f8bf383a36ddb563846a9b143228505fd1f833aa936fdefbc9cb1e46855" dependencies = [ "itertools 0.14.0", "metrique-writer-core", @@ -5341,9 +5366,9 @@ dependencies = [ [[package]] name = "metrique-macro" -version = "0.1.14" +version = "0.1.15" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "83adb8929ae9b2f7a4ec07a04c3af569ffe22f96f02c89063e4a78895d6af760" +checksum = "07c50c41313eaa762e16c251aa3aa7f1eff170ab403fac8fc688840e93f39b18" dependencies = [ "Inflector", "darling 0.23.0", @@ -5354,24 +5379,24 @@ dependencies = [ [[package]] name = "metrique-service-metrics" -version = "0.1.18" +version = "0.1.19" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "4d01f36f47452cd6e33f66fc8185bb32f320aaa5721b6ad7230776442d3cf180" +checksum = "07a742784bddd4a8636cf4e952ec4bbd49d4dab34b7dd24d11026ab16adbfe19" dependencies = [ "metrique-writer", ] [[package]] name = "metrique-timesource" -version = "0.1.8" +version = "0.1.9" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "c60fb3f2836dffc05146f0dfe7bf2e0789909f3fefd72c729491adaef01acc1a" +checksum = "d607939211e4eaaa8cd35394fa5e57faffb7390d0ac513b39992edcaf3cc526c" [[package]] name = "metrique-writer" -version = "0.1.19" +version = "0.1.20" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "677d9ba4f5a6b5dd821f78315095840e88d244fafbdda3cf1688835cd2a56aec" +checksum = "8ea8f6776cef2fed8ebaea2044971e4bfec96268eb47ccccb0e93797a198b025" dependencies = [ "ahash 0.8.12", "crossbeam-queue", @@ -5390,9 +5415,9 @@ dependencies = [ [[package]] name = "metrique-writer-core" -version = "0.1.13" +version = "0.1.14" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "642989d2c349dfcd705a0b6b63887459f71c8b8deb6dc79e39e12eaa17400aba" +checksum = "5399135c49a096ba9565872ddfef24627a4129501e930cbffe84d66897d4b12c" dependencies = [ "derive-where", "itertools 0.14.0", @@ -5402,9 +5427,9 @@ dependencies = [ [[package]] name = "metrique-writer-macro" -version = "0.1.7" +version = "0.1.8" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "12edafee41e67f90ab2efe2b850e10751f0da3da4aeb61b8eb7e6c31666e8da8" +checksum = "7417c002f7b01c3d96792ff553b0b7e333059048a9e40d4f8d4bab23f570773e" dependencies = [ "darling 0.23.0", "proc-macro2", @@ -5673,9 +5698,9 @@ dependencies = [ [[package]] name = "num-conv" -version = "0.2.0" +version = "0.2.1" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "cf97ec579c3c42f953ef76dbf8d55ac91fb219dde70e49aa4a6b7d74e9919050" +checksum = "c6673768db2d862beb9b39a78fdcb1a69439615d5794a1be50caa9bc92c81967" [[package]] name = "num-format" @@ -6046,9 +6071,9 @@ dependencies = [ [[package]] name = "ordered-float" -version = "5.1.0" +version = "5.3.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "7f4779c6901a562440c3786d08192c6fbda7c1c2060edd10006b05ee35d10f2d" +checksum = "b7d950ca161dc355eaf28f82b11345ed76c6e1f6eb1f4f4479e0323b9e2fbd0e" dependencies = [ "num-traits", ] @@ -6315,7 +6340,7 @@ checksum = "8701b58ea97060d5e5b155d383a69952a60943f0e6dfe30b04c287beb0b27455" dependencies = [ "fixedbitset", "hashbrown 0.15.5", - "indexmap 2.13.0", + "indexmap 2.13.1", "serde", ] @@ -6398,7 +6423,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "986d2e952779af96ea048f160fd9194e1751b4faea78bcf3ceb456efe008088e" dependencies = [ "der 0.8.0", - "spki 0.8.0-rc.4", + "spki 0.8.0", ] [[package]] @@ -6428,7 +6453,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "12922b6296c06eb741b02d7b5161e3aaa22864af38dfa025a1a3ba3f68c84577" dependencies = [ "der 0.8.0", - "spki 0.8.0-rc.4", + "spki 0.8.0", ] [[package]] @@ -6515,9 +6540,9 @@ dependencies = [ [[package]] name = "potential_utf" -version = "0.1.4" +version = "0.1.5" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "b73949432f5e2a09657003c25bca5e19a0e9c84f8058ca374f49e0ebe605af77" +checksum = "0103b1cef7ec0cf76490e969665504990193874ea05c85ff9bab8b911d0a0564" dependencies = [ "zerovec", ] @@ -6745,7 +6770,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "b4aeaa1f2460f1d348eeaeed86aea999ce98c1bded6f089ff8514c9d9dbdc973" dependencies = [ "anyhow", - "indexmap 2.13.0", + "indexmap 2.13.1", "log", "protobuf", "protobuf-support", @@ -6785,9 +6810,9 @@ dependencies = [ [[package]] name = "pulldown-cmark" -version = "0.13.2" +version = "0.13.3" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "14104c5a24d9bcf7eb2c24753e0f49fe14555d8bd565ea3d38e4b4303267259d" +checksum = "7c3a14896dfa883796f1cb410461aef38810ea05f2b2c33c5aded3649095fdad" dependencies = [ "bitflags 2.11.0", "memchr", @@ -6922,7 +6947,7 @@ dependencies = [ "once_cell", "socket2", "tracing", - "windows-sys 0.59.0", + "windows-sys 0.60.2", ] [[package]] @@ -7471,14 +7496,14 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "87ed3e93fc7e473e464b9726f4759659e72bc8665e4b8ea227547024f416d905" dependencies = [ "const-oid 0.10.2", - "crypto-bigint 0.7.1", + "crypto-bigint 0.7.3", "crypto-primes", "digest 0.11.2", "pkcs1 0.8.0-rc.4", "pkcs8 0.11.0-rc.11", "rand_core 0.10.0", "signature 3.0.0-rc.10", - "spki 0.8.0-rc.4", + "spki 0.8.0", "zeroize", ] @@ -7627,6 +7652,7 @@ dependencies = [ "rustfs-madmin", "rustfs-metrics", "rustfs-notify", + "rustfs-object-io", "rustfs-obs", "rustfs-policy", "rustfs-protocols", @@ -7827,6 +7853,7 @@ dependencies = [ "hyper-rustls", "hyper-util", "lazy_static", + "libc", "md-5 0.11.0", "memmap2 0.9.10", "metrics", @@ -7848,6 +7875,7 @@ dependencies = [ "rustfs-config", "rustfs-credentials", "rustfs-filemeta", + "rustfs-io-core", "rustfs-io-metrics", "rustfs-lock", "rustfs-madmin", @@ -7970,6 +7998,8 @@ name = "rustfs-io-core" version = "0.0.5" dependencies = [ "bytes", + "futures-core", + "futures-util", "memmap2 0.9.10", "rustfs-io-metrics", "thiserror 2.0.18", @@ -8139,6 +8169,31 @@ dependencies = [ "wildmatch", ] +[[package]] +name = "rustfs-object-io" +version = "0.0.5" +dependencies = [ + "astral-tokio-tar", + "atoi", + "bytes", + "futures-util", + "http 1.4.0", + "rustfs-concurrency", + "rustfs-ecstore", + "rustfs-io-core", + "rustfs-io-metrics", + "rustfs-rio", + "rustfs-s3select-api", + "rustfs-utils", + "s3s", + "serial_test", + "thiserror 2.0.18", + "time", + "tokio", + "tokio-util", + "uuid", +] + [[package]] name = "rustfs-obs" version = "0.0.5" @@ -8863,9 +8918,9 @@ dependencies = [ [[package]] name = "semver" -version = "1.0.27" +version = "1.0.28" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "d767eb0aabc880b29956c35734170f26ed551a859dbd361d140cdbeca61ab1e2" +checksum = "8a7852d02fc848982e0c167ef163aaff9cd91dc640ba85e263cb1ce46fae51cd" dependencies = [ "serde", "serde_core", @@ -8983,7 +9038,7 @@ dependencies = [ "chrono", "hex", "indexmap 1.9.3", - "indexmap 2.13.0", + "indexmap 2.13.1", "schemars 0.9.0", "schemars 1.2.1", "serde_core", @@ -9164,9 +9219,9 @@ dependencies = [ [[package]] name = "simd-adler32" -version = "0.3.8" +version = "0.3.9" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "e320a6c5ad31d271ad523dcf3ad13e2767ad8b1cb8f047f75a8aeaf8da139da2" +checksum = "703d5c7ef118737c72f1af64ad2f6f8c5e1921f818cdcb97b8fe6fc69bf66214" [[package]] name = "simdutf8" @@ -9357,9 +9412,9 @@ dependencies = [ [[package]] name = "spki" -version = "0.8.0-rc.4" +version = "0.8.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "8baeff88f34ed0691978ec34440140e1572b68c7dd4a495fd14a3dc1944daa80" +checksum = "1d9efca8738c78ee9484207732f728b1ef517bbb1833d6fc0879ca898a522f6f" dependencies = [ "base64ct", "der 0.8.0", @@ -9507,9 +9562,9 @@ dependencies = [ [[package]] name = "symbolic-common" -version = "12.17.2" +version = "12.17.3" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "751a2823d606b5d0a7616499e4130a516ebd01a44f39811be2b9600936509c23" +checksum = "52ca086c1eb5c7ee74b151ba83c6487d5d33f8c08ad991b86f3f58f6629e68d5" dependencies = [ "debugid", "memmap2 0.9.10", @@ -9519,9 +9574,9 @@ dependencies = [ [[package]] name = "symbolic-demangle" -version = "12.17.2" +version = "12.17.3" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "79b237cfbe320601dd24b4ac817a5b68bb28f5508e33f08d42be0682cadc8ac9" +checksum = "baa911a28a62823aaf2cc2e074212492a3ee69d0d926cc8f5b12b4a108ff5c0c" dependencies = [ "cpp_demangle", "rustc-demangle", @@ -9815,9 +9870,9 @@ dependencies = [ [[package]] name = "tinystr" -version = "0.8.2" +version = "0.8.3" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "42d3e9c45c09de15d06dd8acf5f4e0e399e85927b7f00711024eb7ae10fa4869" +checksum = "c8323304221c2a851516f22236c5722a72eaa19749016521d6dff0824447d96d" dependencies = [ "displaydoc", "zerovec", @@ -10002,7 +10057,7 @@ checksum = "ebe5ef63511595f1344e2d5cfa636d973292adc0eec1f0ad45fae9f0851ab1d4" dependencies = [ "futures-core", "futures-util", - "indexmap 2.13.0", + "indexmap 2.13.1", "pin-project-lite", "slab", "sync_wrapper", @@ -10253,9 +10308,9 @@ checksum = "e6e4313cd5fcd3dad5cafa179702e2b244f760991f45397d14d4ebf38247da75" [[package]] name = "unicode-segmentation" -version = "1.12.0" +version = "1.13.2" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "f6ccf251212114b54433ec949fd6a7841275f9ada20dddd2f29e9ceea4501493" +checksum = "9629274872b2bfaf8d66f5f15725007f635594914870f65218920345aa11aa8c" [[package]] name = "unicode-width" @@ -10417,9 +10472,9 @@ dependencies = [ [[package]] name = "wasm-bindgen" -version = "0.2.114" +version = "0.2.117" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "6532f9a5c1ece3798cb1c2cfdba640b9b3ba884f5db45973a6f442510a87d38e" +checksum = "0551fc1bb415591e3372d0bc4780db7e587d84e2a7e79da121051c5c4b89d0b0" dependencies = [ "cfg-if", "once_cell", @@ -10430,23 +10485,19 @@ dependencies = [ [[package]] name = "wasm-bindgen-futures" -version = "0.4.64" +version = "0.4.67" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "e9c5522b3a28661442748e09d40924dfb9ca614b21c00d3fd135720e48b67db8" +checksum = "03623de6905b7206edd0a75f69f747f134b7f0a2323392d664448bf2d3c5d87e" dependencies = [ - "cfg-if", - "futures-util", "js-sys", - "once_cell", "wasm-bindgen", - "web-sys", ] [[package]] name = "wasm-bindgen-macro" -version = "0.2.114" +version = "0.2.117" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "18a2d50fcf105fb33bb15f00e7a77b772945a2ee45dcf454961fd843e74c18e6" +checksum = "7fbdf9a35adf44786aecd5ff89b4563a90325f9da0923236f6104e603c7e86be" dependencies = [ "quote", "wasm-bindgen-macro-support", @@ -10454,9 +10505,9 @@ dependencies = [ [[package]] name = "wasm-bindgen-macro-support" -version = "0.2.114" +version = "0.2.117" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "03ce4caeaac547cdf713d280eda22a730824dd11e6b8c3ca9e42247b25c631e3" +checksum = "dca9693ef2bab6d4e6707234500350d8dad079eb508dca05530c85dc3a529ff2" dependencies = [ "bumpalo", "proc-macro2", @@ -10467,9 +10518,9 @@ dependencies = [ [[package]] name = "wasm-bindgen-shared" -version = "0.2.114" +version = "0.2.117" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "75a326b8c223ee17883a4251907455a2431acc2791c98c26279376490c378c16" +checksum = "39129a682a6d2d841b6c429d0c51e5cb0ed1a03829d8b3d1e69a011e62cb3d3b" dependencies = [ "unicode-ident", ] @@ -10491,7 +10542,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "bb0e353e6a2fbdc176932bbaab493762eb1255a7900fe0fea1a2f96c296cc909" dependencies = [ "anyhow", - "indexmap 2.13.0", + "indexmap 2.13.1", "wasm-encoder", "wasmparser", ] @@ -10517,15 +10568,15 @@ checksum = "47b807c72e1bac69382b3a6fb3dbe8ea4c0ed87ff5629b8685ae6b9a611028fe" dependencies = [ "bitflags 2.11.0", "hashbrown 0.15.5", - "indexmap 2.13.0", + "indexmap 2.13.1", "semver", ] [[package]] name = "web-sys" -version = "0.3.91" +version = "0.3.94" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "854ba17bb104abfb26ba36da9729addc7ce7f06f5c0f90f3c391f8461cca21f9" +checksum = "cd70027e39b12f0849461e08ffc50b9cd7688d942c1c8e3c7b22273236b4dd0a" dependencies = [ "js-sys", "wasm-bindgen", @@ -10750,6 +10801,15 @@ dependencies = [ "windows-targets 0.52.6", ] +[[package]] +name = "windows-sys" +version = "0.60.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f2f500e4d28234f72040990ec9d39e3a6b950f9f22d3dba18416c35882612bcb" +dependencies = [ + "windows-targets 0.53.5", +] + [[package]] name = "windows-sys" version = "0.61.2" @@ -10783,13 +10843,30 @@ dependencies = [ "windows_aarch64_gnullvm 0.52.6", "windows_aarch64_msvc 0.52.6", "windows_i686_gnu 0.52.6", - "windows_i686_gnullvm", + "windows_i686_gnullvm 0.52.6", "windows_i686_msvc 0.52.6", "windows_x86_64_gnu 0.52.6", "windows_x86_64_gnullvm 0.52.6", "windows_x86_64_msvc 0.52.6", ] +[[package]] +name = "windows-targets" +version = "0.53.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "4945f9f551b88e0d65f3db0bc25c33b8acea4d9e41163edf90dcd0b19f9069f3" +dependencies = [ + "windows-link", + "windows_aarch64_gnullvm 0.53.1", + "windows_aarch64_msvc 0.53.1", + "windows_i686_gnu 0.53.1", + "windows_i686_gnullvm 0.53.1", + "windows_i686_msvc 0.53.1", + "windows_x86_64_gnu 0.53.1", + "windows_x86_64_gnullvm 0.53.1", + "windows_x86_64_msvc 0.53.1", +] + [[package]] name = "windows-threading" version = "0.2.1" @@ -10811,6 +10888,12 @@ version = "0.52.6" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "32a4622180e7a0ec044bb555404c800bc9fd9ec262ec147edd5989ccd0c02cd3" +[[package]] +name = "windows_aarch64_gnullvm" +version = "0.53.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a9d8416fa8b42f5c947f8482c43e7d89e73a173cead56d044f6a56104a6d1b53" + [[package]] name = "windows_aarch64_msvc" version = "0.42.2" @@ -10823,6 +10906,12 @@ version = "0.52.6" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "09ec2a7bb152e2252b53fa7803150007879548bc709c039df7627cabbd05d469" +[[package]] +name = "windows_aarch64_msvc" +version = "0.53.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b9d782e804c2f632e395708e99a94275910eb9100b2114651e04744e9b125006" + [[package]] name = "windows_i686_gnu" version = "0.42.2" @@ -10835,12 +10924,24 @@ version = "0.52.6" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "8e9b5ad5ab802e97eb8e295ac6720e509ee4c243f69d781394014ebfe8bbfa0b" +[[package]] +name = "windows_i686_gnu" +version = "0.53.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "960e6da069d81e09becb0ca57a65220ddff016ff2d6af6a223cf372a506593a3" + [[package]] name = "windows_i686_gnullvm" version = "0.52.6" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "0eee52d38c090b3caa76c563b86c3a4bd71ef1a819287c19d586d7334ae8ed66" +[[package]] +name = "windows_i686_gnullvm" +version = "0.53.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "fa7359d10048f68ab8b09fa71c3daccfb0e9b559aed648a8f95469c27057180c" + [[package]] name = "windows_i686_msvc" version = "0.42.2" @@ -10853,6 +10954,12 @@ version = "0.52.6" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "240948bc05c5e7c6dabba28bf89d89ffce3e303022809e73deaefe4f6ec56c66" +[[package]] +name = "windows_i686_msvc" +version = "0.53.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1e7ac75179f18232fe9c285163565a57ef8d3c89254a30685b57d83a38d326c2" + [[package]] name = "windows_x86_64_gnu" version = "0.42.2" @@ -10865,6 +10972,12 @@ version = "0.52.6" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "147a5c80aabfbf0c7d901cb5895d1de30ef2907eb21fbbab29ca94c5b08b1a78" +[[package]] +name = "windows_x86_64_gnu" +version = "0.53.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9c3842cdd74a865a8066ab39c8a7a473c0778a3f29370b5fd6b4b9aa7df4a499" + [[package]] name = "windows_x86_64_gnullvm" version = "0.42.2" @@ -10877,6 +10990,12 @@ version = "0.52.6" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "24d5b23dc417412679681396f2b49f3de8c1473deb516bd34410872eff51ed0d" +[[package]] +name = "windows_x86_64_gnullvm" +version = "0.53.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0ffa179e2d07eee8ad8f57493436566c7cc30ac536a3379fdf008f47f6bb7ae1" + [[package]] name = "windows_x86_64_msvc" version = "0.42.2" @@ -10889,6 +11008,12 @@ version = "0.52.6" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "589f6da84c646204747d1270a2a5661ea66ed1cced2631d546fdfb155959f9ec" +[[package]] +name = "windows_x86_64_msvc" +version = "0.53.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d6bbff5f0aada427a1e5a6da5f1f98158182f26556f345ac9e04d36d0ebed650" + [[package]] name = "wit-bindgen" version = "0.51.0" @@ -10917,7 +11042,7 @@ checksum = "b7c566e0f4b284dd6561c786d9cb0142da491f46a9fbed79ea69cdad5db17f21" dependencies = [ "anyhow", "heck", - "indexmap 2.13.0", + "indexmap 2.13.1", "prettyplease", "syn 2.0.117", "wasm-metadata", @@ -10948,7 +11073,7 @@ checksum = "9d66ea20e9553b30172b5e831994e35fbde2d165325bec84fc43dbf6f4eb9cb2" dependencies = [ "anyhow", "bitflags 2.11.0", - "indexmap 2.13.0", + "indexmap 2.13.1", "log", "serde", "serde_derive", @@ -10967,7 +11092,7 @@ checksum = "ecc8ac4bc1dc3381b7f59c34f00b67e18f910c2c0f50015669dde7def656a736" dependencies = [ "anyhow", "id-arena", - "indexmap 2.13.0", + "indexmap 2.13.1", "log", "semver", "serde", @@ -10991,9 +11116,9 @@ dependencies = [ [[package]] name = "writeable" -version = "0.6.2" +version = "0.6.3" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "9edde0db4769d2dc68579893f2306b26c6ecfbe0ef499b013d731b7b9247e0b9" +checksum = "1ffae5123b2d3fc086436f8834ae3ab053a283cfac8fe0a0b8eaae044768a4c4" [[package]] name = "x509-parser" @@ -11076,9 +11201,9 @@ dependencies = [ [[package]] name = "yoke" -version = "0.8.1" +version = "0.8.2" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "72d6e5c6afb84d73944e5cedb052c4680d5657337201555f9f2a16b7406d4954" +checksum = "abe8c5fda708d9ca3df187cae8bfb9ceda00dd96231bed36e445a1a48e66f9ca" dependencies = [ "stable_deref_trait", "yoke-derive", @@ -11087,9 +11212,9 @@ dependencies = [ [[package]] name = "yoke-derive" -version = "0.8.1" +version = "0.8.2" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "b659052874eb698efe5b9e8cf382204678a0086ebf46982b79d6ca3182927e5d" +checksum = "de844c262c8848816172cef550288e7dc6c7b7814b4ee56b3e1553f275f1858e" dependencies = [ "proc-macro2", "quote", @@ -11099,18 +11224,18 @@ dependencies = [ [[package]] name = "zerocopy" -version = "0.8.47" +version = "0.8.48" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "efbb2a062be311f2ba113ce66f697a4dc589f85e78a4aea276200804cea0ed87" +checksum = "eed437bf9d6692032087e337407a86f04cd8d6a16a37199ed57949d415bd68e9" dependencies = [ "zerocopy-derive", ] [[package]] name = "zerocopy-derive" -version = "0.8.47" +version = "0.8.48" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "0e8bc7269b54418e7aeeef514aa68f8690b8c0489a06b0136e5f57c4c5ccab89" +checksum = "70e3cd084b1788766f53af483dd21f93881ff30d7320490ec3ef7526d203bad4" dependencies = [ "proc-macro2", "quote", @@ -11119,18 +11244,18 @@ dependencies = [ [[package]] name = "zerofrom" -version = "0.1.6" +version = "0.1.7" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "50cc42e0333e05660c3587f3bf9d0478688e15d870fab3346451ce7f8c9fbea5" +checksum = "69faa1f2a1ea75661980b013019ed6687ed0e83d069bc1114e2cc74c6c04c4df" dependencies = [ "zerofrom-derive", ] [[package]] name = "zerofrom-derive" -version = "0.1.6" +version = "0.1.7" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "d71e5d6e06ab090c67b5e44993ec16b72dcbaabc526db883a360057678b48502" +checksum = "11532158c46691caf0f2593ea8358fed6bbf68a0315e80aae9bd41fbade684a1" dependencies = [ "proc-macro2", "quote", @@ -11160,9 +11285,9 @@ dependencies = [ [[package]] name = "zerotrie" -version = "0.2.3" +version = "0.2.4" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "2a59c17a5562d507e4b54960e8569ebee33bee890c70aa3fe7b97e85a9fd7851" +checksum = "0f9152d31db0792fa83f70fb2f83148effb5c1f5b8c7686c3459e361d9bc20bf" dependencies = [ "displaydoc", "yoke", @@ -11171,9 +11296,9 @@ dependencies = [ [[package]] name = "zerovec" -version = "0.11.5" +version = "0.11.6" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "6c28719294829477f525be0186d13efa9a3c602f7ec202ca9e353d310fb9a002" +checksum = "90f911cbc359ab6af17377d242225f4d75119aec87ea711a880987b18cd7b239" dependencies = [ "yoke", "zerofrom", @@ -11182,9 +11307,9 @@ dependencies = [ [[package]] name = "zerovec-derive" -version = "0.11.2" +version = "0.11.3" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "eadce39539ca5cb3985590102671f2567e659fca9666581ad3411d59207951f3" +checksum = "625dc425cab0dca6dc3c3319506e6593dcb08a9f387ea3b284dbd52a92c40555" dependencies = [ "proc-macro2", "quote", @@ -11205,7 +11330,7 @@ dependencies = [ "flate2", "getrandom 0.4.2", "hmac 0.12.1", - "indexmap 2.13.0", + "indexmap 2.13.1", "lzma-rust2", "memchr", "pbkdf2 0.12.2", diff --git a/Cargo.toml b/Cargo.toml index 1bc8933b2..3cd70f6c6 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -51,6 +51,7 @@ members = [ "crates/workers", # Worker thread pools and task scheduling "crates/io-metrics", # Zero-copy metrics collection for performance analysis "crates/io-core", # Zero-copy core reader and writer implementations + "crates/object-io", # Object I/O policy and zero-copy helper primitives "crates/zip", # ZIP file handling and compression ] resolver = "3" @@ -97,6 +98,7 @@ rustfs-metrics = { path = "crates/metrics", version = "0.0.5" } rustfs-notify = { path = "crates/notify", version = "0.0.5" } rustfs-io-metrics = { path = "crates/io-metrics", version = "0.0.5" } rustfs-io-core = { path = "crates/io-core", version = "0.0.5" } +rustfs-object-io = { path = "crates/object-io", version = "0.0.5" } rustfs-obs = { path = "crates/obs", version = "0.0.5" } rustfs-policy = { path = "crates/policy", version = "0.0.5" } rustfs-protos = { path = "crates/protos", version = "0.0.5" } @@ -284,7 +286,7 @@ zstd = "0.13.3" # Observability and Metrics metrics = "0.24.3" -dial9-tokio-telemetry = "0.2" +dial9-tokio-telemetry = { version = "0.2", git = "https://github.com/dial9-rs/dial9-tokio-telemetry.git", rev = "60502082601b647c4a51962595721f631b7bbce1" } opentelemetry = { version = "0.31.0" } opentelemetry-appender-tracing = { version = "0.31.1", features = ["experimental_use_tracing_span_context", "experimental_metadata_attributes", "spec_unstable_logs_enabled"] } opentelemetry-otlp = { version = "0.31.1", features = ["gzip-http", "reqwest-rustls"] } diff --git a/_typos.toml b/_typos.toml index 2d1aa7e51..6ddfb7ab1 100644 --- a/_typos.toml +++ b/_typos.toml @@ -39,6 +39,7 @@ abd = "abd" mak = "mak" gae = "gae" GAE = "GAE" +writeable = "writeable" # s3-tests original test names (cannot be changed) nonexisted = "nonexisted" consts = "consts" diff --git a/crates/config/src/constants/zero_copy.rs b/crates/config/src/constants/zero_copy.rs index e931bd02c..ff23d6059 100644 --- a/crates/config/src/constants/zero_copy.rs +++ b/crates/config/src/constants/zero_copy.rs @@ -49,6 +49,34 @@ pub const ENV_OBJECT_ZERO_COPY_ENABLE: &str = "RUSTFS_OBJECT_ZERO_COPY_ENABLE"; /// to regular I/O without errors. pub const DEFAULT_OBJECT_ZERO_COPY_ENABLE: bool = true; +/// Environment variable for zero-copy read operating mode. +/// +/// Supported values: +/// - `off`: disable mmap-backed chunk fast path and always use the compatibility path +/// - `conservative`: allow a single mmap window per request +/// - `balanced`: allow multiple mmap windows with the default size guardrails +/// - `aggressive`: allow multi-window mmap and relax the small-object cutoff +pub const ENV_OBJECT_ZERO_COPY_MODE: &str = "RUSTFS_OBJECT_ZERO_COPY_MODE"; + +/// Default zero-copy read mode. +pub const DEFAULT_OBJECT_ZERO_COPY_MODE: &str = "balanced"; + +/// Environment variable for the maximum mmap window size used by the chunk fast path. +/// +/// This controls the visible bytes per mapped chunk before the implementation emits a new window. +pub const ENV_OBJECT_ZERO_COPY_MMAP_WINDOW_BYTES: &str = "RUSTFS_OBJECT_ZERO_COPY_MMAP_WINDOW_BYTES"; + +/// Default mmap window size for chunk fast path reads: 8 MiB. +pub const DEFAULT_OBJECT_ZERO_COPY_MMAP_WINDOW_BYTES: usize = 8 * 1024 * 1024; + +/// Environment variable for the maximum total active mmap bytes. +/// +/// Requests that would exceed this active window budget fall back to the compatibility path. +pub const ENV_OBJECT_ZERO_COPY_MAX_ACTIVE_MMAP_BYTES: &str = "RUSTFS_OBJECT_ZERO_COPY_MAX_ACTIVE_MMAP_BYTES"; + +/// Default maximum active mmap bytes across concurrent local chunk fast-path reads: 256 MiB. +pub const DEFAULT_OBJECT_ZERO_COPY_MAX_ACTIVE_MMAP_BYTES: usize = 256 * 1024 * 1024; + // ============================================================================= // Direct I/O Configuration // ============================================================================= diff --git a/crates/e2e_test/Cargo.toml b/crates/e2e_test/Cargo.toml index bc4d27721..cea060a54 100644 --- a/crates/e2e_test/Cargo.toml +++ b/crates/e2e_test/Cargo.toml @@ -20,6 +20,11 @@ license.workspace = true repository.workspace = true rust-version.workspace = true +[[bin]] +name = "small_put_bench" +path = "src/bin/small_put_bench.rs" +test = false + [lints] workspace = true @@ -41,6 +46,7 @@ serde_json.workspace = true tonic = { workspace = true } tokio = { workspace = true } tokio-stream = { workspace = true } +clap = { workspace = true } rustfs-madmin.workspace = true rustfs-filemeta.workspace = true bytes.workspace = true @@ -52,6 +58,8 @@ async-compression = { workspace = true, features = ["tokio", "bzip2", "xz"] } async-trait = { workspace = true } flate2.workspace = true http.workspace = true +http-body.workspace = true +http-body-util.workspace = true reqwest = { workspace = true } rustfs-signer.workspace = true tracing = { workspace = true } diff --git a/crates/e2e_test/src/bin/small_put_bench.rs b/crates/e2e_test/src/bin/small_put_bench.rs new file mode 100644 index 000000000..75b5864c5 --- /dev/null +++ b/crates/e2e_test/src/bin/small_put_bench.rs @@ -0,0 +1,441 @@ +// 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 anyhow::{Context, Result, anyhow, bail}; +use aws_sdk_s3::config::{Credentials, Region}; +use aws_sdk_s3::primitives::ByteStream; +use aws_sdk_s3::types::{Delete, ObjectIdentifier}; +use aws_sdk_s3::{Client, Config}; +use aws_smithy_http_client::Builder as SmithyHttpClientBuilder; +use bytes::Bytes; +use clap::Parser; +use serde::Serialize; +use std::path::PathBuf; +use std::time::{Duration, Instant}; + +#[derive(Parser, Debug)] +#[command(name = "small_put_bench")] +#[command(about = "Rust-native small PUT benchmark for RustFS-compatible S3 endpoints")] +struct Args { + #[arg(long, env = "RUSTFS_BENCH_ENDPOINT")] + endpoint: String, + + #[arg(long, env = "RUSTFS_BENCH_ACCESS_KEY", default_value = "rustfsadmin")] + access_key: String, + + #[arg(long, env = "RUSTFS_BENCH_SECRET_KEY", default_value = "rustfsadmin")] + secret_key: String, + + #[arg(long, env = "RUSTFS_BENCH_REGION", default_value = "us-east-1")] + region: String, + + #[arg(long, env = "RUSTFS_BENCH_BUCKET", default_value = "small-put-benchmark")] + bucket: String, + + #[arg(long, env = "RUSTFS_BENCH_SIZES", default_value = "4KiB,16KiB,64KiB,256KiB,1MiB")] + sizes: String, + + #[arg(long, env = "RUSTFS_BENCH_CONCURRENCY", default_value_t = 8)] + concurrency: usize, + + #[arg(long, env = "RUSTFS_BENCH_DURATION_SECS", default_value_t = 10)] + duration_secs: u64, + + #[arg(long, env = "RUSTFS_BENCH_TIMEOUT_SECS", default_value_t = 15)] + timeout_secs: u64, + + #[arg(long, env = "RUSTFS_BENCH_PREFIX")] + prefix: Option, + + #[arg(long)] + output_json: Option, + + #[arg(long, default_value_t = false)] + cleanup: bool, +} + +#[derive(Clone, Debug)] +struct SizeSpec { + label: String, + slug: String, + bytes: usize, +} + +#[derive(Debug)] +struct Sample { + ok: bool, + duration_ms: f64, +} + +#[derive(Debug, Serialize)] +struct SizeSummary { + label: String, + bytes: usize, + total: usize, + succeeded: usize, + failed: usize, + wall_secs: f64, + object_rate: f64, + throughput_mib_per_sec: f64, + avg_ms: Option, + p50_ms: Option, + p90_ms: Option, + p99_ms: Option, +} + +#[derive(Debug, Serialize)] +struct RunSummary { + run_id: String, + endpoint: String, + bucket: String, + concurrency: usize, + duration_secs: u64, + timeout_secs: u64, + sizes: Vec, +} + +fn main() -> Result<()> { + let runtime = tokio::runtime::Builder::new_multi_thread() + .enable_all() + .build() + .context("failed to build tokio runtime")?; + runtime.block_on(async_main()) +} + +async fn async_main() -> Result<()> { + let args = Args::parse(); + validate_args(&args)?; + + let sizes = parse_size_list(&args.sizes)?; + let run_id = args.prefix.clone().unwrap_or_else(default_run_id); + let client = build_s3_client(&args.endpoint, &args.access_key, &args.secret_key, &args.region); + + ensure_bucket(&client, &args.bucket).await?; + + let mut size_summaries = Vec::with_capacity(sizes.len()); + for size in &sizes { + let summary = run_size_benchmark( + client.clone(), + args.bucket.clone(), + run_id.clone(), + size.clone(), + args.concurrency, + Duration::from_secs(args.duration_secs), + Duration::from_secs(args.timeout_secs), + ) + .await?; + print_size_summary(&summary); + size_summaries.push(summary); + } + + if args.cleanup { + cleanup_prefix(&client, &args.bucket, &run_id).await?; + } + + let summary = RunSummary { + run_id, + endpoint: args.endpoint, + bucket: args.bucket, + concurrency: args.concurrency, + duration_secs: args.duration_secs, + timeout_secs: args.timeout_secs, + sizes: size_summaries, + }; + + if let Some(path) = args.output_json { + let json = serde_json::to_vec_pretty(&summary).context("failed to serialize benchmark summary")?; + std::fs::write(&path, json).with_context(|| format!("failed to write benchmark summary to {}", path.display()))?; + println!("Wrote summary to {}", path.display()); + } + + Ok(()) +} + +fn validate_args(args: &Args) -> Result<()> { + if args.concurrency == 0 { + bail!("--concurrency must be greater than zero"); + } + if args.duration_secs == 0 { + bail!("--duration-secs must be greater than zero"); + } + if args.timeout_secs == 0 { + bail!("--timeout-secs must be greater than zero"); + } + Ok(()) +} + +fn build_s3_client(endpoint: &str, access_key: &str, secret_key: &str, region: &str) -> Client { + let credentials = Credentials::new(access_key, secret_key, None, None, "small-put-bench"); + let mut config = Config::builder() + .credentials_provider(credentials) + .region(Region::new(region.to_string())) + .endpoint_url(endpoint) + .force_path_style(true) + .behavior_version_latest(); + + if endpoint.starts_with("http://") { + config = config.http_client(SmithyHttpClientBuilder::new().build_http()); + } + + Client::from_conf(config.build()) +} + +async fn ensure_bucket(client: &Client, bucket: &str) -> Result<()> { + if client.head_bucket().bucket(bucket).send().await.is_ok() { + return Ok(()); + } + + match client.create_bucket().bucket(bucket).send().await { + Ok(_) => Ok(()), + Err(err) => { + let rendered = err.to_string(); + if rendered.contains("BucketAlreadyOwnedByYou") || rendered.contains("BucketAlreadyExists") { + Ok(()) + } else { + Err(err).with_context(|| format!("failed to create benchmark bucket {bucket}")) + } + } + } +} + +async fn run_size_benchmark( + client: Client, + bucket: String, + run_id: String, + size: SizeSpec, + concurrency: usize, + duration: Duration, + timeout: Duration, +) -> Result { + let payload = Bytes::from(vec![0_u8; size.bytes]); + let deadline = Instant::now() + duration; + let wall_start = Instant::now(); + + let mut handles = Vec::with_capacity(concurrency); + for worker in 0..concurrency { + let client = client.clone(); + let bucket = bucket.clone(); + let payload = payload.clone(); + let prefix = format!("{run_id}/{}/worker-{worker}", size.slug); + handles.push(tokio::spawn(async move { + let mut samples = Vec::new(); + let mut idx = 0usize; + + while Instant::now() < deadline { + let key = format!("{prefix}/obj-{idx}.bin"); + let started_at = Instant::now(); + let request = client + .put_object() + .bucket(&bucket) + .key(key) + .body(ByteStream::from(payload.clone())) + .content_type("application/octet-stream"); + + let ok = matches!(tokio::time::timeout(timeout, request.send()).await, Ok(Ok(_))); + samples.push(Sample { + ok, + duration_ms: started_at.elapsed().as_secs_f64() * 1000.0, + }); + idx += 1; + } + + samples + })); + } + + let mut samples = Vec::new(); + for handle in handles { + samples.extend(handle.await.map_err(|err| anyhow!("benchmark worker join error: {err}"))?); + } + + Ok(build_size_summary(&size, samples, wall_start.elapsed())) +} + +fn build_size_summary(size: &SizeSpec, mut samples: Vec, wall_elapsed: Duration) -> SizeSummary { + let total = samples.len(); + let succeeded = samples.iter().filter(|sample| sample.ok).count(); + let failed = total.saturating_sub(succeeded); + let wall_secs = wall_elapsed.as_secs_f64(); + let object_rate = if wall_secs > 0.0 { succeeded as f64 / wall_secs } else { 0.0 }; + let throughput_mib_per_sec = if wall_secs > 0.0 { + ((size.bytes * succeeded) as f64 / (1024.0 * 1024.0)) / wall_secs + } else { + 0.0 + }; + + let avg_ms = if total > 0 { + Some(samples.iter().map(|sample| sample.duration_ms).sum::() / total as f64) + } else { + None + }; + + samples.sort_by(|lhs, rhs| lhs.duration_ms.total_cmp(&rhs.duration_ms)); + let durations: Vec = samples.into_iter().map(|sample| sample.duration_ms).collect(); + + SizeSummary { + label: size.label.clone(), + bytes: size.bytes, + total, + succeeded, + failed, + wall_secs, + object_rate, + throughput_mib_per_sec, + avg_ms, + p50_ms: percentile(&durations, 0.50), + p90_ms: percentile(&durations, 0.90), + p99_ms: percentile(&durations, 0.99), + } +} + +async fn cleanup_prefix(client: &Client, bucket: &str, prefix: &str) -> Result<()> { + let mut continuation_token = None; + loop { + let response = client + .list_objects_v2() + .bucket(bucket) + .prefix(prefix) + .set_continuation_token(continuation_token.clone()) + .send() + .await + .with_context(|| format!("failed to list objects for cleanup under {bucket}/{prefix}"))?; + + let objects: Vec = response + .contents + .unwrap_or_default() + .into_iter() + .filter_map(|object| object.key.map(|key| ObjectIdentifier::builder().key(key).build().ok())) + .flatten() + .collect(); + + for chunk in objects.chunks(1_000) { + if chunk.is_empty() { + continue; + } + + client + .delete_objects() + .bucket(bucket) + .delete( + Delete::builder() + .set_objects(Some(chunk.to_vec())) + .quiet(true) + .build() + .context("failed to build delete request")?, + ) + .send() + .await + .with_context(|| format!("failed to delete cleanup batch under {bucket}/{prefix}"))?; + } + + if response.is_truncated.unwrap_or(false) { + continuation_token = response.next_continuation_token; + } else { + break; + } + } + + Ok(()) +} + +fn parse_size_list(input: &str) -> Result> { + input + .split(',') + .map(str::trim) + .filter(|item| !item.is_empty()) + .map(parse_size_spec) + .collect() +} + +fn parse_size_spec(input: &str) -> Result { + let normalized = input.trim(); + let lower = normalized.to_ascii_lowercase(); + + let (number_part, multiplier) = if let Some(value) = lower.strip_suffix("kib") { + (value, 1024usize) + } else if let Some(value) = lower.strip_suffix("mib") { + (value, 1024usize * 1024usize) + } else if let Some(value) = lower.strip_suffix('b') { + (value, 1usize) + } else { + (lower.as_str(), 1usize) + }; + + let value = number_part + .trim() + .parse::() + .with_context(|| format!("invalid size component: {input}"))?; + let bytes = value + .checked_mul(multiplier) + .ok_or_else(|| anyhow!("size overflow for {input}"))?; + + Ok(SizeSpec { + label: normalized.to_string(), + slug: normalized + .chars() + .filter(|ch| ch.is_ascii_alphanumeric()) + .collect::() + .to_ascii_lowercase(), + bytes, + }) +} + +fn percentile(values: &[f64], percentile: f64) -> Option { + if values.is_empty() { + return None; + } + + let index = ((values.len() - 1) as f64 * percentile).floor() as usize; + values.get(index).copied() +} + +fn default_run_id() -> String { + format!("small-put-bench-{}", chrono::Utc::now().format("%Y%m%d-%H%M%S")) +} + +fn print_size_summary(summary: &SizeSummary) { + println!( + "{}: success={} failed={} obj/s={:.3} MiB/s={:.3} avg={:.3?} p50={:.3?} p90={:.3?} p99={:.3?}", + summary.label, + summary.succeeded, + summary.failed, + summary.object_rate, + summary.throughput_mib_per_sec, + summary.avg_ms, + summary.p50_ms, + summary.p90_ms, + summary.p99_ms, + ); +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn parse_size_spec_supports_binary_units() { + let four_kib = parse_size_spec("4KiB").expect("4KiB should parse"); + assert_eq!(four_kib.bytes, 4 * 1024); + + let one_mib = parse_size_spec("1MiB").expect("1MiB should parse"); + assert_eq!(one_mib.bytes, 1024 * 1024); + } + + #[test] + fn percentile_returns_expected_bucket() { + let values = vec![10.0, 20.0, 30.0, 40.0, 50.0]; + assert_eq!(percentile(&values, 0.50), Some(30.0)); + assert_eq!(percentile(&values, 0.90), Some(40.0)); + } +} diff --git a/crates/e2e_test/src/checksum_upload_test.rs b/crates/e2e_test/src/checksum_upload_test.rs index 6be69df85..c15237a15 100644 --- a/crates/e2e_test/src/checksum_upload_test.rs +++ b/crates/e2e_test/src/checksum_upload_test.rs @@ -19,9 +19,13 @@ mod tests { use crate::common::{RustFSTestEnvironment, init_logging}; use aws_sdk_s3::Client; - use aws_sdk_s3::primitives::ByteStream; + use aws_sdk_s3::primitives::{ByteStream, SdkBody}; use aws_sdk_s3::types::{ChecksumAlgorithm, ChecksumMode, CompletedMultipartUpload, CompletedPart}; use base64::Engine; + use bytes::Bytes; + use futures::StreamExt; + use http_body::Frame; + use http_body_util::StreamBody; use rustfs_rio::{Checksum, ChecksumType as RioChecksumType}; use serial_test::serial; use sha2::{Digest, Sha256}; @@ -64,6 +68,53 @@ mod tests { .encoded } + fn streamed_body_70kib_of_a() -> ByteStream { + let bytes = Bytes::from_static(&[b'a'; 1024]); + let stream = futures::stream::repeat_with(move || { + let frame = Frame::data(bytes.clone()); + Ok::<_, std::io::Error>(frame) + }); + let body = WithSizeHint::new(StreamBody::new(stream.take(70)), 70 * 1024); + ByteStream::new(SdkBody::from_body_1_x(body)) + } + + struct WithSizeHint { + inner: T, + size_hint: usize, + } + + impl WithSizeHint { + fn new(inner: T, size_hint: usize) -> Self { + Self { inner, size_hint } + } + } + + impl http_body::Body for WithSizeHint + where + T: http_body::Body + Unpin, + { + type Data = T::Data; + type Error = T::Error; + + fn poll_frame( + self: std::pin::Pin<&mut Self>, + cx: &mut std::task::Context<'_>, + ) -> std::task::Poll, Self::Error>>> { + let this = self.get_mut(); + std::pin::Pin::new(&mut this.inner).poll_frame(cx) + } + + fn is_end_stream(&self) -> bool { + self.inner.is_end_stream() + } + + fn size_hint(&self) -> http_body::SizeHint { + let mut hint = self.inner.size_hint(); + hint.set_exact(self.size_hint as u64); + hint + } + } + /// PutObject with Content-MD5: upload succeeds and GetObject returns same content. #[tokio::test] #[serial] @@ -136,6 +187,121 @@ mod tests { info!("PASSED: PutObject with checksum_sha256 and GetObject content match"); } + /// Mirrors `s3s-e2e` behavior: only request `checksum_algorithm`, then expect + /// both PutObject and GetObject(checksum_mode=enabled) to expose the same checksum. + #[tokio::test] + #[serial] + async fn test_put_object_with_checksum_algorithm_only() { + init_logging(); + info!("TEST: PutObject with checksum_algorithm only"); + + let mut env = RustFSTestEnvironment::new().await.expect("Failed to create test environment"); + env.start_rustfs_server(vec![]).await.expect("Failed to start RustFS"); + + let client = create_s3_client(&env); + let bucket = "test-checksum-algorithm-only"; + create_bucket(&client, bucket).await.expect("Failed to create bucket"); + + let key = "obj-with-checksum-algorithm-only.txt"; + let content = vec![b'a'; 70 * 1024]; + + let put_resp = client + .put_object() + .bucket(bucket) + .key(key) + .checksum_algorithm(ChecksumAlgorithm::Crc32) + .body(ByteStream::from(content.clone())) + .send() + .await + .expect("PutObject with checksum_algorithm should succeed"); + + let put_checksum = put_resp + .checksum_crc32() + .expect("PutObject should return checksum_crc32 when checksum_algorithm is used") + .to_string(); + + let mut get_resp = client + .get_object() + .bucket(bucket) + .key(key) + .checksum_mode(ChecksumMode::Enabled) + .send() + .await + .expect("GetObject should succeed"); + + let body_bytes = std::mem::replace(&mut get_resp.body, ByteStream::new(aws_sdk_s3::primitives::SdkBody::empty())) + .collect() + .await + .expect("collect body") + .into_bytes(); + + assert_eq!(body_bytes.as_ref(), content.as_slice(), "GetObject body must match uploaded content"); + assert_eq!( + get_resp.checksum_crc32().map(str::to_string), + Some(put_checksum), + "GetObject(checksum_mode=enabled) should expose the stored CRC32 checksum" + ); + } + + /// Matches the `s3s-e2e` streaming upload shape more closely than `ByteStream::from(Vec)`. + #[tokio::test] + #[serial] + async fn test_put_object_with_checksum_algorithm_only_streaming_body() { + init_logging(); + info!("TEST: PutObject with checksum_algorithm only using streaming body"); + + let mut env = RustFSTestEnvironment::new().await.expect("Failed to create test environment"); + env.start_rustfs_server(vec![]).await.expect("Failed to start RustFS"); + + let client = create_s3_client(&env); + let bucket = "test-checksum-algorithm-streaming"; + create_bucket(&client, bucket).await.expect("Failed to create bucket"); + + let key = "obj-with-checksum-algorithm-streaming.txt"; + let expected_content = vec![b'a'; 70 * 1024]; + + let put_resp = client + .put_object() + .bucket(bucket) + .key(key) + .checksum_algorithm(ChecksumAlgorithm::Crc32) + .body(streamed_body_70kib_of_a()) + .send() + .await + .expect("PutObject with streaming checksum_algorithm should succeed"); + + let put_checksum = put_resp + .checksum_crc32() + .expect("PutObject should return checksum_crc32 for streaming checksum_algorithm uploads") + .to_string(); + + let mut get_resp = client + .get_object() + .bucket(bucket) + .key(key) + .checksum_mode(ChecksumMode::Enabled) + .send() + .await + .expect("GetObject should succeed"); + + let body_bytes = std::mem::replace(&mut get_resp.body, ByteStream::new(SdkBody::empty())) + .collect() + .await + .expect("collect body") + .into_bytes(); + + assert_eq!( + body_bytes.as_ref(), + expected_content.as_slice(), + "GetObject body must match uploaded content" + ); + assert_eq!( + get_resp.checksum_crc32().map(str::to_string), + Some(put_checksum), + "GetObject(checksum_mode=enabled) should expose the stored CRC32 checksum for streaming uploads" + ); + } + /// Multipart upload with checksum: CreateMultipartUpload, UploadPart(s) with checksum_sha256, CompleteMultipartUpload; then GetObject verifies content. /// Uses part size >= 5MB (server minimum) for two parts. #[tokio::test] @@ -234,6 +400,114 @@ mod tests { info!("PASSED: MultipartUpload with checksum and GetObject content match"); } + /// Mirrors `s3s-e2e` multipart behavior: request checksum algorithm at MPU creation, + /// rely on auto checksum handling during UploadPart, and expect CompleteMultipartUpload to succeed. + #[tokio::test] + #[serial] + async fn test_multipart_upload_with_crc32_algorithm_only() { + init_logging(); + info!("TEST: MultipartUpload with checksum_algorithm only (CRC32)"); + + let mut env = RustFSTestEnvironment::new().await.expect("Failed to create test environment"); + env.start_rustfs_server(vec![]).await.expect("Failed to start RustFS"); + + let client = create_s3_client(&env); + let bucket = "test-multipart-checksum-crc32-auto"; + create_bucket(&client, bucket).await.expect("Failed to create bucket"); + + let key = "multipart-with-crc32-auto.bin"; + let part1_content = "a".repeat(5 * 1024 * 1024 + 1); + let part2_content = "b".repeat(1024); + + let create_resp = client + .create_multipart_upload() + .bucket(bucket) + .key(key) + .checksum_algorithm(ChecksumAlgorithm::Crc32) + .send() + .await + .expect("CreateMultipartUpload should succeed"); + + let upload_id = create_resp.upload_id().expect("upload_id should be present"); + + let part1_resp = client + .upload_part() + .bucket(bucket) + .key(key) + .upload_id(upload_id) + .part_number(1) + .body(ByteStream::from(part1_content.clone().into_bytes())) + .send() + .await + .expect("UploadPart 1 should succeed"); + let part1_checksum = part1_resp + .checksum_crc32() + .expect("UploadPart 1 should return checksum_crc32") + .to_string(); + + let part2_resp = client + .upload_part() + .bucket(bucket) + .key(key) + .upload_id(upload_id) + .part_number(2) + .body(ByteStream::from(part2_content.clone().into_bytes())) + .send() + .await + .expect("UploadPart 2 should succeed"); + let part2_checksum = part2_resp + .checksum_crc32() + .expect("UploadPart 2 should return checksum_crc32") + .to_string(); + + let completed_upload = CompletedMultipartUpload::builder() + .parts( + CompletedPart::builder() + .part_number(1) + .e_tag(part1_resp.e_tag().expect("etag part 1")) + .checksum_crc32(part1_checksum) + .build(), + ) + .parts( + CompletedPart::builder() + .part_number(2) + .e_tag(part2_resp.e_tag().expect("etag part 2")) + .checksum_crc32(part2_checksum) + .build(), + ) + .build(); + + client + .complete_multipart_upload() + .bucket(bucket) + .key(key) + .upload_id(upload_id) + .multipart_upload(completed_upload) + .send() + .await + .expect("CompleteMultipartUpload should succeed"); + + let body_bytes = client + .get_object() + .bucket(bucket) + .key(key) + .send() + .await + .expect("GetObject should succeed") + .body + .collect() + .await + .expect("collect body") + .into_bytes(); + + let expected_content = format!("{part1_content}{part2_content}"); + assert_eq!( + body_bytes.as_ref(), + expected_content.as_bytes(), + "completed multipart object must match concatenated parts" + ); + } + /// Regression test for issue #2282: /// CRC64NVME full-object checksum should match between direct PutObject and multipart upload. #[tokio::test] diff --git a/crates/e2e_test/src/kms/kms_local_test.rs b/crates/e2e_test/src/kms/kms_local_test.rs index d3a24404b..208dbc6a2 100644 --- a/crates/e2e_test/src/kms/kms_local_test.rs +++ b/crates/e2e_test/src/kms/kms_local_test.rs @@ -344,6 +344,9 @@ async fn test_local_kms_multipart_upload() { test_multipart_upload_with_sse_c(&s3_client, TEST_BUCKET) .await .expect("SSE-C multipart upload test failed"); + test_multipart_download_with_wrong_sse_c_key_fails(&s3_client, TEST_BUCKET) + .await + .expect("SSE-C multipart wrong-key download test failed"); // Test 4: Large multipart upload (test streaming encryption with multiple blocks) // TODO: Re-enable after fixing streaming encryption issues with large files @@ -648,6 +651,99 @@ async fn test_multipart_upload_with_sse_c( Ok(()) } +async fn test_multipart_download_with_wrong_sse_c_key_fails( + s3_client: &aws_sdk_s3::Client, + bucket: &str, +) -> Result<(), Box> { + let object_key = "multipart-sse-c-bad-download-test"; + let part_size = 5 * 1024 * 1024; + let total_parts = 2; + let total_size = part_size * total_parts; + + let encryption_key = "01234567890123456789012345678901"; + let key_b64 = base64::Engine::encode(&base64::engine::general_purpose::STANDARD, encryption_key); + let key_md5 = sse_customer_key_md5_base64(encryption_key); + + let wrong_key = "abcdefghijklmnopqrstuvwxyz012345"; + let wrong_key_b64 = base64::Engine::encode(&base64::engine::general_purpose::STANDARD, wrong_key); + let wrong_key_md5 = sse_customer_key_md5_base64(wrong_key); + + let test_data: Vec = (0..total_size).map(|i| ((i * 5) % 256) as u8).collect(); + + let create_multipart_output = s3_client + .create_multipart_upload() + .bucket(bucket) + .key(object_key) + .sse_customer_algorithm("AES256") + .sse_customer_key(&key_b64) + .sse_customer_key_md5(&key_md5) + .send() + .await?; + + let upload_id = create_multipart_output.upload_id().unwrap(); + let mut completed_parts = Vec::new(); + + for part_number in 1..=total_parts { + let start = (part_number - 1) * part_size; + let end = std::cmp::min(start + part_size, total_size); + let part_data = &test_data[start..end]; + + let upload_part_output = s3_client + .upload_part() + .bucket(bucket) + .key(object_key) + .upload_id(upload_id) + .part_number(part_number as i32) + .body(aws_sdk_s3::primitives::ByteStream::from(part_data.to_vec())) + .sse_customer_algorithm("AES256") + .sse_customer_key(&key_b64) + .sse_customer_key_md5(&key_md5) + .send() + .await?; + + completed_parts.push( + aws_sdk_s3::types::CompletedPart::builder() + .part_number(part_number as i32) + .e_tag(upload_part_output.e_tag().unwrap()) + .build(), + ); + } + + let completed_multipart_upload = aws_sdk_s3::types::CompletedMultipartUpload::builder() + .set_parts(Some(completed_parts)) + .build(); + + s3_client + .complete_multipart_upload() + .bucket(bucket) + .key(object_key) + .upload_id(upload_id) + .multipart_upload(completed_multipart_upload) + .send() + .await?; + + let err = s3_client + .get_object() + .bucket(bucket) + .key(object_key) + .sse_customer_algorithm("AES256") + .sse_customer_key(&wrong_key_b64) + .sse_customer_key_md5(&wrong_key_md5) + .send() + .await + .expect_err("multipart SSE-C download with the wrong key should fail"); + + let service_err = err.into_service_error(); + assert_eq!( + service_err.meta().code(), + Some("InvalidRequest"), + "wrong-key multipart SSE-C download should return InvalidRequest, got {:?}", + service_err.meta().code() + ); + + Ok(()) +} + /// Test large multipart upload to verify streaming encryption works correctly #[allow(dead_code)] async fn test_large_multipart_upload( diff --git a/crates/e2e_test/src/lib.rs b/crates/e2e_test/src/lib.rs index be3f3758f..37b73d888 100644 --- a/crates/e2e_test/src/lib.rs +++ b/crates/e2e_test/src/lib.rs @@ -97,6 +97,10 @@ mod cluster_concurrency_test; #[cfg(test)] mod checksum_upload_test; +// Range request regression tests +#[cfg(test)] +mod range_request_test; + // Group deletion tests #[cfg(test)] mod group_delete_test; diff --git a/crates/e2e_test/src/range_request_test.rs b/crates/e2e_test/src/range_request_test.rs new file mode 100644 index 000000000..b5283129c --- /dev/null +++ b/crates/e2e_test/src/range_request_test.rs @@ -0,0 +1,80 @@ +// 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. + +//! End-to-end regression test for invalid GET object ranges. + +#[cfg(test)] +mod tests { + use crate::common::{RustFSTestEnvironment, init_logging}; + use aws_sdk_s3::Client; + use aws_sdk_s3::error::SdkError; + use aws_sdk_s3::primitives::ByteStream; + use serial_test::serial; + use tracing::info; + + fn create_s3_client(env: &RustFSTestEnvironment) -> Client { + env.create_s3_client() + } + + async fn create_bucket(client: &Client, bucket: &str) -> Result<(), Box> { + match client.create_bucket().bucket(bucket).send().await { + Ok(_) => Ok(()), + Err(err) => { + if err.to_string().contains("BucketAlreadyOwnedByYou") || err.to_string().contains("BucketAlreadyExists") { + Ok(()) + } else { + Err(Box::new(err)) + } + } + } + } + + #[tokio::test] + #[serial] + async fn test_get_object_invalid_range_returns_416_issue_s3_implemented_tests() { + init_logging(); + info!("TEST: GetObject invalid range should return InvalidRange/416"); + + let mut env = RustFSTestEnvironment::new().await.expect("Failed to create test environment"); + env.start_rustfs_server(vec![]).await.expect("Failed to start RustFS"); + + let client = create_s3_client(&env); + let bucket = "test-invalid-range"; + let key = "range.txt"; + let content = b"testcontent"; + + create_bucket(&client, bucket).await.expect("Failed to create bucket"); + client + .put_object() + .bucket(bucket) + .key(key) + .body(ByteStream::from_static(content)) + .send() + .await + .expect("PutObject should succeed"); + + let result = client.get_object().bucket(bucket).key(key).range("bytes=40-50").send().await; + + let err = result.expect_err("GetObject with an unsatisfiable range should fail"); + match err { + SdkError::ServiceError(service_err) => { + assert_eq!(service_err.raw().status().as_u16(), 416, "invalid range should return HTTP 416"); + + let s3_err = service_err.into_err(); + assert_eq!(s3_err.meta().code(), Some("InvalidRange"), "invalid range should map to InvalidRange"); + } + other_err => panic!("Expected S3 service error, got: {other_err:?}"), + } + } +} diff --git a/crates/ecstore/Cargo.toml b/crates/ecstore/Cargo.toml index e76563b09..1de8e70b3 100644 --- a/crates/ecstore/Cargo.toml +++ b/crates/ecstore/Cargo.toml @@ -70,6 +70,7 @@ reed-solomon-erasure = { workspace = true } reed-solomon-simd = { workspace = true } lazy_static.workspace = true rustfs-lock.workspace = true +rustfs-io-core.workspace = true rustfs-io-metrics.workspace = true regex = { workspace = true } path-absolutize = { workspace = true } @@ -97,6 +98,7 @@ rand.workspace = true pin-project-lite.workspace = true md-5.workspace = true memmap2 = { workspace = true } +libc.workspace = true rustfs-madmin.workspace = true rustfs-workers.workspace = true reqwest = { workspace = true } @@ -138,5 +140,17 @@ harness = false name = "comparison_benchmark" harness = false +[[bench]] +name = "direct_chunk_benchmark" +harness = false + +[[bench]] +name = "reconstructed_chunk_benchmark" +harness = false + +[[bench]] +name = "bitrot_chunk_benchmark" +harness = false + [lib] doctest = false diff --git a/crates/ecstore/README.md b/crates/ecstore/README.md index bcca1fa8a..4ced3a106 100644 --- a/crates/ecstore/README.md +++ b/crates/ecstore/README.md @@ -32,6 +32,43 @@ For comprehensive documentation, examples, and usage guides, please visit the main [RustFS repository](https://github.com/rustfs/rustfs). +## 📈 Benchmarks + +ECStore ships several Criterion benchmarks under [`crates/ecstore/benches/`](./benches/). + +### Direct Chunk Path + +Use the direct chunk benchmark to compare the current slice-forwarding path against the previous assembled-copy path: + +```bash +cargo bench -p rustfs-ecstore --bench direct_chunk_benchmark +``` + +To run only the end-to-end ECStore range-read benchmark: + +```bash +cargo bench -p rustfs-ecstore --bench direct_chunk_benchmark ecstore_get_object_chunks +``` + +To run the reconstructed multi-disk range-read benchmark: + +```bash +cargo bench -p rustfs-ecstore --bench reconstructed_chunk_benchmark +``` + +### Saved Comparison Points + +Latest local measurements on this branch: + +- `direct_chunk_path/slice_forwarding/single_block_aligned`: about `477 ns` +- `direct_chunk_path/assembled_copy/single_block_aligned`: about `3.25 us` +- `direct_chunk_path/slice_forwarding/multi_block_unaligned`: about `963 ns` +- `direct_chunk_path/assembled_copy/multi_block_unaligned`: about `7.25 us` +- `ecstore_get_object_chunks/drain/multi_disk_range`: about `644-654 us`, throughput about `2.86-2.90 GiB/s` +- `reconstructed_chunk_path/drain/multi_disk_missing_shard`: about `1.292-1.304 ms`, throughput about `1.43-1.45 GiB/s` + +These numbers are intended as branch-local reference points. Re-run the benchmark on your target machine before treating them as a regression baseline. + ## 📄 License This project is licensed under the Apache License 2.0 - see the [LICENSE](../../LICENSE) file for details. diff --git a/crates/ecstore/benches/bitrot_chunk_benchmark.rs b/crates/ecstore/benches/bitrot_chunk_benchmark.rs new file mode 100644 index 000000000..76165c260 --- /dev/null +++ b/crates/ecstore/benches/bitrot_chunk_benchmark.rs @@ -0,0 +1,106 @@ +// 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 bytes::Bytes; +use criterion::{BenchmarkId, Criterion, Throughput, criterion_group, criterion_main}; +use rustfs_ecstore::bitrot::decode_bitrot_chunk_source_for_bench; +use rustfs_io_core::IoChunk; +use rustfs_utils::HashAlgorithm; +use std::hint::black_box; + +struct BitrotChunkBenchCase { + name: &'static str, + source_chunks: Vec, + shard_size: usize, + expected_decoded_len: usize, + expected_copied: bool, +} + +fn encode_shard(checksum_algo: HashAlgorithm, shard: &[u8]) -> Vec { + let mut encoded = Vec::with_capacity(checksum_algo.size() + shard.len()); + encoded.extend_from_slice(checksum_algo.hash_encode(shard).as_ref()); + encoded.extend_from_slice(shard); + encoded +} + +fn bitrot_chunk_bench_cases() -> [BitrotChunkBenchCase; 2] { + let checksum_algo = HashAlgorithm::Md5; + let shard_one = b"abcd"; + let shard_two = b"efgh"; + let encoded_one = encode_shard(checksum_algo.clone(), shard_one); + let encoded_two = encode_shard(checksum_algo.clone(), shard_two); + + let mut cross_chunk = Vec::with_capacity(encoded_one.len() + encoded_two.len()); + cross_chunk.extend_from_slice(&encoded_one); + cross_chunk.extend_from_slice(&encoded_two); + let split = checksum_algo.size() + 2; + + [ + BitrotChunkBenchCase { + name: "aligned_multi_chunk_no_copy", + source_chunks: vec![ + IoChunk::Shared(Bytes::from(encoded_one)), + IoChunk::Shared(Bytes::from(encoded_two)), + ], + shard_size: shard_one.len(), + expected_decoded_len: shard_one.len() + shard_two.len(), + expected_copied: false, + }, + BitrotChunkBenchCase { + name: "cross_chunk_frame_copy", + source_chunks: vec![ + IoChunk::Shared(Bytes::copy_from_slice(&cross_chunk[..split])), + IoChunk::Shared(Bytes::copy_from_slice(&cross_chunk[split..])), + ], + shard_size: shard_one.len(), + expected_decoded_len: shard_one.len() + shard_two.len(), + expected_copied: true, + }, + ] +} + +fn bench_bitrot_chunk_decode(c: &mut Criterion) { + let checksum_algo = HashAlgorithm::Md5; + let mut group = c.benchmark_group("bitrot_chunk_decode"); + group.sample_size(20); + + for case in bitrot_chunk_bench_cases() { + let (decoded, copied) = + decode_bitrot_chunk_source_for_bench(&case.source_chunks, case.shard_size, checksum_algo.clone(), false) + .expect("decode bitrot source"); + let decoded_len: usize = decoded.iter().map(IoChunk::len).sum(); + + assert_eq!(decoded_len, case.expected_decoded_len); + assert_eq!(copied, case.expected_copied); + + group.throughput(Throughput::Bytes(case.expected_decoded_len as u64)); + group.bench_with_input(BenchmarkId::new("decode", case.name), &case, |b, case| { + b.iter(|| { + let result = decode_bitrot_chunk_source_for_bench( + black_box(&case.source_chunks), + black_box(case.shard_size), + checksum_algo.clone(), + false, + ) + .expect("decode bitrot source"); + black_box(result); + }); + }); + } + + group.finish(); +} + +criterion_group!(benches, bench_bitrot_chunk_decode); +criterion_main!(benches); diff --git a/crates/ecstore/benches/direct_chunk_benchmark.rs b/crates/ecstore/benches/direct_chunk_benchmark.rs new file mode 100644 index 000000000..7d45841f3 --- /dev/null +++ b/crates/ecstore/benches/direct_chunk_benchmark.rs @@ -0,0 +1,390 @@ +// 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 bytes::{Bytes, BytesMut}; +use criterion::{BenchmarkId, Criterion, Throughput, criterion_group, criterion_main}; +use futures_util::StreamExt; +use futures_util::stream; +use http::HeaderMap; +use rustfs_ecstore::bucket::metadata_sys; +use rustfs_ecstore::disk::endpoint::Endpoint; +use rustfs_ecstore::endpoints::{EndpointServerPools, Endpoints, PoolEndpoints}; +use rustfs_ecstore::global::{GLOBAL_LOCAL_DISK_ID_MAP, GLOBAL_LOCAL_DISK_MAP, GLOBAL_LOCAL_DISK_SET_DRIVES}; +use rustfs_ecstore::set_disk::collect_direct_data_shard_chunks_for_benchmark; +use rustfs_ecstore::store::{ECStore, init_local_disks}; +use rustfs_ecstore::store_api::{ + BucketOperations, BucketOptions, ChunkNativePutData, GetObjectChunkCopyMode, HTTPRangeSpec, MakeBucketOptions, ObjectIO, + ObjectOptions, +}; +use rustfs_io_core::{BoxChunkStream, IoChunk, MappedChunk}; +use std::hint::black_box; +use std::net::SocketAddr; +use std::sync::Arc; +use std::sync::atomic::{AtomicU16, Ordering}; +use tempfile::TempDir; +use tokio::runtime::Runtime; +use tokio_util::sync::CancellationToken; + +#[derive(Clone)] +struct BenchCase { + name: &'static str, + data_shards: usize, + block_size: usize, + blocks: usize, + offset: usize, + length: usize, +} + +fn bench_cases() -> [BenchCase; 2] { + [ + BenchCase { + name: "single_block_aligned", + data_shards: 4, + block_size: 256 * 1024, + blocks: 1, + offset: 0, + length: 256 * 1024, + }, + BenchCase { + name: "multi_block_unaligned", + data_shards: 4, + block_size: 256 * 1024, + blocks: 8, + offset: 123_457, + length: 2 * 256 * 1024 + 33_333, + }, + ] +} + +#[derive(Clone)] +struct EcstoreBenchCase { + name: &'static str, + disk_count: usize, + payload_len: usize, + range: HTTPRangeSpec, + expected_copy_mode: GetObjectChunkCopyMode, +} + +struct EcstoreBenchEnv { + _temp_dir: TempDir, + store: Arc, + bucket: String, + key: String, + range: HTTPRangeSpec, + opts: ObjectOptions, + expected_len: usize, + expected_copy_mode: GetObjectChunkCopyMode, +} + +fn ecstore_bench_cases() -> [EcstoreBenchCase; 1] { + [EcstoreBenchCase { + name: "multi_disk_range", + disk_count: 4, + payload_len: 3 * 1024 * 1024 + 137, + range: HTTPRangeSpec { + is_suffix_length: false, + start: 123_457, + end: 2 * 1024 * 1024 + 33_333, + }, + expected_copy_mode: expected_direct_copy_mode(), + }] +} + +#[cfg(unix)] +const fn expected_direct_copy_mode() -> GetObjectChunkCopyMode { + GetObjectChunkCopyMode::TrueZeroCopy +} + +#[cfg(not(unix))] +const fn expected_direct_copy_mode() -> GetObjectChunkCopyMode { + GetObjectChunkCopyMode::SharedBytes +} + +fn clone_range_spec(range: &HTTPRangeSpec) -> HTTPRangeSpec { + HTTPRangeSpec { + is_suffix_length: range.is_suffix_length, + start: range.start, + end: range.end, + } +} + +fn next_loopback_addr() -> SocketAddr { + static NEXT_PORT: AtomicU16 = AtomicU16::new(39013); + let port = NEXT_PORT.fetch_add(1, Ordering::Relaxed); + SocketAddr::from(([127, 0, 0, 1], port)) +} + +fn build_endpoint_pools(paths: &[std::path::PathBuf]) -> EndpointServerPools { + let mut endpoints = Vec::with_capacity(paths.len()); + for (idx, disk_path) in paths.iter().enumerate() { + let mut endpoint = Endpoint::try_from(disk_path.to_str().expect("utf8 path")).expect("endpoint"); + endpoint.set_pool_index(0); + endpoint.set_set_index(0); + endpoint.set_disk_index(idx); + endpoints.push(endpoint); + } + + EndpointServerPools(vec![PoolEndpoints { + legacy: false, + set_count: 1, + drives_per_set: paths.len(), + endpoints: Endpoints::from(endpoints), + cmd_line: "bench".to_string(), + platform: format!("OS: {} | Arch: {}", std::env::consts::OS, std::env::consts::ARCH), + }]) +} + +async fn build_ecstore_bench_env(case: &EcstoreBenchCase) -> EcstoreBenchEnv { + let temp_dir = tempfile::tempdir().expect("tempdir"); + let mut disk_paths = Vec::with_capacity(case.disk_count); + for idx in 0..case.disk_count { + let path = temp_dir.path().join(format!("disk{}", idx + 1)); + tokio::fs::create_dir_all(&path).await.expect("create disk dir"); + disk_paths.push(path); + } + + let endpoint_pools = build_endpoint_pools(&disk_paths); + GLOBAL_LOCAL_DISK_MAP.write().await.clear(); + GLOBAL_LOCAL_DISK_ID_MAP.write().await.clear(); + GLOBAL_LOCAL_DISK_SET_DRIVES.write().await.clear(); + init_local_disks(endpoint_pools.clone()).await.expect("init local disks"); + + let store = ECStore::new(next_loopback_addr(), endpoint_pools, CancellationToken::new()) + .await + .expect("create ecstore"); + + let buckets = store + .list_bucket(&BucketOptions { + no_metadata: true, + ..Default::default() + }) + .await + .expect("list buckets") + .into_iter() + .map(|bucket| bucket.name) + .collect(); + metadata_sys::init_bucket_metadata_sys(store.clone(), buckets).await; + + let object_id = case.name.replace('_', "-"); + let bucket = format!("bench-direct-{object_id}"); + let key = format!("objects/{object_id}.bin"); + store + .make_bucket( + &bucket, + &MakeBucketOptions { + versioning_enabled: true, + ..Default::default() + }, + ) + .await + .expect("make bucket"); + + let payload: Vec = (0..case.payload_len).map(|idx| (idx % 251) as u8).collect(); + let mut reader = ChunkNativePutData::from_vec(payload); + let put_info = store + .put_object(&bucket, &key, &mut reader, &ObjectOptions::default()) + .await + .expect("put object"); + if case.disk_count > 1 { + assert!(put_info.data_blocks > 1, "expected multi-data-shard object"); + } + + let (_, expected_len) = case.range.get_offset_length(case.payload_len as i64).expect("range length"); + + EcstoreBenchEnv { + _temp_dir: temp_dir, + store, + bucket, + key, + range: clone_range_spec(&case.range), + opts: ObjectOptions::default(), + expected_len: expected_len as usize, + expected_copy_mode: case.expected_copy_mode, + } +} + +async fn run_ecstore_get_object_chunks_bench(env: &EcstoreBenchEnv) -> (GetObjectChunkCopyMode, usize, usize) { + let mut result = env + .store + .get_object_chunks(&env.bucket, &env.key, Some(clone_range_spec(&env.range)), HeaderMap::new(), &env.opts) + .await + .expect("get object chunks"); + let copy_mode = result.copy_mode; + let mut total_len = 0usize; + let mut chunk_count = 0usize; + while let Some(chunk) = result.stream.next().await { + let chunk = chunk.expect("chunk"); + total_len += chunk.len(); + chunk_count += 1; + } + + (copy_mode, total_len, chunk_count) +} + +fn build_shard_bytes(case: &BenchCase) -> Vec> { + let total_len = case.block_size * case.blocks; + let payload: Vec = (0..total_len).map(|idx| (idx % 251) as u8).collect(); + let mut shards = vec![Vec::with_capacity(case.blocks); case.data_shards]; + + for block in 0..case.blocks { + let block_start = block * case.block_size; + let block_slice = &payload[block_start..block_start + case.block_size]; + let shard_width = case.block_size / case.data_shards; + for (shard_idx, shard) in shards.iter_mut().enumerate().take(case.data_shards) { + let shard_start = shard_idx * shard_width; + let shard_end = shard_start + shard_width; + shard.push(Bytes::copy_from_slice(&block_slice[shard_start..shard_end])); + } + } + + shards +} + +fn build_mapped_streams(shards: &[Vec]) -> Vec { + shards + .iter() + .map(|shard_chunks| { + let chunks: Vec<_> = shard_chunks + .iter() + .map(|chunk| { + let mapped = MappedChunk::new(chunk.clone(), 0, chunk.len()).expect("mapped chunk"); + Ok::(IoChunk::Mapped(mapped)) + }) + .collect(); + Box::pin(stream::iter(chunks)) as BoxChunkStream + }) + .collect() +} + +fn collect_old_assembly( + shards: &[Vec], + data_shards: usize, + block_size: usize, + offset: usize, + length: usize, +) -> Vec { + if length == 0 { + return Vec::new(); + } + + let start_block = offset / block_size; + let end_block = offset.saturating_add(length - 1) / block_size; + let mut result = Vec::with_capacity(end_block - start_block + 1); + + for block_index in start_block..=end_block { + let block_offset = if block_index == start_block { offset % block_size } else { 0 }; + let block_length = if start_block == end_block { + length + } else if block_index == start_block { + block_size - (offset % block_size) + } else if block_index == end_block { + (offset + length) % block_size + } else { + block_size + }; + + if block_length == 0 { + break; + } + + let mut block = BytesMut::with_capacity(block_length); + let mut write_left = block_length; + let mut skip = block_offset; + + for shard in shards.iter().take(data_shards) { + let shard_chunk = &shard[block_index]; + if skip >= shard_chunk.len() { + skip -= shard_chunk.len(); + continue; + } + + let available = &shard_chunk[skip..]; + skip = 0; + let take = available.len().min(write_left); + block.extend_from_slice(&available[..take]); + write_left -= take; + + if write_left == 0 { + break; + } + } + + result.push(IoChunk::Shared(block.freeze())); + } + + result +} + +fn bench_direct_chunk_path(c: &mut Criterion) { + let runtime = Runtime::new().expect("tokio runtime"); + let mut group = c.benchmark_group("direct_chunk_path"); + + for case in bench_cases() { + let shard_bytes = build_shard_bytes(&case); + group.throughput(Throughput::Bytes(case.length as u64)); + + group.bench_with_input(BenchmarkId::new("slice_forwarding", case.name), &case, |b, case| { + b.iter(|| { + let streams = build_mapped_streams(&shard_bytes); + let chunks = runtime + .block_on(collect_direct_data_shard_chunks_for_benchmark( + streams, + case.data_shards, + case.block_size, + case.blocks * case.block_size, + false, + case.offset, + case.length, + )) + .expect("collect direct chunks"); + black_box(chunks); + }); + }); + + group.bench_with_input(BenchmarkId::new("assembled_copy", case.name), &case, |b, case| { + b.iter(|| { + let chunks = collect_old_assembly(&shard_bytes, case.data_shards, case.block_size, case.offset, case.length); + black_box(chunks); + }); + }); + } + + group.finish(); +} + +fn bench_ecstore_get_object_chunks(c: &mut Criterion) { + let runtime = Runtime::new().expect("tokio runtime"); + let mut group = c.benchmark_group("ecstore_get_object_chunks"); + group.sample_size(10); + + for case in ecstore_bench_cases() { + let env = runtime.block_on(build_ecstore_bench_env(&case)); + let (copy_mode, total_len, _) = runtime.block_on(run_ecstore_get_object_chunks_bench(&env)); + assert_eq!(copy_mode, env.expected_copy_mode); + assert_eq!(total_len, env.expected_len); + + group.throughput(Throughput::Bytes(env.expected_len as u64)); + group.bench_with_input(BenchmarkId::new("drain", case.name), &env, |b, env| { + b.iter(|| { + let result = runtime.block_on(run_ecstore_get_object_chunks_bench(env)); + black_box(result); + }); + }); + } + + group.finish(); +} + +criterion_group!(benches, bench_direct_chunk_path, bench_ecstore_get_object_chunks); +criterion_main!(benches); diff --git a/crates/ecstore/benches/reconstructed_chunk_benchmark.rs b/crates/ecstore/benches/reconstructed_chunk_benchmark.rs new file mode 100644 index 000000000..d77b0d9f3 --- /dev/null +++ b/crates/ecstore/benches/reconstructed_chunk_benchmark.rs @@ -0,0 +1,393 @@ +// 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 criterion::{BenchmarkId, Criterion, Throughput, criterion_group, criterion_main}; +use futures_util::StreamExt; +use http::HeaderMap; +use rustfs_ecstore::bucket::metadata_sys; +use rustfs_ecstore::disk::endpoint::Endpoint; +use rustfs_ecstore::endpoints::{EndpointServerPools, Endpoints, PoolEndpoints}; +use rustfs_ecstore::global::{GLOBAL_LOCAL_DISK_ID_MAP, GLOBAL_LOCAL_DISK_MAP, GLOBAL_LOCAL_DISK_SET_DRIVES}; +use rustfs_ecstore::store::{ECStore, init_local_disks}; +use rustfs_ecstore::store_api::{ + BucketOperations, BucketOptions, ChunkNativePutData, CompletePart, GetObjectChunkCopyMode, HTTPRangeSpec, MakeBucketOptions, + MultipartOperations, ObjectOperations, ObjectOptions, +}; +use std::hint::black_box; +use std::net::SocketAddr; +use std::path::{Path, PathBuf}; +use std::sync::Arc; +use std::sync::atomic::{AtomicU16, Ordering}; +use tempfile::TempDir; +use tokio::runtime::Runtime; +use tokio_util::sync::CancellationToken; + +#[derive(Clone)] +enum ReconstructedReadSpec { + PartNumber { + part_number: usize, + missing_part_name: &'static str, + }, + Range { + start: u64, + end: u64, + missing_part_name: &'static str, + }, +} + +#[derive(Clone)] +struct ReconstructedBenchCase { + object_id: &'static str, + name: &'static str, + read_spec: ReconstructedReadSpec, +} + +struct ReconstructedBenchEnv { + store: Arc, + bucket: String, + key: String, + range: HTTPRangeSpec, + opts: ObjectOptions, + expected_len: usize, +} + +struct ReconstructedBenchSuite { + _temp_dir: TempDir, + store: Arc, + disk_paths: Vec, +} + +const MULTIPART_PART_ONE_LEN: usize = 5 * 1024 * 1024; +const MULTIPART_PART_TWO_LEN: usize = 5 * 1024 * 1024 + 137; +const MULTIPART_PART_THREE_LEN: usize = 1024 * 1024 + 77; + +fn reconstructed_bench_cases() -> [ReconstructedBenchCase; 4] { + [ + ReconstructedBenchCase { + object_id: "mp-part2", + name: "multi_disk_missing_shard_multipart_part2", + read_spec: ReconstructedReadSpec::PartNumber { + part_number: 2, + missing_part_name: "part.2", + }, + }, + ReconstructedBenchCase { + object_id: "mp-cross-range", + name: "multi_disk_missing_shard_multipart_cross_part_range", + read_spec: ReconstructedReadSpec::Range { + start: (MULTIPART_PART_ONE_LEN - 32 * 1024) as u64, + end: (MULTIPART_PART_ONE_LEN + 96 * 1024) as u64, + missing_part_name: "part.2", + }, + }, + ReconstructedBenchCase { + object_id: "mp-part3", + name: "multi_disk_missing_shard_multipart_part3", + read_spec: ReconstructedReadSpec::PartNumber { + part_number: 3, + missing_part_name: "part.3", + }, + }, + ReconstructedBenchCase { + object_id: "mp-cross-final-range", + name: "multi_disk_missing_shard_multipart_cross_final_part_range", + read_spec: ReconstructedReadSpec::Range { + start: (MULTIPART_PART_ONE_LEN + MULTIPART_PART_TWO_LEN - 32 * 1024) as u64, + end: (MULTIPART_PART_ONE_LEN + MULTIPART_PART_TWO_LEN + 96 * 1024) as u64, + missing_part_name: "part.3", + }, + }, + ] +} + +fn next_loopback_addr() -> SocketAddr { + static NEXT_PORT: AtomicU16 = AtomicU16::new(39113); + let port = NEXT_PORT.fetch_add(1, Ordering::Relaxed); + SocketAddr::from(([127, 0, 0, 1], port)) +} + +fn clone_range_spec(range: &HTTPRangeSpec) -> HTTPRangeSpec { + HTTPRangeSpec { + is_suffix_length: range.is_suffix_length, + start: range.start, + end: range.end, + } +} + +fn build_endpoint_pools(paths: &[PathBuf]) -> EndpointServerPools { + let mut endpoints = Vec::with_capacity(paths.len()); + for (idx, disk_path) in paths.iter().enumerate() { + let mut endpoint = Endpoint::try_from(disk_path.to_str().expect("utf8 path")).expect("endpoint"); + endpoint.set_pool_index(0); + endpoint.set_set_index(0); + endpoint.set_disk_index(idx); + endpoints.push(endpoint); + } + + EndpointServerPools(vec![PoolEndpoints { + legacy: false, + set_count: 1, + drives_per_set: paths.len(), + endpoints: Endpoints::from(endpoints), + cmd_line: "bench".to_string(), + platform: format!("OS: {} | Arch: {}", std::env::consts::OS, std::env::consts::ARCH), + }]) +} + +fn find_part_files(root: &Path, part_name: &str, out: &mut Vec) { + let Ok(entries) = std::fs::read_dir(root) else { + return; + }; + for entry in entries.flatten() { + let path = entry.path(); + if path.is_dir() { + find_part_files(&path, part_name, out); + continue; + } + + if path.file_name().and_then(|name| name.to_str()) == Some(part_name) { + out.push(path); + } + } +} + +async fn remove_part_files(root: &Path, part_name: &str) -> Vec<(PathBuf, Vec)> { + let mut paths = Vec::new(); + find_part_files(root, part_name, &mut paths); + + let mut removed = Vec::with_capacity(paths.len()); + for path in paths { + let content = tokio::fs::read(&path).await.expect("read part file before removal"); + tokio::fs::remove_file(&path).await.expect("remove part file"); + removed.push((path, content)); + } + + removed +} + +async fn restore_part_files(files: Vec<(PathBuf, Vec)>) { + for (path, content) in files { + if let Some(parent) = path.parent() { + tokio::fs::create_dir_all(parent) + .await + .expect("ensure parent for part restore"); + } + tokio::fs::write(&path, content).await.expect("restore part file"); + } +} + +async fn run_reconstructed_get_object_chunks( + store: &Arc, + bucket: &str, + key: &str, + range: &HTTPRangeSpec, + opts: &ObjectOptions, +) -> (GetObjectChunkCopyMode, usize, usize) { + let mut result = store + .get_object_chunks(bucket, key, Some(clone_range_spec(range)), HeaderMap::new(), opts) + .await + .expect("get object chunks"); + + let copy_mode = result.copy_mode; + let mut total_len = 0usize; + let mut chunk_count = 0usize; + while let Some(chunk) = result.stream.next().await { + let chunk = chunk.expect("chunk"); + total_len += chunk.len(); + chunk_count += 1; + } + + (copy_mode, total_len, chunk_count) +} + +async fn create_multipart_object(store: &Arc, bucket: &str, key: &str, parts: &[Vec]) -> usize { + let upload = store + .new_multipart_upload(bucket, key, &ObjectOptions::default()) + .await + .expect("new multipart upload"); + + let mut completed_parts = Vec::with_capacity(parts.len()); + for (idx, part) in parts.iter().enumerate() { + let mut reader = ChunkNativePutData::from_vec(part.clone()); + let part_info = store + .put_object_part(bucket, key, &upload.upload_id, idx + 1, &mut reader, &ObjectOptions::default()) + .await + .expect("put object part"); + completed_parts.push(CompletePart { + part_num: idx + 1, + etag: part_info.etag, + ..Default::default() + }); + } + + store + .clone() + .complete_multipart_upload(bucket, key, &upload.upload_id, completed_parts, &ObjectOptions::default()) + .await + .expect("complete multipart upload"); + + parts.iter().map(Vec::len).sum() +} + +async fn build_reconstructed_bench_suite() -> ReconstructedBenchSuite { + let temp_dir = tempfile::tempdir().expect("tempdir"); + let disk_paths: Vec<_> = (1..=4).map(|idx| temp_dir.path().join(format!("disk{idx}"))).collect(); + for disk_path in &disk_paths { + tokio::fs::create_dir_all(disk_path).await.expect("create disk dir"); + } + + let endpoint_pools = build_endpoint_pools(&disk_paths); + GLOBAL_LOCAL_DISK_MAP.write().await.clear(); + GLOBAL_LOCAL_DISK_ID_MAP.write().await.clear(); + GLOBAL_LOCAL_DISK_SET_DRIVES.write().await.clear(); + init_local_disks(endpoint_pools.clone()).await.expect("init local disks"); + + let store = ECStore::new(next_loopback_addr(), endpoint_pools, CancellationToken::new()) + .await + .expect("create ecstore"); + + let buckets = store + .list_bucket(&BucketOptions { + no_metadata: true, + ..Default::default() + }) + .await + .expect("list buckets") + .into_iter() + .map(|bucket| bucket.name) + .collect(); + metadata_sys::init_bucket_metadata_sys(store.clone(), buckets).await; + + ReconstructedBenchSuite { + _temp_dir: temp_dir, + store, + disk_paths, + } +} + +async fn build_reconstructed_bench_env(suite: &ReconstructedBenchSuite, case: &ReconstructedBenchCase) -> ReconstructedBenchEnv { + let bucket = format!("bench-r-{}", case.object_id); + let key = format!("objects/{}.bin", case.object_id); + suite + .store + .make_bucket( + &bucket, + &MakeBucketOptions { + versioning_enabled: true, + ..Default::default() + }, + ) + .await + .expect("make bucket"); + + let part_one: Vec = (0..MULTIPART_PART_ONE_LEN).map(|idx| (idx % 251) as u8).collect(); + let part_two: Vec = (0..MULTIPART_PART_TWO_LEN).map(|idx| ((idx + 11) % 251) as u8).collect(); + let part_three: Vec = (0..MULTIPART_PART_THREE_LEN).map(|idx| ((idx + 29) % 251) as u8).collect(); + let payload_len = create_multipart_object(&suite.store, &bucket, &key, &[part_one, part_two, part_three]).await; + let info = suite + .store + .get_object_info(&bucket, &key, &ObjectOptions::default()) + .await + .expect("get object info"); + let (range, opts, missing_part_name) = match case.read_spec.clone() { + ReconstructedReadSpec::PartNumber { + part_number, + missing_part_name, + } => { + let range = HTTPRangeSpec::from_object_info(&info, part_number).expect("part_number range"); + let opts = ObjectOptions { + part_number: Some(part_number), + ..Default::default() + }; + (range, opts, missing_part_name) + } + ReconstructedReadSpec::Range { + start, + end, + missing_part_name, + } => ( + HTTPRangeSpec { + is_suffix_length: false, + start: start as i64, + end: end as i64, + }, + ObjectOptions::default(), + missing_part_name, + ), + }; + + let (_, expected_len) = range.get_offset_length(payload_len as i64).expect("range length"); + let mut selected_reconstructed = false; + for disk_path in &suite.disk_paths { + let object_root = disk_path + .join(&bucket) + .join("objects") + .join(format!("{}.bin", case.object_id)); + let removed_parts = remove_part_files(&object_root, missing_part_name).await; + if removed_parts.is_empty() { + continue; + } + + let (copy_mode, total_len, _) = run_reconstructed_get_object_chunks(&suite.store, &bucket, &key, &range, &opts).await; + if copy_mode == GetObjectChunkCopyMode::Reconstructed && total_len == expected_len as usize { + selected_reconstructed = true; + break; + } + + restore_part_files(removed_parts).await; + } + assert!( + selected_reconstructed, + "failed to select a missing shard placement that triggers reconstructed copy mode for benchmark case {}", + case.name + ); + + ReconstructedBenchEnv { + store: suite.store.clone(), + bucket, + key, + range: clone_range_spec(&range), + opts, + expected_len: expected_len as usize, + } +} + +async fn run_reconstructed_get_object_chunks_bench(env: &ReconstructedBenchEnv) -> (GetObjectChunkCopyMode, usize, usize) { + run_reconstructed_get_object_chunks(&env.store, &env.bucket, &env.key, &env.range, &env.opts).await +} + +fn bench_reconstructed_chunk_path(c: &mut Criterion) { + let runtime = Runtime::new().expect("tokio runtime"); + let suite = runtime.block_on(build_reconstructed_bench_suite()); + let mut group = c.benchmark_group("reconstructed_chunk_path"); + group.sample_size(10); + for case in reconstructed_bench_cases() { + let env = runtime.block_on(build_reconstructed_bench_env(&suite, &case)); + let (copy_mode, total_len, _) = runtime.block_on(run_reconstructed_get_object_chunks_bench(&env)); + assert_eq!(copy_mode, GetObjectChunkCopyMode::Reconstructed); + assert_eq!(total_len, env.expected_len); + + group.throughput(Throughput::Bytes(env.expected_len as u64)); + group.bench_with_input(BenchmarkId::new("drain", case.name), &env, |b, env| { + b.iter(|| { + let result = runtime.block_on(run_reconstructed_get_object_chunks_bench(env)); + black_box(result); + }); + }); + } + group.finish(); +} + +criterion_group!(benches, bench_reconstructed_chunk_path); +criterion_main!(benches); diff --git a/crates/ecstore/run_benchmarks.sh b/crates/ecstore/run_benchmarks.sh index 7e5266c3e..8a2e86d9b 100755 --- a/crates/ecstore/run_benchmarks.sh +++ b/crates/ecstore/run_benchmarks.sh @@ -119,6 +119,16 @@ run_large_data_test() { print_success "Large-dataset tests completed" } +# Run direct chunk path benchmarks +run_direct_chunk_benchmark() { + print_info "📦 Starting direct chunk path benchmarks..." + echo "================================================" + + cargo bench --bench direct_chunk_benchmark + + print_success "Direct chunk path benchmarks completed" +} + # Generate comparison report generate_comparison_report() { print_info "📊 Generating performance report..." @@ -168,6 +178,7 @@ show_help() { echo " full Run the full benchmark suite" echo " performance Run detailed performance tests" echo " simd Run the SIMD-only tests" + echo " direct Run the direct chunk path benchmarks" echo " large Run large-dataset tests" echo " clean Remove previous results" echo " help Show this help message" @@ -177,6 +188,7 @@ show_help() { echo " $0 performance # Detailed performance test" echo " $0 full # Full benchmark suite" echo " $0 simd # SIMD-only benchmark" + echo " $0 direct # Direct chunk path benchmark" echo " $0 large # Large-dataset benchmark" echo "" echo "Features:" @@ -240,6 +252,11 @@ main() { run_simd_benchmark generate_comparison_report ;; + "direct") + cleanup + run_direct_chunk_benchmark + generate_comparison_report + ;; "large") cleanup run_large_data_test @@ -263,4 +280,4 @@ main() { } # Launch script -main "$@" \ No newline at end of file +main "$@" diff --git a/crates/ecstore/src/bitrot.rs b/crates/ecstore/src/bitrot.rs index 6969edf73..eb52009cf 100644 --- a/crates/ecstore/src/bitrot.rs +++ b/crates/ecstore/src/bitrot.rs @@ -14,13 +14,328 @@ use crate::disk::{self, DiskAPI as _, DiskStore, error::DiskError}; use crate::erasure_coding::{BitrotReader, BitrotWriterWrapper, CustomWriter}; -use bytes::Bytes; +use crate::store_api::{GetObjectChunkCopyMode, GetObjectChunkPath, GetObjectChunkResult}; +use bytes::{Bytes, BytesMut}; +use futures_util::{StreamExt, stream}; +use rustfs_io_core::{BoxChunkStream, IoChunk}; use rustfs_utils::HashAlgorithm; +use std::collections::VecDeque; use std::io::Cursor; use std::time::Instant; use tokio::io::AsyncRead; use tracing::debug; +const BITROT_READ_OPERATION: &str = "bitrot_read"; + +fn classify_chunk_copy_mode(source_direct: bool, copied: bool) -> GetObjectChunkCopyMode { + if copied { + GetObjectChunkCopyMode::SingleCopy + } else if source_direct { + GetObjectChunkCopyMode::TrueZeroCopy + } else { + GetObjectChunkCopyMode::SharedBytes + } +} + +struct ChunkSpan { + bytes: Bytes, + chunk: IoChunk, + copied: bool, +} + +fn take_contiguous_chunk_span(chunk: &IoChunk, offset: usize, len: usize) -> std::io::Result { + match chunk { + IoChunk::Shared(bytes) => { + let bytes = bytes.slice(offset..offset + len); + Ok(ChunkSpan { + bytes: bytes.clone(), + chunk: IoChunk::Shared(bytes), + copied: false, + }) + } + IoChunk::Mapped(mapped) => { + let chunk = IoChunk::Mapped(mapped.slice(offset, len)?); + let bytes = chunk.as_bytes(); + Ok(ChunkSpan { + bytes, + chunk, + copied: false, + }) + } + IoChunk::Pooled(pooled) => { + let chunk = IoChunk::Pooled(pooled.slice(offset, len)?); + let bytes = chunk.as_bytes(); + Ok(ChunkSpan { + bytes, + chunk, + copied: false, + }) + } + } +} + +struct BitrotChunkSource { + source_stream: BoxChunkStream, + source_chunks: VecDeque, + source_chunk_offset: usize, + source_buffered_bytes: usize, + source_done: bool, +} + +struct BitrotChunkStreamState { + source: BitrotChunkSource, + decoded_remaining: usize, + trim_prefix: usize, + output_remaining: usize, + shard_size: usize, + checksum_algo: HashAlgorithm, + skip_verify: bool, +} + +struct ChunkCursor<'a> { + chunks: &'a [IoChunk], + chunk_index: usize, + chunk_offset: usize, + consumed: usize, + total_len: usize, +} + +impl<'a> ChunkCursor<'a> { + fn new(chunks: &'a [IoChunk]) -> Self { + Self { + chunks, + chunk_index: 0, + chunk_offset: 0, + consumed: 0, + total_len: chunks.iter().map(IoChunk::len).sum(), + } + } + + fn remaining(&self) -> usize { + self.total_len.saturating_sub(self.consumed) + } + + fn skip_empty_chunks(&mut self) { + while let Some(chunk) = self.chunks.get(self.chunk_index) { + if self.chunk_offset < chunk.len() { + break; + } + self.chunk_index += 1; + self.chunk_offset = 0; + } + } + + fn advance(&mut self, len: usize) { + self.consumed += len; + self.chunk_offset += len; + self.skip_empty_chunks(); + } + + fn take_span(&mut self, len: usize) -> std::io::Result { + self.skip_empty_chunks(); + if self.remaining() < len { + return Err(std::io::Error::new(std::io::ErrorKind::UnexpectedEof, "truncated bitrot chunk source")); + } + + let Some(chunk) = self.chunks.get(self.chunk_index) else { + return Err(std::io::Error::new(std::io::ErrorKind::UnexpectedEof, "missing bitrot chunk source")); + }; + let available = chunk.len().saturating_sub(self.chunk_offset); + + if len <= available { + let span = take_contiguous_chunk_span(chunk, self.chunk_offset, len)?; + self.advance(len); + return Ok(span); + } + + let mut aggregate = BytesMut::with_capacity(len); + let mut remaining = len; + while remaining > 0 { + self.skip_empty_chunks(); + let Some(chunk) = self.chunks.get(self.chunk_index) else { + return Err(std::io::Error::new(std::io::ErrorKind::UnexpectedEof, "truncated bitrot chunk source")); + }; + let available = chunk.len().saturating_sub(self.chunk_offset); + let take = available.min(remaining); + aggregate.extend_from_slice(&chunk.as_bytes()[self.chunk_offset..self.chunk_offset + take]); + self.advance(take); + remaining -= take; + } + + let bytes = aggregate.freeze(); + Ok(ChunkSpan { + bytes: bytes.clone(), + chunk: IoChunk::Shared(bytes), + copied: true, + }) + } +} + +impl BitrotChunkSource { + fn new(source_stream: BoxChunkStream, source_chunks: VecDeque, source_done: bool) -> Self { + let source_buffered_bytes = source_chunks.iter().map(IoChunk::len).sum(); + Self { + source_stream, + source_chunks, + source_chunk_offset: 0, + source_buffered_bytes, + source_done, + } + } + + fn skip_empty_chunks(&mut self) { + while let Some(chunk) = self.source_chunks.front() { + if self.source_chunk_offset < chunk.len() { + break; + } + self.source_chunks.pop_front(); + self.source_chunk_offset = 0; + } + } + + async fn fill(&mut self, min_bytes: usize) -> std::io::Result<()> { + while self.source_buffered_bytes < min_bytes && !self.source_done { + match self.source_stream.next().await { + Some(Ok(chunk)) => { + self.source_buffered_bytes += chunk.len(); + self.source_chunks.push_back(chunk); + } + Some(Err(err)) => return Err(err), + None => self.source_done = true, + } + } + + if self.source_buffered_bytes < min_bytes { + return Err(std::io::Error::new(std::io::ErrorKind::UnexpectedEof, "truncated bitrot chunk source")); + } + + Ok(()) + } + + fn advance(&mut self, len: usize) { + self.source_buffered_bytes = self.source_buffered_bytes.saturating_sub(len); + self.source_chunk_offset += len; + self.skip_empty_chunks(); + } + + fn take_span(&mut self, len: usize) -> std::io::Result { + self.skip_empty_chunks(); + if self.source_buffered_bytes < len { + return Err(std::io::Error::new(std::io::ErrorKind::UnexpectedEof, "truncated bitrot chunk source")); + } + + let Some(chunk) = self.source_chunks.front() else { + return Err(std::io::Error::new(std::io::ErrorKind::UnexpectedEof, "missing bitrot chunk source")); + }; + let available = chunk.len().saturating_sub(self.source_chunk_offset); + + if len <= available { + let span = take_contiguous_chunk_span(chunk, self.source_chunk_offset, len)?; + self.advance(len); + return Ok(span); + } + + let mut aggregate = BytesMut::with_capacity(len); + let mut remaining = len; + while remaining > 0 { + self.skip_empty_chunks(); + let Some(chunk) = self.source_chunks.front() else { + return Err(std::io::Error::new(std::io::ErrorKind::UnexpectedEof, "truncated bitrot chunk source")); + }; + let available = chunk.len().saturating_sub(self.source_chunk_offset); + let take = available.min(remaining); + aggregate.extend_from_slice(&chunk.as_bytes()[self.source_chunk_offset..self.source_chunk_offset + take]); + self.advance(take); + remaining -= take; + } + + let bytes = aggregate.freeze(); + Ok(ChunkSpan { + bytes: bytes.clone(), + chunk: IoChunk::Shared(bytes), + copied: true, + }) + } +} + +impl BitrotChunkStreamState { + #[allow(clippy::too_many_arguments)] + fn new( + source_stream: BoxChunkStream, + source_chunks: VecDeque, + source_done: bool, + decoded_remaining: usize, + trim_prefix: usize, + output_remaining: usize, + shard_size: usize, + checksum_algo: HashAlgorithm, + skip_verify: bool, + ) -> Self { + Self { + source: BitrotChunkSource::new(source_stream, source_chunks, source_done), + decoded_remaining, + trim_prefix, + output_remaining, + shard_size, + checksum_algo, + skip_verify, + } + } + + fn hash_size(&self) -> usize { + self.checksum_algo.size() + } + + async fn next_verified_chunk(&mut self) -> std::io::Result> { + let hash_size = self.hash_size(); + + while self.output_remaining > 0 && self.decoded_remaining > 0 { + let data_len = self.shard_size.min(self.decoded_remaining); + + let expected_hash = if hash_size > 0 { + self.source.fill(hash_size).await?; + Some(self.source.take_span(hash_size)?) + } else { + None + }; + + self.source.fill(data_len).await?; + let data_span = self.source.take_span(data_len)?; + + if let Some(expected_hash) = expected_hash + && !self.skip_verify + && self.checksum_algo.hash_encode(data_span.bytes.as_ref()).as_ref() != expected_hash.bytes.as_ref() + { + return Err(std::io::Error::new(std::io::ErrorKind::InvalidData, "bitrot hash mismatch")); + } + + self.decoded_remaining -= data_len; + + if self.trim_prefix >= data_len { + self.trim_prefix -= data_len; + continue; + } + + let start = self.trim_prefix; + self.trim_prefix = 0; + let take = (data_len - start).min(self.output_remaining); + self.output_remaining -= take; + + let chunk = if start == 0 && take == data_len { + data_span.chunk + } else { + data_span.chunk.slice(start, take)? + }; + + if !chunk.is_empty() { + return Ok(Some(chunk)); + } + } + + Ok(None) + } +} + /// Create a BitrotReader from either inline data or disk file stream /// /// # Parameters @@ -65,20 +380,39 @@ pub async fn create_bitrot_reader( } else if let Some(disk) = disk { // Read from disk if use_zero_copy { + if !disk.is_local() { + rustfs_io_metrics::record_io_path_selected(BITROT_READ_OPERATION, rustfs_io_metrics::IoPath::Legacy); + rustfs_io_metrics::record_io_fallback( + rustfs_io_metrics::IoStage::ReadSetup, + rustfs_io_metrics::FallbackReason::NonLocalBackend, + ); + + let rd = disk.read_file_stream(bucket, path, offset, length).await?; + let reader = BitrotReader::new(rd, shard_size, checksum_algo, skip_verify); + return Ok(Some(reader)); + } + // Try zero-copy read first (uses mmap on Unix) let start = Instant::now(); match disk.read_file_zero_copy(bucket, path, offset, length).await { Ok(bytes) => { let duration_ms = start.elapsed().as_secs_f64() * 1000.0; - // Record zero-copy metrics - rustfs_io_metrics::record_zero_copy_read(bytes.len(), duration_ms); + rustfs_io_metrics::record_io_path_selected(BITROT_READ_OPERATION, rustfs_io_metrics::IoPath::Fast); + // `read_file_zero_copy()` returns a shared `Bytes` view, but it may still + // internally aggregate multiple chunk windows. The exact chunk-native copy + // mode is only preserved by `create_bitrot_chunk_stream()`. + rustfs_io_metrics::record_io_copy_mode( + BITROT_READ_OPERATION, + rustfs_io_metrics::CopyMode::SharedBytes, + bytes.len(), + ); - // Log successful zero-copy read debug!( size = bytes.len(), + duration_ms, path = %path, - "zero_copy_read_success" + "bitrot_fast_read_success" ); // Wrap Bytes in Cursor for AsyncRead @@ -93,14 +427,16 @@ pub async fn create_bitrot_reader( Ok(Some(reader)) } Err(e) => { - // Record zero-copy fallback - rustfs_io_metrics::record_zero_copy_fallback(&format!("{:?}", e)); + rustfs_io_metrics::record_io_path_selected(BITROT_READ_OPERATION, rustfs_io_metrics::IoPath::Legacy); + rustfs_io_metrics::record_io_fallback( + rustfs_io_metrics::IoStage::ReadSetup, + rustfs_io_metrics::FallbackReason::Unknown, + ); - // Log zero-copy fallback debug!( - reason = %format!("{:?}", e), + reason = %e, path = %path, - "zero_copy_fallback" + "bitrot_fast_read_fallback" ); // Fall back to regular stream read on error @@ -117,6 +453,7 @@ pub async fn create_bitrot_reader( } } } else { + rustfs_io_metrics::record_io_path_selected(BITROT_READ_OPERATION, rustfs_io_metrics::IoPath::Legacy); // Use regular stream read match disk.read_file_stream(bucket, path, offset, length).await { Ok(rd) => { @@ -132,6 +469,198 @@ pub async fn create_bitrot_reader( } } +/// Create a chunk stream from bitrot-encoded data, preserving source chunk provenance when possible. +#[allow(clippy::too_many_arguments)] +pub async fn create_bitrot_chunk_stream( + inline_data: Option<&[u8]>, + disk: Option<&DiskStore>, + bucket: &str, + path: &str, + offset: usize, + length: usize, + total_data_size: usize, + shard_size: usize, + checksum_algo: HashAlgorithm, + skip_verify: bool, + use_zero_copy: bool, +) -> disk::error::Result> { + let fetch_start = (offset / shard_size) * shard_size; + let fetch_end = (offset + length).div_ceil(shard_size) * shard_size; + let fetch_end = fetch_end.min(total_data_size); + let fetch_length = fetch_end.saturating_sub(fetch_start); + let trim_prefix = offset.saturating_sub(fetch_start); + let hash_size = checksum_algo.size(); + let encoded_length = fetch_length.div_ceil(shard_size) * hash_size + fetch_length; + let encoded_offset = fetch_start.div_ceil(shard_size) * hash_size + fetch_start; + + let mut source_done = false; + let (source_stream, mut prefetched_chunks, source_direct) = if let Some(data) = inline_data { + source_done = true; + let mut chunks = VecDeque::new(); + chunks.push_back(IoChunk::Shared( + Bytes::copy_from_slice(data).slice(encoded_offset..encoded_offset + encoded_length), + )); + let source_stream: BoxChunkStream = Box::pin(stream::empty::>()); + (source_stream, chunks, false) + } else if let Some(disk) = disk { + if use_zero_copy { + let mut source_stream = disk.read_file_chunks(bucket, path, encoded_offset, encoded_length).await?; + let mut prefetched_chunks = VecDeque::new(); + let mut direct = true; + while prefetched_chunks.len() < 2 { + let Some(chunk) = source_stream.next().await else { + source_done = true; + break; + }; + let chunk = chunk?; + direct &= matches!(chunk, IoChunk::Mapped(_)); + prefetched_chunks.push_back(chunk); + } + (source_stream, prefetched_chunks, direct) + } else { + source_done = true; + let bytes = disk.read_file_zero_copy(bucket, path, encoded_offset, encoded_length).await?; + let mut chunks = VecDeque::new(); + chunks.push_back(IoChunk::Shared(bytes)); + let source_stream: BoxChunkStream = Box::pin(stream::empty::>()); + (source_stream, chunks, false) + } + } else { + return Ok(None); + }; + + let copied = predicted_stream_copy(encoded_length, shard_size, checksum_algo.size(), &prefetched_chunks, source_done); + let state = BitrotChunkStreamState::new( + source_stream, + std::mem::take(&mut prefetched_chunks), + source_done, + fetch_length, + trim_prefix, + length, + shard_size, + checksum_algo, + skip_verify, + ); + let stream = stream::unfold(Some(state), |state| async move { + let mut state = match state { + Some(state) => state, + None => return None, + }; + + match state.next_verified_chunk().await { + Ok(Some(chunk)) => { + let next_state = if state.output_remaining == 0 { None } else { Some(state) }; + Some((Ok::(chunk), next_state)) + } + Ok(None) => None, + Err(err) => Some((Err(err), None)), + } + }); + Ok(Some(GetObjectChunkResult { + stream: Box::pin(stream), + path: GetObjectChunkPath::Direct, + copy_mode: classify_chunk_copy_mode(source_direct, copied), + })) +} + +fn predicted_stream_copy( + encoded_length: usize, + shard_size: usize, + hash_size: usize, + prefetched_chunks: &VecDeque, + source_done: bool, +) -> bool { + if prefetched_chunks.is_empty() { + return false; + } + + if source_done && prefetched_chunks.len() == 1 { + return false; + } + + let full_frame_len = hash_size + shard_size; + if full_frame_len == 0 { + return false; + } + + let first_window_len = prefetched_chunks.front().map(IoChunk::len).unwrap_or(encoded_length); + encoded_length > first_window_len && !first_window_len.is_multiple_of(full_frame_len) +} + +fn trim_chunk_vec(chunks: Vec, offset: usize, length: usize) -> std::io::Result> { + let mut skip = offset; + let mut remaining = length; + let mut result = Vec::new(); + + for chunk in chunks { + if remaining == 0 { + break; + } + + let chunk_len = chunk.len(); + if skip >= chunk_len { + skip -= chunk_len; + continue; + } + + let start = skip; + let take = (chunk_len - start).min(remaining); + result.push(chunk.slice(start, take)?); + remaining -= take; + skip = 0; + } + + Ok(result) +} + +fn decode_bitrot_chunk_source( + source_chunks: &[IoChunk], + shard_size: usize, + checksum_algo: HashAlgorithm, + skip_verify: bool, +) -> std::io::Result<(Vec, bool)> { + let hash_size = checksum_algo.size(); + let mut cursor = ChunkCursor::new(source_chunks); + let mut result = Vec::new(); + let mut copied = false; + + while cursor.remaining() > 0 { + let expected_hash = if hash_size > 0 { + Some(cursor.take_span(hash_size)?) + } else { + None + }; + + let data_len = shard_size.min(cursor.remaining()); + if data_len == 0 { + break; + } + + let data_span = cursor.take_span(data_len)?; + copied |= data_span.copied; + if let Some(expected_hash) = expected_hash { + copied |= expected_hash.copied; + if !skip_verify && checksum_algo.hash_encode(data_span.bytes.as_ref()).as_ref() != expected_hash.bytes.as_ref() { + return Err(std::io::Error::new(std::io::ErrorKind::InvalidData, "bitrot hash mismatch")); + } + } + + result.push(data_span.chunk); + } + + Ok((result, copied)) +} + +#[doc(hidden)] +pub fn decode_bitrot_chunk_source_for_bench( + source_chunks: &[IoChunk], + shard_size: usize, + checksum_algo: HashAlgorithm, + skip_verify: bool, +) -> std::io::Result<(Vec, bool)> { + decode_bitrot_chunk_source(source_chunks, shard_size, checksum_algo, skip_verify) +} + /// Create a new BitrotWriterWrapper based on the provided parameters /// /// # Parameters @@ -176,6 +705,7 @@ pub async fn create_bitrot_writer( #[cfg(test)] mod tests { use super::*; + use futures_util::StreamExt; #[tokio::test] async fn test_create_bitrot_reader_with_inline_data() { @@ -226,6 +756,246 @@ mod tests { assert!(result.unwrap().is_some()); } + #[tokio::test] + async fn test_create_bitrot_chunk_stream_with_inline_data() { + let shard_size = 4; + let checksum_algo = HashAlgorithm::HighwayHash256S; + let shard1 = b"abcd"; + let shard2 = b"ef"; + + let mut encoded = Vec::new(); + encoded.extend_from_slice(checksum_algo.hash_encode(shard1).as_ref()); + encoded.extend_from_slice(shard1); + encoded.extend_from_slice(checksum_algo.hash_encode(shard2).as_ref()); + encoded.extend_from_slice(shard2); + + let mut stream = create_bitrot_chunk_stream( + Some(&encoded), + None, + "test-bucket", + "test-path", + 0, + shard1.len() + shard2.len(), + shard1.len() + shard2.len(), + shard_size, + checksum_algo, + false, + false, + ) + .await + .unwrap() + .unwrap() + .stream; + + let mut collected = Vec::new(); + while let Some(chunk) = stream.next().await { + collected.extend_from_slice(&chunk.unwrap().as_bytes()); + } + + assert_eq!(collected, b"abcdef"); + } + + #[tokio::test] + async fn test_create_bitrot_chunk_stream_detects_hash_mismatch() { + let shard_size = 4; + let checksum_algo = HashAlgorithm::HighwayHash256S; + let shard = b"abcd"; + + let mut encoded = Vec::new(); + let mut bad_hash = checksum_algo.hash_encode(shard).as_ref().to_vec(); + bad_hash[0] ^= 0xFF; + encoded.extend_from_slice(&bad_hash); + encoded.extend_from_slice(shard); + + let result = create_bitrot_chunk_stream( + Some(&encoded), + None, + "test-bucket", + "test-path", + 0, + shard.len(), + shard.len(), + shard_size, + checksum_algo, + false, + false, + ) + .await; + + let mut stream = result.unwrap().unwrap().stream; + let err = stream.next().await.unwrap().unwrap_err(); + assert_eq!(err.kind(), std::io::ErrorKind::InvalidData); + assert!(err.to_string().contains("bitrot hash mismatch")); + } + + #[tokio::test] + async fn test_create_bitrot_chunk_stream_trims_range_after_decode() { + let shard_size = 4; + let checksum_algo = HashAlgorithm::HighwayHash256S; + let shard1 = b"abcd"; + let shard2 = b"efgh"; + + let mut encoded = Vec::new(); + encoded.extend_from_slice(checksum_algo.hash_encode(shard1).as_ref()); + encoded.extend_from_slice(shard1); + encoded.extend_from_slice(checksum_algo.hash_encode(shard2).as_ref()); + encoded.extend_from_slice(shard2); + + let mut stream = create_bitrot_chunk_stream( + Some(&encoded), + None, + "test-bucket", + "test-path", + 1, + 5, + shard1.len() + shard2.len(), + shard_size, + checksum_algo, + false, + false, + ) + .await + .unwrap() + .unwrap() + .stream; + + let mut collected = Vec::new(); + while let Some(chunk) = stream.next().await { + collected.extend_from_slice(&chunk.unwrap().as_bytes()); + } + + assert_eq!(collected, b"bcdef"); + } + + #[test] + fn test_decode_bitrot_chunk_source_preserves_aligned_multi_chunk_slices() { + let shard_size = 4; + let checksum_algo = HashAlgorithm::Md5; + let shard1 = b"abcd"; + let shard2 = b"efgh"; + + let mut encoded_chunk_one = Vec::new(); + encoded_chunk_one.extend_from_slice(checksum_algo.hash_encode(shard1).as_ref()); + encoded_chunk_one.extend_from_slice(shard1); + + let mut encoded_chunk_two = Vec::new(); + encoded_chunk_two.extend_from_slice(checksum_algo.hash_encode(shard2).as_ref()); + encoded_chunk_two.extend_from_slice(shard2); + + let source_chunks = vec![ + IoChunk::Shared(Bytes::from(encoded_chunk_one)), + IoChunk::Shared(Bytes::from(encoded_chunk_two)), + ]; + let (decoded, copied) = decode_bitrot_chunk_source(&source_chunks, shard_size, checksum_algo, false).unwrap(); + + assert!(!copied, "frame-aligned multi-chunk source should not require aggregate copies"); + assert_eq!(decoded.len(), 2); + assert_eq!(decoded[0].as_bytes(), Bytes::from_static(b"abcd")); + assert_eq!(decoded[1].as_bytes(), Bytes::from_static(b"efgh")); + } + + #[test] + fn test_decode_bitrot_chunk_source_marks_cross_chunk_frame_as_copied() { + let shard_size = 4; + let checksum_algo = HashAlgorithm::Md5; + let shard1 = b"abcd"; + let shard2 = b"efgh"; + + let hash1 = checksum_algo.hash_encode(shard1).as_ref().to_vec(); + let hash2 = checksum_algo.hash_encode(shard2).as_ref().to_vec(); + let mut encoded = Vec::new(); + encoded.extend_from_slice(&hash1); + encoded.extend_from_slice(shard1); + encoded.extend_from_slice(&hash2); + encoded.extend_from_slice(shard2); + + let split = hash1.len() + 2; + let source_chunks = vec![ + IoChunk::Shared(Bytes::copy_from_slice(&encoded[..split])), + IoChunk::Shared(Bytes::copy_from_slice(&encoded[split..])), + ]; + let (decoded, copied) = decode_bitrot_chunk_source(&source_chunks, shard_size, checksum_algo, false).unwrap(); + + assert!(copied, "cross-chunk frame should be classified as requiring a copy"); + assert_eq!(decoded.len(), 2); + assert_eq!(decoded[0].as_bytes(), Bytes::from_static(b"abcd")); + assert_eq!(decoded[1].as_bytes(), Bytes::from_static(b"efgh")); + } + + #[test] + fn test_decode_bitrot_chunk_source_preserves_pooled_single_chunk_slice() { + let shard_size = 4; + let checksum_algo = HashAlgorithm::Md5; + let shard = b"abcd"; + + let mut encoded = Vec::new(); + encoded.extend_from_slice(checksum_algo.hash_encode(shard).as_ref()); + encoded.extend_from_slice(shard); + + let source_chunks = vec![IoChunk::Pooled(rustfs_io_core::PooledChunk::from_vec(encoded))]; + let (decoded, copied) = decode_bitrot_chunk_source(&source_chunks, shard_size, checksum_algo, false).unwrap(); + + assert!(!copied, "single pooled chunk slice should preserve provenance without copy"); + assert_eq!(decoded.len(), 1); + assert!(matches!(&decoded[0], IoChunk::Pooled(_))); + assert_eq!(decoded[0].as_bytes(), Bytes::from_static(b"abcd")); + } + + #[tokio::test] + async fn test_bitrot_chunk_source_marks_cross_chunk_take_as_copied() { + let source_stream: BoxChunkStream = Box::pin(stream::iter(vec![ + Ok(IoChunk::Shared(Bytes::from_static(b"ab"))), + Ok(IoChunk::Shared(Bytes::from_static(b"cd"))), + ])); + let mut source = BitrotChunkSource::new(source_stream, VecDeque::new(), false); + + source.fill(4).await.expect("source fill should succeed"); + let span = source.take_span(4).expect("cross-chunk take should succeed"); + + assert!(span.copied, "cross-chunk take should be classified as copied"); + assert_eq!(span.bytes, Bytes::from_static(b"abcd")); + assert_eq!(span.chunk.as_bytes(), Bytes::from_static(b"abcd")); + } + + #[tokio::test] + async fn test_bitrot_chunk_stream_state_yields_verified_prefix_before_later_truncation() { + let shard_size = 4; + let checksum_algo = HashAlgorithm::Md5; + let shard1 = b"abcd"; + let shard2 = b"efgh"; + + let mut first_frame = Vec::new(); + first_frame.extend_from_slice(checksum_algo.hash_encode(shard1).as_ref()); + first_frame.extend_from_slice(shard1); + + let mut second_frame_prefix = Vec::new(); + second_frame_prefix.extend_from_slice(checksum_algo.hash_encode(shard2).as_ref()); + second_frame_prefix.extend_from_slice(&shard2[..2]); + + let source_stream: BoxChunkStream = Box::pin(stream::iter(vec![ + Ok(IoChunk::Shared(Bytes::from(first_frame))), + Ok(IoChunk::Shared(Bytes::from(second_frame_prefix))), + ])); + let mut state = BitrotChunkStreamState::new( + source_stream, + VecDeque::new(), + false, + shard1.len() + shard2.len(), + 0, + shard1.len() + shard2.len(), + shard_size, + checksum_algo, + false, + ); + + let first = state.next_verified_chunk().await.unwrap().unwrap(); + assert_eq!(first.as_bytes(), Bytes::from_static(b"abcd")); + + let err = state.next_verified_chunk().await.unwrap_err(); + assert_eq!(err.kind(), std::io::ErrorKind::UnexpectedEof); + assert!(err.to_string().contains("truncated bitrot chunk source")); + } + #[tokio::test] async fn test_create_bitrot_reader_with_inline_offset_starts_at_requested_shard() { let shard_size = 4; diff --git a/crates/ecstore/src/bucket/migration.rs b/crates/ecstore/src/bucket/migration.rs index bf7809bc8..15d01a45b 100644 --- a/crates/ecstore/src/bucket/migration.rs +++ b/crates/ecstore/src/bucket/migration.rs @@ -17,7 +17,7 @@ use crate::bucket::metadata::BUCKET_METADATA_FILE; use crate::bucket::replication::{decode_resync_file, encode_resync_file}; use crate::disk::{BUCKET_META_PREFIX, MIGRATING_META_BUCKET, RUSTFS_META_BUCKET}; -use crate::store_api::{BucketOptions, ObjectOptions, PutObjReader, StorageAPI}; +use crate::store_api::{BucketOptions, ChunkNativePutData, ObjectOptions, StorageAPI}; use http::HeaderMap; use rustfs_policy::auth::UserIdentity; use rustfs_policy::policy::PolicyDoc; @@ -263,10 +263,8 @@ async fn migrate_one_if_missing( } }; - if let Err(e) = store - .put_object(RUSTFS_META_BUCKET, path, &mut PutObjReader::from_vec(data), opts) - .await - { + let mut put_data = ChunkNativePutData::from_vec(data); + if let Err(e) = store.put_object(RUSTFS_META_BUCKET, path, &mut put_data, opts).await { warn!("write {label}: {e}"); } else { info!("Migrated {label}"); @@ -343,10 +341,8 @@ pub async fn try_migrate_iam_config(store: Arc) { continue; } }; - if let Err(e) = store - .put_object(RUSTFS_META_BUCKET, path, &mut PutObjReader::from_vec(data), &opts) - .await - { + let mut put_data = ChunkNativePutData::from_vec(data); + if let Err(e) = store.put_object(RUSTFS_META_BUCKET, path, &mut put_data, &opts).await { warn!("write IAM config {path}: {e}"); } else { info!("Migrated IAM config: {path}"); diff --git a/crates/ecstore/src/config/com.rs b/crates/ecstore/src/config/com.rs index 12a534cbd..f2e48c2fe 100644 --- a/crates/ecstore/src/config/com.rs +++ b/crates/ecstore/src/config/com.rs @@ -16,7 +16,7 @@ use crate::config::{Config, GLOBAL_STORAGE_CLASS, KVS, audit, notify, oidc, stor use crate::disk::{MIGRATING_META_BUCKET, RUSTFS_META_BUCKET}; use crate::error::{Error, Result}; use crate::global::is_first_cluster_node_local; -use crate::store_api::{ObjectInfo, ObjectOptions, PutObjReader, StorageAPI}; +use crate::store_api::{ChunkNativePutData, ObjectInfo, ObjectOptions, StorageAPI}; use http::HeaderMap; use rustfs_config::audit::{AUDIT_MQTT_KEYS, AUDIT_MQTT_SUB_SYS, AUDIT_WEBHOOK_KEYS, AUDIT_WEBHOOK_SUB_SYS}; use rustfs_config::notify::{NOTIFY_MQTT_KEYS, NOTIFY_MQTT_SUB_SYS, NOTIFY_WEBHOOK_KEYS, NOTIFY_WEBHOOK_SUB_SYS}; @@ -128,10 +128,8 @@ pub async fn delete_config(api: Arc, file: &str) -> Result<()> } pub async fn save_config_with_opts(api: Arc, file: &str, data: Vec, opts: &ObjectOptions) -> Result<()> { - if let Err(err) = api - .put_object(RUSTFS_META_BUCKET, file, &mut PutObjReader::from_vec(data), opts) - .await - { + let mut put_data = ChunkNativePutData::from_vec(data); + if let Err(err) = api.put_object(RUSTFS_META_BUCKET, file, &mut put_data, opts).await { error!("save_config_with_opts: err: {:?}, file: {}", err, file); return Err(err); } @@ -1078,10 +1076,10 @@ mod tests { use crate::global::{is_dist_erasure, is_erasure, is_erasure_sd, update_erasure_type}; use crate::set_disk::SetDisks; use crate::store_api::{ - BucketInfo, BucketOperations, BucketOptions, CompletePart, DeleteBucketOptions, DeletedObject, GetObjectReader, - HTTPRangeSpec, HealOperations, ListMultipartsInfo, ListObjectVersionsInfo, ListObjectsV2Info, ListOperations, - MakeBucketOptions, MultipartInfo, MultipartOperations, MultipartUploadResult, ObjectIO, ObjectInfo, ObjectOperations, - ObjectOptions, ObjectToDelete, PartInfo, PutObjReader, StorageAPI, WalkOptions, + BucketInfo, BucketOperations, BucketOptions, ChunkNativePutData, CompletePart, DeleteBucketOptions, DeletedObject, + GetObjectReader, HTTPRangeSpec, HealOperations, ListMultipartsInfo, ListObjectVersionsInfo, ListObjectsV2Info, + ListOperations, MakeBucketOptions, MultipartInfo, MultipartOperations, MultipartUploadResult, ObjectIO, ObjectInfo, + ObjectOperations, ObjectOptions, ObjectToDelete, PartInfo, StorageAPI, WalkOptions, }; use http::HeaderMap; use rustfs_config::audit::{AUDIT_MQTT_SUB_SYS, AUDIT_WEBHOOK_SUB_SYS}; @@ -1304,7 +1302,7 @@ mod tests { &self, _bucket: &str, _object: &str, - _data: &mut PutObjReader, + _data: &mut ChunkNativePutData, _opts: &ObjectOptions, ) -> Result { panic!("unused in test") @@ -1491,7 +1489,7 @@ mod tests { _object: &str, _upload_id: &str, _part_id: usize, - _data: &mut PutObjReader, + _data: &mut ChunkNativePutData, _opts: &ObjectOptions, ) -> Result { panic!("unused in test") diff --git a/crates/ecstore/src/config/storageclass.rs b/crates/ecstore/src/config/storageclass.rs index 5fe70f4a0..2f6f92147 100644 --- a/crates/ecstore/src/config/storageclass.rs +++ b/crates/ecstore/src/config/storageclass.rs @@ -150,18 +150,7 @@ impl Config { return false; } - let shard_size = shard_size as usize; - - let mut inline_block = DEFAULT_INLINE_BLOCK; - if self.initialized { - inline_block = self.inline_block; - } - - if versioned { - shard_size <= inline_block / 8 - } else { - shard_size <= inline_block - } + shard_size as usize <= self.inline_shard_limit_bytes(versioned) } pub fn inline_block(&self) -> usize { @@ -172,6 +161,15 @@ impl Config { } } + pub fn inline_shard_limit_bytes(&self, versioned: bool) -> usize { + let inline_block = self.inline_block(); + if versioned { inline_block / 8 } else { inline_block } + } + + pub fn inline_object_limit_bytes(&self, data_shards: usize, versioned: bool) -> usize { + self.inline_shard_limit_bytes(versioned).saturating_mul(data_shards.max(1)) + } + pub fn capacity_optimized(&self) -> bool { if !self.initialized { false @@ -336,3 +334,32 @@ pub fn validate_parity_inner(ss_parity: usize, rrs_parity: usize, set_drive_coun } Ok(()) } + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn inline_object_limit_matches_default_non_versioned_budget() { + let cfg = Config { + initialized: true, + inline_block: DEFAULT_INLINE_BLOCK, + ..Default::default() + }; + + assert_eq!(cfg.inline_shard_limit_bytes(false), DEFAULT_INLINE_BLOCK); + assert_eq!(cfg.inline_object_limit_bytes(8, false), DEFAULT_INLINE_BLOCK * 8); + } + + #[test] + fn inline_object_limit_scales_down_for_versioned_objects() { + let cfg = Config { + initialized: true, + inline_block: DEFAULT_INLINE_BLOCK, + ..Default::default() + }; + + assert_eq!(cfg.inline_shard_limit_bytes(true), DEFAULT_INLINE_BLOCK / 8); + assert_eq!(cfg.inline_object_limit_bytes(8, true), DEFAULT_INLINE_BLOCK); + } +} diff --git a/crates/ecstore/src/data_movement.rs b/crates/ecstore/src/data_movement.rs index f40840624..7d2ef7d5f 100644 --- a/crates/ecstore/src/data_movement.rs +++ b/crates/ecstore/src/data_movement.rs @@ -14,9 +14,11 @@ use crate::error::{Error, Result}; use crate::store::ECStore; -use crate::store_api::{CompletePart, GetObjectReader, MultipartOperations, ObjectIO, ObjectInfo, ObjectOptions, PutObjReader}; +use crate::store_api::{ + ChunkNativePutData, CompletePart, GetObjectReader, MultipartOperations, ObjectIO, ObjectInfo, ObjectOptions, +}; use bytes::Bytes; -use rustfs_rio::{EtagResolvable, HashReader, HashReaderDetector, Index, TryGetIndex}; +use rustfs_rio::{BlockReadable, BoxReadBlockFuture, EtagResolvable, HashReader, HashReaderDetector, Index, TryGetIndex}; use std::io::Cursor; use std::pin::Pin; use std::sync::{ @@ -54,6 +56,11 @@ impl TryGetIndex for IndexedDataMovementRead } } +impl BlockReadable for IndexedDataMovementReader { + fn read_block<'a>(&'a mut self, buf: &'a mut [u8]) -> BoxReadBlockFuture<'a> { + Box::pin(rustfs_utils::read_full(self, buf)) + } +} pub fn decode_part_index(index: Option<&Bytes>) -> Option { let bytes = index?; let mut decoded = Index::new(); @@ -64,7 +71,7 @@ pub fn decode_part_index(index: Option<&Bytes>) -> Option { } } -pub fn put_obj_reader_from_chunk(chunk: Vec, size: i64, actual_size: i64, index: Option) -> Result { +pub fn put_data_from_chunk(chunk: Vec, size: i64, actual_size: i64, index: Option) -> Result { use sha2::{Digest, Sha256}; let sha256hex = if !chunk.is_empty() { @@ -74,8 +81,8 @@ pub fn put_obj_reader_from_chunk(chunk: Vec, size: i64, actual_size: i64, in }; 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)) + let hash_reader = HashReader::from_reader(reader, size, actual_size, None, sha256hex, false)?; + Ok(ChunkNativePutData::new(hash_reader)) } pub fn new_multipart_abort_flag() -> Arc { @@ -172,7 +179,7 @@ pub(crate) async fn migrate_object( let part_size = i64::try_from(part.size).map_err(|_| Error::other("part size overflow"))?; let part_actual_size = if part.actual_size > 0 { part.actual_size } else { part_size }; let index = decode_part_index(part.index.as_ref()); - let mut data = put_obj_reader_from_chunk(chunk, part_size, part_actual_size, index)?; + let mut data = put_data_from_chunk(chunk, part_size, part_actual_size, index)?; let pi = match store .put_object_part( @@ -254,8 +261,8 @@ pub(crate) async fn migrate_object( .first() .and_then(|part| decode_part_index(part.index.as_ref())); 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); + let hrd = HashReader::from_reader(reader, object_info.size, actual_size, object_info.etag.clone(), None, false)?; + let mut data = ChunkNativePutData::new(hrd); if let Err(err) = store .put_object( diff --git a/crates/ecstore/src/disk/disk_store.rs b/crates/ecstore/src/disk/disk_store.rs index 958d2c07a..347e79c72 100644 --- a/crates/ecstore/src/disk/disk_store.rs +++ b/crates/ecstore/src/disk/disk_store.rs @@ -21,6 +21,7 @@ use crate::disk::{ use crate::global::GLOBAL_LOCAL_DISK_ID_MAP; use bytes::Bytes; use rustfs_filemeta::{FileInfo, ObjectPartInfo, RawFileInfo}; +use rustfs_io_core::BoxChunkStream; use std::{ path::PathBuf, sync::{ @@ -738,6 +739,14 @@ impl DiskAPI for LocalDiskWrapper { .await } + async fn read_file_chunks(&self, volume: &str, path: &str, offset: usize, length: usize) -> Result { + self.track_disk_health( + || async { self.disk.read_file_chunks(volume, path, offset, length).await }, + get_max_timeout_duration(), + ) + .await + } + async fn append_file(&self, volume: &str, path: &str) -> Result { self.track_disk_health(|| async { self.disk.append_file(volume, path).await }, Duration::ZERO) .await diff --git a/crates/ecstore/src/disk/local.rs b/crates/ecstore/src/disk/local.rs index 868ab6cd6..d023a121c 100644 --- a/crates/ecstore/src/disk/local.rs +++ b/crates/ecstore/src/disk/local.rs @@ -30,12 +30,19 @@ use crate::disk::{ }; use crate::erasure_coding::bitrot_verify; use crate::global::{GLOBAL_IsErasureSD, GLOBAL_RootDiskThreshold}; -use bytes::Bytes; +use bytes::{Bytes, BytesMut}; +use futures_util::{StreamExt, stream}; use parking_lot::RwLock as ParkingLotRwLock; +use rustfs_config::{ + DEFAULT_OBJECT_ZERO_COPY_ENABLE, DEFAULT_OBJECT_ZERO_COPY_MAX_ACTIVE_MMAP_BYTES, DEFAULT_OBJECT_ZERO_COPY_MMAP_WINDOW_BYTES, + DEFAULT_OBJECT_ZERO_COPY_MODE, ENV_OBJECT_ZERO_COPY_ENABLE, ENV_OBJECT_ZERO_COPY_MAX_ACTIVE_MMAP_BYTES, + ENV_OBJECT_ZERO_COPY_MMAP_WINDOW_BYTES, ENV_OBJECT_ZERO_COPY_MODE, +}; use rustfs_filemeta::{ Cache, FileInfo, FileInfoOpts, FileMeta, MetaCacheEntry, MetacacheWriter, ObjectPartInfo, Opts, RawFileInfo, UpdateFn, get_file_info, read_xl_meta_no_data, }; +use rustfs_io_core::{BoxChunkStream, BytesPool, IoChunk, MappedChunk, PooledChunk}; use rustfs_utils::HashAlgorithm; use rustfs_utils::os::get_info; use rustfs_utils::path::{ @@ -46,7 +53,7 @@ use std::collections::HashMap; use std::collections::HashSet; use std::fmt::Debug; use std::io::SeekFrom; -use std::sync::atomic::{AtomicU32, Ordering}; +use std::sync::atomic::{AtomicU32, AtomicUsize, Ordering}; use std::sync::{Arc, OnceLock}; use std::time::Duration; use std::{ @@ -61,6 +68,11 @@ use tokio::time::interval; use tracing::{debug, error, info, warn}; use uuid::Uuid; +#[cfg(test)] +use serial_test::serial; +#[cfg(test)] +use temp_env::with_var; + #[derive(Debug, Clone)] pub struct FormatInfo { pub id: Option, @@ -97,6 +109,410 @@ pub struct LocalDisk { exit_signal: Option>, } +const LOCAL_CHUNK_FAST_PATH_MIN_BYTES: usize = 64 * 1024; +const LOCAL_DISK_POOLED_SOURCE_FALLBACK: &str = "fallback"; +const LOCAL_DISK_POOLED_SOURCE_COMPAT_COLLECT: &str = "compat_collect"; +const LOCAL_DISK_POOLED_SOURCE_COMPAT_DIRECT: &str = "compat_direct"; +const ACTIVE_MMAP_WINDOW_BUDGET_EXCEEDED_MESSAGE: &str = "active mmap window budget exceeded"; + +#[cfg(unix)] +const LOCAL_CHUNK_COMPAT_MAX_MAPPED_WINDOWS: usize = 1; + +static LOCAL_CHUNK_FALLBACK_POOL: OnceLock = OnceLock::new(); + +fn local_chunk_fallback_pool() -> &'static BytesPool { + LOCAL_CHUNK_FALLBACK_POOL.get_or_init(BytesPool::new_tiered) +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +enum LocalChunkZeroCopyMode { + Off, + Conservative, + Balanced, + Aggressive, +} + +impl LocalChunkZeroCopyMode { + fn from_env() -> Self { + match rustfs_utils::get_env_str(ENV_OBJECT_ZERO_COPY_MODE, DEFAULT_OBJECT_ZERO_COPY_MODE) + .trim() + .to_ascii_lowercase() + .as_str() + { + "off" => Self::Off, + "conservative" => Self::Conservative, + "aggressive" => Self::Aggressive, + _ => Self::Balanced, + } + } + + fn effective() -> Self { + if !rustfs_utils::get_env_bool(ENV_OBJECT_ZERO_COPY_ENABLE, DEFAULT_OBJECT_ZERO_COPY_ENABLE) { + return Self::Off; + } + + Self::from_env() + } + + const fn fast_path_min_bytes(self) -> usize { + match self { + Self::Aggressive => 1, + Self::Off | Self::Conservative | Self::Balanced => LOCAL_CHUNK_FAST_PATH_MIN_BYTES, + } + } + + const fn allows_multi_window(self) -> bool { + matches!(self, Self::Balanced | Self::Aggressive) + } + + const fn is_disabled(self) -> bool { + matches!(self, Self::Off) + } +} + +#[cfg(unix)] +static ACTIVE_LOCAL_MMAP_BYTES: AtomicUsize = AtomicUsize::new(0); + +#[cfg(unix)] +#[derive(Debug)] +struct ActiveMmapWindow { + mmap: memmap2::Mmap, + accounted_len: usize, +} + +#[cfg(unix)] +impl AsRef<[u8]> for ActiveMmapWindow { + fn as_ref(&self) -> &[u8] { + &self.mmap[..] + } +} + +#[cfg(unix)] +impl Drop for ActiveMmapWindow { + fn drop(&mut self) { + let remaining = ACTIVE_LOCAL_MMAP_BYTES + .fetch_sub(self.accounted_len, Ordering::AcqRel) + .saturating_sub(self.accounted_len); + rustfs_io_metrics::record_local_disk_active_mmap_bytes(remaining); + } +} + +#[cfg(unix)] +#[allow(unsafe_code)] +fn mmap_page_size() -> usize { + static PAGE_SIZE: OnceLock = OnceLock::new(); + + *PAGE_SIZE.get_or_init(|| { + let page_size = unsafe { libc::sysconf(libc::_SC_PAGESIZE) }; + if page_size <= 0 { 4096 } else { page_size as usize } + }) +} + +#[cfg(unix)] +fn configured_local_chunk_window_bytes() -> usize { + let page_size = mmap_page_size(); + rustfs_utils::get_env_usize(ENV_OBJECT_ZERO_COPY_MMAP_WINDOW_BYTES, DEFAULT_OBJECT_ZERO_COPY_MMAP_WINDOW_BYTES) + .max(page_size) + .div_ceil(page_size) + * page_size +} + +#[cfg(unix)] +fn configured_local_chunk_max_active_mmap_bytes() -> usize { + rustfs_utils::get_env_usize(ENV_OBJECT_ZERO_COPY_MAX_ACTIVE_MMAP_BYTES, DEFAULT_OBJECT_ZERO_COPY_MAX_ACTIVE_MMAP_BYTES) + .max(configured_local_chunk_window_bytes()) +} + +#[cfg(unix)] +fn should_prefer_pooled_zero_copy_compat(mode: LocalChunkZeroCopyMode, length: usize, window_bytes: usize) -> bool { + if mode.is_disabled() || !mode.allows_multi_window() || length < mode.fast_path_min_bytes() || window_bytes == 0 { + return false; + } + + length.div_ceil(window_bytes) > LOCAL_CHUNK_COMPAT_MAX_MAPPED_WINDOWS +} + +fn fallback_reason_for_local_mmap_error(err: &DiskError) -> rustfs_io_metrics::FallbackReason { + match err { + DiskError::Io(io_error) if io_error.to_string().contains(ACTIVE_MMAP_WINDOW_BUDGET_EXCEEDED_MESSAGE) => { + rustfs_io_metrics::FallbackReason::WindowLimitExceeded + } + _ => rustfs_io_metrics::FallbackReason::MmapUnavailable, + } +} + +#[cfg(unix)] +fn try_reserve_active_mmap_bytes(accounted_len: usize, max_active_bytes: usize) -> bool { + loop { + let current = ACTIVE_LOCAL_MMAP_BYTES.load(Ordering::Acquire); + let Some(next) = current.checked_add(accounted_len) else { + return false; + }; + if next > max_active_bytes { + return false; + } + + if ACTIVE_LOCAL_MMAP_BYTES + .compare_exchange_weak(current, next, Ordering::AcqRel, Ordering::Acquire) + .is_ok() + { + rustfs_io_metrics::record_local_disk_active_mmap_bytes(next); + return true; + } + } +} + +#[cfg(unix)] +#[allow(unsafe_code)] +fn map_file_region_bytes(file_path: &Path, offset: usize, length: usize, max_active_bytes: usize) -> Result { + use memmap2::MmapOptions; + + let aligned_offset = offset / mmap_page_size() * mmap_page_size(); + let logical_offset = offset - aligned_offset; + let map_length = logical_offset.checked_add(length).ok_or(DiskError::FileCorrupt)?; + let visible_end = logical_offset.checked_add(length).ok_or(DiskError::FileCorrupt)?; + if !try_reserve_active_mmap_bytes(map_length, max_active_bytes) { + return Err(DiskError::other(ACTIVE_MMAP_WINDOW_BUDGET_EXCEEDED_MESSAGE)); + } + let file = std::fs::File::open(file_path).map_err(DiskError::from)?; + + let mmap_result = + unsafe { MmapOptions::new().offset(aligned_offset as u64).len(map_length).map(&file) }.map_err(DiskError::other); + let mmap = match mmap_result { + Ok(mmap) => mmap, + Err(err) => { + let remaining = ACTIVE_LOCAL_MMAP_BYTES + .fetch_sub(map_length, Ordering::AcqRel) + .saturating_sub(map_length); + rustfs_io_metrics::record_local_disk_active_mmap_bytes(remaining); + return Err(err); + } + }; + let bytes = Bytes::from_owner(ActiveMmapWindow { + mmap, + accounted_len: map_length, + }); + + Ok(bytes.slice(logical_offset..visible_end)) +} + +#[cfg(unix)] +#[allow(unsafe_code)] +fn map_file_region_chunk(file_path: &Path, offset: usize, length: usize, max_active_bytes: usize) -> Result { + let bytes = map_file_region_bytes(file_path, offset, length, max_active_bytes)?; + MappedChunk::new(bytes, 0, length).map_err(DiskError::other) +} + +#[cfg(unix)] +#[derive(Debug)] +struct LocalMappedChunkStreamState { + file_path: PathBuf, + next_offset: usize, + remaining: usize, + window_bytes: usize, +} + +async fn read_file_pooled_chunk_from_path(file_path: PathBuf, offset: usize, length: usize) -> std::io::Result { + read_file_pooled_chunk_from_path_with_source(file_path, offset, length, LOCAL_DISK_POOLED_SOURCE_FALLBACK).await +} + +async fn read_file_pooled_chunk_from_path_with_source( + file_path: PathBuf, + offset: usize, + length: usize, + metric_source: &'static str, +) -> std::io::Result { + let mut file = File::open(file_path).await?; + if offset > 0 { + file.seek(SeekFrom::Start(offset as u64)).await?; + } + + let mut buffer = local_chunk_fallback_pool().acquire_buffer(length).await; + buffer.resize(length, 0); + file.read_exact(&mut buffer[..length]).await?; + rustfs_io_metrics::record_local_disk_pooled_chunk(metric_source, length); + Ok(IoChunk::Pooled(PooledChunk::new(buffer, length).map_err(std::io::Error::other)?)) +} + +async fn prepare_read_file_request(disk: &LocalDisk, volume: &str, path: &str) -> Result<(PathBuf, PathBuf, Metadata)> { + let volume_dir = disk.get_bucket_path(volume)?; + if !skip_access_checks(volume) { + access(&volume_dir) + .await + .map_err(|e| to_access_error(e, DiskError::VolumeAccessDenied))?; + } + + let file_path = disk.get_object_path(volume, path)?; + check_path_length(file_path.to_string_lossy().as_ref())?; + + let file_path_clone = file_path.clone(); + let meta = tokio::task::spawn_blocking(move || std::fs::metadata(&file_path_clone).map_err(DiskError::from)) + .await + .map_err(DiskError::from)??; + + Ok((volume_dir, file_path, meta)) +} + +fn validate_read_file_bounds(meta: &Metadata, offset: usize, length: usize) -> Result<()> { + let end_offset = offset.checked_add(length).ok_or(DiskError::FileCorrupt)?; + if meta.len() < end_offset as u64 { + error!( + "read_file: file size is less than offset + length {} + {} = {}", + offset, + length, + meta.len() + ); + return Err(DiskError::FileCorrupt); + } + + Ok(()) +} + +#[cfg(unix)] +fn build_lazy_mapped_chunk_stream( + file_path: PathBuf, + offset: usize, + length: usize, + window_bytes: usize, + max_active_bytes: usize, +) -> BoxChunkStream { + let state = LocalMappedChunkStreamState { + file_path, + next_offset: offset, + remaining: length, + window_bytes, + }; + + Box::pin(stream::unfold(Some(state), move |state| async move { + let mut state = match state { + Some(state) => state, + None => return None, + }; + + if state.remaining == 0 { + return None; + } + + let visible_len = state.remaining.min(state.window_bytes); + let window_offset = state.next_offset; + let file_path = state.file_path.clone(); + let mmap_result = + tokio::task::spawn_blocking(move || map_file_region_chunk(&file_path, window_offset, visible_len, max_active_bytes)) + .await; + + match mmap_result { + Ok(Ok(chunk)) => { + state.next_offset += visible_len; + state.remaining -= visible_len; + let next_state = if state.remaining == 0 { None } else { Some(state) }; + Some((Ok(IoChunk::Mapped(chunk)), next_state)) + } + Ok(Err(err)) => { + rustfs_io_metrics::record_io_fallback( + rustfs_io_metrics::IoStage::LocalDiskChunk, + fallback_reason_for_local_mmap_error(&err), + ); + debug!( + error = %err, + offset = window_offset, + len = visible_len, + "local disk lazy mmap window failed, falling back to buffered remainder" + ); + let fallback = + read_file_pooled_chunk_from_path(state.file_path.clone(), state.next_offset, state.remaining).await; + Some((fallback, None)) + } + Err(err) => { + rustfs_io_metrics::record_io_fallback( + rustfs_io_metrics::IoStage::LocalDiskChunk, + rustfs_io_metrics::FallbackReason::MmapUnavailable, + ); + debug!( + error = %err, + offset = window_offset, + len = visible_len, + "local disk lazy mmap task failed, falling back to buffered remainder" + ); + let fallback = + read_file_pooled_chunk_from_path(state.file_path.clone(), state.next_offset, state.remaining).await; + Some((fallback, None)) + } + } + })) +} + +async fn read_file_pooled_chunk_fallback( + disk: &LocalDisk, + volume_dir: &Path, + file_path: PathBuf, + offset: usize, + length: usize, +) -> Result { + read_file_pooled_chunk_fallback_with_source(disk, volume_dir, file_path, offset, length, LOCAL_DISK_POOLED_SOURCE_FALLBACK) + .await +} + +async fn read_file_pooled_chunk_fallback_with_source( + disk: &LocalDisk, + volume_dir: &Path, + file_path: PathBuf, + offset: usize, + length: usize, + metric_source: &'static str, +) -> Result { + let mut f = disk.open_file(file_path, O_RDONLY, volume_dir).await?; + + if offset > 0 { + f.seek(SeekFrom::Start(offset as u64)).await?; + } + + let mut buffer = local_chunk_fallback_pool().acquire_buffer(length).await; + buffer.resize(length, 0); + f.read_exact(&mut buffer[..length]).await?; + rustfs_io_metrics::record_local_disk_pooled_chunk(metric_source, length); + Ok(IoChunk::Pooled(PooledChunk::new(buffer, length).map_err(DiskError::other)?)) +} + +async fn collect_chunk_stream_bytes(mut stream: BoxChunkStream, expected_len: usize) -> Result { + let Some(first) = stream.next().await else { + return Ok(Bytes::new()); + }; + let first = first.map_err(DiskError::from)?; + let first_len = first.len(); + if matches!(first, IoChunk::Pooled(_)) { + rustfs_io_metrics::record_local_disk_pooled_chunk(LOCAL_DISK_POOLED_SOURCE_COMPAT_COLLECT, first_len); + } + let first_bytes = first.as_bytes(); + + let Some(second) = stream.next().await else { + return Ok(first_bytes); + }; + let second = second.map_err(DiskError::from)?; + let second_len = second.len(); + if matches!(second, IoChunk::Pooled(_)) { + rustfs_io_metrics::record_local_disk_pooled_chunk(LOCAL_DISK_POOLED_SOURCE_COMPAT_COLLECT, second_len); + } + let mut chunk_count = 2usize; + let mut total_bytes = first_len + second_len; + let mut buffer = BytesMut::with_capacity(expected_len); + buffer.extend_from_slice(first_bytes.as_ref()); + buffer.extend_from_slice(second.as_bytes().as_ref()); + + while let Some(chunk) = stream.next().await { + let chunk = chunk.map_err(DiskError::from)?; + let chunk_len = chunk.len(); + if matches!(chunk, IoChunk::Pooled(_)) { + rustfs_io_metrics::record_local_disk_pooled_chunk(LOCAL_DISK_POOLED_SOURCE_COMPAT_COLLECT, chunk_len); + } + chunk_count += 1; + total_bytes += chunk_len; + buffer.extend_from_slice(chunk.as_bytes().as_ref()); + } + + rustfs_io_metrics::record_local_disk_compat_collect(chunk_count, total_bytes); + Ok(buffer.freeze()) +} + impl Drop for LocalDisk { fn drop(&mut self) { if let Some(exit_signal) = self.exit_signal.take() { @@ -1835,87 +2251,96 @@ impl DiskAPI for LocalDisk { use std::time::Instant; let start = Instant::now(); - let volume_dir = self.get_bucket_path(volume)?; - if !skip_access_checks(volume) { - access(&volume_dir) - .await - .map_err(|e| to_access_error(e, DiskError::VolumeAccessDenied))?; - } - - let file_path = self.get_object_path(volume, path)?; - check_path_length(file_path.to_string_lossy().as_ref())?; - - // Verify file exists and get metadata - let file_path_clone = file_path.clone(); - let meta = tokio::task::spawn_blocking(move || std::fs::metadata(&file_path_clone).map_err(DiskError::from)) - .await - .map_err(DiskError::from)??; - - let end_offset = offset.checked_add(length).ok_or(DiskError::FileCorrupt)?; - if meta.len() < end_offset as u64 { - error!( - "read_file_zero_copy: file size is less than offset + length {} + {} = {}", - offset, - length, - meta.len() - ); - return Err(DiskError::FileCorrupt); - } - - // Unix: use mmap to read the data (copies into Bytes for safe ownership) - // Non-Unix: fall back to efficient read #[cfg(unix)] { - use memmap2::MmapOptions; - let file_path_clone = file_path.clone(); - let offset_u64 = offset as u64; - - let bytes = tokio::task::spawn_blocking(move || { - let file = std::fs::File::open(&file_path_clone).map_err(DiskError::from)?; - - // Create memory map for the specified region - // SAFETY: The file is opened as read-only, and we're mapping a region - // that we've already verified exists and is within file bounds. - let mmap = unsafe { MmapOptions::new().offset(offset_u64).len(length).map(&file) }.map_err(DiskError::other)?; - - // Copy the mapped region into a Bytes buffer. This avoids undefined - // behavior from treating OS-managed mmap memory as allocator-managed - // Vec storage, at the cost of an extra copy. - Ok::(Bytes::copy_from_slice(&mmap)) - }) - .await - .map_err(DiskError::from)??; - - // Log successful mmap read metrics - let duration_ms = start.elapsed().as_secs_f64() * 1000.0; - - // Record mmap read metrics - rustfs_io_metrics::record_zero_copy_read(length, duration_ms); - - debug!(size = length, duration_ms = duration_ms, "mmap_read_success"); - - return Ok(bytes); + let zero_copy_mode = LocalChunkZeroCopyMode::effective(); + let window_bytes = configured_local_chunk_window_bytes(); + if should_prefer_pooled_zero_copy_compat(zero_copy_mode, length, window_bytes) { + let (volume_dir, file_path, meta) = prepare_read_file_request(self, volume, path).await?; + validate_read_file_bounds(&meta, offset, length)?; + let chunk = read_file_pooled_chunk_fallback_with_source( + self, + &volume_dir, + file_path, + offset, + length, + LOCAL_DISK_POOLED_SOURCE_COMPAT_DIRECT, + ) + .await?; + let bytes = collect_chunk_stream_bytes(Box::pin(stream::iter(vec![Ok(chunk)])), length).await?; + debug!( + size = bytes.len(), + duration_ms = start.elapsed().as_secs_f64() * 1000.0, + "chunk_compat_read_pooled_success" + ); + return Ok(bytes); + } } - // Non-Unix fallback: efficient read into Bytes - #[cfg(not(unix))] + let bytes = collect_chunk_stream_bytes(self.read_file_chunks(volume, path, offset, length).await?, length).await?; + debug!( + size = bytes.len(), + duration_ms = start.elapsed().as_secs_f64() * 1000.0, + "chunk_compat_read_success" + ); + Ok(bytes) + } + + #[allow(unsafe_code)] + #[tracing::instrument(level = "debug", skip(self))] + async fn read_file_chunks(&self, volume: &str, path: &str, offset: usize, length: usize) -> Result { + let (volume_dir, file_path, meta) = prepare_read_file_request(self, volume, path).await?; + validate_read_file_bounds(&meta, offset, length)?; + + let zero_copy_mode = LocalChunkZeroCopyMode::effective(); + if zero_copy_mode.is_disabled() { + rustfs_io_metrics::record_io_fallback( + rustfs_io_metrics::IoStage::LocalDiskChunk, + rustfs_io_metrics::FallbackReason::MmapDisabled, + ); + let chunk = read_file_pooled_chunk_fallback(self, &volume_dir, file_path, offset, length).await?; + return Ok(Box::pin(stream::iter(vec![Ok(chunk)]))); + } + + if length < zero_copy_mode.fast_path_min_bytes() { + rustfs_io_metrics::record_io_fallback( + rustfs_io_metrics::IoStage::LocalDiskChunk, + rustfs_io_metrics::FallbackReason::SmallObject, + ); + let chunk = read_file_pooled_chunk_fallback(self, &volume_dir, file_path, offset, length).await?; + return Ok(Box::pin(stream::iter(vec![Ok(chunk)]))); + } + + #[cfg(unix)] { - // Record zero-copy fallback - rustfs_io_metrics::record_zero_copy_fallback("non_unix_platform"); + let window_bytes = configured_local_chunk_window_bytes(); - debug!(reason = "non_unix_platform", "zero_copy_fallback"); - - let mut f = self.open_file(file_path, O_RDONLY, volume_dir).await?; - - if offset > 0 { - f.seek(SeekFrom::Start(offset as u64)).await?; + if !zero_copy_mode.allows_multi_window() && length > window_bytes { + rustfs_io_metrics::record_io_fallback( + rustfs_io_metrics::IoStage::LocalDiskChunk, + rustfs_io_metrics::FallbackReason::WindowLimitExceeded, + ); + let chunk = read_file_pooled_chunk_fallback(self, &volume_dir, file_path, offset, length).await?; + return Ok(Box::pin(stream::iter(vec![Ok(chunk)]))); } - let mut buffer = Vec::with_capacity(length); - buffer.resize(length, 0); - f.read_exact(&mut buffer).await?; + return Ok(build_lazy_mapped_chunk_stream( + file_path, + offset, + length, + window_bytes, + configured_local_chunk_max_active_mmap_bytes(), + )); + } - Ok(Bytes::from(buffer)) + #[cfg(not(unix))] + { + rustfs_io_metrics::record_io_fallback( + rustfs_io_metrics::IoStage::LocalDiskChunk, + rustfs_io_metrics::FallbackReason::MmapUnavailable, + ); + let chunk = read_file_pooled_chunk_fallback(self, &volume_dir, file_path, offset, length).await?; + Ok(Box::pin(stream::iter(vec![Ok(chunk)]))) } } @@ -2674,6 +3099,7 @@ async fn get_disk_info(drive_path: PathBuf) -> Result<(rustfs_utils::os::DiskInf #[cfg(test)] mod test { use super::*; + use futures_util::StreamExt; #[tokio::test] async fn test_skip_access_checks() { @@ -2960,6 +3386,22 @@ mod test { assert!(matches!(result, Err(DiskError::FileCorrupt))); } + #[tokio::test] + async fn test_read_file_zero_copy_supports_non_zero_offset() { + use tempfile::tempdir; + + let dir = tempdir().unwrap(); + let endpoint = Endpoint::try_from(dir.path().to_str().unwrap()).unwrap(); + let disk = LocalDisk::new(&endpoint, false).await.unwrap(); + + disk.make_volume("test-volume").await.unwrap(); + let content = Bytes::from_static(b"0123456789abcdef"); + disk.write_all("test-volume", "test-file.txt", content.clone()).await.unwrap(); + + let result = disk.read_file_zero_copy("test-volume", "test-file.txt", 3, 7).await.unwrap(); + assert_eq!(result, Bytes::from_static(b"3456789")); + } + #[test] fn test_is_valid_volname() { // Valid volume names (length >= 3) @@ -3057,6 +3499,274 @@ mod test { let _ = fs::remove_file(test_file).await; } + #[tokio::test] + #[serial] + async fn test_read_file_chunks_returns_pooled_chunk_for_local_fallback() { + let dir = tempfile::tempdir().unwrap(); + let bucket = "chunk-bucket"; + let object = "obj.txt"; + let content = b"chunk-data"; + + fs::create_dir_all(dir.path().join(bucket)).await.unwrap(); + fs::write(dir.path().join(bucket).join(object), content).await.unwrap(); + + let endpoint = Endpoint::try_from(dir.path().to_str().unwrap()).unwrap(); + let disk = LocalDisk::new(&endpoint, false).await.unwrap(); + + let mut stream = disk.read_file_chunks(bucket, object, 0, content.len()).await.unwrap(); + let first = stream.next().await.unwrap().unwrap(); + assert!(matches!(first, IoChunk::Pooled(_))); + assert_eq!(first.as_bytes(), Bytes::from_static(content)); + assert!(stream.next().await.is_none()); + } + + #[tokio::test] + #[serial] + async fn test_read_file_chunks_prefers_mapped_chunk_when_eligible() { + let dir = tempfile::tempdir().unwrap(); + let bucket = "chunk-bucket"; + let object = "obj-large.txt"; + let content = vec![7u8; LOCAL_CHUNK_FAST_PATH_MIN_BYTES]; + + fs::create_dir_all(dir.path().join(bucket)).await.unwrap(); + fs::write(dir.path().join(bucket).join(object), &content).await.unwrap(); + + let endpoint = Endpoint::try_from(dir.path().to_str().unwrap()).unwrap(); + let disk = LocalDisk::new(&endpoint, false).await.unwrap(); + + let mut stream = disk.read_file_chunks(bucket, object, 0, content.len()).await.unwrap(); + let first = stream.next().await.unwrap().unwrap(); + #[cfg(unix)] + assert!(matches!(first, IoChunk::Mapped(_))); + #[cfg(not(unix))] + assert!(matches!(first, IoChunk::Pooled(_))); + assert_eq!(first.as_bytes(), Bytes::from(content)); + assert!(stream.next().await.is_none()); + } + + #[tokio::test] + #[serial] + async fn test_read_file_chunks_falls_back_to_pooled_for_small_object() { + let dir = tempfile::tempdir().unwrap(); + let bucket = "chunk-bucket"; + let object = "obj-small.txt"; + let content = b"small-object"; + + fs::create_dir_all(dir.path().join(bucket)).await.unwrap(); + fs::write(dir.path().join(bucket).join(object), content).await.unwrap(); + + let endpoint = Endpoint::try_from(dir.path().to_str().unwrap()).unwrap(); + let disk = LocalDisk::new(&endpoint, false).await.unwrap(); + + let mut stream = disk.read_file_chunks(bucket, object, 0, content.len()).await.unwrap(); + let first = stream.next().await.unwrap().unwrap(); + assert!(matches!(first, IoChunk::Pooled(_))); + assert_eq!(first.as_bytes(), Bytes::from_static(content)); + assert!(stream.next().await.is_none()); + } + + #[tokio::test] + #[serial] + async fn test_read_file_chunks_supports_non_zero_offset() { + let dir = tempfile::tempdir().unwrap(); + let bucket = "chunk-bucket"; + let object = "obj-offset.txt"; + let content = vec![3u8; LOCAL_CHUNK_FAST_PATH_MIN_BYTES + 16]; + + fs::create_dir_all(dir.path().join(bucket)).await.unwrap(); + fs::write(dir.path().join(bucket).join(object), &content).await.unwrap(); + + let endpoint = Endpoint::try_from(dir.path().to_str().unwrap()).unwrap(); + let disk = LocalDisk::new(&endpoint, false).await.unwrap(); + + let mut stream = disk + .read_file_chunks(bucket, object, 1, LOCAL_CHUNK_FAST_PATH_MIN_BYTES) + .await + .unwrap(); + let first = stream.next().await.unwrap().unwrap(); + #[cfg(unix)] + assert!(matches!(first, IoChunk::Mapped(_))); + #[cfg(not(unix))] + assert!(matches!(first, IoChunk::Pooled(_))); + assert_eq!(first.as_bytes(), Bytes::copy_from_slice(&content[1..1 + LOCAL_CHUNK_FAST_PATH_MIN_BYTES])); + assert!(stream.next().await.is_none()); + } + + #[cfg(unix)] + #[tokio::test] + #[serial] + async fn test_read_file_chunks_splits_large_reads_into_multiple_windows() { + let dir = tempfile::tempdir().unwrap(); + let bucket = "chunk-bucket"; + let object = "obj-windowed.txt"; + let content = vec![5u8; DEFAULT_OBJECT_ZERO_COPY_MMAP_WINDOW_BYTES + 32]; + + fs::create_dir_all(dir.path().join(bucket)).await.unwrap(); + fs::write(dir.path().join(bucket).join(object), &content).await.unwrap(); + + let endpoint = Endpoint::try_from(dir.path().to_str().unwrap()).unwrap(); + let disk = LocalDisk::new(&endpoint, false).await.unwrap(); + + let mut stream = disk.read_file_chunks(bucket, object, 0, content.len()).await.unwrap(); + let first = stream.next().await.unwrap().unwrap(); + let second = stream.next().await.unwrap().unwrap(); + + assert!(matches!(first, IoChunk::Mapped(_))); + assert!(matches!(second, IoChunk::Mapped(_))); + assert_eq!(first.len(), DEFAULT_OBJECT_ZERO_COPY_MMAP_WINDOW_BYTES); + assert_eq!(second.len(), 32); + assert_eq!(first.as_bytes(), Bytes::from(vec![5u8; DEFAULT_OBJECT_ZERO_COPY_MMAP_WINDOW_BYTES])); + assert_eq!(second.as_bytes(), Bytes::from(vec![5u8; 32])); + assert!(stream.next().await.is_none()); + } + + #[cfg(unix)] + #[tokio::test] + #[serial] + async fn test_read_file_zero_copy_collects_multi_window_chunk_stream() { + let dir = tempfile::tempdir().unwrap(); + let bucket = "chunk-bucket"; + let object = "obj-zero-copy-compat.txt"; + let content = vec![6u8; DEFAULT_OBJECT_ZERO_COPY_MMAP_WINDOW_BYTES + 48]; + + fs::create_dir_all(dir.path().join(bucket)).await.unwrap(); + fs::write(dir.path().join(bucket).join(object), &content).await.unwrap(); + + let endpoint = Endpoint::try_from(dir.path().to_str().unwrap()).unwrap(); + let disk = LocalDisk::new(&endpoint, false).await.unwrap(); + + let bytes = disk.read_file_zero_copy(bucket, object, 0, content.len()).await.unwrap(); + assert_eq!(bytes, Bytes::from(content)); + } + + #[cfg(unix)] + #[test] + #[serial] + fn test_read_file_chunks_lazy_windows_reuse_single_window_budget() { + let page_size = mmap_page_size(); + let window_bytes = page_size.to_string(); + let max_active_bytes = page_size.to_string(); + + with_var(ENV_OBJECT_ZERO_COPY_ENABLE, Some("true"), || { + with_var(ENV_OBJECT_ZERO_COPY_MODE, Some("aggressive"), || { + with_var(ENV_OBJECT_ZERO_COPY_MMAP_WINDOW_BYTES, Some(window_bytes.clone()), || { + with_var(ENV_OBJECT_ZERO_COPY_MAX_ACTIVE_MMAP_BYTES, Some(max_active_bytes.clone()), || { + ACTIVE_LOCAL_MMAP_BYTES.store(0, Ordering::Release); + + let runtime = tokio::runtime::Runtime::new().unwrap(); + runtime.block_on(async { + let dir = tempfile::tempdir().unwrap(); + let bucket = "chunk-bucket"; + let object = "obj-budgeted.txt"; + let content = vec![9u8; page_size * 2]; + + fs::create_dir_all(dir.path().join(bucket)).await.unwrap(); + fs::write(dir.path().join(bucket).join(object), &content).await.unwrap(); + + let endpoint = Endpoint::try_from(dir.path().to_str().unwrap()).unwrap(); + let disk = LocalDisk::new(&endpoint, false).await.unwrap(); + + let mut stream = disk.read_file_chunks(bucket, object, 0, content.len()).await.unwrap(); + let first = stream.next().await.unwrap().unwrap(); + assert!(matches!(first, IoChunk::Mapped(_))); + assert_eq!(first.len(), page_size); + assert_eq!(ACTIVE_LOCAL_MMAP_BYTES.load(Ordering::Acquire), page_size); + + drop(first); + assert_eq!(ACTIVE_LOCAL_MMAP_BYTES.load(Ordering::Acquire), 0); + + let second = stream.next().await.unwrap().unwrap(); + assert!(matches!(second, IoChunk::Mapped(_))); + assert_eq!(second.len(), page_size); + assert_eq!(ACTIVE_LOCAL_MMAP_BYTES.load(Ordering::Acquire), page_size); + + drop(second); + assert_eq!(ACTIVE_LOCAL_MMAP_BYTES.load(Ordering::Acquire), 0); + assert!(stream.next().await.is_none()); + }); + + assert_eq!(ACTIVE_LOCAL_MMAP_BYTES.load(Ordering::Acquire), 0); + }); + }); + }); + }); + } + + #[test] + #[serial] + fn test_local_chunk_zero_copy_mode_respects_enable_and_mode_env() { + with_var(ENV_OBJECT_ZERO_COPY_ENABLE, Some("false"), || { + assert_eq!(LocalChunkZeroCopyMode::effective(), LocalChunkZeroCopyMode::Off); + }); + + with_var(ENV_OBJECT_ZERO_COPY_ENABLE, Some("true"), || { + with_var(ENV_OBJECT_ZERO_COPY_MODE, Some("aggressive"), || { + assert_eq!(LocalChunkZeroCopyMode::effective(), LocalChunkZeroCopyMode::Aggressive); + }); + }); + } + + #[cfg(unix)] + #[test] + #[serial] + fn test_should_prefer_pooled_zero_copy_compat_for_multi_window_requests() { + let window_bytes = 1024; + let balanced_multi_window_len = LOCAL_CHUNK_FAST_PATH_MIN_BYTES.max(window_bytes * 2); + + assert!(!should_prefer_pooled_zero_copy_compat( + LocalChunkZeroCopyMode::Off, + window_bytes * 2, + window_bytes + )); + assert!(!should_prefer_pooled_zero_copy_compat( + LocalChunkZeroCopyMode::Conservative, + window_bytes * 2, + window_bytes + )); + assert!(!should_prefer_pooled_zero_copy_compat( + LocalChunkZeroCopyMode::Balanced, + window_bytes, + window_bytes + )); + assert!(should_prefer_pooled_zero_copy_compat( + LocalChunkZeroCopyMode::Balanced, + balanced_multi_window_len, + window_bytes + )); + assert!(should_prefer_pooled_zero_copy_compat( + LocalChunkZeroCopyMode::Aggressive, + window_bytes * 2, + window_bytes + )); + } + + #[test] + fn test_fallback_reason_for_local_mmap_error_distinguishes_budget_limit() { + let budget_err = DiskError::other(ACTIVE_MMAP_WINDOW_BUDGET_EXCEEDED_MESSAGE); + let generic_err = DiskError::other("mmap failed"); + + assert_eq!( + fallback_reason_for_local_mmap_error(&budget_err), + rustfs_io_metrics::FallbackReason::WindowLimitExceeded + ); + assert_eq!( + fallback_reason_for_local_mmap_error(&generic_err), + rustfs_io_metrics::FallbackReason::MmapUnavailable + ); + } + + #[cfg(unix)] + #[test] + #[serial] + fn test_configured_local_chunk_window_bytes_aligns_to_page_size() { + with_var(ENV_OBJECT_ZERO_COPY_MMAP_WINDOW_BYTES, Some("12345"), || { + let page_size = mmap_page_size(); + let window_bytes = configured_local_chunk_window_bytes(); + assert!(window_bytes >= 12345); + assert_eq!(window_bytes % page_size, 0); + }); + } + #[test] fn test_is_root_path() { // Unix root path diff --git a/crates/ecstore/src/disk/mod.rs b/crates/ecstore/src/disk/mod.rs index b00a036ce..3bd2857c6 100644 --- a/crates/ecstore/src/disk/mod.rs +++ b/crates/ecstore/src/disk/mod.rs @@ -41,6 +41,7 @@ use error::DiskError; use error::{Error, Result}; use local::LocalDisk; use rustfs_filemeta::{FileInfo, ObjectPartInfo, RawFileInfo}; +use rustfs_io_core::BoxChunkStream; use rustfs_madmin::info_commands::DiskMetrics; use serde::{Deserialize, Serialize}; use std::{fmt::Debug, path::PathBuf, sync::Arc}; @@ -295,6 +296,14 @@ impl DiskAPI for Disk { } } + #[tracing::instrument(skip(self))] + async fn read_file_chunks(&self, volume: &str, path: &str, offset: usize, length: usize) -> Result { + match self { + Disk::Local(local_disk) => local_disk.read_file_chunks(volume, path, offset, length).await, + Disk::Remote(remote_disk) => remote_disk.read_file_chunks(volume, path, offset, length).await, + } + } + #[tracing::instrument(skip(self))] async fn append_file(&self, volume: &str, path: &str) -> Result { match self { @@ -505,6 +514,9 @@ pub trait DiskAPI: Debug + Send + Sync + 'static { /// On other platforms, falls back to efficient read operations. async fn read_file_zero_copy(&self, volume: &str, path: &str, offset: usize, length: usize) -> Result; + /// Chunk-based file read compatibility layer for the zero-copy data plane. + async fn read_file_chunks(&self, volume: &str, path: &str, offset: usize, length: usize) -> Result; + async fn append_file(&self, volume: &str, path: &str) -> Result; async fn create_file(&self, origvolume: &str, volume: &str, path: &str, file_size: i64) -> Result; // ReadFileStream diff --git a/crates/ecstore/src/erasure_coding/bitrot.rs b/crates/ecstore/src/erasure_coding/bitrot.rs index 01e2f4c3c..7a02abf81 100644 --- a/crates/ecstore/src/erasure_coding/bitrot.rs +++ b/crates/ecstore/src/erasure_coding/bitrot.rs @@ -172,6 +172,52 @@ where } } +impl BitrotWriter { + fn write_inline_sync(&mut self, buf: &[u8]) -> std::io::Result { + if buf.is_empty() { + return Ok(0); + } + + if self.finished { + return Err(std::io::Error::new(std::io::ErrorKind::InvalidInput, "bitrot writer already finished")); + } + + if buf.len() > self.shard_size { + return Err(std::io::Error::new( + std::io::ErrorKind::InvalidInput, + format!("data size {} exceeds shard size {}", buf.len(), self.shard_size), + )); + } + + if buf.len() < self.shard_size { + self.finished = true; + } + + match &mut self.inner { + CustomWriter::InlineBuffer(data) => { + if self.hash_algo.size() > 0 { + let hash = self.hash_algo.hash_encode(buf); + if hash.as_ref().is_empty() { + error!("bitrot writer write hash error: hash is empty"); + return Err(std::io::Error::new(std::io::ErrorKind::InvalidInput, "hash is empty")); + } + data.extend_from_slice(hash.as_ref()); + } + data.extend_from_slice(buf); + Ok(buf.len()) + } + CustomWriter::Other(_) => Err(std::io::Error::other("inline sync write requires inline buffer writer")), + } + } + + fn shutdown_inline_sync(&mut self) -> std::io::Result<()> { + match self.inner { + CustomWriter::InlineBuffer(_) => Ok(()), + CustomWriter::Other(_) => Err(std::io::Error::other("inline sync shutdown requires inline buffer writer")), + } + } +} + async fn write_all_vectored(writer: &mut W, hash: &[u8], data: &[u8]) -> std::io::Result<()> where W: AsyncWrite + Unpin, @@ -280,6 +326,10 @@ impl CustomWriter { Self::Other(_) => None, } } + + pub fn is_inline_buffer(&self) -> bool { + matches!(self, Self::InlineBuffer(_)) + } } impl AsyncWrite for CustomWriter { @@ -397,6 +447,24 @@ impl BitrotWriterWrapper { self.bitrot_writer.shutdown().await } + pub fn is_inline_buffer(&self) -> bool { + matches!(self.writer_type, WriterType::InlineBuffer) + } + + pub fn write_inline_sync(&mut self, buf: &[u8]) -> std::io::Result { + if !self.is_inline_buffer() { + return Err(std::io::Error::other("inline sync write requires inline buffer writer")); + } + self.bitrot_writer.write_inline_sync(buf) + } + + pub fn shutdown_inline_sync(&mut self) -> std::io::Result<()> { + if !self.is_inline_buffer() { + return Err(std::io::Error::other("inline sync shutdown requires inline buffer writer")); + } + self.bitrot_writer.shutdown_inline_sync() + } + /// Extract the inline buffer data, consuming the wrapper pub fn into_inline_data(self) -> Option> { match self.writer_type { diff --git a/crates/ecstore/src/erasure_coding/decode.rs b/crates/ecstore/src/erasure_coding/decode.rs index 0e5d03ed4..e0a5790ae 100644 --- a/crates/ecstore/src/erasure_coding/decode.rs +++ b/crates/ecstore/src/erasure_coding/decode.rs @@ -17,6 +17,7 @@ use crate::disk::error_reduce::reduce_errs; use crate::erasure_coding::{BitrotReader, Erasure}; use futures::stream::{FuturesUnordered, StreamExt}; use pin_project_lite::pin_project; +use rustfs_io_core::{IoChunk, PooledChunk}; use std::io; use std::io::ErrorKind; use tokio::io::AsyncRead; @@ -155,6 +156,68 @@ fn get_data_block_len(shards: &[Option>], data_blocks: usize) -> usize { size } +fn block_window( + offset: usize, + length: usize, + block_size: usize, + block_index: usize, + start_block: usize, + end_block: usize, +) -> (usize, usize) { + let end_remainder = offset.saturating_add(length) % block_size; + if start_block == end_block { + (offset % block_size, length) + } else if block_index == start_block { + (offset % block_size, block_size - (offset % block_size)) + } else if block_index == end_block { + (0, if end_remainder == 0 { block_size } else { end_remainder }) + } else { + (0, block_size) + } +} + +fn take_data_blocks_as_chunks( + shards: &mut [Option>], + data_blocks: usize, + mut offset: usize, + length: usize, +) -> io::Result> { + if get_data_block_len(shards, data_blocks) < length { + error!("take_data_blocks_as_chunks get_data_block_len < length"); + return Err(io::Error::new(ErrorKind::UnexpectedEof, "Not enough data blocks to write")); + } + + let mut chunks = Vec::new(); + let mut remaining = length; + for block_op in shards.iter_mut().take(data_blocks) { + let Some(block) = block_op.take() else { + error!("take_data_blocks_as_chunks block_op.is_none()"); + return Err(io::Error::new(ErrorKind::UnexpectedEof, "Missing data block")); + }; + + if offset >= block.len() { + offset -= block.len(); + continue; + } + + let start = offset; + offset = 0; + let take = (block.len() - start).min(remaining); + let chunk = if start == 0 && take == block.len() { + IoChunk::Pooled(PooledChunk::from_vec(block)) + } else { + IoChunk::Pooled(PooledChunk::from_vec(block).slice(start, take)?) + }; + chunks.push(chunk); + remaining -= take; + if remaining == 0 { + break; + } + } + + Ok(chunks) +} + /// Write data blocks from encoded blocks to target, supporting offset and length async fn write_data_blocks( writer: &mut W, @@ -213,6 +276,134 @@ where Ok(total_written) } +pub(crate) struct ErasureChunkDecoder { + erasure: Erasure, + reader: ParallelReader, + offset: usize, + length: usize, + start_block: usize, + end_block: usize, + current_block: usize, + written: usize, + healable_error: Option, + finished: bool, +} + +impl ErasureChunkDecoder +where + R: AsyncRead + Unpin + Send + Sync, +{ + pub(crate) fn new( + erasure: Erasure, + readers: Vec>>, + offset: usize, + length: usize, + total_length: usize, + ) -> io::Result { + if readers.len() != erasure.data_shards + erasure.parity_shards { + return Err(io::Error::new(ErrorKind::InvalidInput, "Invalid number of readers")); + } + + let end_offset = offset + .checked_add(length) + .ok_or_else(|| io::Error::new(ErrorKind::InvalidInput, "offset + length exceeds total length"))?; + if end_offset > total_length { + return Err(io::Error::new(ErrorKind::InvalidInput, "offset + length exceeds total length")); + } + + let start_block = offset / erasure.block_size; + let end_block = if length == 0 { + start_block + } else { + end_offset.saturating_sub(1) / erasure.block_size + }; + let reader = ParallelReader::new(readers, erasure.clone(), offset, total_length); + + Ok(Self { + erasure, + reader, + offset, + length, + start_block, + end_block, + current_block: start_block, + written: 0, + healable_error: None, + finished: length == 0, + }) + } + + pub(crate) async fn next_chunks(&mut self) -> io::Result>> { + if self.finished { + return Ok(None); + } + + if self.current_block > self.end_block { + self.finished = true; + return Ok(None); + } + + let block_index = self.current_block; + self.current_block += 1; + + let (block_offset, block_length) = block_window( + self.offset, + self.length, + self.erasure.block_size, + block_index, + self.start_block, + self.end_block, + ); + if block_length == 0 { + self.finished = true; + return Ok(None); + } + + let (mut shards, errs) = self.reader.read().await; + + if self.healable_error.is_none() + && let (_, Some(err)) = reduce_errs(&errs, &[]) + && (err == Error::FileNotFound || err == Error::FileCorrupt) + { + self.healable_error = Some(err); + } + + if !self.reader.can_decode(&shards) { + self.finished = true; + error!("reconstructed chunk decoder can_decode errs: {:?}", &errs); + return Err(Error::ErasureReadQuorum.into()); + } + + if let Err(err) = self.erasure.decode_data(&mut shards) { + self.finished = true; + error!("reconstructed chunk decoder decode_data err: {:?}", err); + return Err(err); + } + + let chunks = take_data_blocks_as_chunks(&mut shards, self.erasure.data_shards, block_offset, block_length)?; + self.written += chunks.iter().map(IoChunk::len).sum::(); + Ok(Some(chunks)) + } + + pub(crate) fn written(&self) -> usize { + self.written + } + + pub(crate) fn finish_error(&self) -> Option { + if self.written < self.length { + Some(Error::LessData.into()) + } else { + None + } + } + + pub(crate) fn take_healable_error(&mut self) -> Option { + self.healable_error.take() + } +} + +pub(crate) type ReconstructedChunkDecoder = ErasureChunkDecoder; + impl Erasure { pub async fn decode( &self, @@ -230,7 +421,10 @@ impl Erasure { return (0, Some(io::Error::new(ErrorKind::InvalidInput, "Invalid number of readers"))); } - if offset + length > total_length { + let Some(end_offset) = offset.checked_add(length) else { + return (0, Some(io::Error::new(ErrorKind::InvalidInput, "offset + length exceeds total length"))); + }; + if end_offset > total_length { return (0, Some(io::Error::new(ErrorKind::InvalidInput, "offset + length exceeds total length"))); } @@ -245,7 +439,7 @@ impl Erasure { let mut reader = ParallelReader::new(readers, self.clone(), offset, total_length); let start = offset / self.block_size; - let end = (offset + length) / self.block_size; + let end = end_offset.saturating_sub(1) / self.block_size; for i in start..=end { let (block_offset, block_length) = if start == end { @@ -253,7 +447,8 @@ impl Erasure { } else if i == start { (offset % self.block_size, self.block_size - (offset % self.block_size)) } else if i == end { - (0, (offset + length) % self.block_size) + let end_remainder = end_offset % self.block_size; + (0, if end_remainder == 0 { self.block_size } else { end_remainder }) } else { (0, self.block_size) }; @@ -316,6 +511,7 @@ mod tests { disk::error::DiskError, erasure_coding::{BitrotReader, BitrotWriter}, }; + use bytes::Bytes; use rustfs_utils::HashAlgorithm; use std::io::Cursor; @@ -456,4 +652,47 @@ mod tests { let reader_cursor = Cursor::new(buf); BitrotReader::new(reader_cursor, shard_size, hash_algo.clone(), false) } + + async fn create_bitrot_reader_from_shard( + shard: Bytes, + shard_size: usize, + hash_algo: &HashAlgorithm, + ) -> BitrotReader>> { + let writer = Cursor::new(Vec::new()); + let mut writer = BitrotWriter::new(writer, shard_size, hash_algo.clone()); + writer.write(shard.as_ref()).await.unwrap(); + let reader_cursor = Cursor::new(writer.into_inner().into_inner()); + BitrotReader::new(reader_cursor, shard_size, hash_algo.clone(), false) + } + + #[tokio::test] + async fn test_erasure_chunk_decoder_reconstructs_missing_data_shard_as_pooled_chunks() { + let erasure = Erasure::new(2, 1, 4); + let original = b"abcd"; + let encoded = erasure.encode_data(original).unwrap(); + let shard_size = erasure.shard_size(); + let hash_algo = HashAlgorithm::None; + + let readers = vec![ + None, + Some(create_bitrot_reader_from_shard(encoded[1].clone(), shard_size, &hash_algo).await), + Some(create_bitrot_reader_from_shard(encoded[2].clone(), shard_size, &hash_algo).await), + ]; + + let mut decoder = ErasureChunkDecoder::new(erasure, readers, 0, original.len(), original.len()).unwrap(); + let first_batch = decoder.next_chunks().await.unwrap().unwrap(); + assert!( + first_batch.iter().all(|chunk| matches!(chunk, IoChunk::Pooled(_))), + "reconstructed decoder should produce pooled chunks" + ); + let collected = first_batch + .into_iter() + .flat_map(|chunk| chunk.as_bytes().to_vec()) + .collect::>(); + + assert_eq!(collected, original); + assert!(decoder.next_chunks().await.unwrap().is_none()); + assert_eq!(decoder.written(), original.len()); + assert!(decoder.finish_error().is_none()); + } } diff --git a/crates/ecstore/src/erasure_coding/encode.rs b/crates/ecstore/src/erasure_coding/encode.rs index e029f64a4..86b4a6930 100644 --- a/crates/ecstore/src/erasure_coding/encode.rs +++ b/crates/ecstore/src/erasure_coding/encode.rs @@ -17,11 +17,12 @@ use crate::disk::error_reduce::count_errs; use crate::disk::error_reduce::{OBJECT_OP_IGNORED_ERRS, reduce_write_quorum_errs}; use crate::erasure_coding::BitrotWriterWrapper; use crate::erasure_coding::Erasure; +use crate::erasure_coding::erasure::{EncodeBlockBuffer, EncodedShardBlock, EncodedShardBufferPool}; use bytes::Bytes; use futures::StreamExt; use futures::stream::FuturesUnordered; +use rustfs_rio::BlockReadable; use std::sync::Arc; -use std::vec; use tokio::io::AsyncRead; use tokio::sync::mpsc; use tracing::error; @@ -32,6 +33,164 @@ pub(crate) struct MultiWriter<'a> { errs: Vec>, } +pub(crate) struct BlockAssembler { + reader: R, + block_buffer: EncodeBlockBuffer, + total_bytes: usize, +} + +impl BlockAssembler +where + R: AsyncRead + BlockReadable + Send + Sync + Unpin + 'static, +{ + pub(crate) fn new(reader: R, block_size: usize) -> Self { + Self { + reader, + block_buffer: EncodeBlockBuffer::new(block_size), + total_bytes: 0, + } + } + + pub(crate) async fn next_block(&mut self) -> std::io::Result>> { + match self.block_buffer.read_from_block(&mut self.reader).await { + Ok(n) if n > 0 => { + self.total_bytes += n; + Ok(Some(self.block_buffer.filled(n).to_vec())) + } + Ok(_) => Ok(None), + Err(e) if e.kind() == std::io::ErrorKind::UnexpectedEof => { + if let Some(inner) = e.get_ref() + && rustfs_rio::is_checksum_mismatch(inner) + { + return Err(std::io::Error::new(std::io::ErrorKind::InvalidData, e.to_string())); + } + Ok(None) + } + Err(e) => Err(e), + } + } + + pub(crate) fn total_bytes(&self) -> usize { + self.total_bytes + } + + pub(crate) fn into_inner(self) -> R { + self.reader + } +} + +#[derive(Clone)] +pub(crate) struct ErasureChunkEncoder { + erasure: Arc, + buffer_pool: EncodedShardBufferPool, +} + +impl ErasureChunkEncoder { + pub(crate) async fn new(erasure: Arc) -> Self { + let reusable_capacity = erasure.shard_size() * erasure.total_shard_count(); + Self { + erasure, + buffer_pool: EncodedShardBufferPool::with_prefill(reusable_capacity, 2).await, + } + } + + pub(crate) async fn encode_block(&self, block: &[u8]) -> std::io::Result { + let reusable_buffer = self.buffer_pool.acquire().await; + self.erasure.encode_data_block_with_buffer(block, reusable_buffer) + } + + pub(crate) async fn release(&self, block: EncodedShardBlock) { + self.buffer_pool.release(block).await; + } +} + +pub(crate) struct ErasureWritePipeline { + erasure: Arc, + write_quorum: usize, +} + +impl ErasureWritePipeline { + pub(crate) fn new(erasure: Arc, write_quorum: usize) -> Self { + Self { erasure, write_quorum } + } + + pub(crate) async fn run(&self, reader: R, writers: &mut [Option]) -> std::io::Result<(R, usize)> + where + R: AsyncRead + BlockReadable + Send + Sync + Unpin + 'static, + { + let (tx, mut rx) = mpsc::channel::(8); + let producer = ErasureChunkEncoder::new(self.erasure.clone()).await; + let writer_pool = producer.clone(); + let block_size = self.erasure.block_size; + + let task = tokio::spawn(async move { + let mut assembler = BlockAssembler::new(reader, block_size); + while let Some(block) = assembler.next_block().await? { + let res = producer.encode_block(&block).await?; + if let Err(err) = tx.send(res).await { + return Err(std::io::Error::other(format!("Failed to send encoded data : {err}"))); + } + } + + let total = assembler.total_bytes(); + Ok((assembler.into_inner(), total)) + }); + + let mut writers = MultiWriter::new(writers, self.write_quorum); + let mut write_err = None; + + while let Some(block) = rx.recv().await { + if block.is_empty() { + break; + } + let write_result = writers.write(&block).await; + writer_pool.release(block).await; + if let Err(err) = write_result { + write_err = Some(err); + break; + } + } + + if let Some(err) = write_err { + task.abort(); + let _ = task.await; + if let Err(shutdown_err) = writers.shutdown().await { + error!("failed to shutdown erasure writers after write error: {:?}", shutdown_err); + } + return Err(err); + } + + let (reader, total) = task.await??; + writers.shutdown().await?; + Ok((reader, total)) + } +} + +pub(crate) trait ShardSource { + fn shard_count(&self) -> usize; + fn shard(&self, idx: usize) -> Bytes; +} + +impl ShardSource for EncodedShardBlock { + fn shard_count(&self) -> usize { + self.shard_count() + } + + fn shard(&self, idx: usize) -> Bytes { + self.shard(idx) + } +} + +impl ShardSource for Vec { + fn shard_count(&self) -> usize { + self.len() + } + + fn shard(&self, idx: usize) -> Bytes { + self[idx].clone() + } +} + impl<'a> MultiWriter<'a> { pub fn new(writers: &'a mut [Option], write_quorum: usize) -> Self { let length = writers.len(); @@ -42,10 +201,10 @@ impl<'a> MultiWriter<'a> { } } - async fn write_shard(writer_opt: &mut Option, err: &mut Option, shard: &Bytes) { + async fn write_shard(writer_opt: &mut Option, err: &mut Option, shard: Bytes) { match writer_opt { Some(writer) => { - match writer.write(shard).await { + match writer.write(&shard).await { Ok(n) => { if n < shard.len() { *err = Some(Error::ShortWrite); @@ -65,16 +224,40 @@ impl<'a> MultiWriter<'a> { } } - pub async fn write(&mut self, data: Vec) -> std::io::Result<()> { - assert_eq!(data.len(), self.writers.len()); + fn write_shard_inline(writer_opt: &mut Option, err: &mut Option, shard: Bytes) { + match writer_opt { + Some(writer) => match writer.write_inline_sync(&shard) { + Ok(n) => { + if n < shard.len() { + *err = Some(Error::ShortWrite); + *writer_opt = None; + } else { + *err = None; + } + } + Err(e) => { + *err = Some(Error::from(e)); + } + }, + None => { + *err = Some(Error::DiskNotFound); + } + } + } + + pub async fn write(&mut self, data: &T) -> std::io::Result<()> + where + T: ShardSource, + { + assert_eq!(data.shard_count(), self.writers.len()); { let mut futures = FuturesUnordered::new(); - for ((writer_opt, err), shard) in self.writers.iter_mut().zip(self.errs.iter_mut()).zip(data.iter()) { + for (idx, (writer_opt, err)) in self.writers.iter_mut().zip(self.errs.iter_mut()).enumerate() { if err.is_some() { continue; // Skip if we already have an error for this writer } - futures.push(Self::write_shard(writer_opt, err, shard)); + futures.push(Self::write_shard(writer_opt, err, data.shard(idx))); } while let Some(()) = futures.next().await {} } @@ -112,6 +295,45 @@ impl<'a> MultiWriter<'a> { ))) } + pub fn write_inline(&mut self, data: &T) -> std::io::Result<()> + where + T: ShardSource, + { + assert_eq!(data.shard_count(), self.writers.len()); + + for (idx, (writer_opt, err)) in self.writers.iter_mut().zip(self.errs.iter_mut()).enumerate() { + if err.is_some() { + continue; + } + Self::write_shard_inline(writer_opt, err, data.shard(idx)); + } + + let nil_count = self.errs.iter().filter(|&e| e.is_none()).count(); + if nil_count >= self.write_quorum { + return Ok(()); + } + + if let Some(write_err) = reduce_write_quorum_errs(&self.errs, OBJECT_OP_IGNORED_ERRS, self.write_quorum) { + return Err(std::io::Error::other(format!( + "Failed to write inline data: {} (offline-disks={}/{})", + write_err, + count_errs(&self.errs, &Error::DiskNotFound), + self.writers.len() + ))); + } + + Err(std::io::Error::other(format!( + "Failed to write inline data: (offline-disks={}/{}): {}", + count_errs(&self.errs, &Error::DiskNotFound), + self.writers.len(), + self.errs + .iter() + .map(|e| e.as_ref().map_or("".to_string(), |e| e.to_string())) + .collect::>() + .join(", ") + ))) + } + async fn shutdown_writer(writer_opt: &mut Option, err: &mut Option) { match writer_opt { Some(writer) => match writer.shutdown().await { @@ -129,6 +351,23 @@ impl<'a> MultiWriter<'a> { } } + fn shutdown_writer_inline(writer_opt: &mut Option, err: &mut Option) { + match writer_opt { + Some(writer) => match writer.shutdown_inline_sync() { + Ok(()) => { + *err = None; + } + Err(e) => { + *err = Some(Error::from(e)); + *writer_opt = None; + } + }, + None => { + *err = Some(Error::DiskNotFound); + } + } + } + pub async fn shutdown(&mut self) -> std::io::Result<()> { { let mut futures = FuturesUnordered::new(); @@ -173,66 +412,53 @@ impl<'a> MultiWriter<'a> { .join(", ") ))) } + + pub fn shutdown_inline(&mut self) -> std::io::Result<()> { + for (writer_opt, err) in self.writers.iter_mut().zip(self.errs.iter_mut()) { + if err.is_some() { + continue; + } + Self::shutdown_writer_inline(writer_opt, err); + } + + let nil_count = self.errs.iter().filter(|&e| e.is_none()).count(); + if nil_count >= self.write_quorum { + return Ok(()); + } + + if let Some(write_err) = reduce_write_quorum_errs(&self.errs, OBJECT_OP_IGNORED_ERRS, self.write_quorum) { + return Err(std::io::Error::other(format!( + "Failed to shutdown inline writers: {} (offline-disks={}/{})", + write_err, + count_errs(&self.errs, &Error::DiskNotFound), + self.writers.len() + ))); + } + + Err(std::io::Error::other(format!( + "Failed to shutdown inline writers: (offline-disks={}/{}): {}", + count_errs(&self.errs, &Error::DiskNotFound), + self.writers.len(), + self.errs + .iter() + .map(|e| e.as_ref().map_or("".to_string(), |e| e.to_string())) + .collect::>() + .join(", ") + ))) + } } impl Erasure { pub async fn encode( self: Arc, - mut reader: R, + reader: R, writers: &mut [Option], quorum: usize, ) -> std::io::Result<(R, usize)> where - R: AsyncRead + Send + Sync + Unpin + 'static, + R: AsyncRead + BlockReadable + Send + Sync + Unpin + 'static, { - let (tx, mut rx) = mpsc::channel::>(8); - - let task = tokio::spawn(async move { - let block_size = self.block_size; - let mut total = 0; - let mut buf = vec![0u8; block_size]; - loop { - match rustfs_utils::read_full(&mut reader, &mut buf).await { - Ok(n) if n > 0 => { - total += n; - let res = self.encode_data(&buf[..n])?; - if let Err(err) = tx.send(res).await { - return Err(std::io::Error::other(format!("Failed to send encoded data : {err}"))); - } - } - Ok(_) => { - break; - } - Err(e) if e.kind() == std::io::ErrorKind::UnexpectedEof => { - // Check if the inner error is a checksum mismatch - if so, propagate it - if let Some(inner) = e.get_ref() - && rustfs_rio::is_checksum_mismatch(inner) - { - return Err(std::io::Error::new(std::io::ErrorKind::InvalidData, e.to_string())); - } - break; - } - Err(e) => { - return Err(e); - } - } - } - - Ok((reader, total)) - }); - - let mut writers = MultiWriter::new(writers, quorum); - - while let Some(block) = rx.recv().await { - if block.is_empty() { - break; - } - writers.write(block).await?; - } - - let (reader, total) = task.await??; - writers.shutdown().await?; - Ok((reader, total)) + ErasureWritePipeline::new(self, quorum).run(reader, writers).await } } @@ -241,6 +467,7 @@ mod tests { use super::*; use crate::erasure_coding::{BitrotWriterWrapper, CustomWriter}; use rustfs_utils::HashAlgorithm; + use std::io::Cursor; use std::pin::Pin; use std::sync::{Arc, Mutex}; use std::task::{Context, Poll}; @@ -296,4 +523,28 @@ mod tests { assert_eq!(written, b"small payload".len()); assert!(!committed.lock().unwrap().is_empty()); } + + #[tokio::test] + async fn block_assembler_splits_input_into_erasure_blocks() { + let reader = tokio::io::BufReader::new(Cursor::new(b"abcdefghijkl".to_vec())); + let mut assembler = BlockAssembler::new(reader, 4); + + assert_eq!(assembler.next_block().await.unwrap(), Some(b"abcd".to_vec())); + assert_eq!(assembler.next_block().await.unwrap(), Some(b"efgh".to_vec())); + assert_eq!(assembler.next_block().await.unwrap(), Some(b"ijkl".to_vec())); + assert_eq!(assembler.next_block().await.unwrap(), None); + assert_eq!(assembler.total_bytes(), 12); + } + + #[tokio::test] + async fn erasure_chunk_encoder_produces_full_shard_block() { + let erasure = Arc::new(Erasure::new(2, 1, 4)); + let encoder = ErasureChunkEncoder::new(erasure.clone()).await; + let block = encoder.encode_block(b"abcd").await.unwrap(); + + assert_eq!(block.shard_count(), 3); + assert_eq!(block.shard(0).len(), erasure.shard_size()); + + encoder.release(block).await; + } } diff --git a/crates/ecstore/src/erasure_coding/erasure.rs b/crates/ecstore/src/erasure_coding/erasure.rs index 8942ab4e7..63a1350d2 100644 --- a/crates/ecstore/src/erasure_coding/erasure.rs +++ b/crates/ecstore/src/erasure_coding/erasure.rs @@ -19,12 +19,131 @@ use bytes::{Bytes, BytesMut}; use reed_solomon_erasure::galois_8::ReedSolomon; use reed_solomon_simd; +use rustfs_rio::BlockReadable; use smallvec::SmallVec; use std::io; +use std::sync::Arc; use tokio::io::AsyncRead; +use tokio::sync::Mutex; use tracing::warn; use uuid::Uuid; +pub(crate) struct EncodeBlockBuffer { + buf: Vec, +} + +impl EncodeBlockBuffer { + pub(crate) fn new(block_size: usize) -> Self { + Self { + buf: vec![0u8; block_size], + } + } + + pub(crate) async fn read_from(&mut self, reader: &mut R) -> io::Result + where + R: AsyncRead + Send + Sync + Unpin, + { + rustfs_utils::read_full(&mut *reader, &mut self.buf).await + } + + pub(crate) async fn read_from_block(&mut self, reader: &mut R) -> io::Result + where + R: BlockReadable + Send + Sync + Unpin, + { + reader.read_block(&mut self.buf).await + } + + pub(crate) fn filled(&self, len: usize) -> &[u8] { + &self.buf[..len] + } +} + +pub struct EncodedShardBlock { + data: Bytes, + shard_size: usize, + shard_count: usize, +} + +impl EncodedShardBlock { + pub(crate) fn new(data: Bytes, shard_size: usize, shard_count: usize) -> Self { + Self { + data, + shard_size, + shard_count, + } + } + + pub fn shard_count(&self) -> usize { + self.shard_count + } + + pub fn len(&self) -> usize { + self.shard_count + } + + pub fn is_empty(&self) -> bool { + self.shard_count == 0 + } + + pub fn shard(&self, idx: usize) -> Bytes { + let start = idx * self.shard_size; + let end = start + self.shard_size; + self.data.slice(start..end) + } + + pub fn iter(&self) -> impl Iterator + '_ { + (0..self.shard_count).map(|idx| self.shard(idx)) + } + + pub fn into_vec(self) -> Vec { + (0..self.shard_count).map(|idx| self.shard(idx)).collect() + } + + pub fn into_reusable_buffer(self) -> BytesMut { + match self.data.try_into_mut() { + Ok(mut buf) => { + buf.clear(); + buf + } + Err(data) => BytesMut::with_capacity(data.len()), + } + } +} + +#[derive(Clone)] +pub(crate) struct EncodedShardBufferPool { + capacity: usize, + free: Arc>>, +} + +impl EncodedShardBufferPool { + pub(crate) async fn with_prefill(capacity: usize, initial: usize) -> Self { + let mut free = Vec::with_capacity(initial); + for _ in 0..initial { + free.push(BytesMut::with_capacity(capacity)); + } + + Self { + capacity, + free: Arc::new(Mutex::new(free)), + } + } + + pub(crate) async fn acquire(&self) -> BytesMut { + let mut free = self.free.lock().await; + free.pop().unwrap_or_else(|| BytesMut::with_capacity(self.capacity)) + } + + pub(crate) async fn release(&self, block: EncodedShardBlock) { + let mut free = self.free.lock().await; + let mut buf = block.into_reusable_buffer(); + if buf.capacity() < self.capacity { + buf.reserve(self.capacity - buf.capacity()); + } + free.push(buf); + } +} + /// Legacy calc_shard_size formula: (block_size.div_ceil(data_shards) + 1) & !1 /// Matches main branch and filemeta::ErasureInfo for old-version files. pub fn calc_shard_size_legacy(block_size: usize, data_shards: usize) -> usize { @@ -351,6 +470,25 @@ impl Erasure { /// A vector of encoded shards as `Bytes`. #[tracing::instrument(level = "debug", skip_all, fields(data_len=data.len()))] pub fn encode_data(&self, data: &[u8]) -> io::Result> { + Ok(self.encode_data_block(data)?.into_vec()) + } + + /// Encode one logical block into an `EncodedShardBlock` using a caller-provided backing buffer. + /// + /// This is the explicit reuse-oriented variant for non-hot paths that want to + /// thread a reusable `BytesMut` across multiple encode calls. + #[tracing::instrument(level = "debug", skip_all, fields(data_len=data.len()))] + pub fn encode_data_with_buffer(&self, data: &[u8], data_buffer: BytesMut) -> io::Result { + self.encode_data_block_with_buffer(data, data_buffer) + } + + #[tracing::instrument(level = "debug", skip_all, fields(data_len=data.len()))] + pub(crate) fn encode_data_block(&self, data: &[u8]) -> io::Result { + self.encode_data_block_with_buffer(data, BytesMut::with_capacity(self.shard_size() * self.total_shard_count())) + } + + #[tracing::instrument(level = "debug", skip_all, fields(data_len=data.len()))] + pub(crate) fn encode_data_block_with_buffer(&self, data: &[u8], mut data_buffer: BytesMut) -> io::Result { let shard_size_fn = if self.uses_legacy { calc_shard_size_legacy } else { @@ -359,7 +497,10 @@ impl Erasure { let per_shard_size = shard_size_fn(data.len(), self.data_shards); let need_total_size = per_shard_size * self.total_shard_count(); - let mut data_buffer = BytesMut::with_capacity(need_total_size); + data_buffer.clear(); + if data_buffer.capacity() < need_total_size { + data_buffer.reserve(need_total_size - data_buffer.capacity()); + } data_buffer.extend_from_slice(data); data_buffer.resize(need_total_size, 0u8); @@ -382,14 +523,7 @@ impl Erasure { } // Zero-copy split, all shards reference data_buffer - let mut data_buffer = data_buffer.freeze(); - let mut shards = Vec::with_capacity(self.total_shard_count()); - for _ in 0..self.total_shard_count() { - let shard = data_buffer.split_to(per_shard_size); - shards.push(shard); - } - - Ok(shards) + Ok(EncodedShardBlock::new(data_buffer.freeze(), per_shard_size, self.total_shard_count())) } /// Decode and reconstruct missing shards in-place. @@ -478,8 +612,8 @@ impl Erasure { /// /// # Arguments /// * `reader` - An async reader implementing AsyncRead + Send + Sync + Unpin - /// * `mut on_block` - Async callback that receives encoded blocks and returns a Result - /// * `F` - Callback type: FnMut(Result, std::io::Error>) -> Future> + Send + /// * `mut on_block` - Async callback that receives encoded blocks and returns the block for reuse + /// * `F` - Callback type: FnMut(Result) -> Future, E>> + Send /// * `Fut` - Future type returned by the callback /// * `E` - Error type returned by the callback /// * `R` - Reader type implementing AsyncRead + Send + Sync + Unpin @@ -489,26 +623,31 @@ impl Erasure { /// /// # Errors /// Returns error if reading from reader fails or if callback returns error - pub async fn encode_stream_callback_async( + pub(crate) async fn encode_stream_callback_async( self: std::sync::Arc, reader: &mut R, mut on_block: F, ) -> Result where R: AsyncRead + Send + Sync + Unpin, - F: FnMut(std::io::Result>) -> Fut + Send, - Fut: std::future::Future> + Send, + F: FnMut(std::io::Result) -> Fut + Send, + Fut: std::future::Future, E>> + Send, { let block_size = self.block_size; let mut total = 0; + let mut block_buffer = EncodeBlockBuffer::new(block_size); + let reusable_capacity = self.shard_size() * self.total_shard_count(); + let buffer_pool = EncodedShardBufferPool::with_prefill(reusable_capacity, 1).await; loop { - let mut buf = vec![0u8; block_size]; - match rustfs_utils::read_full(&mut *reader, &mut buf).await { + match block_buffer.read_from(&mut *reader).await { Ok(n) if n > 0 => { warn!("encode_stream_callback_async read n={}", n); total += n; - let res = self.encode_data(&buf[..n]); - on_block(res).await? + let reusable_buffer = buffer_pool.acquire().await; + let res = self.encode_data_block_with_buffer(block_buffer.filled(n), reusable_buffer); + if let Some(block) = on_block(res).await? { + buffer_pool.release(block).await; + } } Ok(_) => { warn!("encode_stream_callback_async read unexpected ok"); @@ -520,11 +659,10 @@ impl Erasure { } Err(e) => { warn!("encode_stream_callback_async read error={:?}", e); - on_block(Err(e)).await?; + let _ = on_block(Err(e)).await?; break; } } - buf.clear(); } Ok(total) } @@ -747,8 +885,8 @@ mod tests { let tx = tx.clone(); async move { let shards = res.unwrap(); - tx.send(shards).await.unwrap(); - Ok(()) + tx.send(shards.iter().collect()).await.unwrap(); + Ok(Some(shards)) } }) .await @@ -760,6 +898,36 @@ mod tests { assert_eq!(collected_shards.len(), data_shards + parity_shards); } + #[test] + fn test_encode_data_with_buffer_supports_explicit_reuse() { + let erasure = Erasure::new(4, 2, 1024); + let reusable_capacity = erasure.shard_size() * erasure.total_shard_count(); + + let first_data = b"explicit reusable buffer path".repeat(32); + let first_block = erasure + .encode_data_with_buffer(&first_data, BytesMut::with_capacity(reusable_capacity)) + .expect("first encode should succeed"); + let reusable_buffer = first_block.into_reusable_buffer(); + assert!(reusable_buffer.capacity() >= reusable_capacity); + + let second_data = b"second encode through same reusable buffer".repeat(24); + let second_block = erasure + .encode_data_with_buffer(&second_data, reusable_buffer) + .expect("second encode should succeed"); + + let mut shards_opt: Vec>> = second_block.iter().map(|shard| Some(shard.to_vec())).collect(); + shards_opt[1] = None; + shards_opt[5] = None; + erasure.decode_data(&mut shards_opt).expect("decode should succeed"); + + let mut recovered = Vec::new(); + for shard in shards_opt.iter().take(erasure.data_shards) { + recovered.extend_from_slice(shard.as_ref().expect("data shard should exist after decode")); + } + recovered.truncate(second_data.len()); + assert_eq!(&recovered, &second_data); + } + #[tokio::test] async fn test_encode_stream_callback_async_channel_decode() { use std::io::Cursor; @@ -786,8 +954,8 @@ mod tests { let tx = tx.clone(); async move { let shards = res.unwrap(); - tx.send(shards).await.unwrap(); - Ok(()) + tx.send(shards.iter().collect()).await.unwrap(); + Ok(Some(shards)) } }) .await @@ -800,8 +968,8 @@ mod tests { // Test decode using the old API that operates in-place let mut decode_input: Vec>> = vec![None; data_shards + parity_shards]; - for i in 0..data_shards { - decode_input[i] = Some(shards[i].to_vec()); + for (i, shard) in shards.iter().enumerate().take(data_shards) { + decode_input[i] = Some(shard.to_vec()); } erasure.decode_data(&mut decode_input).unwrap(); @@ -1198,8 +1366,8 @@ mod tests { let tx = tx.clone(); async move { let shards = res.unwrap(); - tx.send(shards).await.unwrap(); - Ok(()) + tx.send(shards.iter().collect()).await.unwrap(); + Ok(Some(shards)) } }) .await @@ -1233,5 +1401,63 @@ mod tests { recovered.truncate(data_clone.len()); assert_eq!(&recovered, &data_clone); } + + #[tokio::test] + #[ignore] + async fn stress_simd_stream_callback_reuses_backing_buffers_across_many_blocks() { + use std::io::Cursor; + use std::sync::Arc; + use std::sync::atomic::{AtomicUsize, Ordering}; + use tokio::sync::Mutex; + + let data_shards = 4; + let parity_shards = 2; + let block_size = 1024; + let erasure = Arc::new(Erasure::new(data_shards, parity_shards, block_size)); + + let sample = + b"SIMD stress callback test payload that intentionally spans many blocks to exercise reusable backing buffers."; + let data = sample.repeat((4 * 1024 * 1024 / sample.len()).max(1)); + let data_clone = data.clone(); + let mut reader = Cursor::new(data); + + let recovered = Arc::new(Mutex::new(Vec::with_capacity(data_clone.len()))); + let block_count = Arc::new(AtomicUsize::new(0)); + let erasure_for_callback = erasure.clone(); + let recovered_for_callback = recovered.clone(); + let block_count_for_callback = block_count.clone(); + + erasure + .clone() + .encode_stream_callback_async::<_, _, (), _>(&mut reader, move |res| { + let erasure_for_callback = erasure_for_callback.clone(); + let recovered_for_callback = recovered_for_callback.clone(); + let block_count_for_callback = block_count_for_callback.clone(); + async move { + let shards = res.unwrap(); + block_count_for_callback.fetch_add(1, Ordering::Relaxed); + + let mut shards_opt: Vec>> = shards.iter().map(|b| Some(b.to_vec())).collect(); + shards_opt[1] = None; + shards_opt[5] = None; + erasure_for_callback.decode_data(&mut shards_opt).unwrap(); + + let mut recovered = recovered_for_callback.lock().await; + for shard in shards_opt.iter().take(data_shards) { + recovered.extend_from_slice(shard.as_ref().unwrap()); + } + + Ok(Some(shards)) + } + }) + .await + .unwrap(); + + assert!(block_count.load(Ordering::Relaxed) > 1024); + + let mut recovered = recovered.lock().await; + recovered.truncate(data_clone.len()); + assert_eq!(&*recovered, &data_clone); + } } } diff --git a/crates/ecstore/src/erasure_coding/heal.rs b/crates/ecstore/src/erasure_coding/heal.rs index 422ae7dac..66b6d0802 100644 --- a/crates/ecstore/src/erasure_coding/heal.rs +++ b/crates/ecstore/src/erasure_coding/heal.rs @@ -77,7 +77,7 @@ impl super::Erasure { let available_writers = writers.iter().filter(|w| w.is_some()).count(); let write_quorum = available_writers.max(1); // At least 1 writer must succeed let mut writers = MultiWriter::new(writers, write_quorum); - writers.write(shards).await?; + writers.write(&shards).await?; } Ok(()) diff --git a/crates/ecstore/src/rpc/remote_disk.rs b/crates/ecstore/src/rpc/remote_disk.rs index dd7945f60..781091ac2 100644 --- a/crates/ecstore/src/rpc/remote_disk.rs +++ b/crates/ecstore/src/rpc/remote_disk.rs @@ -30,8 +30,10 @@ use crate::{ }; use bytes::Bytes; use futures::lock::Mutex; +use futures_util::StreamExt; use http::{HeaderMap, HeaderValue, Method, header::CONTENT_TYPE}; use rustfs_filemeta::{FileInfo, ObjectPartInfo, RawFileInfo}; +use rustfs_io_core::{BoxChunkStream, IoChunk}; use rustfs_protos::proto_gen::node_service::RenamePartRequest; use rustfs_protos::proto_gen::node_service::{ CheckPartsRequest, DeletePathsRequest, DeleteRequest, DeleteVersionRequest, DeleteVersionsRequest, DeleteVolumeRequest, @@ -40,7 +42,7 @@ use rustfs_protos::proto_gen::node_service::{ RenameFileRequest, StatVolumeRequest, UpdateMetadataRequest, VerifyFileRequest, WriteAllRequest, WriteMetadataRequest, node_service_client::NodeServiceClient, }; -use rustfs_rio::{HttpReader, HttpWriter}; +use rustfs_rio::{HttpReader, HttpWriter, open_http_byte_stream}; use serde::{Serialize, de::DeserializeOwned}; use std::{ io::Cursor, @@ -1071,6 +1073,33 @@ impl DiskAPI for RemoteDisk { Ok(Bytes::from(buffer)) } + #[tracing::instrument(level = "debug", skip(self))] + async fn read_file_chunks(&self, volume: &str, path: &str, offset: usize, length: usize) -> Result { + if self.health.is_faulty() { + return Err(DiskError::FaultyDisk); + } + let disk = self.disk_ref().await; + + let url = format!( + "{}/rustfs/rpc/read_file_stream?disk={}&volume={}&path={}&offset={}&length={}", + self.endpoint.grid_host(), + urlencoding::encode(&disk), + urlencoding::encode(volume), + urlencoding::encode(path), + offset, + length + ); + + let mut headers = HeaderMap::new(); + headers.insert(CONTENT_TYPE, HeaderValue::from_static("application/json")); + build_auth_headers(&url, &Method::GET, &mut headers); + + let stream = open_http_byte_stream(url, Method::GET, headers, None) + .await? + .map(|result| result.map(IoChunk::Shared)); + Ok(Box::pin(stream)) + } + #[tracing::instrument(level = "debug", skip(self))] async fn append_file(&self, volume: &str, path: &str) -> Result { info!("append_file {}/{}", volume, path); diff --git a/crates/ecstore/src/set_disk.rs b/crates/ecstore/src/set_disk.rs index 7a150245b..b474d0c59 100644 --- a/crates/ecstore/src/set_disk.rs +++ b/crates/ecstore/src/set_disk.rs @@ -52,9 +52,9 @@ use crate::{ event_notification::{EventArgs, send_event}, global::{GLOBAL_LOCAL_DISK_MAP, GLOBAL_LOCAL_DISK_SET_DRIVES, get_global_deployment_id, is_dist_erasure}, store_api::{ - BucketInfo, BucketOperations, BucketOptions, CompletePart, DeleteBucketOptions, DeletedObject, GetObjectReader, - HTTPRangeSpec, HealOperations, ListMultipartsInfo, ListObjectsV2Info, ListOperations, MakeBucketOptions, MultipartInfo, - MultipartOperations, MultipartUploadResult, ObjectIO, ObjectInfo, ObjectOperations, PartInfo, PutObjReader, StorageAPI, + BucketInfo, BucketOperations, BucketOptions, ChunkNativePutData, CompletePart, DeleteBucketOptions, DeletedObject, + GetObjectReader, HTTPRangeSpec, HealOperations, ListMultipartsInfo, ListObjectsV2Info, ListOperations, MakeBucketOptions, + MultipartInfo, MultipartOperations, MultipartUploadResult, ObjectIO, ObjectInfo, ObjectOperations, PartInfo, StorageAPI, }, store_init::load_format_erasure, }; @@ -118,6 +118,66 @@ use tokio::{ time::{interval, timeout}, }; use tokio_util::sync::CancellationToken; + +const ENV_RUSTFS_PUT_INLINE_OBJECT_MAX_BYTES: &str = "RUSTFS_PUT_INLINE_OBJECT_MAX_BYTES"; +const ENV_RUSTFS_PUT_FORCE_DISABLE_INLINE: &str = "RUSTFS_PUT_FORCE_DISABLE_INLINE"; +const SLOW_PUT_STORAGE_PHASE_DEBUG_THRESHOLD_MS: u64 = 100; +const SLOW_PUT_STORAGE_PHASE_WARN_THRESHOLD_MS: u64 = 1_000; +const SLOW_PUT_STORAGE_PHASE_ERROR_THRESHOLD_MS: u64 = 5_000; + +fn env_flag_enabled(name: &str) -> bool { + rustfs_utils::get_env_bool(name, false) +} + +fn env_non_negative_usize(name: &str) -> Option { + rustfs_utils::get_env_opt_usize(name) +} + +fn resolved_put_inline_buffer_enabled(object_size: i64, inline_by_topology: bool) -> bool { + if !inline_by_topology || object_size < 0 { + return false; + } + + if env_flag_enabled(ENV_RUSTFS_PUT_FORCE_DISABLE_INLINE) { + return false; + } + + env_non_negative_usize(ENV_RUSTFS_PUT_INLINE_OBJECT_MAX_BYTES) + .map(|value| usize::try_from(object_size).is_ok_and(|size| size <= value)) + .unwrap_or(inline_by_topology) +} + +fn log_put_storage_phase( + bucket: &str, + object: &str, + phase: &str, + elapsed: Duration, + object_size: i64, + inline_selected: bool, + write_quorum: usize, +) { + let duration_ms = elapsed.as_millis() as u64; + if duration_ms < SLOW_PUT_STORAGE_PHASE_DEBUG_THRESHOLD_MS { + return; + } + + if duration_ms >= SLOW_PUT_STORAGE_PHASE_ERROR_THRESHOLD_MS { + error!( + phase, + duration_ms, object_size, inline_selected, write_quorum, bucket, object, "PUT storage phase is critically slow" + ); + } else if duration_ms >= SLOW_PUT_STORAGE_PHASE_WARN_THRESHOLD_MS { + warn!( + phase, + duration_ms, object_size, inline_selected, write_quorum, bucket, object, "PUT storage phase is slow" + ); + } else { + debug!( + phase, + duration_ms, object_size, inline_selected, write_quorum, bucket, object, "PUT storage phase exceeded debug threshold" + ); + } +} use tracing::error; use tracing::{debug, info, warn}; use uuid::Uuid; @@ -151,6 +211,9 @@ mod read; mod replication; mod write; +#[doc(hidden)] +pub use read::collect_direct_data_shard_chunks_for_benchmark; + /// Get lock acquire timeout from environment variable RUSTFS_LOCK_ACQUIRE_TIMEOUT (in seconds) /// Defaults to 30 seconds if not set or invalid pub fn get_lock_acquire_timeout() -> Duration { @@ -692,7 +755,13 @@ impl ObjectIO for SetDisks { } #[tracing::instrument(skip(self, data,))] - async fn put_object(&self, bucket: &str, object: &str, data: &mut PutObjReader, opts: &ObjectOptions) -> Result { + async fn put_object( + &self, + bucket: &str, + object: &str, + data: &mut ChunkNativePutData, + opts: &ObjectOptions, + ) -> Result { let disks = self.get_disks_internal().await; let mut object_lock_guard = None; @@ -774,13 +843,18 @@ impl ObjectIO for SetDisks { let erasure = erasure_coding::Erasure::new(fi.erasure.data_blocks, fi.erasure.parity_blocks, fi.erasure.block_size); let is_inline_buffer = { - if let Some(sc) = GLOBAL_STORAGE_CLASS.get() { + let inline_by_topology = if let Some(sc) = GLOBAL_STORAGE_CLASS.get() { sc.should_inline(erasure.shard_file_size(data.size()), opts.versioned) } else { false - } + }; + resolved_put_inline_buffer_enabled(data.size(), inline_by_topology) }; + if is_inline_buffer { + rustfs_io_metrics::record_put_inline_selected(data.size(), opts.versioned); + } + let writer_setup_start = Instant::now(); let mut writers = Vec::with_capacity(shuffle_disks.len()); let mut errors = Vec::with_capacity(shuffle_disks.len()); for disk_op in shuffle_disks.iter() { @@ -814,6 +888,15 @@ impl ObjectIO for SetDisks { writers.push(None); } } + log_put_storage_phase( + bucket, + object, + "writer_setup", + writer_setup_start.elapsed(), + data.size(), + is_inline_buffer, + write_quorum, + ); let nil_count = errors.iter().filter(|&e| e.is_none()).count(); if nil_count < write_quorum { @@ -825,23 +908,33 @@ impl ObjectIO for SetDisks { return Err(Error::other(format!("not enough disks to write: {errors:?}"))); } - let stream = mem::replace( - &mut data.stream, - 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 { - Ok((r, w)) => (r, w), + let object_size = data.size(); + let encode_write_start = Instant::now(); + let w_size = match Self::write_chunk_native_put_data(data, Arc::new(erasure), &mut writers, write_quorum).await { + Ok(written) => written, Err(e) => { + log_put_storage_phase( + bucket, + object, + "encode_write", + encode_write_start.elapsed(), + object_size, + is_inline_buffer, + write_quorum, + ); error!("encode err {:?}", e); return Err(e.into()); } }; // TODO: delete temporary directory on error - - let _ = mem::replace(&mut data.stream, reader); - // if let Err(err) = close_bitrot_writers(&mut writers).await { - // error!("close_bitrot_writers err {:?}", err); - // } + log_put_storage_phase( + bucket, + object, + "encode_write", + encode_write_start.elapsed(), + data.size(), + is_inline_buffer, + write_quorum, + ); if (w_size as i64) < data.size() { warn!("put_object write size < data.size(), w_size={}, data.size={}", w_size, data.size()); @@ -856,11 +949,11 @@ impl ObjectIO for SetDisks { insert_str(&mut user_defined, SUFFIX_COMPRESSION_SIZE, w_size.to_string()); } - let index_op = data.stream.try_get_index().map(|v| v.clone().into_vec()); + let index_op = data.index_bytes(); //TODO: userDefined - let etag = data.stream.try_resolve_etag().unwrap_or_default(); + let etag = data.resolve_etag().unwrap_or_default(); user_defined.insert("etag".to_owned(), etag.clone()); @@ -877,9 +970,9 @@ impl ObjectIO for SetDisks { } if fi.checksum.is_none() - && let Some(content_hash) = data.as_hash_reader().content_hash() + && let Some(content_hash) = data.content_hash_bytes()? { - fi.checksum = Some(content_hash.to_bytes(&[])); + fi.checksum = Some(content_hash); } if let Some(sc) = user_defined.get(AMZ_STORAGE_CLASS) @@ -917,6 +1010,7 @@ impl ObjectIO for SetDisks { drop(writers); // drop writers to close all files, this is to prevent FileAccessDenied errors when renaming data + let post_write_lock_start = Instant::now(); if !opts.no_lock && object_lock_guard.is_none() { let ns_lock = self.new_ns_lock(bucket, object).await?; object_lock_guard = Some(ns_lock.get_write_lock(get_lock_acquire_timeout()).await.map_err(|e| { @@ -926,7 +1020,17 @@ impl ObjectIO for SetDisks { )) })?); } + log_put_storage_phase( + bucket, + object, + "post_write_lock", + post_write_lock_start.elapsed(), + data.size(), + is_inline_buffer, + write_quorum, + ); + let finalize_start = Instant::now(); let (online_disks, _, op_old_dir) = Self::rename_data( &shuffle_disks, RUSTFS_META_TMP_BUCKET, @@ -936,7 +1040,18 @@ impl ObjectIO for SetDisks { object, write_quorum, ) - .await?; + .await + .inspect_err(|_| { + log_put_storage_phase( + bucket, + object, + "finalize", + finalize_start.elapsed(), + data.size(), + is_inline_buffer, + write_quorum, + ); + })?; if let Some(old_dir) = op_old_dir { self.commit_rename_data_dir(&online_disks, bucket, object, &old_dir.to_string(), write_quorum) @@ -946,6 +1061,15 @@ impl ObjectIO for SetDisks { drop(object_lock_guard); // drop object lock guard to release the lock self.delete_all(RUSTFS_META_TMP_BUCKET, &tmp_dir).await?; + log_put_storage_phase( + bucket, + object, + "finalize", + finalize_start.elapsed(), + data.size(), + is_inline_buffer, + write_quorum, + ); for (i, op_disk) in online_disks.iter().enumerate() { if let Some(disk) = op_disk @@ -1864,6 +1988,9 @@ impl ObjectOperations for SetDisks { if let Some(ref version_id) = opts.version_id { fi.version_id = Uuid::parse_str(version_id).ok(); } + if let Some(checksum) = &opts.resolved_checksum { + fi.checksum = Some(checksum.clone()); + } self.update_object_meta(bucket, object, fi.clone(), &online_disks) .await @@ -2090,7 +2217,7 @@ impl ObjectOperations for SetDisks { let gr = gr.unwrap(); let reader = BufReader::new(gr.stream); 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); + let mut p_reader = ChunkNativePutData::new(hash_reader); return match self_.clone().put_object(bucket, object, &mut p_reader, &ropts).await { Ok(restored_info) => { send_event(EventArgs { @@ -2158,7 +2285,7 @@ impl ObjectOperations for SetDisks { }; let reader = BufReader::new(gr.stream); 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 mut p_reader = ChunkNativePutData::new(hash_reader); let p_info = self_ .clone() .put_object_part(bucket, object, &res.upload_id, part_info.number, &mut p_reader, &ObjectOptions::default()) @@ -2382,7 +2509,7 @@ impl MultipartOperations for SetDisks { object: &str, upload_id: &str, part_id: usize, - data: &mut PutObjReader, + data: &mut ChunkNativePutData, opts: &ObjectOptions, ) -> Result { let upload_id_path = Self::get_upload_id_dir(bucket, object, upload_id); @@ -2394,9 +2521,8 @@ impl MultipartOperations for SetDisks { if let Some(checksum) = fi.metadata.get(rustfs_rio::RUSTFS_MULTIPART_CHECKSUM) && !checksum.is_empty() && data - .as_hash_reader() .content_crc_type() - .is_none_or(|v| v.to_string() != *checksum) + .is_none_or(|v: rustfs_rio::ChecksumType| v.to_string() != *checksum) { return Err(Error::other(format!("checksum mismatch: {checksum}"))); } @@ -2461,14 +2587,7 @@ impl MultipartOperations for SetDisks { return Err(Error::other(format!("not enough disks to write: {errors:?}"))); } - let stream = mem::replace( - &mut data.stream, - 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 - - let _ = mem::replace(&mut data.stream, reader); + let w_size = Self::write_chunk_native_put_data(data, Arc::new(erasure), &mut writers, write_quorum).await?; // TODO: delete temporary directory on error if (w_size as i64) < data.size() { warn!("put_object_part write size < data.size(), w_size={}, data.size={}", w_size, data.size()); @@ -2479,9 +2598,9 @@ impl MultipartOperations for SetDisks { ))); } - let index_op = data.stream.try_get_index().map(|v| v.clone().into_vec()); + let index_op = data.index_bytes(); - let mut etag = data.stream.try_resolve_etag().unwrap_or_default(); + let mut etag = data.resolve_etag().unwrap_or_default(); if let Some(ref tag) = opts.preserve_etag { etag = tag.clone(); @@ -2495,7 +2614,7 @@ impl MultipartOperations for SetDisks { } } - let checksums = data.as_hash_reader().content_crc(); + let checksums = data.content_crc(); let part_info = ObjectPartInfo { etag: etag.clone(), @@ -4212,6 +4331,31 @@ mod tests { } } + #[test] + #[serial] + fn resolved_put_inline_buffer_enabled_honors_disable_env() { + temp_env::with_var(ENV_RUSTFS_PUT_FORCE_DISABLE_INLINE, Some("true"), || { + assert!(!resolved_put_inline_buffer_enabled(4096, true)); + }); + } + + #[test] + #[serial] + fn resolved_put_inline_buffer_enabled_honors_max_bytes_override() { + temp_env::with_var(ENV_RUSTFS_PUT_INLINE_OBJECT_MAX_BYTES, Some("4096"), || { + assert!(resolved_put_inline_buffer_enabled(4096, true)); + assert!(!resolved_put_inline_buffer_enabled(4097, true)); + }); + } + + #[test] + #[serial] + fn resolved_put_inline_buffer_enabled_ignores_invalid_override() { + temp_env::with_var(ENV_RUSTFS_PUT_INLINE_OBJECT_MAX_BYTES, Some("invalid"), || { + assert!(resolved_put_inline_buffer_enabled(4096, true)); + }); + } + async fn current_setup_type() -> SetupType { if is_dist_erasure().await { SetupType::DistErasure diff --git a/crates/ecstore/src/set_disk/read.rs b/crates/ecstore/src/set_disk/read.rs index 08f6ec838..c527a7769 100644 --- a/crates/ecstore/src/set_disk/read.rs +++ b/crates/ecstore/src/set_disk/read.rs @@ -13,9 +13,879 @@ // limitations under the License. use super::*; +use crate::bitrot::create_bitrot_chunk_stream; +use crate::erasure_coding::decode::ErasureChunkDecoder; +use crate::erasure_coding::{calc_shard_size, calc_shard_size_legacy}; +use crate::store_api::{GetObjectChunkCopyMode, GetObjectChunkPath, GetObjectChunkResult}; +use bytes::BytesMut; +use futures_util::{Stream, StreamExt, stream}; use rustfs_config::{DEFAULT_OBJECT_ZERO_COPY_ENABLE, ENV_OBJECT_ZERO_COPY_ENABLE}; +use rustfs_io_core::{BoxChunkStream, IoChunk}; +use std::io; +use std::pin::Pin; +use std::sync::Mutex; +use std::task::{Context, Poll}; +use tokio::sync::mpsc::{UnboundedReceiver, UnboundedSender, unbounded_channel}; +use tokio_util::io::ReaderStream; + +struct ChannelChunkStream { + receiver: Mutex>>, +} + +impl ChannelChunkStream { + fn new(receiver: UnboundedReceiver>) -> Self { + Self { + receiver: Mutex::new(receiver), + } + } +} + +struct DirectShardCursor { + stream: BoxChunkStream, + current_chunk: Option, + current_offset: usize, +} + +impl DirectShardCursor { + fn new(stream: BoxChunkStream) -> Self { + Self { + stream, + current_chunk: None, + current_offset: 0, + } + } + + fn current_remaining(&self) -> usize { + self.current_chunk + .as_ref() + .map(|chunk| chunk.len().saturating_sub(self.current_offset)) + .unwrap_or(0) + } + + fn consume_current(&mut self, len: usize) { + self.current_offset += len; + if self.current_remaining() == 0 { + self.current_chunk = None; + self.current_offset = 0; + } + } + + async fn ensure_chunk(&mut self, shard_index: usize) -> io::Result { + if self.current_chunk.is_some() && self.current_remaining() > 0 { + return Ok(true); + } + + match self.stream.next().await { + Some(Ok(chunk)) => { + self.current_chunk = Some(chunk); + self.current_offset = 0; + Ok(true) + } + Some(Err(err)) => Err(err), + None => { + debug!(shard_index, "direct shard cursor reached EOF"); + Ok(false) + } + } + } +} + +impl Stream for ChannelChunkStream { + type Item = io::Result; + + fn poll_next(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll> { + let mut receiver = self.receiver.lock().unwrap(); + receiver.poll_recv(cx) + } +} + +struct ChannelChunkWriter { + sender: UnboundedSender>, + buffer: BytesMut, + chunk_size: usize, +} + +impl ChannelChunkWriter { + fn new(sender: UnboundedSender>, chunk_size: usize) -> Self { + Self { + sender, + buffer: BytesMut::with_capacity(chunk_size.max(1)), + chunk_size: chunk_size.max(1), + } + } + + fn push_buffer(&mut self) -> io::Result<()> { + if self.buffer.is_empty() { + return Ok(()); + } + + let bytes = self.buffer.split().freeze(); + self.sender + .send(Ok(IoChunk::Shared(bytes))) + .map_err(|_| io::Error::new(io::ErrorKind::BrokenPipe, "chunk stream receiver dropped")) + } + + fn finish(&mut self) -> io::Result<()> { + self.push_buffer() + } + + fn send_error(&self, err: io::Error) { + let _ = self.sender.send(Err(err)); + } +} + +impl AsyncWrite for ChannelChunkWriter { + fn poll_write(mut self: Pin<&mut Self>, _cx: &mut Context<'_>, buf: &[u8]) -> Poll> { + self.buffer.extend_from_slice(buf); + let chunk_size = self.chunk_size; + while self.buffer.len() >= chunk_size { + let chunk = self.buffer.split_to(chunk_size).freeze(); + if self.sender.send(Ok(IoChunk::Shared(chunk))).is_err() { + return Poll::Ready(Err(io::Error::new(io::ErrorKind::BrokenPipe, "chunk stream receiver dropped"))); + } + } + + Poll::Ready(Ok(buf.len())) + } + + fn poll_flush(mut self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll> { + Poll::Ready(self.push_buffer()) + } + + fn poll_shutdown(mut self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll> { + Poll::Ready(self.finish()) + } +} + +fn merge_chunk_copy_mode(current: GetObjectChunkCopyMode, next: GetObjectChunkCopyMode) -> GetObjectChunkCopyMode { + use GetObjectChunkCopyMode::{Reconstructed, SharedBytes, SingleCopy, TrueZeroCopy}; + + match (current, next) { + (Reconstructed, _) | (_, Reconstructed) => Reconstructed, + (SingleCopy, _) | (_, SingleCopy) => SingleCopy, + (SharedBytes, _) | (_, SharedBytes) => SharedBytes, + (TrueZeroCopy, TrueZeroCopy) => TrueZeroCopy, + } +} + +fn multipart_logical_part_size(fi: &FileInfo, part_index: usize) -> usize { + let part = &fi.parts[part_index]; + if part.actual_size > 0 { + part.actual_size as usize + } else { + part.size + } +} + +fn multipart_logical_total_size(fi: &FileInfo) -> usize { + if fi.parts.is_empty() { + return fi.size.max(0) as usize; + } + + fi.parts + .iter() + .map(|part| { + if part.actual_size > 0 { + part.actual_size as usize + } else { + part.size + } + }) + .sum() +} + +fn multipart_stored_part_size(fi: &FileInfo, part_index: usize) -> usize { + fi.parts[part_index].size +} + +fn multipart_stored_total_size(fi: &FileInfo) -> usize { + if fi.parts.is_empty() { + return fi.size.max(0) as usize; + } + + fi.parts.iter().map(|part| part.size).sum() +} + +fn multipart_to_logical_part_offset(fi: &FileInfo, offset: usize) -> Result<(usize, usize)> { + if offset == 0 { + return Ok((0, 0)); + } + + let mut part_offset = offset; + for (i, _) in fi.parts.iter().enumerate() { + let logical_part_size = multipart_logical_part_size(fi, i); + if part_offset < logical_part_size { + return Ok((i, part_offset)); + } + + part_offset -= logical_part_size; + } + + Err(Error::other("part not found")) +} + +fn multipart_to_stored_part_offset(fi: &FileInfo, offset: usize) -> Result<(usize, usize)> { + if offset == 0 { + return Ok((0, 0)); + } + + let mut part_offset = offset; + for (i, part) in fi.parts.iter().enumerate() { + if part_offset < part.size { + return Ok((i, part_offset)); + } + + part_offset -= part.size; + } + + Err(Error::other("part not found")) +} + +fn block_window( + offset: usize, + length: usize, + block_size: usize, + block_index: usize, + start_block: usize, + end_block: usize, +) -> (usize, usize) { + let end_remainder = offset.saturating_add(length) % block_size; + if start_block == end_block { + (offset % block_size, length) + } else if block_index == start_block { + (offset % block_size, block_size - (offset % block_size)) + } else if block_index == end_block { + (0, if end_remainder == 0 { block_size } else { end_remainder }) + } else { + (0, block_size) + } +} + +fn direct_block_shard_size( + total_size: usize, + block_size: usize, + data_shards: usize, + block_index: usize, + uses_legacy: bool, +) -> usize { + let block_start = block_index.saturating_mul(block_size); + let logical_block_size = total_size.saturating_sub(block_start).min(block_size); + if uses_legacy { + calc_shard_size_legacy(logical_block_size, data_shards) + } else { + calc_shard_size(logical_block_size, data_shards) + } +} + +#[allow(clippy::too_many_arguments)] +async fn send_direct_data_shard_chunks( + sender: UnboundedSender>, + shard_streams: Vec, + data_shards: usize, + block_size: usize, + total_size: usize, + uses_legacy: bool, + offset: usize, + length: usize, +) { + if length == 0 { + return; + } + + let start_block = offset / block_size; + let end_block = offset.saturating_add(length.saturating_sub(1)) / block_size; + let mut shard_cursors = shard_streams.into_iter().map(DirectShardCursor::new).collect::>(); + + for block_index in start_block..=end_block { + let (block_offset, block_length) = block_window(offset, length, block_size, block_index, start_block, end_block); + if block_length == 0 { + break; + } + + let shard_block_size = direct_block_shard_size(total_size, block_size, data_shards, block_index, uses_legacy); + let mut write_left = block_length; + let mut skip = block_offset; + + for (shard_index, shard_cursor) in shard_cursors.iter_mut().enumerate().take(data_shards) { + let mut shard_block_left = shard_block_size; + + while shard_block_left > 0 { + let has_chunk = match shard_cursor.ensure_chunk(shard_index).await { + Ok(has_chunk) => has_chunk, + Err(err) => { + let _ = sender.send(Err(err)); + return; + } + }; + + if !has_chunk { + let _ = sender.send(Err(io::Error::new( + io::ErrorKind::UnexpectedEof, + format!("missing chunk for data shard {shard_index}"), + ))); + return; + } + + let chunk = shard_cursor.current_chunk.as_ref().expect("chunk should exist after ensure"); + let chunk_remaining = chunk.len().saturating_sub(shard_cursor.current_offset); + let take_from_chunk = chunk_remaining.min(shard_block_left); + if skip >= take_from_chunk { + skip -= take_from_chunk; + shard_cursor.consume_current(take_from_chunk); + shard_block_left -= take_from_chunk; + continue; + } + + let start = shard_cursor.current_offset + skip; + let available = take_from_chunk.saturating_sub(skip); + let take = available.min(write_left); + let out_chunk = if start == shard_cursor.current_offset && take == take_from_chunk { + chunk.slice(start, take).expect("full remaining slice should succeed") + } else { + match chunk.slice(start, take) { + Ok(chunk) => chunk, + Err(err) => { + let _ = sender.send(Err(err)); + return; + } + } + }; + + let consumed = skip + take; + skip = 0; + shard_cursor.consume_current(consumed); + shard_block_left -= consumed; + if sender.send(Ok(out_chunk)).is_err() { + return; + } + write_left -= take; + + if write_left == 0 { + break; + } + } + + if write_left == 0 { + break; + } + } + + if write_left != 0 { + let _ = sender.send(Err(io::Error::new( + io::ErrorKind::UnexpectedEof, + "not enough decoded shard data for requested block", + ))); + return; + } + } +} + +#[doc(hidden)] +pub async fn collect_direct_data_shard_chunks_for_benchmark( + shard_streams: Vec, + data_shards: usize, + block_size: usize, + total_size: usize, + uses_legacy: bool, + offset: usize, + length: usize, +) -> io::Result> { + let (tx, rx) = unbounded_channel(); + send_direct_data_shard_chunks(tx, shard_streams, data_shards, block_size, total_size, uses_legacy, offset, length).await; + + let mut stream = ChannelChunkStream::new(rx); + let mut chunks = Vec::new(); + while let Some(chunk) = stream.next().await { + chunks.push(chunk?); + } + + Ok(chunks) +} + +#[allow(clippy::too_many_arguments)] +async fn build_reconstructed_part_stream( + bucket: &str, + object: &str, + part_number: usize, + part_offset: usize, + part_length: usize, + part_size: usize, + read_offset: usize, + till_offset: usize, + files: &[FileInfo], + disks: &[Option], + erasure: &erasure_coding::Erasure, + checksum_algo: rustfs_utils::HashAlgorithm, + skip_verify_bitrot: bool, + use_zero_copy: bool, +) -> Result> { + let mut readers = Vec::with_capacity(disks.len()); + let mut errors = Vec::with_capacity(disks.len()); + for (idx, disk_op) in disks.iter().enumerate() { + match create_bitrot_reader( + files[idx].data.as_deref(), + disk_op.as_ref(), + bucket, + &format!("{}/{}/part.{}", object, files[idx].data_dir.unwrap_or_default(), part_number), + read_offset, + till_offset, + erasure.shard_size(), + checksum_algo.clone(), + skip_verify_bitrot, + use_zero_copy, + ) + .await + { + Ok(Some(reader)) => { + readers.push(Some(reader)); + errors.push(None); + } + Ok(None) => { + readers.push(None); + errors.push(Some(DiskError::DiskNotFound)); + } + Err(err) => { + readers.push(None); + errors.push(Some(err)); + } + } + } + + let available_shards = errors.iter().filter(|error| error.is_none()).count(); + if available_shards < erasure.data_shards { + return Ok(None); + } + + let missing_shards = readers.len().saturating_sub(available_shards); + if missing_shards > 0 { + debug!( + bucket, + object, + part_number, + missing_shards, + available_shards, + data_shards = erasure.data_shards, + parity_shards = erasure.parity_shards, + "using reconstructed part stream for missing shards" + ); + } + + let (tx, rx) = unbounded_channel(); + let bucket = bucket.to_string(); + let object = object.to_string(); + let erasure = erasure.clone(); + tokio::spawn(async move { + let mut decoder = match ErasureChunkDecoder::new(erasure, readers, part_offset, part_length, part_size) { + Ok(decoder) => decoder, + Err(err) => { + let _ = tx.send(Err(err)); + return; + } + }; + + loop { + match decoder.next_chunks().await { + Ok(Some(chunks)) => { + for chunk in chunks { + if tx.send(Ok(chunk)).is_err() { + return; + } + } + } + Ok(None) => break, + Err(err) => { + let _ = tx.send(Err(err)); + return; + } + } + } + + if let Some(err) = decoder.finish_error() { + let _ = tx.send(Err(err)); + return; + } + + if let Some(disk_err) = decoder.take_healable_error() { + let allow_heal_only = + decoder.written() == part_length && matches!(disk_err, DiskError::FileNotFound | DiskError::FileCorrupt); + if !allow_heal_only { + let _ = tx.send(Err(io::Error::other(disk_err.to_string()))); + return; + } + + debug!( + bucket, + object, + part_number, + bytes_written = decoder.written(), + error = %disk_err, + "reconstructed part completed with healable shard error" + ); + } + }); + + Ok(Some(Box::pin(ChannelChunkStream::new(rx)))) +} impl SetDisks { + #[tracing::instrument(level = "debug", skip(self, h, opts))] + pub(crate) async fn get_object_chunks( + &self, + bucket: &str, + object: &str, + range: Option, + h: HeaderMap, + opts: &ObjectOptions, + ) -> Result { + let lock_optimization_enabled = is_lock_optimization_enabled(); + + let read_lock_guard = if !opts.no_lock { + let acquire_start = Instant::now(); + + if is_deadlock_detection_enabled() { + debug!( + lock_id = format!("{}:{}", bucket, object), + lock_type = "read", + resource = format!("{}/{}", bucket, object), + "Waiting for read lock" + ); + } + + let guard = self + .new_ns_lock(bucket, object) + .await? + .get_read_lock(get_lock_acquire_timeout()) + .await + .map_err(|e| { + Error::other(format!( + "Failed to acquire read lock: {}", + self.format_lock_error_from_error(bucket, object, "read", &e) + )) + })?; + + let _lock_id = record_lock_acquire(bucket, object, "read"); + metrics::counter!("rustfs.lock.acquire.total", "type" => "read").increment(1); + metrics::histogram!("rustfs.lock.acquire.duration.seconds").record(acquire_start.elapsed().as_secs_f64()); + + Some(guard) + } else { + None + }; + + let (fi, files, disks) = self + .get_object_fileinfo(bucket, object, opts, true) + .await + .map_err(|err| to_object_err(err, vec![bucket, object]))?; + let object_info = ObjectInfo::from_file_info(&fi, bucket, object, opts.versioned || opts.version_suspended); + + if object_info.delete_marker { + if opts.version_id.is_none() { + return Err(to_object_err(Error::FileNotFound, vec![bucket, object])); + } + return Err(to_object_err(Error::MethodNotAllowed, vec![bucket, object])); + } + + if object_info.size == 0 { + return Ok(GetObjectChunkResult { + stream: Box::pin(stream::iter(Vec::>::new())), + path: GetObjectChunkPath::Direct, + copy_mode: GetObjectChunkCopyMode::SharedBytes, + }); + } + + let (bridge_offset, bridge_length) = if fi.parts.is_empty() { + (0, fi.size) + } else { + let total_size = multipart_logical_total_size(&fi); + if let Some(range) = &range { + let (offset, length) = range + .get_offset_length(total_size as i64) + .map_err(|err| to_object_err(err, vec![bucket, object]))?; + (offset, length) + } else { + (0, total_size as i64) + } + }; + + if object_info.is_remote() { + let mut opts = opts.clone(); + if object_info.parts.len() == 1 { + opts.part_number = Some(1); + } + let gr = get_transitioned_object_reader(bucket, object, &range, &h, &object_info, &opts).await?; + let stream = ReaderStream::new(gr.stream).map(|result| result.map(IoChunk::Shared)); + return Ok(GetObjectChunkResult { + stream: Box::pin(stream), + path: GetObjectChunkPath::Bridge, + copy_mode: GetObjectChunkCopyMode::SingleCopy, + }); + } + + if fi.erasure.data_blocks > 0 { + let (disks, files) = Self::shuffle_disks_and_parts_metadata_by_index(&disks, &files, &fi); + let total_size = multipart_logical_total_size(&fi); + let requested_length = if let Some(range) = &range { + let (offset, length) = range + .get_offset_length(total_size as i64) + .map_err(|err| to_object_err(err, vec![bucket, object]))?; + (offset, length as usize) + } else { + (0, total_size) + }; + + let (part_index, mut part_offset) = multipart_to_logical_part_offset(&fi, requested_length.0)?; + let mut end_offset = requested_length.0; + if requested_length.1 > 0 { + end_offset += requested_length.1 - 1; + } + let (last_part_index, _) = multipart_to_logical_part_offset(&fi, end_offset)?; + + let use_zero_copy = rustfs_utils::get_env_bool(ENV_OBJECT_ZERO_COPY_ENABLE, DEFAULT_OBJECT_ZERO_COPY_ENABLE); + let single_shard_file = if fi.erasure.data_blocks == 1 { + Some( + files + .first() + .ok_or_else(|| Error::other("single-shard multipart metadata missing"))?, + ) + } else { + None + }; + let single_shard_disk = if fi.erasure.data_blocks == 1 { + Some( + disks + .first() + .ok_or_else(|| Error::other("single-shard multipart disk slot missing"))?, + ) + } else { + None + }; + let mut part_streams = Vec::new(); + let mut part_total_read = 0usize; + let mut merged_copy_mode = GetObjectChunkCopyMode::TrueZeroCopy; + + for current_part in part_index..=last_part_index { + let part_number = fi.parts[current_part].number; + let part_size = multipart_logical_part_size(&fi, current_part); + let mut part_length = part_size - part_offset; + if part_length > (requested_length.1 - part_total_read) { + part_length = requested_length.1 - part_total_read; + } + let checksum_info = fi.erasure.get_checksum_info(part_number); + let checksum_algo = + if fi.uses_legacy_checksum && checksum_info.algorithm == rustfs_utils::HashAlgorithm::HighwayHash256S { + rustfs_utils::HashAlgorithm::HighwayHash256SLegacy + } else { + checksum_info.algorithm.clone() + }; + + if fi.erasure.data_blocks == 1 { + let single_shard_file = single_shard_file.expect("single-shard multipart metadata must exist"); + let single_shard_disk = single_shard_disk.expect("single-shard multipart disk slot must exist"); + let data_dir = single_shard_file + .data_dir + .as_ref() + .map(uuid::Uuid::to_string) + .unwrap_or_default(); + let data_path = format!("{}/{}/part.{}", object, data_dir, part_number); + let chunk_result = create_bitrot_chunk_stream( + single_shard_file.data.as_deref(), + single_shard_disk.as_ref(), + bucket, + &data_path, + part_offset, + part_length, + part_size, + fi.erasure.shard_size(), + checksum_algo, + opts.skip_verify_bitrot, + use_zero_copy, + ) + .await?; + + let Some(chunk_result) = chunk_result else { + part_streams.clear(); + break; + }; + merged_copy_mode = merge_chunk_copy_mode(merged_copy_mode, chunk_result.copy_mode); + part_streams.push(chunk_result.stream); + } else { + let erasure = erasure_coding::Erasure::new_with_options( + fi.erasure.data_blocks, + fi.erasure.parity_blocks, + fi.erasure.block_size, + fi.uses_legacy_checksum, + ); + let read_offset = (part_offset / erasure.block_size) * erasure.shard_size(); + let till_offset = erasure.shard_file_offset(part_offset, part_length, part_size); + let shard_length = till_offset.saturating_sub(read_offset); + let shard_total_size = erasure.shard_file_size(part_size as i64) as usize; + let mut shard_streams = Vec::with_capacity(erasure.data_shards); + let mut part_copy_mode = GetObjectChunkCopyMode::TrueZeroCopy; + let mut needs_reconstruct = false; + + for shard_index in 0..erasure.data_shards { + let data_path = + format!("{}/{}/part.{}", object, files[shard_index].data_dir.unwrap_or_default(), part_number); + let chunk_result = match create_bitrot_chunk_stream( + files[shard_index].data.as_deref(), + disks[shard_index].as_ref(), + bucket, + &data_path, + read_offset, + shard_length, + shard_total_size, + erasure.shard_size(), + checksum_algo.clone(), + opts.skip_verify_bitrot, + use_zero_copy, + ) + .await + { + Ok(Some(chunk_result)) => chunk_result, + Ok(None) => { + needs_reconstruct = true; + shard_streams.clear(); + break; + } + Err(err) => { + debug!( + bucket, + object, + part_number, + shard_index, + error = %err, + "multi-shard direct chunk path unavailable, falling back to decoded read path" + ); + needs_reconstruct = true; + shard_streams.clear(); + break; + } + }; + part_copy_mode = merge_chunk_copy_mode(part_copy_mode, chunk_result.copy_mode); + shard_streams.push(chunk_result.stream); + } + + if needs_reconstruct { + let reconstructed_stream = match build_reconstructed_part_stream( + bucket, + object, + part_number, + part_offset, + part_length, + part_size, + read_offset, + till_offset, + &files, + &disks, + &erasure, + checksum_algo, + opts.skip_verify_bitrot, + use_zero_copy, + ) + .await? + { + Some(stream) => stream, + None => { + part_streams.clear(); + break; + } + }; + + merged_copy_mode = merge_chunk_copy_mode(merged_copy_mode, GetObjectChunkCopyMode::Reconstructed); + part_streams.push(reconstructed_stream); + part_total_read += part_length; + part_offset = 0; + continue; + } + + if shard_streams.len() != erasure.data_shards { + part_streams.clear(); + break; + } + + let (tx, rx) = unbounded_channel(); + tokio::spawn(send_direct_data_shard_chunks( + tx, + shard_streams, + erasure.data_shards, + erasure.block_size, + part_size, + fi.uses_legacy_checksum, + part_offset, + part_length, + )); + + merged_copy_mode = merge_chunk_copy_mode(merged_copy_mode, part_copy_mode); + part_streams.push(Box::pin(ChannelChunkStream::new(rx))); + } + part_total_read += part_length; + part_offset = 0; + } + + if !part_streams.is_empty() { + return Ok(GetObjectChunkResult { + stream: Box::pin(stream::iter(part_streams).flatten()), + path: GetObjectChunkPath::Direct, + copy_mode: merged_copy_mode, + }); + } + } + + let read_lock_guard = if lock_optimization_enabled { + if read_lock_guard.is_some() { + let lock_id = format!("{}:{}", bucket, object); + record_lock_release(bucket, object, &lock_id, "read"); + metrics::counter!("rustfs.lock.release.early.total", "type" => "read").increment(1); + } + drop(read_lock_guard); + debug!(bucket, object, "Lock optimization: released read lock after metadata read"); + None + } else { + read_lock_guard + }; + + let chunk_size = get_duplex_buffer_size(); + let bucket = bucket.to_owned(); + let object = object.to_owned(); + let set_index = self.set_index; + let pool_index = self.pool_index; + let skip_verify = opts.skip_verify_bitrot; + let (tx, rx) = unbounded_channel(); + + tokio::spawn(async move { + let _guard = read_lock_guard; + let mut writer = ChannelChunkWriter::new(tx, chunk_size); + if let Err(err) = Self::get_object_with_fileinfo( + &bucket, + &object, + bridge_offset, + bridge_length, + &mut writer, + fi, + files, + &disks, + set_index, + pool_index, + skip_verify, + ) + .await + { + error!("get_object_with_fileinfo {bucket}/{object} err {:?}", err); + writer.send_error(io::Error::other(err.to_string())); + } + + if let Err(err) = writer.finish() { + debug!(bucket, object, error = %err, "failed to flush chunk writer"); + } + }); + + Ok(GetObjectChunkResult { + stream: Box::pin(ChannelChunkStream::new(rx)), + path: GetObjectChunkPath::Bridge, + copy_mode: GetObjectChunkCopyMode::SingleCopy, + }) + } + pub(super) async fn read_parts( disks: &[Option], bucket: &str, @@ -583,27 +1453,42 @@ impl SetDisks { debug!(bucket, object, requested_length = length, offset, "get_object_with_fileinfo start"); let (disks, files) = Self::shuffle_disks_and_parts_metadata_by_index(disks, &files, &fi); - let total_size = fi.size as usize; - - let length = if length < 0 { - fi.size as usize - offset + let logical_total_size = multipart_logical_total_size(&fi); + let use_stored_part_sizes = length > logical_total_size as i64; + let total_size = if use_stored_part_sizes { + multipart_stored_total_size(&fi) } else { - length as usize + logical_total_size }; - if offset > total_size || offset + length > total_size { + let length = if length < 0 { total_size - offset } else { length as usize }; + + let Some(end_offset_exclusive) = offset.checked_add(length) else { + error!("get_object_with_fileinfo offset overflow: {}, length: {}", offset, length); + return Err(Error::other("offset out of range")); + }; + + if offset > total_size || end_offset_exclusive > total_size { error!("get_object_with_fileinfo offset out of range: {}, total_size: {}", offset, total_size); return Err(Error::other("offset out of range")); } - let (part_index, mut part_offset) = fi.to_part_offset(offset)?; + let (part_index, mut part_offset) = if use_stored_part_sizes { + multipart_to_stored_part_offset(&fi, offset)? + } else { + multipart_to_logical_part_offset(&fi, offset)? + }; let mut end_offset = offset; if length > 0 { - end_offset += length - 1 + end_offset = end_offset_exclusive - 1; } - let (last_part_index, last_part_relative_offset) = fi.to_part_offset(end_offset)?; + let (last_part_index, last_part_relative_offset) = if use_stored_part_sizes { + multipart_to_stored_part_offset(&fi, end_offset)? + } else { + multipart_to_logical_part_offset(&fi, end_offset)? + }; debug!( bucket, @@ -635,7 +1520,11 @@ impl SetDisks { } let part_number = fi.parts[current_part].number; - let part_size = fi.parts[current_part].size; + let part_size = if use_stored_part_sizes { + multipart_stored_part_size(&fi, current_part) + } else { + multipart_logical_part_size(&fi, current_part) + }; let mut part_length = part_size - part_offset; if part_length > (length - total_read) { part_length = length - total_read @@ -833,3 +1722,103 @@ impl SetDisks { Ok(()) } } + +#[cfg(test)] +mod tests { + use super::*; + use bytes::Bytes; + use futures_util::StreamExt; + + #[tokio::test] + async fn send_direct_data_shard_chunks_reassembles_multi_block_range() { + let data_shards = 4; + let block_size = 16; + let shard_streams: Vec = vec![ + Box::pin(stream::iter(vec![ + Ok(IoChunk::Shared(Bytes::copy_from_slice(&[0, 1, 2, 3]))), + Ok(IoChunk::Shared(Bytes::copy_from_slice(&[16, 17, 18, 19]))), + ])), + Box::pin(stream::iter(vec![ + Ok(IoChunk::Shared(Bytes::copy_from_slice(&[4, 5, 6, 7]))), + Ok(IoChunk::Shared(Bytes::copy_from_slice(&[20, 21, 22, 23]))), + ])), + Box::pin(stream::iter(vec![ + Ok(IoChunk::Shared(Bytes::copy_from_slice(&[8, 9, 10, 11]))), + Ok(IoChunk::Shared(Bytes::copy_from_slice(&[24, 25, 26, 27]))), + ])), + Box::pin(stream::iter(vec![ + Ok(IoChunk::Shared(Bytes::copy_from_slice(&[12, 13, 14, 15]))), + Ok(IoChunk::Shared(Bytes::copy_from_slice(&[28, 29, 30, 31]))), + ])), + ]; + let (tx, rx) = unbounded_channel(); + + send_direct_data_shard_chunks(tx, shard_streams, data_shards, block_size, 32, false, 3, 18).await; + + let mut stream = ChannelChunkStream::new(rx); + let mut collected = Vec::new(); + while let Some(chunk) = stream.next().await { + collected.extend_from_slice(&chunk.unwrap().as_bytes()); + } + + assert_eq!(collected, (3u8..21).collect::>()); + } + + #[tokio::test] + async fn send_direct_data_shard_chunks_keeps_block_boundaries_with_cross_block_chunks() { + let data_shards = 4; + let block_size = 16; + let shard_streams: Vec = vec![ + Box::pin(stream::iter(vec![Ok(IoChunk::Shared(Bytes::copy_from_slice(&[ + 0, 1, 2, 3, 16, 17, 18, 19, + ])))])), + Box::pin(stream::iter(vec![Ok(IoChunk::Shared(Bytes::copy_from_slice(&[ + 4, 5, 6, 7, 20, 21, 22, 23, + ])))])), + Box::pin(stream::iter(vec![Ok(IoChunk::Shared(Bytes::copy_from_slice(&[ + 8, 9, 10, 11, 24, 25, 26, 27, + ])))])), + Box::pin(stream::iter(vec![Ok(IoChunk::Shared(Bytes::copy_from_slice(&[ + 12, 13, 14, 15, 28, 29, 30, 31, + ])))])), + ]; + let (tx, rx) = unbounded_channel(); + + send_direct_data_shard_chunks(tx, shard_streams, data_shards, block_size, 32, false, 3, 18).await; + + let mut stream = ChannelChunkStream::new(rx); + let mut collected = Vec::new(); + while let Some(chunk) = stream.next().await { + collected.extend_from_slice(&chunk.unwrap().as_bytes()); + } + + assert_eq!(collected, (3u8..21).collect::>()); + } + + #[tokio::test] + async fn send_direct_data_shard_chunks_keeps_final_full_block_when_length_is_block_aligned() { + let data_shards = 2; + let block_size = 16; + let shard_streams: Vec = vec![ + Box::pin(stream::iter(vec![ + Ok(IoChunk::Shared(Bytes::copy_from_slice(&[0, 1, 2, 3, 4, 5, 6, 7]))), + Ok(IoChunk::Shared(Bytes::copy_from_slice(&[16, 17, 18, 19, 20, 21, 22, 23]))), + ])), + Box::pin(stream::iter(vec![ + Ok(IoChunk::Shared(Bytes::copy_from_slice(&[8, 9, 10, 11, 12, 13, 14, 15]))), + Ok(IoChunk::Shared(Bytes::copy_from_slice(&[24, 25, 26, 27, 28, 29, 30, 31]))), + ])), + ]; + let (tx, rx) = unbounded_channel(); + + send_direct_data_shard_chunks(tx, shard_streams, data_shards, block_size, 32, false, 0, 32).await; + + let mut stream = ChannelChunkStream::new(rx); + let mut collected = Vec::new(); + while let Some(chunk) = stream.next().await { + collected.extend_from_slice(&chunk.unwrap().as_bytes()); + } + + assert_eq!(collected, (0u8..32).collect::>()); + } +} diff --git a/crates/ecstore/src/set_disk/write.rs b/crates/ecstore/src/set_disk/write.rs index d937ede3e..7c0122ee4 100644 --- a/crates/ecstore/src/set_disk/write.rs +++ b/crates/ecstore/src/set_disk/write.rs @@ -13,12 +13,87 @@ // limitations under the License. use super::*; +use crate::store_api::ChunkNativePutData; impl SetDisks { + fn all_inline_bitrot_writers(writers: &[Option]) -> bool { + writers.iter().all(|writer| { + writer + .as_ref() + .is_some_and(crate::erasure_coding::BitrotWriterWrapper::is_inline_buffer) + }) + } + + async fn write_chunk_native_put_data_inline( + data: &mut ChunkNativePutData, + erasure: Arc, + writers: &mut [Option], + write_quorum: usize, + ) -> std::io::Result { + let stream = data.take_stream()?; + let mut assembler = erasure_coding::encode::BlockAssembler::new(stream, erasure.block_size); + let encoder = erasure_coding::encode::ErasureChunkEncoder::new(erasure).await; + let mut writer_group = erasure_coding::encode::MultiWriter::new(writers, write_quorum); + + loop { + let block = match assembler.next_block().await { + Ok(Some(block)) => block, + Ok(None) => break, + Err(err) => { + data.restore_stream(assembler.into_inner()); + return Err(err); + } + }; + + let encoded = match encoder.encode_block(&block).await { + Ok(encoded) => encoded, + Err(err) => { + data.restore_stream(assembler.into_inner()); + return Err(err); + } + }; + + if let Err(err) = writer_group.write_inline(&encoded) { + encoder.release(encoded).await; + data.restore_stream(assembler.into_inner()); + return Err(err); + } + + encoder.release(encoded).await; + } + + let total_bytes = assembler.total_bytes(); + let stream = assembler.into_inner(); + if let Err(err) = writer_group.shutdown_inline() { + data.restore_stream(stream); + return Err(err); + } + data.restore_stream(stream); + Ok(total_bytes) + } + pub(super) fn default_read_quorum(&self) -> usize { self.set_drive_count - self.default_parity_count } + pub(super) async fn write_chunk_native_put_data( + data: &mut ChunkNativePutData, + erasure: Arc, + writers: &mut [Option], + write_quorum: usize, + ) -> std::io::Result { + if Self::all_inline_bitrot_writers(writers) { + return Self::write_chunk_native_put_data_inline(data, erasure, writers, write_quorum).await; + } + + let stream = data.take_stream()?; + let (stream, written) = erasure_coding::encode::ErasureWritePipeline::new(erasure, write_quorum) + .run(stream, writers) + .await?; + data.restore_stream(stream); + Ok(written) + } + pub(super) fn default_write_quorum(&self) -> usize { let mut data_count = self.set_drive_count - self.default_parity_count; if data_count == self.default_parity_count { @@ -627,3 +702,60 @@ impl SetDisks { None } } + +#[cfg(test)] +mod tests { + use super::*; + use crate::erasure_coding::{BitrotWriterWrapper, CustomWriter}; + use crate::store_api::ChunkNativePutData; + + #[tokio::test] + async fn write_chunk_native_put_data_restores_reader_state_after_encoding() { + let payload = b"chunk-native-put-payload".repeat(8); + let erasure = Arc::new(erasure_coding::Erasure::new(2, 1, 8)); + let mut reader = ChunkNativePutData::from_vec(payload.clone()); + let mut writers: Vec> = (0..erasure.total_shard_count()) + .map(|_| { + Some(BitrotWriterWrapper::new( + CustomWriter::new_inline_buffer(), + erasure.shard_size(), + HashAlgorithm::HighwayHash256S, + )) + }) + .collect(); + + let written = SetDisks::write_chunk_native_put_data(&mut reader, erasure.clone(), &mut writers, 2) + .await + .unwrap(); + + assert_eq!(written, payload.len()); + assert_eq!(reader.size(), payload.len() as i64); + assert_eq!(reader.actual_size(), payload.len() as i64); + assert!( + reader.resolve_etag().is_some(), + "restored reader should preserve computed etag state after chunk-native encode" + ); + + let inline_lengths: Vec = writers + .into_iter() + .map(|writer| writer.expect("writer").into_inline_data().expect("inline data").len()) + .collect(); + assert!(inline_lengths.iter().all(|len| *len > 0)); + } + + #[tokio::test] + async fn write_chunk_native_put_data_detects_inline_writer_set() { + let erasure = Arc::new(erasure_coding::Erasure::new(2, 1, 8)); + let writers: Vec> = (0..erasure.total_shard_count()) + .map(|_| { + Some(BitrotWriterWrapper::new( + CustomWriter::new_inline_buffer(), + erasure.shard_size(), + HashAlgorithm::HighwayHash256S, + )) + }) + .collect(); + + assert!(SetDisks::all_inline_bitrot_writers(&writers)); + } +} diff --git a/crates/ecstore/src/sets.rs b/crates/ecstore/src/sets.rs index 48502a4b6..3d32847a8 100644 --- a/crates/ecstore/src/sets.rs +++ b/crates/ecstore/src/sets.rs @@ -15,7 +15,7 @@ use crate::disk::error_reduce::count_errs; use crate::error::{Error, Result}; -use crate::store_api::{ListPartsInfo, ObjectInfoOrErr, WalkOptions}; +use crate::store_api::{GetObjectChunkResult, ListPartsInfo, ObjectInfoOrErr, WalkOptions}; use crate::{ disk::{ DiskAPI, DiskInfo, DiskOption, DiskStore, @@ -28,10 +28,10 @@ use crate::{ global::{GLOBAL_LOCAL_DISK_SET_DRIVES, get_global_lock_clients, is_dist_erasure}, set_disk::SetDisks, store_api::{ - BucketInfo, BucketOperations, BucketOptions, CompletePart, DeleteBucketOptions, DeletedObject, GetObjectReader, - HTTPRangeSpec, HealOperations, ListMultipartsInfo, ListObjectVersionsInfo, ListObjectsV2Info, ListOperations, - MakeBucketOptions, MultipartInfo, MultipartOperations, MultipartUploadResult, ObjectIO, ObjectInfo, ObjectOperations, - ObjectOptions, ObjectToDelete, PartInfo, PutObjReader, StorageAPI, + BucketInfo, BucketOperations, BucketOptions, ChunkNativePutData, CompletePart, DeleteBucketOptions, DeletedObject, + GetObjectReader, HTTPRangeSpec, HealOperations, ListMultipartsInfo, ListObjectVersionsInfo, ListObjectsV2Info, + ListOperations, MakeBucketOptions, MultipartInfo, MultipartOperations, MultipartUploadResult, ObjectIO, ObjectInfo, + ObjectOperations, ObjectOptions, ObjectToDelete, PartInfo, StorageAPI, }, store_init::{check_format_erasure_values, get_format_erasure_in_quorum, load_format_erasure_all, save_format_file}, }; @@ -287,6 +287,19 @@ impl Sets { self.get_disks(self.get_hashed_set_index(key)) } + pub async fn get_object_chunks( + &self, + bucket: &str, + object: &str, + range: Option, + h: HeaderMap, + opts: &ObjectOptions, + ) -> Result { + self.get_disks_by_key(object) + .get_object_chunks(bucket, object, range, h, opts) + .await + } + fn get_hashed_set_index(&self, input: &str) -> usize { match self.distribution_algo { DistributionAlgoVersion::V1 => crc_hash(input, self.disk_set.len()), @@ -375,7 +388,13 @@ impl ObjectIO for Sets { .await } #[tracing::instrument(level = "debug", skip(self, data))] - async fn put_object(&self, bucket: &str, object: &str, data: &mut PutObjReader, opts: &ObjectOptions) -> Result { + async fn put_object( + &self, + bucket: &str, + object: &str, + data: &mut ChunkNativePutData, + opts: &ObjectOptions, + ) -> Result { self.get_disks_by_key(object).put_object(bucket, object, data, opts).await } } @@ -688,7 +707,7 @@ impl MultipartOperations for Sets { object: &str, upload_id: &str, part_id: usize, - data: &mut PutObjReader, + data: &mut ChunkNativePutData, opts: &ObjectOptions, ) -> Result { self.get_disks_by_key(object) diff --git a/crates/ecstore/src/store.rs b/crates/ecstore/src/store.rs index 07e16baa4..a90e19aef 100644 --- a/crates/ecstore/src/store.rs +++ b/crates/ecstore/src/store.rs @@ -59,9 +59,10 @@ use crate::{ rpc::S3PeerSys, sets::Sets, store_api::{ - BucketInfo, BucketOperations, BucketOptions, CompletePart, DeleteBucketOptions, DeletedObject, GetObjectReader, - HTTPRangeSpec, HealOperations, ListObjectsV2Info, ListOperations, MakeBucketOptions, MultipartOperations, - MultipartUploadResult, ObjectInfo, ObjectOperations, ObjectOptions, ObjectToDelete, PartInfo, PutObjReader, StorageAPI, + BucketInfo, BucketOperations, BucketOptions, ChunkNativePutData, CompletePart, DeleteBucketOptions, DeletedObject, + GetObjectChunkResult, GetObjectReader, HTTPRangeSpec, HealOperations, ListObjectsV2Info, ListOperations, + MakeBucketOptions, MultipartOperations, MultipartUploadResult, ObjectInfo, ObjectOperations, ObjectOptions, + ObjectToDelete, PartInfo, StorageAPI, }, store_init, }; @@ -259,11 +260,31 @@ impl ObjectIO for ECStore { self.handle_get_object_reader(bucket, object, range, h, opts).await } #[instrument(level = "debug", skip(self, data))] - async fn put_object(&self, bucket: &str, object: &str, data: &mut PutObjReader, opts: &ObjectOptions) -> Result { + async fn put_object( + &self, + bucket: &str, + object: &str, + data: &mut ChunkNativePutData, + opts: &ObjectOptions, + ) -> Result { enqueue_transition_after_write(self.handle_put_object(bucket, object, data, opts).await, LcEventSrc::S3PutObject).await } } +impl ECStore { + #[instrument(level = "debug", skip(self))] + pub async fn get_object_chunks( + &self, + bucket: &str, + object: &str, + range: Option, + h: HeaderMap, + opts: &ObjectOptions, + ) -> Result { + self.handle_get_object_chunks(bucket, object, range, h, opts).await + } +} + lazy_static! { static ref enableObjcetLockConfig: ObjectLockConfiguration = ObjectLockConfiguration { object_lock_enabled: Some(ObjectLockEnabled::from_static(ObjectLockEnabled::ENABLED)), @@ -495,7 +516,7 @@ impl MultipartOperations for ECStore { object: &str, upload_id: &str, part_id: usize, - data: &mut PutObjReader, + data: &mut ChunkNativePutData, opts: &ObjectOptions, ) -> Result { self.handle_put_object_part(bucket, object, upload_id, part_id, data, opts) diff --git a/crates/ecstore/src/store/multipart.rs b/crates/ecstore/src/store/multipart.rs index dfeb1c143..12ef5712f 100644 --- a/crates/ecstore/src/store/multipart.rs +++ b/crates/ecstore/src/store/multipart.rs @@ -176,7 +176,7 @@ impl ECStore { object: &str, upload_id: &str, part_id: usize, - data: &mut PutObjReader, + data: &mut ChunkNativePutData, opts: &ObjectOptions, ) -> Result { check_put_object_part_args(bucket, object, upload_id)?; diff --git a/crates/ecstore/src/store/object.rs b/crates/ecstore/src/store/object.rs index e648e0e02..e7f044b50 100644 --- a/crates/ecstore/src/store/object.rs +++ b/crates/ecstore/src/store/object.rs @@ -13,6 +13,7 @@ // limitations under the License. use super::*; +use crate::store_api::GetObjectChunkResult; fn select_data_movement_target_pool( existing_pool_idx: Result, @@ -213,12 +214,40 @@ impl ECStore { .await } + #[instrument(level = "debug", skip(self))] + pub(super) async fn handle_get_object_chunks( + &self, + bucket: &str, + object: &str, + range: Option, + h: HeaderMap, + opts: &ObjectOptions, + ) -> Result { + check_get_obj_args(bucket, object)?; + + let object = encode_dir_object(object); + + if self.single_pool() { + return self.pools[0].get_object_chunks(bucket, object.as_str(), range, h, opts).await; + } + + let mut opts = opts.clone(); + opts.no_lock = true; + + let (_, idx) = self + .get_latest_accessible_object_info_with_idx(bucket, &object, &opts) + .await?; + self.pools[idx] + .get_object_chunks(bucket, object.as_str(), range, h, &opts) + .await + } + #[instrument(level = "debug", skip(self, data))] pub(super) async fn handle_put_object( &self, bucket: &str, object: &str, - data: &mut PutObjReader, + data: &mut ChunkNativePutData, opts: &ObjectOptions, ) -> Result { check_put_object_args(bucket, object)?; diff --git a/crates/ecstore/src/store_api.rs b/crates/ecstore/src/store_api.rs index cdfde143b..cdaf86a1e 100644 --- a/crates/ecstore/src/store_api.rs +++ b/crates/ecstore/src/store_api.rs @@ -31,6 +31,7 @@ use rustfs_filemeta::{ RestoreStatusOps as _, VersionPurgeStatusType, parse_restore_obj_status, replication_statuses_map, version_purge_statuses_map, }; +use rustfs_io_core::BoxChunkStream; use rustfs_lock::NamespaceLockWrapper; use rustfs_madmin::heal_commands::HealResultItem; use rustfs_rio::Checksum; diff --git a/crates/ecstore/src/store_api/readers.rs b/crates/ecstore/src/store_api/readers.rs index 461e8ff7e..867d5d9ae 100644 --- a/crates/ecstore/src/store_api/readers.rs +++ b/crates/ecstore/src/store_api/readers.rs @@ -1,7 +1,101 @@ use super::*; +use rustfs_rio::TryGetIndex; + +pub struct ChunkNativePutData { + stream: Option, + size: i64, + actual_size: i64, +} + +impl Debug for ChunkNativePutData { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.debug_struct("ChunkNativePutData").finish() + } +} + +impl ChunkNativePutData { + pub fn new(stream: HashReader) -> Self { + let size = stream.size(); + let actual_size = stream.actual_size(); + Self { + stream: Some(stream), + size, + actual_size, + } + } + + pub fn from_vec(data: Vec) -> Self { + use sha2::{Digest, Sha256}; + + let content_length = data.len() as i64; + let sha256hex = if content_length > 0 { + Some(hex_simd::encode_to_string(Sha256::digest(&data), hex_simd::AsciiCase::Lower)) + } else { + None + }; + Self::new(HashReader::from_stream(Cursor::new(data), content_length, content_length, None, sha256hex, false).unwrap()) + } + + pub fn take_stream(&mut self) -> std::io::Result { + self.stream + .take() + .ok_or_else(|| std::io::Error::other("ChunkNativePutData stream already taken")) + } + + pub fn restore_stream(&mut self, stream: HashReader) { + self.size = stream.size(); + self.actual_size = stream.actual_size(); + self.stream = Some(stream); + } + + pub fn as_hash_reader(&self) -> Option<&HashReader> { + self.stream.as_ref() + } + + pub fn as_hash_reader_mut(&mut self) -> Option<&mut HashReader> { + self.stream.as_mut() + } + + pub fn index_bytes(&self) -> Option { + self.as_hash_reader() + .and_then(|reader| reader.try_get_index().map(|index| index.clone().into_vec())) + } + + pub fn resolve_etag(&mut self) -> Option { + self.as_hash_reader_mut() + .and_then(rustfs_rio::EtagResolvable::try_resolve_etag) + } + + pub fn content_hash_bytes(&mut self) -> std::io::Result> { + let Some(reader) = self.as_hash_reader_mut() else { + return Ok(None); + }; + + Ok(reader + .finalize_content_hash()? + .as_ref() + .map(|checksum| checksum.to_bytes(&[]))) + } + + pub fn content_crc_type(&self) -> Option { + self.as_hash_reader().and_then(HashReader::content_crc_type) + } + + pub fn content_crc(&self) -> HashMap { + self.as_hash_reader().map_or_else(HashMap::new, HashReader::content_crc) + } + + pub fn size(&self) -> i64 { + self.size + } + + pub fn actual_size(&self) -> i64 { + self.actual_size + } +} pub struct PutObjReader { - pub stream: HashReader, + data: ChunkNativePutData, } impl Debug for PutObjReader { @@ -12,32 +106,81 @@ impl Debug for PutObjReader { impl PutObjReader { pub fn new(stream: HashReader) -> Self { - PutObjReader { stream } + Self { + data: ChunkNativePutData::new(stream), + } } - pub fn as_hash_reader(&self) -> &HashReader { - &self.stream + pub fn chunk_native_data(&self) -> &ChunkNativePutData { + &self.data + } + + pub fn chunk_native_data_mut(&mut self) -> &mut ChunkNativePutData { + &mut self.data + } + + pub fn take_stream(&mut self) -> std::io::Result { + self.data.take_stream() + } + + pub fn restore_stream(&mut self, stream: HashReader) { + self.data.restore_stream(stream); + } + + pub fn as_hash_reader(&self) -> Option<&HashReader> { + self.data.as_hash_reader() + } + + pub fn as_hash_reader_mut(&mut self) -> Option<&mut HashReader> { + self.data.as_hash_reader_mut() + } + + pub fn index_bytes(&self) -> Option { + self.data.index_bytes() + } + + pub fn resolve_etag(&mut self) -> Option { + self.data.resolve_etag() + } + + pub fn content_hash_bytes(&mut self) -> std::io::Result> { + self.data.content_hash_bytes() + } + + pub fn content_crc_type(&self) -> Option { + self.data.content_crc_type() + } + + pub fn content_crc(&self) -> HashMap { + self.data.content_crc() } pub fn from_vec(data: Vec) -> Self { - use sha2::{Digest, Sha256}; - let content_length = data.len() as i64; - let sha256hex = if content_length > 0 { - Some(hex_simd::encode_to_string(Sha256::digest(&data), hex_simd::AsciiCase::Lower)) - } else { - None - }; - PutObjReader { - stream: HashReader::from_stream(Cursor::new(data), content_length, content_length, None, sha256hex, false).unwrap(), + Self { + data: ChunkNativePutData::from_vec(data), } } pub fn size(&self) -> i64 { - self.stream.size() + self.data.size() } pub fn actual_size(&self) -> i64 { - self.stream.actual_size() + self.data.actual_size() + } +} + +impl std::ops::Deref for PutObjReader { + type Target = ChunkNativePutData; + + fn deref(&self) -> &Self::Target { + &self.data + } +} + +impl std::ops::DerefMut for PutObjReader { + fn deref_mut(&mut self) -> &mut Self::Target { + &mut self.data } } @@ -46,6 +189,26 @@ pub struct GetObjectReader { pub object_info: ObjectInfo, } +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum GetObjectChunkPath { + Direct, + Bridge, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum GetObjectChunkCopyMode { + TrueZeroCopy, + SharedBytes, + SingleCopy, + Reconstructed, +} + +pub struct GetObjectChunkResult { + pub stream: BoxChunkStream, + pub path: GetObjectChunkPath, + pub copy_mode: GetObjectChunkCopyMode, +} + impl GetObjectReader { #[tracing::instrument(level = "debug", skip(reader, rs, opts, _h))] pub fn new( @@ -63,31 +226,34 @@ impl GetObjectReader { rs = HTTPRangeSpec::from_object_info(oi, part_number); } - // TODO:Encrypted + let logical_size = oi.get_actual_size()?; + let encrypted_object = oi.user_defined.contains_key("x-rustfs-encryption-key") + || oi + .user_defined + .contains_key("x-amz-server-side-encryption-customer-algorithm"); let (algo, is_compressed) = oi.is_compressed_ok()?; // TODO: check TRANSITION if is_compressed { - let actual_size = oi.get_actual_size()?; let (off, length, dec_off, dec_length) = if let Some(rs) = rs { // Support range requests for compressed objects - let (dec_off, dec_length) = rs.get_offset_length(actual_size)?; + let (dec_off, dec_length) = rs.get_offset_length(logical_size)?; (0, oi.size, dec_off, dec_length) } else { - (0, oi.size, 0, actual_size) + (0, oi.size, 0, logical_size) }; let dec_reader = DecompressReader::new(reader, algo); - let actual_size_usize = if actual_size > 0 { - actual_size as usize + let actual_size_usize = if logical_size > 0 { + logical_size as usize } else { - return Err(Error::other(format!("invalid decompressed size {actual_size}"))); + return Err(Error::other(format!("invalid decompressed size {logical_size}"))); }; - let final_reader: Box = if dec_off > 0 || dec_length != actual_size { + let final_reader: Box = if dec_off > 0 || dec_length != logical_size { // Use RangedDecompressReader for streaming range processing // The new implementation supports any offset size by streaming and skipping data match RangedDecompressReader::new(dec_reader, dec_off, dec_length, actual_size_usize) { @@ -122,8 +288,19 @@ impl GetObjectReader { )); } + if encrypted_object && rs.is_none() { + return Ok(( + GetObjectReader { + stream: reader, + object_info: oi.clone(), + }, + 0, + oi.size, + )); + } + if let Some(rs) = rs { - let (off, length) = rs.get_offset_length(oi.size)?; + let (off, length) = rs.get_offset_length(logical_size)?; Ok(( GetObjectReader { @@ -140,7 +317,7 @@ impl GetObjectReader { object_info: oi.clone(), }, 0, - oi.size, + logical_size, )) } } diff --git a/crates/ecstore/src/store_api/traits.rs b/crates/ecstore/src/store_api/traits.rs index facabeae3..a351976e8 100644 --- a/crates/ecstore/src/store_api/traits.rs +++ b/crates/ecstore/src/store_api/traits.rs @@ -11,7 +11,13 @@ pub trait ObjectIO: Send + Sync + Debug + 'static { opts: &ObjectOptions, ) -> Result; - async fn put_object(&self, bucket: &str, object: &str, data: &mut PutObjReader, opts: &ObjectOptions) -> Result; + async fn put_object( + &self, + bucket: &str, + object: &str, + data: &mut ChunkNativePutData, + opts: &ObjectOptions, + ) -> Result; } /// Bucket-level storage operations. @@ -126,7 +132,7 @@ pub trait MultipartOperations: Send + Sync + Debug { object: &str, upload_id: &str, part_id: usize, - data: &mut PutObjReader, + data: &mut ChunkNativePutData, opts: &ObjectOptions, ) -> Result; async fn get_multipart_info( diff --git a/crates/ecstore/src/store_api/types.rs b/crates/ecstore/src/store_api/types.rs index 829f81ce5..efbcf396d 100644 --- a/crates/ecstore/src/store_api/types.rs +++ b/crates/ecstore/src/store_api/types.rs @@ -70,6 +70,7 @@ pub struct ObjectOptions { pub eval_metadata: Option>, + pub resolved_checksum: Option, pub want_checksum: Option, pub skip_verify_bitrot: bool, } @@ -283,7 +284,7 @@ pub struct ObjectInfo { pub expires: Option, pub num_versions: usize, pub successor_mod_time: Option, - pub put_object_reader: Option, + pub put_object_reader: Option, pub etag: Option, pub inlined: bool, pub metadata_only: bool, @@ -509,6 +510,18 @@ impl ObjectInfo { }) .collect(); + let actual_size = fi + .parts + .iter() + .map(|part| { + if part.actual_size > 0 { + part.actual_size + } else { + i64::try_from(part.size).unwrap_or_default() + } + }) + .sum(); + // TODO: part checksums ObjectInfo { @@ -521,6 +534,7 @@ impl ObjectInfo { delete_marker: fi.deleted, mod_time: fi.mod_time, size: fi.size, + actual_size, parts, is_latest: fi.is_latest, user_tags, diff --git a/crates/ecstore/src/tier/tier.rs b/crates/ecstore/src/tier/tier.rs index 8cf6e6062..97b546f38 100644 --- a/crates/ecstore/src/tier/tier.rs +++ b/crates/ecstore/src/tier/tier.rs @@ -51,7 +51,7 @@ use crate::{ disk::{MIGRATING_META_BUCKET, RUSTFS_META_BUCKET}, global::is_first_cluster_node_local, store::ECStore, - store_api::{ObjectIO as _, ObjectOptions, PutObjReader}, + store_api::{ChunkNativePutData, ObjectIO as _, ObjectOptions}, }; use rustfs_rio::HashReader; use rustfs_utils::path::{SLASH_SEPARATOR, path_join}; @@ -1046,9 +1046,8 @@ impl TierConfigMgr { opts: &ObjectOptions, ) -> std::result::Result<(), std::io::Error> { debug!("save tier config:{}", file); - let _ = api - .put_object(RUSTFS_META_BUCKET, file, &mut PutObjReader::from_vec(data.to_vec()), opts) - .await?; + let mut put_data = ChunkNativePutData::from_vec(data.to_vec()); + let _ = api.put_object(RUSTFS_META_BUCKET, file, &mut put_data, opts).await?; Ok(()) } @@ -1104,11 +1103,12 @@ async fn load_tier_config(api: Arc) -> std::result::Result { let cfg = TierConfigMgr::unmarshal(&data)?; let normalized = encode_external_tiering_config_blob(&cfg)?; + let mut put_data = ChunkNativePutData::from_vec(normalized.to_vec()); let _ = api .put_object( RUSTFS_META_BUCKET, &config_file, - &mut PutObjReader::from_vec(normalized.to_vec()), + &mut put_data, &ObjectOptions { max_parity: true, ..Default::default() @@ -1158,10 +1158,11 @@ async fn read_tier_config_from_bucket( } async fn write_tier_config_to_rustfs(api: Arc, path: &str, data: Bytes) -> io::Result<()> { + let mut put_data = ChunkNativePutData::from_vec(data.to_vec()); api.put_object( RUSTFS_META_BUCKET, path, - &mut PutObjReader::from_vec(data.to_vec()), + &mut put_data, &ObjectOptions { max_parity: true, ..Default::default() diff --git a/crates/filemeta/src/filemeta.rs b/crates/filemeta/src/filemeta.rs index 33eee875a..80e5f4653 100644 --- a/crates/filemeta/src/filemeta.rs +++ b/crates/filemeta/src/filemeta.rs @@ -24,7 +24,7 @@ use rustfs_utils::http::headers::{ AMZ_STORAGE_CLASS, }; use rustfs_utils::http::{ - AMZ_BUCKET_REPLICATION_STATUS, SUFFIX_DATA_MOV, SUFFIX_HEALING, SUFFIX_PURGESTATUS, SUFFIX_REPLICA_STATUS, + AMZ_BUCKET_REPLICATION_STATUS, SUFFIX_CRC, SUFFIX_DATA_MOV, SUFFIX_HEALING, SUFFIX_PURGESTATUS, SUFFIX_REPLICA_STATUS, SUFFIX_REPLICA_TIMESTAMP, SUFFIX_REPLICATION_STATUS, SUFFIX_REPLICATION_TIMESTAMP, has_internal_suffix, insert_bytes, is_internal_key, }; @@ -230,6 +230,10 @@ impl FileMeta { } } + if let Some(checksum) = fi.checksum.as_ref() { + insert_bytes(&mut obj.meta_sys, SUFFIX_CRC, checksum.to_vec()); + } + if let Some(mod_time) = fi.mod_time { obj.mod_time = Some(mod_time); } diff --git a/crates/heal/src/heal/storage.rs b/crates/heal/src/heal/storage.rs index 33f6210bc..e1699c192 100644 --- a/crates/heal/src/heal/storage.rs +++ b/crates/heal/src/heal/storage.rs @@ -208,7 +208,7 @@ impl HealStorageAPI for ECStoreHealStorage { async fn put_object_data(&self, bucket: &str, object: &str, data: &[u8]) -> Result<()> { debug!("Putting object data: {}/{} ({} bytes)", bucket, object, data.len()); - let mut reader = rustfs_ecstore::store_api::PutObjReader::from_vec(data.to_vec()); + let mut reader = rustfs_ecstore::store_api::ChunkNativePutData::from_vec(data.to_vec()); match (*self.ecstore) .put_object(bucket, object, &mut reader, &Default::default()) .await diff --git a/crates/heal/tests/heal_integration_test.rs b/crates/heal/tests/heal_integration_test.rs index e7ba323ef..765221408 100644 --- a/crates/heal/tests/heal_integration_test.rs +++ b/crates/heal/tests/heal_integration_test.rs @@ -18,7 +18,7 @@ use rustfs_ecstore::{ disk::endpoint::Endpoint, endpoints::{EndpointServerPools, Endpoints, PoolEndpoints}, store::ECStore, - store_api::{BucketOperations, ObjectIO, ObjectOperations, ObjectOptions, PutObjReader}, + store_api::{BucketOperations, ChunkNativePutData, ObjectIO, ObjectOperations, ObjectOptions}, }; use rustfs_heal::heal::{ manager::{HealConfig, HealManager}, @@ -162,7 +162,7 @@ async fn create_test_bucket(ecstore: &Arc, bucket_name: &str) { /// Test helper: Upload test object async fn upload_test_object(ecstore: &Arc, bucket: &str, object: &str, data: &[u8]) { - let mut reader = PutObjReader::from_vec(data.to_vec()); + let mut reader = ChunkNativePutData::from_vec(data.to_vec()); let object_info = (**ecstore) .put_object(bucket, object, &mut reader, &ObjectOptions::default()) .await diff --git a/crates/io-core/Cargo.toml b/crates/io-core/Cargo.toml index 81f932532..fd4fb14b5 100644 --- a/crates/io-core/Cargo.toml +++ b/crates/io-core/Cargo.toml @@ -29,12 +29,14 @@ workspace = true [dependencies] bytes = { workspace = true } +futures-core = { workspace = true } thiserror = { workspace = true } tokio = { workspace = true, features = ["io-util", "fs", "rt", "sync"] } memmap2 = { workspace = true } rustfs-io-metrics = { workspace = true } [dev-dependencies] +futures-util = { workspace = true } tokio = { workspace = true, features = ["rt-multi-thread", "macros"] } [lib] diff --git a/crates/io-core/src/adapter.rs b/crates/io-core/src/adapter.rs new file mode 100644 index 000000000..b2548382c --- /dev/null +++ b/crates/io-core/src/adapter.rs @@ -0,0 +1,124 @@ +// 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. + +//! Compatibility adapter that exposes chunk streams as `AsyncRead`. + +use crate::chunk::BoxChunkStream; +use bytes::Bytes; +use std::io; +use std::pin::Pin; +use std::task::{Context, Poll}; +use tokio::io::{AsyncRead, ReadBuf}; + +/// `AsyncRead` adapter for boxed chunk streams. +pub struct ChunkStreamReader { + stream: BoxChunkStream, + current: Option, + offset: usize, +} + +impl ChunkStreamReader { + #[must_use] + pub fn new(stream: BoxChunkStream) -> Self { + Self { + stream, + current: None, + offset: 0, + } + } +} + +impl AsyncRead for ChunkStreamReader { + fn poll_read(mut self: Pin<&mut Self>, cx: &mut Context<'_>, buf: &mut ReadBuf<'_>) -> Poll> { + loop { + if let Some(current) = &self.current { + if self.offset < current.len() { + let remaining = ¤t[self.offset..]; + let to_read = remaining.len().min(buf.remaining()); + buf.put_slice(&remaining[..to_read]); + self.offset += to_read; + return Poll::Ready(Ok(())); + } + + self.current = None; + self.offset = 0; + continue; + } + + match self.stream.as_mut().poll_next(cx) { + Poll::Pending => return Poll::Pending, + Poll::Ready(Some(Ok(chunk))) => { + let next = chunk.as_bytes(); + if next.is_empty() { + continue; + } + self.current = Some(next); + self.offset = 0; + } + Poll::Ready(Some(Err(err))) => return Poll::Ready(Err(err)), + Poll::Ready(None) => return Poll::Ready(Ok(())), + } + } + } +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::chunk::{IoChunk, MappedChunk, PooledChunk}; + use bytes::Bytes; + use futures_util::stream; + use tokio::io::AsyncReadExt; + + #[tokio::test] + async fn test_chunk_stream_reader_reads_single_chunk() { + let stream: BoxChunkStream = Box::pin(stream::iter(vec![Ok(IoChunk::Shared(Bytes::from_static(b"hello")))])); + let mut reader = ChunkStreamReader::new(stream); + let mut out = Vec::new(); + reader.read_to_end(&mut out).await.unwrap(); + assert_eq!(out, b"hello"); + } + + #[tokio::test] + async fn test_chunk_stream_reader_reads_multiple_chunks() { + let stream: BoxChunkStream = Box::pin(stream::iter(vec![ + Ok(IoChunk::Shared(Bytes::from_static(b"he"))), + Ok(IoChunk::Mapped(MappedChunk::new(Bytes::from_static(b"llo!"), 0, 4).unwrap())), + Ok(IoChunk::Pooled(PooledChunk::from_bytes(Bytes::from_static(b" world")).unwrap())), + ])); + let mut reader = ChunkStreamReader::new(stream); + let mut out = Vec::new(); + reader.read_to_end(&mut out).await.unwrap(); + assert_eq!(out, b"hello! world"); + } + + #[tokio::test] + async fn test_chunk_stream_reader_handles_empty_stream() { + let stream: BoxChunkStream = Box::pin(stream::iter(Vec::>::new())); + let mut reader = ChunkStreamReader::new(stream); + let mut out = Vec::new(); + reader.read_to_end(&mut out).await.unwrap(); + assert!(out.is_empty()); + } + + #[tokio::test] + async fn test_chunk_stream_reader_propagates_stream_error() { + let stream: BoxChunkStream = Box::pin(stream::iter(vec![Err(io::Error::other("chunk stream failure"))])); + let mut reader = ChunkStreamReader::new(stream); + let mut out = Vec::new(); + let err = reader.read_to_end(&mut out).await.unwrap_err(); + assert_eq!(err.kind(), io::ErrorKind::Other); + assert!(err.to_string().contains("chunk stream failure")); + } +} diff --git a/crates/io-core/src/chunk.rs b/crates/io-core/src/chunk.rs new file mode 100644 index 000000000..4a3441cc5 --- /dev/null +++ b/crates/io-core/src/chunk.rs @@ -0,0 +1,276 @@ +// 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. + +//! Core chunk ownership abstractions for the zero-copy data plane. + +use crate::pool::PooledBuffer; +use bytes::Bytes; +use futures_core::Stream; +use std::io; +use std::pin::Pin; + +/// Boxed asynchronous stream of I/O chunks. +pub type BoxChunkStream = Pin> + Send + Sync + 'static>>; + +/// Source of chunked data. +pub trait ChunkSource { + fn into_chunk_stream(self) -> BoxChunkStream + where + Self: Sized; +} + +/// Owned chunk variants used by the zero-copy data plane. +#[derive(Debug)] +pub enum IoChunk { + Shared(Bytes), + Mapped(MappedChunk), + Pooled(PooledChunk), +} + +impl IoChunk { + /// Returns the visible length of this chunk. + #[must_use] + pub fn len(&self) -> usize { + match self { + Self::Shared(bytes) => bytes.len(), + Self::Mapped(chunk) => chunk.len(), + Self::Pooled(chunk) => chunk.len(), + } + } + + /// Returns true when the chunk has no visible data. + #[must_use] + pub fn is_empty(&self) -> bool { + self.len() == 0 + } + + /// Returns a shared read-only view of the visible bytes. + #[must_use] + pub fn as_bytes(&self) -> Bytes { + match self { + Self::Shared(bytes) => bytes.clone(), + Self::Mapped(chunk) => chunk.as_bytes(), + Self::Pooled(chunk) => chunk.as_bytes(), + } + } + + /// Returns a sliced view relative to the currently visible bytes. + pub fn slice(&self, offset: usize, len: usize) -> io::Result { + match self { + Self::Shared(bytes) => { + validate_slice_bounds(bytes.len(), offset, len)?; + Ok(Self::Shared(bytes.slice(offset..offset + len))) + } + Self::Mapped(chunk) => chunk.slice(offset, len).map(Self::Mapped), + Self::Pooled(chunk) => chunk.slice(offset, len).map(Self::Pooled), + } + } +} + +/// Logical view into mapped file bytes. +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct MappedChunk { + bytes: Bytes, + logical_offset: usize, + logical_len: usize, +} + +impl MappedChunk { + pub fn new(bytes: Bytes, logical_offset: usize, logical_len: usize) -> io::Result { + validate_slice_bounds(bytes.len(), logical_offset, logical_len)?; + Ok(Self { + bytes, + logical_offset, + logical_len, + }) + } + + /// Returns the visible length of this mapped chunk. + #[must_use] + pub const fn len(&self) -> usize { + self.logical_len + } + + /// Returns true when the chunk has no visible data. + #[must_use] + pub const fn is_empty(&self) -> bool { + self.logical_len == 0 + } + + /// Returns the visible bytes for this logical view. + #[must_use] + pub fn as_bytes(&self) -> Bytes { + self.bytes + .slice(self.logical_offset..self.logical_offset.saturating_add(self.logical_len)) + } + + /// Returns a sliced logical view relative to the current logical view. + pub fn slice(&self, offset: usize, len: usize) -> io::Result { + validate_slice_bounds(self.logical_len, offset, len)?; + Self::new(self.bytes.clone(), self.logical_offset + offset, len) + } +} + +/// Placeholder pooled chunk variant for commit 4. +/// +/// This is backed by a `PooledBuffer` and exposes a visible read-only window. +#[derive(Debug)] +pub struct PooledChunk { + bytes: Bytes, +} + +#[derive(Debug)] +struct PooledChunkOwner { + buffer: PooledBuffer, + visible_len: usize, +} + +impl AsRef<[u8]> for PooledChunkOwner { + fn as_ref(&self) -> &[u8] { + &self.buffer[..self.visible_len] + } +} + +#[derive(Debug)] +struct DetachedVecChunkOwner { + bytes: Vec, +} + +impl AsRef<[u8]> for DetachedVecChunkOwner { + fn as_ref(&self) -> &[u8] { + &self.bytes + } +} + +impl PooledChunk { + pub fn new(buffer: PooledBuffer, len: usize) -> io::Result { + validate_slice_bounds(buffer.len(), 0, len)?; + Ok(Self { + bytes: Bytes::from_owner(PooledChunkOwner { + buffer, + visible_len: len, + }), + }) + } + + /// Convenience constructor for detached test and compatibility values. + pub fn from_bytes(bytes: Bytes) -> io::Result { + let len = bytes.len(); + Self::new(PooledBuffer::from_bytes(bytes), len) + } + + /// Detached constructor that takes ownership of an existing `Vec` + /// without introducing an additional copy. + pub fn from_vec(bytes: Vec) -> Self { + Self { + bytes: Bytes::from_owner(DetachedVecChunkOwner { bytes }), + } + } + + /// Returns the visible length of this pooled chunk. + #[must_use] + pub fn len(&self) -> usize { + self.bytes.len() + } + + /// Returns true when the chunk has no visible data. + #[must_use] + pub fn is_empty(&self) -> bool { + self.bytes.is_empty() + } + + /// Returns the visible bytes for this pooled chunk. + #[must_use] + pub fn as_bytes(&self) -> Bytes { + self.bytes.clone() + } + + /// Returns a sliced pooled chunk relative to the current visible view. + pub fn slice(&self, offset: usize, len: usize) -> io::Result { + validate_slice_bounds(self.bytes.len(), offset, len)?; + Ok(Self { + bytes: self.bytes.slice(offset..offset + len), + }) + } +} + +fn validate_slice_bounds(visible_len: usize, offset: usize, len: usize) -> io::Result<()> { + let end = offset + .checked_add(len) + .ok_or_else(|| io::Error::new(io::ErrorKind::InvalidInput, "chunk slice overflows"))?; + if end > visible_len { + return Err(io::Error::new(io::ErrorKind::InvalidInput, "chunk slice exceeds visible length")); + } + Ok(()) +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::pool::BytesPool; + + #[test] + fn test_shared_chunk_len_and_slice() { + let chunk = IoChunk::Shared(Bytes::from_static(b"abcdef")); + assert_eq!(chunk.len(), 6); + assert!(!chunk.is_empty()); + assert_eq!(chunk.as_bytes(), Bytes::from_static(b"abcdef")); + assert_eq!(chunk.slice(1, 3).unwrap().as_bytes(), Bytes::from_static(b"bcd")); + } + + #[test] + fn test_mapped_chunk_len_and_slice() { + let chunk = MappedChunk::new(Bytes::from_static(b"abcdefgh"), 2, 4).unwrap(); + assert_eq!(chunk.len(), 4); + assert_eq!(chunk.as_bytes(), Bytes::from_static(b"cdef")); + assert_eq!(chunk.slice(1, 2).unwrap().as_bytes(), Bytes::from_static(b"de")); + } + + #[test] + fn test_pooled_chunk_len_and_as_bytes() { + let chunk = PooledChunk::from_bytes(Bytes::from_static(b"hello")).unwrap(); + assert_eq!(chunk.len(), 5); + assert_eq!(chunk.as_bytes(), Bytes::from_static(b"hello")); + assert_eq!(chunk.slice(1, 3).unwrap().as_bytes(), Bytes::from_static(b"ell")); + } + + #[test] + fn test_io_chunk_as_bytes_for_all_variants() { + let shared = IoChunk::Shared(Bytes::from_static(b"s")); + let mapped = IoChunk::Mapped(MappedChunk::new(Bytes::from_static(b"mapped"), 0, 6).unwrap()); + let pooled = IoChunk::Pooled(PooledChunk::from_bytes(Bytes::from_static(b"p")).unwrap()); + + assert_eq!(shared.as_bytes(), Bytes::from_static(b"s")); + assert_eq!(mapped.as_bytes(), Bytes::from_static(b"mapped")); + assert_eq!(pooled.as_bytes(), Bytes::from_static(b"p")); + } + + #[tokio::test] + async fn test_pooled_chunk_keeps_owner_alive_until_last_view_drops() { + let pool = BytesPool::new_tiered(); + let mut buffer = pool.acquire_buffer(16).await; + buffer.extend_from_slice(b"pooled-bytes"); + + let chunk = PooledChunk::new(buffer, "pooled-bytes".len()).unwrap(); + let bytes = chunk.as_bytes(); + + assert_eq!(pool.available_buffers(), 0); + drop(chunk); + assert_eq!(pool.available_buffers(), 0); + assert_eq!(bytes, Bytes::from_static(b"pooled-bytes")); + + drop(bytes); + assert_eq!(pool.available_buffers(), 1); + } +} diff --git a/crates/io-core/src/lib.rs b/crates/io-core/src/lib.rs index fe0ee031f..f8565d50e 100644 --- a/crates/io-core/src/lib.rs +++ b/crates/io-core/src/lib.rs @@ -46,8 +46,10 @@ //! let mut buffer = pool.acquire_buffer(8192).await; //! ``` +pub mod adapter; pub mod backpressure; pub mod bufreader_optimizer; +pub mod chunk; pub mod config; pub mod deadlock_detector; pub mod direct_io; @@ -68,7 +70,9 @@ pub use reader::{ZeroCopyObjectReader, ZeroCopyReadError}; pub use writer::{ZeroCopyObjectWriter, ZeroCopyWriteError}; // BufReader optimizer exports +pub use adapter::ChunkStreamReader; pub use bufreader_optimizer::{BufReaderConfig, BufReaderOptimizer, BufReaderStats, BufferedSource}; +pub use chunk::{BoxChunkStream, ChunkSource, IoChunk, MappedChunk, PooledChunk}; // Shared memory exports pub use shared_memory::{ArcData, ArcMetadata, SharedMemoryConfig, SharedMemoryPool, SharedMemoryStats}; diff --git a/crates/io-core/src/pool.rs b/crates/io-core/src/pool.rs index aa08fcfde..fc5e642f3 100644 --- a/crates/io-core/src/pool.rs +++ b/crates/io-core/src/pool.rs @@ -17,7 +17,7 @@ //! Migrated from rustfs-ecstore to provide unified buffer pooling //! across rustfs and rustfs-ecstore without cyclic dependencies. -use bytes::BytesMut; +use bytes::{Bytes, BytesMut}; use std::mem::ManuallyDrop; use std::sync::atomic::{AtomicU64, Ordering}; use std::sync::{Arc, Mutex}; @@ -108,6 +108,7 @@ pub struct BytesPoolMetrics { /// A buffer managed by the BytesPool. /// /// When dropped, the buffer is automatically returned to the pool for reuse. +#[derive(Debug)] pub struct PooledBuffer { /// The underlying buffer (ManuallyDrop to allow taking on drop) pub buffer: ManuallyDrop, @@ -117,6 +118,45 @@ pub struct PooledBuffer { _permit: Option, } +impl PooledBuffer { + /// Create a detached pooled buffer from bytes. + /// + /// This is primarily used for tests and transitional adapters where the + /// chunk abstraction needs a pool-shaped owner before a real pool-backed + /// producer exists. + #[must_use] + pub fn from_bytes(bytes: Bytes) -> Self { + Self { + buffer: ManuallyDrop::new(BytesMut::from(bytes.as_ref())), + tier: None, + _permit: None, + } + } + + /// Current visible length of the underlying buffer. + #[must_use] + pub fn len(&self) -> usize { + self.buffer.len() + } + + /// Total buffer capacity. + #[must_use] + pub fn capacity(&self) -> usize { + self.buffer.capacity() + } + + /// Clear the visible contents while preserving capacity. + pub fn clear(&mut self) { + self.buffer.clear(); + } + + /// Returns true when the visible buffer is empty. + #[must_use] + pub fn is_empty(&self) -> bool { + self.buffer.is_empty() + } +} + /// BytesPool configuration. /// /// Allows customization of buffer sizes and limits for each tier. diff --git a/crates/io-metrics/src/lib.rs b/crates/io-metrics/src/lib.rs index 3a0f1ce46..345a9d98a 100644 --- a/crates/io-metrics/src/lib.rs +++ b/crates/io-metrics/src/lib.rs @@ -118,8 +118,129 @@ pub use config::{ // Re-exports for convenience pub use collector::MetricsCollector; +pub use metric_names::data_plane; pub use performance::PerformanceMetrics; +/// High-level request path selected for an I/O operation. +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum IoPath { + Fast, + Legacy, +} + +impl IoPath { + #[must_use] + pub const fn as_str(self) -> &'static str { + match self { + Self::Fast => "fast", + Self::Legacy => "legacy", + } + } +} + +/// Effective copy mode observed for an I/O operation. +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum CopyMode { + TrueZeroCopy, + SharedBytes, + SingleCopy, + Reconstructed, + Transformed, +} + +impl CopyMode { + #[must_use] + pub const fn as_str(self) -> &'static str { + match self { + Self::TrueZeroCopy => "true_zero_copy", + Self::SharedBytes => "shared_bytes", + Self::SingleCopy => "single_copy", + Self::Reconstructed => "reconstructed", + Self::Transformed => "transformed", + } + } +} + +/// Stage where a data plane decision or fallback happened. +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum IoStage { + Unknown, + ReadSetup, + HttpBridge, + CacheWriteback, + LocalDiskChunk, + RangeGuard, + PutTransform, +} + +impl IoStage { + #[must_use] + pub const fn as_str(self) -> &'static str { + match self { + Self::Unknown => "unknown", + Self::ReadSetup => "read_setup", + Self::HttpBridge => "http_bridge", + Self::CacheWriteback => "cache_writeback", + Self::LocalDiskChunk => "local_disk_chunk", + Self::RangeGuard => "range_guard", + Self::PutTransform => "put_transform", + } + } +} + +/// Reason why the data plane fell back from a preferred path. +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum FallbackReason { + Unknown, + MmapDisabled, + MmapUnavailable, + SmallObject, + WindowLimitExceeded, + UnalignedWindow, + RangeNotSupported, + EncryptionEnabled, + CompressionEnabled, + TransformEncryptionLegacy, + TransformCompressionLegacy, + TransformCompressionEncryptionLegacy, + ChunkBridgeUnavailable, + NonLocalBackend, +} + +impl FallbackReason { + #[must_use] + pub const fn as_str(self) -> &'static str { + match self { + Self::Unknown => "unknown", + Self::MmapDisabled => "mmap_disabled", + Self::MmapUnavailable => "mmap_unavailable", + Self::SmallObject => "small_object", + Self::WindowLimitExceeded => "window_limit_exceeded", + Self::UnalignedWindow => "unaligned_window", + Self::RangeNotSupported => "range_not_supported", + Self::EncryptionEnabled => "encryption_enabled", + Self::CompressionEnabled => "compression_enabled", + Self::TransformEncryptionLegacy => "transform_encryption_legacy", + Self::TransformCompressionLegacy => "transform_compression_legacy", + Self::TransformCompressionEncryptionLegacy => "transform_compression_encryption_legacy", + Self::ChunkBridgeUnavailable => "chunk_bridge_unavailable", + Self::NonLocalBackend => "non_local_backend", + } + } +} + +#[inline(always)] +fn put_size_bucket_label(size_bytes: i64) -> &'static str { + match size_bytes { + ..=0 => "unknown", + 1..=16_384 => "le_16kib", + 16_385..=65_536 => "le_64kib", + 65_537..=262_144 => "le_256kib", + 262_145..=1_048_576 => "le_1mib", + _ => "gt_1mib", + } +} + /// Record GetObject request start. #[inline(always)] pub fn record_get_object_request_start(concurrent_requests: usize) { @@ -206,40 +327,138 @@ pub fn record_object_cache_writeback() { counter!("rustfs_io_object_cache_writeback_total").increment(1); } -/// Record a zero-copy read operation. -/// -/// # Arguments -/// -/// * `size_bytes` - Size of the data read in bytes -/// * `duration_ms` - Time taken for the read operation in milliseconds +/// Record which request path was selected for an operation. #[inline(always)] -pub fn record_zero_copy_read(size_bytes: usize, duration_ms: f64) { - counter!("rustfs.zero_copy.reads.total").increment(1); - histogram!("rustfs.zero_copy.read.size.bytes").record(size_bytes as f64); - histogram!("rustfs.zero_copy.read.duration.ms").record(duration_ms); +pub fn record_io_path_selected(operation: &'static str, io_path: IoPath) { + counter!( + metric_names::data_plane::PATH_SELECTED_TOTAL, + "path" => operation, + "mode" => io_path.as_str() + ) + .increment(1); } -/// Record memory copies avoided by using zero-copy. -/// -/// # Arguments -/// -/// * `bytes_saved` - Number of bytes that would have been copied without zero-copy +/// Record the effective copy mode for an operation. #[inline(always)] -pub fn record_memory_copy_saved(bytes_saved: usize) { - counter!("rustfs.zero_copy.memory.saved.bytes").increment(bytes_saved as u64); +pub fn record_io_copy_mode(operation: &'static str, copy_mode: CopyMode, size_bytes: usize) { + counter!( + metric_names::data_plane::COPY_MODE_BYTES_TOTAL, + "path" => operation, + "mode" => copy_mode.as_str() + ) + .increment(size_bytes as u64); } -/// Record a fallback from zero-copy to regular read. -/// -/// This happens when zero-copy read fails (e.g., mmap not available, -/// file too large, etc.) and the system falls back to regular I/O. -/// -/// # Arguments -/// -/// * `reason` - Reason for the fallback (e.g., "mmap_unavailable", "file_too_large") +/// Record a data plane fallback decision. #[inline(always)] -pub fn record_zero_copy_fallback(reason: &str) { - counter!("rustfs.zero_copy.fallback.total", "reason" => reason.to_string()).increment(1); +pub fn record_io_fallback(stage: IoStage, reason: FallbackReason) { + counter!( + metric_names::data_plane::FALLBACK_TOTAL, + "stage" => stage.as_str(), + "reason" => reason.as_str() + ) + .increment(1); +} + +/// Record the currently active mmap bytes held by LocalDisk chunk streams. +#[inline(always)] +pub fn record_local_disk_active_mmap_bytes(active_bytes: usize) { + gauge!(metric_names::data_plane::LOCAL_DISK_ACTIVE_MMAP_BYTES).set(active_bytes as f64); +} + +/// Record pooled chunk usage in LocalDisk compatibility paths. +#[inline(always)] +pub fn record_local_disk_pooled_chunk(source: &'static str, size_bytes: usize) { + counter!( + metric_names::data_plane::LOCAL_DISK_POOLED_CHUNKS_TOTAL, + "source" => source + ) + .increment(1); + counter!( + metric_names::data_plane::LOCAL_DISK_POOLED_BYTES_TOTAL, + "source" => source + ) + .increment(size_bytes as u64); +} + +/// Record a compatibility chunk-stream aggregation performed by `read_file_zero_copy()`. +#[inline(always)] +pub fn record_local_disk_compat_collect(chunk_count: usize, total_bytes: usize) { + counter!(metric_names::data_plane::LOCAL_DISK_COMPAT_COLLECT_TOTAL).increment(1); + histogram!(metric_names::data_plane::LOCAL_DISK_COMPAT_COLLECT_CHUNKS).record(chunk_count as f64); + histogram!(metric_names::data_plane::LOCAL_DISK_COMPAT_COLLECT_BYTES).record(total_bytes as f64); +} + +/// Record an attempted PUT fast path. +#[inline(always)] +pub fn record_put_object_attempted_fast_path(size_bytes: i64) { + counter!(metric_names::data_plane::PUT_FAST_PATH_ATTEMPTS_TOTAL).increment(1); + + if size_bytes > 0 { + histogram!(metric_names::data_plane::PUT_FAST_PATH_ATTEMPT_SIZE_BYTES).record(size_bytes as f64); + } +} + +/// Record which transformed PUT pipeline was selected. +#[inline(always)] +pub fn record_put_transform_selected(kind: &'static str, io_path: IoPath, size_bytes: usize) { + counter!( + metric_names::data_plane::PUT_TRANSFORM_SELECTED_TOTAL, + "kind" => kind, + "mode" => io_path.as_str() + ) + .increment(1); + + histogram!( + metric_names::data_plane::PUT_TRANSFORM_SIZE_BYTES, + "kind" => kind, + "mode" => io_path.as_str() + ) + .record(size_bytes as f64); +} + +/// Record PUT path selection with size-bucket context. +#[inline(always)] +pub fn record_put_path_selected(size_bytes: i64, io_path: IoPath) { + counter!( + "rustfs.s3.put_object.path.selected.total", + "mode" => io_path.as_str(), + "size_bucket" => put_size_bucket_label(size_bytes) + ) + .increment(1); +} + +/// Record PUT copy mode with size-bucket context. +#[inline(always)] +pub fn record_put_copy_mode(size_bytes: i64, copy_mode: CopyMode) { + counter!( + "rustfs.s3.put_object.copy_mode.total", + "mode" => copy_mode.as_str(), + "size_bucket" => put_size_bucket_label(size_bytes) + ) + .increment(1); +} + +/// Record PUT fallback with size-bucket context. +#[inline(always)] +pub fn record_put_fallback(size_bytes: i64, reason: FallbackReason) { + counter!( + "rustfs.s3.put_object.fallback.total", + "reason" => reason.as_str(), + "size_bucket" => put_size_bucket_label(size_bytes) + ) + .increment(1); +} + +/// Record inline-object selection for PUT with size-bucket context. +#[inline(always)] +pub fn record_put_inline_selected(size_bytes: i64, versioned: bool) { + counter!( + "rustfs.s3.put_object.inline.selected.total", + "versioned" => if versioned { "true" } else { "false" }, + "size_bucket" => put_size_bucket_label(size_bytes) + ) + .increment(1); } // ============================================================================ @@ -297,41 +516,6 @@ pub fn record_bytes_pool_hit_rate(tier: &str, hit_rate: f64) { gauge!("rustfs.bytes.pool.hit.rate", "tier" => tier.to_string()).set(hit_rate * 100.0); } -/// Record zero-copy write operation. -/// -/// # Arguments -/// -/// * `size_bytes` - Size of the data written in bytes -/// * `duration_ms` - Time taken for the write operation in milliseconds -#[inline(always)] -pub fn record_zero_copy_write(size_bytes: usize, duration_ms: f64) { - counter!("rustfs.zero_copy.write.total").increment(1); - histogram!("rustfs.zero_copy.write.size.bytes").record(size_bytes as f64); - histogram!("rustfs.zero_copy.write.duration.ms").record(duration_ms); -} - -/// Record zero-copy write fallback. -/// -/// This happens when zero-copy write fails and the system falls back to regular I/O. -/// -/// # Arguments -/// -/// * `reason` - Reason for the fallback -#[inline(always)] -pub fn record_zero_copy_write_fallback(reason: &str) { - counter!("rustfs.zero_copy.write.fallback.total", "reason" => reason.to_string()).increment(1); -} - -/// Record bytes saved from zero-copy. -/// -/// # Arguments -/// -/// * `size_bytes` - Number of bytes saved from zero-copy -#[inline(always)] -pub fn record_bytes_saved(size_bytes: usize) { - counter!("rustfs.zero_copy.bytes.saved.total").increment(size_bytes as u64); -} - // ============================================================================ // S3 Operation Metrics (GetObject, PutObject, etc.) // ============================================================================ @@ -343,6 +527,9 @@ pub fn record_bytes_saved(size_bytes: usize) { /// * `duration_ms` - Operation duration in milliseconds /// * `size_bytes` - Object size in bytes /// * `from_cache` - Whether the object was served from cache +/// +/// Note: this function records aggregate S3 GET metrics only. It must not be +/// interpreted as the definitive source of truth for data-plane copy mode. #[inline(always)] pub fn record_get_object(duration_ms: f64, size_bytes: i64, from_cache: bool) { counter!("rustfs.s3.get_object.total").increment(1); @@ -365,14 +552,33 @@ pub fn record_get_object(duration_ms: f64, size_bytes: i64, from_cache: bool) { /// /// * `duration_ms` - Operation duration in milliseconds /// * `size_bytes` - Object size in bytes -/// * `zero_copy_enabled` - Whether zero-copy was enabled for this operation +/// * `zero_copy_enabled` - Legacy aggregate flag preserved for compatibility +/// +/// Note: this function records aggregate S3 PUT metrics only. The definitive +/// outcome of request-level fast-path attempts must be tracked separately via +/// ADR 0001 data-plane helpers. #[inline(always)] pub fn record_put_object(duration_ms: f64, size_bytes: i64, zero_copy_enabled: bool) { counter!("rustfs.s3.put_object.total").increment(1); histogram!("rustfs.s3.put_object.duration.ms").record(duration_ms); + counter!( + "rustfs.s3.put_object.bucketed.total", + "size_bucket" => put_size_bucket_label(size_bytes) + ) + .increment(1); + histogram!( + "rustfs.s3.put_object.bucketed.duration.ms", + "size_bucket" => put_size_bucket_label(size_bytes) + ) + .record(duration_ms); if size_bytes > 0 { histogram!("rustfs.s3.put_object.size.bytes").record(size_bytes as f64); + histogram!( + "rustfs.s3.put_object.bucketed.size.bytes", + "size_bucket" => put_size_bucket_label(size_bytes) + ) + .record(size_bytes as f64); } if zero_copy_enabled { @@ -717,10 +923,110 @@ mod tests { use super::*; #[test] - fn test_record_zero_copy_read() { - record_zero_copy_read(1024, 10.5); - record_memory_copy_saved(1024); - record_zero_copy_fallback("test"); + fn test_io_path_as_str_values_stable() { + assert_eq!(IoPath::Fast.as_str(), "fast"); + assert_eq!(IoPath::Legacy.as_str(), "legacy"); + } + + #[test] + fn test_copy_mode_as_str_values_stable() { + assert_eq!(CopyMode::TrueZeroCopy.as_str(), "true_zero_copy"); + assert_eq!(CopyMode::SharedBytes.as_str(), "shared_bytes"); + assert_eq!(CopyMode::SingleCopy.as_str(), "single_copy"); + assert_eq!(CopyMode::Reconstructed.as_str(), "reconstructed"); + assert_eq!(CopyMode::Transformed.as_str(), "transformed"); + } + + #[test] + fn test_fallback_reason_as_str_values_stable() { + assert_eq!(FallbackReason::Unknown.as_str(), "unknown"); + assert_eq!(FallbackReason::MmapDisabled.as_str(), "mmap_disabled"); + assert_eq!(FallbackReason::MmapUnavailable.as_str(), "mmap_unavailable"); + assert_eq!(FallbackReason::SmallObject.as_str(), "small_object"); + assert_eq!(FallbackReason::WindowLimitExceeded.as_str(), "window_limit_exceeded"); + assert_eq!(FallbackReason::UnalignedWindow.as_str(), "unaligned_window"); + assert_eq!(FallbackReason::RangeNotSupported.as_str(), "range_not_supported"); + assert_eq!(FallbackReason::EncryptionEnabled.as_str(), "encryption_enabled"); + assert_eq!(FallbackReason::CompressionEnabled.as_str(), "compression_enabled"); + assert_eq!(FallbackReason::TransformEncryptionLegacy.as_str(), "transform_encryption_legacy"); + assert_eq!(FallbackReason::TransformCompressionLegacy.as_str(), "transform_compression_legacy"); + assert_eq!( + FallbackReason::TransformCompressionEncryptionLegacy.as_str(), + "transform_compression_encryption_legacy" + ); + assert_eq!(FallbackReason::ChunkBridgeUnavailable.as_str(), "chunk_bridge_unavailable"); + assert_eq!(FallbackReason::NonLocalBackend.as_str(), "non_local_backend"); + } + + #[test] + fn test_record_io_path_selected() { + record_io_path_selected("get", IoPath::Fast); + record_io_path_selected("put", IoPath::Legacy); + } + + #[test] + fn test_record_io_copy_mode() { + record_io_copy_mode("get", CopyMode::SharedBytes, 1024); + record_io_copy_mode("put", CopyMode::Transformed, 2048); + } + + #[test] + fn test_record_io_fallback() { + record_io_fallback(IoStage::ReadSetup, FallbackReason::MmapUnavailable); + record_io_fallback(IoStage::HttpBridge, FallbackReason::ChunkBridgeUnavailable); + } + + #[test] + fn test_record_local_disk_active_mmap_bytes() { + record_local_disk_active_mmap_bytes(4096); + record_local_disk_active_mmap_bytes(0); + } + + #[test] + fn test_record_local_disk_pooled_chunk() { + record_local_disk_pooled_chunk("fallback", 4096); + record_local_disk_pooled_chunk("compat_collect", 8192); + } + + #[test] + fn test_record_local_disk_compat_collect() { + record_local_disk_compat_collect(3, 16384); + } + + #[test] + fn test_record_put_object_attempted_fast_path() { + record_put_object_attempted_fast_path(1024 * 1024); + record_put_object_attempted_fast_path(0); + } + + #[test] + fn test_record_put_transform_selected() { + record_put_transform_selected("compression", IoPath::Fast, 2048); + record_put_transform_selected("compression_encryption", IoPath::Legacy, 4096); + } + + #[test] + fn test_record_put_path_selected() { + record_put_path_selected(8 * 1024, IoPath::Fast); + record_put_path_selected(2 * 1024 * 1024, IoPath::Legacy); + } + + #[test] + fn test_record_put_copy_mode() { + record_put_copy_mode(8 * 1024, CopyMode::SingleCopy); + record_put_copy_mode(512 * 1024, CopyMode::Transformed); + } + + #[test] + fn test_record_put_fallback() { + record_put_fallback(32 * 1024, FallbackReason::CompressionEnabled); + record_put_fallback(2 * 1024 * 1024, FallbackReason::EncryptionEnabled); + } + + #[test] + fn test_record_put_inline_selected() { + record_put_inline_selected(8 * 1024, false); + record_put_inline_selected(32 * 1024, true); } #[test] @@ -731,13 +1037,6 @@ mod tests { record_bytes_pool_hit_rate("small", 0.85); } - #[test] - fn test_record_zero_copy_write() { - record_zero_copy_write(1024, 10.5); - record_zero_copy_write_fallback("test"); - record_bytes_saved(1024); - } - // S3 Operation Metrics Tests #[test] fn test_record_get_object() { @@ -857,157 +1156,6 @@ mod tests { } } -// ============================================================================ -// Zero-Copy Optimization Metrics (Phase 1 Extension) -// ============================================================================ - pub mod bandwidth; pub mod global_metrics; pub mod metric_names; - -pub use metric_names::zero_copy; - -/// Record a zero-copy buffer operation. -/// -/// This function records metrics for zero-copy buffer operations, -/// including the operation type and size. -#[inline(always)] -pub fn record_zero_copy_buffer_operation(operation: &str, size: usize) { - counter!( - zero_copy::BUFFER_OPERATIONS_TOTAL, - "operation" => operation.to_string() - ) - .increment(1); - - counter!( - zero_copy::BUFFER_BYTES_TOTAL, - "operation" => operation.to_string() - ) - .increment(size as u64); -} - -/// Record memory copy operations. -/// -/// This function tracks the number and size of memory copies, -/// which should be minimized in zero-copy paths. -#[inline(always)] -pub fn record_memory_copy(count: u32, size: usize) { - counter!(zero_copy::MEMORY_COPY_TOTAL).increment(count as u64); - - counter!(zero_copy::MEMORY_COPY_BYTES_TOTAL).increment(size as u64); - - histogram!("rustfs_memory_copy_size_bytes").record(size as f64); -} - -/// Record a shared reference operation. -/// -/// This function tracks operations that create or use shared references -/// for zero-copy data sharing. -#[inline(always)] -pub fn record_shared_ref_operation(operation: &str) { - counter!( - zero_copy::SHARED_REF_OPERATIONS_TOTAL, - "operation" => operation.to_string() - ) - .increment(1); -} - -/// Record BufReader optimization. -/// -/// This function tracks BufReader layer elimination and buffer size -/// adjustments. -#[inline(always)] -pub fn record_bufreader_optimization(layers_eliminated: u32, buffer_size: usize) { - counter!(zero_copy::BUFREADER_LAYERS_ELIMINATED_TOTAL).increment(layers_eliminated as u64); - - histogram!(zero_copy::BUFREADER_BUFFER_SIZE_BYTES).record(buffer_size as f64); -} - -/// Record Direct I/O operation. -/// -/// This function tracks Direct I/O operations and their success/fallback -/// status. -#[inline(always)] -pub fn record_direct_io_operation(operation: &str, size: usize, success: bool) { - let status = if success { "success" } else { "fallback" }; - - counter!( - zero_copy::DIRECT_IO_OPERATIONS_TOTAL, - "operation" => operation.to_string(), - "status" => status.to_string() - ) - .increment(1); - - counter!( - zero_copy::DIRECT_IO_BYTES_TOTAL, - "operation" => operation.to_string(), - "status" => status.to_string() - ) - .increment(size as u64); -} - -/// Update zero-copy performance metrics. -/// -/// This function updates gauge metrics for overall zero-copy performance. -#[inline(always)] -pub fn update_zero_copy_performance_metrics(copy_count: u32, throughput_mbps: f64, memory_saved: u64) { - gauge!(zero_copy::AVG_COPY_COUNT).set(copy_count as f64); - - gauge!(zero_copy::THROUGHPUT_MBPS).set(throughput_mbps); - - gauge!(zero_copy::MEMORY_SAVED_BYTES).set(memory_saved as f64); -} - -// ============================================================================ -// Zero-Copy Metrics Tests -// ============================================================================ - -#[cfg(test)] -mod zero_copy_tests { - use super::*; - - #[test] - fn test_record_zero_copy_buffer_operation() { - // This test verifies the function compiles and runs - // Actual metric verification requires a metrics recorder - record_zero_copy_buffer_operation("read", 1024); - record_zero_copy_buffer_operation("write", 2048); - } - - #[test] - fn test_record_memory_copy() { - record_memory_copy(1, 1024); - record_memory_copy(2, 2048); - } - - #[test] - fn test_record_shared_ref_operation() { - record_shared_ref_operation("create"); - record_shared_ref_operation("share"); - } - - #[test] - fn test_record_bufreader_optimization() { - record_bufreader_optimization(1, 8192); - record_bufreader_optimization(2, 65536); - } - - #[test] - fn test_record_direct_io_operation() { - record_direct_io_operation("read", 4096, true); - record_direct_io_operation("write", 8192, false); - } - - #[test] - fn test_update_zero_copy_performance_metrics() { - update_zero_copy_performance_metrics(2, 150.5, 1024 * 1024); - } - - #[test] - fn test_metric_names() { - // Verify metric names are defined - assert!(!zero_copy::BUFFER_OPERATIONS_TOTAL.is_empty()); - assert!(!zero_copy::MEMORY_COPY_TOTAL.is_empty()); - assert!(!zero_copy::THROUGHPUT_MBPS.is_empty()); - } -} diff --git a/crates/io-metrics/src/metric_names.rs b/crates/io-metrics/src/metric_names.rs index e7581ff8f..6c611b6aa 100644 --- a/crates/io-metrics/src/metric_names.rs +++ b/crates/io-metrics/src/metric_names.rs @@ -14,41 +14,44 @@ //! Metric name constants for consistent naming across the codebase. -/// Zero-copy operation metric names. -pub mod zero_copy { - /// Total number of zero-copy buffer operations - pub const BUFFER_OPERATIONS_TOTAL: &str = "rustfs_zero_copy_buffer_operations_total"; +/// Request-level data plane metric names introduced by ADR 0001. +pub mod data_plane { + /// Total number of selected request paths. + pub const PATH_SELECTED_TOTAL: &str = "rustfs.io.path.selected_total"; - /// Total bytes processed by zero-copy buffer operations - pub const BUFFER_BYTES_TOTAL: &str = "rustfs_zero_copy_buffer_bytes_total"; + /// Total bytes observed for a given effective copy mode. + pub const COPY_MODE_BYTES_TOTAL: &str = "rustfs.io.copy_mode.bytes_total"; - /// Total number of memory copies - pub const MEMORY_COPY_TOTAL: &str = "rustfs_memory_copy_total"; + /// Total number of data plane fallbacks. + pub const FALLBACK_TOTAL: &str = "rustfs.io.zero_copy.fallback_total"; - /// Total bytes copied in memory - pub const MEMORY_COPY_BYTES_TOTAL: &str = "rustfs_memory_copy_bytes_total"; + /// Current active local-disk mmap bytes held by chunk fast paths. + pub const LOCAL_DISK_ACTIVE_MMAP_BYTES: &str = "rustfs.io.local_disk.active_mmap.bytes"; - /// Total number of shared reference operations - pub const SHARED_REF_OPERATIONS_TOTAL: &str = "rustfs_shared_ref_operations_total"; + /// Total pooled chunks produced or consumed by LocalDisk compatibility paths. + pub const LOCAL_DISK_POOLED_CHUNKS_TOTAL: &str = "rustfs.io.local_disk.pooled_chunks.total"; - /// Total number of BufReader layers eliminated - pub const BUFREADER_LAYERS_ELIMINATED_TOTAL: &str = "rustfs_bufreader_layers_eliminated_total"; + /// Total pooled bytes produced or consumed by LocalDisk compatibility paths. + pub const LOCAL_DISK_POOLED_BYTES_TOTAL: &str = "rustfs.io.local_disk.pooled_bytes.total"; - /// BufReader buffer size distribution - pub const BUFREADER_BUFFER_SIZE_BYTES: &str = "rustfs_bufreader_buffer_size_bytes"; + /// Total number of compatibility chunk-stream aggregations performed for LocalDisk reads. + pub const LOCAL_DISK_COMPAT_COLLECT_TOTAL: &str = "rustfs.io.local_disk.compat_collect.total"; - /// Total number of Direct I/O operations - pub const DIRECT_IO_OPERATIONS_TOTAL: &str = "rustfs_direct_io_operations_total"; + /// Chunk count distribution for LocalDisk compatibility chunk aggregation. + pub const LOCAL_DISK_COMPAT_COLLECT_CHUNKS: &str = "rustfs.io.local_disk.compat_collect.chunks"; - /// Total bytes processed by Direct I/O - pub const DIRECT_IO_BYTES_TOTAL: &str = "rustfs_direct_io_bytes_total"; + /// Byte distribution for LocalDisk compatibility chunk aggregation. + pub const LOCAL_DISK_COMPAT_COLLECT_BYTES: &str = "rustfs.io.local_disk.compat_collect.bytes"; - /// Average copy count per operation - pub const AVG_COPY_COUNT: &str = "rustfs_zero_copy_avg_copy_count"; + /// Total number of attempted PUT fast paths. + pub const PUT_FAST_PATH_ATTEMPTS_TOTAL: &str = "rustfs.io.put.fast_path.attempts_total"; - /// Throughput in MB/s - pub const THROUGHPUT_MBPS: &str = "rustfs_zero_copy_throughput_mbps"; + /// Size distribution for attempted PUT fast paths. + pub const PUT_FAST_PATH_ATTEMPT_SIZE_BYTES: &str = "rustfs.io.put.fast_path.attempt.size.bytes"; - /// Memory saved by zero-copy in bytes - pub const MEMORY_SAVED_BYTES: &str = "rustfs_zero_copy_memory_saved_bytes"; + /// Total number of transformed PUT selections grouped by transform kind and ingress path. + pub const PUT_TRANSFORM_SELECTED_TOTAL: &str = "rustfs.io.put.transform.selected_total"; + + /// Size distribution for transformed PUT selections. + pub const PUT_TRANSFORM_SIZE_BYTES: &str = "rustfs.io.put.transform.size.bytes"; } diff --git a/crates/object-io/Cargo.toml b/crates/object-io/Cargo.toml new file mode 100644 index 000000000..797c8bfcd --- /dev/null +++ b/crates/object-io/Cargo.toml @@ -0,0 +1,37 @@ +[package] +name = "rustfs-object-io" +version.workspace = true +edition.workspace = true +license.workspace = true +repository.workspace = true +rust-version.workspace = true +homepage.workspace = true +description = "Object I/O policy helpers and zero-copy support primitives for RustFS." +keywords.workspace = true +categories.workspace = true + +[lints] +workspace = true + +[dependencies] +atoi = { workspace = true } +bytes = { workspace = true } +futures-util = { workspace = true } +rustfs-ecstore = { workspace = true } +rustfs-concurrency = { workspace = true } +rustfs-io-core = { workspace = true } +rustfs-io-metrics = { workspace = true } +rustfs-rio = { workspace = true } +rustfs-s3select-api = { workspace = true } +rustfs-utils = { workspace = true ,features = ["http"]} +http = { workspace = true } +s3s.workspace = true +thiserror = { workspace = true } +time = { workspace = true, features = ["parsing", "formatting"] } +tokio = { workspace = true, features = ["io-util"] } +tokio-util = { workspace = true, features = ["io"] } +astral-tokio-tar = { workspace = true } +uuid = { workspace = true } + +[dev-dependencies] +serial_test = { workspace = true } diff --git a/crates/object-io/src/get.rs b/crates/object-io/src/get.rs new file mode 100644 index 000000000..90e49b992 --- /dev/null +++ b/crates/object-io/src/get.rs @@ -0,0 +1,1703 @@ +// 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 bytes::Bytes; +use futures_util::StreamExt; +use http::HeaderMap; +use http::header::{CACHE_CONTROL, CONTENT_DISPOSITION, CONTENT_LANGUAGE}; +use rustfs_concurrency::GetObjectCacheEligibility; +use rustfs_ecstore::bucket::lifecycle::lifecycle::TRANSITION_COMPLETE; +use rustfs_ecstore::client::object_api_utils::to_s3s_etag; +use rustfs_ecstore::error::StorageError; +use rustfs_ecstore::store_api::{GetObjectChunkCopyMode, GetObjectChunkPath, GetObjectChunkResult, HTTPRangeSpec, ObjectInfo}; +use rustfs_io_core::BoxChunkStream; +use rustfs_rio::Reader; +use rustfs_s3select_api::object_store::bytes_stream; +use rustfs_utils::http::{AMZ_CHECKSUM_MODE, AMZ_CHECKSUM_TYPE}; +use s3s::dto::{ + ChecksumType, ContentType, GetObjectOutput, SSECustomerAlgorithm, SSECustomerKeyMD5, SSEKMSKeyId, ServerSideEncryption, + StreamingBlob, Timestamp, +}; +use std::collections::HashMap; +use std::str::FromStr; +use std::sync::Arc; +use time::{OffsetDateTime, format_description::well_known::Rfc3339}; +use tokio::io::{AsyncRead, AsyncSeek, ReadBuf}; +use tokio_util::io::ReaderStream; + +pub struct InMemoryAsyncReader { + cursor: std::io::Cursor, +} + +impl InMemoryAsyncReader { + pub fn new(data: Bytes) -> 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 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::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())) + } +} + +pub fn build_memory_blob(buf: Bytes, response_content_length: i64, optimal_buffer_size: usize) -> Option { + let mem_reader = InMemoryAsyncReader::new(buf); + Some(StreamingBlob::wrap(bytes_stream( + ReaderStream::with_capacity(Box::new(mem_reader), optimal_buffer_size), + response_content_length as usize, + ))) +} + +#[derive(Clone)] +pub struct FrozenGetObjectBody { + body: Arc, +} + +impl FrozenGetObjectBody { + pub fn new(body: Bytes) -> Self { + Self { body: Arc::new(body) } + } + + pub fn shared_body(&self) -> &Arc { + &self.body + } + + pub fn into_shared_body(self) -> Arc { + self.body + } + + pub fn build_blob(&self, response_content_length: i64, optimal_buffer_size: usize) -> Option { + build_memory_blob((*self.body).clone(), response_content_length, optimal_buffer_size) + } +} + +pub fn build_reader_blob(reader: R, response_content_length: i64, optimal_buffer_size: usize) -> Option +where + R: AsyncRead + Send + Sync + 'static, +{ + Some(StreamingBlob::wrap(bytes_stream( + ReaderStream::with_capacity(reader, optimal_buffer_size), + response_content_length as usize, + ))) +} + +pub fn build_chunk_blob(chunk_stream: BoxChunkStream) -> Option { + Some(StreamingBlob::wrap(chunk_stream.map(|result| result.map(|chunk| chunk.as_bytes())))) +} + +/// ADR-facing alias for the chunk-stream to HTTP body bridge. +pub fn build_chunk_http_body(chunk_stream: BoxChunkStream) -> Option { + build_chunk_blob(chunk_stream) +} + +pub fn map_chunk_copy_mode(copy_mode: GetObjectChunkCopyMode) -> rustfs_io_metrics::CopyMode { + match copy_mode { + GetObjectChunkCopyMode::TrueZeroCopy => rustfs_io_metrics::CopyMode::TrueZeroCopy, + GetObjectChunkCopyMode::SharedBytes => rustfs_io_metrics::CopyMode::SharedBytes, + GetObjectChunkCopyMode::SingleCopy => rustfs_io_metrics::CopyMode::SingleCopy, + GetObjectChunkCopyMode::Reconstructed => rustfs_io_metrics::CopyMode::Reconstructed, + } +} + +pub fn chunk_body_data_plane_labels( + path: GetObjectChunkPath, + copy_mode: rustfs_io_metrics::CopyMode, +) -> (rustfs_io_metrics::IoPath, rustfs_io_metrics::CopyMode) { + ( + match path { + GetObjectChunkPath::Direct | GetObjectChunkPath::Bridge => rustfs_io_metrics::IoPath::Fast, + }, + copy_mode, + ) +} + +pub fn get_object_chunk_fast_path_guard( + has_sse_customer_key: bool, + has_sse_customer_key_md5: bool, +) -> Result<(), ChunkReadFallback> { + if has_sse_customer_key || has_sse_customer_key_md5 { + return Err(ChunkReadFallback::read_setup(rustfs_io_metrics::FallbackReason::EncryptionEnabled)); + } + + Ok(()) +} + +pub fn get_object_sequential_hint(rs: Option<&HTTPRangeSpec>) -> bool { + if rs.is_none() { + true + } else if let Some(range_spec) = rs { + range_spec.start == 0 && !range_spec.is_suffix_length + } else { + false + } +} + +pub trait CachedGetObjectSource { + fn body(&self) -> &Arc; + fn content_length(&self) -> i64; + fn content_type(&self) -> Option<&str>; + fn e_tag(&self) -> Option<&str>; + fn last_modified(&self) -> Option<&str>; + fn cache_control(&self) -> Option<&str>; + fn content_disposition(&self) -> Option<&str>; + fn content_encoding(&self) -> Option<&str>; + fn content_language(&self) -> Option<&str>; + fn storage_class(&self) -> Option<&str>; + fn version_id(&self) -> Option<&str>; + fn delete_marker(&self) -> bool; + fn tag_count(&self) -> Option; + fn user_metadata(&self) -> &HashMap; + fn checksum_crc32(&self) -> Option<&str>; + fn checksum_crc32c(&self) -> Option<&str>; + fn checksum_sha1(&self) -> Option<&str>; + fn checksum_sha256(&self) -> Option<&str>; + fn checksum_crc64nvme(&self) -> Option<&str>; + fn checksum_type(&self) -> Option<&ChecksumType>; +} + +pub struct GetObjectOutputContext { + pub output: GetObjectOutput, + pub event_info: ObjectInfo, + pub response_content_length: i64, + pub optimal_buffer_size: usize, + pub copy_mode_override: Option, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub struct GetObjectStrategyLayout { + pub is_sequential_hint: bool, + pub optimal_buffer_size: usize, +} + +pub enum GetObjectBodySource { + Reader(Box), + Chunk { + stream: BoxChunkStream, + path: GetObjectChunkPath, + copy_mode: rustfs_io_metrics::CopyMode, + }, +} + +pub struct GetObjectReadSetup { + pub info: ObjectInfo, + pub event_info: ObjectInfo, + pub body_source: GetObjectBodySource, + pub rs: Option, + pub content_type: Option, + pub last_modified: Option, + pub response_content_length: i64, + pub content_range: Option, + pub server_side_encryption: Option, + pub sse_customer_algorithm: Option, + pub sse_customer_key_md5: Option, + pub ssekms_key_id: Option, + pub encryption_applied: bool, +} + +pub struct LegacyReadPlan { + pub rs: Option, + pub content_type: Option, + pub last_modified: Option, + pub response_content_length: i64, + pub content_range: Option, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum GetObjectBodyPlan { + CacheWriteback, + BufferEncrypted, + BufferSeekable, + Stream, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum GetObjectDataPlaneRequestSource { + CacheServed, + Disk, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub struct GetObjectDataPlaneMetricContract { + pub request_source: GetObjectDataPlaneRequestSource, + pub io_path: rustfs_io_metrics::IoPath, + pub copy_mode: rustfs_io_metrics::CopyMode, + pub record_cache_served_metric: bool, + pub record_cache_writeback_metric: bool, +} + +impl GetObjectDataPlaneMetricContract { + pub fn cache_served() -> Self { + Self { + request_source: GetObjectDataPlaneRequestSource::CacheServed, + io_path: rustfs_io_metrics::IoPath::Fast, + copy_mode: rustfs_io_metrics::CopyMode::SharedBytes, + record_cache_served_metric: true, + record_cache_writeback_metric: false, + } + } + + pub fn disk( + io_path: rustfs_io_metrics::IoPath, + copy_mode: rustfs_io_metrics::CopyMode, + body_plan: GetObjectBodyPlan, + ) -> Self { + Self { + request_source: GetObjectDataPlaneRequestSource::Disk, + io_path, + copy_mode, + record_cache_served_metric: false, + record_cache_writeback_metric: matches!(body_plan, GetObjectBodyPlan::CacheWriteback), + } + } +} + +pub struct GetObjectCacheWriteback { + pub body: Arc, + pub content_length: i64, + pub content_type: Option, + pub content_encoding: Option, + pub cache_control: Option, + pub content_disposition: Option, + pub content_language: Option, + pub expires: Option, + pub storage_class: Option, + pub version_id: Option, + pub delete_marker: bool, + pub user_metadata: HashMap, + pub e_tag: Option, + pub last_modified: Option, + pub checksum_crc32: Option, + pub checksum_crc32c: Option, + pub checksum_sha1: Option, + pub checksum_sha256: Option, + pub checksum_crc64nvme: Option, + pub checksum_type: Option, +} + +pub struct GetObjectBodyMaterialization { + pub body: Option, + pub cache_writeback: Option, + pub plan: GetObjectBodyPlan, +} + +#[derive(Default)] +pub struct GetObjectEncryptionState { + pub server_side_encryption: Option, + pub sse_customer_algorithm: Option, + pub sse_customer_key_md5: Option, + pub ssekms_key_id: Option, + pub encryption_applied: bool, + pub response_content_length_override: Option, +} + +pub struct ChunkReadSetupResult { + pub read_setup: GetObjectReadSetup, + pub io_path: rustfs_io_metrics::IoPath, +} + +#[derive(Debug, thiserror::Error)] +pub enum MaterializeGetObjectBodyError { + #[error("failed to read object for caching: {0}")] + CacheRead(std::io::Error), + #[error("failed to read decrypted object: {0}")] + EncryptedRead(std::io::Error), +} + +pub enum GetObjectResponseMode { + Plain, + CorsWrapped, +} + +pub struct GetObjectFlowResult { + pub output: GetObjectOutput, + pub event_info: ObjectInfo, + pub version_id_for_event: String, + pub response_mode: GetObjectResponseMode, +} + +pub fn build_get_object_flow_result( + output: GetObjectOutput, + event_info: ObjectInfo, + version_id_for_event: String, + response_mode: GetObjectResponseMode, +) -> GetObjectFlowResult { + GetObjectFlowResult { + output, + event_info, + version_id_for_event, + response_mode, + } +} + +pub fn build_cached_get_object_flow_result_from_source( + bucket: &str, + key: &str, + cached: &T, + version_id_for_event: String, +) -> GetObjectFlowResult +where + T: CachedGetObjectSource, +{ + build_get_object_flow_result( + build_cached_get_object_output_from_source(cached), + build_cached_get_object_event_info_from_source(bucket, key, cached), + version_id_for_event, + GetObjectResponseMode::Plain, + ) +} + +pub fn build_cors_wrapped_get_object_flow_result( + output_context: GetObjectOutputContext, + version_id_for_event: String, +) -> GetObjectFlowResult { + let GetObjectOutputContext { + output, + event_info, + response_content_length: _, + optimal_buffer_size: _, + copy_mode_override: _, + } = output_context; + build_get_object_flow_result(output, event_info, version_id_for_event, GetObjectResponseMode::CorsWrapped) +} + +pub fn build_cached_get_object_output_from_source(cached: &T) -> GetObjectOutput +where + T: CachedGetObjectSource, +{ + let body_data = Arc::clone(cached.body()); + let body = Some(StreamingBlob::wrap::<_, std::convert::Infallible>(futures_util::stream::once( + async move { Ok((*body_data).clone()) }, + ))); + + let last_modified = cached + .last_modified() + .and_then(|s| OffsetDateTime::parse(s, &Rfc3339).ok()) + .map(Timestamp::from); + + let content_type = cached.content_type().and_then(|ct| ContentType::from_str(ct).ok()); + + let metadata = (!cached.user_metadata().is_empty()).then(|| cached.user_metadata().clone()); + + GetObjectOutput { + body, + content_length: Some(cached.content_length()), + accept_ranges: Some("bytes".to_string()), + e_tag: cached.e_tag().map(to_s3s_etag), + last_modified, + content_type, + cache_control: cached.cache_control().map(str::to_string), + content_disposition: cached.content_disposition().map(str::to_string), + content_encoding: cached.content_encoding().map(str::to_string), + content_language: cached.content_language().map(str::to_string), + version_id: cached.version_id().map(str::to_string), + delete_marker: Some(cached.delete_marker()), + tag_count: cached.tag_count(), + metadata, + checksum_crc32: cached.checksum_crc32().map(str::to_string), + checksum_crc32c: cached.checksum_crc32c().map(str::to_string), + checksum_sha1: cached.checksum_sha1().map(str::to_string), + checksum_sha256: cached.checksum_sha256().map(str::to_string), + checksum_crc64nvme: cached.checksum_crc64nvme().map(str::to_string), + checksum_type: cached.checksum_type().cloned(), + ..Default::default() + } +} +pub fn build_cached_get_object_event_info_from_source(bucket: &str, key: &str, cached: &T) -> ObjectInfo +where + T: CachedGetObjectSource, +{ + let last_modified = cached.last_modified().and_then(|s| OffsetDateTime::parse(s, &Rfc3339).ok()); + let version_id = cached.version_id().and_then(|v| uuid::Uuid::parse_str(v).ok()); + + ObjectInfo { + bucket: bucket.to_string(), + name: key.to_string(), + storage_class: cached.storage_class().map(str::to_string), + mod_time: last_modified, + size: cached.content_length(), + actual_size: cached.content_length(), + is_dir: false, + user_defined: cached.user_metadata().clone(), + version_id, + delete_marker: cached.delete_marker(), + content_type: cached.content_type().map(str::to_string), + content_encoding: cached.content_encoding().map(str::to_string), + etag: cached.e_tag().map(str::to_string), + ..Default::default() + } +} + +#[derive(Debug, Default, Clone, PartialEq, Eq)] +pub struct GetObjectChecksums { + pub crc32: Option, + pub crc32c: Option, + pub sha1: Option, + pub sha256: Option, + pub crc64nvme: Option, + pub checksum_type: Option, +} + +fn decode_get_object_checksums(decrypted_checksums: HashMap) -> GetObjectChecksums { + let mut checksums = GetObjectChecksums::default(); + + for (key, checksum) in decrypted_checksums { + if key == AMZ_CHECKSUM_TYPE { + checksums.checksum_type = Some(ChecksumType::from(checksum)); + continue; + } + + match rustfs_rio::ChecksumType::from_string(key.as_str()) { + rustfs_rio::ChecksumType::CRC32 => checksums.crc32 = Some(checksum), + rustfs_rio::ChecksumType::CRC32C => checksums.crc32c = Some(checksum), + rustfs_rio::ChecksumType::SHA1 => checksums.sha1 = Some(checksum), + rustfs_rio::ChecksumType::SHA256 => checksums.sha256 = Some(checksum), + rustfs_rio::ChecksumType::CRC64_NVME => checksums.crc64nvme = Some(checksum), + _ => (), + } + } + + checksums +} + +fn read_object_checksums( + info: &ObjectInfo, + headers: &HeaderMap, + part_number: Option, +) -> std::io::Result { + let (decrypted_checksums, _is_multipart) = info + .decrypt_checksums(part_number.unwrap_or(0), headers) + .map_err(|e| std::io::Error::other(e.to_string()))?; + + Ok(decode_get_object_checksums(decrypted_checksums)) +} + +pub fn build_get_object_checksums( + info: &ObjectInfo, + headers: &HeaderMap, + part_number: Option, + rs: Option<&HTTPRangeSpec>, +) -> std::io::Result { + if let Some(checksum_mode) = headers.get(AMZ_CHECKSUM_MODE) + && checksum_mode.to_str().unwrap_or_default() == "ENABLED" + && rs.is_none() + { + return read_object_checksums(info, headers, part_number); + } + + Ok(GetObjectChecksums::default()) +} + +pub fn build_output_version_id(versioned: bool, version_id: Option<&uuid::Uuid>) -> Option { + if !versioned { + return None; + } + + version_id.map(|vid| { + if *vid == uuid::Uuid::nil() { + "null".to_string() + } else { + vid.to_string() + } + }) +} + +pub fn plan_get_object_strategy_layout( + rs: Option<&HTTPRangeSpec>, + response_content_length: i64, + suggested_buffer_size: usize, + fallback_buffer_size: usize, +) -> GetObjectStrategyLayout { + let is_sequential_hint = get_object_sequential_hint(rs); + let optimal_buffer_size = if suggested_buffer_size > 0 { + suggested_buffer_size.min(fallback_buffer_size) + } else { + rustfs_io_core::get_concurrency_aware_buffer_size(response_content_length, fallback_buffer_size) + }; + + GetObjectStrategyLayout { + is_sequential_hint, + optimal_buffer_size, + } +} + +pub fn plan_get_object_body( + cache_eligibility: GetObjectCacheEligibility, + seekable_object_size_threshold: usize, +) -> GetObjectBodyPlan { + if cache_eligibility.should_cache() { + return GetObjectBodyPlan::CacheWriteback; + } + + let should_buffer_for_seek = cache_eligibility.response_size > 0 + && cache_eligibility.response_size <= seekable_object_size_threshold as i64 + && !cache_eligibility.is_part_request + && !cache_eligibility.is_range_request; + + if cache_eligibility.encryption_applied && should_buffer_for_seek { + GetObjectBodyPlan::BufferEncrypted + } else if should_buffer_for_seek { + GetObjectBodyPlan::BufferSeekable + } else { + GetObjectBodyPlan::Stream + } +} + +pub fn build_get_object_cache_writeback(info: &ObjectInfo, body: Bytes, content_length: i64) -> GetObjectCacheWriteback { + let checksums = read_object_checksums(info, &HeaderMap::new(), None).unwrap_or_default(); + let body = FrozenGetObjectBody::new(body); + GetObjectCacheWriteback { + body: body.into_shared_body(), + content_length, + content_type: info.content_type.clone(), + content_encoding: info.content_encoding.clone(), + cache_control: None, + content_disposition: None, + content_language: None, + expires: None, + storage_class: info.storage_class.clone(), + version_id: info.version_id.map(|vid| { + if vid == uuid::Uuid::nil() { + "null".to_string() + } else { + vid.to_string() + } + }), + delete_marker: info.delete_marker, + user_metadata: HashMap::new(), + e_tag: info.etag.clone(), + last_modified: info.mod_time.and_then(|t| t.format(&Rfc3339).ok()), + checksum_crc32: checksums.crc32, + checksum_crc32c: checksums.crc32c, + checksum_sha1: checksums.sha1, + checksum_sha256: checksums.sha256, + checksum_crc64nvme: checksums.crc64nvme, + checksum_type: checksums.checksum_type, + } +} + +pub fn finalize_get_object_cache_writeback( + info: &ObjectInfo, + writeback: GetObjectCacheWriteback, + user_metadata: HashMap, +) -> GetObjectCacheWriteback { + GetObjectCacheWriteback { + cache_control: info.user_defined.get(CACHE_CONTROL.as_str()).cloned(), + content_disposition: info.user_defined.get(CONTENT_DISPOSITION.as_str()).cloned(), + content_language: info.user_defined.get(CONTENT_LANGUAGE.as_str()).cloned(), + expires: info.expires.and_then(|t| t.format(&Rfc3339).ok()), + user_metadata, + ..writeback + } +} + +pub async fn materialize_get_object_body( + mut final_stream: R, + info: &ObjectInfo, + plan: GetObjectBodyPlan, + response_content_length: i64, + optimal_buffer_size: usize, +) -> Result +where + R: AsyncRead + Send + Sync + Unpin + 'static, +{ + match plan { + GetObjectBodyPlan::CacheWriteback => { + let mut buf = Vec::with_capacity(response_content_length as usize); + tokio::io::AsyncReadExt::read_to_end(&mut final_stream, &mut buf) + .await + .map_err(MaterializeGetObjectBodyError::CacheRead)?; + let body = FrozenGetObjectBody::new(Bytes::from(buf)); + + Ok(GetObjectBodyMaterialization { + body: body.build_blob(response_content_length, optimal_buffer_size), + cache_writeback: Some(build_get_object_cache_writeback( + info, + body.shared_body().as_ref().clone(), + response_content_length, + )), + plan, + }) + } + GetObjectBodyPlan::BufferEncrypted => { + let mut buf = Vec::with_capacity(response_content_length as usize); + tokio::io::AsyncReadExt::read_to_end(&mut final_stream, &mut buf) + .await + .map_err(MaterializeGetObjectBodyError::EncryptedRead)?; + let body = FrozenGetObjectBody::new(Bytes::from(buf)); + + Ok(GetObjectBodyMaterialization { + body: body.build_blob(response_content_length, optimal_buffer_size), + cache_writeback: None, + plan, + }) + } + GetObjectBodyPlan::BufferSeekable => { + let mut buf = Vec::with_capacity(response_content_length as usize); + let body = match tokio::io::AsyncReadExt::read_to_end(&mut final_stream, &mut buf).await { + Ok(_) => FrozenGetObjectBody::new(Bytes::from(buf)).build_blob(response_content_length, optimal_buffer_size), + Err(_) => build_reader_blob(final_stream, response_content_length, optimal_buffer_size), + }; + + Ok(GetObjectBodyMaterialization { + body, + cache_writeback: None, + plan, + }) + } + GetObjectBodyPlan::Stream => Ok(GetObjectBodyMaterialization { + body: build_reader_blob(final_stream, response_content_length, optimal_buffer_size), + cache_writeback: None, + plan, + }), + } +} + +fn resolve_requested_range( + info: &ObjectInfo, + mut rs: Option, + part_number: Option, +) -> Option { + if let Some(part_number) = part_number + && rs.is_none() + { + rs = HTTPRangeSpec::from_object_info(info, part_number); + } + + rs +} + +fn resolve_response_range( + total_size: i64, + rs: Option, +) -> std::io::Result<(Option, i64, Option)> { + let Some(range_spec) = rs else { + return Ok((None, total_size, None)); + }; + + let (start, length) = range_spec.get_offset_length(total_size)?; + let content_range = Some(format!("bytes {}-{}/{}", start, start as i64 + length - 1, total_size)); + + Ok((Some(range_spec), length, content_range)) +} + +pub fn plan_legacy_read( + info: &ObjectInfo, + rs: Option, + part_number: Option, +) -> std::io::Result { + let content_type = info + .content_type + .as_ref() + .and_then(|content_type| ContentType::from_str(content_type).ok()); + let last_modified = info.mod_time.map(Timestamp::from); + let rs = resolve_requested_range(info, rs, part_number); + let total_size = info.get_actual_size()?; + let (rs, response_content_length, content_range) = resolve_response_range(total_size, rs)?; + + Ok(LegacyReadPlan { + rs, + content_type, + last_modified, + response_content_length, + content_range, + }) +} + +pub fn build_reader_read_setup( + info: ObjectInfo, + event_info: ObjectInfo, + final_stream: Box, + plan: LegacyReadPlan, + encryption_state: GetObjectEncryptionState, +) -> GetObjectReadSetup { + let LegacyReadPlan { + rs, + content_type, + last_modified, + response_content_length, + content_range, + } = plan; + + let GetObjectEncryptionState { + server_side_encryption, + sse_customer_algorithm, + sse_customer_key_md5, + ssekms_key_id, + encryption_applied, + response_content_length_override, + } = encryption_state; + + GetObjectReadSetup { + info, + event_info, + body_source: GetObjectBodySource::Reader(final_stream), + rs, + content_type, + last_modified, + response_content_length: response_content_length_override.unwrap_or(response_content_length), + content_range, + server_side_encryption, + sse_customer_algorithm, + sse_customer_key_md5, + ssekms_key_id, + encryption_applied, + } +} + +#[allow(clippy::too_many_arguments)] +pub fn build_get_object_output_context( + body: Option, + info: ObjectInfo, + event_info: ObjectInfo, + content_type: Option, + last_modified: Option, + response_content_length: i64, + content_range: Option, + server_side_encryption: Option, + sse_customer_algorithm: Option, + sse_customer_key_md5: Option, + ssekms_key_id: Option, + checksums: &GetObjectChecksums, + filtered_metadata: Option>, + versioned: bool, + optimal_buffer_size: usize, + copy_mode_override: Option, +) -> GetObjectOutputContext { + let output_version_id = build_output_version_id(versioned, info.version_id.as_ref()); + let output = build_get_object_output( + body, + &info, + content_type, + last_modified, + response_content_length, + content_range, + server_side_encryption, + sse_customer_algorithm, + sse_customer_key_md5, + ssekms_key_id, + checksums, + output_version_id, + filtered_metadata, + ); + + GetObjectOutputContext { + output, + event_info, + response_content_length, + optimal_buffer_size, + copy_mode_override, + } +} + +#[allow(clippy::too_many_arguments)] +fn build_chunk_read_setup( + info: ObjectInfo, + event_info: ObjectInfo, + path: GetObjectChunkPath, + copy_mode: rustfs_io_metrics::CopyMode, + stream: BoxChunkStream, + plan: ChunkReadPlan, +) -> GetObjectReadSetup { + let ChunkReadPlan { + rs, + content_type, + last_modified, + response_content_length, + content_range, + } = plan; + + GetObjectReadSetup { + info, + event_info, + body_source: GetObjectBodySource::Chunk { stream, path, copy_mode }, + rs, + content_type, + last_modified, + response_content_length, + content_range, + server_side_encryption: None, + sse_customer_algorithm: None, + sse_customer_key_md5: None, + ssekms_key_id: None, + encryption_applied: false, + } +} + +pub fn finalize_chunk_read_setup( + info: ObjectInfo, + event_info: ObjectInfo, + chunk_result: GetObjectChunkResult, + plan: ChunkReadPlan, +) -> ChunkReadSetupResult { + let copy_mode = map_chunk_copy_mode(chunk_result.copy_mode); + let (io_path, _) = chunk_body_data_plane_labels(chunk_result.path, copy_mode); + + ChunkReadSetupResult { + io_path, + read_setup: build_chunk_read_setup(info, event_info, chunk_result.path, copy_mode, chunk_result.stream, plan), + } +} + +#[allow(clippy::too_many_arguments)] +pub fn build_get_object_output( + body: Option, + info: &ObjectInfo, + content_type: Option, + last_modified: Option, + response_content_length: i64, + content_range: Option, + server_side_encryption: Option, + sse_customer_algorithm: Option, + sse_customer_key_md5: Option, + ssekms_key_id: Option, + checksums: &GetObjectChecksums, + output_version_id: Option, + filtered_metadata: Option>, +) -> GetObjectOutput { + GetObjectOutput { + body, + content_length: Some(response_content_length), + last_modified, + content_type, + content_encoding: info.content_encoding.clone(), + accept_ranges: Some("bytes".to_string()), + content_range, + e_tag: info.etag.as_ref().map(|etag| to_s3s_etag(etag)), + metadata: filtered_metadata, + server_side_encryption, + sse_customer_algorithm, + sse_customer_key_md5, + ssekms_key_id, + checksum_crc32: checksums.crc32.clone(), + checksum_crc32c: checksums.crc32c.clone(), + checksum_sha1: checksums.sha1.clone(), + checksum_sha256: checksums.sha256.clone(), + checksum_crc64nvme: checksums.crc64nvme.clone(), + checksum_type: checksums.checksum_type.clone(), + version_id: output_version_id, + ..Default::default() + } +} + +#[derive(Debug)] +pub struct ChunkReadPlan { + pub rs: Option, + pub content_type: Option, + pub last_modified: Option, + pub response_content_length: i64, + pub content_range: Option, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub struct ChunkReadFallback { + pub stage: rustfs_io_metrics::IoStage, + pub reason: rustfs_io_metrics::FallbackReason, +} + +impl ChunkReadFallback { + pub const fn new(stage: rustfs_io_metrics::IoStage, reason: rustfs_io_metrics::FallbackReason) -> Self { + Self { stage, reason } + } + + pub const fn read_setup(reason: rustfs_io_metrics::FallbackReason) -> Self { + Self::new(rustfs_io_metrics::IoStage::ReadSetup, reason) + } + + pub const fn range_guard(reason: rustfs_io_metrics::FallbackReason) -> Self { + Self::new(rustfs_io_metrics::IoStage::RangeGuard, reason) + } +} + +#[derive(Debug)] +pub enum ChunkReadDecision { + Eligible(ChunkReadPlan), + Fallback(ChunkReadFallback), +} + +#[derive(Debug)] +pub enum ChunkReadPlanError { + NoSuchKey, + MethodNotAllowed, + Io(std::io::Error), +} + +impl From for ChunkReadPlanError { + fn from(value: std::io::Error) -> Self { + Self::Io(value) + } +} + +impl From for ChunkReadPlanError { + fn from(value: StorageError) -> Self { + Self::Io(std::io::Error::other(value.to_string())) + } +} + +pub fn get_object_chunk_range_guard(rs: Option<&HTTPRangeSpec>) -> Result<(), ChunkReadFallback> { + let Some(range_spec) = rs else { + return Ok(()); + }; + + let unsupported = if range_spec.is_suffix_length { + range_spec.end != -1 + } else { + range_spec.start < 0 || range_spec.end < -1 || (range_spec.end != -1 && range_spec.end < range_spec.start) + }; + + if unsupported { + return Err(ChunkReadFallback::range_guard(rustfs_io_metrics::FallbackReason::RangeNotSupported)); + } + + Ok(()) +} + +pub fn plan_chunk_read( + info: &ObjectInfo, + version_id_missing: bool, + rs: Option, + part_number: Option, +) -> Result { + if info.delete_marker { + if version_id_missing { + return Err(ChunkReadPlanError::NoSuchKey); + } + return Err(ChunkReadPlanError::MethodNotAllowed); + } + + let (_, is_compressed) = info.is_compressed_ok()?; + if is_compressed { + return Ok(ChunkReadDecision::Fallback(ChunkReadFallback::read_setup( + rustfs_io_metrics::FallbackReason::CompressionEnabled, + ))); + } + + if info.transitioned_object.status == TRANSITION_COMPLETE { + return Ok(ChunkReadDecision::Fallback(ChunkReadFallback::read_setup( + rustfs_io_metrics::FallbackReason::NonLocalBackend, + ))); + } + + let has_encryption_metadata = info.user_defined.contains_key("x-rustfs-encryption-key") + || info.user_defined.contains_key("x-amz-server-side-encryption") + || info + .user_defined + .contains_key("x-amz-server-side-encryption-customer-algorithm"); + if has_encryption_metadata { + return Ok(ChunkReadDecision::Fallback(ChunkReadFallback::read_setup( + rustfs_io_metrics::FallbackReason::EncryptionEnabled, + ))); + } + + let rs = resolve_requested_range(info, rs, part_number); + if let Err(fallback) = get_object_chunk_range_guard(rs.as_ref()) { + return Ok(ChunkReadDecision::Fallback(fallback)); + } + + let content_type = info + .content_type + .as_ref() + .and_then(|content_type| ContentType::from_str(content_type).ok()); + let last_modified = info.mod_time.map(Timestamp::from); + let total_size = info.get_actual_size()?; + let (rs, response_content_length, content_range) = resolve_response_range(total_size, rs)?; + + Ok(ChunkReadDecision::Eligible(ChunkReadPlan { + rs, + content_type, + last_modified, + response_content_length, + content_range, + })) +} + +#[cfg(test)] +mod tests { + use super::*; + + struct MockCachedSource { + body: Arc, + content_length: i64, + content_type: Option, + e_tag: Option, + last_modified: Option, + cache_control: Option, + content_disposition: Option, + content_encoding: Option, + content_language: Option, + storage_class: Option, + version_id: Option, + delete_marker: bool, + tag_count: Option, + user_metadata: HashMap, + checksum_crc32: Option, + checksum_crc32c: Option, + checksum_sha1: Option, + checksum_sha256: Option, + checksum_crc64nvme: Option, + checksum_type: Option, + } + + impl CachedGetObjectSource for MockCachedSource { + fn body(&self) -> &Arc { + &self.body + } + + fn content_length(&self) -> i64 { + self.content_length + } + + fn content_type(&self) -> Option<&str> { + self.content_type.as_deref() + } + + fn e_tag(&self) -> Option<&str> { + self.e_tag.as_deref() + } + + fn last_modified(&self) -> Option<&str> { + self.last_modified.as_deref() + } + + fn cache_control(&self) -> Option<&str> { + self.cache_control.as_deref() + } + + fn content_disposition(&self) -> Option<&str> { + self.content_disposition.as_deref() + } + + fn content_encoding(&self) -> Option<&str> { + self.content_encoding.as_deref() + } + + fn content_language(&self) -> Option<&str> { + self.content_language.as_deref() + } + + fn storage_class(&self) -> Option<&str> { + self.storage_class.as_deref() + } + + fn version_id(&self) -> Option<&str> { + self.version_id.as_deref() + } + + fn delete_marker(&self) -> bool { + self.delete_marker + } + + fn tag_count(&self) -> Option { + self.tag_count + } + + fn user_metadata(&self) -> &HashMap { + &self.user_metadata + } + + fn checksum_crc32(&self) -> Option<&str> { + self.checksum_crc32.as_deref() + } + + fn checksum_crc32c(&self) -> Option<&str> { + self.checksum_crc32c.as_deref() + } + + fn checksum_sha1(&self) -> Option<&str> { + self.checksum_sha1.as_deref() + } + + fn checksum_sha256(&self) -> Option<&str> { + self.checksum_sha256.as_deref() + } + + fn checksum_crc64nvme(&self) -> Option<&str> { + self.checksum_crc64nvme.as_deref() + } + + fn checksum_type(&self) -> Option<&ChecksumType> { + self.checksum_type.as_ref() + } + } + + #[test] + fn map_chunk_copy_mode_uses_expected_metric_modes() { + assert_eq!( + map_chunk_copy_mode(GetObjectChunkCopyMode::TrueZeroCopy), + rustfs_io_metrics::CopyMode::TrueZeroCopy + ); + assert_eq!( + map_chunk_copy_mode(GetObjectChunkCopyMode::SharedBytes), + rustfs_io_metrics::CopyMode::SharedBytes + ); + assert_eq!( + map_chunk_copy_mode(GetObjectChunkCopyMode::SingleCopy), + rustfs_io_metrics::CopyMode::SingleCopy + ); + assert_eq!( + map_chunk_copy_mode(GetObjectChunkCopyMode::Reconstructed), + rustfs_io_metrics::CopyMode::Reconstructed + ); + } + + #[test] + fn chunk_body_labels_keep_fast_path_and_copy_mode() { + let (path, copy_mode) = chunk_body_data_plane_labels(GetObjectChunkPath::Bridge, rustfs_io_metrics::CopyMode::SingleCopy); + assert_eq!(path, rustfs_io_metrics::IoPath::Fast); + assert_eq!(copy_mode, rustfs_io_metrics::CopyMode::SingleCopy); + } + + #[test] + fn get_object_chunk_fast_path_guard_rejects_ssec_requests() { + let err = get_object_chunk_fast_path_guard(true, false).unwrap_err(); + assert_eq!( + err, + ChunkReadFallback { + stage: rustfs_io_metrics::IoStage::ReadSetup, + reason: rustfs_io_metrics::FallbackReason::EncryptionEnabled, + } + ); + } + + #[test] + fn get_object_chunk_fast_path_guard_allows_plain_request() { + assert!(get_object_chunk_fast_path_guard(false, false).is_ok()); + } + + #[test] + fn get_object_chunk_range_guard_rejects_invalid_suffix_range() { + let err = get_object_chunk_range_guard(Some(&HTTPRangeSpec { + is_suffix_length: true, + start: 4, + end: 0, + })) + .unwrap_err(); + + assert_eq!( + err, + ChunkReadFallback { + stage: rustfs_io_metrics::IoStage::RangeGuard, + reason: rustfs_io_metrics::FallbackReason::RangeNotSupported, + } + ); + } + + #[test] + fn build_get_object_checksums_returns_default_when_mode_absent() { + let checksums = build_get_object_checksums(&ObjectInfo::default(), &HeaderMap::new(), None, None).unwrap(); + assert_eq!(checksums, GetObjectChecksums::default()); + } + + #[test] + fn plan_chunk_read_returns_fallback_for_compressed_object() { + let mut info = ObjectInfo::default(); + rustfs_utils::http::insert_str( + &mut info.user_defined, + rustfs_utils::http::SUFFIX_COMPRESSION, + rustfs_utils::CompressionAlgorithm::Zstd.to_string(), + ); + + let decision = plan_chunk_read(&info, true, None, None).unwrap(); + assert!( + matches!( + decision, + ChunkReadDecision::Fallback(ChunkReadFallback { + stage: rustfs_io_metrics::IoStage::ReadSetup, + reason: rustfs_io_metrics::FallbackReason::CompressionEnabled, + }) + ), + "unexpected decision: {:?}", + decision + ); + } + + #[test] + fn plan_chunk_read_returns_fallback_for_transitioned_object() { + let mut info = ObjectInfo::default(); + info.transitioned_object.status = TRANSITION_COMPLETE.to_string(); + + let decision = plan_chunk_read(&info, true, None, None).unwrap(); + assert!(matches!( + decision, + ChunkReadDecision::Fallback(ChunkReadFallback { + stage: rustfs_io_metrics::IoStage::ReadSetup, + reason: rustfs_io_metrics::FallbackReason::NonLocalBackend, + }) + )); + } + + #[test] + fn plan_chunk_read_returns_range_guard_fallback_for_invalid_range() { + let info = ObjectInfo { + size: 16, + actual_size: 16, + ..Default::default() + }; + + let decision = plan_chunk_read( + &info, + true, + Some(HTTPRangeSpec { + is_suffix_length: false, + start: -1, + end: 4, + }), + None, + ) + .unwrap(); + + assert!(matches!( + decision, + ChunkReadDecision::Fallback(ChunkReadFallback { + stage: rustfs_io_metrics::IoStage::RangeGuard, + reason: rustfs_io_metrics::FallbackReason::RangeNotSupported, + }) + )); + } + + #[test] + fn plan_chunk_read_allows_suffix_range() { + let info = ObjectInfo { + size: 16, + actual_size: 16, + ..Default::default() + }; + + let decision = plan_chunk_read( + &info, + true, + Some(HTTPRangeSpec { + is_suffix_length: true, + start: 4, + end: -1, + }), + None, + ) + .unwrap(); + + let ChunkReadDecision::Eligible(plan) = decision else { + panic!("expected eligible plan"); + }; + let rs = plan.rs.expect("suffix range should be preserved"); + assert!(rs.is_suffix_length); + assert_eq!(rs.start, 4); + assert_eq!(plan.response_content_length, 4); + assert_eq!(plan.content_range.as_deref(), Some("bytes 12-15/16")); + } + + #[test] + fn plan_chunk_read_uses_part_number_range_when_available() { + let mut info = ObjectInfo { + size: 12, + actual_size: 12, + content_type: Some("application/octet-stream".to_string()), + ..Default::default() + }; + info.parts = vec![Default::default(), Default::default()]; + info.parts[0].number = 1; + info.parts[0].size = 5; + info.parts[0].actual_size = 5; + info.parts[1].number = 2; + info.parts[1].size = 7; + info.parts[1].actual_size = 7; + + let decision = plan_chunk_read(&info, true, None, Some(2)).unwrap(); + let ChunkReadDecision::Eligible(plan) = decision else { + panic!("expected eligible plan"); + }; + let rs = plan.rs.expect("range from part number"); + assert_eq!(rs.start, 5); + assert_eq!(rs.end, 11); + assert_eq!(plan.response_content_length, 7); + assert_eq!(plan.content_range.as_deref(), Some("bytes 5-11/12")); + assert!(plan.content_type.is_some()); + } + + #[test] + fn plan_chunk_read_returns_delete_marker_errors() { + let info = ObjectInfo { + delete_marker: true, + ..Default::default() + }; + + let err = plan_chunk_read(&info, true, None, None).unwrap_err(); + assert!(matches!(err, ChunkReadPlanError::NoSuchKey)); + + let err = plan_chunk_read(&info, false, None, None).unwrap_err(); + assert!(matches!(err, ChunkReadPlanError::MethodNotAllowed)); + } + + #[test] + fn plan_legacy_read_uses_part_number_range_when_available() { + let mut info = ObjectInfo { + size: 12, + actual_size: 12, + content_type: Some("application/octet-stream".to_string()), + ..Default::default() + }; + info.parts = vec![Default::default(), Default::default()]; + info.parts[0].number = 1; + info.parts[0].size = 5; + info.parts[0].actual_size = 5; + info.parts[1].number = 2; + info.parts[1].size = 7; + info.parts[1].actual_size = 7; + + let plan = plan_legacy_read(&info, None, Some(2)).unwrap(); + + let rs = plan.rs.expect("range from part number"); + assert_eq!(rs.start, 5); + assert_eq!(rs.end, 11); + assert_eq!(plan.response_content_length, 7); + assert_eq!(plan.content_range.as_deref(), Some("bytes 5-11/12")); + assert!(plan.content_type.is_some()); + } + + #[test] + fn build_reader_read_setup_uses_encryption_length_override() { + let plan = LegacyReadPlan { + rs: None, + content_type: None, + last_modified: None, + response_content_length: 16, + content_range: None, + }; + let encryption_state = GetObjectEncryptionState { + encryption_applied: true, + response_content_length_override: Some(12), + ..Default::default() + }; + let reader = Box::new(rustfs_rio::WarpReader::new(tokio::io::empty())) as Box; + + let setup = build_reader_read_setup(ObjectInfo::default(), ObjectInfo::default(), reader, plan, encryption_state); + + assert!(setup.encryption_applied); + assert_eq!(setup.response_content_length, 12); + match setup.body_source { + GetObjectBodySource::Reader(_) => {} + GetObjectBodySource::Chunk { .. } => panic!("expected reader body source"), + } + } + + #[test] + fn plan_get_object_body_prefers_cache_writeback_when_cacheable() { + let plan = plan_get_object_body( + GetObjectCacheEligibility { + cache_enabled: true, + cache_writeback_enabled: true, + is_part_request: false, + is_range_request: false, + encryption_applied: false, + response_size: 1024, + max_cacheable_size: 2048, + }, + 4096, + ); + + assert_eq!(plan, GetObjectBodyPlan::CacheWriteback); + } + + #[test] + fn cache_served_metric_contract_is_mutually_exclusive_with_cache_writeback() { + let contract = GetObjectDataPlaneMetricContract::cache_served(); + + assert_eq!(contract.request_source, GetObjectDataPlaneRequestSource::CacheServed); + assert_eq!(contract.io_path, rustfs_io_metrics::IoPath::Fast); + assert_eq!(contract.copy_mode, rustfs_io_metrics::CopyMode::SharedBytes); + assert!(contract.record_cache_served_metric); + assert!(!contract.record_cache_writeback_metric); + } + + #[test] + fn disk_metric_contract_can_mark_cache_writeback_without_reclassifying_request_source() { + let contract = GetObjectDataPlaneMetricContract::disk( + rustfs_io_metrics::IoPath::Legacy, + rustfs_io_metrics::CopyMode::SingleCopy, + GetObjectBodyPlan::CacheWriteback, + ); + + assert_eq!(contract.request_source, GetObjectDataPlaneRequestSource::Disk); + assert_eq!(contract.io_path, rustfs_io_metrics::IoPath::Legacy); + assert_eq!(contract.copy_mode, rustfs_io_metrics::CopyMode::SingleCopy); + assert!(!contract.record_cache_served_metric); + assert!(contract.record_cache_writeback_metric); + } + + #[test] + fn plan_get_object_body_uses_encrypted_buffer_for_small_plain_request() { + let plan = plan_get_object_body( + GetObjectCacheEligibility { + cache_enabled: false, + cache_writeback_enabled: false, + is_part_request: false, + is_range_request: false, + encryption_applied: true, + response_size: 1024, + max_cacheable_size: 0, + }, + 4096, + ); + + assert_eq!(plan, GetObjectBodyPlan::BufferEncrypted); + } + + #[test] + fn get_object_sequential_hint_distinguishes_prefix_and_suffix_ranges() { + assert!(get_object_sequential_hint(None)); + assert!(get_object_sequential_hint(Some(&HTTPRangeSpec { + is_suffix_length: false, + start: 0, + end: -1, + }))); + assert!(!get_object_sequential_hint(Some(&HTTPRangeSpec { + is_suffix_length: false, + start: 4, + end: 8, + }))); + assert!(!get_object_sequential_hint(Some(&HTTPRangeSpec { + is_suffix_length: true, + start: 4, + end: -1, + }))); + } + + #[test] + fn plan_get_object_strategy_layout_caps_buffer_to_fallback() { + let layout = plan_get_object_strategy_layout(None, 1024, 8192, 4096); + + assert!(layout.is_sequential_hint); + assert_eq!(layout.optimal_buffer_size, 4096); + } + + #[test] + fn build_get_object_cache_writeback_formats_metadata() { + let info = ObjectInfo { + content_type: Some("application/octet-stream".to_string()), + etag: Some("abc123".to_string()), + mod_time: Some(OffsetDateTime::UNIX_EPOCH), + checksum: rustfs_rio::Checksum::new_from_data(rustfs_rio::ChecksumType::CRC32, b"abc") + .map(|checksum| checksum.to_bytes(&[])), + ..Default::default() + }; + + let writeback = build_get_object_cache_writeback(&info, Bytes::from_static(b"abc"), 3); + + assert_eq!(*writeback.body, Bytes::from_static(b"abc")); + assert_eq!(writeback.content_length, 3); + assert_eq!(writeback.content_type.as_deref(), Some("application/octet-stream")); + assert_eq!(writeback.e_tag.as_deref(), Some("abc123")); + assert_eq!(writeback.last_modified.as_deref(), Some("1970-01-01T00:00:00Z")); + assert_eq!(writeback.checksum_crc32.as_deref(), Some("NSRBwg==")); + } + + #[test] + fn finalize_get_object_cache_writeback_applies_http_metadata_and_user_metadata() { + let mut info = ObjectInfo { + expires: Some(OffsetDateTime::UNIX_EPOCH), + ..Default::default() + }; + info.user_defined + .insert("cache-control".to_string(), "max-age=3600".to_string()); + info.user_defined + .insert("content-disposition".to_string(), "attachment".to_string()); + info.user_defined.insert("content-language".to_string(), "en-US".to_string()); + + let writeback = finalize_get_object_cache_writeback( + &info, + GetObjectCacheWriteback { + body: Arc::new(Bytes::from_static(b"abc")), + content_length: 3, + content_type: None, + content_encoding: None, + cache_control: None, + content_disposition: None, + content_language: None, + expires: None, + storage_class: None, + version_id: None, + delete_marker: false, + user_metadata: HashMap::new(), + e_tag: None, + last_modified: None, + checksum_crc32: None, + checksum_crc32c: None, + checksum_sha1: None, + checksum_sha256: None, + checksum_crc64nvme: None, + checksum_type: None, + }, + HashMap::from([(String::from("custom"), String::from("value"))]), + ); + + assert_eq!(writeback.cache_control.as_deref(), Some("max-age=3600")); + assert_eq!(writeback.content_disposition.as_deref(), Some("attachment")); + assert_eq!(writeback.content_language.as_deref(), Some("en-US")); + assert_eq!(writeback.expires.as_deref(), Some("1970-01-01T00:00:00Z")); + assert_eq!(writeback.user_metadata.get("custom").map(String::as_str), Some("value")); + } + + #[test] + fn frozen_get_object_body_reuses_same_shared_bytes_for_cache_writeback() { + let frozen = FrozenGetObjectBody::new(Bytes::from_static(b"abc")); + let shared = Arc::clone(frozen.shared_body()); + assert_eq!(*shared, Bytes::from_static(b"abc")); + assert!(Arc::ptr_eq(&shared, frozen.shared_body())); + } + + #[test] + fn build_output_version_id_maps_nil_uuid_to_null() { + let nil = uuid::Uuid::nil(); + let version_id = build_output_version_id(true, Some(&nil)); + + assert_eq!(version_id.as_deref(), Some("null")); + } + + #[test] + fn build_get_object_output_context_preserves_copy_mode_override() { + let info = ObjectInfo { + version_id: Some(uuid::Uuid::nil()), + ..Default::default() + }; + let output_context = build_get_object_output_context( + None, + info, + ObjectInfo::default(), + None, + None, + 8, + None, + None, + None, + None, + None, + &GetObjectChecksums::default(), + None, + true, + 4096, + Some(rustfs_io_metrics::CopyMode::Reconstructed), + ); + + assert_eq!(output_context.output.version_id.as_deref(), Some("null")); + assert_eq!(output_context.copy_mode_override, Some(rustfs_io_metrics::CopyMode::Reconstructed)); + assert_eq!(output_context.optimal_buffer_size, 4096); + } + + #[test] + fn build_cached_get_object_flow_result_from_source_builds_plain_mode() { + let result = build_cached_get_object_flow_result_from_source( + "bucket", + "key", + &MockCachedSource { + body: Arc::new(Bytes::from_static(b"abc")), + content_length: 3, + content_type: None, + e_tag: None, + last_modified: None, + cache_control: None, + content_disposition: None, + content_encoding: None, + content_language: None, + storage_class: None, + version_id: None, + delete_marker: false, + tag_count: None, + user_metadata: HashMap::new(), + checksum_crc32: Some("crc32".to_string()), + checksum_crc32c: None, + checksum_sha1: None, + checksum_sha256: None, + checksum_crc64nvme: None, + checksum_type: Some(ChecksumType::from_static(ChecksumType::FULL_OBJECT)), + }, + "vid".to_string(), + ); + + assert!(matches!(result.response_mode, GetObjectResponseMode::Plain)); + assert_eq!(result.version_id_for_event, "vid"); + assert_eq!(result.event_info.bucket, "bucket"); + assert_eq!(result.event_info.name, "key"); + assert_eq!(result.output.checksum_crc32.as_deref(), Some("crc32")); + assert_eq!(result.output.checksum_type, Some(ChecksumType::from_static(ChecksumType::FULL_OBJECT))); + } + + #[test] + fn build_cors_wrapped_get_object_flow_result_uses_wrapped_mode() { + let result = build_cors_wrapped_get_object_flow_result( + GetObjectOutputContext { + output: GetObjectOutput::default(), + event_info: ObjectInfo::default(), + response_content_length: 1, + optimal_buffer_size: 1024, + copy_mode_override: None, + }, + "vid".to_string(), + ); + + assert!(matches!(result.response_mode, GetObjectResponseMode::CorsWrapped)); + assert_eq!(result.version_id_for_event, "vid"); + } + + #[test] + fn finalize_chunk_read_setup_preserves_body_source_and_io_path() { + let chunk_result = GetObjectChunkResult { + stream: Box::pin(futures_util::stream::empty::>()), + path: GetObjectChunkPath::Direct, + copy_mode: GetObjectChunkCopyMode::Reconstructed, + }; + let plan = ChunkReadPlan { + rs: Some(HTTPRangeSpec { + is_suffix_length: false, + start: 0, + end: 7, + }), + content_type: None, + last_modified: None, + response_content_length: 8, + content_range: Some("bytes 0-7/8".to_string()), + }; + + let result = finalize_chunk_read_setup(ObjectInfo::default(), ObjectInfo::default(), chunk_result, plan); + + assert_eq!(result.io_path, rustfs_io_metrics::IoPath::Fast); + match result.read_setup.body_source { + GetObjectBodySource::Chunk { path, copy_mode, .. } => { + assert_eq!(path, GetObjectChunkPath::Direct); + assert_eq!(copy_mode, rustfs_io_metrics::CopyMode::Reconstructed); + } + GetObjectBodySource::Reader(_) => panic!("expected chunk body source"), + } + assert_eq!(result.read_setup.response_content_length, 8); + assert_eq!(result.read_setup.content_range.as_deref(), Some("bytes 0-7/8")); + } +} diff --git a/crates/object-io/src/lib.rs b/crates/object-io/src/lib.rs new file mode 100644 index 000000000..9de88dbe3 --- /dev/null +++ b/crates/object-io/src/lib.rs @@ -0,0 +1,16 @@ +// 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. + +pub mod get; +pub mod put; diff --git a/crates/object-io/src/put.rs b/crates/object-io/src/put.rs new file mode 100644 index 000000000..7085a6468 --- /dev/null +++ b/crates/object-io/src/put.rs @@ -0,0 +1,1564 @@ +// 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 bytes::Buf; +use futures_util::{Stream, StreamExt}; +use http::HeaderMap; +use rustfs_ecstore::compress::{MIN_COMPRESSIBLE_SIZE, is_compressible}; +use rustfs_ecstore::store_api::ObjectOptions; +use rustfs_rio::{ + BlockReadable, BoxReadBlockFuture, Checksum, EtagResolvable, HashReader, HashReaderDetector, Reader, TryGetIndex, WarpReader, +}; +use rustfs_utils::http::AMZ_SERVER_SIDE_ENCRYPTION_CUSTOMER_ALGORITHM; +use rustfs_utils::http::headers::{ + AMZ_DECODED_CONTENT_LENGTH, AMZ_MINIO_SNOWBALL_IGNORE_DIRS, AMZ_MINIO_SNOWBALL_IGNORE_ERRORS, AMZ_MINIO_SNOWBALL_PREFIX, + AMZ_RUSTFS_SNOWBALL_IGNORE_DIRS, AMZ_RUSTFS_SNOWBALL_IGNORE_ERRORS, AMZ_RUSTFS_SNOWBALL_PREFIX, AMZ_SERVER_SIDE_ENCRYPTION, + AMZ_SERVER_SIDE_ENCRYPTION_KMS_ID, AMZ_SNOWBALL_EXTRACT, AMZ_SNOWBALL_IGNORE_DIRS, AMZ_SNOWBALL_IGNORE_ERRORS, + AMZ_SNOWBALL_PREFIX, +}; +use s3s::dto::{ChecksumAlgorithm, PutObjectInput, ServerSideEncryption}; +use s3s::{S3Error, s3_error}; +use std::collections::HashMap; +use std::pin::Pin; +use tokio::io::AsyncRead; +use tokio_tar::Archive; +use tokio_util::io::StreamReader; + +pub const AMZ_SNOWBALL_EXTRACT_COMPAT: &str = "X-Amz-Snowball-Auto-Extract"; +pub const AMZ_SNOWBALL_PREFIX_INTERNAL: &str = "X-Amz-Meta-Rustfs-Snowball-Prefix"; +pub const AMZ_SNOWBALL_IGNORE_DIRS_INTERNAL: &str = "X-Amz-Meta-Rustfs-Snowball-Ignore-Dirs"; +pub const AMZ_SNOWBALL_IGNORE_ERRORS_INTERNAL: &str = "X-Amz-Meta-Rustfs-Snowball-Ignore-Errors"; + +const AMZ_META_PREFIX_LOWER: &str = "x-amz-meta-"; +const SNOWBALL_PREFIX_SUFFIX_LOWER: &str = "snowball-prefix"; +const SNOWBALL_IGNORE_DIRS_SUFFIX_LOWER: &str = "snowball-ignore-dirs"; +const SNOWBALL_IGNORE_ERRORS_SUFFIX_LOWER: &str = "snowball-ignore-errors"; +const SNOWBALL_PREFIX_HEADER_KEYS: &[&str] = &[AMZ_MINIO_SNOWBALL_PREFIX, AMZ_SNOWBALL_PREFIX, AMZ_RUSTFS_SNOWBALL_PREFIX]; +const SNOWBALL_IGNORE_DIRS_HEADER_KEYS: &[&str] = &[ + AMZ_MINIO_SNOWBALL_IGNORE_DIRS, + AMZ_SNOWBALL_IGNORE_DIRS, + AMZ_RUSTFS_SNOWBALL_IGNORE_DIRS, +]; +const SNOWBALL_IGNORE_ERRORS_HEADER_KEYS: &[&str] = &[ + AMZ_MINIO_SNOWBALL_IGNORE_ERRORS, + AMZ_SNOWBALL_IGNORE_ERRORS, + AMZ_RUSTFS_SNOWBALL_IGNORE_ERRORS, +]; +pub const PUT_REDUCED_COPY_MIN_SIZE_BYTES: i64 = 1024 * 1024; + +#[derive(Debug, Clone, Default, PartialEq, Eq)] +pub struct PutObjectChecksums { + pub crc32: Option, + pub crc32c: Option, + pub sha1: Option, + pub sha256: Option, + pub crc64nvme: Option, +} + +impl PutObjectChecksums { + pub fn merge_from_map(&mut self, checksums: &HashMap) { + for (key, checksum) in checksums { + match rustfs_rio::ChecksumType::from_string(key.as_str()) { + rustfs_rio::ChecksumType::CRC32 => self.crc32 = Some(checksum.clone()), + rustfs_rio::ChecksumType::CRC32C => self.crc32c = Some(checksum.clone()), + rustfs_rio::ChecksumType::SHA1 => self.sha1 = Some(checksum.clone()), + rustfs_rio::ChecksumType::SHA256 => self.sha256 = Some(checksum.clone()), + rustfs_rio::ChecksumType::CRC64_NVME => self.crc64nvme = Some(checksum.clone()), + _ => {} + } + } + } +} + +#[derive(Debug, Clone, Copy, Default, PartialEq, Eq)] +pub struct PutObjectTransformStage { + compression_applied: bool, + encryption_applied: bool, +} + +impl PutObjectTransformStage { + pub fn mark_compression(&mut self) { + self.compression_applied = true; + } + + pub fn mark_encryption(&mut self) { + self.encryption_applied = true; + } + + #[must_use] + pub const fn compression_applied(self) -> bool { + self.compression_applied + } + + #[must_use] + pub const fn encryption_applied(self) -> bool { + self.encryption_applied + } + + #[must_use] + pub fn effective_copy_mode(self) -> rustfs_io_metrics::CopyMode { + resolve_put_effective_copy_mode(self.compression_applied, self.encryption_applied) + } + + #[must_use] + pub fn metric_kind(self) -> Option<&'static str> { + resolve_put_transform_metric_kind(self.compression_applied, self.encryption_applied) + } +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum PutObjectCompatIngressKind { + BufferedStreamCompat, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub struct PutObjectCompatIngressPlan { + pub kind: PutObjectCompatIngressKind, + pub buffer_size: usize, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum PutObjectIngressKind { + LegacyCompat, + ReducedCopyCandidate, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub struct PutObjectIngressPlan { + pub kind: PutObjectIngressKind, + pub compat: PutObjectCompatIngressPlan, + pub enable_zero_copy: bool, +} + +pub type BoxIngressStream = Pin> + Send + Sync + 'static>>; +pub type PutObjectCompatIngressStream = futures_util::stream::Map) -> std::io::Result>; +pub type PutObjectCompatIngress = tokio::io::BufReader, B>>; + +pub fn box_put_object_ingress_stream(body: S) -> BoxIngressStream +where + S: Stream> + Send + Sync + 'static, +{ + Box::pin(body) +} + +pub struct PutObjectReducedCopyIngress { + body: BoxIngressStream, + compat: PutObjectCompatIngressPlan, +} + +impl PutObjectReducedCopyIngress { + pub const fn new(body: BoxIngressStream, compat: PutObjectCompatIngressPlan) -> Self { + Self { body, compat } + } + + pub fn from_stream(body: S, compat: PutObjectCompatIngressPlan) -> Self + where + S: Stream> + Send + Sync + 'static, + { + Self::new(box_put_object_ingress_stream(body), compat) + } + + pub const fn compat_plan(&self) -> PutObjectCompatIngressPlan { + self.compat + } + + pub fn into_body(self) -> BoxIngressStream { + self.body + } +} + +impl PutObjectReducedCopyIngress +where + B: Buf, + E: std::fmt::Display, + PutObjectCompatIngress, B, E>: Send + Sync + Unpin + 'static, +{ + pub fn into_compat_reader(self) -> Box { + build_put_object_compat_reader(self.body, self.compat) + } +} + +pub enum PutObjectIngressSource { + LegacyCompat(Box), + ReducedCopyCandidate(PutObjectReducedCopyIngress), +} + +struct PutObjectReducedCopyReader { + body: S, + current_chunk: Option, + pending_error: Option, +} + +impl PutObjectReducedCopyReader { + fn new(body: S) -> Self { + Self { + body, + current_chunk: None, + pending_error: None, + } + } + + fn copy_chunk_into_slice(chunk: &mut B, buf: &mut [u8]) -> usize + where + B: Buf, + { + let mut copied = 0; + let to_copy = chunk.remaining().min(buf.len()); + + while copied < to_copy { + let slice = chunk.chunk(); + if !slice.is_empty() { + let len = (to_copy - copied).min(slice.len()); + buf[copied..copied + len].copy_from_slice(&slice[..len]); + chunk.advance(len); + copied += len; + continue; + } + + let dest = &mut buf[copied..to_copy]; + chunk.copy_to_slice(dest); + copied = to_copy; + } + + copied + } + + fn copy_chunk_into_read_buf(chunk: &mut B, buf: &mut tokio::io::ReadBuf<'_>) -> usize + where + B: Buf, + { + let to_copy = chunk.remaining().min(buf.remaining()); + let dest = &mut buf.initialize_unfilled()[..to_copy]; + let copied = Self::copy_chunk_into_slice(chunk, dest); + buf.advance(copied); + copied + } + + async fn read_into_slice(&mut self, buf: &mut [u8]) -> std::io::Result + where + S: Stream> + Unpin, + B: Buf + Unpin, + E: std::fmt::Display, + { + if let Some(err) = self.pending_error.take() { + return Err(err); + } + + if buf.is_empty() { + return Ok(0); + } + + let mut copied = 0; + + loop { + if copied == buf.len() { + return Ok(copied); + } + + if let Some(chunk) = self.current_chunk.as_mut() { + if !chunk.has_remaining() { + self.current_chunk = None; + continue; + } + + copied += Self::copy_chunk_into_slice(chunk, &mut buf[copied..]); + if chunk.has_remaining() || copied == buf.len() { + return Ok(copied); + } + + self.current_chunk = None; + continue; + } + + match self.body.next().await { + Some(Ok(chunk)) => { + if chunk.remaining() == 0 { + continue; + } + self.current_chunk = Some(chunk); + } + Some(Err(err)) => { + let err = std::io::Error::other(err.to_string()); + if copied > 0 { + self.pending_error = Some(err); + return Ok(copied); + } + return Err(err); + } + None => return Ok(copied), + } + } + } +} + +impl AsyncRead for PutObjectReducedCopyReader +where + S: Stream> + Unpin, + B: Buf + Unpin, + E: std::fmt::Display, +{ + fn poll_read( + self: std::pin::Pin<&mut Self>, + cx: &mut std::task::Context<'_>, + buf: &mut tokio::io::ReadBuf<'_>, + ) -> std::task::Poll> { + let this = self.get_mut(); + + if let Some(err) = this.pending_error.take() { + return std::task::Poll::Ready(Err(err)); + } + + let mut read_any = false; + + loop { + if buf.remaining() == 0 { + return std::task::Poll::Ready(Ok(())); + } + + if let Some(chunk) = this.current_chunk.as_mut() { + if !chunk.has_remaining() { + this.current_chunk = None; + continue; + } + + if Self::copy_chunk_into_read_buf(chunk, buf) > 0 { + read_any = true; + } + + if chunk.has_remaining() || buf.remaining() == 0 { + return std::task::Poll::Ready(Ok(())); + } + + this.current_chunk = None; + continue; + } + + match std::pin::Pin::new(&mut this.body).poll_next(cx) { + std::task::Poll::Ready(Some(Ok(chunk))) => { + if chunk.remaining() == 0 { + continue; + } + this.current_chunk = Some(chunk); + } + std::task::Poll::Ready(Some(Err(err))) => { + let err = std::io::Error::other(err.to_string()); + if read_any { + this.pending_error = Some(err); + return std::task::Poll::Ready(Ok(())); + } + return std::task::Poll::Ready(Err(err)); + } + std::task::Poll::Ready(None) => return std::task::Poll::Ready(Ok(())), + std::task::Poll::Pending => { + if read_any { + return std::task::Poll::Ready(Ok(())); + } + return std::task::Poll::Pending; + } + } + } + } +} + +impl BlockReadable for PutObjectReducedCopyReader +where + S: Stream> + Unpin + Send + Sync, + B: Buf + Unpin + Send + Sync, + E: std::fmt::Display + Send + Sync, +{ + fn read_block<'a>(&'a mut self, buf: &'a mut [u8]) -> BoxReadBlockFuture<'a> { + Box::pin(async move { self.read_into_slice(buf).await }) + } +} + +impl EtagResolvable for PutObjectReducedCopyReader {} + +impl HashReaderDetector for PutObjectReducedCopyReader {} + +impl TryGetIndex for PutObjectReducedCopyReader {} + +fn build_put_object_reduced_copy_reader(candidate: PutObjectReducedCopyIngress) -> Box +where + B: Buf + Send + Sync + Unpin + 'static, + E: std::fmt::Display + Send + Sync + 'static, +{ + Box::new(PutObjectReducedCopyReader::new(candidate.into_body())) +} + +impl PutObjectIngressSource +where + B: Buf, + E: std::fmt::Display, + PutObjectCompatIngress, B, E>: Send + Sync + Unpin + 'static, +{ + pub fn into_compat_reader(self) -> Box { + match self { + Self::LegacyCompat(reader) => reader, + Self::ReducedCopyCandidate(candidate) => candidate.into_compat_reader(), + } + } + + pub fn into_reader(self) -> Box + where + B: Buf + Send + Sync + Unpin + 'static, + E: std::fmt::Display + Send + Sync + 'static, + { + match self { + Self::LegacyCompat(reader) => reader, + Self::ReducedCopyCandidate(candidate) => build_put_object_reduced_copy_reader(candidate), + } + } +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum PutObjectPlainBodyKind { + LegacyCompat, + ReducedCopyCandidate, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum PutObjectBodyKind { + Plain(PutObjectPlainBodyKind), + Compressed, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub struct PutObjectBodyPlan { + pub ingress: PutObjectIngressPlan, + pub kind: PutObjectBodyKind, +} + +impl PutObjectBodyPlan { + pub const fn should_compress(&self) -> bool { + matches!(self.kind, PutObjectBodyKind::Compressed) + } + + pub const fn plain_body_kind(&self) -> Option { + match self.kind { + PutObjectBodyKind::Plain(kind) => Some(kind), + PutObjectBodyKind::Compressed => None, + } + } +} + +#[derive(Debug, Clone, Default, PartialEq, Eq)] +pub struct PutObjectExtractOptions { + pub prefix: Option, + pub ignore_dirs: bool, + pub ignore_errors: bool, +} + +#[derive(Debug, Clone, Default, PartialEq, Eq)] +pub struct PutObjectLegacyHashValues { + pub md5hex: Option, + pub sha256hex: Option, +} + +impl PutObjectLegacyHashValues { + pub fn clear_for_transformed_body(&mut self) { + self.md5hex = None; + self.sha256hex = None; + } +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub struct PutObjectLegacyHashStagePlan { + pub size: i64, + pub actual_size: i64, + pub apply_s3_checksum: bool, + pub ignore_s3_checksum_value: bool, +} + +pub struct PutObjectHashStage { + pub reader: HashReader, + pub want_checksum: Option, + pub ingress_kind: PutObjectIngressKind, +} + +fn build_put_object_hash_stage( + reader: Box, + ingress_kind: PutObjectIngressKind, + hash_values: PutObjectLegacyHashValues, + plan: PutObjectLegacyHashStagePlan, + headers: &HeaderMap, + trailing_headers: Option, +) -> std::io::Result { + let mut reader = HashReader::new(reader, plan.size, plan.actual_size, hash_values.md5hex, hash_values.sha256hex, false)?; + let requested_checksum_type = rustfs_rio::ChecksumType::from_header(headers); + let want_checksum = if plan.apply_s3_checksum { + reader.add_checksum_from_s3s(headers, trailing_headers, plan.ignore_s3_checksum_value)?; + if requested_checksum_type.is_set() && reader.checksum().is_none() { + reader.enable_auto_checksum(requested_checksum_type)?; + } + reader.checksum() + } else { + None + }; + + Ok(PutObjectHashStage { + reader, + want_checksum, + ingress_kind, + }) +} + +pub fn apply_trailing_checksums( + algorithm: Option<&str>, + trailing_headers: &Option, + checksums: &mut PutObjectChecksums, +) { + let Some(alg) = algorithm else { return }; + let Some(checksum_str) = trailing_headers.as_ref().and_then(|trailer| { + let key = match alg { + ChecksumAlgorithm::CRC32 => rustfs_rio::ChecksumType::CRC32.key(), + ChecksumAlgorithm::CRC32C => rustfs_rio::ChecksumType::CRC32C.key(), + ChecksumAlgorithm::SHA1 => rustfs_rio::ChecksumType::SHA1.key(), + ChecksumAlgorithm::SHA256 => rustfs_rio::ChecksumType::SHA256.key(), + ChecksumAlgorithm::CRC64NVME => rustfs_rio::ChecksumType::CRC64_NVME.key(), + _ => return None, + }; + trailer.read(|headers| { + headers + .get(key.unwrap_or_default()) + .and_then(|value| value.to_str().ok().map(|s| s.to_string())) + }) + }) else { + return; + }; + + match alg { + ChecksumAlgorithm::CRC32 => checksums.crc32 = checksum_str, + ChecksumAlgorithm::CRC32C => checksums.crc32c = checksum_str, + ChecksumAlgorithm::SHA1 => checksums.sha1 = checksum_str, + ChecksumAlgorithm::SHA256 => checksums.sha256 = checksum_str, + ChecksumAlgorithm::CRC64NVME => checksums.crc64nvme = checksum_str, + _ => (), + } +} + +pub fn resolve_put_body_size(content_length: Option, headers: &HeaderMap) -> s3s::S3Result { + let size = match content_length { + Some(c) => c, + None => { + if let Some(val) = headers.get(AMZ_DECODED_CONTENT_LENGTH) { + match atoi::atoi::(val.as_bytes()) { + Some(x) => x, + None => return Err(s3_error!(UnexpectedContent)), + } + } else { + return Err(s3_error!(UnexpectedContent)); + } + } + }; + + if size == -1 { + return Err(s3_error!(UnexpectedContent)); + } + + Ok(size) +} + +pub fn should_use_zero_copy(size: i64, headers: &HeaderMap) -> bool { + const ZERO_COPY_MIN_SIZE: i64 = 1024 * 1024; + + if size <= ZERO_COPY_MIN_SIZE { + return false; + } + + !has_put_encryption_headers(headers) && !put_request_is_compressible(headers) +} + +fn has_put_encryption_headers(headers: &HeaderMap) -> bool { + headers.get(AMZ_SERVER_SIDE_ENCRYPTION).is_some() + || headers.get(AMZ_SERVER_SIDE_ENCRYPTION_CUSTOMER_ALGORITHM).is_some() + || headers.get(AMZ_SERVER_SIDE_ENCRYPTION_KMS_ID).is_some() +} + +fn put_request_is_compressible(headers: &HeaderMap) -> bool { + if let Some(content_type) = headers.get("content-type") + && let Ok(ct) = content_type.to_str() + { + let compressible_types = [ + "text/plain", + "text/html", + "text/css", + "text/javascript", + "application/javascript", + "application/json", + "application/xml", + "text/xml", + ]; + return compressible_types.iter().any(|ty| ct.contains(ty)); + } + + false +} + +fn should_use_put_reduced_copy_candidate( + size: i64, + headers: &HeaderMap, + encryption_enabled: bool, + compression_enabled: bool, +) -> bool { + if size <= PUT_REDUCED_COPY_MIN_SIZE_BYTES { + return false; + } + + if !encryption_enabled && has_put_encryption_headers(headers) { + return false; + } + + compression_enabled || !put_request_is_compressible(headers) +} + +fn map_put_object_ingress_error(result: Result) -> std::io::Result +where + E: std::fmt::Display, +{ + result.map_err(|err| std::io::Error::other(err.to_string())) +} + +pub fn build_put_object_compat_ingress(body: S, plan: PutObjectCompatIngressPlan) -> PutObjectCompatIngress +where + S: Stream>, + B: Buf, + E: std::fmt::Display, +{ + match plan.kind { + PutObjectCompatIngressKind::BufferedStreamCompat => tokio::io::BufReader::with_capacity( + plan.buffer_size, + StreamReader::new(body.map(map_put_object_ingress_error:: as fn(Result) -> std::io::Result)), + ), + } +} + +pub fn build_put_object_compat_reader(body: S, plan: PutObjectCompatIngressPlan) -> Box +where + S: Stream>, + B: Buf, + E: std::fmt::Display, + PutObjectCompatIngress: Send + Sync + Unpin + 'static, +{ + Box::new(WarpReader::new(build_put_object_compat_ingress(body, plan))) +} + +pub fn build_put_object_ingress_source(body: S, plan: PutObjectBodyPlan) -> PutObjectIngressSource +where + S: Stream> + Send + Sync + 'static, + B: Buf + 'static, + E: std::fmt::Display + 'static, + PutObjectCompatIngress, B, E>: Send + Sync + Unpin + 'static, +{ + match (plan.kind, plan.ingress.kind) { + (PutObjectBodyKind::Compressed, PutObjectIngressKind::ReducedCopyCandidate) + | (PutObjectBodyKind::Plain(PutObjectPlainBodyKind::ReducedCopyCandidate), PutObjectIngressKind::ReducedCopyCandidate) => { + PutObjectIngressSource::ReducedCopyCandidate(PutObjectReducedCopyIngress::from_stream(body, plan.ingress.compat)) + } + (PutObjectBodyKind::Compressed, _) + | (PutObjectBodyKind::Plain(PutObjectPlainBodyKind::LegacyCompat), _) + | (PutObjectBodyKind::Plain(PutObjectPlainBodyKind::ReducedCopyCandidate), PutObjectIngressKind::LegacyCompat) => { + PutObjectIngressSource::LegacyCompat(build_put_object_compat_reader( + box_put_object_ingress_stream(body), + plan.ingress.compat, + )) + } + } +} + +pub fn build_put_object_legacy_hash_stage( + reader: Box, + hash_values: PutObjectLegacyHashValues, + plan: PutObjectLegacyHashStagePlan, + headers: &HeaderMap, + trailing_headers: Option, +) -> std::io::Result { + build_put_object_hash_stage(reader, PutObjectIngressKind::LegacyCompat, hash_values, plan, headers, trailing_headers) +} + +pub fn build_put_object_plain_hash_stage( + ingress: PutObjectIngressSource, + hash_values: PutObjectLegacyHashValues, + plan: PutObjectLegacyHashStagePlan, + headers: &HeaderMap, + trailing_headers: Option, +) -> std::io::Result +where + B: Buf + Send + Sync + Unpin + 'static, + E: std::fmt::Display + Send + Sync + 'static, +{ + match ingress { + PutObjectIngressSource::LegacyCompat(reader) => { + build_put_object_hash_stage(reader, PutObjectIngressKind::LegacyCompat, hash_values, plan, headers, trailing_headers) + } + PutObjectIngressSource::ReducedCopyCandidate(candidate) => build_put_object_hash_stage( + build_put_object_reduced_copy_reader(candidate), + PutObjectIngressKind::ReducedCopyCandidate, + hash_values, + plan, + headers, + trailing_headers, + ), + } +} + +pub fn plan_put_object_ingress(size: i64, headers: &HeaderMap, buffer_size: usize) -> PutObjectIngressPlan { + plan_put_object_ingress_with_transforms(size, headers, buffer_size, false, false) +} + +pub fn plan_put_object_ingress_with_transforms( + size: i64, + headers: &HeaderMap, + buffer_size: usize, + encryption_enabled: bool, + compression_enabled: bool, +) -> PutObjectIngressPlan { + let enable_zero_copy = should_use_put_reduced_copy_candidate(size, headers, encryption_enabled, compression_enabled); + PutObjectIngressPlan { + kind: if enable_zero_copy { + PutObjectIngressKind::ReducedCopyCandidate + } else { + PutObjectIngressKind::LegacyCompat + }, + compat: PutObjectCompatIngressPlan { + kind: PutObjectCompatIngressKind::BufferedStreamCompat, + buffer_size, + }, + enable_zero_copy, + } +} + +pub fn plan_put_object_body(size: i64, headers: &HeaderMap, key: &str, buffer_size: usize) -> PutObjectBodyPlan { + plan_put_object_body_with_transforms(size, headers, key, buffer_size, false) +} + +pub fn plan_put_object_body_with_transforms( + size: i64, + headers: &HeaderMap, + key: &str, + buffer_size: usize, + encryption_enabled: bool, +) -> PutObjectBodyPlan { + let compression_enabled = size > MIN_COMPRESSIBLE_SIZE as i64 && is_compressible(headers, key); + let ingress = plan_put_object_ingress_with_transforms(size, headers, buffer_size, encryption_enabled, compression_enabled); + let kind = if compression_enabled { + PutObjectBodyKind::Compressed + } else { + PutObjectBodyKind::Plain(match ingress.kind { + PutObjectIngressKind::LegacyCompat => PutObjectPlainBodyKind::LegacyCompat, + PutObjectIngressKind::ReducedCopyCandidate => PutObjectPlainBodyKind::ReducedCopyCandidate, + }) + }; + + PutObjectBodyPlan { ingress, kind } +} + +pub fn resolve_put_effective_copy_mode(applied_compression: bool, applied_encryption: bool) -> rustfs_io_metrics::CopyMode { + if applied_compression || applied_encryption { + rustfs_io_metrics::CopyMode::Transformed + } else { + rustfs_io_metrics::CopyMode::SingleCopy + } +} + +pub fn resolve_put_transform_metric_kind(applied_compression: bool, applied_encryption: bool) -> Option<&'static str> { + match (applied_compression, applied_encryption) { + (true, true) => Some("compression_encryption"), + (true, false) => Some("compression"), + (false, true) => Some("encryption"), + (false, false) => None, + } +} + +pub fn resolve_put_transformed_fallback_reason( + ingress_kind: PutObjectIngressKind, + compressed: bool, + encryption_enabled: bool, +) -> Option { + if ingress_kind == PutObjectIngressKind::ReducedCopyCandidate { + return None; + } + + match (compressed, encryption_enabled) { + (true, true) => Some(rustfs_io_metrics::FallbackReason::TransformCompressionEncryptionLegacy), + (true, false) => Some(rustfs_io_metrics::FallbackReason::TransformCompressionLegacy), + (false, true) => Some(rustfs_io_metrics::FallbackReason::TransformEncryptionLegacy), + (false, false) => None, + } +} + +pub fn header_value_is_true(headers: &HeaderMap, key: &str) -> bool { + headers + .get(key) + .and_then(|value| value.to_str().ok()) + .is_some_and(|value| value.trim().eq_ignore_ascii_case("true")) +} + +pub fn is_put_object_extract_requested(headers: &HeaderMap) -> bool { + header_value_is_true(headers, AMZ_SNOWBALL_EXTRACT) || header_value_is_true(headers, AMZ_SNOWBALL_EXTRACT_COMPAT) +} + +fn trimmed_header_value(headers: &HeaderMap, key: &str) -> Option { + headers + .get(key) + .and_then(|value| value.to_str().ok()) + .map(|value| value.trim().to_string()) +} + +fn is_exact_snowball_meta_key(key: &str, exact_keys: &[&str]) -> bool { + exact_keys.iter().any(|exact_key| key.eq_ignore_ascii_case(exact_key)) +} + +fn snowball_meta_value_by_suffix(headers: &HeaderMap, suffix_lower: &str, exact_keys: &[&str]) -> Option { + for (name, value) in headers { + let key = name.as_str(); + if key.starts_with(AMZ_META_PREFIX_LOWER) + && key.ends_with(suffix_lower) + && !is_exact_snowball_meta_key(key, exact_keys) + && let Ok(parsed) = value.to_str() + { + return Some(parsed.trim().to_string()); + } + } + + None +} + +fn snowball_meta_value(headers: &HeaderMap, exact_keys: &[&str], suffix_lower: &str) -> Option { + for key in exact_keys { + if let Some(value) = trimmed_header_value(headers, key) { + return Some(value); + } + } + + snowball_meta_value_by_suffix(headers, suffix_lower, exact_keys) +} + +fn snowball_meta_flag(headers: &HeaderMap, exact_keys: &[&str], suffix_lower: &str) -> bool { + snowball_meta_value(headers, exact_keys, suffix_lower).is_some_and(|value| value.eq_ignore_ascii_case("true")) +} + +pub fn normalize_snowball_prefix(prefix: &str) -> Option { + let normalized = prefix.trim().trim_matches('/'); + if normalized.is_empty() { + return None; + } + + Some(normalized.to_string()) +} + +pub fn normalize_extract_entry_key(path: &str, prefix: Option<&str>, is_dir: bool) -> String { + let path = path.trim_matches('/'); + let mut key = match prefix { + Some(prefix) if !path.is_empty() => format!("{prefix}/{path}"), + Some(prefix) => prefix.to_string(), + None => path.to_string(), + }; + + if is_dir && !key.ends_with('/') { + key.push('/'); + } + + key +} + +pub fn resolve_put_object_extract_options(headers: &HeaderMap) -> PutObjectExtractOptions { + let prefix = snowball_meta_value(headers, SNOWBALL_PREFIX_HEADER_KEYS, SNOWBALL_PREFIX_SUFFIX_LOWER) + .and_then(|value| normalize_snowball_prefix(&value)); + let ignore_dirs = snowball_meta_flag(headers, SNOWBALL_IGNORE_DIRS_HEADER_KEYS, SNOWBALL_IGNORE_DIRS_SUFFIX_LOWER); + let ignore_errors = snowball_meta_flag(headers, SNOWBALL_IGNORE_ERRORS_HEADER_KEYS, SNOWBALL_IGNORE_ERRORS_SUFFIX_LOWER); + + PutObjectExtractOptions { + prefix, + ignore_dirs, + ignore_errors, + } +} + +pub fn map_extract_archive_error(err: impl std::fmt::Display) -> S3Error { + s3_error!(InvalidArgument, "Failed to process archive entry: {}", err) +} + +pub async fn apply_extract_entry_pax_extensions( + entry: &mut tokio_tar::Entry>, + metadata: &mut HashMap, + opts: &mut ObjectOptions, +) -> s3s::S3Result<()> +where + R: AsyncRead + Send + Unpin + 'static, +{ + let Some(extensions) = entry.pax_extensions().await.map_err(map_extract_archive_error)? else { + return Ok(()); + }; + + for ext in extensions { + let ext = ext.map_err(map_extract_archive_error)?; + let key = ext.key().map_err(map_extract_archive_error)?; + let value = ext.value().map_err(map_extract_archive_error)?; + + if let Some(meta_key) = key.strip_prefix("minio.metadata.") { + let meta_key = meta_key.strip_prefix("x-amz-meta-").unwrap_or(meta_key); + if !meta_key.is_empty() { + metadata.insert(meta_key.to_string(), value.to_string()); + } + continue; + } + + if key == "minio.versionId" && !value.is_empty() { + opts.version_id = Some(value.to_string()); + } + } + + Ok(()) +} + +pub fn is_sse_kms_requested(input: &PutObjectInput, headers: &HeaderMap) -> bool { + input + .server_side_encryption + .as_ref() + .is_some_and(|sse| sse.as_str().eq_ignore_ascii_case(ServerSideEncryption::AWS_KMS)) + || input.ssekms_key_id.is_some() + || headers + .get(AMZ_SERVER_SIDE_ENCRYPTION) + .and_then(|value| value.to_str().ok()) + .is_some_and(|value| value.trim().eq_ignore_ascii_case(ServerSideEncryption::AWS_KMS)) + || headers.contains_key(AMZ_SERVER_SIDE_ENCRYPTION_KMS_ID) +} + +pub fn is_post_object_sse_kms_requested(input: &PutObjectInput, headers: &HeaderMap) -> bool { + is_sse_kms_requested(input, headers) +} + +#[cfg(test)] +mod tests { + use super::*; + use bytes::Bytes; + use http::{HeaderMap, HeaderName, HeaderValue}; + use std::io::Cursor; + use tokio::io::AsyncReadExt; + + #[test] + fn should_use_zero_copy_accepts_large_unencrypted_binary_payload() { + let mut headers = HeaderMap::new(); + headers.insert("content-type", HeaderValue::from_static("application/octet-stream")); + assert!(should_use_zero_copy(2 * 1024 * 1024, &headers)); + } + + #[test] + fn should_use_zero_copy_rejects_small_payloads() { + assert!(!should_use_zero_copy(512 * 1024, &HeaderMap::new())); + } + + #[test] + fn resolve_put_body_size_uses_content_length_when_present() { + assert_eq!(resolve_put_body_size(Some(123), &HeaderMap::new()).unwrap(), 123); + } + + #[test] + fn resolve_put_body_size_uses_decoded_content_length_header() { + let mut headers = HeaderMap::new(); + headers.insert(AMZ_DECODED_CONTENT_LENGTH, "456".parse().unwrap()); + assert_eq!(resolve_put_body_size(None, &headers).unwrap(), 456); + } + + #[test] + fn should_use_zero_copy_rejects_encrypted_payloads() { + let mut headers = HeaderMap::new(); + headers.insert("x-amz-server-side-encryption", HeaderValue::from_static("AES256")); + assert!(!should_use_zero_copy(2 * 1024 * 1024, &headers)); + } + + #[test] + fn should_use_zero_copy_rejects_compressible_payloads() { + let mut headers = HeaderMap::new(); + headers.insert("content-type", HeaderValue::from_static("application/json")); + assert!(!should_use_zero_copy(2 * 1024 * 1024, &headers)); + } + + #[test] + fn should_use_zero_copy_rejects_boundary_at_1mb() { + let headers = HeaderMap::new(); + + assert!(!should_use_zero_copy(1024 * 1024, &headers)); + } + + #[test] + fn should_use_zero_copy_rejects_small_objects() { + let headers = HeaderMap::new(); + + assert!(!should_use_zero_copy(1024 * 1024 - 1, &headers)); + } + + #[test] + fn should_use_zero_copy_rejects_one_megabyte() { + let headers = HeaderMap::new(); + + assert!(!should_use_zero_copy(1024 * 1024, &headers)); + } + + #[test] + fn should_use_zero_copy_rejects_encrypted_requests() { + let mut headers = HeaderMap::new(); + headers.insert(AMZ_SERVER_SIDE_ENCRYPTION, HeaderValue::from_static("AES256")); + + assert!(!should_use_zero_copy(2 * 1024 * 1024, &headers)); + } + + #[test] + fn should_use_zero_copy_rejects_encrypted_requests_with_sse_customer_algorithm() { + let mut headers = HeaderMap::new(); + headers.insert(AMZ_SERVER_SIDE_ENCRYPTION_CUSTOMER_ALGORITHM, HeaderValue::from_static("AES256")); + + assert!(!should_use_zero_copy(2 * 1024 * 1024, &headers)); + } + + #[test] + fn should_use_zero_copy_rejects_encrypted_requests_with_kms_key_id() { + let mut headers = HeaderMap::new(); + headers.insert(AMZ_SERVER_SIDE_ENCRYPTION_KMS_ID, HeaderValue::from_static("test-kms-key-id")); + + assert!(!should_use_zero_copy(2 * 1024 * 1024, &headers)); + } + + #[test] + fn should_use_zero_copy_rejects_compressible_content_types() { + let mut headers = HeaderMap::new(); + headers.insert( + rustfs_utils::http::CONTENT_TYPE, + HeaderValue::from_static("application/json; charset=utf-8"), + ); + + assert!(!should_use_zero_copy(2 * 1024 * 1024, &headers)); + } + + #[test] + fn should_use_zero_copy_allows_large_unencrypted_binary_objects() { + let mut headers = HeaderMap::new(); + headers.insert(rustfs_utils::http::CONTENT_TYPE, HeaderValue::from_static("application/octet-stream")); + + assert!(should_use_zero_copy(2 * 1024 * 1024, &headers)); + } + + #[test] + fn plan_put_object_ingress_preserves_buffer_size_and_fast_path_decision() { + let mut headers = HeaderMap::new(); + headers.insert("content-type", HeaderValue::from_static("application/octet-stream")); + + let plan = plan_put_object_ingress(2 * 1024 * 1024, &headers, 256 * 1024); + + assert_eq!(plan.kind, PutObjectIngressKind::ReducedCopyCandidate); + assert_eq!(plan.compat.kind, PutObjectCompatIngressKind::BufferedStreamCompat); + assert_eq!(plan.compat.buffer_size, 256 * 1024); + assert!(plan.enable_zero_copy); + } + + #[test] + fn plan_put_object_body_disables_compression_for_small_payloads() { + let plan = plan_put_object_body(1024, &HeaderMap::new(), "small.bin", 64 * 1024); + + assert_eq!(plan.ingress.kind, PutObjectIngressKind::LegacyCompat); + assert_eq!(plan.ingress.compat.buffer_size, 64 * 1024); + assert!(!plan.ingress.enable_zero_copy); + assert_eq!(plan.kind, PutObjectBodyKind::Plain(PutObjectPlainBodyKind::LegacyCompat)); + assert!(!plan.should_compress()); + } + + #[test] + fn plan_put_object_body_marks_large_plain_payload_as_reduced_copy_candidate() { + let mut headers = HeaderMap::new(); + headers.insert("content-type", HeaderValue::from_static("application/octet-stream")); + + let plan = plan_put_object_body(2 * 1024 * 1024, &headers, "large.bin", 256 * 1024); + + assert_eq!(plan.kind, PutObjectBodyKind::Plain(PutObjectPlainBodyKind::ReducedCopyCandidate)); + assert_eq!(plan.plain_body_kind(), Some(PutObjectPlainBodyKind::ReducedCopyCandidate)); + assert!(!plan.should_compress()); + } + + #[test] + fn plan_put_object_body_with_transforms_allows_encrypted_large_binary_payloads() { + let mut headers = HeaderMap::new(); + headers.insert("content-type", HeaderValue::from_static("application/octet-stream")); + headers.insert(AMZ_SERVER_SIDE_ENCRYPTION, HeaderValue::from_static("AES256")); + + let plan = plan_put_object_body_with_transforms(2 * 1024 * 1024, &headers, "large.bin", 256 * 1024, true); + + assert_eq!(plan.kind, PutObjectBodyKind::Plain(PutObjectPlainBodyKind::ReducedCopyCandidate)); + assert_eq!(plan.plain_body_kind(), Some(PutObjectPlainBodyKind::ReducedCopyCandidate)); + assert!(plan.ingress.enable_zero_copy); + } + + #[test] + fn build_put_object_ingress_source_preserves_reduced_copy_candidate_for_compressed_body() { + let stream = futures_util::stream::iter([Ok::(Bytes::from_static(b"compressed"))]); + let plan = PutObjectBodyPlan { + ingress: PutObjectIngressPlan { + kind: PutObjectIngressKind::ReducedCopyCandidate, + compat: PutObjectCompatIngressPlan { + kind: PutObjectCompatIngressKind::BufferedStreamCompat, + buffer_size: 256 * 1024, + }, + enable_zero_copy: true, + }, + kind: PutObjectBodyKind::Compressed, + }; + + let source = build_put_object_ingress_source(stream, plan); + assert!(matches!(source, PutObjectIngressSource::ReducedCopyCandidate(_))); + } + + #[test] + fn put_object_body_plan_reports_compressed_kind() { + let plan = PutObjectBodyPlan { + ingress: PutObjectIngressPlan { + kind: PutObjectIngressKind::LegacyCompat, + compat: PutObjectCompatIngressPlan { + kind: PutObjectCompatIngressKind::BufferedStreamCompat, + buffer_size: 256 * 1024, + }, + enable_zero_copy: false, + }, + kind: PutObjectBodyKind::Compressed, + }; + + assert_eq!(plan.plain_body_kind(), None); + assert!(plan.should_compress()); + } + + #[test] + fn resolve_put_effective_copy_mode_marks_transformed_paths() { + assert_eq!(resolve_put_effective_copy_mode(false, false), rustfs_io_metrics::CopyMode::SingleCopy); + assert_eq!(resolve_put_effective_copy_mode(true, false), rustfs_io_metrics::CopyMode::Transformed); + assert_eq!(resolve_put_effective_copy_mode(false, true), rustfs_io_metrics::CopyMode::Transformed); + } + + #[test] + fn put_object_transform_stage_tracks_transform_shape() { + let mut stage = PutObjectTransformStage::default(); + assert_eq!(stage.effective_copy_mode(), rustfs_io_metrics::CopyMode::SingleCopy); + assert_eq!(stage.metric_kind(), None); + + stage.mark_compression(); + assert!(stage.compression_applied()); + assert_eq!(stage.effective_copy_mode(), rustfs_io_metrics::CopyMode::Transformed); + assert_eq!(stage.metric_kind(), Some("compression")); + + stage.mark_encryption(); + assert!(stage.encryption_applied()); + assert_eq!(stage.metric_kind(), Some("compression_encryption")); + } + + #[test] + fn resolve_put_transform_metric_kind_reports_transform_shape() { + assert_eq!(resolve_put_transform_metric_kind(false, false), None); + assert_eq!(resolve_put_transform_metric_kind(true, false), Some("compression")); + assert_eq!(resolve_put_transform_metric_kind(false, true), Some("encryption")); + assert_eq!(resolve_put_transform_metric_kind(true, true), Some("compression_encryption")); + } + + #[test] + fn resolve_put_transformed_fallback_reason_isolated_from_plain_path() { + assert_eq!( + resolve_put_transformed_fallback_reason(PutObjectIngressKind::LegacyCompat, true, false), + Some(rustfs_io_metrics::FallbackReason::TransformCompressionLegacy) + ); + assert_eq!( + resolve_put_transformed_fallback_reason(PutObjectIngressKind::LegacyCompat, false, true), + Some(rustfs_io_metrics::FallbackReason::TransformEncryptionLegacy) + ); + assert_eq!( + resolve_put_transformed_fallback_reason(PutObjectIngressKind::LegacyCompat, true, true), + Some(rustfs_io_metrics::FallbackReason::TransformCompressionEncryptionLegacy) + ); + assert_eq!( + resolve_put_transformed_fallback_reason(PutObjectIngressKind::ReducedCopyCandidate, true, true), + None + ); + } + + #[test] + fn is_put_object_extract_requested_accepts_meta_header() { + let mut headers = HeaderMap::new(); + headers.insert(AMZ_SNOWBALL_EXTRACT, HeaderValue::from_static("true")); + assert!(is_put_object_extract_requested(&headers)); + } + + #[test] + fn is_put_object_extract_requested_accepts_compat_header_case_insensitive() { + let mut headers = HeaderMap::new(); + headers.insert(AMZ_SNOWBALL_EXTRACT_COMPAT, HeaderValue::from_static(" TRUE ")); + assert!(is_put_object_extract_requested(&headers)); + } + + #[test] + fn is_put_object_extract_requested_rejects_missing_or_false_value() { + let mut headers = HeaderMap::new(); + assert!(!is_put_object_extract_requested(&headers)); + headers.insert(AMZ_SNOWBALL_EXTRACT, HeaderValue::from_static("false")); + assert!(!is_put_object_extract_requested(&headers)); + } + + #[test] + fn normalize_snowball_prefix_trims_slashes_and_whitespace() { + assert_eq!(normalize_snowball_prefix(" /batch/incoming/ "), Some("batch/incoming".to_string())); + assert_eq!(normalize_snowball_prefix("///"), None); + } + + #[test] + fn normalize_extract_entry_key_applies_prefix_and_directory_suffix() { + assert_eq!( + normalize_extract_entry_key("nested/path.txt", Some("imports"), false), + "imports/nested/path.txt" + ); + assert_eq!(normalize_extract_entry_key("nested/dir/", Some("imports"), true), "imports/nested/dir/"); + assert_eq!(normalize_extract_entry_key("top-level", None, false), "top-level"); + } + + #[test] + fn resolve_put_object_extract_options_defaults_when_headers_missing() { + let headers = HeaderMap::new(); + let options = resolve_put_object_extract_options(&headers); + assert_eq!( + options, + PutObjectExtractOptions { + prefix: None, + ignore_dirs: false, + ignore_errors: false + } + ); + } + + #[test] + fn resolve_put_object_extract_options_accepts_internal_headers() { + let mut headers = HeaderMap::new(); + headers.insert(AMZ_SNOWBALL_PREFIX_INTERNAL, HeaderValue::from_static("/internal/prefix/")); + headers.insert(AMZ_SNOWBALL_IGNORE_DIRS_INTERNAL, HeaderValue::from_static("true")); + headers.insert(AMZ_SNOWBALL_IGNORE_ERRORS_INTERNAL, HeaderValue::from_static("TRUE")); + + let options = resolve_put_object_extract_options(&headers); + assert_eq!(options.prefix.as_deref(), Some("internal/prefix")); + assert!(options.ignore_dirs); + assert!(options.ignore_errors); + } + + #[test] + fn resolve_put_object_extract_options_accepts_standard_headers() { + let mut headers = HeaderMap::new(); + headers.insert(AMZ_SNOWBALL_PREFIX, HeaderValue::from_static(" /standard/prefix/ ")); + headers.insert(AMZ_SNOWBALL_IGNORE_DIRS, HeaderValue::from_static(" true ")); + headers.insert(AMZ_SNOWBALL_IGNORE_ERRORS, HeaderValue::from_static("TRUE")); + + let options = resolve_put_object_extract_options(&headers); + assert_eq!(options.prefix.as_deref(), Some("standard/prefix")); + assert!(options.ignore_dirs); + assert!(options.ignore_errors); + } + + #[test] + fn resolve_put_object_extract_options_accepts_suffix_compatible_headers() { + let mut headers = HeaderMap::new(); + headers.insert( + HeaderName::from_static("x-amz-meta-acme-snowball-prefix"), + HeaderValue::from_static(" /partner/import "), + ); + headers.insert( + HeaderName::from_static("x-amz-meta-acme-snowball-ignore-dirs"), + HeaderValue::from_static(" true "), + ); + headers.insert( + HeaderName::from_static("x-amz-meta-acme-snowball-ignore-errors"), + HeaderValue::from_static("TRUE"), + ); + + let options = resolve_put_object_extract_options(&headers); + assert_eq!(options.prefix.as_deref(), Some("partner/import")); + assert!(options.ignore_dirs); + assert!(options.ignore_errors); + } + + #[test] + fn resolve_put_object_extract_options_prefers_exact_headers_over_suffix_fallback() { + let mut headers = HeaderMap::new(); + headers.insert("x-amz-meta-acme-snowball-prefix", HeaderValue::from_static("/fallback/prefix/")); + headers.insert(AMZ_RUSTFS_SNOWBALL_PREFIX, HeaderValue::from_static("/internal/prefix/")); + headers.insert(AMZ_SNOWBALL_PREFIX, HeaderValue::from_static("/standard/prefix/")); + headers.insert(AMZ_MINIO_SNOWBALL_PREFIX, HeaderValue::from_static("/minio/prefix/")); + + let options = resolve_put_object_extract_options(&headers); + assert_eq!(options.prefix.as_deref(), Some("minio/prefix")); + } + + #[test] + fn resolve_put_object_extract_options_exact_flags_override_suffix_fallback() { + let mut headers = HeaderMap::new(); + headers.insert(AMZ_SNOWBALL_IGNORE_DIRS, HeaderValue::from_static("false")); + headers.insert("x-amz-meta-acme-snowball-ignore-dirs", HeaderValue::from_static("true")); + headers.insert(AMZ_RUSTFS_SNOWBALL_IGNORE_ERRORS, HeaderValue::from_static("false")); + headers.insert("x-amz-meta-acme-snowball-ignore-errors", HeaderValue::from_static("true")); + + let options = resolve_put_object_extract_options(&headers); + assert!(!options.ignore_dirs); + assert!(!options.ignore_errors); + } + + #[tokio::test] + async fn build_put_object_compat_ingress_reads_buffered_stream() { + let stream = futures_util::stream::iter([ + Ok::(Bytes::from_static(b"abc")), + Ok::(Bytes::from_static(b"def")), + ]); + let plan = PutObjectCompatIngressPlan { + kind: PutObjectCompatIngressKind::BufferedStreamCompat, + buffer_size: 8, + }; + + let mut reader = build_put_object_compat_ingress(stream, plan); + let mut buf = Vec::new(); + reader.read_to_end(&mut buf).await.unwrap(); + + assert_eq!(buf, b"abcdef"); + } + + #[tokio::test] + async fn build_put_object_compat_ingress_maps_stream_errors() { + let stream = futures_util::stream::iter([Err::("boom")]); + let plan = PutObjectCompatIngressPlan { + kind: PutObjectCompatIngressKind::BufferedStreamCompat, + buffer_size: 8, + }; + + let mut reader = build_put_object_compat_ingress(stream, plan); + let err = reader.read_to_end(&mut Vec::new()).await.unwrap_err(); + + assert_eq!(err.kind(), std::io::ErrorKind::Other); + assert!(err.to_string().contains("boom")); + } + + #[tokio::test] + async fn build_put_object_compat_reader_wraps_ingress_as_reader() { + let stream = futures_util::stream::iter([Ok::(Bytes::from_static(b"reader"))]); + let plan = PutObjectCompatIngressPlan { + kind: PutObjectCompatIngressKind::BufferedStreamCompat, + buffer_size: 16, + }; + + let mut reader = build_put_object_compat_reader(stream, plan); + let mut buf = Vec::new(); + reader.read_to_end(&mut buf).await.unwrap(); + + assert_eq!(buf, b"reader"); + } + + #[tokio::test] + async fn build_put_object_ingress_source_keeps_plain_candidate_as_reduced_copy_source() { + let stream = futures_util::stream::iter([Ok::(Bytes::from_static(b"candidate"))]); + let mut headers = HeaderMap::new(); + headers.insert("content-type", HeaderValue::from_static("application/octet-stream")); + let plan = plan_put_object_body(2 * 1024 * 1024, &headers, "large.bin", 16); + + let source = build_put_object_ingress_source(stream, plan); + let candidate = match source { + PutObjectIngressSource::ReducedCopyCandidate(candidate) => candidate, + PutObjectIngressSource::LegacyCompat(_) => panic!("expected reduced-copy candidate"), + }; + + assert_eq!(candidate.compat_plan().kind, PutObjectCompatIngressKind::BufferedStreamCompat); + assert_eq!(candidate.compat_plan().buffer_size, 16); + + let mut reader = candidate.into_compat_reader(); + let mut buf = Vec::new(); + reader.read_to_end(&mut buf).await.unwrap(); + + assert_eq!(buf, b"candidate"); + } + + #[tokio::test] + async fn build_put_object_ingress_source_routes_compressed_body_to_legacy_compat_reader() { + let stream = futures_util::stream::iter([Ok::(Bytes::from_static(b"compressed"))]); + let mut headers = HeaderMap::new(); + headers.insert("content-type", HeaderValue::from_static("application/json")); + let plan = plan_put_object_body(2 * 1024 * 1024, &headers, "large.json", 32); + + let source = build_put_object_ingress_source(stream, plan); + let mut reader = match source { + PutObjectIngressSource::LegacyCompat(reader) => reader, + PutObjectIngressSource::ReducedCopyCandidate(_) => panic!("expected legacy compat reader"), + }; + + let mut buf = Vec::new(); + reader.read_to_end(&mut buf).await.unwrap(); + + assert_eq!(buf, b"compressed"); + } + + #[tokio::test] + async fn build_put_object_reduced_copy_reader_reads_across_multiple_chunks() { + let stream = futures_util::stream::iter([ + Ok::(Bytes::from_static(b"ab")), + Ok::(Bytes::from_static(b"cd")), + Ok::(Bytes::from_static(b"ef")), + ]); + + let mut reader = build_put_object_reduced_copy_reader(PutObjectReducedCopyIngress::from_stream( + stream, + PutObjectCompatIngressPlan { + kind: PutObjectCompatIngressKind::BufferedStreamCompat, + buffer_size: 16, + }, + )); + let mut buf = Vec::new(); + reader.read_to_end(&mut buf).await.unwrap(); + + assert_eq!(buf, b"abcdef"); + } + + #[tokio::test] + async fn build_put_object_reduced_copy_reader_defers_stream_error_until_next_read() { + let stream = futures_util::stream::iter([Ok::(Bytes::from_static(b"ab")), Err::("boom")]); + + let mut reader = build_put_object_reduced_copy_reader(PutObjectReducedCopyIngress::from_stream( + stream, + PutObjectCompatIngressPlan { + kind: PutObjectCompatIngressKind::BufferedStreamCompat, + buffer_size: 16, + }, + )); + + let mut buf = [0_u8; 8]; + let n = reader.read(&mut buf).await.unwrap(); + assert_eq!(n, 2); + assert_eq!(&buf[..n], b"ab"); + + let err = reader.read(&mut buf).await.unwrap_err(); + assert_eq!(err.kind(), std::io::ErrorKind::Other); + assert!(err.to_string().contains("boom")); + } + + #[tokio::test] + async fn build_put_object_reduced_copy_reader_supports_direct_block_reads() { + let stream = futures_util::stream::iter([ + Ok::(Bytes::from_static(b"ab")), + Ok::(Bytes::from_static(b"cd")), + Ok::(Bytes::from_static(b"ef")), + ]); + + let mut reader = build_put_object_reduced_copy_reader(PutObjectReducedCopyIngress::from_stream( + stream, + PutObjectCompatIngressPlan { + kind: PutObjectCompatIngressKind::BufferedStreamCompat, + buffer_size: 16, + }, + )); + + let mut first = [0_u8; 4]; + let mut second = [0_u8; 4]; + assert_eq!(reader.read_block(&mut first).await.unwrap(), 4); + assert_eq!(&first, b"abcd"); + assert_eq!(reader.read_block(&mut second).await.unwrap(), 2); + assert_eq!(&second[..2], b"ef"); + assert_eq!(reader.read_block(&mut second).await.unwrap(), 0); + } + + #[tokio::test] + async fn build_put_object_plain_hash_stage_reports_reduced_copy_candidate_boundary() { + let stream = futures_util::stream::iter([Ok::(Bytes::from_static(b"plain"))]); + let mut headers = HeaderMap::new(); + headers.insert("content-type", HeaderValue::from_static("application/octet-stream")); + let plan = plan_put_object_body(2 * 1024 * 1024, &headers, "large.bin", 32); + + let stage = build_put_object_plain_hash_stage( + build_put_object_ingress_source(stream, plan), + PutObjectLegacyHashValues::default(), + PutObjectLegacyHashStagePlan { + size: 5, + actual_size: 5, + apply_s3_checksum: false, + ignore_s3_checksum_value: false, + }, + &headers, + None, + ) + .unwrap(); + + assert_eq!(stage.ingress_kind, PutObjectIngressKind::ReducedCopyCandidate); + + let mut reader = stage.reader; + let mut buf = [0_u8; 8]; + let n = reader.read_block(&mut buf).await.unwrap(); + assert_eq!(n, 5); + assert_eq!(&buf[..n], b"plain"); + assert_eq!(reader.read_block(&mut buf).await.unwrap(), 0); + } + + #[tokio::test] + async fn build_put_object_plain_hash_stage_preserves_legacy_compat_boundary() { + let stream = futures_util::stream::iter([Ok::(Bytes::from_static(b"legacy"))]); + let plan = plan_put_object_body(1024, &HeaderMap::new(), "small.bin", 32); + + let stage = build_put_object_plain_hash_stage( + build_put_object_ingress_source(stream, plan), + PutObjectLegacyHashValues::default(), + PutObjectLegacyHashStagePlan { + size: 6, + actual_size: 6, + apply_s3_checksum: false, + ignore_s3_checksum_value: false, + }, + &HeaderMap::new(), + None, + ) + .unwrap(); + + assert_eq!(stage.ingress_kind, PutObjectIngressKind::LegacyCompat); + + let mut reader = stage.reader; + let mut buf = Vec::new(); + reader.read_to_end(&mut buf).await.unwrap(); + + assert_eq!(buf, b"legacy"); + } + + #[test] + fn put_object_legacy_hash_values_clear_for_transformed_body_resets_hashes() { + let mut hash_values = PutObjectLegacyHashValues { + md5hex: Some("md5".to_string()), + sha256hex: Some("sha256".to_string()), + }; + + hash_values.clear_for_transformed_body(); + + assert_eq!(hash_values, PutObjectLegacyHashValues::default()); + } + + #[test] + fn build_put_object_legacy_hash_stage_preserves_legacy_compat_boundary() { + let reader: Box = Box::new(WarpReader::new(Cursor::new(Vec::from(&b"abc"[..])))); + let stage = build_put_object_legacy_hash_stage( + reader, + PutObjectLegacyHashValues::default(), + PutObjectLegacyHashStagePlan { + size: 3, + actual_size: 3, + apply_s3_checksum: false, + ignore_s3_checksum_value: false, + }, + &HeaderMap::new(), + None, + ) + .unwrap(); + + assert_eq!(stage.reader.size(), 3); + assert_eq!(stage.reader.actual_size(), 3); + assert!(stage.want_checksum.is_none()); + } +} diff --git a/crates/protocols/src/swift/object.rs b/crates/protocols/src/swift/object.rs index 1c59da07f..763a8f202 100644 --- a/crates/protocols/src/swift/object.rs +++ b/crates/protocols/src/swift/object.rs @@ -55,7 +55,7 @@ use super::{SwiftError, SwiftResult}; 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_ecstore::store_api::{BucketOperations, BucketOptions, ChunkNativePutData, ObjectIO, ObjectOperations, ObjectOptions}; use rustfs_rio::HashReader; use std::collections::HashMap; use tracing::debug; @@ -381,8 +381,8 @@ where 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); + // 15. Hand the hash reader to the chunk-native PUT data wrapper + let mut put_reader = ChunkNativePutData::new(hash_reader); // 16. Upload object to storage let obj_info = store @@ -464,8 +464,8 @@ where 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); + // Hand the hash reader to the chunk-native PUT data wrapper + let mut put_reader = ChunkNativePutData::new(hash_reader); // Upload object to storage let obj_info = store diff --git a/crates/rio/src/checksum.rs b/crates/rio/src/checksum.rs index 86867cc15..b09916bf8 100644 --- a/crates/rio/src/checksum.rs +++ b/crates/rio/src/checksum.rs @@ -29,11 +29,21 @@ pub const RUSTFS_MULTIPART_CHECKSUM: &str = "x-rustfs-multipart-checksum"; /// RustFS multipart checksum type metadata key pub const RUSTFS_MULTIPART_CHECKSUM_TYPE: &str = "x-rustfs-multipart-checksum-type"; +const AMZ_CHECKSUM_ALGORITHM: &str = "x-amz-checksum-algorithm"; +const AMZ_SDK_CHECKSUM_ALGORITHM: &str = "x-amz-sdk-checksum-algorithm"; + /// Checksum type enumeration with flags #[derive(Debug, Clone, Copy, PartialEq, Eq, Default)] pub struct ChecksumType(pub u32); impl ChecksumType { + fn algorithm_from_headers(headers: &HeaderMap) -> Option<&str> { + headers + .get(AMZ_CHECKSUM_ALGORITHM) + .and_then(|v| v.to_str().ok()) + .or_else(|| headers.get(AMZ_SDK_CHECKSUM_ALGORITHM).and_then(|v| v.to_str().ok())) + } + /// Checksum will be sent in trailing header pub const TRAILING: ChecksumType = ChecksumType(1 << 0); @@ -156,10 +166,7 @@ impl ChecksumType { pub fn from_header(headers: &HeaderMap) -> Self { Self::from_string_with_obj_type( - headers - .get("x-amz-checksum-algorithm") - .and_then(|v| v.to_str().ok()) - .unwrap_or(""), + Self::algorithm_from_headers(headers).unwrap_or(""), headers.get("x-amz-checksum-type").and_then(|v| v.to_str().ok()).unwrap_or(""), ) } @@ -573,7 +580,7 @@ pub fn get_content_checksum(headers: &HeaderMap) -> Result, std fn get_content_checksum_direct(headers: &HeaderMap) -> (ChecksumType, String) { let mut checksum_type = ChecksumType::NONE; - if let Some(alg) = headers.get("x-amz-checksum-algorithm").and_then(|v| v.to_str().ok()) { + if let Some(alg) = ChecksumType::algorithm_from_headers(headers) { checksum_type = ChecksumType::from_string_with_obj_type( alg, headers.get("x-amz-checksum-type").and_then(|s| s.to_str().ok()).unwrap_or(""), @@ -1131,7 +1138,8 @@ fn crc64_combine(poly: u64, crc1: u64, crc2: u64, len2: i64) -> u64 { #[cfg(test)] mod tests { - use super::{Checksum, ChecksumType}; + use super::{AMZ_SDK_CHECKSUM_ALGORITHM, Checksum, ChecksumType, get_content_checksum_direct}; + use http::{HeaderMap, HeaderValue}; #[test] fn crc64_nvme_add_part_matches_full_object_checksum() { @@ -1186,4 +1194,24 @@ mod tests { assert_eq!(combined.encoded, expected.encoded); assert_eq!(combined.raw, expected.raw); } + + #[test] + fn checksum_type_from_header_supports_sdk_checksum_algorithm_header() { + let mut headers = HeaderMap::new(); + headers.insert(AMZ_SDK_CHECKSUM_ALGORITHM, HeaderValue::from_static("CRC32")); + + assert_eq!(ChecksumType::from_header(&headers), ChecksumType::CRC32); + } + + #[test] + fn get_content_checksum_direct_supports_sdk_checksum_algorithm_header() { + let mut headers = HeaderMap::new(); + headers.insert(AMZ_SDK_CHECKSUM_ALGORITHM, HeaderValue::from_static("CRC32")); + headers.insert("x-amz-checksum-crc32", HeaderValue::from_static("nct/nQ==")); + + let (checksum_type, checksum_value) = get_content_checksum_direct(&headers); + + assert_eq!(checksum_type, ChecksumType::CRC32); + assert_eq!(checksum_value, "nct/nQ=="); + } } diff --git a/crates/rio/src/compress_reader.rs b/crates/rio/src/compress_reader.rs index 418373a89..e0d961234 100644 --- a/crates/rio/src/compress_reader.rs +++ b/crates/rio/src/compress_reader.rs @@ -13,6 +13,7 @@ // limitations under the License. use crate::compress_index::{Index, TryGetIndex}; +use crate::{BlockReadable, BoxReadBlockFuture, Reader}; use pin_project_lite::pin_project; use rustfs_utils::compress::{CompressionAlgorithm, compress_block, decompress_block}; use rustfs_utils::{put_uvarint, uvarint}; @@ -85,6 +86,31 @@ where read_buffer: vec![0u8; block_size], } } + + fn copy_buffered(&mut self, buf: &mut [u8]) -> usize { + if self.pos >= self.buffer.len() || buf.is_empty() { + return 0; + } + + let to_copy = min(buf.len(), self.buffer.len() - self.pos); + buf[..to_copy].copy_from_slice(&self.buffer[self.pos..self.pos + to_copy]); + self.pos += to_copy; + if self.pos == self.buffer.len() { + self.buffer.clear(); + self.pos = 0; + } + to_copy + } + + fn queue_compressed_block(&mut self, uncompressed_data: &[u8]) -> io::Result<()> { + let out = build_compressed_block(uncompressed_data, self.compression_algorithm); + self.written += out.len(); + self.uncomp_written += uncompressed_data.len(); + self.index.add(self.written as i64, self.uncomp_written as i64)?; + self.buffer = out; + self.pos = 0; + Ok(()) + } } impl TryGetIndex for CompressReader { @@ -170,6 +196,50 @@ where delegate_reader_capabilities_generic_no_index!(CompressReader, inner); +impl BlockReadable for CompressReader +where + R: Reader, +{ + fn read_block<'a>(&'a mut self, buf: &'a mut [u8]) -> BoxReadBlockFuture<'a> { + Box::pin(async move { + if buf.is_empty() { + return Ok(0); + } + + let mut written = self.copy_buffered(buf); + while written < buf.len() { + if self.done { + break; + } + + self.temp_buffer.resize(self.block_size, 0); + let n = { + let inner = &mut self.inner; + let temp = &mut self.temp_buffer[..self.block_size]; + match inner.read_block(temp).await { + Ok(n) => n, + Err(err) if err.kind() == io::ErrorKind::UnexpectedEof => 0, + Err(err) => return Err(err), + } + }; + + if n == 0 { + self.done = true; + self.temp_buffer.clear(); + break; + } + + let block = self.temp_buffer[..n].to_vec(); + self.temp_buffer.clear(); + self.queue_compressed_block(&block)?; + written += self.copy_buffered(&mut buf[written..]); + } + + Ok(written) + }) + } +} + pin_project! { /// A reader wrapper that decompresses data on the fly using DEFLATE algorithm. /// Header format: @@ -390,6 +460,7 @@ fn build_compressed_block(uncompressed_data: &[u8], compression_algorithm: Compr #[cfg(test)] mod tests { use super::*; + use crate::{BlockReadable, WarpReader}; use rand::RngExt; use std::io::Cursor; use tokio::io::{AsyncReadExt, BufReader}; @@ -479,4 +550,27 @@ mod tests { assert_eq!(&decompressed, &data); } + + #[tokio::test] + async fn test_compress_reader_read_block_round_trips() { + let data = b"hello world, hello world, hello world!"; + let reader = Cursor::new(data.to_vec()); + let mut compress_reader = CompressReader::new(WarpReader::new(reader), CompressionAlgorithm::Gzip); + let mut compressed = Vec::new(); + let mut buf = [0u8; 19]; + + loop { + let n = compress_reader.read_block(&mut buf).await.unwrap(); + if n == 0 { + break; + } + compressed.extend_from_slice(&buf[..n]); + } + + let mut decompress_reader = DecompressReader::new(Cursor::new(compressed), CompressionAlgorithm::Gzip); + let mut decompressed = Vec::new(); + decompress_reader.read_to_end(&mut decompressed).await.unwrap(); + + assert_eq!(&decompressed, data); + } } diff --git a/crates/rio/src/encrypt_reader.rs b/crates/rio/src/encrypt_reader.rs index 4b8e275cf..d815a75d3 100644 --- a/crates/rio/src/encrypt_reader.rs +++ b/crates/rio/src/encrypt_reader.rs @@ -13,6 +13,7 @@ // limitations under the License. use crate::compress_index::{Index, TryGetIndex}; +use crate::{BlockReadable, BoxReadBlockFuture, Reader}; use aes_gcm::aead::Aead; use aes_gcm::{Aes256Gcm, KeyInit, Nonce}; use pin_project_lite::pin_project; @@ -59,6 +60,69 @@ where } } +fn encrypt_segment_bytes(cipher: &Aes256Gcm, nonce_bytes: &[u8; 12], plaintext: &[u8]) -> std::io::Result> { + let nonce = Nonce::try_from(nonce_bytes.as_slice()).map_err(|_| Error::other("invalid nonce length"))?; + 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 + .encrypt(&nonce, plaintext) + .map_err(|e| Error::other(format!("encrypt error: {e}")))?; + let int_len = put_uvarint_len(plaintext_len as u64); + let clen = int_len + ciphertext.len() + 4; + let mut header = [0u8; 8]; + header[0] = 0x00; + 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; + debug!( + "encrypt block header typ=0 len={} header={:?} plaintext_len={} ciphertext_len={}", + clen, + header, + plaintext_len, + ciphertext.len() + ); + 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); + out.extend_from_slice(&ciphertext); + Ok(out) +} + +impl EncryptReader +where + R: Reader, +{ + fn copy_buffered(&mut self, buf: &mut [u8]) -> usize { + if self.buffer_pos >= self.buffer.len() || buf.is_empty() { + return 0; + } + + let to_copy = buf.len().min(self.buffer.len() - self.buffer_pos); + buf[..to_copy].copy_from_slice(&self.buffer[self.buffer_pos..self.buffer_pos + to_copy]); + self.buffer_pos += to_copy; + if self.buffer_pos == self.buffer.len() { + self.buffer.clear(); + self.buffer_pos = 0; + } + to_copy + } + + fn encrypt_segment(&self, plaintext: &[u8]) -> std::io::Result> { + let nonce = derive_block_nonce(&self.base_nonce, self.block_index); + encrypt_segment_bytes(&self.cipher, &nonce, plaintext) + } +} + impl AsyncRead for EncryptReader where R: AsyncRead + Unpin + Send + Sync, @@ -97,49 +161,8 @@ where *this.buffer_pos += to_copy; Poll::Ready(Ok(())) } else { - // Encrypt the chunk 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 = this - .cipher - .encrypt(&nonce, plaintext) - .map_err(|e| Error::other(format!("encrypt error: {e}")))?; - let int_len = put_uvarint_len(plaintext_len as u64); - let clen = int_len + ciphertext.len() + 4; - // Header: 8 bytes - // 0: type (0 = encrypted, 0xFF = end) - // 1-3: length (little endian u24, ciphertext length) - // 4-7: CRC32 of ciphertext (little endian u32) - let mut header = [0u8; 8]; - header[0] = 0x00; // 0 = encrypted - 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; - debug!( - "encrypt block header typ=0 len={} header={:?} plaintext_len={} ciphertext_len={}", - clen, - header, - plaintext_len, - ciphertext.len() - ); - let mut out = Vec::with_capacity(8 + int_len + ciphertext.len()); - out.extend_from_slice(&header); - 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 = encrypt_segment_bytes(this.cipher, &block_nonce, &this.read_buffer[..n])?; *this.buffer_pos = 0; *this.block_index += 1; let to_copy = std::cmp::min(buf.remaining(), this.buffer.len()); @@ -164,6 +187,50 @@ where } } +impl BlockReadable for EncryptReader +where + R: Reader, +{ + fn read_block<'a>(&'a mut self, buf: &'a mut [u8]) -> BoxReadBlockFuture<'a> { + Box::pin(async move { + if buf.is_empty() { + return Ok(0); + } + + let mut written = self.copy_buffered(buf); + while written < buf.len() { + if self.finished { + break; + } + + let mut plaintext = vec![0u8; 8 * 1024]; + let n = match self.inner.read_block(&mut plaintext).await { + Ok(n) => n, + Err(err) if err.kind() == std::io::ErrorKind::UnexpectedEof => 0, + Err(err) => return Err(err), + }; + if n == 0 { + self.buffer = [0xFF, 0, 0, 0, 0, 0, 0, 0].to_vec(); + self.buffer_pos = 0; + self.finished = true; + } else { + self.buffer = self.encrypt_segment(&plaintext[..n])?; + self.buffer_pos = 0; + } + + let copied = self.copy_buffered(&mut buf[written..]); + written += copied; + + if copied == 0 && self.finished { + break; + } + } + + Ok(written) + }) + } +} + 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. @@ -486,7 +553,7 @@ mod tests { use std::pin::Pin; use std::task::{Context, Poll}; - use crate::HardLimitReader; + use crate::{BlockReadable, HardLimitReader, WarpReader}; use super::*; use futures::StreamExt; @@ -674,6 +741,35 @@ mod tests { assert_eq!(&decrypted, &data); } + #[tokio::test] + async fn test_encrypt_reader_read_block_round_trips() { + let data = b"hello sse encrypt via blocks"; + 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.to_vec()); + let mut encrypt_reader = EncryptReader::new(WarpReader::new(reader), key, nonce); + let mut encrypted = Vec::new(); + let mut buf = [0u8; 17]; + + loop { + let n = encrypt_reader.read_block(&mut buf).await.unwrap(); + if n == 0 { + break; + } + encrypted.extend_from_slice(&buf[..n]); + } + + let reader = Cursor::new(encrypted); + let mut decrypt_reader = DecryptReader::new(WarpReader::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_large_with_small_chunks() { let size = 1024 * 1024; diff --git a/crates/rio/src/etag_reader.rs b/crates/rio/src/etag_reader.rs index ba1638069..d6f410ae9 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}; +use crate::{BlockReadable, BoxReadBlockFuture, EtagResolvable, HashReaderDetector, HashReaderMut, Reader}; use md5::{Digest, Md5}; use pin_project_lite::pin_project; use std::pin::Pin; @@ -135,9 +135,41 @@ where } } +impl BlockReadable for EtagReader +where + R: Reader, +{ + fn read_block<'a>(&'a mut self, buf: &'a mut [u8]) -> BoxReadBlockFuture<'a> { + Box::pin(async move { + let n = match self.inner.read_block(buf).await { + Ok(n) => n, + Err(err) if err.kind() == std::io::ErrorKind::UnexpectedEof => 0, + Err(err) => return Err(err), + }; + if n > 0 { + self.md5.update(&buf[..n]); + return Ok(n); + } + + self.finished = true; + if let Some(checksum) = &self.checksum { + let etag = self.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 Err(std::io::Error::new(std::io::ErrorKind::InvalidData, "Checksum mismatch")); + } + } + + Ok(0) + }) + } +} + #[cfg(test)] mod tests { use super::*; + use crate::{BlockReadable, WarpReader}; use rand::RngExt; use std::io::Cursor; use tokio::io::{AsyncReadExt, BufReader}; @@ -180,6 +212,24 @@ mod tests { assert_eq!(etag, Some(expected)); } + #[tokio::test] + async fn test_etag_reader_read_block_updates_checksum() { + let data = b"hello world"; + let mut hasher = Md5::new(); + hasher.update(data); + let expected = faster_hex::hex_string(hasher.finalize().as_slice()).to_string(); + 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 = [0_u8; 32]; + let n = etag_reader.read_block(&mut buf).await.unwrap(); + assert_eq!(n, data.len()); + assert_eq!(&buf[..n], data); + assert_eq!(etag_reader.read_block(&mut buf).await.unwrap(), 0); + assert_eq!(etag_reader.try_resolve_etag(), Some(expected)); + } + #[tokio::test] async fn test_etag_reader_multiple_get() { let data = b"abc123"; diff --git a/crates/rio/src/hardlimit_reader.rs b/crates/rio/src/hardlimit_reader.rs index e50b052f5..c74f37d89 100644 --- a/crates/rio/src/hardlimit_reader.rs +++ b/crates/rio/src/hardlimit_reader.rs @@ -12,6 +12,7 @@ // See the License for the specific language governing permissions and // limitations under the License. +use crate::{BlockReadable, BoxReadBlockFuture, Reader}; use pin_project_lite::pin_project; use std::io::{Error, Result}; use std::pin::Pin; @@ -61,11 +62,51 @@ where delegate_reader_capabilities_generic!(HardLimitReader, inner); +impl BlockReadable for HardLimitReader +where + R: Reader, +{ + fn read_block<'a>(&'a mut self, buf: &'a mut [u8]) -> BoxReadBlockFuture<'a> { + Box::pin(async move { + if self.remaining < 0 { + return Err(Error::other("input provided more bytes than specified")); + } + + let max_len = match usize::try_from(self.remaining) { + Ok(remaining) => remaining.min(buf.len()), + Err(_) => buf.len(), + }; + + if max_len == 0 { + let mut probe = [0_u8; 1]; + match self.inner.read_block(&mut probe).await { + Ok(0) => return Ok(0), + Ok(n) => { + self.remaining -= n as i64; + return Err(Error::other("input provided more bytes than specified")); + } + Err(err) if err.kind() == std::io::ErrorKind::UnexpectedEof => return Ok(0), + Err(err) => return Err(err), + } + } + + let n = self.inner.read_block(&mut buf[..max_len]).await?; + self.remaining -= n as i64; + if self.remaining < 0 { + return Err(Error::other("input provided more bytes than specified")); + } + + Ok(n) + }) + } +} + #[cfg(test)] mod tests { use std::vec; use super::*; + use crate::{BlockReadable, WarpReader}; use rustfs_utils::read_full; use tokio::io::{AsyncReadExt, BufReader}; @@ -128,4 +169,20 @@ mod tests { assert_eq!(n, 0); assert_eq!(&buf, data); } + + #[tokio::test] + async fn test_hardlimit_reader_read_block_enforces_limit() { + let data = b"abcdef"; + let reader = BufReader::new(&data[..]); + let reader = Box::new(WarpReader::new(reader)); + let mut hardlimit = HardLimitReader::new(reader, 3); + + let mut buf = [0_u8; 8]; + let n = hardlimit.read_block(&mut buf).await.unwrap(); + assert_eq!(n, 3); + assert_eq!(&buf[..n], b"abc"); + + let err = hardlimit.read_block(&mut buf).await.unwrap_err(); + assert_eq!(err.kind(), std::io::ErrorKind::Other); + } } diff --git a/crates/rio/src/hash_reader.rs b/crates/rio/src/hash_reader.rs index aee0a50d6..092eb3786 100644 --- a/crates/rio/src/hash_reader.rs +++ b/crates/rio/src/hash_reader.rs @@ -90,7 +90,10 @@ use crate::ChecksumType; use crate::Sha256Hasher; use crate::compress_index::{Index, TryGetIndex}; use crate::get_content_checksum; -use crate::{DynReader, EtagReader, EtagResolvable, HardLimitReader, HashReaderDetector, WarpReader, boxed_reader, wrap_reader}; +use crate::{ + BlockReadable, BoxReadBlockFuture, DynReader, EtagReader, EtagResolvable, HardLimitReader, HashReaderDetector, WarpReader, + boxed_reader, wrap_reader, +}; use base64::Engine; use base64::engine::general_purpose; use http::HeaderMap; @@ -408,6 +411,23 @@ impl HashReader { Ok(()) } + pub fn enable_auto_checksum(&mut self, checksum_type: ChecksumType) -> Result<(), std::io::Error> { + if !checksum_type.is_set() { + return Ok(()); + } + + let Some(hasher) = checksum_type.hasher() else { + return Err(std::io::Error::new(std::io::ErrorKind::InvalidData, "Invalid checksum type")); + }; + + self.content_hash = Some(Checksum { + checksum_type, + ..Default::default() + }); + self.content_hasher = Some(hasher); + Ok(()) + } + pub fn checksum(&self) -> Option { if self .content_hash @@ -449,6 +469,97 @@ impl HashReader { } map } + + pub fn finalize_content_hash(&mut self) -> std::io::Result> { + self.finish_checksum_validation()?; + Ok(self.content_hash.clone()) + } + + fn update_read_state(&mut self, data: &[u8]) -> std::io::Result<()> { + self.bytes_read += data.len() as u64; + + if data.is_empty() { + return Ok(()); + } + + if let Some(hasher) = self.content_sha256_hasher.as_mut() { + hasher.write_all(data)?; + } + + if let Some(hasher) = self.content_hasher.as_mut() { + hasher.write_all(data)?; + } + + Ok(()) + } + + fn finish_checksum_validation(&mut self) -> std::io::Result<()> { + if self.checksum_on_finish { + return Ok(()); + } + + if let (Some(mut hasher), Some(expected_sha256)) = (self.content_sha256_hasher.take(), self.content_sha256.as_ref()) { + let sha256 = hex_simd::encode_to_string(hasher.finalize(), hex_simd::AsciiCase::Lower); + if sha256 != *expected_sha256 { + error!("SHA256 mismatch, expected={:?}, actual={:?}", expected_sha256, sha256); + return Err(std::io::Error::new(std::io::ErrorKind::InvalidData, "SHA256 mismatch")); + } + } + + if let Some(mut expected_content_hash) = self.content_hash.clone() + && let Some(mut hasher) = self.content_hasher.take() + { + if expected_content_hash.checksum_type.trailing() + && let Some(trailer) = self.trailer_s3s.as_ref() + && let Some(Some(checksum_str)) = trailer.read(|headers| { + expected_content_hash + .checksum_type + .key() + .and_then(|key| headers.get(key).and_then(|value| value.to_str().ok().map(|s| s.to_string()))) + }) + { + expected_content_hash.encoded = checksum_str; + expected_content_hash.raw = general_purpose::STANDARD + .decode(&expected_content_hash.encoded) + .map_err(|_| std::io::Error::other("Invalid base64 checksum"))?; + + if expected_content_hash.raw.is_empty() { + return Err(std::io::Error::other("Content hash mismatch")); + } + } + + let content_hash = hasher.finalize(); + if expected_content_hash.encoded.is_empty() { + expected_content_hash.raw = content_hash.clone(); + expected_content_hash.encoded = general_purpose::STANDARD.encode(&content_hash); + self.content_hash = Some(expected_content_hash); + } else if content_hash != expected_content_hash.raw { + let expected_hex = hex_simd::encode_to_string(&expected_content_hash.raw, hex_simd::AsciiCase::Lower); + let actual_hex = hex_simd::encode_to_string(content_hash, hex_simd::AsciiCase::Lower); + error!( + "Content hash mismatch, type={:?}, encoded={:?}, expected={:?}, actual={:?}", + expected_content_hash.checksum_type, expected_content_hash.encoded, expected_hex, actual_hex + ); + let checksum_err = crate::errors::ChecksumMismatch { + want: expected_hex, + got: actual_hex, + }; + return Err(std::io::Error::new(std::io::ErrorKind::InvalidData, checksum_err)); + } + } + + self.checksum_on_finish = true; + Ok(()) + } + + pub async fn read_block(&mut self, buf: &mut [u8]) -> std::io::Result { + let n = self.inner.read_block(buf).await?; + self.update_read_state(&buf[..n])?; + if n == 0 { + self.finish_checksum_validation()?; + } + Ok(n) + } } impl HashReaderMut for HashReader { @@ -508,84 +619,23 @@ impl HashReaderMut for HashReader { impl AsyncRead for HashReader { fn poll_read(self: Pin<&mut Self>, cx: &mut Context<'_>, buf: &mut ReadBuf<'_>) -> Poll> { - let this = self.project(); + let this = self.get_mut(); let before = buf.filled().len(); - match this.inner.poll_read(cx, buf) { + match Pin::new(&mut this.inner).poll_read(cx, buf) { Poll::Pending => Poll::Pending, Poll::Ready(Ok(())) => { let data = &buf.filled()[before..]; let filled = data.len(); - - *this.bytes_read += filled as u64; - - if filled > 0 { - // Update SHA256 hasher - if let Some(hasher) = this.content_sha256_hasher - && let Err(e) = hasher.write_all(data) - { - error!("SHA256 hasher write error, error={:?}", e); - return Poll::Ready(Err(std::io::Error::other(e))); - } - - // Update content hasher - if let Some(hasher) = this.content_hasher - && let Err(e) = hasher.write_all(data) - { - return Poll::Ready(Err(std::io::Error::other(e))); - } + if let Err(e) = this.update_read_state(data) { + error!("hash reader state update error, error={:?}", e); + return Poll::Ready(Err(std::io::Error::other(e))); } - if filled == 0 && !*this.checksum_on_finish { - // check SHA256 - if let (Some(hasher), Some(expected_sha256)) = (this.content_sha256_hasher, this.content_sha256) { - let sha256 = hex_simd::encode_to_string(hasher.finalize(), hex_simd::AsciiCase::Lower); - if sha256 != *expected_sha256 { - error!("SHA256 mismatch, expected={:?}, actual={:?}", expected_sha256, sha256); - return Poll::Ready(Err(std::io::Error::new(std::io::ErrorKind::InvalidData, "SHA256 mismatch"))); - } - } - - // check content hasher - if let (Some(hasher), Some(expected_content_hash)) = (this.content_hasher, this.content_hash) { - if expected_content_hash.checksum_type.trailing() - && let Some(trailer) = this.trailer_s3s.as_ref() - && let Some(Some(checksum_str)) = trailer.read(|headers| { - expected_content_hash - .checksum_type - .key() - .and_then(|key| headers.get(key).and_then(|value| value.to_str().ok().map(|s| s.to_string()))) - }) - { - expected_content_hash.encoded = checksum_str; - expected_content_hash.raw = general_purpose::STANDARD - .decode(&expected_content_hash.encoded) - .map_err(|_| std::io::Error::other("Invalid base64 checksum"))?; - - if expected_content_hash.raw.is_empty() { - return Poll::Ready(Err(std::io::Error::other("Content hash mismatch"))); - } - } - - let content_hash = hasher.finalize(); - - if content_hash != expected_content_hash.raw { - let expected_hex = hex_simd::encode_to_string(&expected_content_hash.raw, hex_simd::AsciiCase::Lower); - let actual_hex = hex_simd::encode_to_string(content_hash, hex_simd::AsciiCase::Lower); - error!( - "Content hash mismatch, type={:?}, encoded={:?}, expected={:?}, actual={:?}", - expected_content_hash.checksum_type, expected_content_hash.encoded, expected_hex, actual_hex - ); - // Use ChecksumMismatch error so that API layer can return BadDigest - let checksum_err = crate::errors::ChecksumMismatch { - want: expected_hex, - got: actual_hex, - }; - return Poll::Ready(Err(std::io::Error::new(std::io::ErrorKind::InvalidData, checksum_err))); - } - } - - *this.checksum_on_finish = true; + if filled == 0 + && let Err(e) = this.finish_checksum_validation() + { + return Poll::Ready(Err(e)); } Poll::Ready(Ok(())) } @@ -623,13 +673,39 @@ impl TryGetIndex for HashReader { } } +impl BlockReadable for HashReader { + fn read_block<'a>(&'a mut self, buf: &'a mut [u8]) -> BoxReadBlockFuture<'a> { + Box::pin(async move { self.read_block(buf).await }) + } +} + #[cfg(test)] mod tests { use super::*; use crate::{DecryptReader, EncryptReader, encrypt_reader, wrap_reader}; use rand::RngExt; use std::io::Cursor; - use tokio::io::{AsyncReadExt, BufReader}; + use std::pin::Pin; + use std::task::{Context, Poll}; + use tokio::io::{AsyncRead, AsyncReadExt, BufReader, ReadBuf}; + + struct UnexpectedEofReader; + + impl AsyncRead for UnexpectedEofReader { + fn poll_read(self: Pin<&mut Self>, _cx: &mut Context<'_>, _buf: &mut ReadBuf<'_>) -> Poll> { + Poll::Ready(Ok(())) + } + } + + impl BlockReadable for UnexpectedEofReader { + fn read_block<'a>(&'a mut self, _buf: &'a mut [u8]) -> BoxReadBlockFuture<'a> { + Box::pin(async { Err(std::io::Error::new(std::io::ErrorKind::UnexpectedEof, "synthetic unexpected eof")) }) + } + } + + impl EtagResolvable for UnexpectedEofReader {} + impl HashReaderDetector for UnexpectedEofReader {} + impl TryGetIndex for UnexpectedEofReader {} #[tokio::test] async fn test_hashreader_wrapping_logic() { @@ -742,6 +818,37 @@ mod tests { assert_eq!(buf, data); } + #[tokio::test] + async fn test_hashreader_read_block_reads_full_and_tail_block() { + let data = b"hello block reader"; + 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 first = [0_u8; 8]; + let n1 = hash_reader.read_block(&mut first).await.unwrap(); + assert_eq!(n1, 8); + assert_eq!(&first[..n1], b"hello bl"); + + let mut second = [0_u8; 32]; + let n2 = hash_reader.read_block(&mut second).await.unwrap(); + assert_eq!(n2, data.len() - n1); + assert_eq!(&second[..n2], b"ock reader"); + + let n3 = hash_reader.read_block(&mut second).await.unwrap(); + assert_eq!(n3, 0); + } + + #[tokio::test] + async fn test_hashreader_read_block_propagates_unexpected_eof() { + let mut hash_reader = HashReader::new(Box::new(UnexpectedEofReader), 0, 0, None, None, true).unwrap(); + let mut buf = [0_u8; 8]; + + let err = hash_reader.read_block(&mut buf).await.unwrap_err(); + + assert_eq!(err.kind(), std::io::ErrorKind::UnexpectedEof); + } + #[tokio::test] async fn test_hashreader_new_logic() { let data = b"test data"; diff --git a/crates/rio/src/http_reader.rs b/crates/rio/src/http_reader.rs index 39ef43ecd..b207fbc95 100644 --- a/crates/rio/src/http_reader.rs +++ b/crates/rio/src/http_reader.rs @@ -23,7 +23,6 @@ use rustfs_utils::get_env_opt_str; use std::io::IoSlice; use std::io::{self, Error}; use std::net::IpAddr; -use std::ops::Not as _; use std::pin::Pin; use std::sync::LazyLock; use std::task::{Context, Poll}; @@ -137,6 +136,71 @@ fn get_http_client(url: &str) -> Client { CLIENT.clone() } +type HttpByteStream = Pin> + Send + Sync>>; + +async fn request_http_byte_stream( + url: String, + method: Method, + headers: HeaderMap, + body: Option>, + meter_stream_recv_bytes: bool, +) -> io::Result<(bool, HttpByteStream)> { + let track_internode_metrics = is_internode_rpc_url(&url); + let client = get_http_client(&url); + let mut request: RequestBuilder = client.request(method, url.clone()).headers(headers); + if let Some(body) = body { + request = request.body(body); + } + + let resp = request.send().await.map_err(|e| { + if track_internode_metrics { + global_internode_metrics().record_error(); + } + Error::other(format!("HttpReader HTTP request error: {e}")) + })?; + + if !resp.status().is_success() { + if track_internode_metrics { + global_internode_metrics().record_error(); + } + return Err(Error::other(format!( + "HttpReader HTTP request failed with non-200 status {}", + resp.status() + ))); + } + + if track_internode_metrics { + global_internode_metrics().record_outgoing_request(); + } + + let stream = resp + .bytes_stream() + .map_ok(move |bytes| { + if track_internode_metrics && meter_stream_recv_bytes { + global_internode_metrics().record_recv_bytes(bytes.len()); + } + bytes + }) + .map_err(move |e| { + if track_internode_metrics { + global_internode_metrics().record_error(); + } + Error::other(format!("HttpReader stream error: {e}")) + }); + + Ok((track_internode_metrics, Box::pin(stream))) +} + +pub async fn open_http_byte_stream( + url: String, + method: Method, + headers: HeaderMap, + body: Option>, +) -> io::Result { + let (_track_internode_metrics, stream) = request_http_byte_stream(url, method, headers, body, true).await?; + Ok(stream) +} + pin_project! { pub struct HttpReader { url:String, @@ -161,43 +225,11 @@ impl HttpReader { body: Option>, _read_buf_size: usize, ) -> io::Result { - let track_internode_metrics = is_internode_rpc_url(&url); - let client = get_http_client(&url); - let mut request: RequestBuilder = client.request(method.clone(), url.clone()).headers(headers.clone()); - if let Some(body) = body { - request = request.body(body); - } - - let resp = request.send().await.map_err(|e| { - if track_internode_metrics { - global_internode_metrics().record_error(); - } - Error::other(format!("HttpReader HTTP request error: {e}")) - })?; - - if resp.status().is_success().not() { - if track_internode_metrics { - global_internode_metrics().record_error(); - } - return Err(Error::other(format!( - "HttpReader HTTP request failed with non-200 status {}", - resp.status() - ))); - } - - if track_internode_metrics { - global_internode_metrics().record_outgoing_request(); - } - - let stream = resp.bytes_stream().map_err(move |e| { - if track_internode_metrics { - global_internode_metrics().record_error(); - } - Error::other(format!("HttpReader stream error: {e}")) - }); + let (track_internode_metrics, stream) = + request_http_byte_stream(url.clone(), method.clone(), headers.clone(), body, false).await?; Ok(Self { - inner: StreamReader::new(Box::pin(stream)), + inner: StreamReader::new(stream), url, method, headers, diff --git a/crates/rio/src/lib.rs b/crates/rio/src/lib.rs index 9663f133d..6e9040382 100644 --- a/crates/rio/src/lib.rs +++ b/crates/rio/src/lib.rs @@ -15,6 +15,10 @@ // Default encryption block size - aligned with system default read buffer size (1MB) pub const DEFAULT_ENCRYPTION_BLOCK_SIZE: usize = 1024 * 1024; +use std::future::Future; +use std::pin::Pin; +use tokio::io::AsyncReadExt; + macro_rules! delegate_reader_capabilities_generic { ($name:ident<$inner_ty:ident>, $inner:ident) => { impl<$inner_ty> crate::EtagResolvable for $name<$inner_ty> @@ -114,14 +118,39 @@ pub use compress_index::{Index, TryGetIndex}; mod etag; +pub type BoxReadBlockFuture<'a> = Pin> + Send + 'a>>; + +pub trait BlockReadable { + fn read_block<'a>(&'a mut self, buf: &'a mut [u8]) -> BoxReadBlockFuture<'a>; +} + +fn read_block_via_async_read<'a, R>(reader: &'a mut R, buf: &'a mut [u8]) -> BoxReadBlockFuture<'a> +where + R: tokio::io::AsyncRead + Unpin + Send + Sync + 'a, +{ + Box::pin(async move { + let mut total = 0; + + while total < buf.len() { + match reader.read(&mut buf[total..]).await { + Ok(0) => return Ok(total), + Ok(n) => total += n, + Err(err) => return Err(err), + } + } + + Ok(total) + }) +} + 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 trait Reader: ReadStream + ReaderCapabilities + BlockReadable {} +impl Reader for T where T: ReadStream + ReaderCapabilities + BlockReadable {} pub type DynReader = Box; @@ -154,6 +183,42 @@ pub trait HashReaderDetector { } } +impl BlockReadable for crate::WarpReader +where + R: tokio::io::AsyncRead + Unpin + Send + Sync, +{ + fn read_block<'a>(&'a mut self, buf: &'a mut [u8]) -> BoxReadBlockFuture<'a> { + read_block_via_async_read(self, buf) + } +} + +impl BlockReadable for tokio::io::BufReader +where + R: tokio::io::AsyncRead + Unpin + Send + Sync, +{ + fn read_block<'a>(&'a mut self, buf: &'a mut [u8]) -> BoxReadBlockFuture<'a> { + read_block_via_async_read(self, buf) + } +} + +impl BlockReadable for crate::LimitReader +where + R: Reader, +{ + fn read_block<'a>(&'a mut self, buf: &'a mut [u8]) -> BoxReadBlockFuture<'a> { + read_block_via_async_read(self, buf) + } +} + +impl BlockReadable for crate::DecryptReader +where + R: Reader, +{ + fn read_block<'a>(&'a mut self, buf: &'a mut [u8]) -> BoxReadBlockFuture<'a> { + read_block_via_async_read(self, buf) + } +} + pub fn boxed_reader(reader: R) -> DynReader where R: Reader + 'static, @@ -198,3 +263,74 @@ where self.as_ref().try_get_index() } } + +impl BlockReadable for Box +where + T: BlockReadable + ?Sized, +{ + fn read_block<'a>(&'a mut self, buf: &'a mut [u8]) -> BoxReadBlockFuture<'a> { + self.as_mut().read_block(buf) + } +} + +#[cfg(test)] +mod tests { + use super::*; + use std::collections::VecDeque; + use std::io::{self, ErrorKind}; + use std::pin::Pin; + use std::task::{Context, Poll}; + use tokio::io::{AsyncRead, ReadBuf}; + + enum ReadStep { + Data(Vec), + Error(ErrorKind), + Eof, + } + + struct StepReader { + steps: VecDeque, + } + + impl StepReader { + fn new(steps: impl IntoIterator) -> Self { + Self { + steps: steps.into_iter().collect(), + } + } + } + + impl AsyncRead for StepReader { + fn poll_read(mut self: Pin<&mut Self>, _cx: &mut Context<'_>, buf: &mut ReadBuf<'_>) -> Poll> { + match self.steps.pop_front().unwrap_or(ReadStep::Eof) { + ReadStep::Data(data) => { + buf.put_slice(&data); + Poll::Ready(Ok(())) + } + ReadStep::Error(kind) => Poll::Ready(Err(io::Error::new(kind, "synthetic read failure"))), + ReadStep::Eof => Poll::Ready(Ok(())), + } + } + } + + #[tokio::test] + async fn test_read_block_via_async_read_preserves_midstream_error_kind() { + let reader = StepReader::new([ReadStep::Data(b"ab".to_vec()), ReadStep::Error(ErrorKind::ConnectionReset)]); + let mut reader = WarpReader::new(reader); + let mut buf = [0_u8; 4]; + + let err = reader.read_block(&mut buf).await.unwrap_err(); + + assert_eq!(err.kind(), ErrorKind::ConnectionReset); + } + + #[tokio::test] + async fn test_read_block_via_async_read_returns_zero_on_initial_eof() { + let mut reader = WarpReader::new(StepReader::new([ReadStep::Eof])); + let mut buf = [0_u8; 4]; + + let n = reader.read_block(&mut buf).await.unwrap(); + + assert_eq!(n, 0); + } +} diff --git a/crates/scanner/tests/lifecycle_integration_test.rs b/crates/scanner/tests/lifecycle_integration_test.rs index b8e8304f7..ca0d475ed 100644 --- a/crates/scanner/tests/lifecycle_integration_test.rs +++ b/crates/scanner/tests/lifecycle_integration_test.rs @@ -24,7 +24,7 @@ use rustfs_ecstore::{ pools::path2_bucket_object_with_base_path, store::ECStore, store_api::{ - BucketOperations, MakeBucketOptions, MultipartOperations, ObjectIO, ObjectOperations, ObjectOptions, PutObjReader, + BucketOperations, ChunkNativePutData, MakeBucketOptions, MultipartOperations, ObjectIO, ObjectOperations, ObjectOptions, }, tier::{ tier_config::{TierConfig, TierMinIO, TierType}, @@ -235,7 +235,7 @@ async fn create_test_lock_bucket(ecstore: &Arc, bucket_name: &str) { /// Test helper: Upload test object async fn upload_test_object(ecstore: &Arc, bucket: &str, object: &str, data: &[u8]) { - let mut reader = PutObjReader::from_vec(data.to_vec()); + let mut reader = ChunkNativePutData::from_vec(data.to_vec()); let object_info = (**ecstore) .put_object(bucket, object, &mut reader, &ObjectOptions::default()) .await @@ -781,7 +781,7 @@ mod serial_tests { .await .expect("Failed to set lifecycle configuration"); - let mut reader = PutObjReader::from_vec(put_payload.to_vec()); + let mut reader = ChunkNativePutData::from_vec(put_payload.to_vec()); let mut metadata = HashMap::new(); metadata.insert("content-type".to_string(), "text/plain".to_string()); ecstore @@ -838,7 +838,7 @@ mod serial_tests { .expect("Failed to create multipart upload"); let part_data = b"multipart immediate transition"; - let mut reader = PutObjReader::from_vec(part_data.to_vec()); + let mut reader = ChunkNativePutData::from_vec(part_data.to_vec()); let part = ecstore .put_object_part( multipart_bucket.as_str(), @@ -903,7 +903,7 @@ mod serial_tests { .get_object_info(src_bucket.as_str(), src_object, &ObjectOptions::default()) .await .expect("Failed to load source object info"); - src_info.put_object_reader = Some(PutObjReader::from_vec(payload.to_vec())); + src_info.put_object_reader = Some(ChunkNativePutData::from_vec(payload.to_vec())); ecstore .copy_object( @@ -969,7 +969,7 @@ mod serial_tests { .await .expect("Failed to create multipart upload"); - let mut part1_reader = PutObjReader::from_vec(part1); + let mut part1_reader = ChunkNativePutData::from_vec(part1); let uploaded_part1 = ecstore .put_object_part( bucket_name.as_str(), @@ -982,7 +982,7 @@ mod serial_tests { .await .expect("Failed to upload first multipart part"); - let mut part2_reader = PutObjReader::from_vec(part2); + let mut part2_reader = ChunkNativePutData::from_vec(part2); let uploaded_part2 = ecstore .put_object_part( bucket_name.as_str(), diff --git a/rustfs/Cargo.toml b/rustfs/Cargo.toml index 16a27178c..ff923e174 100644 --- a/rustfs/Cargo.toml +++ b/rustfs/Cargo.toml @@ -86,6 +86,7 @@ rustfs-utils = { workspace = true, features = ["full"] } rustfs-zip = { workspace = true } rustfs-io-core = { workspace = true } rustfs-io-metrics = { workspace = true } +rustfs-object-io = { workspace = true } rustfs-concurrency = { workspace = true } rustfs-scanner = { workspace = true } diff --git a/rustfs/src/app/lifecycle_transition_api_test.rs b/rustfs/src/app/lifecycle_transition_api_test.rs index ae144b880..8ed844525 100644 --- a/rustfs/src/app/lifecycle_transition_api_test.rs +++ b/rustfs/src/app/lifecycle_transition_api_test.rs @@ -27,8 +27,8 @@ use rustfs_ecstore::{ global::GLOBAL_TierConfigMgr, store::ECStore, store_api::{ - BucketOperations, BucketOptions, MakeBucketOptions, MultipartOperations, ObjectIO, ObjectOperations, ObjectOptions, - PutObjReader, + BucketOperations, BucketOptions, ChunkNativePutData, MakeBucketOptions, MultipartOperations, ObjectIO, ObjectOperations, + ObjectOptions, }, tier::{ tier_config::{TierConfig, TierType}, @@ -148,7 +148,7 @@ async fn upload_test_object( object: &str, data: &[u8], ) -> rustfs_ecstore::store_api::ObjectInfo { - let mut reader = PutObjReader::from_vec(data.to_vec()); + let mut reader = ChunkNativePutData::from_vec(data.to_vec()); (**ecstore) .put_object(bucket, object, &mut reader, &ObjectOptions::default()) .await @@ -446,7 +446,7 @@ async fn complete_multipart_upload_transitions_immediately_via_usecase() { .await .expect("Failed to create multipart upload"); - let mut reader = PutObjReader::from_vec(payload.to_vec()); + let mut reader = ChunkNativePutData::from_vec(payload.to_vec()); let uploaded_part = ecstore .put_object_part(bucket.as_str(), object, &upload.upload_id, 1, &mut reader, &ObjectOptions::default()) .await diff --git a/rustfs/src/app/multipart_usecase.rs b/rustfs/src/app/multipart_usecase.rs index 0f97f1fe1..b370fc89e 100644 --- a/rustfs/src/app/multipart_usecase.rs +++ b/rustfs/src/app/multipart_usecase.rs @@ -44,10 +44,13 @@ use rustfs_ecstore::compress::is_compressible; use rustfs_ecstore::error::{StorageError, is_err_object_not_found, is_err_version_not_found}; use rustfs_ecstore::new_object_layer_fn; 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::{ + ChunkNativePutData, CompletePart, HTTPRangeSpec, MultipartUploadResult, ObjectIO, ObjectOptions, +}; use rustfs_ecstore::store_api::{MultipartOperations, ObjectOperations}; use rustfs_filemeta::{ReplicationStatusType, ReplicationType}; -use rustfs_rio::{CompressReader, HashReader}; +use rustfs_object_io::put::PutObjectChecksums; +use rustfs_rio::{CompressReader, HashReader, Reader, WarpReader}; use rustfs_s3_common::S3Operation; use rustfs_targets::EventName; use rustfs_utils::CompressionAlgorithm; @@ -719,6 +722,13 @@ impl DefaultMultipartUsecase { .map_err(ApiError::from)?; let mut size = size.ok_or_else(|| s3_error!(UnexpectedContent))?; + let mut requested_checksum_type = rustfs_rio::ChecksumType::from_header(&req.headers); + if !requested_checksum_type.is_set() + && let Some(checksum_algo) = fi.user_defined.get(rustfs_rio::RUSTFS_MULTIPART_CHECKSUM) + && let Some(checksum_type) = fi.user_defined.get(rustfs_rio::RUSTFS_MULTIPART_CHECKSUM_TYPE) + { + requested_checksum_type = rustfs_rio::ChecksumType::from_string_with_obj_type(checksum_algo, checksum_type); + } // Apply adaptive buffer sizing based on part size for optimal streaming performance. // Uses workload profile configuration (enabled by default) to select appropriate buffer size. @@ -731,6 +741,8 @@ 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 { @@ -744,31 +756,32 @@ impl DefaultMultipartUsecase { let mut sha256hex = get_content_sha256_with_query(&req.headers, req.uri.query()); - 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 is_compressible { + let mut hrd = + HashReader::new(reader, 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()); } + if requested_checksum_type.is_set() && hrd.checksum().is_none() { + hrd.enable_auto_checksum(requested_checksum_type).map_err(ApiError::from)?; + } + let compress_reader = CompressReader::new(hrd, CompressionAlgorithm::default()); + reader = Box::new(compress_reader); 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(body, size, actual_size, md5hex, sha256hex, false).map_err(ApiError::from)? - }; + md5hex = None; + sha256hex = None; + } + + let mut reader = HashReader::new(reader, 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()); } + if requested_checksum_type.is_set() && reader.checksum().is_none() { + reader.enable_auto_checksum(requested_checksum_type).map_err(ApiError::from)?; + } let has_ssec = sse_customer_algorithm.is_some(); // When SSE-C headers are present, skip managed-encryption metadata to avoid @@ -818,9 +831,8 @@ impl DefaultMultipartUsecase { let requested_kms_key_id = material.kms_key_id.clone(); let encrypted_reader = material.wrap_reader(reader); - reader = - HashReader::from_reader(encrypted_reader, HashReader::SIZE_PRESERVE_LAYER, actual_size, None, None, false) - .map_err(ApiError::from)?; + reader = HashReader::new(encrypted_reader, HashReader::SIZE_PRESERVE_LAYER, actual_size, None, None, false) + .map_err(ApiError::from)?; fi.user_defined.extend(material.metadata); @@ -829,18 +841,20 @@ impl DefaultMultipartUsecase { None => (None, None), }; - let mut reader = PutObjReader::new(reader); + let mut reader = ChunkNativePutData::new(reader); let info = store .put_object_part(&bucket, &key, &upload_id, part_id, &mut reader, &opts) .await .map_err(ApiError::from)?; - let mut checksum_crc32 = input.checksum_crc32; - let mut checksum_crc32c = input.checksum_crc32c; - let mut checksum_sha1 = input.checksum_sha1; - let mut checksum_sha256 = input.checksum_sha256; - let mut checksum_crc64nvme = input.checksum_crc64nvme; + let mut checksums = PutObjectChecksums { + crc32: input.checksum_crc32, + crc32c: input.checksum_crc32c, + sha1: input.checksum_sha1, + sha256: input.checksum_sha256, + crc64nvme: input.checksum_crc64nvme, + }; if let Some(alg) = &input.checksum_algorithm && let Some(Some(checksum_str)) = req.trailing_headers.as_ref().map(|trailer| { @@ -860,25 +874,26 @@ impl DefaultMultipartUsecase { }) { match alg.as_str() { - ChecksumAlgorithm::CRC32 => checksum_crc32 = checksum_str, - ChecksumAlgorithm::CRC32C => checksum_crc32c = checksum_str, - ChecksumAlgorithm::SHA1 => checksum_sha1 = checksum_str, - ChecksumAlgorithm::SHA256 => checksum_sha256 = checksum_str, - ChecksumAlgorithm::CRC64NVME => checksum_crc64nvme = checksum_str, + ChecksumAlgorithm::CRC32 => checksums.crc32 = checksum_str, + ChecksumAlgorithm::CRC32C => checksums.crc32c = checksum_str, + ChecksumAlgorithm::SHA1 => checksums.sha1 = checksum_str, + ChecksumAlgorithm::SHA256 => checksums.sha256 = checksum_str, + ChecksumAlgorithm::CRC64NVME => checksums.crc64nvme = checksum_str, _ => (), } } + checksums.merge_from_map(&reader.content_crc()); let output = UploadPartOutput { server_side_encryption: requested_sse, ssekms_key_id: requested_kms_key_id, sse_customer_algorithm, sse_customer_key_md5, - checksum_crc32, - checksum_crc32c, - checksum_sha1, - checksum_sha256, - checksum_crc64nvme, + checksum_crc32: checksums.crc32, + checksum_crc32c: checksums.crc32c, + checksum_sha1: checksums.sha1, + checksum_sha256: checksums.sha256, + checksum_crc64nvme: checksums.crc64nvme, e_tag: info.etag.map(|etag| to_s3s_etag(&etag)), ..Default::default() }; @@ -1116,6 +1131,8 @@ 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, @@ -1127,74 +1144,23 @@ 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; - let mut reader = match sse_decryption(src_decryption_request).await? { - Some(material) => { - if let Some(original) = material.original_size { - src_info.actual_size = original; - } + 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; + } - 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 mut reader = HashReader::new(reader, size, actual_size, None, None, false).map_err(ApiError::from)?; let server_side_encryption = mp_info .user_defined @@ -1235,9 +1201,8 @@ impl DefaultMultipartUsecase { let requested_kms_key_id = material.kms_key_id.clone(); let encrypted_reader = material.wrap_reader(reader); - reader = - HashReader::from_reader(encrypted_reader, HashReader::SIZE_PRESERVE_LAYER, actual_size, None, None, false) - .map_err(ApiError::from)?; + reader = HashReader::new(encrypted_reader, HashReader::SIZE_PRESERVE_LAYER, actual_size, None, None, false) + .map_err(ApiError::from)?; mp_info.user_defined.extend(material.metadata); @@ -1246,7 +1211,7 @@ impl DefaultMultipartUsecase { None => (None, None), }; - let mut reader = PutObjReader::new(reader); + let mut reader = ChunkNativePutData::new(reader); let dst_opts = ObjectOptions { user_defined: mp_info.user_defined.clone(), diff --git a/rustfs/src/app/object_usecase.rs b/rustfs/src/app/object_usecase.rs index efebe94e4..c94347b5f 100644 --- a/rustfs/src/app/object_usecase.rs +++ b/rustfs/src/app/object_usecase.rs @@ -14,22 +14,31 @@ //! Object application use-case contracts. +mod app_adapters; +mod get_object_flow; +mod get_object_zero_copy; +mod put_object_extract; +mod put_object_flow; +mod types; +#[cfg(test)] +mod zero_copy_tests; +use self::app_adapters::*; +use self::get_object_flow::{GetObjectBootstrap, GetObjectFlowRuntime}; +use self::types::*; + use crate::app::context::{AppContext, default_notify_interface, get_global_app_context}; use crate::capacity::capacity_manager::get_capacity_manager; use crate::config::RustFSBufferConfig; use crate::error::ApiError; use crate::storage::access::{PostObjectRequestMarker, authorize_request, has_bypass_governance_header, req_info_mut}; -use crate::storage::concurrency::{ - CachedGetObject, ConcurrencyManager, GetObjectGuard, get_concurrency_aware_buffer_size, get_concurrency_manager, -}; +use crate::storage::concurrency::{CachedGetObject, ConcurrencyManager, GetObjectGuard, get_concurrency_manager}; use crate::storage::ecfs::*; use crate::storage::head_prefix::{head_prefix_not_found_message, probe_prefix_has_children}; -use crate::storage::helper::{OperationHelper, spawn_background, spawn_background_with_context}; +use crate::storage::helper::OperationHelper; use crate::storage::options::{ copy_dst_opts, copy_src_opts, del_opts, extract_metadata, extract_metadata_from_mime_with_object_name, filter_object_metadata, get_content_sha256_with_query, get_opts, normalize_content_encoding_for_storage, put_opts, }; -use crate::storage::request_context::spawn_traced; use crate::storage::s3_api::multipart::parse_list_parts_params; use crate::storage::s3_api::{acl, restore, select}; use crate::storage::timeout_wrapper::{RequestTimeoutWrapper, TimeoutConfig}; @@ -43,8 +52,6 @@ use http::{HeaderMap, HeaderValue, StatusCode}; use md5::Context as Md5Context; use metrics::{counter, histogram}; use pin_project_lite::pin_project; -// Performance metrics recording (with zero-copy-metrics integration) -use rustfs_concurrency::GetObjectQueueSnapshot; use rustfs_ecstore::bucket::quota::checker::QuotaChecker; use rustfs_ecstore::bucket::{ lifecycle::{ @@ -77,8 +84,8 @@ use rustfs_ecstore::error::{StorageError, is_err_bucket_not_found, is_err_object use rustfs_ecstore::new_object_layer_fn; use rustfs_ecstore::set_disk::is_valid_storage_class; use rustfs_ecstore::store_api::{ - BucketOperations, BucketOptions, HTTPRangeSpec, ObjectIO, ObjectInfo, ObjectOperations, ObjectOptions, ObjectToDelete, - PutObjReader, + BucketOperations, BucketOptions, ChunkNativePutData, HTTPRangeSpec, ObjectIO, ObjectInfo, ObjectOperations, ObjectOptions, + ObjectToDelete, }; use rustfs_filemeta::{ REPLICATE_INCOMING_DELETE, ReplicationStatusType, ReplicationType, RestoreStatusOps, VersionPurgeStatusType, @@ -86,26 +93,22 @@ use rustfs_filemeta::{ }; use rustfs_io_metrics; use rustfs_notify::EventArgsBuilder; -use rustfs_policy::policy::action::{Action, S3Action}; -use rustfs_rio::{CompressReader, DynReader, HashReader, wrap_reader}; -use rustfs_s3_common::S3Operation; -use rustfs_s3select_api::{ - object_store::bytes_stream, - query::{Context, Query}, +use rustfs_object_io::put::{ + is_post_object_sse_kms_requested, is_put_object_extract_requested, is_sse_kms_requested, resolve_put_body_size, }; +use rustfs_policy::policy::action::{Action, S3Action}; +use rustfs_rio::{CompressReader, EtagReader, HashReader, Reader, WarpReader}; +use rustfs_s3_common::S3Operation; +use rustfs_s3select_api::query::{Context, Query}; use rustfs_s3select_query::get_global_db; use rustfs_targets::EventName; use rustfs_utils::http::{ AMZ_BUCKET_REPLICATION_STATUS, AMZ_CHECKSUM_MODE, AMZ_CHECKSUM_TYPE, AMZ_WEBSITE_REDIRECT_LOCATION, CONTENT_TYPE, SUFFIX_ACTUAL_SIZE, SUFFIX_COMPRESSION, SUFFIX_COMPRESSION_SIZE, SUFFIX_REPLICATION_STATUS, SUFFIX_REPLICATION_TIMESTAMP, headers::{ - AMZ_DECODED_CONTENT_LENGTH, AMZ_MINIO_SNOWBALL_IGNORE_DIRS, AMZ_MINIO_SNOWBALL_IGNORE_ERRORS, AMZ_MINIO_SNOWBALL_PREFIX, AMZ_OBJECT_LOCK_LEGAL_HOLD, AMZ_OBJECT_LOCK_LEGAL_HOLD_LOWER, AMZ_OBJECT_LOCK_MODE, AMZ_OBJECT_LOCK_MODE_LOWER, AMZ_OBJECT_LOCK_RETAIN_UNTIL_DATE, AMZ_OBJECT_LOCK_RETAIN_UNTIL_DATE_LOWER, AMZ_OBJECT_TAGGING, AMZ_RESTORE_EXPIRY_DAYS, - AMZ_RESTORE_REQUEST_DATE, AMZ_RUSTFS_SNOWBALL_IGNORE_DIRS, AMZ_RUSTFS_SNOWBALL_IGNORE_ERRORS, AMZ_RUSTFS_SNOWBALL_PREFIX, - AMZ_SERVER_SIDE_ENCRYPTION, AMZ_SERVER_SIDE_ENCRYPTION_CUSTOMER_ALGORITHM, AMZ_SERVER_SIDE_ENCRYPTION_KMS_ID, - AMZ_SNOWBALL_EXTRACT, AMZ_SNOWBALL_IGNORE_DIRS, AMZ_SNOWBALL_IGNORE_ERRORS, AMZ_SNOWBALL_PREFIX, AMZ_STORAGE_CLASS, - AMZ_TAG_COUNT, + AMZ_RESTORE_REQUEST_DATE, AMZ_STORAGE_CLASS, AMZ_TAG_COUNT, }, insert_str, remove_str, }; @@ -119,7 +122,6 @@ use s3s::dto::*; use s3s::header::{X_AMZ_RESTORE, X_AMZ_RESTORE_OUTPUT_PATH}; use s3s::{S3Error, S3ErrorCode, S3Request, S3Response, S3Result, s3_error}; use std::collections::HashMap; -use std::convert::Infallible; use std::ops::Add; use std::path::Path; use std::str::FromStr; @@ -131,7 +133,7 @@ use tokio::sync::RwLock; use tokio::sync::mpsc; use tokio_stream::wrappers::ReceiverStream; use tokio_tar::Archive; -use tokio_util::io::{ReaderStream, StreamReader}; +use tokio_util::io::StreamReader; use tracing::{debug, error, info, instrument, warn}; use uuid::Uuid; @@ -155,70 +157,6 @@ impl Drop for DeadlockRequestGuard { } } -struct GetObjectBootstrap { - timeout_config: TimeoutConfig, - wrapper: RequestTimeoutWrapper, - request_start: std::time::Instant, - request_guard: GetObjectGuard, - _deadlock_request_guard: DeadlockRequestGuard, - concurrent_requests: usize, -} - -struct GetObjectIoPlanning<'a> { - _disk_permit: tokio::sync::SemaphorePermit<'a>, - permit_wait_duration: Duration, - queue_status: concurrency::IoQueueStatus, - queue_utilization: f64, -} - -struct GetObjectRequestContext { - bucket: String, - key: String, - cache_key: String, - version_id_for_event: String, - part_number: Option, - rs: Option, - opts: ObjectOptions, -} - -struct GetObjectReadSetup { - info: ObjectInfo, - event_info: ObjectInfo, - final_stream: DynReader, - rs: Option, - content_type: Option, - last_modified: Option, - response_content_length: i64, - content_range: Option, - server_side_encryption: Option, - sse_customer_algorithm: Option, - sse_customer_key_md5: Option, - ssekms_key_id: Option, - encryption_applied: bool, -} - -struct GetObjectPreparedRead<'a> { - io_planning: GetObjectIoPlanning<'a>, - read_setup: GetObjectReadSetup, -} - -struct GetObjectStrategyContext { - io_strategy: concurrency::IoStrategy, - optimal_buffer_size: usize, -} - -struct GetObjectCachedHit { - output: GetObjectOutput, - event_info: ObjectInfo, -} - -struct GetObjectOutputContext { - output: GetObjectOutput, - event_info: ObjectInfo, - response_content_length: i64, - optimal_buffer_size: usize, -} - async fn enqueue_transitioned_delete_cleanup(bucket: &str, object: &str, opts: &ObjectOptions, existing: Option<&ObjectInfo>) { let Some(existing) = existing else { return; @@ -303,61 +241,6 @@ impl AsyncRead for ExtractArchiveEtagReader { } } -/// Determine if zero-copy write should be used for this PutObject operation. -/// -/// Zero-copy is beneficial for large objects without encryption or compression. -/// -/// # Arguments -/// -/// * `size` - Object size in bytes -/// * `headers` - HTTP headers (to check for encryption/compression) -/// -/// # Returns -/// -/// `true` if zero-copy should be used, `false` otherwise -fn should_use_zero_copy(size: i64, headers: &HeaderMap) -> bool { - // Only use zero-copy for objects larger than 1MB - const ZERO_COPY_MIN_SIZE: i64 = 1024 * 1024; - - if size <= ZERO_COPY_MIN_SIZE { - return false; - } - - // Don't use zero-copy if encryption is requested - if headers.get(AMZ_SERVER_SIDE_ENCRYPTION).is_some() - || headers.get(AMZ_SERVER_SIDE_ENCRYPTION_CUSTOMER_ALGORITHM).is_some() - || headers.get(AMZ_SERVER_SIDE_ENCRYPTION_KMS_ID).is_some() - { - return false; - } - - // Don't use zero-copy if compression is likely (compressible content types) - // The compression check happens later in the flow - if let Some(content_type) = headers.get(CONTENT_TYPE) - && let Ok(ct) = content_type.to_str() - { - // Skip zero-copy for easily compressible content types - // since compression will be applied - let compressible_types = [ - "text/plain", - "text/html", - "text/css", - "text/javascript", - "application/javascript", - "application/json", - "application/xml", - "text/xml", - ]; - for ct_type in compressible_types { - if ct.contains(ct_type) { - return false; - } - } - } - - true -} - #[cfg(test)] mod deadlock_request_guard_tests { use super::DeadlockRequestGuard; @@ -387,60 +270,6 @@ async fn maybe_enqueue_transition_immediate(obj_info: &ObjectInfo, src: LcEventS enqueue_transition_immediate(obj_info, src).await; } -/// Extract trailing-header checksum values, overriding the corresponding input fields. -fn apply_trailing_checksums( - algorithm: Option<&str>, - trailing_headers: &Option, - checksums: &mut PutObjectChecksums, -) { - let Some(alg) = algorithm else { return }; - let Some(checksum_str) = trailing_headers.as_ref().and_then(|trailer| { - let key = match alg { - ChecksumAlgorithm::CRC32 => rustfs_rio::ChecksumType::CRC32.key(), - ChecksumAlgorithm::CRC32C => rustfs_rio::ChecksumType::CRC32C.key(), - ChecksumAlgorithm::SHA1 => rustfs_rio::ChecksumType::SHA1.key(), - ChecksumAlgorithm::SHA256 => rustfs_rio::ChecksumType::SHA256.key(), - ChecksumAlgorithm::CRC64NVME => rustfs_rio::ChecksumType::CRC64_NVME.key(), - _ => return None, - }; - trailer.read(|headers| { - headers - .get(key.unwrap_or_default()) - .and_then(|value| value.to_str().ok().map(|s| s.to_string())) - }) - }) else { - return; - }; - - match alg { - ChecksumAlgorithm::CRC32 => checksums.crc32 = checksum_str, - ChecksumAlgorithm::CRC32C => checksums.crc32c = checksum_str, - ChecksumAlgorithm::SHA1 => checksums.sha1 = checksum_str, - ChecksumAlgorithm::SHA256 => checksums.sha256 = checksum_str, - ChecksumAlgorithm::CRC64NVME => checksums.crc64nvme = checksum_str, - _ => (), - } -} - -#[derive(Default)] -struct GetObjectChecksums { - crc32: Option, - crc32c: Option, - sha1: Option, - sha256: Option, - crc64nvme: Option, - checksum_type: Option, -} - -#[derive(Default)] -struct PutObjectChecksums { - crc32: Option, - crc32c: Option, - sha1: Option, - sha256: Option, - crc64nvme: Option, -} - fn normalize_delete_objects_version_id(version_id: Option) -> Result<(Option, Option), String> { let version_id = version_id.map(|v| v.trim().to_string()).filter(|v| !v.is_empty()); match version_id { @@ -471,148 +300,6 @@ fn build_put_object_expiration_header(event: &lifecycle::Event) -> Option, - ignore_dirs: bool, - ignore_errors: bool, -} - -fn header_value_is_true(headers: &HeaderMap, key: &str) -> bool { - headers - .get(key) - .and_then(|value| value.to_str().ok()) - .is_some_and(|value| value.trim().eq_ignore_ascii_case("true")) -} - -fn is_put_object_extract_requested(headers: &HeaderMap) -> bool { - header_value_is_true(headers, AMZ_SNOWBALL_EXTRACT) || header_value_is_true(headers, AMZ_SNOWBALL_EXTRACT_COMPAT) -} - -fn trimmed_header_value(headers: &HeaderMap, key: &str) -> Option { - headers - .get(key) - .and_then(|value| value.to_str().ok()) - .map(|value| value.trim().to_string()) -} - -fn is_exact_snowball_meta_key(key: &str, exact_keys: &[&str]) -> bool { - exact_keys.iter().any(|exact_key| key.eq_ignore_ascii_case(exact_key)) -} - -fn snowball_meta_value_by_suffix(headers: &HeaderMap, suffix_lower: &str, exact_keys: &[&str]) -> Option { - for (name, value) in headers { - let key = name.as_str(); - if key.starts_with(AMZ_META_PREFIX_LOWER) - && key.ends_with(suffix_lower) - && !is_exact_snowball_meta_key(key, exact_keys) - && let Ok(parsed) = value.to_str() - { - return Some(parsed.trim().to_string()); - } - } - - None -} - -fn snowball_meta_value(headers: &HeaderMap, exact_keys: &[&str], suffix_lower: &str) -> Option { - for key in exact_keys { - if let Some(value) = trimmed_header_value(headers, key) { - return Some(value); - } - } - - snowball_meta_value_by_suffix(headers, suffix_lower, exact_keys) -} - -fn snowball_meta_flag(headers: &HeaderMap, exact_keys: &[&str], suffix_lower: &str) -> bool { - snowball_meta_value(headers, exact_keys, suffix_lower).is_some_and(|value| value.eq_ignore_ascii_case("true")) -} - -fn normalize_snowball_prefix(prefix: &str) -> Option { - let normalized = prefix.trim().trim_matches('/'); - if normalized.is_empty() { - return None; - } - - Some(normalized.to_string()) -} - -fn normalize_extract_entry_key(path: &str, prefix: Option<&str>, is_dir: bool) -> String { - let path = path.trim_matches('/'); - let mut key = match prefix { - Some(prefix) if !path.is_empty() => format!("{prefix}/{path}"), - Some(prefix) => prefix.to_string(), - None => path.to_string(), - }; - - if is_dir && !key.ends_with('/') { - key.push('/'); - } - - key -} - -fn map_extract_archive_error(err: impl std::fmt::Display) -> S3Error { - s3_error!(InvalidArgument, "Failed to process archive entry: {}", err) -} - -async fn apply_extract_entry_pax_extensions( - entry: &mut tokio_tar::Entry>, - metadata: &mut HashMap, - opts: &mut ObjectOptions, -) -> S3Result<()> -where - R: AsyncRead + Send + Unpin + 'static, -{ - let Some(extensions) = entry.pax_extensions().await.map_err(map_extract_archive_error)? else { - return Ok(()); - }; - - for ext in extensions { - let ext = ext.map_err(map_extract_archive_error)?; - let key = ext.key().map_err(map_extract_archive_error)?; - let value = ext.value().map_err(map_extract_archive_error)?; - - if let Some(meta_key) = key.strip_prefix("minio.metadata.") { - let meta_key = meta_key.strip_prefix("x-amz-meta-").unwrap_or(meta_key); - if !meta_key.is_empty() { - metadata.insert(meta_key.to_string(), value.to_string()); - } - continue; - } - - if key == "minio.versionId" && !value.is_empty() { - opts.version_id = Some(value.to_string()); - } - } - - Ok(()) -} - #[allow(clippy::too_many_arguments)] fn apply_put_request_metadata( metadata: &mut HashMap, @@ -643,7 +330,7 @@ fn apply_put_request_metadata( metadata.insert("content-language".to_string(), content_language.to_string()); } if let Some(content_type) = content_type { - metadata.insert("content-type".to_string(), content_type.to_string()); + metadata.insert(CONTENT_TYPE.to_string(), content_type.to_string()); } if let Some(expires) = expires { let mut formatted = Vec::new(); @@ -831,36 +518,6 @@ fn delete_creates_delete_marker(opts: &ObjectOptions) -> bool { opts.version_id.is_none() && opts.versioned && !opts.version_suspended } -fn resolve_put_object_extract_options(headers: &HeaderMap) -> PutObjectExtractOptions { - let prefix = snowball_meta_value(headers, SNOWBALL_PREFIX_HEADER_KEYS, SNOWBALL_PREFIX_SUFFIX_LOWER) - .and_then(|value| normalize_snowball_prefix(&value)); - let ignore_dirs = snowball_meta_flag(headers, SNOWBALL_IGNORE_DIRS_HEADER_KEYS, SNOWBALL_IGNORE_DIRS_SUFFIX_LOWER); - let ignore_errors = snowball_meta_flag(headers, SNOWBALL_IGNORE_ERRORS_HEADER_KEYS, SNOWBALL_IGNORE_ERRORS_SUFFIX_LOWER); - - PutObjectExtractOptions { - prefix, - ignore_dirs, - ignore_errors, - } -} - -fn is_sse_kms_requested(input: &PutObjectInput, headers: &HeaderMap) -> bool { - input - .server_side_encryption - .as_ref() - .is_some_and(|sse| sse.as_str().eq_ignore_ascii_case(ServerSideEncryption::AWS_KMS)) - || input.ssekms_key_id.is_some() - || headers - .get(AMZ_SERVER_SIDE_ENCRYPTION) - .and_then(|value| value.to_str().ok()) - .is_some_and(|value| value.trim().eq_ignore_ascii_case(ServerSideEncryption::AWS_KMS)) - || headers.contains_key(AMZ_SERVER_SIDE_ENCRYPTION_KMS_ID) -} - -fn is_post_object_sse_kms_requested(input: &PutObjectInput, headers: &HeaderMap) -> bool { - is_sse_kms_requested(input, headers) -} - async fn resolve_put_object_expiration(bucket: &str, obj_info: &ObjectInfo) -> Option { let Ok((lifecycle_config, _)) = metadata_sys::get_lifecycle_config(bucket).await else { debug!("resolve_put_object_expiration: lifecycle config not found for bucket {bucket}"); @@ -929,711 +586,22 @@ impl DefaultObjectUsecase { fn spawn_cache_invalidation(bucket: String, key: String, version_id: Option) { let manager = get_concurrency_manager(); - spawn_traced(async move { + crate::storage::request_context::spawn_traced(async move { manager.invalidate_cache_versioned(&bucket, &key, version_id.as_deref()).await; }); } - fn build_cached_get_object_output(cached: &CachedGetObject) -> GetObjectOutput { - let body_data = cached.body.clone(); - let body = Some(StreamingBlob::wrap::<_, Infallible>(futures::stream::once(async move { - Ok((*body_data).clone()) - }))); - - let last_modified = cached - .last_modified - .as_ref() - .and_then(|s| match OffsetDateTime::parse(s, &Rfc3339) { - Ok(dt) => Some(Timestamp::from(dt)), - Err(e) => { - warn!("Failed to parse cached last_modified '{}': {}", s, e); - None - } - }); - - let content_type = cached.content_type.as_ref().and_then(|ct| ContentType::from_str(ct).ok()); - - GetObjectOutput { - body, - content_length: Some(cached.content_length), - accept_ranges: Some("bytes".to_string()), - e_tag: cached.e_tag.as_ref().map(|etag| to_s3s_etag(etag)), - last_modified, - content_type, - cache_control: cached.cache_control.clone(), - content_disposition: cached.content_disposition.clone(), - content_encoding: cached.content_encoding.clone(), - content_language: cached.content_language.clone(), - version_id: cached.version_id.clone(), - delete_marker: Some(cached.delete_marker), - tag_count: cached.tag_count, - metadata: if cached.user_metadata.is_empty() { - None - } else { - Some(cached.user_metadata.clone()) - }, - ..Default::default() - } - } - - fn build_cached_get_object_event_info(bucket: &str, key: &str, cached: &CachedGetObject) -> ObjectInfo { - ObjectInfo { - bucket: bucket.to_string(), - name: key.to_string(), - storage_class: cached.storage_class.clone(), - mod_time: cached - .last_modified - .as_ref() - .and_then(|s| OffsetDateTime::parse(s, &Rfc3339).ok()), - size: cached.content_length, - actual_size: cached.content_length, - is_dir: false, - user_defined: cached.user_metadata.clone(), - version_id: cached.version_id.as_ref().and_then(|v| Uuid::parse_str(v).ok()), - delete_marker: cached.delete_marker, - content_type: cached.content_type.clone(), - content_encoding: cached.content_encoding.clone(), - etag: cached.e_tag.clone(), - ..Default::default() - } - } - - fn build_memory_blob(buf: Vec, response_content_length: i64, optimal_buffer_size: usize) -> Option { - let mem_reader = InMemoryAsyncReader::new(buf); - Some(StreamingBlob::wrap(bytes_stream( - ReaderStream::with_capacity(Box::new(mem_reader), optimal_buffer_size), - response_content_length as usize, - ))) - } - - fn build_reader_blob(reader: R, response_content_length: i64, optimal_buffer_size: usize) -> Option - where - R: AsyncRead + Send + Sync + 'static, - { - Some(StreamingBlob::wrap(bytes_stream( - ReaderStream::with_capacity(reader, optimal_buffer_size), - response_content_length as usize, - ))) - } - - fn init_get_object_bootstrap(bucket: &str, key: &str, request_id: &str) -> S3Result { - let timeout_config = TimeoutConfig::from_env(); - let wrapper = RequestTimeoutWrapper::with_request_id(timeout_config.clone(), request_id.to_string()); - let request_start = std::time::Instant::now(); - let request_guard = ConcurrencyManager::track_request(); - let concurrent_requests = GetObjectGuard::concurrent_requests(); - - let deadlock_detector = deadlock_detector::get_deadlock_detector(); - let request_id = wrapper.request_id().to_string(); - deadlock_detector.register_request(&request_id, format!("GetObject {bucket}/{key}")); - let deadlock_request_guard = DeadlockRequestGuard::new(deadlock_detector, request_id); - - if wrapper.is_timeout() { - warn!( - bucket = %bucket, - key = %key, - timeout_secs = timeout_config.get_object_timeout.as_secs(), - elapsed_ms = wrapper.elapsed().as_millis(), - "GetObject request timed out before processing" - ); - return Err(s3_error!(InternalError, "Request timeout before processing")); - } - - rustfs_io_metrics::record_get_object_request_start(concurrent_requests); - - debug!( - "GetObject request started with {} concurrent requests, timeout={:?}", - concurrent_requests, timeout_config.get_object_timeout - ); - - Ok(GetObjectBootstrap { - timeout_config, - wrapper, - request_start, - request_guard, - _deadlock_request_guard: deadlock_request_guard, - concurrent_requests, - }) - } - - async fn acquire_get_object_io_planning<'a>( - manager: &'a ConcurrencyManager, - wrapper: &RequestTimeoutWrapper, - timeout_config: &TimeoutConfig, - bucket: &str, - key: &str, - ) -> S3Result> { - let permit_wait_start = std::time::Instant::now(); - let disk_permit = manager - .acquire_disk_read_permit() - .await - .map_err(|_| s3_error!(InternalError, "disk read semaphore closed"))?; - let permit_wait_duration = permit_wait_start.elapsed(); - - if wrapper.is_timeout() { - warn!( - bucket = %bucket, - key = %key, - wait_ms = permit_wait_duration.as_millis(), - timeout_secs = timeout_config.get_object_timeout.as_secs(), - elapsed_ms = wrapper.elapsed().as_millis(), - "GetObject request timed out while waiting for disk permit" - ); - - rustfs_io_metrics::record_get_object_timeout(Some("disk_permit"), Some(wrapper.elapsed().as_secs_f64())); - return Err(s3_error!(InternalError, "Request timeout while waiting for disk permit")); - } - - let queue_status = manager.io_queue_status(); - let queue_snapshot = GetObjectQueueSnapshot::from_available_permits( - queue_status.total_permits, - queue_status.total_permits.saturating_sub(queue_status.permits_in_use), - ); - let queue_utilization = queue_snapshot.utilization_percent(); - - if queue_snapshot.is_congested(80.0) { - warn!( - bucket = %bucket, - key = %key, - queue_utilization = format!("{:.1}%", queue_utilization), - permits_in_use = queue_status.permits_in_use, - total_permits = queue_status.total_permits, - "I/O queue congestion detected" - ); - - rustfs_io_metrics::record_io_queue_congestion(); - } - - if wrapper.is_timeout() { - warn!( - bucket = %bucket, - key = %key, - timeout_secs = timeout_config.get_object_timeout.as_secs(), - elapsed_ms = wrapper.elapsed().as_millis(), - "GetObject request timed out before reading object" - ); - rustfs_io_metrics::record_get_object_timeout(Some("before_read"), Some(wrapper.elapsed().as_secs_f64())); - return Err(s3_error!(InternalError, "Request timeout before reading object")); - } - - Ok(GetObjectIoPlanning { - _disk_permit: disk_permit, - permit_wait_duration, - queue_status, - queue_utilization, - }) - } - - async fn prepare_get_object_request_context(req: &S3Request) -> S3Result { - let GetObjectInput { - bucket, - key, - version_id, - part_number, - range, - .. - } = req.input.clone(); - - validate_object_key(&key, "GET")?; - - let part_number = part_number.map(|v| v as usize); - - if let Some(part_num) = part_number - && part_num == 0 - { - return Err(s3_error!(InvalidArgument, "Invalid part number: part number must be greater than 0")); - } - - let rs = range.map(|v| match v { - Range::Int { first, last } => HTTPRangeSpec { - is_suffix_length: false, - start: first as i64, - end: if let Some(last) = last { last as i64 } else { -1 }, - }, - Range::Suffix { length } => HTTPRangeSpec { - is_suffix_length: true, - start: length as i64, - end: -1, - }, - }); - - if rs.is_some() && part_number.is_some() { - return Err(s3_error!(InvalidArgument, "range and part_number invalid")); - } - - let opts: ObjectOptions = get_opts(&bucket, &key, version_id.clone(), part_number, &req.headers) - .await - .map_err(ApiError::from)?; - - Ok(GetObjectRequestContext { - cache_key: ConcurrencyManager::make_cache_key(&bucket, &key, version_id.as_deref()), - version_id_for_event: version_id.unwrap_or_default(), - bucket, - key, - part_number, - rs, - opts, - }) - } - #[allow(clippy::too_many_arguments)] - async fn prepare_get_object_read_execution<'a>( - req: &S3Request, - manager: &'a ConcurrencyManager, - wrapper: &RequestTimeoutWrapper, - timeout_config: &TimeoutConfig, - bucket: &str, - key: &str, - rs: Option, - opts: &ObjectOptions, - part_number: Option, - ) -> S3Result> { - let h = HeaderMap::new(); - let io_planning = Self::acquire_get_object_io_planning(manager, wrapper, timeout_config, bucket, key).await?; - let store = get_validated_store(bucket).await?; - - let read_start = std::time::Instant::now(); - let read_setup = - Self::prepare_get_object_read(req, &store, manager, bucket, key, rs, h, opts, part_number, read_start).await?; - - Ok(GetObjectPreparedRead { io_planning, read_setup }) - } - - #[allow(clippy::too_many_arguments)] - async fn prepare_get_object_read( - req: &S3Request, - store: &rustfs_ecstore::store::ECStore, - manager: &ConcurrencyManager, - bucket: &str, - key: &str, - mut rs: Option, - h: HeaderMap, - opts: &ObjectOptions, - part_number: Option, - read_start: std::time::Instant, - ) -> S3Result { - let reader = store - .get_object_reader(bucket, key, rs.clone(), h, opts) - .await - .map_err(ApiError::from)?; - - let info = reader.object_info; - - use rustfs_io_metrics::{record_memory_copy_saved, record_zero_copy_read}; - let read_duration = read_start.elapsed(); - let estimated_saved = (info.size * 2) as usize; - record_zero_copy_read(info.size as usize, read_duration.as_secs_f64() * 1000.0); - record_memory_copy_saved(estimated_saved); - - manager.record_disk_operation(info.size as u64, read_duration, true).await; - - check_preconditions(&req.headers, &info)?; - - debug!(object_size = info.size, part_count = info.parts.len(), "GET object metadata snapshot"); - for part in &info.parts { - debug!( - part_number = part.number, - part_size = part.size, - part_actual_size = part.actual_size, - "GET object part details" - ); - } - - let event_info = info.clone(); - let content_type = if let Some(content_type) = &info.content_type { - match ContentType::from_str(content_type) { - Ok(res) => Some(res), - Err(err) => { - error!("parse content-type err {} {:?}", content_type, err); - None - } - } - } else { - None - }; - let last_modified = info.mod_time.map(Timestamp::from); - - if let Some(part_number) = part_number - && rs.is_none() - { - rs = HTTPRangeSpec::from_object_info(&info, part_number); - } - - validate_sse_headers_for_read(&info.user_defined, &req.headers)?; - - let mut content_length = info.get_actual_size().map_err(ApiError::from)?; - let content_range = if let Some(rs) = &rs { - let total_size = content_length; - let (start, length) = rs.get_offset_length(total_size).map_err(ApiError::from)?; - content_length = length; - Some(format!("bytes {}-{}/{}", start, start as i64 + length - 1, total_size)) - } else { - None - }; - - debug!( - "GET object metadata check: parts={}, provided_sse_key={:?}", - info.parts.len(), - req.input.sse_customer_key.is_some() - ); - - let decryption_request = DecryptionRequest { - bucket, - key, - metadata: &info.user_defined, - sse_customer_key: req.input.sse_customer_key.as_ref(), - sse_customer_key_md5: req.input.sse_customer_key_md5.as_ref(), - part_number: None, - parts: &info.parts, - etag: info.etag.as_deref(), - }; - - let mut response_content_length = content_length; - let encrypted_stream = reader.stream; - - let ( - server_side_encryption, - sse_customer_algorithm, - sse_customer_key_md5, - ssekms_key_id, - encryption_applied, - final_stream, - ) = match sse_decryption(decryption_request).await? { - Some(material) => { - let server_side_encryption = Some(material.server_side_encryption.clone()); - let sse_customer_algorithm = Some(material.algorithm.clone()); - let sse_customer_key_md5 = material.customer_key_md5.clone(); - let ssekms_key_id = material.kms_key_id.clone(); - - let (decrypted_stream, plaintext_size) = material - .wrap_reader(encrypted_stream, content_length) - .await - .map_err(ApiError::from)?; - - response_content_length = plaintext_size; - - ( - server_side_encryption, - sse_customer_algorithm, - sse_customer_key_md5, - ssekms_key_id, - true, - decrypted_stream, - ) - } - None => (None, None, None, None, false, wrap_reader(encrypted_stream)), - }; - - Ok(GetObjectReadSetup { - info, - event_info, - final_stream, - rs, - content_type, - last_modified, - response_content_length, - content_range, - server_side_encryption, - sse_customer_algorithm, - sse_customer_key_md5, - ssekms_key_id, - encryption_applied, - }) - } - #[allow(clippy::too_many_arguments)] - fn finalize_get_object_strategy( - &self, - manager: &ConcurrencyManager, - bucket: &str, - key: &str, - info: &ObjectInfo, - rs: Option<&HTTPRangeSpec>, - response_content_length: i64, - permit_wait_duration: Duration, - queue_utilization: f64, - queue_status: &concurrency::IoQueueStatus, - concurrent_requests: usize, - ) -> GetObjectStrategyContext { - let base_buffer_size = self.base_buffer_size(); - - let is_sequential_hint = if rs.is_none() { - true - } else if let Some(range_spec) = rs { - range_spec.start == 0 && !range_spec.is_suffix_length - } else { - false - }; - - if let Some(range_spec) = rs - && range_spec.start >= 0 - { - manager.record_access(range_spec.start as u64, response_content_length as u64); - } - - if response_content_length > 0 { - manager.record_transfer(response_content_length as u64, permit_wait_duration); - } - - let io_strategy = - manager.calculate_io_strategy_with_context(info.size, base_buffer_size, permit_wait_duration, is_sequential_hint); - - debug!( - wait_ms = permit_wait_duration.as_millis() as u64, - load_level = ?io_strategy.load_level, - buffer_size = io_strategy.buffer_size, - buffer_multiplier = io_strategy.buffer_multiplier, - readahead = io_strategy.enable_readahead, - cache_wb = io_strategy.cache_writeback_enabled, - storage_media = ?io_strategy.storage_media, - access_pattern = ?io_strategy.access_pattern, - bandwidth_tier = ?io_strategy.bandwidth_tier, - concurrent_requests = io_strategy.concurrent_requests, - file_size = info.size, - is_sequential = is_sequential_hint, - "Enhanced multi-factor I/O strategy calculated" - ); - - let io_priority = manager.get_io_priority(response_content_length); - - if manager.is_priority_scheduling_enabled() { - debug!( - bucket = %bucket, - key = %key, - priority = %io_priority, - request_size = response_content_length, - "I/O priority assigned (based on actual request size)" - ); - - rustfs_io_metrics::record_io_priority_assignment(io_priority.as_str()); - } - - rustfs_io_metrics::record_get_object_io_state( - permit_wait_duration.as_secs_f64(), - queue_utilization, - queue_status.permits_in_use, - queue_status.total_permits.saturating_sub(queue_status.permits_in_use), - io_strategy.load_level.as_str(), - io_strategy.buffer_multiplier, - ); - rustfs_io_metrics::record_io_priority_assignment(io_priority.as_str()); - - debug!( - actual_request_size = response_content_length, - priority = %io_priority.as_str(), - "I/O priority finalized with actual request size" - ); - - let base_buffer_size = get_buffer_size_opt_in(response_content_length); - let optimal_buffer_size = if io_strategy.buffer_size > 0 { - io_strategy.buffer_size.min(base_buffer_size) - } else { - get_concurrency_aware_buffer_size(response_content_length, base_buffer_size) - }; - - debug!( - "GetObject buffer sizing: file_size={}, base={}, optimal={}, concurrent_requests={}, io_strategy={:?}", - response_content_length, base_buffer_size, optimal_buffer_size, concurrent_requests, io_strategy.load_level - ); - - GetObjectStrategyContext { - io_strategy, - optimal_buffer_size, - } - } - - fn build_get_object_checksums( - info: &ObjectInfo, - headers: &HeaderMap, - part_number: Option, - rs: Option<&HTTPRangeSpec>, - ) -> S3Result { - let mut checksums = GetObjectChecksums::default(); - - if let Some(checksum_mode) = headers.get(AMZ_CHECKSUM_MODE) - && checksum_mode.to_str().unwrap_or_default() == "ENABLED" - && rs.is_none() - { - let (decrypted_checksums, _is_multipart) = - info.decrypt_checksums(part_number.unwrap_or(0), headers).map_err(|e| { - error!("decrypt_checksums error: {}", e); - ApiError::from(e) - })?; - - for (key, checksum) in decrypted_checksums { - if key == AMZ_CHECKSUM_TYPE { - checksums.checksum_type = Some(ChecksumType::from(checksum)); - continue; - } - - match rustfs_rio::ChecksumType::from_string(key.as_str()) { - rustfs_rio::ChecksumType::CRC32 => checksums.crc32 = Some(checksum), - rustfs_rio::ChecksumType::CRC32C => checksums.crc32c = Some(checksum), - rustfs_rio::ChecksumType::SHA1 => checksums.sha1 = Some(checksum), - rustfs_rio::ChecksumType::SHA256 => checksums.sha256 = Some(checksum), - rustfs_rio::ChecksumType::CRC64_NVME => checksums.crc64nvme = Some(checksum), - _ => (), - } - } - } - - Ok(checksums) - } - #[allow(clippy::too_many_arguments)] - async fn build_get_object_body( - mut final_stream: R, - info: &ObjectInfo, - cache_key: &str, - response_content_length: i64, - optimal_buffer_size: usize, - part_number: Option, - has_range: bool, - encryption_applied: bool, - cache_writeback_enabled: bool, - ) -> S3Result> - where - R: AsyncRead + Send + Sync + Unpin + 'static, - { - let manager = get_concurrency_manager(); - let cache_eligibility = manager.get_object_cache_eligibility( - cache_writeback_enabled, - part_number.is_some(), - has_range, - encryption_applied, - response_content_length, - ); - let should_cache = cache_eligibility.should_cache(); - - let body = if should_cache { - debug!( - "Reading object into memory for caching: key={} size={}", - cache_key, response_content_length - ); - - let mut buf = Vec::with_capacity(response_content_length as usize); - if let Err(e) = tokio::io::AsyncReadExt::read_to_end(&mut final_stream, &mut buf).await { - error!("Failed to read object into memory for caching: {}", e); - return Err(ApiError::from(StorageError::other(format!("Failed to read object for caching: {e}"))).into()); - } - - if buf.len() != response_content_length as usize { - warn!( - "Object size mismatch during cache read: expected={} actual={}", - response_content_length, - buf.len() - ); - } - - let last_modified_str = info.mod_time.and_then(|t| match t.format(&Rfc3339) { - Ok(s) => Some(s), - Err(e) => { - warn!("Failed to format last_modified for cache writeback: {}", e); - None - } - }); - - let cached_response = CachedGetObject::new(Bytes::from(buf.clone()), response_content_length) - .with_content_type(info.content_type.clone().unwrap_or_default()) - .with_e_tag(info.etag.clone().unwrap_or_default()) - .with_last_modified(last_modified_str.unwrap_or_default()); - - let cache_key_clone = cache_key.to_string(); - spawn_traced(async move { - let manager = get_concurrency_manager(); - manager.put_cached_object(cache_key_clone.clone(), cached_response).await; - debug!("Object cached successfully with metadata: {}", cache_key_clone); - }); - - rustfs_io_metrics::record_object_cache_writeback(); - Self::build_memory_blob(buf, response_content_length, optimal_buffer_size) - } else if encryption_applied { - let seekable_object_size_threshold = rustfs_config::DEFAULT_OBJECT_SEEK_SUPPORT_THRESHOLD; - let should_buffer_encrypted_object = response_content_length > 0 - && response_content_length <= seekable_object_size_threshold as i64 - && part_number.is_none() - && !has_range; - - if should_buffer_encrypted_object { - let mut buf = Vec::with_capacity(response_content_length as usize); - if let Err(e) = tokio::io::AsyncReadExt::read_to_end(&mut final_stream, &mut buf).await { - error!("Failed to read decrypted object into memory: {}", e); - return Err(ApiError::from(StorageError::other(format!("Failed to read decrypted object: {e}"))).into()); - } - - if buf.len() != response_content_length as usize { - warn!( - "Encrypted object size mismatch during read: expected={} actual={}", - response_content_length, - buf.len() - ); - } - - Self::build_memory_blob(buf, response_content_length, optimal_buffer_size) - } else { - info!( - "Encrypted object: Using unlimited stream for decryption with buffer size {}", - optimal_buffer_size - ); - Self::build_reader_blob(final_stream, response_content_length, optimal_buffer_size) - } - } else { - let seekable_object_size_threshold = rustfs_config::DEFAULT_OBJECT_SEEK_SUPPORT_THRESHOLD; - - let should_provide_seek_support = response_content_length > 0 - && response_content_length <= seekable_object_size_threshold as i64 - && part_number.is_none() - && !has_range; - - if should_provide_seek_support { - debug!( - "Reading small object into memory for seek support: key={} size={}", - cache_key, response_content_length - ); - - let mut buf = Vec::with_capacity(response_content_length as usize); - match tokio::io::AsyncReadExt::read_to_end(&mut final_stream, &mut buf).await { - Ok(_) => { - if buf.len() != response_content_length as usize { - warn!( - "Object size mismatch during seek support read: expected={} actual={}", - response_content_length, - buf.len() - ); - } - - Self::build_memory_blob(buf, response_content_length, optimal_buffer_size) - } - Err(e) => { - error!("Failed to read object into memory for seek support: {}", e); - Self::build_reader_blob(final_stream, response_content_length, optimal_buffer_size) - } - } - } else { - Self::build_reader_blob(final_stream, response_content_length, optimal_buffer_size) - } - }; - - Ok(body) - } - - fn put_object_execution_context(req: &S3Request) -> (EventName, QuotaOperation, &'static str) { - if req.extensions.get::().is_some() { - (EventName::ObjectCreatedPost, QuotaOperation::PostObject, "POST") - } else { - (EventName::ObjectCreatedPut, QuotaOperation::PutObject, "PUT") - } - } - #[instrument(level = "debug", skip(self, _fs, req))] pub async fn execute_put_object(&self, _fs: &FS, req: S3Request) -> S3Result> { - let start_time = std::time::Instant::now(); - let mut req = req; - if let Some(context) = &self.context { let _ = context.object_store(); } - let (event_name, quota_operation, request_method_name) = Self::put_object_execution_context(&req); - if req.extensions.get::().is_some() && is_post_object_sse_kms_requested(&req.input, &req.headers) - { + let request_context = prepare_put_object_request_context(&req); + let (event_name, quota_operation, request_method_name) = put_object_execution_context(&req); + let helper = new_operation_helper(&req, event_name, S3Operation::PutObject, false); + + if request_context.is_post_object && is_post_object_sse_kms_requested(&req.input, &request_context.headers) { return Err(s3_error!(NotImplemented, "SSE-KMS is not supported for POST object uploads")); } if let Some(ref storage_class) = req.input.storage_class @@ -1641,380 +609,19 @@ impl DefaultObjectUsecase { { return Err(s3_error!(InvalidStorageClass)); } - if is_put_object_extract_requested(&req.headers) { + if is_put_object_extract_requested(&request_context.headers) { return self.execute_put_object_extract(req).await; } - let input = std::mem::take(&mut req.input); + let resolved_size = resolve_put_body_size(req.input.content_length, &request_context.headers)?; + self.check_bucket_quota(&req.input.bucket, quota_operation, resolved_size as u64) + .await?; - let PutObjectInput { - body, - bucket, - cache_control, - key, - content_length, - content_disposition, - content_encoding, - content_language, - content_type, - expires, - tagging, - metadata, - version_id, - server_side_encryption, - sse_customer_algorithm, - sse_customer_key, - sse_customer_key_md5, - ssekms_key_id, - content_md5, - object_lock_legal_hold_status, - object_lock_mode, - object_lock_retain_until_date, - storage_class, - website_redirect_location, - .. - } = input; - - // Merge SSE-C params from headers (fallback when S3 layer does not populate input) - let (h_algo, h_key, h_md5) = extract_ssec_params_from_headers(&req.headers)?; - let sse_customer_algorithm = sse_customer_algorithm.or(h_algo); - let sse_customer_key = sse_customer_key.or(h_key); - let sse_customer_key_md5 = sse_customer_key_md5.or(h_md5); - - // Merge server_side_encryption from headers (fallback when S3 layer does not populate input) - let server_side_encryption = server_side_encryption.or(extract_server_side_encryption_from_headers(&req.headers)?); - - // Validate object key - validate_object_key(&key, request_method_name)?; - - if let Some(size) = content_length { - self.check_bucket_quota(&bucket, quota_operation, size as u64).await?; - } - - let Some(body) = body else { return Err(s3_error!(IncompleteBody)) }; - - let mut size = match content_length { - Some(c) => c, - None => { - if let Some(val) = req.headers.get(AMZ_DECODED_CONTENT_LENGTH) { - match atoi::atoi::(val.as_bytes()) { - Some(x) => x, - None => return Err(s3_error!(UnexpectedContent)), - } - } else { - return Err(s3_error!(UnexpectedContent)); - } - } - }; - - if size == -1 { - return Err(s3_error!(UnexpectedContent)); - } - - // Apply adaptive buffer sizing based on file size for optimal streaming performance. - // Uses workload profile configuration (enabled by default) to select appropriate buffer size. - // Buffer sizes range from 32KB to 4MB depending on file size and configured workload profile. - let buffer_size = get_buffer_size_opt_in(size); - - // Detect zero-copy opportunity before encryption/compression decisions - // Zero-copy is beneficial for large unencrypted, uncompressed objects - let enable_zero_copy = should_use_zero_copy(size, &req.headers); - - if enable_zero_copy { - // Record zero-copy write attempt - counter!("rustfs.zero_copy.write.attempts.total").increment(1); - histogram!("rustfs.zero_copy.write.size.bytes").record(size as f64); - debug!("Zero-copy write enabled for {} byte object (bucket={}, key={})", size, bucket, key); - } - - let body = tokio::io::BufReader::with_capacity( - buffer_size, - StreamReader::new(body.map(|f| f.map_err(|e| std::io::Error::other(e.to_string())))), - ); - - let store = get_validated_store(&bucket).await?; - - // TDD: Get bucket default encryption configuration - let bucket_sse_config = metadata_sys::get_sse_config(&bucket).await.ok(); - debug!("TDD: bucket_sse_config={:?}", bucket_sse_config); - - // TDD: Determine effective encryption configuration (request overrides bucket default) - let original_sse = server_side_encryption.clone(); - let mut effective_sse = server_side_encryption.or_else(|| { - bucket_sse_config.as_ref().and_then(|(config, _timestamp)| { - debug!("TDD: Processing bucket SSE config: {:?}", config); - config.rules.first().and_then(|rule| { - debug!("TDD: Processing SSE rule: {:?}", rule); - rule.apply_server_side_encryption_by_default.as_ref().map(|sse| { - debug!("TDD: Found SSE default: {:?}", sse); - match sse.sse_algorithm.as_str() { - "AES256" => ServerSideEncryption::from_static(ServerSideEncryption::AES256), - "aws:kms" => ServerSideEncryption::from_static(ServerSideEncryption::AWS_KMS), - _ => ServerSideEncryption::from_static(ServerSideEncryption::AES256), // fallback to AES256 - } - }) - }) - }) - }); - debug!("TDD: effective_sse={:?} (original={:?})", effective_sse, original_sse); - - let mut effective_kms_key_id = ssekms_key_id.or_else(|| { - bucket_sse_config.as_ref().and_then(|(config, _timestamp)| { - config.rules.first().and_then(|rule| { - rule.apply_server_side_encryption_by_default - .as_ref() - .and_then(|sse| sse.kms_master_key_id.clone()) - }) - }) - }); - - // Validate SSE-C headers early: reject partial/invalid combinations per S3 spec - validate_sse_headers_for_write( - effective_sse.as_ref(), - effective_kms_key_id.as_ref(), - sse_customer_algorithm.as_ref(), - sse_customer_key.as_ref(), - sse_customer_key_md5.as_ref(), - true, // PutObject requires all three: algorithm, key, key_md5 - )?; - - let mut metadata = metadata.unwrap_or_default(); - apply_put_request_metadata( - &mut metadata, - &req.headers, - &key, - cache_control, - content_disposition, - content_encoding, - content_language, - content_type, - expires, - website_redirect_location, - tagging, - storage_class.clone(), - )?; - - let mut opts: ObjectOptions = put_opts(&bucket, &key, version_id.clone(), &req.headers, metadata.clone()) - .await - .map_err(ApiError::from)?; - apply_put_request_object_lock_opts( - &bucket, - object_lock_legal_hold_status, - object_lock_mode, - object_lock_retain_until_date, - &mut opts, - ) - .await?; - - let current_opts: ObjectOptions = get_opts(&bucket, &key, version_id.clone(), None, &req.headers) - .await - .map_err(ApiError::from)?; - match store.get_object_info(&bucket, &key, ¤t_opts).await { - Ok(existing_obj_info) => validate_existing_object_lock_for_write(&existing_obj_info)?, - Err(err) => { - if !is_err_object_not_found(&err) && !is_err_version_not_found(&err) { - return Err(ApiError::from(err).into()); - } - } - } - - let actual_size = size; - - let mut md5hex = if let Some(base64_md5) = content_md5 { - let md5 = base64_simd::STANDARD - .decode_to_vec(base64_md5.as_bytes()) - .map_err(|e| ApiError::from(StorageError::other(format!("Invalid content MD5: {e}"))))?; - Some(hex_simd::encode_to_string(&md5, hex_simd::AsciiCase::Lower)) - } else { - None - }; - - let mut sha256hex = get_content_sha256_with_query(&req.headers, req.uri.query()); - - 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::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()); - } - - opts.want_checksum = hrd.checksum(); - insert_str(&mut opts.user_defined, SUFFIX_COMPRESSION, algorithm.to_string()); - insert_str(&mut opts.user_defined, SUFFIX_ACTUAL_SIZE, size.to_string()); - - size = HashReader::SIZE_PRESERVE_LAYER; - 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) { - return Err(ApiError::from(err).into()); - } - - opts.want_checksum = reader.checksum(); - } - - let mut helper = OperationHelper::new(&req, event_name, S3Operation::PutObject); - - // Apply encryption using unified SSE API. - let encryption_request = EncryptionRequest { - bucket: &bucket, - key: &key, - server_side_encryption: effective_sse.clone(), - ssekms_key_id: effective_kms_key_id.clone(), - sse_customer_algorithm: sse_customer_algorithm.clone(), - sse_customer_key, - sse_customer_key_md5: sse_customer_key_md5.clone(), - content_size: actual_size, - part_number: None, - part_key: None, - part_nonce: None, - }; - - let encryption_material = match sse_encryption(encryption_request).await { - Ok(material) => material, - Err(err) => { - let result = Err(err.into()); - let _ = helper.complete(&result); - return result; - } - }; - - if let Some(material) = encryption_material { - effective_sse = Some(material.server_side_encryption.clone()); - effective_kms_key_id = material.kms_key_id.clone(); - - let encrypted_reader = material.wrap_reader(reader); - reader = HashReader::from_reader(encrypted_reader, HashReader::SIZE_PRESERVE_LAYER, actual_size, None, None, false) - .map_err(ApiError::from)?; - - let encryption_metadata = material.metadata; - metadata.extend(encryption_metadata.clone()); - opts.user_defined.extend(encryption_metadata); - } - - let mut reader = PutObjReader::new(reader); - - let mt2 = metadata.clone(); - opts.user_defined.extend(metadata); - - let repoptions = - get_must_replicate_options(&mt2, "".to_string(), ReplicationStatusType::Empty, ReplicationType::Object, opts.clone()); - - let dsc = must_replicate(&bucket, &key, repoptions).await; - - if dsc.replicate_any() { - insert_str(&mut opts.user_defined, SUFFIX_REPLICATION_TIMESTAMP, jiff::Zoned::now().to_string()); - insert_str( - &mut opts.user_defined, - SUFFIX_REPLICATION_STATUS, - dsc.pending_status().unwrap_or_default(), - ); - } - - let obj_info = match store - .put_object(&bucket, &key, &mut reader, &opts) - .await - .map_err(ApiError::from) - { - Ok(obj_info) => obj_info, - Err(err) => { - let result: S3Result> = Err(err.into()); - let _ = helper.complete(&result); - return result; - } - }; - - maybe_enqueue_transition_immediate(&obj_info, LcEventSrc::S3PutObject).await; - - // Fast in-memory update for immediate quota consistency - rustfs_ecstore::data_usage::increment_bucket_usage_memory(&bucket, obj_info.size as u64).await; - - let raw_version = obj_info.version_id.map(|v| v.to_string()); - - helper = helper.object(obj_info.clone()); - if let Some(version_id) = &raw_version { - helper = helper.version_id(version_id.clone()); - } - - Self::spawn_cache_invalidation(bucket.clone(), key.clone(), raw_version.clone()); - - let put_version = if BucketVersioningSys::prefix_enabled(&bucket, &key).await { - raw_version - } else { - None - }; - - let e_tag = obj_info.etag.clone().map(|etag| to_s3s_etag(&etag)); - - let repoptions = - get_must_replicate_options(&mt2, "".to_string(), ReplicationStatusType::Empty, ReplicationType::Object, opts); - - let dsc = must_replicate(&bucket, &key, repoptions).await; - let expiration = resolve_put_object_expiration(&bucket, &obj_info).await; - - if dsc.replicate_any() { - schedule_replication(obj_info.clone(), store, dsc, ReplicationType::Object).await; - } - - let mut checksums = PutObjectChecksums { - crc32: input.checksum_crc32, - crc32c: input.checksum_crc32c, - sha1: input.checksum_sha1, - sha256: input.checksum_sha256, - crc64nvme: input.checksum_crc64nvme, - }; - apply_trailing_checksums( - input.checksum_algorithm.as_ref().map(|a| a.as_str()), - &req.trailing_headers, - &mut checksums, - ); - - let output = PutObjectOutput { - e_tag, - server_side_encryption: effective_sse, - sse_customer_algorithm: sse_customer_algorithm.clone(), - sse_customer_key_md5: sse_customer_key_md5.clone(), - ssekms_key_id: effective_kms_key_id, - expiration, - checksum_crc32: checksums.crc32, - checksum_crc32c: checksums.crc32c, - checksum_sha1: checksums.sha1, - checksum_sha256: checksums.sha256, - checksum_crc64nvme: checksums.crc64nvme, - version_id: put_version, - ..Default::default() - }; - - // For browser-based POST uploads (multipart/form-data), response status/body handling - // is decided by s3s PostObject serializer (success_action_status / redirect semantics). - - let result = Ok(S3Response::new(output)); - let _ = helper.complete(&result); - - // Record write operation for capacity management (inline to avoid per-request tokio::spawn overhead) - let manager = get_capacity_manager(); - manager.record_write_operation().await; - - // Record PutObject metrics via zero-copy-metrics - { - let duration_ms = start_time.elapsed().as_millis() as f64; - rustfs_io_metrics::record_put_object( - duration_ms, - size, - enable_zero_copy, // Track if zero-copy was enabled - ); - } - - result + let input = req.input; + let flow_result = + DefaultObjectUsecase::run_put_object_flow(input, request_context, request_method_name, resolved_size).await?; + let helper = bind_helper_object(helper, flow_result.helper_object, flow_result.helper_version_id); + complete_put_response(helper, flow_result.output) } pub async fn execute_put_object_acl(&self, req: S3Request) -> S3Result> { @@ -2370,7 +977,7 @@ impl DefaultObjectUsecase { let cache_key = ConcurrencyManager::make_cache_key(&bucket, &object, version_id.clone().as_deref()); let cache_bucket = bucket.clone(); let cache_object = object.clone(); - spawn_traced(async move { + crate::storage::request_context::spawn_traced(async move { manager .invalidate_cache_versioned(&cache_bucket, &cache_object, version_id.as_deref()) .await; @@ -2405,191 +1012,6 @@ impl DefaultObjectUsecase { result } - async fn maybe_get_cached_get_object( - manager: &ConcurrencyManager, - bucket: &str, - key: &str, - cache_key: &str, - part_number: Option, - rs: Option<&HTTPRangeSpec>, - request_start: std::time::Instant, - ) -> Option { - if !manager.is_cache_enabled() || part_number.is_some() || rs.is_some() { - return None; - } - - let cached = manager.get_cached_object(cache_key).await?; - let cache_serve_duration = request_start.elapsed(); - - debug!("Serving object from response cache: {} (latency: {:?})", cache_key, cache_serve_duration); - - rustfs_io_metrics::record_get_object_cache_served(cache_serve_duration.as_secs_f64(), cached.body.len()); - - use rustfs_io_metrics::{record_memory_copy_saved, record_zero_copy_read}; - record_zero_copy_read(cached.body.len(), cache_serve_duration.as_secs_f64() * 1000.0); - record_memory_copy_saved(cached.body.len()); - - manager.record_transfer(cached.content_length as u64, Duration::from_micros(1)); - - let output = Self::build_cached_get_object_output(&cached); - let event_info = Self::build_cached_get_object_event_info(bucket, key, &cached); - - rustfs_io_metrics::record_get_object(request_start.elapsed().as_millis() as f64, cached.content_length, true); - - Some(GetObjectCachedHit { output, event_info }) - } - - fn finalize_get_object_completion( - cache_key: &str, - wrapper: &RequestTimeoutWrapper, - timeout_config: &TimeoutConfig, - total_duration: Duration, - response_content_length: i64, - optimal_buffer_size: usize, - ) { - rustfs_io_metrics::record_get_object_completion( - total_duration.as_secs_f64(), - response_content_length, - optimal_buffer_size, - ); - - rustfs_io_metrics::record_get_object(total_duration.as_millis() as f64, response_content_length, false); - - if wrapper.is_timeout() { - warn!( - "GetObject request exceeded timeout: key={} duration={:?} timeout={:?}", - cache_key, - wrapper.elapsed(), - timeout_config.get_object_timeout - ); - rustfs_io_metrics::record_get_object_timeout(None, Some(wrapper.elapsed().as_secs_f64())); - } - - debug!( - "GetObject completed: key={} size={} duration={:?} buffer={}", - cache_key, response_content_length, total_duration, optimal_buffer_size - ); - } - - async fn finalize_get_object_response( - helper: OperationHelper, - bucket: &str, - method: &hyper::Method, - headers: &HeaderMap, - event_info: ObjectInfo, - version_id_for_event: String, - output: GetObjectOutput, - ) -> S3Result> { - let helper = helper.object(event_info).version_id(version_id_for_event); - let response = wrap_response_with_cors(bucket, method, headers, output).await; - let result = Ok(response); - let _ = helper.complete(&result); - result - } - #[allow(clippy::too_many_arguments)] - async fn build_get_object_output_context( - &self, - req: &S3Request, - cache_key: &str, - manager: &ConcurrencyManager, - bucket: &str, - key: &str, - info: ObjectInfo, - event_info: ObjectInfo, - final_stream: DynReader, - rs: Option, - content_type: Option, - last_modified: Option, - response_content_length: i64, - content_range: Option, - server_side_encryption: Option, - sse_customer_algorithm: Option, - sse_customer_key_md5: Option, - ssekms_key_id: Option, - encryption_applied: bool, - permit_wait_duration: Duration, - queue_utilization: f64, - queue_status: &concurrency::IoQueueStatus, - concurrent_requests: usize, - part_number: Option, - versioned: bool, - ) -> S3Result { - let strategy = self.finalize_get_object_strategy( - manager, - bucket, - key, - &info, - rs.as_ref(), - response_content_length, - permit_wait_duration, - queue_utilization, - queue_status, - concurrent_requests, - ); - let GetObjectStrategyContext { - io_strategy, - optimal_buffer_size, - } = strategy; - - let body = Self::build_get_object_body( - final_stream, - &info, - cache_key, - response_content_length, - optimal_buffer_size, - part_number, - rs.is_some(), - encryption_applied, - io_strategy.cache_writeback_enabled, - ) - .await?; - - let checksums = Self::build_get_object_checksums(&info, &req.headers, part_number, rs.as_ref())?; - - let output_version_id = if versioned { - info.version_id.map(|vid| { - if vid == Uuid::nil() { - "null".to_string() - } else { - vid.to_string() - } - }) - } else { - None - }; - - let output = GetObjectOutput { - body, - content_length: Some(response_content_length), - last_modified, - content_type, - content_encoding: info.content_encoding.clone(), - accept_ranges: Some("bytes".to_string()), - content_range, - e_tag: info.etag.map(|etag| to_s3s_etag(&etag)), - metadata: filter_object_metadata(&info.user_defined), - server_side_encryption, - sse_customer_algorithm, - sse_customer_key_md5, - ssekms_key_id, - checksum_crc32: checksums.crc32, - checksum_crc32c: checksums.crc32c, - checksum_sha1: checksums.sha1, - checksum_sha256: checksums.sha256, - checksum_crc64nvme: checksums.crc64nvme, - checksum_type: checksums.checksum_type, - version_id: output_version_id, - ..Default::default() - }; - - Ok(GetObjectOutputContext { - output, - event_info, - response_content_length, - optimal_buffer_size, - }) - } - #[instrument( level = "debug", skip(self, req), @@ -2605,130 +1027,29 @@ impl DefaultObjectUsecase { .get::() .map(|ctx| ctx.request_id.clone()) .unwrap_or_else(|| crate::storage::request_context::RequestContext::fallback().request_id); - let bootstrap = Self::init_get_object_bootstrap(&req.input.bucket, &req.input.key, &request_id)?; - let timeout_config = bootstrap.timeout_config; - let wrapper = bootstrap.wrapper; - let request_start = bootstrap.request_start; - let concurrent_requests = bootstrap.concurrent_requests; - let mut request_guard = bootstrap.request_guard; - - let mut helper = OperationHelper::new(&req, EventName::ObjectAccessedGet, S3Operation::GetObject).suppress_event(); - // mc get 3 - - let request_context = Self::prepare_get_object_request_context(&req).await?; - let GetObjectRequestContext { - bucket, - key, - cache_key, - version_id_for_event, - part_number, - rs, - opts, - } = request_context; - - // Try to get from cache for small, frequently accessed objects + let bootstrap = init_get_object_bootstrap(&req.input.bucket, &req.input.key, &request_id)?; + let request_context = prepare_get_object_request_context(&req).await?; + let base_buffer_size = self.base_buffer_size(); let manager = get_concurrency_manager(); - - if let Some(cached_hit) = - Self::maybe_get_cached_get_object(manager, &bucket, &key, &cache_key, part_number, rs.as_ref(), request_start).await - { - let GetObjectCachedHit { output, event_info } = cached_hit; - helper = helper.object(event_info).version_id(version_id_for_event.clone()); - - let result = Ok(S3Response::new(output)); - let _ = helper.complete(&result); - return result; - } - - let prepared_read = Self::prepare_get_object_read_execution( - &req, + let flow_runtime = GetObjectFlowRuntime { manager, - &wrapper, - &timeout_config, - &bucket, - &key, - rs, - &opts, - part_number, - ) - .await?; - let GetObjectPreparedRead { io_planning, read_setup } = prepared_read; - let permit_wait_duration = io_planning.permit_wait_duration; - let queue_status = io_planning.queue_status; - let queue_utilization = io_planning.queue_utilization; + bootstrap: &bootstrap, + base_buffer_size, + }; + let helper = new_operation_helper(&req, EventName::ObjectAccessedGet, S3Operation::GetObject, true); + let flow_result = get_object_flow::run_get_object_flow(request_context.clone(), flow_runtime).await; - let GetObjectReadSetup { - info, - event_info, - final_stream, - rs, - content_type, - last_modified, - response_content_length, - content_range, - server_side_encryption, - sse_customer_algorithm, - sse_customer_key_md5, - ssekms_key_id, - encryption_applied, - } = read_setup; + let GetObjectBootstrap { + mut request_guard, + _deadlock_request_guard, + .. + } = bootstrap; - let versioned = BucketVersioningSys::prefix_enabled(&bucket, &key).await; - let output_context = self - .build_get_object_output_context( - &req, - &cache_key, - manager, - &bucket, - &key, - info, - event_info, - final_stream, - rs, - content_type, - last_modified, - response_content_length, - content_range, - server_side_encryption, - sse_customer_algorithm, - sse_customer_key_md5, - ssekms_key_id, - encryption_applied, - permit_wait_duration, - queue_utilization, - &queue_status, - concurrent_requests, - part_number, - versioned, - ) - .await?; - let GetObjectOutputContext { - output, - event_info, - response_content_length, - optimal_buffer_size, - } = output_context; + let result = match flow_result { + Ok(flow_result) => complete_get_flow_result(helper, &request_context, flow_result).await, + Err(err) => Err(err), + }; - let total_duration = request_start.elapsed(); - Self::finalize_get_object_completion( - &cache_key, - &wrapper, - &timeout_config, - total_duration, - response_content_length, - optimal_buffer_size, - ); - - let result = Self::finalize_get_object_response( - helper, - &bucket, - &req.method, - &req.headers, - event_info, - version_id_for_event, - output, - ) - .await; if result.is_ok() { request_guard.finish_ok(); } else { @@ -3336,6 +1657,8 @@ 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, @@ -3347,12 +1670,11 @@ impl DefaultObjectUsecase { etag: src_info.etag.as_deref(), }; - 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; + 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; + } } strip_managed_encryption_metadata(&mut src_info.user_defined); @@ -3363,11 +1685,16 @@ impl DefaultObjectUsecase { let mut compress_metadata = HashMap::new(); - let should_compress = is_compressible(&req.headers, &key) && actual_size > MIN_COMPRESSIBLE_SIZE as i64; - - if should_compress { + if is_compressible(&req.headers, &key) && actual_size > MIN_COMPRESSIBLE_SIZE as i64 { 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); @@ -3384,7 +1711,7 @@ impl DefaultObjectUsecase { } if let Some(ct) = content_type { src_info.content_type = Some(ct.clone()); - src_info.user_defined.insert("content-type".to_string(), ct); + src_info.user_defined.insert(CONTENT_TYPE.to_string(), ct); } } @@ -3399,68 +1726,7 @@ impl DefaultObjectUsecase { src_info.user_defined.extend(object_lock_metadata); } - 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 mut reader = HashReader::new(reader, length, actual_size, None, None, false).map_err(ApiError::from)?; let encryption_request = EncryptionRequest { bucket: &bucket, @@ -3481,13 +1747,13 @@ impl DefaultObjectUsecase { effective_kms_key_id = material.kms_key_id.clone(); let encrypted_reader = material.wrap_reader(reader); - reader = HashReader::from_reader(encrypted_reader, HashReader::SIZE_PRESERVE_LAYER, actual_size, None, None, false) + reader = HashReader::new(encrypted_reader, HashReader::SIZE_PRESERVE_LAYER, actual_size, None, None, false) .map_err(ApiError::from)?; src_info.user_defined.extend(material.metadata); } - src_info.put_object_reader = Some(PutObjReader::new(reader)); + src_info.put_object_reader = Some(ChunkNativePutData::new(reader)); // check quota @@ -3717,7 +1983,7 @@ impl DefaultObjectUsecase { let manager = get_concurrency_manager(); let bucket_clone = bucket.clone(); let deleted_objects = dobjs.clone(); - spawn_traced(async move { + crate::storage::request_context::spawn_traced(async move { for dobj in deleted_objects { manager .invalidate_cache_versioned( @@ -3830,7 +2096,7 @@ impl DefaultObjectUsecase { .as_ref() .map(|context| context.notify()) .unwrap_or_else(default_notify_interface); - spawn_background(async move { + crate::storage::helper::spawn_background(async move { for res in delete_results { if let Some(dobj) = res.delete_object { let event_name = if dobj.delete_marker { @@ -4002,7 +2268,6 @@ impl DefaultObjectUsecase { }) .await; } - // Prefix/force-delete returns empty ObjectInfo; still emit bucket notification so webhooks match S3 DELETE. helper = helper .event_name(EventName::ObjectRemovedDelete) .object(ObjectInfo { @@ -4012,7 +2277,6 @@ impl DefaultObjectUsecase { }) .version_id(String::new()); let result = Ok(S3Response::with_status(DeleteObjectOutput::default(), StatusCode::NO_CONTENT)); - // Match non-empty delete path: capacity manager write-op telemetry. let manager = get_capacity_manager(); manager.record_write_operation().await; let _ = helper.complete(&result); @@ -4120,7 +2384,7 @@ impl DefaultObjectUsecase { let version_id_clone = version_id.clone(); let cache_bucket = bucket.clone(); let cache_object = object.clone(); - spawn_traced(async move { + crate::storage::request_context::spawn_traced(async move { manager .invalidate_cache_versioned(&cache_bucket, &cache_object, version_id_clone.as_deref()) .await; @@ -4632,7 +2896,7 @@ impl DefaultObjectUsecase { let rreq_clone = rreq.clone(); let version_id_clone = version_id.clone(); - spawn_traced(async move { + crate::storage::request_context::spawn_traced(async move { let opts = ObjectOptions { transition: TransitionOptions { restore_request: rreq_clone, @@ -4725,7 +2989,7 @@ impl DefaultObjectUsecase { let (tx, rx) = mpsc::channel::>(2); let stream = ReceiverStream::new(rx); - spawn_traced(async move { + crate::storage::request_context::spawn_traced(async move { let _ = tx .send(Ok(SelectObjectContentEvent::Cont(ContinuationEvent::default()))) .await; @@ -4746,419 +3010,22 @@ impl DefaultObjectUsecase { #[instrument(level = "debug", skip(self, req))] pub async fn execute_put_object_extract(&self, req: S3Request) -> S3Result> { - let helper = OperationHelper::new(&req, EventName::ObjectCreatedPut, S3Operation::PutObject).suppress_event(); - let auth_method = req.method.clone(); - let auth_uri = req.uri.clone(); - let auth_headers = req.headers.clone(); - let auth_extensions = req.extensions.clone(); - let auth_credentials = req.credentials.clone(); - let auth_region = req.region.clone(); - let auth_service = req.service.clone(); - let auth_trailing_headers = req.trailing_headers.clone(); - if is_sse_kms_requested(&req.input, &req.headers) { + let request_context = prepare_put_object_request_context(&req); + let helper = new_operation_helper(&req, EventName::ObjectCreatedPut, S3Operation::PutObject, true); + if is_sse_kms_requested(&req.input, &request_context.headers) { return Err(s3_error!(NotImplemented, "SSE-KMS is not supported for extract uploads")); } - let input = req.input; - - let PutObjectInput { - body, - bucket, - key, - version_id, - cache_control, - content_disposition, - content_encoding, - content_length, - content_language, - content_type, - content_md5, - expires, - object_lock_legal_hold_status, - object_lock_mode, - object_lock_retain_until_date, - server_side_encryption, - sse_customer_algorithm, - sse_customer_key, - sse_customer_key_md5, - ssekms_key_id, - storage_class, - tagging, - website_redirect_location, - .. - } = input; - - let event_version_id = version_id; - let (h_algo, h_key, h_md5) = extract_ssec_params_from_headers(&req.headers)?; - let sse_customer_algorithm = sse_customer_algorithm.or(h_algo); - let sse_customer_key = sse_customer_key.or(h_key); - let sse_customer_key_md5 = sse_customer_key_md5.or(h_md5); - - let original_sse = server_side_encryption.or(extract_server_side_encryption_from_headers(&req.headers)?); - let bucket_sse_config = metadata_sys::get_sse_config(&bucket).await.ok(); - let mut effective_sse = original_sse.or_else(|| { - bucket_sse_config.as_ref().and_then(|(config, _timestamp)| { - config.rules.first().and_then(|rule| { - rule.apply_server_side_encryption_by_default - .as_ref() - .map(|sse| match sse.sse_algorithm.as_str() { - "AES256" => ServerSideEncryption::from_static(ServerSideEncryption::AES256), - "aws:kms" => ServerSideEncryption::from_static(ServerSideEncryption::AWS_KMS), - _ => ServerSideEncryption::from_static(ServerSideEncryption::AES256), - }) - }) - }) - }); - let mut effective_kms_key_id = ssekms_key_id.or_else(|| { - bucket_sse_config.as_ref().and_then(|(config, _timestamp)| { - config.rules.first().and_then(|rule| { - rule.apply_server_side_encryption_by_default - .as_ref() - .and_then(|sse| sse.kms_master_key_id.clone()) - }) - }) - }); - if effective_sse - .as_ref() - .is_some_and(|sse| sse.as_str().eq_ignore_ascii_case(ServerSideEncryption::AWS_KMS)) - { - return Err(s3_error!(NotImplemented, "SSE-KMS is not supported for extract uploads")); - } - validate_sse_headers_for_write( - effective_sse.as_ref(), - effective_kms_key_id.as_ref(), - sse_customer_algorithm.as_ref(), - sse_customer_key.as_ref(), - sse_customer_key_md5.as_ref(), - true, - )?; - let Some(body) = body else { return Err(s3_error!(IncompleteBody)) }; - - let size = match content_length { - Some(c) => c, - None => { - if let Some(val) = req.headers.get(AMZ_DECODED_CONTENT_LENGTH) { - match atoi::atoi::(val.as_bytes()) { - Some(x) => x, - None => return Err(s3_error!(UnexpectedContent)), - } - } else { - return Err(s3_error!(UnexpectedContent)); - } - } - }; - if size == -1 { - return Err(s3_error!(UnexpectedContent)); - } - validate_object_key(&key, "PUT")?; - self.check_bucket_quota(&bucket, QuotaOperation::PutObject, size as u64) + let resolved_size = resolve_put_body_size(req.input.content_length, &request_context.headers)?; + self.check_bucket_quota(&req.input.bucket, QuotaOperation::PutObject, resolved_size as u64) .await?; - - // Apply adaptive buffer sizing based on file size for optimal streaming performance. - // Uses workload profile configuration (enabled by default) to select appropriate buffer size. - // Buffer sizes range from 32KB to 4MB depending on file size and configured workload profile. - let buffer_size = get_buffer_size_opt_in(size); - let body = tokio::io::BufReader::with_capacity( - buffer_size, - StreamReader::new(body.map(|f| f.map_err(|e| std::io::Error::other(e.to_string())))), - ); - - let Some(ext) = Path::new(&key).extension().and_then(|s| s.to_str()) else { - return Err(s3_error!(InvalidArgument, "key extension not found")); - }; - - let ext = ext.to_owned(); - - let md5hex = if let Some(base64_md5) = content_md5 { - let md5 = base64_simd::STANDARD - .decode_to_vec(base64_md5.as_bytes()) - .map_err(|e| ApiError::from(StorageError::other(format!("Invalid content MD5: {e}"))))?; - Some(hex_simd::encode_to_string(&md5, hex_simd::AsciiCase::Lower)) - } else { - None - }; - - let sha256hex = get_content_sha256_with_query(&req.headers, req.uri.query()); - let actual_size = size; - - 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()); - } - - let archive_etag = Arc::new(Mutex::new(None)); - let decoder = CompressionFormat::from_extension(&ext) - .get_decoder(ExtractArchiveEtagReader::new(archive_reader, archive_etag.clone())) - .map_err(|e| { - error!("get_decoder err {:?}", e); - s3_error!(InvalidArgument, "get_decoder err") - })?; - - let mut ar = Archive::new(decoder); - let mut entries = ar.entries().map_err(|e| { - error!("get entries err {:?}", e); - s3_error!(InvalidArgument, "get entries err") - })?; - - let Some(store) = new_object_layer_fn() else { - return Err(S3Error::with_message(S3ErrorCode::InternalError, "Not init".to_string())); - }; - - let extract_options = resolve_put_object_extract_options(&req.headers); - let version_id = match event_version_id { - Some(v) => v.to_string(), - None => String::new(), - }; - let notify = self .context .as_ref() .map(|context| context.notify()) .unwrap_or_else(default_notify_interface); - let req_params = extract_params_header(&req.headers); - let host = get_request_host(&req.headers); - let port = get_request_port(&req.headers); - let user_agent = get_request_user_agent(&req.headers); - - while let Some(entry) = entries.next().await { - let mut f = match entry { - Ok(f) => f, - Err(e) => { - if extract_options.ignore_errors { - warn!("Skipping archive entry because read failed and ignore-errors is enabled: {e}"); - continue; - } - error!("Failed to read archive entry: {}", e); - return Err(s3_error!(InvalidArgument, "Failed to read archive entry: {:?}", e)); - } - }; - - let fpath = match f.path() { - Ok(path) => path, - Err(e) => { - if extract_options.ignore_errors { - warn!("Skipping archive entry because path decode failed and ignore-errors is enabled: {e}"); - continue; - } - return Err(s3_error!(InvalidArgument, "Failed to decode archive entry path")); - } - }; - - let is_dir = f.header().entry_type().is_dir(); - let fpath = normalize_extract_entry_key(&fpath.to_string_lossy(), extract_options.prefix.as_deref(), is_dir); - - let mut auth_req = S3Request { - input: PutObjectInput::default(), - method: auth_method.clone(), - uri: auth_uri.clone(), - headers: auth_headers.clone(), - extensions: auth_extensions.clone(), - credentials: auth_credentials.clone(), - region: auth_region.clone(), - service: auth_service.clone(), - trailing_headers: auth_trailing_headers.clone(), - }; - { - let req_info = req_info_mut(&mut auth_req)?; - req_info.bucket = Some(bucket.clone()); - req_info.object = Some(fpath.clone()); - req_info.version_id = None; - } - authorize_request(&mut auth_req, Action::S3Action(S3Action::PutObjectAction)).await?; - - let mut size = f.header().size().unwrap_or_default() as i64; - let archive_entry_mod_time = f - .header() - .mtime() - .ok() - .and_then(|modified_at_secs| OffsetDateTime::from_unix_timestamp(modified_at_secs as i64).ok()); - let mut metadata = HashMap::new(); - apply_put_request_metadata( - &mut metadata, - &req.headers, - &fpath, - cache_control.clone(), - content_disposition.clone(), - content_encoding.clone(), - content_language.clone(), - content_type.clone(), - expires.clone(), - website_redirect_location.clone(), - tagging.clone(), - storage_class.clone(), - )?; - let mut opts = put_opts(&bucket, &fpath, None, &req.headers, metadata.clone()) - .await - .map_err(ApiError::from)?; - apply_extract_entry_pax_extensions(&mut f, &mut metadata, &mut opts).await?; - if archive_entry_mod_time.is_some() { - opts.mod_time = archive_entry_mod_time; - } - - debug!("Extracting file: {}, size: {} bytes", fpath, size); - - if is_dir { - if extract_options.ignore_dirs { - debug!("Skipping directory entry during archive extract: {}", fpath); - continue; - } - size = 0; - } - - let actual_size = size; - - 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::from_stream(f, 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(f, size, actual_size, None, None, false).map_err(ApiError::from)? - }; - apply_put_request_object_lock_opts( - &bucket, - object_lock_legal_hold_status.clone(), - object_lock_mode.clone(), - object_lock_retain_until_date.clone(), - &mut opts, - ) - .await?; - if let Some(material) = sse_encryption(EncryptionRequest { - bucket: &bucket, - key: &fpath, - server_side_encryption: effective_sse.clone(), - ssekms_key_id: effective_kms_key_id.clone(), - sse_customer_algorithm: sse_customer_algorithm.clone(), - sse_customer_key: sse_customer_key.clone(), - sse_customer_key_md5: sse_customer_key_md5.clone(), - content_size: actual_size, - part_number: None, - part_key: None, - part_nonce: None, - }) - .await? - { - effective_sse = Some(material.server_side_encryption.clone()); - effective_kms_key_id = material.kms_key_id.clone(); - - let encrypted_reader = material.wrap_reader(hrd); - hrd = HashReader::from_reader(encrypted_reader, HashReader::SIZE_PRESERVE_LAYER, actual_size, None, None, false) - .map_err(ApiError::from)?; - - let encryption_metadata = material.metadata; - metadata.extend(encryption_metadata.clone()); - opts.user_defined.extend(encryption_metadata); - } - opts.user_defined.extend(metadata); - let mut reader = PutObjReader::new(hrd); - - let obj_info = match store.put_object(&bucket, &fpath, &mut reader, &opts).await { - Ok(info) => info, - Err(e) => { - if extract_options.ignore_errors { - warn!("Skipping archive entry because object write failed and ignore-errors is enabled: {e}"); - continue; - } - return Err(ApiError::from(e).into()); - } - }; - - let manager = get_concurrency_manager(); - let fpath_clone = fpath.clone(); - let bucket_clone = bucket.clone(); - spawn_traced(async move { - manager.invalidate_cache_versioned(&bucket_clone, &fpath_clone, None).await; - }); - - let e_tag = obj_info.etag.clone().map(|etag| to_s3s_etag(&etag)); - - let output = PutObjectOutput { - e_tag, - ..Default::default() - }; - - let event_args = rustfs_notify::EventArgs { - event_name: EventName::ObjectCreatedPut, - bucket_name: bucket.clone(), - object: obj_info.clone(), - req_params: req_params.clone(), - resp_elements: extract_resp_elements(&S3Response::new(output.clone())), - version_id: version_id.clone(), - host: host.clone(), - port, - user_agent: user_agent.clone(), - }; - - let notify = notify.clone(); - let request_context = req - .extensions - .get::() - .cloned(); - spawn_background_with_context(request_context, async move { - notify.notify(event_args).await; - }); - } - - let mut checksums = PutObjectChecksums { - crc32: input.checksum_crc32, - crc32c: input.checksum_crc32c, - sha1: input.checksum_sha1, - sha256: input.checksum_sha256, - crc64nvme: input.checksum_crc64nvme, - }; - apply_trailing_checksums( - input.checksum_algorithm.as_ref().map(|a| a.as_str()), - &req.trailing_headers, - &mut checksums, - ); - - warn!( - "put object extract checksum_crc32={:?}, checksum_crc32c={:?}, checksum_sha1={:?}, checksum_sha256={:?}, checksum_crc64nvme={:?}", - checksums.crc32, checksums.crc32c, checksums.sha1, checksums.sha256, checksums.crc64nvme, - ); - - drop(entries); - let mut decoder = match ar.into_inner() { - Ok(decoder) => decoder, - Err(_) => return Err(s3_error!(InvalidArgument, "Failed to finalize archive reader")), - }; - tokio::io::copy(&mut decoder, &mut tokio::io::sink()) - .await - .map_err(map_extract_archive_error)?; - let archive_etag = archive_etag - .lock() - .ok() - .and_then(|etag| etag.clone()) - .map(|etag| to_s3s_etag(&etag)); - - let output = PutObjectOutput { - e_tag: archive_etag, - checksum_crc32: checksums.crc32, - checksum_crc32c: checksums.crc32c, - checksum_sha1: checksums.sha1, - checksum_sha256: checksums.sha256, - checksum_crc64nvme: checksums.crc64nvme, - ..Default::default() - }; - let result = Ok(S3Response::new(output)); - let _ = helper.complete(&result); - result + let input = req.input; + let output = DefaultObjectUsecase::run_put_object_extract_flow(input, request_context, notify, resolved_size).await?; + complete_put_response(helper, output) } } @@ -5174,7 +3041,7 @@ fn object_attributes_requested(object_attributes: &[ObjectAttributes], name: &'s #[cfg(test)] mod tests { use super::*; - use http::{Extensions, HeaderMap, HeaderName, HeaderValue, Method, Uri}; + use http::{Extensions, HeaderMap, Method, Uri}; fn build_request(input: T, method: Method) -> S3Request { S3Request { @@ -5199,7 +3066,7 @@ mod tests { .unwrap(); let req = build_request(input, Method::PUT); - let (event_name, quota_operation, method_name) = DefaultObjectUsecase::put_object_execution_context(&req); + let (event_name, quota_operation, method_name) = put_object_execution_context(&req); assert_eq!(event_name, EventName::ObjectCreatedPut); assert!(matches!(quota_operation, QuotaOperation::PutObject)); assert_eq!(method_name, "PUT"); @@ -5215,298 +3082,12 @@ mod tests { let mut req = build_request(input, Method::POST); req.extensions.insert(PostObjectRequestMarker); - let (event_name, quota_operation, method_name) = DefaultObjectUsecase::put_object_execution_context(&req); + let (event_name, quota_operation, method_name) = put_object_execution_context(&req); assert_eq!(event_name, EventName::ObjectCreatedPost); assert!(matches!(quota_operation, QuotaOperation::PostObject)); assert_eq!(method_name, "POST"); } - #[test] - fn is_put_object_extract_requested_accepts_meta_header() { - let mut headers = HeaderMap::new(); - headers.insert(AMZ_SNOWBALL_EXTRACT, HeaderValue::from_static("true")); - - assert!(is_put_object_extract_requested(&headers)); - } - - #[test] - fn is_put_object_extract_requested_accepts_compat_header_case_insensitive() { - let mut headers = HeaderMap::new(); - headers.insert(AMZ_SNOWBALL_EXTRACT_COMPAT, HeaderValue::from_static(" TRUE ")); - - assert!(is_put_object_extract_requested(&headers)); - } - - #[test] - fn is_put_object_extract_requested_rejects_missing_or_false_value() { - let mut headers = HeaderMap::new(); - assert!(!is_put_object_extract_requested(&headers)); - - headers.insert(AMZ_SNOWBALL_EXTRACT, HeaderValue::from_static("false")); - assert!(!is_put_object_extract_requested(&headers)); - } - - #[test] - fn normalize_snowball_prefix_trims_slashes_and_whitespace() { - assert_eq!(normalize_snowball_prefix(" /batch/incoming/ "), Some("batch/incoming".to_string())); - assert_eq!(normalize_snowball_prefix("///"), None); - } - - #[test] - fn normalize_extract_entry_key_applies_prefix_and_directory_suffix() { - assert_eq!( - normalize_extract_entry_key("nested/path.txt", Some("imports"), false), - "imports/nested/path.txt" - ); - assert_eq!(normalize_extract_entry_key("nested/dir/", Some("imports"), true), "imports/nested/dir/"); - assert_eq!(normalize_extract_entry_key("top-level", None, false), "top-level"); - } - - #[test] - fn should_use_zero_copy_rejects_boundary_at_1mb() { - let headers = HeaderMap::new(); - - assert!(!should_use_zero_copy(1024 * 1024, &headers)); - } - - #[test] - fn should_use_zero_copy_rejects_small_objects() { - let headers = HeaderMap::new(); - - assert!(!should_use_zero_copy(1024 * 1024 - 1, &headers)); - } - - #[test] - fn should_use_zero_copy_rejects_one_megabyte() { - let headers = HeaderMap::new(); - - assert!(!should_use_zero_copy(1024 * 1024, &headers)); - } - - #[test] - fn should_use_zero_copy_rejects_encrypted_requests() { - let mut headers = HeaderMap::new(); - headers.insert(AMZ_SERVER_SIDE_ENCRYPTION, HeaderValue::from_static("AES256")); - - assert!(!should_use_zero_copy(2 * 1024 * 1024, &headers)); - } - - #[test] - fn should_use_zero_copy_rejects_encrypted_requests_with_sse_customer_algorithm() { - let mut headers = HeaderMap::new(); - headers.insert(AMZ_SERVER_SIDE_ENCRYPTION_CUSTOMER_ALGORITHM, HeaderValue::from_static("AES256")); - - assert!(!should_use_zero_copy(2 * 1024 * 1024, &headers)); - } - - #[test] - fn should_use_zero_copy_rejects_encrypted_requests_with_kms_key_id() { - let mut headers = HeaderMap::new(); - headers.insert(AMZ_SERVER_SIDE_ENCRYPTION_KMS_ID, HeaderValue::from_static("test-kms-key-id")); - - assert!(!should_use_zero_copy(2 * 1024 * 1024, &headers)); - } - - #[test] - fn should_use_zero_copy_rejects_compressible_content_types() { - let mut headers = HeaderMap::new(); - headers.insert(CONTENT_TYPE, HeaderValue::from_static("application/json; charset=utf-8")); - - assert!(!should_use_zero_copy(2 * 1024 * 1024, &headers)); - } - - #[test] - fn should_use_zero_copy_allows_large_unencrypted_binary_objects() { - let mut headers = HeaderMap::new(); - headers.insert(CONTENT_TYPE, HeaderValue::from_static("application/octet-stream")); - - assert!(should_use_zero_copy(2 * 1024 * 1024, &headers)); - } - - #[test] - fn resolve_put_object_extract_options_defaults_when_headers_missing() { - let headers = HeaderMap::new(); - let options = resolve_put_object_extract_options(&headers); - assert_eq!( - options, - PutObjectExtractOptions { - prefix: None, - ignore_dirs: false, - ignore_errors: false - } - ); - } - - #[test] - fn resolve_put_object_extract_options_accepts_internal_headers() { - let mut headers = HeaderMap::new(); - headers.insert(AMZ_SNOWBALL_PREFIX_INTERNAL, HeaderValue::from_static("/internal/prefix/")); - headers.insert(AMZ_SNOWBALL_IGNORE_DIRS_INTERNAL, HeaderValue::from_static("true")); - headers.insert(AMZ_SNOWBALL_IGNORE_ERRORS_INTERNAL, HeaderValue::from_static("TRUE")); - - let options = resolve_put_object_extract_options(&headers); - assert_eq!(options.prefix.as_deref(), Some("internal/prefix")); - assert!(options.ignore_dirs); - assert!(options.ignore_errors); - } - - #[test] - fn resolve_put_object_extract_options_accepts_standard_headers() { - let mut headers = HeaderMap::new(); - headers.insert(AMZ_SNOWBALL_PREFIX, HeaderValue::from_static(" /standard/prefix/ ")); - headers.insert(AMZ_SNOWBALL_IGNORE_DIRS, HeaderValue::from_static(" true ")); - headers.insert(AMZ_SNOWBALL_IGNORE_ERRORS, HeaderValue::from_static("TRUE")); - - let options = resolve_put_object_extract_options(&headers); - assert_eq!(options.prefix.as_deref(), Some("standard/prefix")); - assert!(options.ignore_dirs); - assert!(options.ignore_errors); - } - - #[test] - fn resolve_put_object_extract_options_accepts_suffix_compatible_headers() { - let mut headers = HeaderMap::new(); - headers.insert( - HeaderName::from_static("x-amz-meta-acme-snowball-prefix"), - HeaderValue::from_static(" /partner/import "), - ); - headers.insert( - HeaderName::from_static("x-amz-meta-acme-snowball-ignore-dirs"), - HeaderValue::from_static(" true "), - ); - headers.insert( - HeaderName::from_static("x-amz-meta-acme-snowball-ignore-errors"), - HeaderValue::from_static("TRUE"), - ); - - let options = resolve_put_object_extract_options(&headers); - assert_eq!(options.prefix.as_deref(), Some("partner/import")); - assert!(options.ignore_dirs); - assert!(options.ignore_errors); - } - - #[test] - fn resolve_put_object_extract_options_prefers_exact_headers_over_suffix_fallback() { - let mut headers = HeaderMap::new(); - headers.insert("x-amz-meta-acme-snowball-prefix", HeaderValue::from_static("/fallback/prefix/")); - headers.insert(AMZ_RUSTFS_SNOWBALL_PREFIX, HeaderValue::from_static("/internal/prefix/")); - headers.insert(AMZ_SNOWBALL_PREFIX, HeaderValue::from_static("/standard/prefix/")); - headers.insert(AMZ_MINIO_SNOWBALL_PREFIX, HeaderValue::from_static("/minio/prefix/")); - - let options = resolve_put_object_extract_options(&headers); - assert_eq!(options.prefix.as_deref(), Some("minio/prefix")); - } - - #[test] - fn resolve_put_object_extract_options_exact_flags_override_suffix_fallback() { - let mut headers = HeaderMap::new(); - headers.insert(AMZ_SNOWBALL_IGNORE_DIRS, HeaderValue::from_static("false")); - headers.insert("x-amz-meta-acme-snowball-ignore-dirs", HeaderValue::from_static("true")); - headers.insert(AMZ_RUSTFS_SNOWBALL_IGNORE_ERRORS, HeaderValue::from_static("false")); - headers.insert("x-amz-meta-acme-snowball-ignore-errors", HeaderValue::from_static("true")); - - let options = resolve_put_object_extract_options(&headers); - assert!(!options.ignore_dirs); - assert!(!options.ignore_errors); - } - - #[tokio::test] - async fn execute_put_object_rejects_post_object_sse_kms_from_input() { - let input = PutObjectInput::builder() - .bucket("test-bucket".to_string()) - .key("test-key".to_string()) - .server_side_encryption(Some(ServerSideEncryption::from_static(ServerSideEncryption::AWS_KMS))) - .build() - .unwrap(); - - let mut req = build_request(input, Method::POST); - req.extensions.insert(PostObjectRequestMarker); - - let usecase = DefaultObjectUsecase::without_context(); - let fs = FS::new(); - - let err = usecase.execute_put_object(&fs, req).await.unwrap_err(); - assert_eq!(err.code(), &S3ErrorCode::NotImplemented); - } - - #[tokio::test] - async fn execute_put_object_rejects_extract_sse_kms() { - let input = PutObjectInput::builder() - .bucket("test-bucket".to_string()) - .key("archive.tar".to_string()) - .server_side_encryption(Some(ServerSideEncryption::from_static(ServerSideEncryption::AWS_KMS))) - .build() - .unwrap(); - - let mut req = build_request(input, Method::PUT); - req.headers.insert(AMZ_SNOWBALL_EXTRACT, HeaderValue::from_static("true")); - - let usecase = DefaultObjectUsecase::without_context(); - let fs = FS::new(); - - let err = usecase.execute_put_object(&fs, req).await.unwrap_err(); - assert_eq!(err.code(), &S3ErrorCode::NotImplemented); - } - - #[tokio::test] - async fn execute_put_object_extract_rejects_invalid_storage_class() { - let input = PutObjectInput::builder() - .bucket("test-bucket".to_string()) - .key("archive.tar".to_string()) - .storage_class(Some(StorageClass::from_static("INVALID"))) - .build() - .unwrap(); - - let mut req = build_request(input, Method::PUT); - req.headers.insert(AMZ_SNOWBALL_EXTRACT, HeaderValue::from_static("true")); - - let usecase = DefaultObjectUsecase::without_context(); - let fs = FS::new(); - - let err = usecase.execute_put_object(&fs, req).await.unwrap_err(); - assert_eq!(err.code(), &S3ErrorCode::InvalidStorageClass); - } - - #[tokio::test] - async fn execute_put_object_rejects_post_object_sse_kms_from_headers() { - let input = PutObjectInput::builder() - .bucket("test-bucket".to_string()) - .key("test-key".to_string()) - .build() - .unwrap(); - - let mut req = build_request(input, Method::POST); - req.extensions.insert(PostObjectRequestMarker); - req.headers - .insert(AMZ_SERVER_SIDE_ENCRYPTION, HeaderValue::from_static("aws:kms")); - - let usecase = DefaultObjectUsecase::without_context(); - let fs = FS::new(); - - let err = usecase.execute_put_object(&fs, req).await.unwrap_err(); - assert_eq!(err.code(), &S3ErrorCode::NotImplemented); - } - - #[tokio::test] - async fn execute_put_object_rejects_post_object_sse_kms_key_id_header() { - let input = PutObjectInput::builder() - .bucket("test-bucket".to_string()) - .key("test-key".to_string()) - .build() - .unwrap(); - - let mut req = build_request(input, Method::POST); - req.extensions.insert(PostObjectRequestMarker); - req.headers - .insert(AMZ_SERVER_SIDE_ENCRYPTION_KMS_ID, HeaderValue::from_static("test-kms-key-id")); - - let usecase = DefaultObjectUsecase::without_context(); - let fs = FS::new(); - - let err = usecase.execute_put_object(&fs, req).await.unwrap_err(); - assert_eq!(err.code(), &S3ErrorCode::NotImplemented); - } - #[tokio::test] async fn execute_put_object_rejects_invalid_storage_class() { let input = PutObjectInput::builder() @@ -5524,22 +3105,6 @@ mod tests { assert_eq!(err.code(), &S3ErrorCode::InvalidStorageClass); } - #[tokio::test] - async fn execute_get_object_rejects_zero_part_number() { - let input = GetObjectInput::builder() - .bucket("test-bucket".to_string()) - .key("test-key".to_string()) - .part_number(Some(0)) - .build() - .unwrap(); - - let req = build_request(input, Method::GET); - let usecase = DefaultObjectUsecase::without_context(); - - let err = usecase.execute_get_object(req).await.unwrap_err(); - assert_eq!(err.code(), &S3ErrorCode::InvalidArgument); - } - #[tokio::test] async fn execute_copy_object_rejects_self_copy_without_replace_directive() { let input = CopyObjectInput::builder() diff --git a/rustfs/src/app/object_usecase/app_adapters.rs b/rustfs/src/app/object_usecase/app_adapters.rs new file mode 100644 index 000000000..d424764fb --- /dev/null +++ b/rustfs/src/app/object_usecase/app_adapters.rs @@ -0,0 +1,616 @@ +// 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 super::get_object_flow::GetObjectBootstrap; +use super::*; +use crate::app::context::NotifyInterface; +use crate::storage::concurrency::{self, get_buffer_size_opt_in}; +use hashbrown::HashMap; +use rustfs_object_io::get::{ + CachedGetObjectSource as ObjectIoCachedGetObjectSource, GetObjectBodyPlan as ObjectIoGetObjectBodyPlan, + GetObjectCacheWriteback, GetObjectDataPlaneMetricContract as ObjectIoGetObjectDataPlaneMetricContract, GetObjectFlowResult, + GetObjectResponseMode, MaterializeGetObjectBodyError as ObjectIoMaterializeGetObjectBodyError, + build_cached_get_object_flow_result_from_source as object_io_build_cached_get_object_flow_result_from_source, + finalize_get_object_cache_writeback as object_io_finalize_get_object_cache_writeback, + materialize_get_object_body as object_io_materialize_get_object_body, plan_get_object_body as object_io_plan_get_object_body, + plan_get_object_strategy_layout as object_io_plan_get_object_strategy_layout, +}; + +pub(super) async fn prepare_get_object_request_context(req: &S3Request) -> S3Result { + let GetObjectInput { + bucket, + key, + version_id, + part_number, + range, + .. + } = req.input.clone(); + + validate_object_key(&key, "GET")?; + + let part_number = part_number.map(|v| v as usize); + + if let Some(part_num) = part_number + && part_num == 0 + { + return Err(s3_error!(InvalidArgument, "Invalid part number: part number must be greater than 0")); + } + + let rs = range.map(|v| match v { + Range::Int { first, last } => HTTPRangeSpec { + is_suffix_length: false, + start: first as i64, + end: if let Some(last) = last { last as i64 } else { -1 }, + }, + Range::Suffix { length } => HTTPRangeSpec { + is_suffix_length: true, + start: length as i64, + end: -1, + }, + }); + + if rs.is_some() && part_number.is_some() { + return Err(s3_error!(InvalidArgument, "range and part_number invalid")); + } + + let opts: ObjectOptions = get_opts(&bucket, &key, version_id.clone(), part_number, &req.headers) + .await + .map_err(ApiError::from)?; + + Ok(GetObjectRequestContext { + cache_key: ConcurrencyManager::make_cache_key(&bucket, &key, version_id.as_deref()), + version_id_for_event: version_id.unwrap_or_default(), + bucket, + key, + part_number, + rs, + opts, + headers: req.headers.clone(), + method: req.method.clone(), + sse_customer_key: req.input.sse_customer_key.clone(), + sse_customer_key_md5: req.input.sse_customer_key_md5.clone(), + }) +} + +impl ObjectIoCachedGetObjectSource for CachedGetObject { + fn body(&self) -> &std::sync::Arc { + &self.body + } + + fn content_length(&self) -> i64 { + self.content_length + } + + fn content_type(&self) -> Option<&str> { + self.content_type.as_deref() + } + + fn e_tag(&self) -> Option<&str> { + self.e_tag.as_deref() + } + + fn last_modified(&self) -> Option<&str> { + self.last_modified.as_deref() + } + + fn cache_control(&self) -> Option<&str> { + self.cache_control.as_deref() + } + + fn content_disposition(&self) -> Option<&str> { + self.content_disposition.as_deref() + } + + fn content_encoding(&self) -> Option<&str> { + self.content_encoding.as_deref() + } + + fn content_language(&self) -> Option<&str> { + self.content_language.as_deref() + } + + fn storage_class(&self) -> Option<&str> { + self.storage_class.as_deref() + } + + fn version_id(&self) -> Option<&str> { + self.version_id.as_deref() + } + + fn delete_marker(&self) -> bool { + self.delete_marker + } + + fn tag_count(&self) -> Option { + self.tag_count + } + + fn user_metadata(&self) -> &std::collections::HashMap { + &self.user_metadata + } + + fn checksum_crc32(&self) -> Option<&str> { + self.checksum_crc32.as_deref() + } + + fn checksum_crc32c(&self) -> Option<&str> { + self.checksum_crc32c.as_deref() + } + + fn checksum_sha1(&self) -> Option<&str> { + self.checksum_sha1.as_deref() + } + + fn checksum_sha256(&self) -> Option<&str> { + self.checksum_sha256.as_deref() + } + + fn checksum_crc64nvme(&self) -> Option<&str> { + self.checksum_crc64nvme.as_deref() + } + + fn checksum_type(&self) -> Option<&ChecksumType> { + self.checksum_type.as_ref() + } +} + +pub(super) fn init_get_object_bootstrap(bucket: &str, key: &str, request_id: &str) -> S3Result { + let timeout_config = TimeoutConfig::from_env(); + let wrapper = RequestTimeoutWrapper::with_request_id(timeout_config.clone(), request_id.to_string()); + let request_start = std::time::Instant::now(); + let request_guard = ConcurrencyManager::track_request(); + let concurrent_requests = GetObjectGuard::concurrent_requests(); + + let deadlock_detector = deadlock_detector::get_deadlock_detector(); + deadlock_detector.register_request(request_id, format!("GetObject {bucket}/{key}")); + let deadlock_request_guard = DeadlockRequestGuard::new(deadlock_detector, request_id.to_string()); + + if wrapper.is_timeout() { + warn!( + bucket = %bucket, + key = %key, + timeout_secs = timeout_config.get_object_timeout.as_secs(), + elapsed_ms = wrapper.elapsed().as_millis(), + "GetObject request timed out before processing" + ); + return Err(s3_error!(InternalError, "Request timeout before processing")); + } + + rustfs_io_metrics::record_get_object_request_start(concurrent_requests); + + debug!( + "GetObject request started with {} concurrent requests, timeout={:?}", + concurrent_requests, timeout_config.get_object_timeout + ); + + Ok(GetObjectBootstrap { + timeout_config, + wrapper, + request_start, + request_guard, + _deadlock_request_guard: deadlock_request_guard, + concurrent_requests, + }) +} + +#[allow(clippy::too_many_arguments)] +pub(super) async fn maybe_get_cached_get_object_flow_result( + manager: &ConcurrencyManager, + bucket: &str, + key: &str, + cache_key: &str, + version_id_for_event: String, + part_number: Option, + rs: Option<&HTTPRangeSpec>, + request_start: std::time::Instant, +) -> Option { + if !manager.is_cache_enabled() || part_number.is_some() || rs.is_some() { + return None; + } + + let cached = manager.get_cached_object(cache_key).await?; + let cache_serve_duration = request_start.elapsed(); + let metric_contract = ObjectIoGetObjectDataPlaneMetricContract::cache_served(); + + debug!("Serving object from response cache: {} (latency: {:?})", cache_key, cache_serve_duration); + + if metric_contract.record_cache_served_metric { + rustfs_io_metrics::record_get_object_cache_served(cache_serve_duration.as_secs_f64(), cached.body.len()); + } + rustfs_io_metrics::record_io_path_selected("get", metric_contract.io_path); + rustfs_io_metrics::record_io_copy_mode("get", metric_contract.copy_mode, cached.body.len()); + + manager.record_transfer(cached.content_length as u64, Duration::from_micros(1)); + + rustfs_io_metrics::record_get_object(request_start.elapsed().as_millis() as f64, cached.content_length, true); + + Some(object_io_build_cached_get_object_flow_result_from_source( + bucket, + key, + cached.as_ref(), + version_id_for_event, + )) +} + +pub(super) struct GetObjectBodyAdapterOutput { + pub(super) body: Option, + pub(super) body_plan: ObjectIoGetObjectBodyPlan, + pub(super) cache_writeback: Option, +} + +pub(super) fn spawn_get_object_cache_writeback( + cache_key: &str, + writeback: GetObjectCacheWriteback, + metric_contract: ObjectIoGetObjectDataPlaneMetricContract, +) { + debug_assert_eq!( + metric_contract.request_source, + rustfs_object_io::get::GetObjectDataPlaneRequestSource::Disk + ); + debug_assert!(!metric_contract.record_cache_served_metric); + debug_assert!(metric_contract.record_cache_writeback_metric); + + let cached_response = CachedGetObject::from_get_object_cache_writeback(writeback); + + let cache_key_clone = cache_key.to_string(); + crate::storage::request_context::spawn_traced(async move { + let manager = get_concurrency_manager(); + manager.put_cached_object(cache_key_clone.clone(), cached_response).await; + debug!("Object cached successfully with metadata: {}", cache_key_clone); + }); + + if metric_contract.record_cache_writeback_metric { + rustfs_io_metrics::record_object_cache_writeback(); + } +} + +pub(super) async fn build_get_object_body_adapter( + final_stream: R, + info: &ObjectInfo, + cache_key: &str, + response_content_length: i64, + optimal_buffer_size: usize, + cache_eligibility: rustfs_concurrency::GetObjectCacheEligibility, +) -> S3Result +where + R: AsyncRead + Send + Sync + Unpin + 'static, +{ + let body_plan = object_io_plan_get_object_body(cache_eligibility, rustfs_config::DEFAULT_OBJECT_SEEK_SUPPORT_THRESHOLD); + + match body_plan { + ObjectIoGetObjectBodyPlan::CacheWriteback => { + debug!( + "Reading object into memory for caching: key={} size={}", + cache_key, response_content_length + ); + } + ObjectIoGetObjectBodyPlan::BufferSeekable => { + debug!( + "Reading small object into memory for seek support: key={} size={}", + cache_key, response_content_length + ); + } + ObjectIoGetObjectBodyPlan::Stream if cache_eligibility.encryption_applied => { + info!( + "Encrypted object: Using unlimited stream for decryption with buffer size {}", + optimal_buffer_size + ); + } + _ => {} + } + + let materialized = + object_io_materialize_get_object_body(final_stream, info, body_plan, response_content_length, optimal_buffer_size) + .await + .map_err(|err| match err { + ObjectIoMaterializeGetObjectBodyError::CacheRead(err) => { + error!("Failed to read object into memory for caching: {}", err); + ApiError::from(StorageError::other(format!("Failed to read object for caching: {err}"))) + } + ObjectIoMaterializeGetObjectBodyError::EncryptedRead(err) => { + error!("Failed to read decrypted object into memory: {}", err); + ApiError::from(StorageError::other(format!("Failed to read decrypted object: {err}"))) + } + })?; + + Ok(GetObjectBodyAdapterOutput { + body: materialized.body, + body_plan: materialized.plan, + cache_writeback: materialized.cache_writeback.map(|writeback| { + object_io_finalize_get_object_cache_writeback( + info, + writeback, + filter_object_metadata(&info.user_defined).unwrap_or_default(), + ) + }), + }) +} + +pub(super) fn finalize_get_object_completion( + cache_key: &str, + wrapper: &RequestTimeoutWrapper, + timeout_config: &TimeoutConfig, + total_duration: Duration, + response_content_length: i64, + optimal_buffer_size: usize, + metric_contract: ObjectIoGetObjectDataPlaneMetricContract, +) { + rustfs_io_metrics::record_get_object_completion(total_duration.as_secs_f64(), response_content_length, optimal_buffer_size); + + rustfs_io_metrics::record_get_object(total_duration.as_millis() as f64, response_content_length, false); + rustfs_io_metrics::record_io_copy_mode("get", metric_contract.copy_mode, response_content_length.max(0) as usize); + + if wrapper.is_timeout() { + warn!( + "GetObject request exceeded timeout: key={} duration={:?} timeout={:?}", + cache_key, + wrapper.elapsed(), + timeout_config.get_object_timeout + ); + rustfs_io_metrics::record_get_object_timeout(None, Some(wrapper.elapsed().as_secs_f64())); + } + + debug!( + "GetObject completed: key={} size={} duration={:?} buffer={}", + cache_key, response_content_length, total_duration, optimal_buffer_size + ); +} + +#[allow(clippy::too_many_arguments)] +pub(super) fn finalize_get_object_strategy_runtime( + base_buffer_size: usize, + manager: &ConcurrencyManager, + bucket: &str, + key: &str, + info: &ObjectInfo, + rs: Option<&HTTPRangeSpec>, + response_content_length: i64, + permit_wait_duration: Duration, + queue_utilization: f64, + queue_status: &concurrency::IoQueueStatus, + concurrent_requests: usize, +) -> (concurrency::IoStrategy, usize) { + let strategy_layout = object_io_plan_get_object_strategy_layout( + rs, + response_content_length, + 0, + get_buffer_size_opt_in(response_content_length), + ); + + if let Some(range_spec) = rs + && range_spec.start >= 0 + { + manager.record_access(range_spec.start as u64, response_content_length as u64); + } + + if response_content_length > 0 { + manager.record_transfer(response_content_length as u64, permit_wait_duration); + } + + let io_strategy = manager.calculate_io_strategy_with_context( + info.size, + base_buffer_size, + permit_wait_duration, + strategy_layout.is_sequential_hint, + ); + + debug!( + wait_ms = permit_wait_duration.as_millis() as u64, + load_level = ?io_strategy.load_level, + buffer_size = io_strategy.buffer_size, + buffer_multiplier = io_strategy.buffer_multiplier, + readahead = io_strategy.enable_readahead, + cache_wb = io_strategy.cache_writeback_enabled, + storage_media = ?io_strategy.storage_media, + access_pattern = ?io_strategy.access_pattern, + bandwidth_tier = ?io_strategy.bandwidth_tier, + concurrent_requests = io_strategy.concurrent_requests, + file_size = info.size, + is_sequential = strategy_layout.is_sequential_hint, + "Enhanced multi-factor I/O strategy calculated" + ); + + let io_priority = manager.get_io_priority(response_content_length); + + if manager.is_priority_scheduling_enabled() { + debug!( + bucket = %bucket, + key = %key, + priority = %io_priority, + request_size = response_content_length, + "I/O priority assigned (based on actual request size)" + ); + + rustfs_io_metrics::record_io_priority_assignment(io_priority.as_str()); + } + + rustfs_io_metrics::record_get_object_io_state( + permit_wait_duration.as_secs_f64(), + queue_utilization, + queue_status.permits_in_use, + queue_status.total_permits.saturating_sub(queue_status.permits_in_use), + io_strategy.load_level.as_str(), + io_strategy.buffer_multiplier, + ); + rustfs_io_metrics::record_io_priority_assignment(io_priority.as_str()); + + let strategy_layout = object_io_plan_get_object_strategy_layout( + rs, + response_content_length, + io_strategy.buffer_size, + get_buffer_size_opt_in(response_content_length), + ); + + debug!( + actual_request_size = response_content_length, + priority = %io_priority.as_str(), + "I/O priority finalized with actual request size" + ); + + debug!( + "GetObject buffer sizing: file_size={}, base={}, optimal={}, concurrent_requests={}, io_strategy={:?}", + response_content_length, + get_buffer_size_opt_in(response_content_length), + strategy_layout.optimal_buffer_size, + concurrent_requests, + io_strategy.load_level + ); + + (io_strategy, strategy_layout.optimal_buffer_size) +} + +pub(super) fn prepare_put_object_request_context(req: &S3Request) -> PutObjectRequestContext { + PutObjectRequestContext { + headers: req.headers.clone(), + trailing_headers: req.trailing_headers.clone(), + uri_query: req.uri.query().map(str::to_string), + is_post_object: req.extensions.get::().is_some(), + method: req.method.clone(), + uri: req.uri.clone(), + extensions: req.extensions.clone(), + credentials: req.credentials.clone(), + region: req.region.clone(), + service: req.service.clone(), + } +} + +pub(super) fn put_object_execution_context(req: &S3Request) -> (EventName, QuotaOperation, &'static str) { + if req.extensions.get::().is_some() { + (EventName::ObjectCreatedPost, QuotaOperation::PostObject, "POST") + } else { + (EventName::ObjectCreatedPut, QuotaOperation::PutObject, "PUT") + } +} + +pub(super) fn new_operation_helper( + req: &S3Request, + event_name: EventName, + operation: S3Operation, + suppress_event: bool, +) -> OperationHelper { + let helper = OperationHelper::new(req, event_name, operation); + if suppress_event { helper.suppress_event() } else { helper } +} + +pub(super) fn bind_helper_object( + helper: OperationHelper, + object_info: ObjectInfo, + version_id: Option, +) -> OperationHelper { + let helper = helper.object(object_info); + if let Some(version_id) = version_id { + helper.version_id(version_id) + } else { + helper + } +} + +pub(super) async fn complete_get_flow_result( + helper: OperationHelper, + request_context: &GetObjectRequestContext, + flow_result: GetObjectFlowResult, +) -> S3Result> { + match flow_result.response_mode { + GetObjectResponseMode::Plain => { + let helper = bind_helper_object(helper, flow_result.event_info, Some(flow_result.version_id_for_event)); + let result = Ok(S3Response::new(flow_result.output)); + let _ = helper.complete(&result); + result + } + GetObjectResponseMode::CorsWrapped => { + let helper = helper + .object(flow_result.event_info) + .version_id(flow_result.version_id_for_event); + let response = wrap_response_with_cors( + &request_context.bucket, + &request_context.method, + &request_context.headers, + flow_result.output, + ) + .await; + let result = Ok(response); + let _ = helper.complete(&result); + result + } + } +} + +pub(super) fn complete_put_response(helper: OperationHelper, output: PutObjectOutput) -> S3Result> { + let result = Ok(S3Response::new(output)); + let _ = helper.complete(&result); + result +} + +#[allow(clippy::too_many_arguments)] +pub(super) fn spawn_put_extract_notification( + notify: Arc, + request_context: Option, + bucket: String, + req_params: HashMap, + version_id: String, + host: String, + port: u16, + user_agent: String, + obj_info: ObjectInfo, + output: PutObjectOutput, +) { + let event_args = rustfs_notify::EventArgs { + event_name: EventName::ObjectCreatedPut, + bucket_name: bucket, + object: obj_info, + req_params, + resp_elements: extract_resp_elements(&S3Response::new(output)), + version_id, + host, + port, + user_agent, + }; + + crate::storage::helper::spawn_background_with_context(request_context, async move { + notify.notify(event_args).await; + }); +} + +pub(super) async fn get_validated_store_adapter(bucket: &str) -> S3Result> { + get_validated_store(bucket).await +} + +pub(super) async fn bucket_prefix_versioning_enabled(bucket: &str, key: &str) -> bool { + BucketVersioningSys::prefix_enabled(bucket, key).await +} + +pub(super) async fn authorize_extract_put_target( + request_context: &PutObjectRequestContext, + bucket: &str, + object: &str, +) -> S3Result<()> { + let mut auth_req = S3Request { + input: PutObjectInput::default(), + method: request_context.method.clone(), + uri: request_context.uri.clone(), + headers: request_context.headers.clone(), + extensions: request_context.extensions.clone(), + credentials: request_context.credentials.clone(), + region: request_context.region.clone(), + service: request_context.service.clone(), + trailing_headers: request_context.trailing_headers.clone(), + }; + { + let req_info = req_info_mut(&mut auth_req)?; + req_info.bucket = Some(bucket.to_string()); + req_info.object = Some(object.to_string()); + req_info.version_id = None; + } + authorize_request(&mut auth_req, Action::S3Action(S3Action::PutObjectAction)).await +} diff --git a/rustfs/src/app/object_usecase/get_object_flow.rs b/rustfs/src/app/object_usecase/get_object_flow.rs new file mode 100644 index 000000000..92934e4fa --- /dev/null +++ b/rustfs/src/app/object_usecase/get_object_flow.rs @@ -0,0 +1,280 @@ +// 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 super::DeadlockRequestGuard; +use super::app_adapters::{ + bucket_prefix_versioning_enabled, build_get_object_body_adapter, finalize_get_object_completion, + finalize_get_object_strategy_runtime, maybe_get_cached_get_object_flow_result, spawn_get_object_cache_writeback, +}; +use super::get_object_zero_copy::{GetObjectPreparedRead, prepare_get_object_read_execution}; +use super::types::GetObjectRequestContext; +use crate::error::ApiError; +use crate::storage::concurrency::{self, ConcurrencyManager, GetObjectGuard}; +use crate::storage::options::filter_object_metadata; +use crate::storage::timeout_wrapper::{RequestTimeoutWrapper, TimeoutConfig}; +use rustfs_ecstore::store_api::{HTTPRangeSpec, ObjectInfo}; +use rustfs_object_io::get::{ + GetObjectBodyPlan as ObjectIoGetObjectBodyPlan, GetObjectBodySource, + GetObjectDataPlaneMetricContract as ObjectIoGetObjectDataPlaneMetricContract, GetObjectFlowResult, GetObjectOutputContext, + GetObjectReadSetup, build_chunk_blob as object_io_build_chunk_blob, + build_cors_wrapped_get_object_flow_result as object_io_build_cors_wrapped_get_object_flow_result, + build_get_object_checksums as object_io_build_get_object_checksums, + build_get_object_output_context as object_io_build_get_object_output_context, + chunk_body_data_plane_labels as object_io_chunk_body_data_plane_labels, +}; +use s3s::S3Result; +use s3s::dto::{ContentType, SSECustomerAlgorithm, SSECustomerKeyMD5, SSEKMSKeyId, ServerSideEncryption, Timestamp}; +use std::time::Duration; + +pub(super) struct GetObjectBootstrap { + pub(super) timeout_config: TimeoutConfig, + pub(super) wrapper: RequestTimeoutWrapper, + pub(super) request_start: std::time::Instant, + pub(super) request_guard: GetObjectGuard, + pub(super) _deadlock_request_guard: DeadlockRequestGuard, + pub(super) concurrent_requests: usize, +} + +#[derive(Clone, Copy)] +pub(super) struct GetObjectFlowRuntime<'a> { + pub(super) manager: &'a ConcurrencyManager, + pub(super) bootstrap: &'a GetObjectBootstrap, + pub(super) base_buffer_size: usize, +} + +#[allow(clippy::too_many_arguments)] +pub(super) async fn build_get_object_output_context( + request_context: &GetObjectRequestContext, + cache_key: &str, + manager: &ConcurrencyManager, + bucket: &str, + key: &str, + info: ObjectInfo, + event_info: ObjectInfo, + body_source: GetObjectBodySource, + rs: Option, + content_type: Option, + last_modified: Option, + response_content_length: i64, + content_range: Option, + server_side_encryption: Option, + sse_customer_algorithm: Option, + sse_customer_key_md5: Option, + ssekms_key_id: Option, + encryption_applied: bool, + permit_wait_duration: Duration, + queue_utilization: f64, + queue_status: &concurrency::IoQueueStatus, + concurrent_requests: usize, + base_buffer_size: usize, + part_number: Option, + versioned: bool, +) -> S3Result<(GetObjectOutputContext, ObjectIoGetObjectDataPlaneMetricContract)> { + let (io_strategy, optimal_buffer_size) = finalize_get_object_strategy_runtime( + base_buffer_size, + manager, + bucket, + key, + &info, + rs.as_ref(), + response_content_length, + permit_wait_duration, + queue_utilization, + queue_status, + concurrent_requests, + ); + + let (body, metric_contract) = match body_source { + GetObjectBodySource::Reader(final_stream) => { + let cache_eligibility = manager.get_object_cache_eligibility( + io_strategy.cache_writeback_enabled, + part_number.is_some(), + rs.is_some(), + encryption_applied, + response_content_length, + ); + let adapter_output = build_get_object_body_adapter( + final_stream, + &info, + cache_key, + response_content_length, + optimal_buffer_size, + cache_eligibility, + ) + .await?; + let metric_contract = ObjectIoGetObjectDataPlaneMetricContract::disk( + rustfs_io_metrics::IoPath::Legacy, + rustfs_io_metrics::CopyMode::SingleCopy, + adapter_output.body_plan, + ); + if let Some(writeback) = adapter_output.cache_writeback { + spawn_get_object_cache_writeback(cache_key, writeback, metric_contract); + } + + (adapter_output.body, metric_contract) + } + GetObjectBodySource::Chunk { + stream: chunk_stream, + path, + copy_mode, + } => { + let (io_path, copy_mode) = object_io_chunk_body_data_plane_labels(path, copy_mode); + ( + object_io_build_chunk_blob(chunk_stream), + ObjectIoGetObjectDataPlaneMetricContract::disk(io_path, copy_mode, ObjectIoGetObjectBodyPlan::Stream), + ) + } + }; + + let checksums = object_io_build_get_object_checksums(&info, &request_context.headers, part_number, rs.as_ref()) + .map_err(ApiError::from)?; + let filtered_metadata = filter_object_metadata(&info.user_defined); + + Ok(( + object_io_build_get_object_output_context( + body, + info, + event_info, + content_type, + last_modified, + response_content_length, + content_range, + server_side_encryption, + sse_customer_algorithm, + sse_customer_key_md5, + ssekms_key_id, + &checksums, + filtered_metadata, + versioned, + optimal_buffer_size, + Some(metric_contract.copy_mode), + ), + metric_contract, + )) +} + +pub(super) async fn run_get_object_flow( + request_context: GetObjectRequestContext, + runtime: GetObjectFlowRuntime<'_>, +) -> S3Result { + let GetObjectFlowRuntime { + manager, + bootstrap, + base_buffer_size, + } = runtime; + let timeout_config = &bootstrap.timeout_config; + let wrapper = &bootstrap.wrapper; + let request_start = bootstrap.request_start; + let concurrent_requests = bootstrap.concurrent_requests; + let bucket = request_context.bucket.clone(); + let key = request_context.key.clone(); + let cache_key = request_context.cache_key.clone(); + let version_id_for_event = request_context.version_id_for_event.clone(); + let part_number = request_context.part_number; + let rs = request_context.rs.clone(); + let opts = request_context.opts.clone(); + + if let Some(cached_result) = maybe_get_cached_get_object_flow_result( + manager, + &bucket, + &key, + &cache_key, + version_id_for_event.clone(), + part_number, + rs.as_ref(), + request_start, + ) + .await + { + return Ok(cached_result); + } + + let prepared_read = prepare_get_object_read_execution( + &request_context, + manager, + wrapper, + timeout_config, + &bucket, + &key, + rs, + &opts, + part_number, + ) + .await?; + let GetObjectPreparedRead { io_planning, read_setup } = prepared_read; + let permit_wait_duration = io_planning.permit_wait_duration; + let queue_status = io_planning.queue_status; + let queue_utilization = io_planning.queue_utilization; + + let GetObjectReadSetup { + info, + event_info, + body_source, + rs, + content_type, + last_modified, + response_content_length, + content_range, + server_side_encryption, + sse_customer_algorithm, + sse_customer_key_md5, + ssekms_key_id, + encryption_applied, + } = read_setup; + + let versioned = bucket_prefix_versioning_enabled(&bucket, &key).await; + let (output_context, metric_contract) = build_get_object_output_context( + &request_context, + &cache_key, + manager, + &bucket, + &key, + info, + event_info, + body_source, + rs, + content_type, + last_modified, + response_content_length, + content_range, + server_side_encryption, + sse_customer_algorithm, + sse_customer_key_md5, + ssekms_key_id, + encryption_applied, + permit_wait_duration, + queue_utilization, + &queue_status, + concurrent_requests, + base_buffer_size, + part_number, + versioned, + ) + .await?; + let response_content_length = output_context.response_content_length; + let optimal_buffer_size = output_context.optimal_buffer_size; + + let total_duration = request_start.elapsed(); + finalize_get_object_completion( + &cache_key, + wrapper, + timeout_config, + total_duration, + response_content_length, + optimal_buffer_size, + metric_contract, + ); + + Ok(object_io_build_cors_wrapped_get_object_flow_result(output_context, version_id_for_event)) +} diff --git a/rustfs/src/app/object_usecase/get_object_zero_copy.rs b/rustfs/src/app/object_usecase/get_object_zero_copy.rs new file mode 100644 index 000000000..0df64b82c --- /dev/null +++ b/rustfs/src/app/object_usecase/get_object_zero_copy.rs @@ -0,0 +1,338 @@ +// 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 super::app_adapters::get_validated_store_adapter; +use super::types::GetObjectRequestContext; +use crate::error::ApiError; +use crate::storage::concurrency::{self, ConcurrencyManager}; +use crate::storage::timeout_wrapper::{RequestTimeoutWrapper, TimeoutConfig}; +use crate::storage::{ + DecryptionRequest, check_preconditions, sse_decryption, validate_sse_headers_for_read, validate_ssec_for_read, +}; +use http::HeaderMap; +use rustfs_concurrency::GetObjectQueueSnapshot; +use rustfs_ecstore::store_api::{HTTPRangeSpec, ObjectIO, ObjectOperations, ObjectOptions}; +use rustfs_object_io::get::{ + ChunkReadDecision, ChunkReadPlanError, GetObjectEncryptionState as ObjectIoGetObjectEncryptionState, GetObjectReadSetup, + build_reader_read_setup as object_io_build_reader_read_setup, + finalize_chunk_read_setup as object_io_finalize_chunk_read_setup, + get_object_chunk_fast_path_guard as object_io_get_object_chunk_fast_path_guard, plan_chunk_read as object_io_plan_chunk_read, + plan_legacy_read as object_io_plan_legacy_read, +}; +use rustfs_rio::{Reader, WarpReader}; +use s3s::{S3Error, S3ErrorCode, S3Result, s3_error}; +use std::time::Duration; +use tracing::{debug, warn}; + +pub(super) struct GetObjectIoPlanning<'a> { + pub(super) _disk_permit: tokio::sync::SemaphorePermit<'a>, + pub(super) permit_wait_duration: Duration, + pub(super) queue_status: concurrency::IoQueueStatus, + pub(super) queue_utilization: f64, +} + +pub(super) struct GetObjectPreparedRead<'a> { + pub(super) io_planning: GetObjectIoPlanning<'a>, + pub(super) read_setup: GetObjectReadSetup, +} + +pub(super) async fn acquire_get_object_io_planning<'a>( + manager: &'a ConcurrencyManager, + wrapper: &RequestTimeoutWrapper, + timeout_config: &TimeoutConfig, + bucket: &str, + key: &str, +) -> S3Result> { + let permit_wait_start = std::time::Instant::now(); + let disk_permit = manager + .acquire_disk_read_permit() + .await + .map_err(|_| s3_error!(InternalError, "disk read semaphore closed"))?; + let permit_wait_duration = permit_wait_start.elapsed(); + + if wrapper.is_timeout() { + warn!( + bucket = %bucket, + key = %key, + wait_ms = permit_wait_duration.as_millis(), + timeout_secs = timeout_config.get_object_timeout.as_secs(), + elapsed_ms = wrapper.elapsed().as_millis(), + "GetObject request timed out while waiting for disk permit" + ); + + rustfs_io_metrics::record_get_object_timeout(Some("disk_permit"), Some(wrapper.elapsed().as_secs_f64())); + return Err(s3_error!(InternalError, "Request timeout while waiting for disk permit")); + } + + let queue_status = manager.io_queue_status(); + let queue_snapshot = GetObjectQueueSnapshot::from_available_permits( + queue_status.total_permits, + queue_status.total_permits.saturating_sub(queue_status.permits_in_use), + ); + let queue_utilization = queue_snapshot.utilization_percent(); + + if queue_snapshot.is_congested(80.0) { + warn!( + bucket = %bucket, + key = %key, + queue_utilization = format!("{:.1}%", queue_utilization), + permits_in_use = queue_status.permits_in_use, + total_permits = queue_status.total_permits, + "I/O queue congestion detected" + ); + + rustfs_io_metrics::record_io_queue_congestion(); + } + + if wrapper.is_timeout() { + warn!( + bucket = %bucket, + key = %key, + timeout_secs = timeout_config.get_object_timeout.as_secs(), + elapsed_ms = wrapper.elapsed().as_millis(), + "GetObject request timed out before reading object" + ); + rustfs_io_metrics::record_get_object_timeout(Some("before_read"), Some(wrapper.elapsed().as_secs_f64())); + return Err(s3_error!(InternalError, "Request timeout before reading object")); + } + + Ok(GetObjectIoPlanning { + _disk_permit: disk_permit, + permit_wait_duration, + queue_status, + queue_utilization, + }) +} + +#[allow(clippy::too_many_arguments)] +pub(super) async fn prepare_get_object_read( + request_context: &GetObjectRequestContext, + store: &rustfs_ecstore::store::ECStore, + manager: &ConcurrencyManager, + bucket: &str, + key: &str, + rs: Option, + h: HeaderMap, + opts: &ObjectOptions, + part_number: Option, + read_start: std::time::Instant, +) -> S3Result { + let reader = store + .get_object_reader(bucket, key, rs.clone(), h, opts) + .await + .map_err(ApiError::from)?; + + let info = reader.object_info; + + let read_duration = read_start.elapsed(); + rustfs_io_metrics::record_io_path_selected("get", rustfs_io_metrics::IoPath::Legacy); + + manager.record_disk_operation(info.size as u64, read_duration, true).await; + + check_preconditions(&request_context.headers, &info)?; + + debug!(object_size = info.size, part_count = info.parts.len(), "GET object metadata snapshot"); + for part in &info.parts { + debug!( + part_number = part.number, + part_size = part.size, + part_actual_size = part.actual_size, + "GET object part details" + ); + } + + let event_info = info.clone(); + validate_sse_headers_for_read(&info.user_defined, &request_context.headers)?; + validate_ssec_for_read( + &info.user_defined, + request_context.sse_customer_key.as_ref(), + request_context.sse_customer_key_md5.as_ref(), + )?; + let read_plan = object_io_plan_legacy_read(&info, rs, part_number).map_err(ApiError::from)?; + + debug!( + "GET object metadata check: parts={}, provided_sse_key={:?}", + info.parts.len(), + request_context.sse_customer_key.is_some() + ); + + let decryption_request = DecryptionRequest { + bucket, + key, + metadata: &info.user_defined, + sse_customer_key: request_context.sse_customer_key.as_ref(), + sse_customer_key_md5: request_context.sse_customer_key_md5.as_ref(), + part_number: None, + parts: &info.parts, + etag: info.etag.as_deref(), + }; + + let encrypted_stream = reader.stream; + + let (encryption_state, final_stream) = match sse_decryption(decryption_request).await? { + Some(material) => { + let server_side_encryption = Some(material.server_side_encryption.clone()); + let sse_customer_algorithm = Some(material.algorithm.clone()); + let sse_customer_key_md5 = material.customer_key_md5.clone(); + let ssekms_key_id = material.kms_key_id.clone(); + let (decrypted_stream, plaintext_size) = material + .wrap_reader(encrypted_stream, read_plan.response_content_length) + .await + .map_err(ApiError::from)?; + + ( + ObjectIoGetObjectEncryptionState { + server_side_encryption, + sse_customer_algorithm, + sse_customer_key_md5, + ssekms_key_id, + encryption_applied: true, + response_content_length_override: Some(plaintext_size), + }, + decrypted_stream, + ) + } + None => ( + ObjectIoGetObjectEncryptionState::default(), + Box::new(WarpReader::new(encrypted_stream)) as Box, + ), + }; + + Ok(object_io_build_reader_read_setup( + info, + event_info, + final_stream, + read_plan, + encryption_state, + )) +} + +#[allow(clippy::too_many_arguments)] +pub(super) async fn prepare_get_object_read_execution<'a>( + request_context: &GetObjectRequestContext, + manager: &'a ConcurrencyManager, + wrapper: &RequestTimeoutWrapper, + timeout_config: &TimeoutConfig, + bucket: &str, + key: &str, + rs: Option, + opts: &ObjectOptions, + part_number: Option, +) -> S3Result> { + let h = HeaderMap::new(); + let io_planning = acquire_get_object_io_planning(manager, wrapper, timeout_config, bucket, key).await?; + let store = get_validated_store_adapter(bucket).await?; + + let read_start = std::time::Instant::now(); + let read_setup = match object_io_get_object_chunk_fast_path_guard( + request_context.sse_customer_key.is_some(), + request_context.sse_customer_key_md5.is_some(), + ) { + Ok(()) => match prepare_get_object_chunk_read( + request_context, + &store, + manager, + bucket, + key, + rs.clone(), + part_number, + opts, + read_start, + ) + .await? + { + Some(read_setup) => read_setup, + None => { + prepare_get_object_read(request_context, &store, manager, bucket, key, rs, h, opts, part_number, read_start) + .await? + } + }, + Err(fallback) => { + rustfs_io_metrics::record_io_fallback(fallback.stage, fallback.reason); + prepare_get_object_read(request_context, &store, manager, bucket, key, rs, h, opts, part_number, read_start).await? + } + }; + + Ok(GetObjectPreparedRead { io_planning, read_setup }) +} + +#[allow(clippy::too_many_arguments)] +pub(super) async fn prepare_get_object_chunk_read( + request_context: &GetObjectRequestContext, + store: &rustfs_ecstore::store::ECStore, + manager: &ConcurrencyManager, + bucket: &str, + key: &str, + mut rs: Option, + part_number: Option, + opts: &ObjectOptions, + read_start: std::time::Instant, +) -> S3Result> { + let info = store.get_object_info(bucket, key, opts).await.map_err(ApiError::from)?; + + validate_sse_headers_for_read(&info.user_defined, &request_context.headers)?; + validate_ssec_for_read( + &info.user_defined, + request_context.sse_customer_key.as_ref(), + request_context.sse_customer_key_md5.as_ref(), + )?; + check_preconditions(&request_context.headers, &info)?; + + let encrypted_object = info.user_defined.contains_key("x-rustfs-encryption-key") + || info + .user_defined + .contains_key("x-amz-server-side-encryption-customer-algorithm"); + if encrypted_object { + rustfs_io_metrics::record_io_fallback( + rustfs_io_metrics::IoStage::ReadSetup, + rustfs_io_metrics::FallbackReason::EncryptionEnabled, + ); + return Ok(None); + } + + let plan = match object_io_plan_chunk_read(&info, opts.version_id.is_none(), rs.clone(), part_number) { + Ok(ChunkReadDecision::Eligible(plan)) => plan, + Ok(ChunkReadDecision::Fallback(fallback)) => { + rustfs_io_metrics::record_io_fallback(fallback.stage, fallback.reason); + return Ok(None); + } + Err(ChunkReadPlanError::NoSuchKey) => return Err(S3Error::new(S3ErrorCode::NoSuchKey)), + Err(ChunkReadPlanError::MethodNotAllowed) => return Err(S3Error::new(S3ErrorCode::MethodNotAllowed)), + Err(ChunkReadPlanError::Io(err)) => return Err(ApiError::from(err).into()), + }; + rs = plan.rs.clone(); + + let read_duration = read_start.elapsed(); + manager.record_disk_operation(info.size as u64, read_duration, true).await; + let event_info = info.clone(); + + let chunk_result = match store + .get_object_chunks(bucket, key, rs.clone(), HeaderMap::new(), opts) + .await + .map_err(ApiError::from) + { + Ok(result) => result, + Err(_err) => { + rustfs_io_metrics::record_io_fallback( + rustfs_io_metrics::IoStage::HttpBridge, + rustfs_io_metrics::FallbackReason::ChunkBridgeUnavailable, + ); + return Ok(None); + } + }; + let setup_result = object_io_finalize_chunk_read_setup(info, event_info, chunk_result, plan); + rustfs_io_metrics::record_io_path_selected("get", setup_result.io_path); + + Ok(Some(setup_result.read_setup)) +} diff --git a/rustfs/src/app/object_usecase/put_object_extract.rs b/rustfs/src/app/object_usecase/put_object_extract.rs new file mode 100644 index 000000000..0f8071b50 --- /dev/null +++ b/rustfs/src/app/object_usecase/put_object_extract.rs @@ -0,0 +1,499 @@ +// 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 super::*; +use crate::app::context::NotifyInterface; +use rustfs_object_io::put::{ + apply_extract_entry_pax_extensions, apply_trailing_checksums, is_sse_kms_requested, map_extract_archive_error, + normalize_extract_entry_key, resolve_put_object_extract_options, +}; + +impl DefaultObjectUsecase { + pub(super) async fn run_put_object_extract_flow( + input: PutObjectInput, + request_context: PutObjectRequestContext, + notify: Arc, + resolved_size: i64, + ) -> S3Result { + if is_sse_kms_requested(&input, &request_context.headers) { + return Err(s3_error!(NotImplemented, "SSE-KMS is not supported for extract uploads")); + } + + let PutObjectInput { + body, + bucket, + key, + version_id, + cache_control, + content_disposition, + content_encoding, + content_length: _content_length, + content_language, + content_type, + content_md5, + expires, + object_lock_legal_hold_status, + object_lock_mode, + object_lock_retain_until_date, + server_side_encryption, + sse_customer_algorithm, + sse_customer_key, + sse_customer_key_md5, + ssekms_key_id, + storage_class, + tagging, + website_redirect_location, + .. + } = input; + + let event_version_id = version_id; + let (h_algo, h_key, h_md5) = extract_ssec_params_from_headers(&request_context.headers)?; + let sse_customer_algorithm = sse_customer_algorithm.or(h_algo); + let sse_customer_key = sse_customer_key.or(h_key); + let sse_customer_key_md5 = sse_customer_key_md5.or(h_md5); + + let original_sse = server_side_encryption.or(extract_server_side_encryption_from_headers(&request_context.headers)?); + let bucket_sse_config = metadata_sys::get_sse_config(&bucket).await.ok(); + let mut effective_sse = original_sse.or_else(|| { + bucket_sse_config.as_ref().and_then(|(config, _timestamp)| { + config.rules.first().and_then(|rule| { + rule.apply_server_side_encryption_by_default + .as_ref() + .map(|sse| match sse.sse_algorithm.as_str() { + "AES256" => ServerSideEncryption::from_static(ServerSideEncryption::AES256), + "aws:kms" => ServerSideEncryption::from_static(ServerSideEncryption::AWS_KMS), + _ => ServerSideEncryption::from_static(ServerSideEncryption::AES256), + }) + }) + }) + }); + let mut effective_kms_key_id = ssekms_key_id.or_else(|| { + bucket_sse_config.as_ref().and_then(|(config, _timestamp)| { + config.rules.first().and_then(|rule| { + rule.apply_server_side_encryption_by_default + .as_ref() + .and_then(|sse| sse.kms_master_key_id.clone()) + }) + }) + }); + if effective_sse + .as_ref() + .is_some_and(|sse| sse.as_str().eq_ignore_ascii_case(ServerSideEncryption::AWS_KMS)) + { + return Err(s3_error!(NotImplemented, "SSE-KMS is not supported for extract uploads")); + } + validate_sse_headers_for_write( + effective_sse.as_ref(), + effective_kms_key_id.as_ref(), + sse_customer_algorithm.as_ref(), + sse_customer_key.as_ref(), + sse_customer_key_md5.as_ref(), + true, + )?; + let Some(body) = body else { return Err(s3_error!(IncompleteBody)) }; + + let size = resolved_size; + validate_object_key(&key, "PUT")?; + + let buffer_size = get_buffer_size_opt_in(size); + let body = tokio::io::BufReader::with_capacity( + buffer_size, + StreamReader::new(body.map(|f| f.map_err(|e| std::io::Error::other(e.to_string())))), + ); + + let Some(ext) = Path::new(&key).extension().and_then(|s| s.to_str()) else { + return Err(s3_error!(InvalidArgument, "key extension not found")); + }; + + let ext = ext.to_owned(); + + let md5hex = if let Some(base64_md5) = content_md5 { + let md5 = base64_simd::STANDARD + .decode_to_vec(base64_md5.as_bytes()) + .map_err(|e| ApiError::from(StorageError::other(format!("Invalid content MD5: {e}"))))?; + Some(hex_simd::encode_to_string(&md5, hex_simd::AsciiCase::Lower)) + } else { + None + }; + + let sha256hex = get_content_sha256_with_query(&request_context.headers, request_context.uri_query.as_deref()); + let actual_size = size; + + 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(&request_context.headers, request_context.trailing_headers.clone(), false) + { + return Err(ApiError::from(err).into()); + } + + let archive_etag = Arc::new(Mutex::new(None)); + let decoder = CompressionFormat::from_extension(&ext) + .get_decoder(ExtractArchiveEtagReader::new(archive_reader, archive_etag.clone())) + .map_err(|e| { + error!("get_decoder err {:?}", e); + s3_error!(InvalidArgument, "get_decoder err") + })?; + + let mut ar = Archive::new(decoder); + let mut entries = ar.entries().map_err(|e| { + error!("get entries err {:?}", e); + s3_error!(InvalidArgument, "get entries err") + })?; + + let Some(store) = new_object_layer_fn() else { + return Err(S3Error::with_message(S3ErrorCode::InternalError, "Not init".to_string())); + }; + + let extract_options = resolve_put_object_extract_options(&request_context.headers); + let version_id = match event_version_id { + Some(v) => v.to_string(), + None => String::new(), + }; + + let req_params = extract_params_header(&request_context.headers); + let host = get_request_host(&request_context.headers); + let port = get_request_port(&request_context.headers); + let user_agent = get_request_user_agent(&request_context.headers); + let tracing_context = request_context + .extensions + .get::() + .cloned(); + + while let Some(entry) = entries.next().await { + let mut f = match entry { + Ok(f) => f, + Err(e) => { + if extract_options.ignore_errors { + warn!("Skipping archive entry because read failed and ignore-errors is enabled: {e}"); + continue; + } + error!("Failed to read archive entry: {}", e); + return Err(s3_error!(InvalidArgument, "Failed to read archive entry: {:?}", e)); + } + }; + + let fpath = match f.path() { + Ok(path) => path, + Err(e) => { + if extract_options.ignore_errors { + warn!("Skipping archive entry because path decode failed and ignore-errors is enabled: {e}"); + continue; + } + return Err(s3_error!(InvalidArgument, "Failed to decode archive entry path")); + } + }; + + let is_dir = f.header().entry_type().is_dir(); + let fpath = normalize_extract_entry_key(&fpath.to_string_lossy(), extract_options.prefix.as_deref(), is_dir); + + authorize_extract_put_target(&request_context, &bucket, &fpath).await?; + + let mut size = f.header().size().unwrap_or_default() as i64; + let archive_entry_mod_time = f + .header() + .mtime() + .ok() + .and_then(|modified_at_secs| OffsetDateTime::from_unix_timestamp(modified_at_secs as i64).ok()); + let mut metadata = HashMap::new(); + apply_put_request_metadata( + &mut metadata, + &request_context.headers, + &fpath, + cache_control.clone(), + content_disposition.clone(), + content_encoding.clone(), + content_language.clone(), + content_type.clone(), + expires.clone(), + website_redirect_location.clone(), + tagging.clone(), + storage_class.clone(), + )?; + let mut opts = put_opts(&bucket, &fpath, None, &request_context.headers, metadata.clone()) + .await + .map_err(ApiError::from)?; + apply_extract_entry_pax_extensions(&mut f, &mut metadata, &mut opts).await?; + if archive_entry_mod_time.is_some() { + opts.mod_time = archive_entry_mod_time; + } + + debug!("Extracting file: {}, size: {} bytes", fpath, size); + + if is_dir { + if extract_options.ignore_dirs { + debug!("Skipping directory entry during archive extract: {}", fpath); + continue; + } + size = 0; + } + + let actual_size = size; + 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::from_stream(f, 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(f, size, actual_size, None, None, false).map_err(ApiError::from)? + }; + apply_put_request_object_lock_opts( + &bucket, + object_lock_legal_hold_status.clone(), + object_lock_mode.clone(), + object_lock_retain_until_date.clone(), + &mut opts, + ) + .await?; + if let Some(material) = sse_encryption(EncryptionRequest { + bucket: &bucket, + key: &fpath, + server_side_encryption: effective_sse.clone(), + ssekms_key_id: effective_kms_key_id.clone(), + sse_customer_algorithm: sse_customer_algorithm.clone(), + sse_customer_key: sse_customer_key.clone(), + sse_customer_key_md5: sse_customer_key_md5.clone(), + content_size: actual_size, + part_number: None, + part_key: None, + part_nonce: None, + }) + .await? + { + effective_sse = Some(material.server_side_encryption.clone()); + effective_kms_key_id = material.kms_key_id.clone(); + + let encrypted_reader = material.wrap_reader(hrd); + hrd = HashReader::from_reader(encrypted_reader, HashReader::SIZE_PRESERVE_LAYER, actual_size, None, None, false) + .map_err(ApiError::from)?; + + let encryption_metadata = material.metadata; + metadata.extend(encryption_metadata.clone()); + opts.user_defined.extend(encryption_metadata); + } + opts.user_defined.extend(metadata); + let mut reader = rustfs_ecstore::store_api::ChunkNativePutData::new(hrd); + + let obj_info = match store.put_object(&bucket, &fpath, &mut reader, &opts).await { + Ok(info) => info, + Err(e) => { + if extract_options.ignore_errors { + warn!("Skipping archive entry because object write failed and ignore-errors is enabled: {e}"); + continue; + } + return Err(ApiError::from(e).into()); + } + }; + + let manager = get_concurrency_manager(); + let fpath_clone = fpath.clone(); + let bucket_clone = bucket.clone(); + crate::storage::request_context::spawn_traced(async move { + manager.invalidate_cache_versioned(&bucket_clone, &fpath_clone, None).await; + }); + + let e_tag = obj_info.etag.clone().map(|etag| to_s3s_etag(&etag)); + + let output = PutObjectOutput { + e_tag, + ..Default::default() + }; + + spawn_put_extract_notification( + notify.clone(), + tracing_context.clone(), + bucket.clone(), + req_params.clone(), + version_id.clone(), + host.clone(), + port, + user_agent.clone(), + obj_info.clone(), + output, + ); + } + + let mut checksums = PutObjectChecksums { + crc32: input.checksum_crc32, + crc32c: input.checksum_crc32c, + sha1: input.checksum_sha1, + sha256: input.checksum_sha256, + crc64nvme: input.checksum_crc64nvme, + }; + apply_trailing_checksums( + input.checksum_algorithm.as_ref().map(|a| a.as_str()), + &request_context.trailing_headers, + &mut checksums, + ); + + drop(entries); + let mut decoder = match ar.into_inner() { + Ok(decoder) => decoder, + Err(_) => return Err(s3_error!(InvalidArgument, "Failed to finalize archive reader")), + }; + tokio::io::copy(&mut decoder, &mut tokio::io::sink()) + .await + .map_err(map_extract_archive_error)?; + let archive_etag = archive_etag + .lock() + .ok() + .and_then(|etag| etag.clone()) + .map(|etag| to_s3s_etag(&etag)); + + let output = PutObjectOutput { + e_tag: archive_etag, + checksum_crc32: checksums.crc32, + checksum_crc32c: checksums.crc32c, + checksum_sha1: checksums.sha1, + checksum_sha256: checksums.sha256, + checksum_crc64nvme: checksums.crc64nvme, + ..Default::default() + }; + Ok(output) + } +} + +#[cfg(test)] +mod tests { + use super::*; + use http::{Extensions, HeaderMap, HeaderValue, Method, Uri}; + use rustfs_utils::http::headers::{AMZ_SERVER_SIDE_ENCRYPTION, AMZ_SERVER_SIDE_ENCRYPTION_KMS_ID, AMZ_SNOWBALL_EXTRACT}; + + fn build_request(input: T, method: Method) -> S3Request { + S3Request { + input, + method, + uri: Uri::from_static("/"), + headers: HeaderMap::new(), + extensions: Extensions::new(), + credentials: None, + region: None, + service: None, + trailing_headers: None, + } + } + + #[tokio::test] + async fn execute_put_object_rejects_post_object_sse_kms_from_input() { + let input = PutObjectInput::builder() + .bucket("test-bucket".to_string()) + .key("test-key".to_string()) + .server_side_encryption(Some(ServerSideEncryption::from_static(ServerSideEncryption::AWS_KMS))) + .build() + .unwrap(); + + let mut req = build_request(input, Method::POST); + req.extensions.insert(PostObjectRequestMarker); + + let usecase = DefaultObjectUsecase::without_context(); + let fs = FS::new(); + + let err = usecase.execute_put_object(&fs, req).await.unwrap_err(); + assert_eq!(err.code(), &S3ErrorCode::NotImplemented); + } + + #[tokio::test] + async fn execute_put_object_rejects_extract_sse_kms() { + let input = PutObjectInput::builder() + .bucket("test-bucket".to_string()) + .key("archive.tar".to_string()) + .server_side_encryption(Some(ServerSideEncryption::from_static(ServerSideEncryption::AWS_KMS))) + .build() + .unwrap(); + + let mut req = build_request(input, Method::PUT); + req.headers.insert(AMZ_SNOWBALL_EXTRACT, HeaderValue::from_static("true")); + + let usecase = DefaultObjectUsecase::without_context(); + let fs = FS::new(); + + let err = usecase.execute_put_object(&fs, req).await.unwrap_err(); + assert_eq!(err.code(), &S3ErrorCode::NotImplemented); + } + + #[tokio::test] + async fn execute_put_object_extract_rejects_invalid_storage_class() { + let input = PutObjectInput::builder() + .bucket("test-bucket".to_string()) + .key("archive.tar".to_string()) + .storage_class(Some(StorageClass::from_static("INVALID"))) + .build() + .unwrap(); + + let mut req = build_request(input, Method::PUT); + req.headers.insert(AMZ_SNOWBALL_EXTRACT, HeaderValue::from_static("true")); + + let usecase = DefaultObjectUsecase::without_context(); + let fs = FS::new(); + + let err = usecase.execute_put_object(&fs, req).await.unwrap_err(); + assert_eq!(err.code(), &S3ErrorCode::InvalidStorageClass); + } + + #[tokio::test] + async fn execute_put_object_rejects_post_object_sse_kms_from_headers() { + let input = PutObjectInput::builder() + .bucket("test-bucket".to_string()) + .key("test-key".to_string()) + .build() + .unwrap(); + + let mut req = build_request(input, Method::POST); + req.extensions.insert(PostObjectRequestMarker); + req.headers + .insert(AMZ_SERVER_SIDE_ENCRYPTION, HeaderValue::from_static("aws:kms")); + + let usecase = DefaultObjectUsecase::without_context(); + let fs = FS::new(); + + let err = usecase.execute_put_object(&fs, req).await.unwrap_err(); + assert_eq!(err.code(), &S3ErrorCode::NotImplemented); + } + + #[tokio::test] + async fn execute_put_object_rejects_post_object_sse_kms_key_id_header() { + let input = PutObjectInput::builder() + .bucket("test-bucket".to_string()) + .key("test-key".to_string()) + .build() + .unwrap(); + + let mut req = build_request(input, Method::POST); + req.extensions.insert(PostObjectRequestMarker); + req.headers + .insert(AMZ_SERVER_SIDE_ENCRYPTION_KMS_ID, HeaderValue::from_static("test-kms-key-id")); + + let usecase = DefaultObjectUsecase::without_context(); + let fs = FS::new(); + + let err = usecase.execute_put_object(&fs, req).await.unwrap_err(); + assert_eq!(err.code(), &S3ErrorCode::NotImplemented); + } +} diff --git a/rustfs/src/app/object_usecase/put_object_flow.rs b/rustfs/src/app/object_usecase/put_object_flow.rs new file mode 100644 index 000000000..c64d97e73 --- /dev/null +++ b/rustfs/src/app/object_usecase/put_object_flow.rs @@ -0,0 +1,868 @@ +// 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 super::*; +use bytes::Buf; +use futures::{Stream, StreamExt}; +use rustfs_ecstore::config::GLOBAL_STORAGE_CLASS; +use rustfs_io_core::{BytesPool, PooledBuffer}; +use rustfs_object_io::put::{ + PutObjectChecksums, PutObjectIngressKind, PutObjectLegacyHashStagePlan, PutObjectLegacyHashValues, PutObjectTransformStage, + apply_trailing_checksums, build_put_object_ingress_source, build_put_object_legacy_hash_stage, + build_put_object_plain_hash_stage, plan_put_object_body_with_transforms, resolve_put_transformed_fallback_reason, +}; +use rustfs_rio::{BlockReadable, BoxReadBlockFuture, EtagResolvable, HashReaderDetector, TryGetIndex}; +use rustfs_utils::http::headers::AMZ_TRAILER; + +const DEFAULT_SMALL_PUT_EAGER_MAX_BYTES: i64 = 1024 * 1024; +const ENV_RUSTFS_PUT_SMALL_EAGER_MAX_BYTES: &str = "RUSTFS_PUT_SMALL_EAGER_MAX_BYTES"; +const ENV_RUSTFS_PUT_FORCE_DISABLE_SMALL_EAGER: &str = "RUSTFS_PUT_FORCE_DISABLE_SMALL_EAGER"; +const SLOW_PUT_PHASE_DEBUG_THRESHOLD_MS: u64 = 100; +const SLOW_PUT_PHASE_WARN_THRESHOLD_MS: u64 = 1_000; +const SLOW_PUT_PHASE_ERROR_THRESHOLD_MS: u64 = 5_000; + +fn resolved_checksum_bytes(checksums: &PutObjectChecksums) -> Option { + [ + (rustfs_rio::ChecksumType::CRC32, checksums.crc32.as_deref()), + (rustfs_rio::ChecksumType::CRC32C, checksums.crc32c.as_deref()), + (rustfs_rio::ChecksumType::SHA1, checksums.sha1.as_deref()), + (rustfs_rio::ChecksumType::SHA256, checksums.sha256.as_deref()), + (rustfs_rio::ChecksumType::CRC64_NVME, checksums.crc64nvme.as_deref()), + ] + .into_iter() + .find_map(|(checksum_type, value)| { + value.and_then(|value| rustfs_rio::Checksum::new_with_type(checksum_type, value).map(|checksum| checksum.to_bytes(&[]))) + }) +} + +fn clamp_small_put_eager_max_bytes(inline_object_limit_bytes: Option) -> i64 { + inline_object_limit_bytes + .unwrap_or(DEFAULT_SMALL_PUT_EAGER_MAX_BYTES as usize) + .min(DEFAULT_SMALL_PUT_EAGER_MAX_BYTES as usize) as i64 +} + +fn env_flag_enabled(name: &str) -> bool { + rustfs_utils::get_env_bool(name, false) +} + +fn env_non_negative_i64(name: &str) -> Option { + rustfs_utils::get_env_opt_i64(name).filter(|value| *value >= 0) +} + +fn topology_aware_small_put_eager_max_bytes(store: &rustfs_ecstore::store::ECStore, versioned: bool) -> i64 { + let Some(first_pool) = store.pools.first() else { + return DEFAULT_SMALL_PUT_EAGER_MAX_BYTES; + }; + + let data_shards = first_pool + .set_drive_count + .saturating_sub(first_pool.default_parity_count) + .max(1); + + let inline_object_limit = GLOBAL_STORAGE_CLASS + .get() + .map(|config| config.inline_object_limit_bytes(data_shards, versioned)); + + clamp_small_put_eager_max_bytes(inline_object_limit) +} + +fn resolved_small_put_eager_max_bytes(default_max_bytes: i64) -> i64 { + if env_flag_enabled(ENV_RUSTFS_PUT_FORCE_DISABLE_SMALL_EAGER) { + return 0; + } + + env_non_negative_i64(ENV_RUSTFS_PUT_SMALL_EAGER_MAX_BYTES) + .map(|value| value.min(DEFAULT_SMALL_PUT_EAGER_MAX_BYTES).min(default_max_bytes)) + .unwrap_or(default_max_bytes) +} + +fn should_use_small_put_eager_path(size: i64, eager_max_bytes: i64, compression_enabled: bool, encryption_enabled: bool) -> bool { + size > 0 && size <= eager_max_bytes && !compression_enabled && !encryption_enabled +} + +fn request_uses_trailing_checksum(headers: &HeaderMap, trailing_headers: &Option) -> bool { + trailing_headers.is_some() + || headers.contains_key(AMZ_TRAILER) + || matches!( + rustfs_rio::get_content_checksum(headers), + Ok(Some(checksum)) if checksum.checksum_type.trailing() + ) +} + +fn put_path_label(small_eager: bool, reduced_copy: bool, compressed: bool) -> &'static str { + if small_eager { + "small_eager" + } else if compressed { + "compressed" + } else if reduced_copy { + "reduced_copy" + } else { + "legacy_plain" + } +} + +#[allow(clippy::too_many_arguments)] +fn log_put_flow_phase( + bucket: &str, + key: &str, + phase: &str, + elapsed: std::time::Duration, + object_size: i64, + small_eager: bool, + reduced_copy: bool, + compressed: bool, + encrypted: bool, +) { + let duration_ms = elapsed.as_millis() as u64; + if duration_ms < SLOW_PUT_PHASE_DEBUG_THRESHOLD_MS { + return; + } + + let put_path = put_path_label(small_eager, reduced_copy, compressed); + if duration_ms >= SLOW_PUT_PHASE_ERROR_THRESHOLD_MS { + error!( + phase, + duration_ms, object_size, put_path, compressed, encrypted, bucket, key, "Small PUT phase is critically slow" + ); + } else if duration_ms >= SLOW_PUT_PHASE_WARN_THRESHOLD_MS { + warn!( + phase, + duration_ms, object_size, put_path, compressed, encrypted, bucket, key, "Small PUT phase is slow" + ); + } else { + debug!( + phase, + duration_ms, object_size, put_path, compressed, encrypted, bucket, key, "Small PUT phase exceeded debug threshold" + ); + } +} + +struct PooledBufferReader { + buffer: PooledBuffer, + position: usize, +} + +impl PooledBufferReader { + fn new(buffer: PooledBuffer) -> Self { + Self { buffer, position: 0 } + } +} + +impl tokio::io::AsyncRead for PooledBufferReader { + fn poll_read( + mut self: std::pin::Pin<&mut Self>, + _cx: &mut std::task::Context<'_>, + buf: &mut tokio::io::ReadBuf<'_>, + ) -> std::task::Poll> { + let remaining = &self.buffer[self.position..]; + if remaining.is_empty() { + return std::task::Poll::Ready(Ok(())); + } + + let to_copy = remaining.len().min(buf.remaining()); + buf.put_slice(&remaining[..to_copy]); + self.position += to_copy; + std::task::Poll::Ready(Ok(())) + } +} + +impl BlockReadable for PooledBufferReader { + fn read_block<'a>(&'a mut self, buf: &'a mut [u8]) -> BoxReadBlockFuture<'a> { + Box::pin(async move { + let remaining = &self.buffer[self.position..]; + if remaining.is_empty() { + return Ok(0); + } + + let to_copy = remaining.len().min(buf.len()); + buf[..to_copy].copy_from_slice(&remaining[..to_copy]); + self.position += to_copy; + Ok(to_copy) + }) + } +} + +impl EtagResolvable for PooledBufferReader {} + +impl HashReaderDetector for PooledBufferReader {} + +impl TryGetIndex for PooledBufferReader {} + +async fn read_small_put_body_eager(body: S, size: i64, pool: std::sync::Arc) -> S3Result +where + S: Stream>, + B: Buf, + E: std::fmt::Display, +{ + let expected_len = usize::try_from(size).map_err(|_| s3_error!(InvalidRequest, "Object size overflow"))?; + let mut data = pool.acquire_buffer(expected_len).await; + let mut body = Box::pin(body); + + while let Some(result) = body.next().await { + let mut chunk = result.map_err(|err| S3Error::with_message(S3ErrorCode::IncompleteBody, err.to_string()))?; + let chunk_len = chunk.remaining(); + if chunk_len == 0 { + continue; + } + + let new_len = data + .len() + .checked_add(chunk_len) + .ok_or_else(|| s3_error!(InvalidRequest, "Object size overflow"))?; + if new_len > expected_len { + return Err(s3_error!(IncompleteBody)); + } + + let start = data.len(); + data.resize(new_len, 0); + chunk.copy_to_slice(&mut data[start..new_len]); + + if data.len() == expected_len { + return Ok(data); + } + } + + if data.len() != expected_len { + return Err(s3_error!(IncompleteBody)); + } + + Ok(data) +} + +async fn build_small_put_eager_hash_stage( + body: S, + size: i64, + pool: std::sync::Arc, + hash_values: PutObjectLegacyHashValues, + headers: &HeaderMap, + trailing_headers: Option, +) -> S3Result +where + S: Stream>, + B: Buf, + E: std::fmt::Display, +{ + let data = read_small_put_body_eager(body, size, pool).await?; + build_put_object_legacy_hash_stage( + Box::new(PooledBufferReader::new(data)), + hash_values, + PutObjectLegacyHashStagePlan { + size, + actual_size: size, + apply_s3_checksum: true, + ignore_s3_checksum_value: false, + }, + headers, + trailing_headers, + ) + .map_err(ApiError::from) + .map_err(Into::into) +} + +impl DefaultObjectUsecase { + pub(super) async fn run_put_object_flow( + input: PutObjectInput, + request_context: PutObjectRequestContext, + request_method_name: &'static str, + resolved_size: i64, + ) -> S3Result { + let start_time = std::time::Instant::now(); + + let PutObjectInput { + body, + bucket, + cache_control, + key, + content_length: _content_length, + content_disposition, + content_encoding, + content_language, + content_type, + expires, + tagging, + metadata, + version_id, + server_side_encryption, + sse_customer_algorithm, + sse_customer_key, + sse_customer_key_md5, + ssekms_key_id, + content_md5, + object_lock_legal_hold_status, + object_lock_mode, + object_lock_retain_until_date, + storage_class, + website_redirect_location, + .. + } = input; + + let (h_algo, h_key, h_md5) = extract_ssec_params_from_headers(&request_context.headers)?; + let sse_customer_algorithm = sse_customer_algorithm.or(h_algo); + let sse_customer_key = sse_customer_key.or(h_key); + let sse_customer_key_md5 = sse_customer_key_md5.or(h_md5); + + let server_side_encryption = + server_side_encryption.or(extract_server_side_encryption_from_headers(&request_context.headers)?); + + validate_object_key(&key, request_method_name)?; + + let Some(body) = body else { return Err(s3_error!(IncompleteBody)) }; + + let mut size = resolved_size; + let mut transform_stage = PutObjectTransformStage::default(); + let mut plain_reduced_copy_stage = false; + let mut small_object_eager_stage = false; + let bytes_pool = get_concurrency_manager().bytes_pool(); + + let store = get_validated_store_adapter(&bucket).await?; + + let bucket_sse_config = metadata_sys::get_sse_config(&bucket).await.ok(); + + let mut effective_sse = server_side_encryption.or_else(|| { + bucket_sse_config.as_ref().and_then(|(config, _timestamp)| { + config.rules.first().and_then(|rule| { + rule.apply_server_side_encryption_by_default + .as_ref() + .map(|sse| match sse.sse_algorithm.as_str() { + "AES256" => ServerSideEncryption::from_static(ServerSideEncryption::AES256), + "aws:kms" => ServerSideEncryption::from_static(ServerSideEncryption::AWS_KMS), + _ => ServerSideEncryption::from_static(ServerSideEncryption::AES256), + }) + }) + }) + }); + + let mut effective_kms_key_id = ssekms_key_id.or_else(|| { + bucket_sse_config.as_ref().and_then(|(config, _timestamp)| { + config.rules.first().and_then(|rule| { + rule.apply_server_side_encryption_by_default + .as_ref() + .and_then(|sse| sse.kms_master_key_id.clone()) + }) + }) + }); + + validate_sse_headers_for_write( + effective_sse.as_ref(), + effective_kms_key_id.as_ref(), + sse_customer_algorithm.as_ref(), + sse_customer_key.as_ref(), + sse_customer_key_md5.as_ref(), + true, + )?; + + let encryption_enabled_for_put = effective_sse.is_some() + || effective_kms_key_id.is_some() + || sse_customer_algorithm.is_some() + || sse_customer_key.is_some() + || sse_customer_key_md5.is_some(); + + let body_plan = plan_put_object_body_with_transforms( + size, + &request_context.headers, + &key, + get_buffer_size_opt_in(size), + encryption_enabled_for_put, + ); + if body_plan.ingress.kind == PutObjectIngressKind::ReducedCopyCandidate { + rustfs_io_metrics::record_put_object_attempted_fast_path(size); + debug!( + encryption_enabled = encryption_enabled_for_put, + compressed = body_plan.should_compress(), + "Zero-copy write enabled for {} byte object (bucket={}, key={})", + size, + bucket, + key + ); + } else if let Some(reason) = resolve_put_transformed_fallback_reason( + body_plan.ingress.kind, + body_plan.should_compress(), + encryption_enabled_for_put, + ) { + rustfs_io_metrics::record_io_fallback(rustfs_io_metrics::IoStage::PutTransform, reason); + rustfs_io_metrics::record_put_fallback(size, reason); + } + + let mut metadata = metadata.unwrap_or_default(); + apply_put_request_metadata( + &mut metadata, + &request_context.headers, + &key, + cache_control, + content_disposition, + content_encoding, + content_language, + content_type, + expires, + website_redirect_location, + tagging, + storage_class.clone(), + )?; + + let mut opts: ObjectOptions = put_opts(&bucket, &key, version_id.clone(), &request_context.headers, metadata.clone()) + .await + .map_err(ApiError::from)?; + apply_put_request_object_lock_opts( + &bucket, + object_lock_legal_hold_status, + object_lock_mode, + object_lock_retain_until_date, + &mut opts, + ) + .await?; + let eager_max_bytes = + resolved_small_put_eager_max_bytes(topology_aware_small_put_eager_max_bytes(&store, opts.versioned)); + let can_use_small_put_eager = + !request_uses_trailing_checksum(&request_context.headers, &request_context.trailing_headers); + + let current_opts: ObjectOptions = get_opts(&bucket, &key, version_id.clone(), None, &request_context.headers) + .await + .map_err(ApiError::from)?; + match store.get_object_info(&bucket, &key, ¤t_opts).await { + Ok(existing_obj_info) => validate_existing_object_lock_for_write(&existing_obj_info)?, + Err(err) => { + if !is_err_object_not_found(&err) && !is_err_version_not_found(&err) { + return Err(ApiError::from(err).into()); + } + } + } + + let actual_size = size; + let mut hash_values = PutObjectLegacyHashValues { + md5hex: if let Some(base64_md5) = content_md5 { + let md5 = base64_simd::STANDARD + .decode_to_vec(base64_md5.as_bytes()) + .map_err(|e| ApiError::from(StorageError::other(format!("Invalid content MD5: {e}"))))?; + Some(hex_simd::encode_to_string(&md5, hex_simd::AsciiCase::Lower)) + } else { + None + }, + sha256hex: get_content_sha256_with_query(&request_context.headers, request_context.uri_query.as_deref()), + }; + + let reader_stage_start = std::time::Instant::now(); + let stage = if can_use_small_put_eager + && should_use_small_put_eager_path(size, eager_max_bytes, body_plan.should_compress(), encryption_enabled_for_put) + { + small_object_eager_stage = true; + debug!( + "Plain PUT is using the eager small-object path (bucket={}, key={}, size={}, eager_max={})", + bucket, key, size, eager_max_bytes + ); + build_small_put_eager_hash_stage( + body, + size, + bytes_pool.clone(), + hash_values, + &request_context.headers, + request_context.trailing_headers.clone(), + ) + .await? + } else if body_plan.should_compress() { + transform_stage.mark_compression(); + let algorithm = CompressionAlgorithm::default(); + insert_str(&mut metadata, SUFFIX_COMPRESSION, algorithm.to_string()); + insert_str(&mut metadata, SUFFIX_ACTUAL_SIZE, size.to_string()); + + let ingress_source = build_put_object_ingress_source(body, body_plan); + let stage = build_put_object_plain_hash_stage( + ingress_source, + std::mem::take(&mut hash_values), + PutObjectLegacyHashStagePlan { + size, + actual_size: size, + apply_s3_checksum: true, + ignore_s3_checksum_value: false, + }, + &request_context.headers, + request_context.trailing_headers.clone(), + ) + .map_err(ApiError::from)?; + + if stage.ingress_kind == PutObjectIngressKind::ReducedCopyCandidate { + plain_reduced_copy_stage = true; + } + opts.want_checksum = stage.want_checksum; + insert_str(&mut opts.user_defined, SUFFIX_COMPRESSION, algorithm.to_string()); + insert_str(&mut opts.user_defined, SUFFIX_ACTUAL_SIZE, size.to_string()); + + let reader: Box = Box::new(CompressReader::new(stage.reader, algorithm)); + size = HashReader::SIZE_PRESERVE_LAYER; + hash_values.clear_for_transformed_body(); + build_put_object_legacy_hash_stage( + reader, + hash_values, + PutObjectLegacyHashStagePlan { + size, + actual_size, + apply_s3_checksum: size >= 0, + ignore_s3_checksum_value: false, + }, + &request_context.headers, + request_context.trailing_headers.clone(), + ) + .map_err(ApiError::from)? + } else { + let ingress_source = build_put_object_ingress_source(body, body_plan); + let stage = build_put_object_plain_hash_stage( + ingress_source, + hash_values, + PutObjectLegacyHashStagePlan { + size, + actual_size, + apply_s3_checksum: size >= 0, + ignore_s3_checksum_value: false, + }, + &request_context.headers, + request_context.trailing_headers.clone(), + ) + .map_err(ApiError::from)?; + + if stage.ingress_kind == PutObjectIngressKind::ReducedCopyCandidate { + plain_reduced_copy_stage = true; + debug!( + "Plain PUT is using the reduced-copy Reader + BlockReadable hash path (bucket={}, key={})", + bucket, key + ); + } + + stage + }; + log_put_flow_phase( + &bucket, + &key, + "build_hash_stage", + reader_stage_start.elapsed(), + actual_size, + small_object_eager_stage, + plain_reduced_copy_stage, + transform_stage.compression_applied(), + false, + ); + let mut reader = stage.reader; + if stage.want_checksum.is_some() { + opts.want_checksum = stage.want_checksum; + } + + let encryption_request = EncryptionRequest { + bucket: &bucket, + key: &key, + server_side_encryption: effective_sse.clone(), + ssekms_key_id: effective_kms_key_id.clone(), + sse_customer_algorithm: sse_customer_algorithm.clone(), + sse_customer_key, + sse_customer_key_md5: sse_customer_key_md5.clone(), + content_size: actual_size, + part_number: None, + part_key: None, + part_nonce: None, + }; + + if let Some(material) = sse_encryption(encryption_request).await? { + transform_stage.mark_encryption(); + effective_sse = Some(material.server_side_encryption.clone()); + 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) + .map_err(ApiError::from)?; + + let encryption_metadata = material.metadata; + metadata.extend(encryption_metadata.clone()); + opts.user_defined.extend(encryption_metadata); + } + + let mut reader = ChunkNativePutData::new(reader); + + let mt2 = metadata.clone(); + opts.user_defined.extend(metadata); + + let repoptions = + get_must_replicate_options(&mt2, "".to_string(), ReplicationStatusType::Empty, ReplicationType::Object, opts.clone()); + + let dsc = must_replicate(&bucket, &key, repoptions).await; + + if dsc.replicate_any() { + insert_str(&mut opts.user_defined, SUFFIX_REPLICATION_TIMESTAMP, jiff::Zoned::now().to_string()); + insert_str( + &mut opts.user_defined, + SUFFIX_REPLICATION_STATUS, + dsc.pending_status().unwrap_or_default(), + ); + } + + let store_put_start = std::time::Instant::now(); + let obj_info = store + .put_object(&bucket, &key, &mut reader, &opts) + .await + .map_err(ApiError::from)?; + log_put_flow_phase( + &bucket, + &key, + "store_put_object", + store_put_start.elapsed(), + actual_size, + small_object_eager_stage, + plain_reduced_copy_stage, + transform_stage.compression_applied(), + transform_stage.encryption_applied(), + ); + + maybe_enqueue_transition_immediate(&obj_info, LcEventSrc::S3PutObject).await; + + rustfs_ecstore::data_usage::increment_bucket_usage_memory(&bucket, obj_info.size as u64).await; + + let raw_version = obj_info.version_id.map(|v| v.to_string()); + + Self::spawn_cache_invalidation(bucket.clone(), key.clone(), raw_version.clone()); + + let put_version = if bucket_prefix_versioning_enabled(&bucket, &key).await { + raw_version.clone() + } else { + None + }; + + let e_tag = obj_info.etag.clone().map(|etag| to_s3s_etag(&etag)); + + let repoptions = + get_must_replicate_options(&mt2, "".to_string(), ReplicationStatusType::Empty, ReplicationType::Object, opts); + + let dsc = must_replicate(&bucket, &key, repoptions).await; + let expiration = resolve_put_object_expiration(&bucket, &obj_info).await; + + if dsc.replicate_any() { + schedule_replication(obj_info.clone(), store.clone(), dsc, ReplicationType::Object).await; + } + + let mut checksums = PutObjectChecksums { + crc32: input.checksum_crc32, + crc32c: input.checksum_crc32c, + sha1: input.checksum_sha1, + sha256: input.checksum_sha256, + crc64nvme: input.checksum_crc64nvme, + }; + apply_trailing_checksums( + input.checksum_algorithm.as_ref().map(|a| a.as_str()), + &request_context.trailing_headers, + &mut checksums, + ); + checksums.merge_from_map(&reader.content_crc()); + if let Some(checksum_bytes) = resolved_checksum_bytes(&checksums) + && obj_info + .checksum + .as_ref() + .is_none_or(|stored| rustfs_rio::read_checksums(stored.as_ref(), 0).0.is_empty()) + { + let checksum_update_opts = ObjectOptions { + version_id: raw_version.clone(), + resolved_checksum: Some(checksum_bytes), + ..Default::default() + }; + let _ = store + .put_object_metadata(&bucket, &key, &checksum_update_opts) + .await + .map_err(ApiError::from)?; + } + + let output = PutObjectOutput { + e_tag, + server_side_encryption: effective_sse, + sse_customer_algorithm: sse_customer_algorithm.clone(), + sse_customer_key_md5: sse_customer_key_md5.clone(), + ssekms_key_id: effective_kms_key_id, + expiration, + checksum_crc32: checksums.crc32, + checksum_crc32c: checksums.crc32c, + checksum_sha1: checksums.sha1, + checksum_sha256: checksums.sha256, + checksum_crc64nvme: checksums.crc64nvme, + version_id: put_version, + ..Default::default() + }; + + let manager = get_capacity_manager(); + manager.record_write_operation().await; + + { + let duration_ms = start_time.elapsed().as_millis() as f64; + let fast_path_selected = plain_reduced_copy_stage || small_object_eager_stage; + rustfs_io_metrics::record_put_object(duration_ms, size, fast_path_selected); + let io_path = if fast_path_selected { + rustfs_io_metrics::IoPath::Fast + } else { + rustfs_io_metrics::IoPath::Legacy + }; + rustfs_io_metrics::record_io_path_selected("put", io_path); + rustfs_io_metrics::record_put_path_selected(actual_size, io_path); + let effective_copy_mode = transform_stage.effective_copy_mode(); + rustfs_io_metrics::record_io_copy_mode("put", effective_copy_mode, actual_size.max(0) as usize); + rustfs_io_metrics::record_put_copy_mode(actual_size, effective_copy_mode); + if let Some(transform_kind) = transform_stage.metric_kind() { + rustfs_io_metrics::record_put_transform_selected(transform_kind, io_path, actual_size.max(0) as usize); + } + } + + Ok(PutObjectFlowResult { + output, + helper_object: obj_info, + helper_version_id: raw_version, + }) + } +} + +#[cfg(test)] +mod tests { + use super::*; + use bytes::Bytes; + use futures::{StreamExt, stream}; + use rustfs_io_core::BytesPool; + use serial_test::serial; + use std::sync::Arc; + use tokio::time::{Duration, timeout}; + + #[test] + fn small_put_eager_path_only_targets_plain_small_objects() { + assert!(should_use_small_put_eager_path( + 64 * 1024, + DEFAULT_SMALL_PUT_EAGER_MAX_BYTES, + false, + false + )); + assert!(should_use_small_put_eager_path( + DEFAULT_SMALL_PUT_EAGER_MAX_BYTES, + DEFAULT_SMALL_PUT_EAGER_MAX_BYTES, + false, + false + )); + assert!(!should_use_small_put_eager_path( + DEFAULT_SMALL_PUT_EAGER_MAX_BYTES + 1, + DEFAULT_SMALL_PUT_EAGER_MAX_BYTES, + false, + false + )); + assert!(!should_use_small_put_eager_path( + 64 * 1024, + DEFAULT_SMALL_PUT_EAGER_MAX_BYTES, + true, + false + )); + assert!(!should_use_small_put_eager_path( + 64 * 1024, + DEFAULT_SMALL_PUT_EAGER_MAX_BYTES, + false, + true + )); + assert!(!should_use_small_put_eager_path(0, DEFAULT_SMALL_PUT_EAGER_MAX_BYTES, false, false)); + } + + #[test] + fn clamp_small_put_eager_max_bytes_caps_inline_budget() { + assert_eq!( + clamp_small_put_eager_max_bytes(Some(rustfs_object_io::put::PUT_REDUCED_COPY_MIN_SIZE_BYTES as usize * 2)), + DEFAULT_SMALL_PUT_EAGER_MAX_BYTES + ); + assert_eq!(clamp_small_put_eager_max_bytes(Some(128 * 1024)), 128 * 1024); + assert_eq!(clamp_small_put_eager_max_bytes(None), DEFAULT_SMALL_PUT_EAGER_MAX_BYTES); + } + + #[test] + #[serial] + fn resolved_small_put_eager_max_bytes_honors_disable_env() { + temp_env::with_var(ENV_RUSTFS_PUT_FORCE_DISABLE_SMALL_EAGER, Some("true"), || { + assert_eq!(resolved_small_put_eager_max_bytes(256 * 1024), 0); + }); + } + + #[test] + #[serial] + fn resolved_small_put_eager_max_bytes_narrows_default_budget() { + temp_env::with_var(ENV_RUSTFS_PUT_SMALL_EAGER_MAX_BYTES, Some("4096"), || { + assert_eq!(resolved_small_put_eager_max_bytes(256 * 1024), 4096); + }); + } + + #[test] + #[serial] + fn resolved_small_put_eager_max_bytes_ignores_invalid_override() { + temp_env::with_var(ENV_RUSTFS_PUT_SMALL_EAGER_MAX_BYTES, Some("invalid"), || { + assert_eq!(resolved_small_put_eager_max_bytes(256 * 1024), 256 * 1024); + }); + } + + #[tokio::test] + async fn read_small_put_body_eager_requires_exact_content_length() { + let body = stream::iter(vec![ + Ok::(Bytes::from_static(b"abc")), + Ok::(Bytes::from_static(b"def")), + ]); + + let pool = Arc::new(BytesPool::new_tiered()); + let data = read_small_put_body_eager(body, 6, pool) + .await + .expect("eager read should succeed"); + assert_eq!(data.as_ref(), b"abcdef"); + } + + #[tokio::test] + async fn read_small_put_body_eager_rejects_length_mismatch() { + let body = stream::iter(vec![Ok::(Bytes::from_static(b"abc"))]); + let pool = Arc::new(BytesPool::new_tiered()); + + let err = read_small_put_body_eager(body, 4, pool) + .await + .expect_err("short eager read should fail"); + assert_eq!(err.code(), &S3ErrorCode::IncompleteBody); + } + + #[tokio::test] + async fn read_small_put_body_eager_rejects_overlong_body() { + let body = stream::iter(vec![ + Ok::(Bytes::from_static(b"abc")), + Ok::(Bytes::from_static(b"def")), + ]); + let pool = Arc::new(BytesPool::new_tiered()); + + let err = read_small_put_body_eager(body, 5, pool) + .await + .expect_err("overlong eager read should fail"); + assert_eq!(err.code(), &S3ErrorCode::IncompleteBody); + } + + #[tokio::test] + async fn read_small_put_body_eager_returns_buffer_to_pool_after_drop() { + let body = stream::iter(vec![Ok::(Bytes::from_static(b"abc"))]); + let pool = Arc::new(BytesPool::new_tiered()); + + let data = read_small_put_body_eager(body, 3, pool.clone()) + .await + .expect("pooled eager read should succeed"); + assert_eq!(pool.available_buffers(), 0); + + drop(data); + assert_eq!(pool.available_buffers(), 1); + } + + #[tokio::test] + async fn read_small_put_body_eager_returns_after_expected_bytes_without_waiting_for_eof() { + let body = stream::once(async { Ok::(Bytes::from_static(b"abc")) }).chain(stream::pending()); + let pool = Arc::new(BytesPool::new_tiered()); + + let data = timeout(Duration::from_millis(50), read_small_put_body_eager(body, 3, pool)) + .await + .expect("eager read should not wait for stream termination") + .expect("eager read should succeed once content-length bytes are read"); + + assert_eq!(data.as_ref(), b"abc"); + } +} diff --git a/rustfs/src/app/object_usecase/types.rs b/rustfs/src/app/object_usecase/types.rs new file mode 100644 index 000000000..63444d8d9 --- /dev/null +++ b/rustfs/src/app/object_usecase/types.rs @@ -0,0 +1,52 @@ +// 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 super::*; + +#[derive(Clone)] +pub(super) struct GetObjectRequestContext { + pub(super) bucket: String, + pub(super) key: String, + pub(super) cache_key: String, + pub(super) version_id_for_event: String, + pub(super) part_number: Option, + pub(super) rs: Option, + pub(super) opts: ObjectOptions, + pub(super) headers: HeaderMap, + pub(super) method: hyper::Method, + pub(super) sse_customer_key: Option, + pub(super) sse_customer_key_md5: Option, +} + +pub(super) type PutObjectChecksums = rustfs_object_io::put::PutObjectChecksums; + +#[derive(Clone)] +pub(super) struct PutObjectRequestContext { + pub(super) headers: HeaderMap, + pub(super) trailing_headers: Option, + pub(super) uri_query: Option, + pub(super) is_post_object: bool, + pub(super) method: hyper::Method, + pub(super) uri: hyper::Uri, + pub(super) extensions: http::Extensions, + pub(super) credentials: Option, + pub(super) region: Option, + pub(super) service: Option, +} + +pub(super) struct PutObjectFlowResult { + pub(super) output: PutObjectOutput, + pub(super) helper_object: ObjectInfo, + pub(super) helper_version_id: Option, +} diff --git a/rustfs/src/app/object_usecase/zero_copy_tests.rs b/rustfs/src/app/object_usecase/zero_copy_tests.rs new file mode 100644 index 000000000..59c3dba1c --- /dev/null +++ b/rustfs/src/app/object_usecase/zero_copy_tests.rs @@ -0,0 +1,1222 @@ +// 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 super::*; +use futures::StreamExt; +use http::{Extensions, HeaderMap, Method, Uri}; +use rustfs_ecstore::store_api::GetObjectChunkPath; +use rustfs_ecstore::{ + bucket::metadata_sys, + disk::endpoint::Endpoint, + endpoints::{EndpointServerPools, Endpoints, PoolEndpoints}, + store::ECStore, + store_api::{ + BucketOperations, ChunkNativePutData, CompletePart, MakeBucketOptions, MultipartOperations, ObjectIO, ObjectOptions, + }, +}; +use rustfs_object_io::get::{GetObjectBodySource, GetObjectReadSetup}; +use serial_test::serial; +use std::{ + path::PathBuf, + sync::{Arc, Once, OnceLock}, +}; +use tokio::fs; +use tokio_util::sync::CancellationToken; +use uuid::Uuid; + +static DIRECT_CHUNK_TEST_ENV: OnceLock<(Vec, Arc)> = OnceLock::new(); +static DIRECT_CHUNK_MULTI_DISK_TEST_ENV: OnceLock<(Vec, Arc)> = OnceLock::new(); +static DIRECT_CHUNK_TEST_INIT: Once = Once::new(); + +fn init_direct_chunk_test_tracing() { + DIRECT_CHUNK_TEST_INIT.call_once(|| {}); +} + +async fn setup_direct_chunk_test_env() -> (Vec, Arc) { + init_direct_chunk_test_tracing(); + + if let Some((paths, store)) = DIRECT_CHUNK_TEST_ENV.get() { + return (paths.clone(), store.clone()); + } + + let test_base_dir = format!("/tmp/rustfs_app_chunk_direct_test_{}", Uuid::new_v4()); + let temp_dir = PathBuf::from(&test_base_dir); + if temp_dir.exists() { + fs::remove_dir_all(&temp_dir).await.ok(); + } + fs::create_dir_all(&temp_dir).await.unwrap(); + + let disk_path = temp_dir.join("disk1"); + fs::create_dir_all(&disk_path).await.unwrap(); + + let mut endpoint = Endpoint::try_from(disk_path.to_str().unwrap()).unwrap(); + endpoint.set_pool_index(0); + endpoint.set_set_index(0); + endpoint.set_disk_index(0); + + let pool_endpoints = PoolEndpoints { + legacy: false, + set_count: 1, + drives_per_set: 1, + endpoints: Endpoints::from(vec![endpoint]), + cmd_line: "test".to_string(), + platform: format!("OS: {} | Arch: {}", std::env::consts::OS, std::env::consts::ARCH), + }; + + let endpoint_pools = EndpointServerPools(vec![pool_endpoints]); + + rustfs_ecstore::store::init_local_disks(endpoint_pools.clone()).await.unwrap(); + + let server_addr: std::net::SocketAddr = "127.0.0.1:9013".parse().unwrap(); + let ecstore = ECStore::new(server_addr, endpoint_pools, CancellationToken::new()) + .await + .unwrap(); + + let buckets_list = ecstore + .list_bucket(&rustfs_ecstore::store_api::BucketOptions { + no_metadata: true, + ..Default::default() + }) + .await + .unwrap(); + let buckets = buckets_list.into_iter().map(|v| v.name).collect(); + metadata_sys::init_bucket_metadata_sys(ecstore.clone(), buckets).await; + + let _ = DIRECT_CHUNK_TEST_ENV.set((vec![disk_path], ecstore.clone())); + + (DIRECT_CHUNK_TEST_ENV.get().unwrap().0.clone(), ecstore) +} + +async fn setup_direct_chunk_multi_disk_test_env() -> (Vec, Arc) { + init_direct_chunk_test_tracing(); + + if let Some((paths, store)) = DIRECT_CHUNK_MULTI_DISK_TEST_ENV.get() { + return (paths.clone(), store.clone()); + } + + let test_base_dir = format!("/tmp/rustfs_app_chunk_direct_multi_test_{}", Uuid::new_v4()); + let temp_dir = PathBuf::from(&test_base_dir); + if temp_dir.exists() { + fs::remove_dir_all(&temp_dir).await.ok(); + } + fs::create_dir_all(&temp_dir).await.unwrap(); + + let disk_paths = vec![ + temp_dir.join("disk1"), + temp_dir.join("disk2"), + temp_dir.join("disk3"), + temp_dir.join("disk4"), + ]; + for disk_path in &disk_paths { + fs::create_dir_all(disk_path).await.unwrap(); + } + + let mut endpoints = Vec::new(); + for (i, disk_path) in disk_paths.iter().enumerate() { + let mut endpoint = Endpoint::try_from(disk_path.to_str().unwrap()).unwrap(); + endpoint.set_pool_index(0); + endpoint.set_set_index(0); + endpoint.set_disk_index(i); + endpoints.push(endpoint); + } + + let pool_endpoints = PoolEndpoints { + legacy: false, + set_count: 1, + drives_per_set: 4, + endpoints: Endpoints::from(endpoints), + cmd_line: "test".to_string(), + platform: format!("OS: {} | Arch: {}", std::env::consts::OS, std::env::consts::ARCH), + }; + + let endpoint_pools = EndpointServerPools(vec![pool_endpoints]); + + rustfs_ecstore::store::init_local_disks(endpoint_pools.clone()).await.unwrap(); + + let server_addr: std::net::SocketAddr = "127.0.0.1:9014".parse().unwrap(); + let ecstore = ECStore::new(server_addr, endpoint_pools, CancellationToken::new()) + .await + .unwrap(); + + let buckets_list = ecstore + .list_bucket(&rustfs_ecstore::store_api::BucketOptions { + no_metadata: true, + ..Default::default() + }) + .await + .unwrap(); + let buckets = buckets_list.into_iter().map(|v| v.name).collect(); + metadata_sys::init_bucket_metadata_sys(ecstore.clone(), buckets).await; + + let _ = DIRECT_CHUNK_MULTI_DISK_TEST_ENV.set((disk_paths.clone(), ecstore.clone())); + + (DIRECT_CHUNK_MULTI_DISK_TEST_ENV.get().unwrap().0.clone(), ecstore) +} + +async fn create_direct_chunk_test_bucket(ecstore: &Arc, bucket_name: &str) { + (**ecstore) + .make_bucket( + bucket_name, + &MakeBucketOptions { + versioning_enabled: true, + ..Default::default() + }, + ) + .await + .unwrap(); +} + +async fn create_direct_chunk_test_multipart_object( + ecstore: &Arc, + bucket: &str, + key: &str, + parts: Vec>, +) -> Vec> { + let upload = ecstore + .new_multipart_upload(bucket, key, &ObjectOptions::default()) + .await + .unwrap(); + + let mut completed_parts = Vec::new(); + for (idx, part) in parts.iter().enumerate() { + let mut reader = ChunkNativePutData::from_vec(part.clone()); + let part_info = ecstore + .put_object_part(bucket, key, &upload.upload_id, idx + 1, &mut reader, &ObjectOptions::default()) + .await + .unwrap(); + completed_parts.push(CompletePart { + part_num: idx + 1, + etag: part_info.etag, + ..Default::default() + }); + } + + ecstore + .clone() + .complete_multipart_upload(bucket, key, &upload.upload_id, completed_parts, &ObjectOptions::default()) + .await + .unwrap(); + + parts +} + +fn find_part_file(root: &std::path::Path, part_name: &str) -> Option { + let entries = std::fs::read_dir(root).ok()?; + for entry in entries.flatten() { + let path = entry.path(); + if path.is_dir() { + if let Some(found) = find_part_file(&path, part_name) { + return Some(found); + } + continue; + } + + if path.file_name().and_then(|name| name.to_str()) == Some(part_name) { + return Some(path); + } + } + + None +} + +fn find_part_files(root: &std::path::Path, part_name: &str, out: &mut Vec) { + let Ok(entries) = std::fs::read_dir(root) else { + return; + }; + for entry in entries.flatten() { + let path = entry.path(); + if path.is_dir() { + find_part_files(&path, part_name, out); + continue; + } + + if path.file_name().and_then(|name| name.to_str()) == Some(part_name) { + out.push(path); + } + } +} + +async fn remove_part_files(root: &std::path::Path, part_name: &str) -> Vec<(PathBuf, Vec)> { + let mut paths = Vec::new(); + find_part_files(root, part_name, &mut paths); + + let mut removed = Vec::with_capacity(paths.len()); + for path in paths { + let content = fs::read(&path).await.expect("read part file before removal"); + fs::remove_file(&path).await.expect("remove part file"); + removed.push((path, content)); + } + + removed +} + +async fn restore_part_files(files: Vec<(PathBuf, Vec)>) { + for (path, content) in files { + if let Some(parent) = path.parent() { + fs::create_dir_all(parent).await.expect("ensure parent for part restore"); + } + fs::write(&path, content).await.expect("restore part file"); + } +} + +async fn select_reconstructed_chunk_read( + disk_paths: &[PathBuf], + bucket: &str, + key: &str, + missing_part_name: &str, + ecstore: &Arc, + manager: &ConcurrencyManager, + request_context: &GetObjectRequestContext, +) -> GetObjectReadSetup { + for disk_path in disk_paths { + let object_root = disk_path.join(bucket).join(key); + let removed_parts = remove_part_files(&object_root, missing_part_name).await; + if removed_parts.is_empty() { + continue; + } + + let candidate = get_object_zero_copy::prepare_get_object_chunk_read( + request_context, + ecstore, + manager, + &request_context.bucket, + &request_context.key, + request_context.rs.clone(), + request_context.part_number, + &request_context.opts, + std::time::Instant::now(), + ) + .await + .unwrap() + .expect("expected chunk fast path"); + + let is_reconstructed = matches!( + &candidate.body_source, + GetObjectBodySource::Chunk { + path: GetObjectChunkPath::Direct, + copy_mode: rustfs_io_metrics::CopyMode::Reconstructed, + .. + } + ); + if is_reconstructed { + return candidate; + } + + restore_part_files(removed_parts).await; + } + + panic!("expected reconstructed chunk path after removing one disk shard copy of {missing_part_name}"); +} + +fn build_request(input: T, method: Method) -> S3Request { + S3Request { + input, + method, + uri: Uri::from_static("/"), + headers: HeaderMap::new(), + extensions: Extensions::new(), + credentials: None, + region: None, + service: None, + trailing_headers: None, + } +} + +#[tokio::test] +async fn execute_get_object_rejects_zero_part_number() { + let input = GetObjectInput::builder() + .bucket("test-bucket".to_string()) + .key("test-key".to_string()) + .part_number(Some(0)) + .build() + .unwrap(); + + let req = build_request(input, Method::GET); + let usecase = DefaultObjectUsecase::without_context(); + + let err = usecase.execute_get_object(req).await.unwrap_err(); + assert_eq!(err.code(), &S3ErrorCode::InvalidArgument); +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 1)] +#[serial] +#[ignore = "requires isolated global object layer state"] +async fn prepare_get_object_chunk_read_marks_direct_path_for_single_disk_store() { + let (_disk_paths, ecstore) = setup_direct_chunk_test_env().await; + let bucket = format!("direct-chunk-{}", &Uuid::new_v4().simple().to_string()[..8]); + let key = "test/object.bin"; + let payload = vec![1u8; 128 * 1024]; + + create_direct_chunk_test_bucket(&ecstore, &bucket).await; + + let mut reader = ChunkNativePutData::from_vec(payload.clone()); + ecstore + .put_object(&bucket, key, &mut reader, &ObjectOptions::default()) + .await + .unwrap(); + + let input = GetObjectInput::builder() + .bucket(bucket.clone()) + .key(key.to_string()) + .build() + .unwrap(); + let req = build_request(input, Method::GET); + + let request_context = prepare_get_object_request_context(&req).await.unwrap(); + let manager = get_concurrency_manager(); + + let read_setup = get_object_zero_copy::prepare_get_object_chunk_read( + &request_context, + &ecstore, + manager, + &request_context.bucket, + &request_context.key, + request_context.rs.clone(), + request_context.part_number, + &request_context.opts, + std::time::Instant::now(), + ) + .await + .unwrap() + .expect("expected chunk fast path"); + + match read_setup.body_source { + GetObjectBodySource::Chunk { path, copy_mode, .. } => { + assert!(matches!(path, GetObjectChunkPath::Direct), "expected direct chunk path"); + assert_eq!(copy_mode, rustfs_io_metrics::CopyMode::TrueZeroCopy); + } + GetObjectBodySource::Reader(_) => panic!("expected chunk body source"), + } + + let usecase = DefaultObjectUsecase::without_context(); + let response = usecase.execute_get_object(req).await.unwrap(); + assert_eq!(response.output.content_length, Some(payload.len() as i64)); + let mut body = response.output.body.expect("expected body"); + let mut collected = Vec::new(); + while let Some(chunk) = body.next().await { + collected.extend_from_slice(&chunk.unwrap()); + } + assert_eq!(collected, payload); +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 1)] +#[serial] +#[ignore = "requires isolated global object layer state"] +async fn prepare_get_object_chunk_read_falls_back_to_legacy_when_chunk_bridge_fails() { + let (disk_paths, ecstore) = setup_direct_chunk_test_env().await; + let bucket = format!("direct-chunk-fallback-{}", &Uuid::new_v4().simple().to_string()[..8]); + let key = "test/fallback.bin"; + let payload = vec![3u8; 128 * 1024]; + + create_direct_chunk_test_bucket(&ecstore, &bucket).await; + + let mut reader = ChunkNativePutData::from_vec(payload); + ecstore + .put_object(&bucket, key, &mut reader, &ObjectOptions::default()) + .await + .unwrap(); + + let object_root = disk_paths[0].join(&bucket).join("test").join("fallback.bin"); + let missing_part = find_part_file(&object_root, "part.1").expect("part file on disk"); + fs::remove_file(&missing_part).await.expect("remove single-disk part file"); + + let input = GetObjectInput::builder() + .bucket(bucket.clone()) + .key(key.to_string()) + .build() + .unwrap(); + let req = build_request(input, Method::GET); + let request_context = prepare_get_object_request_context(&req).await.unwrap(); + let manager = get_concurrency_manager(); + + let read_setup = get_object_zero_copy::prepare_get_object_chunk_read( + &request_context, + &ecstore, + manager, + &request_context.bucket, + &request_context.key, + request_context.rs.clone(), + request_context.part_number, + &request_context.opts, + std::time::Instant::now(), + ) + .await + .unwrap(); + + assert!(read_setup.is_none(), "chunk bridge failure should fall back to legacy reader"); +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 1)] +#[serial] +#[ignore = "requires isolated global object layer state"] +async fn execute_get_object_range_marks_direct_path_for_single_disk_store() { + let (_disk_paths, ecstore) = setup_direct_chunk_test_env().await; + let bucket = format!("direct-range-{}", &Uuid::new_v4().simple().to_string()[..8]); + let key = "test/range.bin"; + let payload = vec![2u8; 128 * 1024]; + + create_direct_chunk_test_bucket(&ecstore, &bucket).await; + + let mut reader = ChunkNativePutData::from_vec(payload.clone()); + ecstore + .put_object(&bucket, key, &mut reader, &ObjectOptions::default()) + .await + .unwrap(); + + let input = GetObjectInput::builder() + .bucket(bucket.clone()) + .key(key.to_string()) + .range(Some(Range::Int { + first: 0, + last: Some(63 * 1024), + })) + .build() + .unwrap(); + let req = build_request(input, Method::GET); + + let request_context = prepare_get_object_request_context(&req).await.unwrap(); + let manager = get_concurrency_manager(); + + let read_setup = get_object_zero_copy::prepare_get_object_chunk_read( + &request_context, + &ecstore, + manager, + &request_context.bucket, + &request_context.key, + request_context.rs.clone(), + request_context.part_number, + &request_context.opts, + std::time::Instant::now(), + ) + .await + .unwrap() + .expect("expected chunk fast path"); + + match read_setup.body_source { + GetObjectBodySource::Chunk { path, .. } => { + assert!(matches!(path, GetObjectChunkPath::Direct), "expected direct chunk path") + } + GetObjectBodySource::Reader(_) => panic!("expected chunk body source"), + } + + let usecase = DefaultObjectUsecase::without_context(); + let response = usecase.execute_get_object(req).await.unwrap(); + assert_eq!(response.output.content_length, Some(63 * 1024 + 1)); + let mut body = response.output.body.expect("expected body"); + let mut collected = Vec::new(); + while let Some(chunk) = body.next().await { + collected.extend_from_slice(&chunk.unwrap()); + } + assert_eq!(collected, payload[..(63 * 1024 + 1)].to_vec()); +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 1)] +#[serial] +#[ignore = "requires isolated global object layer state"] +async fn execute_get_object_range_marks_direct_path_for_multi_disk_store_without_rebuild() { + let (_disk_paths, ecstore) = setup_direct_chunk_multi_disk_test_env().await; + let bucket = format!("direct-multi-range-{}", &Uuid::new_v4().simple().to_string()[..8]); + let key = "test/multi-range.bin"; + let payload_len = 3 * 1024 * 1024 + 137; + let payload: Vec = (0..payload_len).map(|idx| (idx % 251) as u8).collect(); + + create_direct_chunk_test_bucket(&ecstore, &bucket).await; + + let mut reader = ChunkNativePutData::from_vec(payload.clone()); + let put_info = ecstore + .put_object(&bucket, key, &mut reader, &ObjectOptions::default()) + .await + .unwrap(); + assert!( + put_info.data_blocks > 1, + "expected multi-data-shard object, got {}+{}", + put_info.data_blocks, + put_info.parity_blocks + ); + + let range_start = 123_457_u64; + let range_end = 2 * 1024 * 1024 + 33_333_u64; + let input = GetObjectInput::builder() + .bucket(bucket.clone()) + .key(key.to_string()) + .range(Some(Range::Int { + first: range_start, + last: Some(range_end), + })) + .build() + .unwrap(); + let req = build_request(input, Method::GET); + + let request_context = prepare_get_object_request_context(&req).await.unwrap(); + let manager = get_concurrency_manager(); + + let read_setup = get_object_zero_copy::prepare_get_object_chunk_read( + &request_context, + &ecstore, + manager, + &request_context.bucket, + &request_context.key, + request_context.rs.clone(), + request_context.part_number, + &request_context.opts, + std::time::Instant::now(), + ) + .await + .unwrap() + .expect("expected chunk fast path"); + + match read_setup.body_source { + GetObjectBodySource::Chunk { path, copy_mode, .. } => { + assert!(matches!(path, GetObjectChunkPath::Direct), "expected direct chunk path"); + #[cfg(unix)] + assert_eq!( + copy_mode, + rustfs_io_metrics::CopyMode::TrueZeroCopy, + "multi-data-shard direct path should preserve shard mmap slices without assembly copies" + ); + #[cfg(not(unix))] + assert_eq!( + copy_mode, + rustfs_io_metrics::CopyMode::SharedBytes, + "multi-data-shard direct path should avoid assembly copies even without mmap" + ); + } + GetObjectBodySource::Reader(_) => panic!("expected chunk body source"), + } + + let usecase = DefaultObjectUsecase::without_context(); + let response = usecase.execute_get_object(req).await.unwrap(); + assert_eq!(response.output.content_length, Some((range_end - range_start + 1) as i64)); + let mut body = response.output.body.expect("expected body"); + let mut collected = Vec::new(); + while let Some(chunk) = body.next().await { + collected.extend_from_slice(&chunk.unwrap()); + } + assert_eq!(collected, payload[range_start as usize..=range_end as usize].to_vec()); +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 1)] +#[serial] +#[ignore = "requires isolated global object layer state"] +async fn execute_get_object_range_marks_reconstructed_path_for_multi_disk_store_with_missing_shard() { + let (disk_paths, ecstore) = setup_direct_chunk_multi_disk_test_env().await; + let bucket = format!("direct-multi-reconstructed-{}", &Uuid::new_v4().simple().to_string()[..8]); + let key = "test/multi-reconstructed.bin"; + let payload_len = 3 * 1024 * 1024 + 137; + let payload: Vec = (0..payload_len).map(|idx| (idx % 251) as u8).collect(); + + create_direct_chunk_test_bucket(&ecstore, &bucket).await; + + let mut reader = ChunkNativePutData::from_vec(payload.clone()); + let put_info = ecstore + .put_object(&bucket, key, &mut reader, &ObjectOptions::default()) + .await + .unwrap(); + assert!(put_info.data_blocks > 1, "expected multi-data-shard object"); + + let object_root = disk_paths[0].join(&bucket).join("test").join("multi-reconstructed.bin"); + let missing_part = find_part_file(&object_root, "part.1").expect("part file on first disk"); + fs::remove_file(&missing_part).await.expect("remove first shard part"); + + let range_start = 123_457_u64; + let range_end = 2 * 1024 * 1024 + 33_333_u64; + let input = GetObjectInput::builder() + .bucket(bucket.clone()) + .key(key.to_string()) + .range(Some(Range::Int { + first: range_start, + last: Some(range_end), + })) + .build() + .unwrap(); + let req = build_request(input, Method::GET); + + let request_context = prepare_get_object_request_context(&req).await.unwrap(); + let manager = get_concurrency_manager(); + + let read_setup = get_object_zero_copy::prepare_get_object_chunk_read( + &request_context, + &ecstore, + manager, + &request_context.bucket, + &request_context.key, + request_context.rs.clone(), + request_context.part_number, + &request_context.opts, + std::time::Instant::now(), + ) + .await + .unwrap() + .expect("expected chunk fast path"); + + match read_setup.body_source { + GetObjectBodySource::Chunk { path, copy_mode, .. } => { + assert!(matches!(path, GetObjectChunkPath::Direct), "expected direct chunk path"); + assert_eq!( + copy_mode, + rustfs_io_metrics::CopyMode::Reconstructed, + "missing data shard should trigger reconstructed chunk path" + ); + } + GetObjectBodySource::Reader(_) => panic!("expected chunk body source"), + } + + let usecase = DefaultObjectUsecase::without_context(); + let response = usecase.execute_get_object(req).await.unwrap(); + assert_eq!(response.output.content_length, Some((range_end - range_start + 1) as i64)); + let mut body = response.output.body.expect("expected body"); + let mut collected = Vec::new(); + while let Some(chunk) = body.next().await { + collected.extend_from_slice(&chunk.unwrap()); + } + assert_eq!(collected, payload[range_start as usize..=range_end as usize].to_vec()); +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 1)] +#[serial] +#[ignore = "requires isolated global object layer state"] +async fn execute_get_object_part_number_marks_direct_path_for_single_disk_store() { + let (_disk_paths, ecstore) = setup_direct_chunk_test_env().await; + let bucket = format!("direct-part-{}", &Uuid::new_v4().simple().to_string()[..8]); + let key = "test/multipart.bin"; + let part_one = vec![11u8; 5 * 1024 * 1024]; + let part_two = vec![22u8; 5 * 1024 * 1024]; + + create_direct_chunk_test_bucket(&ecstore, &bucket).await; + let parts = create_direct_chunk_test_multipart_object(&ecstore, &bucket, key, vec![part_one.clone(), part_two.clone()]).await; + + let input = GetObjectInput::builder() + .bucket(bucket.clone()) + .key(key.to_string()) + .part_number(Some(1)) + .build() + .unwrap(); + let req = build_request(input, Method::GET); + + let request_context = prepare_get_object_request_context(&req).await.unwrap(); + let manager = get_concurrency_manager(); + + let read_setup = get_object_zero_copy::prepare_get_object_chunk_read( + &request_context, + &ecstore, + manager, + &request_context.bucket, + &request_context.key, + request_context.rs.clone(), + request_context.part_number, + &request_context.opts, + std::time::Instant::now(), + ) + .await + .unwrap() + .expect("expected chunk fast path"); + + match read_setup.body_source { + GetObjectBodySource::Chunk { path, copy_mode, .. } => { + assert!(matches!(path, GetObjectChunkPath::Direct), "expected direct chunk path"); + assert_eq!(copy_mode, rustfs_io_metrics::CopyMode::TrueZeroCopy); + } + GetObjectBodySource::Reader(_) => panic!("expected chunk body source"), + } + + let usecase = DefaultObjectUsecase::without_context(); + let response = usecase.execute_get_object(req).await.unwrap(); + assert_eq!(response.output.content_length, Some(parts[0].len() as i64)); + let mut body = response.output.body.expect("expected body"); + let mut collected = Vec::new(); + while let Some(chunk) = body.next().await { + collected.extend_from_slice(&chunk.unwrap()); + } + assert_eq!(collected, parts[0]); +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 1)] +#[serial] +#[ignore = "requires isolated global object layer state"] +async fn execute_get_object_whole_multipart_marks_direct_path_for_single_disk_store() { + let (_disk_paths, ecstore) = setup_direct_chunk_test_env().await; + let bucket = format!("direct-whole-{}", &Uuid::new_v4().simple().to_string()[..8]); + let key = "test/multipart-whole.bin"; + let part_one = vec![11u8; 5 * 1024 * 1024]; + let part_two = vec![22u8; 5 * 1024 * 1024 + 321]; + + create_direct_chunk_test_bucket(&ecstore, &bucket).await; + let parts = create_direct_chunk_test_multipart_object(&ecstore, &bucket, key, vec![part_one.clone(), part_two.clone()]).await; + + let input = GetObjectInput::builder() + .bucket(bucket.clone()) + .key(key.to_string()) + .build() + .unwrap(); + let req = build_request(input, Method::GET); + + let request_context = prepare_get_object_request_context(&req).await.unwrap(); + let manager = get_concurrency_manager(); + + let read_setup = get_object_zero_copy::prepare_get_object_chunk_read( + &request_context, + &ecstore, + manager, + &request_context.bucket, + &request_context.key, + request_context.rs.clone(), + request_context.part_number, + &request_context.opts, + std::time::Instant::now(), + ) + .await + .unwrap() + .expect("expected chunk fast path"); + + match read_setup.body_source { + GetObjectBodySource::Chunk { path, copy_mode, .. } => { + assert!(matches!(path, GetObjectChunkPath::Direct), "expected direct chunk path"); + assert_eq!(copy_mode, rustfs_io_metrics::CopyMode::TrueZeroCopy); + } + GetObjectBodySource::Reader(_) => panic!("expected chunk body source"), + } + + let usecase = DefaultObjectUsecase::without_context(); + let response = usecase.execute_get_object(req).await.unwrap(); + assert_eq!(response.output.content_length, Some(parts.iter().map(|part| part.len() as i64).sum())); + let mut body = response.output.body.expect("expected body"); + let mut collected = Vec::new(); + while let Some(chunk) = body.next().await { + collected.extend_from_slice(&chunk.unwrap()); + } + + let mut expected = Vec::with_capacity(parts.iter().map(Vec::len).sum()); + for part in &parts { + expected.extend_from_slice(part); + } + assert_eq!(collected, expected); +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 1)] +#[serial] +#[ignore = "requires isolated global object layer state"] +async fn execute_get_object_part_number_marks_reconstructed_path_for_multi_disk_store_with_missing_shard() { + let (disk_paths, ecstore) = setup_direct_chunk_multi_disk_test_env().await; + let bucket = format!("direct-part-reconstructed-{}", &Uuid::new_v4().simple().to_string()[..8]); + let key = "test/multipart-reconstructed.bin"; + let part_one: Vec = (0..(5 * 1024 * 1024)).map(|idx| (idx % 251) as u8).collect(); + let part_two: Vec = (0..(5 * 1024 * 1024 + 137)).map(|idx| ((idx + 11) % 251) as u8).collect(); + let part_three: Vec = (0..(1024 * 1024 + 77)).map(|idx| ((idx + 29) % 251) as u8).collect(); + + create_direct_chunk_test_bucket(&ecstore, &bucket).await; + let parts = + create_direct_chunk_test_multipart_object(&ecstore, &bucket, key, vec![part_one.clone(), part_two.clone(), part_three]) + .await; + + let input = GetObjectInput::builder() + .bucket(bucket.clone()) + .key(key.to_string()) + .part_number(Some(1)) + .build() + .unwrap(); + let req = build_request(input, Method::GET); + + let request_context = prepare_get_object_request_context(&req).await.unwrap(); + let manager = get_concurrency_manager(); + let read_setup = + select_reconstructed_chunk_read(&disk_paths, &bucket, key, "part.1", &ecstore, manager, &request_context).await; + + match read_setup.body_source { + GetObjectBodySource::Chunk { path, copy_mode, .. } => { + assert!(matches!(path, GetObjectChunkPath::Direct), "expected direct chunk path"); + assert_eq!( + copy_mode, + rustfs_io_metrics::CopyMode::Reconstructed, + "missing data shard in multipart part should trigger reconstructed chunk path" + ); + } + GetObjectBodySource::Reader(_) => panic!("expected chunk body source"), + } + + let usecase = DefaultObjectUsecase::without_context(); + let response = usecase.execute_get_object(req).await.unwrap(); + assert_eq!(response.output.content_length, Some(parts[0].len() as i64)); + let mut body = response.output.body.expect("expected body"); + let mut collected = Vec::new(); + while let Some(chunk) = body.next().await { + collected.extend_from_slice(&chunk.unwrap()); + } + assert_eq!(collected, parts[0]); +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 1)] +#[serial] +#[ignore = "requires isolated global object layer state"] +async fn execute_get_object_part_number_marks_reconstructed_path_for_second_multipart_part_with_missing_shard() { + let (disk_paths, ecstore) = setup_direct_chunk_multi_disk_test_env().await; + let bucket = format!("direct-part2-reconstructed-{}", &Uuid::new_v4().simple().to_string()[..8]); + let key = "test/multipart-reconstructed-second-part.bin"; + let part_one: Vec = (0..(5 * 1024 * 1024)).map(|idx| (idx % 251) as u8).collect(); + let part_two: Vec = (0..(5 * 1024 * 1024 + 137)).map(|idx| ((idx + 11) % 251) as u8).collect(); + let part_three: Vec = (0..(1024 * 1024 + 77)).map(|idx| ((idx + 29) % 251) as u8).collect(); + + create_direct_chunk_test_bucket(&ecstore, &bucket).await; + let parts = + create_direct_chunk_test_multipart_object(&ecstore, &bucket, key, vec![part_one, part_two.clone(), part_three]).await; + + let input = GetObjectInput::builder() + .bucket(bucket.clone()) + .key(key.to_string()) + .part_number(Some(2)) + .build() + .unwrap(); + let req = build_request(input, Method::GET); + + let request_context = prepare_get_object_request_context(&req).await.unwrap(); + let manager = get_concurrency_manager(); + let read_setup = + select_reconstructed_chunk_read(&disk_paths, &bucket, key, "part.2", &ecstore, manager, &request_context).await; + + match read_setup.body_source { + GetObjectBodySource::Chunk { path, copy_mode, .. } => { + assert!(matches!(path, GetObjectChunkPath::Direct), "expected direct chunk path"); + assert_eq!( + copy_mode, + rustfs_io_metrics::CopyMode::Reconstructed, + "missing data shard in multipart part 2 should trigger reconstructed chunk path" + ); + } + GetObjectBodySource::Reader(_) => panic!("expected chunk body source"), + } + + let usecase = DefaultObjectUsecase::without_context(); + let response = usecase.execute_get_object(req).await.unwrap(); + assert_eq!(response.output.content_length, Some(parts[1].len() as i64)); + let mut body = response.output.body.expect("expected body"); + let mut collected = Vec::new(); + while let Some(chunk) = body.next().await { + collected.extend_from_slice(&chunk.unwrap()); + } + assert_eq!(collected, parts[1]); +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 1)] +#[serial] +#[ignore = "requires isolated global object layer state"] +async fn execute_get_object_part_number_marks_reconstructed_path_for_final_multipart_part_with_missing_shard() { + let (disk_paths, ecstore) = setup_direct_chunk_multi_disk_test_env().await; + let bucket = format!("direct-part3-reconstructed-{}", &Uuid::new_v4().simple().to_string()[..8]); + let key = "test/multipart-reconstructed-final-part.bin"; + let part_one: Vec = (0..(5 * 1024 * 1024)).map(|idx| (idx % 251) as u8).collect(); + let part_two: Vec = (0..(5 * 1024 * 1024 + 137)).map(|idx| ((idx + 11) % 251) as u8).collect(); + let part_three: Vec = (0..(1024 * 1024 + 77)).map(|idx| ((idx + 29) % 251) as u8).collect(); + + create_direct_chunk_test_bucket(&ecstore, &bucket).await; + let parts = + create_direct_chunk_test_multipart_object(&ecstore, &bucket, key, vec![part_one, part_two, part_three.clone()]).await; + + let input = GetObjectInput::builder() + .bucket(bucket.clone()) + .key(key.to_string()) + .part_number(Some(3)) + .build() + .unwrap(); + let req = build_request(input, Method::GET); + + let request_context = prepare_get_object_request_context(&req).await.unwrap(); + let manager = get_concurrency_manager(); + let read_setup = + select_reconstructed_chunk_read(&disk_paths, &bucket, key, "part.3", &ecstore, manager, &request_context).await; + + match read_setup.body_source { + GetObjectBodySource::Chunk { path, copy_mode, .. } => { + assert!(matches!(path, GetObjectChunkPath::Direct), "expected direct chunk path"); + assert_eq!( + copy_mode, + rustfs_io_metrics::CopyMode::Reconstructed, + "missing data shard in the final multipart part should trigger reconstructed chunk path" + ); + } + GetObjectBodySource::Reader(_) => panic!("expected chunk body source"), + } + + let usecase = DefaultObjectUsecase::without_context(); + let response = usecase.execute_get_object(req).await.unwrap(); + assert_eq!(response.output.content_length, Some(parts[2].len() as i64)); + let mut body = response.output.body.expect("expected body"); + let mut collected = Vec::new(); + while let Some(chunk) = body.next().await { + collected.extend_from_slice(&chunk.unwrap()); + } + assert_eq!(collected, part_three); +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 1)] +#[serial] +#[ignore = "requires isolated global object layer state"] +async fn execute_get_object_whole_multipart_marks_reconstructed_path_for_missing_middle_part_shard() { + let (disk_paths, ecstore) = setup_direct_chunk_multi_disk_test_env().await; + let bucket = format!("direct-whole-multipart-reconstructed-{}", &Uuid::new_v4().simple().to_string()[..8]); + let key = "test/multipart-whole-reconstructed.bin"; + let part_one: Vec = (0..(5 * 1024 * 1024)).map(|idx| (idx % 251) as u8).collect(); + let part_two: Vec = (0..(5 * 1024 * 1024 + 257)).map(|idx| ((idx + 17) % 251) as u8).collect(); + let part_three: Vec = (0..(1024 * 1024 + 211)).map(|idx| ((idx + 33) % 251) as u8).collect(); + + create_direct_chunk_test_bucket(&ecstore, &bucket).await; + let parts = create_direct_chunk_test_multipart_object(&ecstore, &bucket, key, vec![part_one, part_two, part_three]).await; + let input = GetObjectInput::builder() + .bucket(bucket.clone()) + .key(key.to_string()) + .build() + .unwrap(); + let req = build_request(input, Method::GET); + + let request_context = prepare_get_object_request_context(&req).await.unwrap(); + let manager = get_concurrency_manager(); + let read_setup = + select_reconstructed_chunk_read(&disk_paths, &bucket, key, "part.2", &ecstore, manager, &request_context).await; + + let mut direct_stream = match read_setup.body_source { + GetObjectBodySource::Chunk { path, copy_mode, stream } => { + assert!(matches!(path, GetObjectChunkPath::Direct), "expected direct chunk path"); + assert_eq!( + copy_mode, + rustfs_io_metrics::CopyMode::Reconstructed, + "whole multipart GET should keep the reconstructed chunk fast path when a middle part is missing a shard" + ); + stream + } + GetObjectBodySource::Reader(_) => panic!("expected chunk body source"), + }; + + let mut expected = Vec::with_capacity(parts.iter().map(Vec::len).sum()); + for part in &parts { + expected.extend_from_slice(part); + } + + let mut direct_collected = Vec::new(); + while let Some(chunk) = direct_stream.next().await { + direct_collected.extend_from_slice(chunk.unwrap().as_bytes().as_ref()); + } + assert_eq!( + direct_collected.len(), + expected.len(), + "prepared chunk stream reconstructed whole multipart length mismatch" + ); + let direct_first_diff = direct_collected + .iter() + .zip(expected.iter()) + .position(|(left, right)| left != right); + assert_eq!( + direct_first_diff, None, + "prepared chunk stream reconstructed whole multipart first diff at {:?}", + direct_first_diff + ); + assert_eq!( + direct_collected, expected, + "prepared chunk stream should cover the whole reconstructed multipart object" + ); + + let usecase = DefaultObjectUsecase::without_context(); + let response = usecase.execute_get_object(req).await.unwrap(); + assert_eq!(response.output.content_length, Some(expected.len() as i64)); + let mut body = response.output.body.expect("expected body"); + let mut collected = Vec::new(); + while let Some(chunk) = body.next().await { + collected.extend_from_slice(&chunk.unwrap()); + } + assert_eq!(collected.len(), expected.len(), "whole multipart reconstructed GET length mismatch"); + let first_diff = collected.iter().zip(expected.iter()).position(|(left, right)| left != right); + assert_eq!(first_diff, None, "whole multipart reconstructed GET first diff at {:?}", first_diff); + assert_eq!(collected, expected); +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 1)] +#[serial] +#[ignore = "requires isolated global object layer state"] +async fn execute_get_object_multipart_range_marks_reconstructed_path_for_missing_shard_part() { + let (disk_paths, ecstore) = setup_direct_chunk_multi_disk_test_env().await; + let bucket = format!("direct-multipart-range-reconstructed-{}", &Uuid::new_v4().simple().to_string()[..8]); + let key = "test/multipart-range-reconstructed.bin"; + let part_one: Vec = (0..(5 * 1024 * 1024)).map(|idx| (idx % 251) as u8).collect(); + let part_two: Vec = (0..(5 * 1024 * 1024 + 257)).map(|idx| ((idx + 17) % 251) as u8).collect(); + let part_three: Vec = (0..(1024 * 1024 + 211)).map(|idx| ((idx + 33) % 251) as u8).collect(); + + create_direct_chunk_test_bucket(&ecstore, &bucket).await; + let parts = + create_direct_chunk_test_multipart_object(&ecstore, &bucket, key, vec![part_one.clone(), part_two.clone(), part_three]) + .await; + + let range_start = 1024_u64; + let range_end = 64 * 1024 + 2048; + let input = GetObjectInput::builder() + .bucket(bucket.clone()) + .key(key.to_string()) + .range(Some(Range::Int { + first: range_start, + last: Some(range_end), + })) + .build() + .unwrap(); + let req = build_request(input, Method::GET); + + let request_context = prepare_get_object_request_context(&req).await.unwrap(); + let manager = get_concurrency_manager(); + let read_setup = + select_reconstructed_chunk_read(&disk_paths, &bucket, key, "part.1", &ecstore, manager, &request_context).await; + + match read_setup.body_source { + GetObjectBodySource::Chunk { path, copy_mode, .. } => { + assert!(matches!(path, GetObjectChunkPath::Direct), "expected direct chunk path"); + assert_eq!( + copy_mode, + rustfs_io_metrics::CopyMode::Reconstructed, + "multipart partial range should stay on reconstructed fast path when the covered part has a missing shard" + ); + } + GetObjectBodySource::Reader(_) => panic!("expected chunk body source"), + } + + let mut expected = Vec::with_capacity(parts.iter().map(Vec::len).sum()); + for part in &parts { + expected.extend_from_slice(part); + } + + let usecase = DefaultObjectUsecase::without_context(); + let response = usecase.execute_get_object(req).await.unwrap(); + assert_eq!(response.output.content_length, Some((range_end - range_start + 1) as i64)); + let mut body = response.output.body.expect("expected body"); + let mut collected = Vec::new(); + while let Some(chunk) = body.next().await { + collected.extend_from_slice(&chunk.unwrap()); + } + assert_eq!(collected, expected[range_start as usize..=range_end as usize].to_vec()); +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 1)] +#[serial] +#[ignore = "requires isolated global object layer state"] +async fn execute_get_object_cross_part_multipart_range_marks_reconstructed_path_for_missing_next_part_shard() { + let (disk_paths, ecstore) = setup_direct_chunk_multi_disk_test_env().await; + let bucket = format!("direct-multipart-cross-range-{}", &Uuid::new_v4().simple().to_string()[..8]); + let key = "test/multipart-cross-range-reconstructed.bin"; + let part_one: Vec = (0..(5 * 1024 * 1024)).map(|idx| (idx % 251) as u8).collect(); + let part_two: Vec = (0..(5 * 1024 * 1024 + 257)).map(|idx| ((idx + 17) % 251) as u8).collect(); + let part_three: Vec = (0..(1024 * 1024 + 211)).map(|idx| ((idx + 33) % 251) as u8).collect(); + + create_direct_chunk_test_bucket(&ecstore, &bucket).await; + let parts = create_direct_chunk_test_multipart_object(&ecstore, &bucket, key, vec![part_one, part_two, part_three]).await; + + let part_one_len = parts[0].len() as u64; + let range_start = part_one_len - 32 * 1024; + let range_end = part_one_len + 96 * 1024; + let input = GetObjectInput::builder() + .bucket(bucket.clone()) + .key(key.to_string()) + .range(Some(Range::Int { + first: range_start, + last: Some(range_end), + })) + .build() + .unwrap(); + let req = build_request(input, Method::GET); + + let request_context = prepare_get_object_request_context(&req).await.unwrap(); + let manager = get_concurrency_manager(); + let read_setup = + select_reconstructed_chunk_read(&disk_paths, &bucket, key, "part.2", &ecstore, manager, &request_context).await; + + match read_setup.body_source { + GetObjectBodySource::Chunk { path, copy_mode, .. } => { + assert!(matches!(path, GetObjectChunkPath::Direct), "expected direct chunk path"); + assert_eq!( + copy_mode, + rustfs_io_metrics::CopyMode::Reconstructed, + "cross-part multipart range should keep the reconstructed chunk fast path when the next part has a missing shard" + ); + } + GetObjectBodySource::Reader(_) => panic!("expected chunk body source"), + } + + let mut expected = Vec::with_capacity(parts.iter().map(Vec::len).sum()); + for part in &parts { + expected.extend_from_slice(part); + } + + let usecase = DefaultObjectUsecase::without_context(); + let response = usecase.execute_get_object(req).await.unwrap(); + assert_eq!(response.output.content_length, Some((range_end - range_start + 1) as i64)); + let mut body = response.output.body.expect("expected body"); + let mut collected = Vec::new(); + while let Some(chunk) = body.next().await { + collected.extend_from_slice(&chunk.unwrap()); + } + assert_eq!(collected, expected[range_start as usize..=range_end as usize].to_vec()); +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 1)] +#[serial] +#[ignore = "requires isolated global object layer state"] +async fn execute_get_object_cross_part_multipart_range_marks_reconstructed_path_for_missing_final_part_shard() { + let (disk_paths, ecstore) = setup_direct_chunk_multi_disk_test_env().await; + let bucket = format!("direct-multipart-final-cross-range-{}", &Uuid::new_v4().simple().to_string()[..8]); + let key = "test/multipart-final-cross-range-reconstructed.bin"; + let part_one: Vec = (0..(5 * 1024 * 1024)).map(|idx| (idx % 251) as u8).collect(); + let part_two: Vec = (0..(5 * 1024 * 1024 + 257)).map(|idx| ((idx + 17) % 251) as u8).collect(); + let part_three: Vec = (0..(1024 * 1024 + 211)).map(|idx| ((idx + 33) % 251) as u8).collect(); + + create_direct_chunk_test_bucket(&ecstore, &bucket).await; + let parts = create_direct_chunk_test_multipart_object(&ecstore, &bucket, key, vec![part_one, part_two, part_three]).await; + + let part_two_end = (parts[0].len() + parts[1].len()) as u64; + let range_start = part_two_end - 32 * 1024; + let range_end = part_two_end + 96 * 1024; + let input = GetObjectInput::builder() + .bucket(bucket.clone()) + .key(key.to_string()) + .range(Some(Range::Int { + first: range_start, + last: Some(range_end), + })) + .build() + .unwrap(); + let req = build_request(input, Method::GET); + + let request_context = prepare_get_object_request_context(&req).await.unwrap(); + let manager = get_concurrency_manager(); + let read_setup = + select_reconstructed_chunk_read(&disk_paths, &bucket, key, "part.3", &ecstore, manager, &request_context).await; + + match read_setup.body_source { + GetObjectBodySource::Chunk { path, copy_mode, .. } => { + assert!(matches!(path, GetObjectChunkPath::Direct), "expected direct chunk path"); + assert_eq!( + copy_mode, + rustfs_io_metrics::CopyMode::Reconstructed, + "cross-part multipart range into the final part should keep the reconstructed chunk fast path" + ); + } + GetObjectBodySource::Reader(_) => panic!("expected chunk body source"), + } + + let mut expected = Vec::with_capacity(parts.iter().map(Vec::len).sum()); + for part in &parts { + expected.extend_from_slice(part); + } + + let usecase = DefaultObjectUsecase::without_context(); + let response = usecase.execute_get_object(req).await.unwrap(); + assert_eq!(response.output.content_length, Some((range_end - range_start + 1) as i64)); + let mut body = response.output.body.expect("expected body"); + let mut collected = Vec::new(); + while let Some(chunk) = body.next().await { + collected.extend_from_slice(&chunk.unwrap()); + } + assert_eq!(collected, expected[range_start as usize..=range_end as usize].to_vec()); +} diff --git a/rustfs/src/error.rs b/rustfs/src/error.rs index 14a2d8352..95b40dfb7 100644 --- a/rustfs/src/error.rs +++ b/rustfs/src/error.rs @@ -259,6 +259,9 @@ impl From for ApiError { fn from(err: std::io::Error) -> Self { // Check if the error is a ChecksumMismatch (BadDigest) if let Some(inner) = err.get_ref() { + if let Some(storage_error) = inner.downcast_ref::() { + return storage_error.clone().into(); + } if inner.downcast_ref::().is_some() { return ApiError { code: S3ErrorCode::BadDigest, @@ -552,4 +555,15 @@ mod tests { // This is expected because ApiError is not a typical Error implementation assert!(error.source().is_none()); } + + #[test] + fn test_api_error_from_io_error_unwraps_invalid_range_storage_error() { + let io_error = std::io::Error::from(StorageError::InvalidRangeSpec("range invalid".to_string())); + + let api_error: ApiError = io_error.into(); + + assert_eq!(api_error.code, S3ErrorCode::InvalidRange); + assert_eq!(api_error.message, ApiError::error_code_to_message(&S3ErrorCode::InvalidRange)); + assert!(api_error.source.is_some()); + } } diff --git a/rustfs/src/server/http.rs b/rustfs/src/server/http.rs index 58a5b7af6..90b940e86 100644 --- a/rustfs/src/server/http.rs +++ b/rustfs/src/server/http.rs @@ -426,7 +426,7 @@ pub async fn start_http_server( } } }; - + #[allow(unused)] let socket_ref = SockRef::from(&socket); // ── POST-ACCEPT SOCKET SYSCALLS ── diff --git a/rustfs/src/storage/concurrency/object_cache.rs b/rustfs/src/storage/concurrency/object_cache.rs index b8f715e3e..676773f72 100644 --- a/rustfs/src/storage/concurrency/object_cache.rs +++ b/rustfs/src/storage/concurrency/object_cache.rs @@ -32,6 +32,7 @@ use hashbrown::HashMap; use moka::future::Cache; use rustfs_config::MI_B; +use rustfs_object_io::get::GetObjectCacheWriteback; use std::sync::Arc; use std::sync::atomic::{AtomicU64, Ordering}; use std::time::{Duration, Instant}; @@ -1063,6 +1064,13 @@ pub struct CachedGetObject { pub replication_status: Option, /// User-defined metadata (x-amz-meta-*) pub user_metadata: std::collections::HashMap, + /// Additional checksum metadata persisted with cached GET responses + pub checksum_crc32: Option, + pub checksum_crc32c: Option, + pub checksum_sha1: Option, + pub checksum_sha256: Option, + pub checksum_crc64nvme: Option, + pub checksum_type: Option, /// When this object was cached (for internal use, automatically set) #[allow(dead_code)] cached_at: Option, @@ -1089,6 +1097,12 @@ impl Default for CachedGetObject { tag_count: None, replication_status: None, user_metadata: std::collections::HashMap::new(), + checksum_crc32: None, + checksum_crc32c: None, + checksum_sha1: None, + checksum_sha256: None, + checksum_crc64nvme: None, + checksum_type: None, cached_at: None, access_count: Arc::new(AtomicU64::new(0)), } @@ -1109,6 +1123,35 @@ impl CachedGetObject { } } + /// Consume a GET cache writeback payload into the cache-owned representation. + pub fn from_get_object_cache_writeback(writeback: GetObjectCacheWriteback) -> Self { + Self { + body: writeback.body, + content_length: writeback.content_length, + content_type: writeback.content_type, + e_tag: writeback.e_tag, + last_modified: writeback.last_modified, + expires: writeback.expires, + cache_control: writeback.cache_control, + content_disposition: writeback.content_disposition, + content_encoding: writeback.content_encoding, + content_language: writeback.content_language, + storage_class: writeback.storage_class, + version_id: writeback.version_id, + delete_marker: writeback.delete_marker, + user_metadata: writeback.user_metadata, + checksum_crc32: writeback.checksum_crc32, + checksum_crc32c: writeback.checksum_crc32c, + checksum_sha1: writeback.checksum_sha1, + checksum_sha256: writeback.checksum_sha256, + checksum_crc64nvme: writeback.checksum_crc64nvme, + checksum_type: writeback.checksum_type, + cached_at: Some(Instant::now()), + access_count: Arc::new(AtomicU64::new(0)), + ..Default::default() + } + } + /// Builder method to set content_type pub fn with_content_type(mut self, content_type: String) -> Self { self.content_type = Some(content_type); @@ -1881,6 +1924,57 @@ mod cached_object_tests { assert_eq!(obj.user_metadata.get("x-amz-meta-custom"), Some(&"value".to_string())); } + #[test] + fn test_cached_get_object_from_get_object_cache_writeback() { + let body = Arc::new(Bytes::from("test data")); + let obj = CachedGetObject::from_get_object_cache_writeback(GetObjectCacheWriteback { + body: Arc::clone(&body), + content_length: 9, + content_type: Some("text/plain".to_string()), + content_encoding: Some("gzip".to_string()), + cache_control: Some("max-age=3600".to_string()), + content_disposition: Some("attachment".to_string()), + content_language: Some("en-US".to_string()), + expires: Some("2024-12-31T23:59:59Z".to_string()), + storage_class: Some("STANDARD".to_string()), + version_id: Some("null".to_string()), + delete_marker: false, + user_metadata: { + let mut metadata = std::collections::HashMap::new(); + metadata.insert("custom-key".to_string(), "value".to_string()); + metadata + }, + e_tag: Some("\"abc123\"".to_string()), + last_modified: Some("2024-01-01T12:00:00Z".to_string()), + checksum_crc32: Some("crc32".to_string()), + checksum_crc32c: None, + checksum_sha1: None, + checksum_sha256: None, + checksum_crc64nvme: None, + checksum_type: Some(s3s::dto::ChecksumType::from_static(s3s::dto::ChecksumType::FULL_OBJECT)), + }); + + assert_eq!(obj.content_length, 9); + assert_eq!(obj.content_type.as_deref(), Some("text/plain")); + assert_eq!(obj.content_encoding.as_deref(), Some("gzip")); + assert_eq!(obj.cache_control.as_deref(), Some("max-age=3600")); + assert_eq!(obj.content_disposition.as_deref(), Some("attachment")); + assert_eq!(obj.content_language.as_deref(), Some("en-US")); + assert_eq!(obj.expires.as_deref(), Some("2024-12-31T23:59:59Z")); + assert_eq!(obj.storage_class.as_deref(), Some("STANDARD")); + assert_eq!(obj.version_id.as_deref(), Some("null")); + assert!(!obj.delete_marker); + assert_eq!(obj.user_metadata.get("custom-key").map(String::as_str), Some("value")); + assert_eq!(obj.e_tag.as_deref(), Some("\"abc123\"")); + assert_eq!(obj.last_modified.as_deref(), Some("2024-01-01T12:00:00Z")); + assert_eq!(obj.checksum_crc32.as_deref(), Some("crc32")); + assert_eq!( + obj.checksum_type, + Some(s3s::dto::ChecksumType::from_static(s3s::dto::ChecksumType::FULL_OBJECT)) + ); + assert!(Arc::ptr_eq(&obj.body, &body)); + } + #[test] fn test_cached_get_object_size() { let obj = CachedGetObject::new(Bytes::from("test"), 4); diff --git a/rustfs/src/storage/ecfs.rs b/rustfs/src/storage/ecfs.rs index 28679492f..d488a0209 100644 --- a/rustfs/src/storage/ecfs.rs +++ b/rustfs/src/storage/ecfs.rs @@ -29,7 +29,6 @@ use rustfs_ecstore::{ use rustfs_s3_common::{S3Operation, record_s3_op}; use s3s::{S3, S3Error, S3ErrorCode, S3Request, S3Response, S3Result, dto::*, s3_error}; use std::fmt::Debug; -use tokio::io::{AsyncRead, AsyncSeek}; use tracing::{debug, error, instrument, warn}; use uuid::Uuid; @@ -44,44 +43,6 @@ pub(crate) struct ListObjectUnorderedQuery { pub(crate) allow_unordered: Option, } -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())) - } -} - impl FS { pub fn new() -> Self { rustfs_s3_common::init_s3_metrics(); diff --git a/scripts/bench-small-put-local.sh b/scripts/bench-small-put-local.sh new file mode 100755 index 000000000..eb10b4a70 --- /dev/null +++ b/scripts/bench-small-put-local.sh @@ -0,0 +1,152 @@ +#!/usr/bin/env bash +set -euo pipefail + +ROOT_DIR="$(cd "$(dirname "$0")/.." && pwd)" +RUSTFS_BIN="${ROOT_DIR}/target/debug/rustfs" +WARP_BIN="${WARP_BIN:-warp}" +MC_BIN="${MC_BIN:-mc}" + +PORT="${RUSTFS_BENCH_PORT:-9900}" +CONSOLE_PORT="${RUSTFS_BENCH_CONSOLE_PORT:-9901}" +HOST="127.0.0.1:${PORT}" +ACCESS_KEY="${RUSTFS_BENCH_ACCESS_KEY:-rustfsadmin}" +SECRET_KEY="${RUSTFS_BENCH_SECRET_KEY:-rustfsadmin}" +BENCH_ROOT="${ROOT_DIR}/target/benchmarks" +RUN_ID="${RUSTFS_BENCH_RUN_ID:-small-put-$(date +%Y%m%d-%H%M%S)}" +RUN_DIR="${BENCH_ROOT}/${RUN_ID}" +VOLUME_DIR="${RUN_DIR}/volumes" +LOG_DIR="${RUN_DIR}/logs" +MC_CONFIG_DIR="${RUN_DIR}/mc-config" +RESULTS_MD="${RUN_DIR}/RESULTS.md" +SERVER_LOG="${LOG_DIR}/rustfs.log" +WARP_BUCKET="${RUSTFS_BENCH_BUCKET:-warp-small-put-benchmark}" +WARP_DURATION="${RUSTFS_BENCH_DURATION:-20s}" +WARP_CONCURRENCY="${RUSTFS_BENCH_CONCURRENCY:-32}" + +mkdir -p "${VOLUME_DIR}" "${LOG_DIR}" "${MC_CONFIG_DIR}" +mkdir -p "${VOLUME_DIR}"/disk{1..4} + +# Verify mc (MinIO Client) is available +if ! command -v "${MC_BIN}" >/dev/null 2>&1; then + echo "Error: 'mc' (MinIO Client) not found. MC_BIN=${MC_BIN}" >&2 + echo "" >&2 + echo "Install mc with one of the following:" >&2 + echo " macOS (Homebrew): brew install minio/stable/mc" >&2 + echo " Go: go install github.com/minio/mc@latest" >&2 + echo " Linux (amd64): curl -sSL https://dl.min.io/client/mc/release/linux-amd64/mc -o /usr/local/bin/mc && chmod +x /usr/local/bin/mc" >&2 + echo " Linux (arm64): curl -sSL https://dl.min.io/client/mc/release/linux-arm64/mc -o /usr/local/bin/mc && chmod +x /usr/local/bin/mc" >&2 + echo "" >&2 + echo "Or set MC_BIN=/path/to/mc before running this script." >&2 + exit 1 +fi + +# Verify warp is available +if ! command -v "${WARP_BIN}" >/dev/null 2>&1; then + echo "Error: 'warp' benchmark tool not found. WARP_BIN=${WARP_BIN}" >&2 + echo "" >&2 + echo "Install warp with one of the following:" >&2 + echo " Go: go install github.com/minio/warp@latest" >&2 + echo " Linux (amd64): curl -sSL https://github.com/minio/warp/releases/latest/download/warp_Linux_x86_64.tar.gz | tar -xz -C /usr/local/bin warp" >&2 + echo " macOS (amd64): curl -sSL https://github.com/minio/warp/releases/latest/download/warp_Darwin_x86_64.tar.gz | tar -xz -C /usr/local/bin warp" >&2 + echo " macOS (arm64): curl -sSL https://github.com/minio/warp/releases/latest/download/warp_Darwin_arm64.tar.gz | tar -xz -C /usr/local/bin warp" >&2 + echo "" >&2 + echo "Or set WARP_BIN=/path/to/warp before running this script." >&2 + exit 1 +fi + +SIZES=( + "4KiB" + "16KiB" + "64KiB" + "256KiB" + "1MiB" +) + +cleanup() { + if [[ -n "${SERVER_PID:-}" ]] && kill -0 "${SERVER_PID}" >/dev/null 2>&1; then + kill "${SERVER_PID}" >/dev/null 2>&1 || true + wait "${SERVER_PID}" >/dev/null 2>&1 || true + fi +} + +trap cleanup EXIT + +echo "Starting local RustFS benchmark server..." +( + cd "${ROOT_DIR}" + export RUST_LOG="${RUST_LOG:-warn}" + export RUSTFS_VOLUMES="${VOLUME_DIR}/disk{1...4}" + export RUSTFS_ADDRESS=":${PORT}" + export RUSTFS_CONSOLE_ENABLE=false + export RUSTFS_CONSOLE_ADDRESS=":${CONSOLE_PORT}" + export RUSTFS_ROOT_USER="${ACCESS_KEY}" + export RUSTFS_ROOT_PASSWORD="${SECRET_KEY}" + export RUSTFS_OBJECT_CACHE_ENABLE=false + "${RUSTFS_BIN}" server >"${SERVER_LOG}" 2>&1 +) & +SERVER_PID=$! + +echo "Waiting for RustFS server to become ready..." +for _ in $(seq 1 60); do + if curl -fsS "http://${HOST}/health/ready" >/dev/null 2>&1 || curl -fsS "http://${HOST}/minio/health/ready" >/dev/null 2>&1; then + break + fi + sleep 1 +done + +if ! curl -fsS "http://${HOST}/health/ready" >/dev/null 2>&1 && ! curl -fsS "http://${HOST}/minio/health/ready" >/dev/null 2>&1; then + echo "RustFS did not become ready in time. See ${SERVER_LOG}" >&2 + exit 1 +fi + +echo "Preparing mc alias..." +MC_CONFIG_DIR="${MC_CONFIG_DIR}" "${MC_BIN}" alias set bench "http://${HOST}" "${ACCESS_KEY}" "${SECRET_KEY}" >/dev/null +MC_CONFIG_DIR="${MC_CONFIG_DIR}" "${MC_BIN}" mb --ignore-existing "bench/${WARP_BUCKET}" >/dev/null + +{ + echo "# Small PUT Benchmark Results" + echo + echo "- Run ID: \`${RUN_ID}\`" + echo "- Host: \`${HOST}\`" + echo "- Duration per size: \`${WARP_DURATION}\`" + echo "- Concurrency: \`${WARP_CONCURRENCY}\`" + echo "- Bucket: \`${WARP_BUCKET}\`" + echo +} >"${RESULTS_MD}" + +for size in "${SIZES[@]}"; do + echo + echo "Running warp PUT benchmark for size ${size}..." + size_slug="$(echo "${size}" | tr '[:upper:]' '[:lower:]')" + OUT_FILE="${LOG_DIR}/warp-put-${size}.log" + BENCHDATA_FILE="${RUN_DIR}/warp-put-${size}.csv.zst" + "${WARP_BIN}" put \ + --no-color \ + --host "${HOST}" \ + --access-key "${ACCESS_KEY}" \ + --secret-key "${SECRET_KEY}" \ + --bucket "${WARP_BUCKET}" \ + --obj.size "${size}" \ + --duration "${WARP_DURATION}" \ + --concurrent "${WARP_CONCURRENCY}" \ + --prefix "${RUN_ID}/${size_slug}" \ + --disable-multipart \ + --noclear \ + --benchdata "${BENCHDATA_FILE}" \ + --insecure \ + >"${OUT_FILE}" 2>&1 + + { + echo "## ${size}" + echo + echo '```text' + tail -n 20 "${OUT_FILE}" + echo '```' + echo + } >>"${RESULTS_MD}" +done + +echo +echo "Benchmark completed." +echo "Results: ${RESULTS_MD}" +echo "Server log: ${SERVER_LOG}" diff --git a/scripts/bench-small-put-mc-local.sh b/scripts/bench-small-put-mc-local.sh new file mode 100755 index 000000000..d7f716a76 --- /dev/null +++ b/scripts/bench-small-put-mc-local.sh @@ -0,0 +1,300 @@ +#!/usr/bin/env bash +set -euo pipefail + +ROOT_DIR="$(cd "$(dirname "$0")/.." && pwd)" +RUSTFS_BIN="${RUSTFS_BIN:-${ROOT_DIR}/target/debug/rustfs}" +MC_BIN="${MC_BIN:-mc}" + +PORT="${RUSTFS_BENCH_PORT:-9910}" +CONSOLE_PORT="${RUSTFS_BENCH_CONSOLE_PORT:-9911}" +HOST="127.0.0.1:${PORT}" +ACCESS_KEY="${RUSTFS_BENCH_ACCESS_KEY:-rustfsadmin}" +SECRET_KEY="${RUSTFS_BENCH_SECRET_KEY:-rustfsadmin}" + +OPS_PER_SIZE="${RUSTFS_BENCH_OPS_PER_SIZE:-0}" +BENCH_SECONDS="${RUSTFS_BENCH_SECONDS:-10}" +CONCURRENCY="${RUSTFS_BENCH_CONCURRENCY:-32}" +PUT_TIMEOUT_SECS="${RUSTFS_BENCH_PUT_TIMEOUT_SECS:-15}" +RUN_ID="${RUSTFS_BENCH_RUN_ID:-small-put-mc-$(date +%Y%m%d-%H%M%S)}" +BENCH_ROOT="${ROOT_DIR}/target/benchmarks" +RUN_DIR="${BENCH_ROOT}/${RUN_ID}" +VOLUME_DIR="${RUN_DIR}/volumes" +LOG_DIR="${RUN_DIR}/logs" +FILE_DIR="${RUN_DIR}/files" +MC_CONFIG_DIR="${RUN_DIR}/mc-config" +RESULTS_MD="${RUN_DIR}/RESULTS.md" +SERVER_LOG="${LOG_DIR}/rustfs.log" +BUCKET="${RUSTFS_BENCH_BUCKET:-small-put-benchmark}" + +# Verify mc (MinIO Client) is available +if ! command -v "${MC_BIN}" >/dev/null 2>&1; then + echo "Error: 'mc' (MinIO Client) not found. MC_BIN=${MC_BIN}" >&2 + echo "" >&2 + echo "Install mc with one of the following:" >&2 + echo " macOS (Homebrew): brew install minio/stable/mc" >&2 + echo " Go: go install github.com/minio/mc@latest" >&2 + echo " Linux (amd64): curl -sSL https://dl.min.io/client/mc/release/linux-amd64/mc -o /usr/local/bin/mc && chmod +x /usr/local/bin/mc" >&2 + echo " Linux (arm64): curl -sSL https://dl.min.io/client/mc/release/linux-arm64/mc -o /usr/local/bin/mc && chmod +x /usr/local/bin/mc" >&2 + echo "" >&2 + echo "Or set MC_BIN=/path/to/mc before running this script." >&2 + exit 1 +fi + +mkdir -p "${VOLUME_DIR}" "${LOG_DIR}" "${FILE_DIR}" "${MC_CONFIG_DIR}" +mkdir -p "${VOLUME_DIR}"/disk{1..4} + +SIZES=( + "4KiB" + "16KiB" + "64KiB" + "256KiB" + "1MiB" +) + +size_to_bytes() { + case "$1" in + 4KiB) echo 4096 ;; + 16KiB) echo 16384 ;; + 64KiB) echo 65536 ;; + 256KiB) echo 262144 ;; + 1MiB) echo 1048576 ;; + *) + echo "unsupported size: $1" >&2 + return 1 + ;; + esac +} + +now_secs() { + perl -MTime::HiRes=time -e 'printf "%.6f\n", time' +} + +prepare_sample_file() { + local size_label="$1" + local bytes="$2" + local file_path="${FILE_DIR}/${size_label}.bin" + if [[ ! -f "${file_path}" ]]; then + dd if=/dev/zero of="${file_path}" bs="${bytes}" count=1 status=none + fi + echo "${file_path}" +} + +run_mc_put_with_timeout() { + local sample_file="$1" + local key="$2" + perl -e ' + use strict; + use warnings; + + my $timeout = shift @ARGV; + my $pid = fork(); + die "fork failed: $!" unless defined $pid; + + if ($pid == 0) { + exec @ARGV; + die "exec failed: $!"; + } + + my $timed_out = 0; + local $SIG{ALRM} = sub { + $timed_out = 1; + kill "TERM", $pid; + select undef, undef, undef, 0.2; + kill "KILL", $pid; + }; + + alarm($timeout); + my $waited = waitpid($pid, 0); + alarm(0); + + exit(1) if $timed_out; + exit(($waited == $pid && $? == 0) ? 0 : 1); + ' "${PUT_TIMEOUT_SECS}" env MC_CONFIG_DIR="${MC_CONFIG_DIR}" "${MC_BIN}" put --quiet --disable-multipart "${sample_file}" "${key}" >/dev/null 2>&1 +} + +export -f run_mc_put_with_timeout + +cleanup() { + if [[ -n "${SERVER_PID:-}" ]] && kill -0 "${SERVER_PID}" >/dev/null 2>&1; then + kill "${SERVER_PID}" >/dev/null 2>&1 || true + wait "${SERVER_PID}" >/dev/null 2>&1 || true + fi +} + +trap cleanup EXIT + +echo "Starting local RustFS benchmark server..." +( + cd "${ROOT_DIR}" + export RUST_LOG="${RUST_LOG:-warn}" + export RUSTFS_VOLUMES="${VOLUME_DIR}/disk{1...4}" + export RUSTFS_ADDRESS=":${PORT}" + export RUSTFS_CONSOLE_ENABLE=false + export RUSTFS_CONSOLE_ADDRESS=":${CONSOLE_PORT}" + export RUSTFS_ROOT_USER="${ACCESS_KEY}" + export RUSTFS_ROOT_PASSWORD="${SECRET_KEY}" + export RUSTFS_OBJECT_CACHE_ENABLE=false + "${RUSTFS_BIN}" server >"${SERVER_LOG}" 2>&1 +) & +SERVER_PID=$! + +echo "Waiting for RustFS server to become ready..." +for _ in $(seq 1 60); do + if curl -fsS "http://${HOST}/health/ready" >/dev/null 2>&1 || curl -fsS "http://${HOST}/minio/health/ready" >/dev/null 2>&1; then + break + fi + sleep 1 +done + +if ! curl -fsS "http://${HOST}/health/ready" >/dev/null 2>&1 && ! curl -fsS "http://${HOST}/minio/health/ready" >/dev/null 2>&1; then + echo "RustFS did not become ready in time. See ${SERVER_LOG}" >&2 + exit 1 +fi + +echo "Preparing mc alias..." +MC_CONFIG_DIR="${MC_CONFIG_DIR}" "${MC_BIN}" alias set bench "http://${HOST}" "${ACCESS_KEY}" "${SECRET_KEY}" >/dev/null +MC_CONFIG_DIR="${MC_CONFIG_DIR}" "${MC_BIN}" mb --ignore-existing "bench/${BUCKET}" >/dev/null + +{ + echo "# Small PUT Benchmark Results" + echo + echo "- Run ID: \`${RUN_ID}\`" + echo "- Host: \`${HOST}\`" + echo "- Bucket: \`${BUCKET}\`" + if [[ "${OPS_PER_SIZE}" -gt 0 ]]; then + echo "- Ops per size: \`${OPS_PER_SIZE}\`" + else + echo "- Duration per size: \`${BENCH_SECONDS}s\`" + fi + echo "- Concurrency: \`${CONCURRENCY}\`" + echo "- Per request timeout: \`${PUT_TIMEOUT_SECS}s\`" + echo +} >"${RESULTS_MD}" + +for size_label in "${SIZES[@]}"; do + bytes="$(size_to_bytes "${size_label}")" + sample_file="$(prepare_sample_file "${size_label}" "${bytes}")" + raw_json="${RUN_DIR}/${size_label}.jsonl" + prefix="$(echo "${size_label}" | tr '[:upper:]' '[:lower:]')" + worker_dir="${RUN_DIR}/${size_label}-workers" + mkdir -p "${worker_dir}" + + echo "Running mc PUT benchmark for size ${size_label}..." + batch_start="$(now_secs)" + if [[ "${OPS_PER_SIZE}" -gt 0 ]]; then + seq 1 "${OPS_PER_SIZE}" | xargs -n1 -P "${CONCURRENCY}" -I{} \ + env MC_CONFIG_DIR="${MC_CONFIG_DIR}" MC_BIN="${MC_BIN}" PUT_TIMEOUT_SECS="${PUT_TIMEOUT_SECS}" SAMPLE_FILE="${sample_file}" BUCKET="${BUCKET}" PREFIX="${prefix}" /bin/bash -lc ' + idx="$1" + start=$(perl -MTime::HiRes=time -e '"'"'printf "%.6f\n", time'"'"') + if run_mc_put_with_timeout "${SAMPLE_FILE}" "bench/${BUCKET}/${PREFIX}/obj-${idx}.bin"; then + finish=$(perl -MTime::HiRes=time -e '"'"'printf "%.6f\n", time'"'"') + perl -e '"'"' + use strict; + use warnings; + my ($idx, $start, $finish) = @ARGV; + my $duration_ms = sprintf("%.3f", ($finish - $start) * 1000); + print "{\"idx\":${idx},\"ok\":true,\"duration_ms\":${duration_ms}}\n"; + '"'"' -- "${idx}" "${start}" "${finish}" + else + finish=$(perl -MTime::HiRes=time -e '"'"'printf "%.6f\n", time'"'"') + perl -e '"'"' + use strict; + use warnings; + my ($idx, $start, $finish) = @ARGV; + my $duration_ms = sprintf("%.3f", ($finish - $start) * 1000); + print "{\"idx\":${idx},\"ok\":false,\"duration_ms\":${duration_ms}}\n"; + '"'"' -- "${idx}" "${start}" "${finish}" + fi + ' _ {} >"${raw_json}" + else + bench_deadline="$(perl -e 'use strict; use warnings; my ($start, $secs) = @ARGV; printf "%.6f", $start + $secs;' -- "${batch_start}" "${BENCH_SECONDS}")" + pids=() + for worker in $(seq 1 "${CONCURRENCY}"); do + ( + out_file="${worker_dir}/worker-${worker}.jsonl" + idx=0 + while true; do + now="$(now_secs)" + if perl -e 'use strict; use warnings; exit(($ARGV[0] >= $ARGV[1]) ? 0 : 1)' -- "${now}" "${bench_deadline}"; then + break + fi + start="${now}" + key="bench/${BUCKET}/${prefix}/worker-${worker}/obj-${idx}.bin" + if run_mc_put_with_timeout "${sample_file}" "${key}"; then + finish="$(now_secs)" + perl -e ' + use strict; + use warnings; + my ($idx, $start, $finish) = @ARGV; + my $duration_ms = sprintf("%.3f", ($finish - $start) * 1000); + print "{\"idx\":${idx},\"ok\":true,\"duration_ms\":${duration_ms}}\n"; + ' -- "${idx}" "${start}" "${finish}" >>"${out_file}" + else + finish="$(now_secs)" + perl -e ' + use strict; + use warnings; + my ($idx, $start, $finish) = @ARGV; + my $duration_ms = sprintf("%.3f", ($finish - $start) * 1000); + print "{\"idx\":${idx},\"ok\":false,\"duration_ms\":${duration_ms}}\n"; + ' -- "${idx}" "${start}" "${finish}" >>"${out_file}" + fi + idx=$((idx + 1)) + done + ) & + pids+=($!) + done + for pid in "${pids[@]}"; do + wait "${pid}" + done + cat "${worker_dir}"/worker-*.jsonl >"${raw_json}" + fi + batch_finish="$(now_secs)" + + jq_summary="$(jq -s ' + def pct(p): + if length == 0 then null + else (sort_by(.duration_ms) | .[((length - 1) * p | floor)].duration_ms) + end; + { + total: length, + succeeded: (map(select(.ok == true)) | length), + failed: (map(select(.ok == false)) | length), + avg_ms: (if length == 0 then null else (map(.duration_ms) | add / length) end), + p50_ms: pct(0.50), + p90_ms: pct(0.90), + p99_ms: pct(0.99) + } + ' "${raw_json}")" + + success_count="$(echo "${jq_summary}" | jq -r '.succeeded')" + avg_ms="$(echo "${jq_summary}" | jq -r '.avg_ms')" + p50_ms="$(echo "${jq_summary}" | jq -r '.p50_ms')" + p90_ms="$(echo "${jq_summary}" | jq -r '.p90_ms')" + p99_ms="$(echo "${jq_summary}" | jq -r '.p99_ms')" + failed_count="$(echo "${jq_summary}" | jq -r '.failed')" + + wall_secs="$(perl -e 'use strict; use warnings; my ($start, $finish) = @ARGV; printf "%.6f", ($finish - $start);' -- "${batch_start}" "${batch_finish}")" + mib_per_sec="$(perl -e 'use strict; use warnings; my ($bytes, $count, $secs) = @ARGV; printf "%.3f", (($bytes * $count) / (1024 * 1024)) / $secs;' -- "${bytes}" "${success_count}" "${wall_secs}")" + obj_per_sec="$(perl -e 'use strict; use warnings; my ($count, $secs) = @ARGV; printf "%.3f", $count / $secs;' -- "${success_count}" "${wall_secs}")" + + { + echo "## ${size_label}" + echo + echo "- Successful PUTs: \`${success_count}\`" + echo "- Failed PUTs: \`${failed_count}\`" + echo "- Wall time: \`${wall_secs}s\`" + echo "- Throughput: \`${mib_per_sec} MiB/s\`" + echo "- Object rate: \`${obj_per_sec} obj/s\`" + echo "- Avg latency: \`${avg_ms} ms\`" + echo "- p50 latency: \`${p50_ms} ms\`" + echo "- p90 latency: \`${p90_ms} ms\`" + echo "- p99 latency: \`${p99_ms} ms\`" + echo + } >>"${RESULTS_MD}" +done + +echo +echo "Benchmark completed." +echo "Results: ${RESULTS_MD}" +echo "Server log: ${SERVER_LOG}" diff --git a/scripts/run.sh b/scripts/run.sh index 36c4f788c..4d3e94dce 100755 --- a/scripts/run.sh +++ b/scripts/run.sh @@ -38,7 +38,8 @@ mkdir -p ./target/volume/test{1..4} if [ -z "$RUST_LOG" ]; then export RUST_BACKTRACE=1 - export RUST_LOG="info,rustfs=debug,rustfs_ecstore=info,s3s=debug,rustfs_iam=info,rustfs_notify=info" +# export RUST_LOG="info,rustfs=debug,rustfs_ecstore=info,s3s=debug,rustfs_iam=info,rustfs_notify=info" + export RUST_LOG="error" fi # export RUSTFS_ERASURE_SET_DRIVE_COUNT=5 @@ -71,7 +72,7 @@ export RUSTFS_OBS_SERVICE_NAME=rustfs # Service name export RUSTFS_OBS_SERVICE_VERSION=0.1.0 # Service version export RUSTFS_OBS_ENVIRONMENT=production # Environment name development, staging, production export RUSTFS_OBS_LOGGER_LEVEL=info # Log level, supports trace, debug, info, warn, error -export RUSTFS_OBS_LOG_STDOUT_ENABLED=true # Whether to enable local stdout logging +export RUSTFS_OBS_LOG_STDOUT_ENABLED=false # Whether to enable local stdout logging export RUSTFS_OBS_LOG_DIRECTORY="$current_dir/deploy/logs" # Log directory export RUSTFS_OBS_LOG_ROTATION_TIME="minutely" # Log rotation time unit, can be "minutely", "hourly", "daily" export RUSTFS_OBS_LOG_KEEP_FILES=10 # Number of log files to keep