diff --git a/crates/ecstore/src/bucket/replication/replication_resyncer.rs b/crates/ecstore/src/bucket/replication/replication_resyncer.rs index 39d20d8df..8136108c1 100644 --- a/crates/ecstore/src/bucket/replication/replication_resyncer.rs +++ b/crates/ecstore/src/bucket/replication/replication_resyncer.rs @@ -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( } 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( 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>, } async fn replicate_object_with_multipart( - ctx: MultipartReplicationContext<'_, S>, + mut ctx: MultipartReplicationContext<'_, S>, ) -> std::io::Result> { + 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>>); + +impl AsyncRead for SharedReplicationReader { + fn poll_read(self: Pin<&mut Self>, cx: &mut Context<'_>, buf: &mut ReadBuf<'_>) -> Poll> { + // 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( 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 = Vec::new(); let mut header_size = replication_put_object_header_size(&put_opts); @@ -5635,7 +5658,13 @@ async fn replicate_multipart_parts_and_complete( )?; 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( ); } + 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, info: ObjectInfo, ranges: StdMutex>, 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, + ) { + 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!("{number}"))); + } + 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, Arc>>, tokio::task::JoinHandle<()>) { + let journal = Arc::new(StdMutex::new(Vec::::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| { + let journal = journal.clone(); + async move { + let (request, body) = request.into_parts(); + let query: HashMap = 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") { + "target-bucketobjectupload-1" + } else if request.method == http::Method::PUT { + "" + } else if request.method == http::Method::POST && query.contains_key("uploadId") { + "http://localhost/objecttarget-bucketobject"target-2"" + } 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::::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| { - let journal = journal.clone(); - async move { - let (request, body) = request.into_parts(); - let query: HashMap = 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") { - "target-bucketobjectupload-1" - } else if request.method == http::Method::PUT { - "" - } else if request.method == http::Method::POST && query.contains_key("uploadId") { - "http://localhost/objecttarget-bucketobject"target-2"" - } 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()),