diff --git a/crates/ecstore/src/rebalance.rs b/crates/ecstore/src/rebalance.rs index aba7aaa27..6cd01b6e2 100644 --- a/crates/ecstore/src/rebalance.rs +++ b/crates/ecstore/src/rebalance.rs @@ -1169,6 +1169,37 @@ fn resolve_rebalance_worker_result( } } +type RebalanceEntryTask = tokio::task::JoinHandle>; + +async fn wait_rebalance_entry_tasks(set_idx: usize, tasks: Arc>>) -> Result<()> { + let tasks = { + let mut tasks = tasks.lock().await; + std::mem::take(&mut *tasks) + }; + + let mut first_error = None; + for task in tasks { + match task.await { + Ok(Ok(())) => {} + Ok(Err(err)) => { + error!("rebalance entry task failed for set {}: {}", set_idx, err); + if first_error.is_none() { + first_error = Some(err); + } + } + Err(err) => { + let err = Error::other(format!("rebalance entry task join error for set {set_idx}: {err}")); + error!("{}", err); + if first_error.is_none() { + first_error = Some(err); + } + } + } + } + + if let Some(err) = first_error { Err(err) } else { Ok(()) } +} + fn resolve_rebalance_save_task_result( pool_idx: usize, save_task_result: std::result::Result, tokio::task::JoinError>, @@ -1691,26 +1722,28 @@ impl ECStore { let mut jobs = Vec::new(); let entry_error = Arc::new(tokio::sync::Mutex::new(None::)); + let entry_workers = Arc::new(tokio::sync::Semaphore::new(pool.disk_set.len().max(1))); - // let wk = Workers::new(pool.disk_set.len() * 2).map_err(Error::other)?; - // wk.clone().take().await; for (set_idx, set) in pool.disk_set.iter().enumerate() { + let entry_tasks = Arc::new(tokio::sync::Mutex::new(Vec::::new())); let rebalance_entry: ListCallback = Arc::new({ let this = Arc::clone(self); let bucket = bucket.clone(); let entry_error = entry_error.clone(); let callback_rx = rx.clone(); - // let wk = wk.clone(); let set = set.clone(); let bucket_configs = bucket_configs.clone(); + let entry_tasks = entry_tasks.clone(); + let entry_workers = entry_workers.clone(); move |entry: MetaCacheEntry| { let this = this.clone(); let bucket = bucket.clone(); let entry_error = entry_error.clone(); let callback_rx = callback_rx.clone(); - // let wk = wk.clone(); let set = set.clone(); let bucket_configs = bucket_configs.clone(); + let entry_tasks = entry_tasks.clone(); + let entry_workers = entry_workers.clone(); Box::pin(async move { if callback_rx.is_cancelled() { return; @@ -1719,20 +1752,38 @@ impl ECStore { return; } - info!("rebalance_entry: rebalance_entry spawn start"); - // wk.take().await; - // tokio::spawn(async move { - info!("rebalance_entry: rebalance_entry spawn start2"); - if let Err(err) = this.rebalance_entry(bucket, pool_index, entry, set, bucket_configs).await { - error!("rebalance_entry: rebalance entry failed: {err}"); - let mut first_err = entry_error.lock().await; - if first_err.is_none() { - *first_err = Some(err); - callback_rx.cancel(); - } + let permit = tokio::select! { + _ = callback_rx.cancelled() => return, + permit = entry_workers.clone().acquire_owned() => match permit { + Ok(permit) => permit, + Err(err) => { + error!("rebalance_entry: worker semaphore closed: {err}"); + return; + } + }, + }; + + if entry_error.lock().await.is_some() { + return; } - info!("rebalance_entry: rebalance_entry spawn done"); - // }); + + let task = tokio::spawn(async move { + let _permit = permit; + info!("rebalance_entry: rebalance entry task start"); + let result = this.rebalance_entry(bucket, pool_index, entry, set, bucket_configs).await; + if let Err(err) = &result { + error!("rebalance_entry: rebalance entry failed: {err}"); + let mut first_err = entry_error.lock().await; + if first_err.is_none() { + *first_err = Some(err.clone()); + callback_rx.cancel(); + } + } + info!("rebalance_entry: rebalance entry task done"); + result + }); + + entry_tasks.lock().await.push(task); }) } }); @@ -1740,23 +1791,23 @@ impl ECStore { let set = set.clone(); let rx = rx.clone(); let bucket = bucket.clone(); - // let wk = wk.clone(); + let entry_tasks = entry_tasks.clone(); let job = tokio::spawn(async move { - let result = set.list_objects_to_rebalance(rx, bucket, rebalance_entry).await; + let list_result = set.list_objects_to_rebalance(rx, bucket, rebalance_entry).await; + let entry_result = wait_rebalance_entry_tasks(set_idx, entry_tasks).await; + let result = list_result.and(entry_result); if let Err(err) = &result { error!("Rebalance worker {} error: {}", set_idx, err); } else { info!("Rebalance worker {} done", set_idx); } - // wk.clone().give().await; result }); jobs.push((set_idx, job)); } - // wk.wait().await; let mut worker_error: Option = None; for (set_idx, job) in jobs { if let Err(err) = resolve_rebalance_worker_result(set_idx, job.await) @@ -2680,6 +2731,29 @@ mod rebalance_unit_tests { assert!(err.to_string().contains("rebalance worker 7 task join error")); } + #[tokio::test] + async fn test_wait_rebalance_entry_tasks_returns_ok_for_successful_tasks() { + let tasks = Arc::new(tokio::sync::Mutex::new(vec![tokio::spawn(async { Ok(()) })])); + + super::wait_rebalance_entry_tasks(1, tasks) + .await + .expect("successful entry tasks should pass"); + } + + #[tokio::test] + async fn test_wait_rebalance_entry_tasks_returns_first_task_error() { + let tasks = Arc::new(tokio::sync::Mutex::new(vec![ + tokio::spawn(async { Ok(()) }), + tokio::spawn(async { Err(Error::other("entry failed")) }), + ])); + + let err = super::wait_rebalance_entry_tasks(1, tasks) + .await + .expect_err("entry task failure should be returned"); + + assert!(err.to_string().contains("entry failed")); + } + #[test] fn test_resolve_rebalance_save_task_result_passthrough() { assert!(resolve_rebalance_save_task_result(0, Ok(Ok(()))).is_ok()); diff --git a/rustfs/src/server/layer.rs b/rustfs/src/server/layer.rs index 3c02d0dfc..ebac6d33a 100644 --- a/rustfs/src/server/layer.rs +++ b/rustfs/src/server/layer.rs @@ -262,6 +262,7 @@ where if should_force_zero_content_length_for_empty_body_route(&req) { req.headers_mut() .insert(http::header::CONTENT_LENGTH, HeaderValue::from_static("0")); + req.headers_mut().remove(http::header::TRANSFER_ENCODING); } let mut inner = self.inner.clone(); @@ -274,14 +275,14 @@ fn should_force_zero_content_length_for_empty_body_route(req: &HttpRequest return false; } - if req.headers().contains_key(http::header::TRANSFER_ENCODING) { - return false; - } - if is_empty_body_admin_path(req.method(), req.uri().path()) { return true; } + if req.headers().contains_key(http::header::TRANSFER_ENCODING) { + return false; + } + is_empty_body_s3_path(req.method(), req.uri()) } @@ -1301,13 +1302,13 @@ mod tests { } #[tokio::test] - async fn empty_body_layer_preserves_admin_transfer_encoding_without_content_length() { + async fn empty_body_layer_normalizes_admin_chunked_request_without_content_length() { let capture = HeaderCaptureService::default(); let headers = capture.headers(); let mut service = EmptyBodyContentLengthCompatLayer.layer(capture); let request = Request::builder() - .method(Method::PUT) - .uri("/rustfs/admin/v3/set-group-status?group=test&status=enabled") + .method(Method::POST) + .uri(format!("{MINIO_ADMIN_V3_PREFIX}/rebalance/start")) .header(http::header::TRANSFER_ENCODING, "chunked") .body(()) .expect("request"); @@ -1315,8 +1316,8 @@ mod tests { let _ = service.call(request).await.expect("service call"); let headers = headers.lock().expect("captured headers").take().expect("captured headers"); - assert!(headers.get(http::header::CONTENT_LENGTH).is_none()); - assert_eq!(headers.get(http::header::TRANSFER_ENCODING).unwrap(), "chunked"); + assert_eq!(headers.get(http::header::CONTENT_LENGTH).unwrap(), "0"); + assert!(headers.get(http::header::TRANSFER_ENCODING).is_none()); } #[test]