fix(replication): preserve compressed multipart payload integrity (#7440)

Read compressed multipart replicas from one decoded stream and verify complete coverage before publication.
This commit is contained in:
Zhengchao An
2026-09-08 07:47:04 +08:00
committed by GitHub
parent 66b2a0f907
commit 3adfe65193
@@ -94,12 +94,14 @@ use rustfs_utils::{DEFAULT_SIP_HASH_KEY, get_env_usize, sip_hash};
use s3s::dto::ReplicationConfiguration;
use std::collections::{HashMap, HashSet};
use std::fmt::Display;
use std::pin::Pin;
use std::sync::atomic::{AtomicBool, Ordering};
use std::sync::{Arc, LazyLock, Mutex as StdMutex};
use std::task::{Context, Poll};
use std::time::Instant;
use time::OffsetDateTime;
use time::format_description::well_known::Rfc3339;
use tokio::io::AsyncRead;
use tokio::io::{AsyncRead, AsyncReadExt, ReadBuf};
use tokio::sync::{OwnedSemaphorePermit, RwLock, Semaphore};
use tokio::task::{JoinHandle, JoinSet};
use tokio::time::Duration as TokioDuration;
@@ -4536,7 +4538,6 @@ impl ReplicateObjectInfoExt for ReplicateObjectInfo {
let source_version_id = self.version_id;
let assigned_version_id = if is_multipart {
drop(gr);
replicate_object_with_multipart(MultipartReplicationContext {
storage: storage.clone(),
cli: tgt_client.clone(),
@@ -4547,6 +4548,7 @@ impl ReplicateObjectInfoExt for ReplicateObjectInfo {
obj_opts: &obj_opts,
arn: &rinfo.arn,
put_opts,
full_stream: Some(gr.stream),
})
.await
} else {
@@ -5283,7 +5285,6 @@ async fn replicate_all_payload_to_target<S: ReplicationObjectIO>(
}
if ctx.is_multipart {
drop(gr);
replicate_object_with_multipart(MultipartReplicationContext {
storage: ctx.storage.clone(),
cli: ctx.tgt_client.clone(),
@@ -5294,6 +5295,7 @@ async fn replicate_all_payload_to_target<S: ReplicationObjectIO>(
obj_opts: ctx.obj_opts,
arn: ctx.arn,
put_opts: ctx.put_opts,
full_stream: Some(gr.stream),
})
.await
} else {
@@ -5365,11 +5367,15 @@ struct MultipartReplicationContext<'a, S: ReplicationObjectIO> {
obj_opts: &'a ObjectOptions,
arn: &'a str,
put_opts: PutObjectOptions,
full_stream: Option<Box<dyn AsyncRead + Unpin + Send + Sync>>,
}
async fn replicate_object_with_multipart<S: ReplicationObjectIO>(
ctx: MultipartReplicationContext<'_, S>,
mut ctx: MultipartReplicationContext<'_, S>,
) -> std::io::Result<Option<String>> {
if ctx.obj_opts.raw_data_movement_read || !ctx.object_info.is_compressed() {
ctx.full_stream = None;
}
let mut attempts = 1;
let upload_id = loop {
match ctx
@@ -5557,6 +5563,21 @@ struct MultipartReplicationReadPlan {
next_offset: i64,
}
#[derive(Clone)]
struct SharedReplicationReader(Arc<StdMutex<Box<dyn AsyncRead + Unpin + Send + Sync>>>);
impl AsyncRead for SharedReplicationReader {
fn poll_read(self: Pin<&mut Self>, cx: &mut Context<'_>, buf: &mut ReadBuf<'_>) -> Poll<std::io::Result<()>> {
// UploadPart owns its body, but successive parts consume one source.
// The mutex only covers polling; no guard survives an await.
let mut reader = self
.0
.lock()
.map_err(|_| std::io::Error::other("replication source stream lock poisoned"))?;
Pin::new(&mut **reader).poll_read(cx, buf)
}
}
fn multipart_replication_read_plan(
object_info: &ObjectInfo,
obj_opts: &ObjectOptions,
@@ -5615,8 +5636,10 @@ async fn replicate_multipart_parts_and_complete<S: ReplicationObjectIO>(
obj_opts,
arn,
put_opts,
full_stream,
} = ctx;
let full_stream = full_stream.map(|reader| SharedReplicationReader(Arc::new(StdMutex::new(reader))));
let mut uploaded_parts: Vec<CompletedPart> = Vec::new();
let mut header_size = replication_put_object_header_size(&put_opts);
@@ -5635,7 +5658,13 @@ async fn replicate_multipart_parts_and_complete<S: ReplicationObjectIO>(
)?;
offset = part_plan.next_offset;
let byte_stream = if let Some(range_spec) = part_plan.range {
let byte_stream = if let Some(reader) = &full_stream {
let limit = u64::try_from(part_plan.part_size)
.map_err(|_| std::io::Error::new(std::io::ErrorKind::InvalidData, "negative multipart replication part size"))?;
let part_stream =
wrap_with_bandwidth_monitor_with_header(Box::new(reader.clone().take(limit)), src_bucket, arn, header_size);
async_read_to_bytestream(part_stream)
} else if let Some(range_spec) = part_plan.range {
let part_reader = storage
.get_object_reader(src_bucket, object, Some(range_spec), HeaderMap::new(), obj_opts)
.await
@@ -5681,6 +5710,25 @@ async fn replicate_multipart_parts_and_complete<S: ReplicationObjectIO>(
);
}
if let Some(mut reader) = full_stream {
// Range planning also trusts the recorded logical part boundaries.
// A full decoded stream avoids repeating or skipping bytes when those
// boundaries are stale, and its EOF must precede publication.
let expected_size = object_info
.get_actual_size()
.map_err(|err| std::io::Error::other(err.to_string()))?;
let mut remaining = [0_u8; 1];
if offset != expected_size || reader.read(&mut remaining).await? != 0 {
return Err(std::io::Error::new(
std::io::ErrorKind::InvalidData,
"compressed replication part lengths do not cover the complete source",
));
}
// Releasing the reader also releases the source read lock, before
// CompleteMultipartUpload or replication status writes can run.
drop(reader);
}
let actual_size = replication_multipart_complete_actual_size(&object_info.user_defined);
let completed = cli
@@ -8181,6 +8229,7 @@ mod tests {
#[derive(Debug)]
struct Source {
body: Bytes,
stored: Option<Bytes>,
info: ObjectInfo,
ranges: StdMutex<Vec<(i64, i64)>>,
full_reads: std::sync::atomic::AtomicUsize,
@@ -8209,6 +8258,18 @@ mod tests {
self.info.version_id.map(|id| id.to_string()),
"every read retains the selected source version"
);
if let Some(stored) = &self.stored {
if let Some(range) = &range {
self.ranges.lock().expect("range journal lock").push((range.start, range.end));
} else {
self.full_reads.fetch_add(1, Ordering::Relaxed);
}
let plan =
crate::object_api::ReadPlan::build_for_request(range, &self.info, opts, &HeaderMap::new(), None).await?;
let start = plan.storage_offset();
let end = start + usize::try_from(plan.storage_length()).expect("nonnegative storage length");
return plan.into_object_reader(Box::new(std::io::Cursor::new(stored.slice(start..end))), &self.info);
}
if range.is_none() {
self.full_reads.fetch_add(1, Ordering::Relaxed);
return Ok(GetObjectReader {
@@ -8282,6 +8343,265 @@ mod tests {
run_transport(4096, None, true).await;
}
#[tokio::test]
async fn multipart_transport_checks_positive_compressed_part_totals() {
const MIB: i64 = 1024 * 1024;
for etag_present in [true, false] {
run_compressed_transport([5 * MIB, MIB], etag_present, true, None).await;
run_compressed_transport([5 * MIB, MIB / 2], etag_present, false, None).await;
}
}
#[tokio::test]
async fn multipart_transport_checks_positive_compressed_part_boundaries() {
const MIB: i64 = 1024 * 1024;
for etag_present in [true, false] {
run_compressed_transport([5 * MIB + MIB / 2, MIB / 2], etag_present, false, None).await;
}
}
#[derive(Clone, Copy)]
enum SourceReadFault {
Short,
Extra,
}
#[tokio::test]
async fn multipart_transport_requires_complete_compressed_stream() {
const MIB: i64 = 1024 * 1024;
for fault in [SourceReadFault::Short, SourceReadFault::Extra] {
run_compressed_transport([5 * MIB, MIB], true, false, Some(fault)).await;
}
}
async fn run_compressed_transport(
actual_sizes: [i64; 2],
etag_present: bool,
valid_layout: bool,
fault: Option<SourceReadFault>,
) {
use crate::io_support::rio::TryGetIndex as _;
use rustfs_utils::CompressionAlgorithm;
use tokio::io::AsyncReadExt as _;
const MIB: usize = 1024 * 1024;
let body = Bytes::from([vec![0x11; 5 * MIB], vec![0x22; MIB / 2], vec![0x33; MIB / 2]].concat());
let mut stored = Vec::new();
let mut parts = Vec::new();
for (index, plaintext) in [body.slice(..5 * MIB), body.slice(5 * MIB..)].into_iter().enumerate() {
let mut compressor = crate::io_support::rio::compression_reader(
std::io::Cursor::new(plaintext),
CompressionAlgorithm::default(),
false,
);
let mut compressed = Vec::new();
compressor.read_to_end(&mut compressed).await.expect("compress source part");
parts.push(ObjectPartInfo {
number: index + 1,
size: compressed.len(),
actual_size: actual_sizes[index],
index: compressor
.try_get_index()
.map(crate::io_support::rio::compression_index_storage_bytes),
..Default::default()
});
stored.extend_from_slice(&compressed);
}
let mut metadata = HashMap::new();
rustfs_utils::http::insert_str(
&mut metadata,
rustfs_utils::http::SUFFIX_COMPRESSION,
crate::io_support::rio::compression_metadata_value(CompressionAlgorithm::default()),
);
rustfs_utils::http::insert_str(&mut metadata, rustfs_utils::http::SUFFIX_ACTUAL_SIZE, body.len().to_string());
let source = Arc::new(Source {
info: ObjectInfo {
bucket: "source".to_string(),
name: "object".to_string(),
size: i64::try_from(stored.len()).expect("stored size"),
actual_size: i64::try_from(body.len()).expect("plaintext size"),
etag: etag_present
.then(|| faster_hex::hex_string(rustfs_utils::hash::HashAlgorithm::Md5.hash_encode(&body).as_ref())),
version_id: Some(Uuid::new_v4()),
user_defined: Arc::new(metadata),
parts: Arc::new(parts),
..Default::default()
},
body: body.clone(),
stored: Some(Bytes::from(stored)),
ranges: StdMutex::new(Vec::new()),
full_reads: std::sync::atomic::AtomicUsize::new(0),
});
let opts = ObjectOptions {
version_id: source.info.version_id.map(|id| id.to_string()),
..Default::default()
};
let mut full = source
.get_object_reader("source", "object", None, HeaderMap::new(), &opts)
.await
.expect("open complete compressed source");
let decoded = full.read_all().await.expect("decode complete source");
assert!(
decoded.as_slice() == body.as_ref(),
"the independent full GET must match every source byte"
);
let (target, journal, server) = start_transport_target().await;
let (put_opts, is_multipart) = replication_put_object_options("STANDARD", &source.info).expect("replication options");
let mut reader = source
.get_object_reader("source", "object", None, HeaderMap::new(), &opts)
.await
.expect("open replication source");
match fault {
Some(SourceReadFault::Short) => {
reader.stream = Box::new(
reader
.stream
.take(u64::try_from(body.len() - 1).expect("short stream length")),
);
}
Some(SourceReadFault::Extra) => {
reader.stream = Box::new(reader.stream.chain(std::io::Cursor::new(vec![0x44])));
}
None => {}
}
let result = tokio::time::timeout(
std::time::Duration::from_secs(30),
replicate_all_payload_to_target(
ReplicateAllPayloadContext {
storage: &source,
tgt_client: &target,
bucket: "source",
object: "object",
object_info: &source.info,
obj_opts: &opts,
arn: &target.arn,
transfer_size: i64::try_from(body.len()).expect("plaintext size"),
is_multipart,
put_opts,
},
reader,
),
)
.await;
server.abort();
assert!(server.await.expect_err("fixture server is stopped").is_cancelled());
let result = result.expect("replication must finish");
let requests = journal.lock().expect("request journal lock");
let completes: Vec<_> = requests
.iter()
.filter(|request| request.method == http::Method::POST && request.query.contains_key("uploadId"))
.collect();
if result.is_err() {
assert!(!valid_layout, "valid compressed layout must replicate: {result:?}");
assert!(completes.is_empty(), "failed replication must not publish a partial object");
assert!(
requests
.iter()
.any(|request| request.method == http::Method::DELETE && request.query.contains_key("uploadId")),
"failed multipart upload must be aborted"
);
return;
}
assert!(fault.is_none(), "a short or oversized source must fail before Complete");
assert_eq!(requests.len(), 4, "initiate, two parts and complete without retries");
assert_eq!(completes.len(), 1, "publish the replica once");
let mut replica = Vec::new();
for number in [1, 2] {
let part = requests
.iter()
.find(|request| request.query.get("partNumber") == Some(&number.to_string()))
.expect("completed source part");
assert_eq!(part.method, http::Method::PUT);
replica.extend_from_slice(&part.body);
let complete_xml = std::str::from_utf8(&completes[0].body).expect("complete XML");
assert!(complete_xml.contains(&format!("<PartNumber>{number}</PartNumber>")));
}
assert!(
replica.as_slice() == body.as_ref(),
"replica differs from complete GET: sizes={actual_sizes:?}, source={}, replica={}",
body.len(),
replica.len()
);
if let Some(etag) = &source.info.etag {
assert_eq!(
rustfs_utils::http::get_header(&completes[0].headers, rustfs_utils::http::SUFFIX_SOURCE_ETAG).as_deref(),
Some(etag.as_str())
);
}
}
async fn start_transport_target() -> (Arc<TargetClient>, Arc<StdMutex<Vec<RequestRecord>>>, tokio::task::JoinHandle<()>) {
let journal = Arc::new(StdMutex::new(Vec::<RequestRecord>::new()));
let listener = tokio::net::TcpListener::bind(("127.0.0.1", 0))
.await
.expect("bind multipart target");
let endpoint = format!("http://{}", listener.local_addr().expect("multipart target address"));
let server_journal = journal.clone();
let server = tokio::spawn(async move {
let mut connections = JoinSet::new();
loop {
let (stream, _) = listener.accept().await.expect("accept multipart request");
let journal = server_journal.clone();
connections.spawn(async move {
let service = hyper::service::service_fn(move |request: hyper::Request<hyper::body::Incoming>| {
let journal = journal.clone();
async move {
let (request, body) = request.into_parts();
let query: HashMap<String, String> = url::form_urlencoded::parse(
request.uri.query().unwrap_or_default().as_bytes(),
).into_owned().collect();
let body = match body.collect().await {
Ok(body) => body.to_bytes(),
Err(_) => return Ok::<_, Infallible>(hyper::Response::builder()
.status(http::StatusCode::BAD_REQUEST)
.body(Full::new(Bytes::new()))
.expect("incomplete upload response")),
};
let response = if request.method == http::Method::POST && query.contains_key("uploads") {
"<InitiateMultipartUploadResult><Bucket>target-bucket</Bucket><Key>object</Key><UploadId>upload-1</UploadId></InitiateMultipartUploadResult>"
} else if request.method == http::Method::PUT {
""
} else if request.method == http::Method::POST && query.contains_key("uploadId") {
"<CompleteMultipartUploadResult><Location>http://localhost/object</Location><Bucket>target-bucket</Bucket><Key>object</Key><ETag>&quot;target-2&quot;</ETag></CompleteMultipartUploadResult>"
} else if request.method == http::Method::DELETE && query.contains_key("uploadId") {
""
} else {
panic!("unexpected multipart request: {} {}", request.method, request.uri)
};
let response_etag = if request.method == http::Method::PUT && !query.contains_key("partNumber") {
format!("\"{}\"", faster_hex::hex_string(rustfs_utils::hash::HashAlgorithm::Md5.hash_encode(&body).as_ref()))
} else {
"\"uploaded-part\"".to_string()
};
journal.lock().expect("request journal lock").push(RequestRecord {
method: request.method, query, headers: request.headers, body,
});
Ok::<_, Infallible>(hyper::Response::builder()
.header("content-type", "application/xml")
.header("etag", response_etag)
.body(Full::new(Bytes::from_static(response.as_bytes())))
.expect("multipart response"))
}
});
hyper::server::conn::http1::Builder::new()
.serve_connection(hyper_util::rt::TokioIo::new(stream), service)
.await.expect("serve multipart connection");
});
}
});
let mut target = test_target_client(endpoint);
let config = target
.client
.config()
.to_builder()
.request_checksum_calculation(aws_sdk_s3::config::RequestChecksumCalculation::WhenRequired)
.force_path_style(true)
.build();
Arc::get_mut(&mut target).expect("unshared test target").client = Arc::new(aws_sdk_s3::Client::from_conf(config));
(target, journal, server)
}
async fn run_transport(tail_size: usize, unknown_part: Option<(usize, i64)>, passthrough: bool) {
const FIRST_SIZE: usize = 5 * 1024 * 1024;
// The stored (compressed ciphertext) bytes of a passthrough part
@@ -8350,70 +8670,11 @@ mod tests {
..Default::default()
},
body: body.clone(),
stored: None,
ranges: StdMutex::new(Vec::new()),
full_reads: std::sync::atomic::AtomicUsize::new(0),
});
let journal = Arc::new(StdMutex::new(Vec::<RequestRecord>::new()));
let listener = tokio::net::TcpListener::bind(("127.0.0.1", 0))
.await
.expect("bind multipart target");
let endpoint = format!("http://{}", listener.local_addr().expect("multipart target address"));
let server_journal = journal.clone();
let server = tokio::spawn(async move {
let mut connections = JoinSet::new();
loop {
let (stream, _) = listener.accept().await.expect("accept multipart request");
let journal = server_journal.clone();
connections.spawn(async move {
let service = hyper::service::service_fn(move |request: hyper::Request<hyper::body::Incoming>| {
let journal = journal.clone();
async move {
let (request, body) = request.into_parts();
let query: HashMap<String, String> = url::form_urlencoded::parse(
request.uri.query().unwrap_or_default().as_bytes(),
).into_owned().collect();
let body = body.collect().await.expect("read complete multipart request body").to_bytes();
let response = if request.method == http::Method::POST && query.contains_key("uploads") {
"<InitiateMultipartUploadResult><Bucket>target-bucket</Bucket><Key>object</Key><UploadId>upload-1</UploadId></InitiateMultipartUploadResult>"
} else if request.method == http::Method::PUT {
""
} else if request.method == http::Method::POST && query.contains_key("uploadId") {
"<CompleteMultipartUploadResult><Location>http://localhost/object</Location><Bucket>target-bucket</Bucket><Key>object</Key><ETag>&quot;target-2&quot;</ETag></CompleteMultipartUploadResult>"
} else if request.method == http::Method::DELETE && query.contains_key("uploadId") {
""
} else {
panic!("unexpected multipart request: {} {}", request.method, request.uri)
};
let response_etag = if request.method == http::Method::PUT && !query.contains_key("partNumber") {
format!("\"{}\"", faster_hex::hex_string(rustfs_utils::hash::HashAlgorithm::Md5.hash_encode(&body).as_ref()))
} else {
"\"uploaded-part\"".to_string()
};
journal.lock().expect("request journal lock").push(RequestRecord {
method: request.method, query, headers: request.headers, body,
});
Ok::<_, Infallible>(hyper::Response::builder()
.header("content-type", "application/xml")
.header("etag", response_etag)
.body(Full::new(Bytes::from_static(response.as_bytes())))
.expect("multipart response"))
}
});
hyper::server::conn::http1::Builder::new()
.serve_connection(hyper_util::rt::TokioIo::new(stream), service)
.await.expect("serve multipart connection");
});
}
});
let mut target = test_target_client(endpoint);
let config = target
.client
.config()
.to_builder()
.request_checksum_calculation(aws_sdk_s3::config::RequestChecksumCalculation::WhenRequired)
.force_path_style(true)
.build();
Arc::get_mut(&mut target).expect("unshared test target").client = Arc::new(aws_sdk_s3::Client::from_conf(config));
let (target, journal, server) = start_transport_target().await;
let (put_opts, is_multipart) = replication_put_object_options("STANDARD", &source.info).expect("replication options");
let opts = ObjectOptions {
version_id: source.info.version_id.map(|id| id.to_string()),