From b46a8164f83153304eb2e163f632aaeaa1b93147 Mon Sep 17 00:00:00 2001 From: cxymds Date: Wed, 2 Sep 2026 09:09:55 +0800 Subject: [PATCH] fix(http): drain request bodies after early responses (#7019) * fix(http): drain request bodies after early responses * fix(http): log early response body drain failures --------- Co-authored-by: houseme --- rustfs/src/server/http.rs | 699 +++++++++++++++++++++++++++++++++++++- 1 file changed, 692 insertions(+), 7 deletions(-) diff --git a/rustfs/src/server/http.rs b/rustfs/src/server/http.rs index a9786bd3f..a3ed6462d 100644 --- a/rustfs/src/server/http.rs +++ b/rustfs/src/server/http.rs @@ -41,7 +41,7 @@ use crate::storage_api::server::http::{ tonic_boot_epoch_response_headers, verify_tonic_rpc_signature_with_bootstrap, }; use bytes::Bytes; -use http::{HeaderMap, Method, Request as HttpRequest, Response, Uri}; +use http::{HeaderMap, HeaderValue, Method, Request as HttpRequest, Response, Uri, Version, header::CONNECTION}; use hyper::body::Incoming; use hyper_util::{ rt::{TokioExecutor, TokioIo, TokioTimer}, @@ -74,7 +74,7 @@ use s3s::{ }; use socket2::{SockRef, TcpKeepalive}; use std::io::{Error, Result}; -use std::net::SocketAddr; +use std::net::{IpAddr, SocketAddr}; use std::pin::Pin; use std::sync::Arc; use std::sync::atomic::{AtomicU64, Ordering}; @@ -109,6 +109,8 @@ const METRIC_HTTP_SERVER_RESPONSE_BODY_CHUNK_SIZE_BYTES: &str = "rustfs_http_ser const METRIC_HTTP_SERVER_RESPONSE_BODY_CHUNK_LATENCY_SECONDS: &str = "rustfs_http_server_response_body_chunk_latency_seconds"; const METRIC_HTTP_SERVER_RESPONSE_BODY_STREAM_DURATION_SECONDS: &str = "rustfs_http_server_response_body_stream_duration_seconds"; const METRIC_HTTP_SERVER_CONNECTION_CAP_SATURATED_TOTAL: &str = "rustfs_http_server_connection_cap_saturated_total"; +const METRIC_HTTP_SERVER_EARLY_RESPONSE_BODY_DRAINS_TOTAL: &str = "rustfs_http_server_early_response_body_drains_total"; +const LABEL_HTTP_BODY_DRAIN_RESULT: &str = "result"; const HTTP_STREAMING_BODY_FAILURE_STAGE_TRANSPORT: &str = "http_transport"; const HTTP_STREAMING_BODY_FAILURE_REASON_TRANSPORT: &str = "transport_failure"; const HTTP_STREAMING_BODY_FAILURE_CLASS_TRANSPORT: &str = "transport"; @@ -178,6 +180,7 @@ const EVENT_CONNECTION_CAP_STATE: &str = "connection_cap_state"; const EVENT_HTTP_TRANSPORT_PARAMETERS: &str = "http_transport_parameters"; const EVENT_HTTP_ACCEPT_LOOP_STATE: &str = "http_accept_loop_state"; const EVENT_HTTP_CONNECTION_DRAIN: &str = "http_connection_drain"; +const EVENT_HTTP_EARLY_RESPONSE_BODY_DRAIN: &str = "http_early_response_body_drain"; const EVENT_PEER_ADDR_UNAVAILABLE: &str = "peer_addr_unavailable"; const EVENT_RPC_SIGNATURE_VERIFICATION_FAILED: &str = "rpc_signature_verification_failed"; const EVENT_GRPC_TRACE_CONTEXT_PROPAGATION_FAILED: &str = "grpc_trace_context_propagation_failed"; @@ -418,6 +421,20 @@ fn log_transport_failed(peer_addr: &str, error_kind: &str, error_message: &str) ); } +fn format_peer_addr(real_ip: Option, remote_addr: Option) -> String { + real_ip + .map(|addr| addr.to_string()) + .or_else(|| remote_addr.map(|addr| addr.to_string())) + .unwrap_or_else(|| "unknown".to_string()) +} + +fn request_peer_addr(request: &HttpRequest) -> String { + format_peer_addr( + request.extensions().get::().map(|info| info.real_ip), + request.extensions().get::().map(|addr| addr.0), + ) +} + #[inline] fn record_active_http_requests(delta: i64) { let next = if delta >= 0 { @@ -531,6 +548,315 @@ where } } +#[derive(Clone)] +struct EarlyResponseBodyService { + inner: S, + idle_timeout: Duration, +} + +impl EarlyResponseBodyService { + fn new(inner: S, idle_timeout: Duration) -> Self { + Self { inner, idle_timeout } + } +} + +struct EarlyResponseBodyState { + abandoned: parking_lot::Mutex>, +} + +impl Default for EarlyResponseBodyState { + fn default() -> Self { + Self { + abandoned: parking_lot::Mutex::new(None), + } + } +} + +impl EarlyResponseBodyState { + fn abandon(&self, body: B) { + let previous = self.abandoned.lock().replace(body); + debug_assert!(previous.is_none(), "a request can abandon its body only once"); + } + + fn take_abandoned(&self) -> Option { + self.abandoned.lock().take() + } +} + +/// Request body passed into s3s while retaining ownership of the raw transport +/// body if an operation returns before consuming it. +/// +/// Hyper closes the HTTP/1 read side as soon as an [`Incoming`] receiver is +/// dropped before EOF. A streaming reverse proxy can still be writing at that +/// point and turns the resulting `EPIPE` into a 502, hiding RustFS's actual S3 +/// error response. Moving the raw body back to the outer service keeps Hyper's +/// receiver alive without running s3s payload hashing or signature transforms. +struct EarlyResponseBody +where + B: http_body::Body + Unpin, +{ + inner: Option, + state: Arc>, + complete: bool, +} + +impl EarlyResponseBody +where + B: http_body::Body + Unpin, +{ + fn new(body: B, state: Arc>) -> Self { + Self { + inner: Some(body), + state, + complete: false, + } + } +} + +impl http_body::Body for EarlyResponseBody +where + B: http_body::Body + Unpin, +{ + type Data = Bytes; + type Error = B::Error; + + fn is_end_stream(&self) -> bool { + self.complete || self.inner.as_ref().is_none_or(|body| body.is_end_stream()) + } + + fn poll_frame( + mut self: Pin<&mut Self>, + cx: &mut Context<'_>, + ) -> Poll, Self::Error>>> { + let Some(body) = self.inner.as_mut() else { + self.complete = true; + return Poll::Ready(None); + }; + + match Pin::new(body).poll_frame(cx) { + Poll::Ready(None) => { + self.complete = true; + Poll::Ready(None) + } + Poll::Ready(Some(Err(err))) => { + self.complete = true; + Poll::Ready(Some(Err(err))) + } + other => other, + } + } + + fn size_hint(&self) -> http_body::SizeHint { + self.inner + .as_ref() + .map_or_else(http_body::SizeHint::default, |body| body.size_hint()) + } +} + +impl Drop for EarlyResponseBody +where + B: http_body::Body + Unpin, +{ + fn drop(&mut self) { + if self.complete { + return; + } + + let Some(body) = self.inner.take() else { + return; + }; + if body.is_end_stream() { + return; + } + + self.state.abandon(body); + } +} + +#[derive(Clone, Copy)] +enum EarlyResponseBodyDrainResult { + Completed, + BodyError, + IdleTimeout, +} + +impl EarlyResponseBodyDrainResult { + fn as_str(self) -> &'static str { + match self { + Self::Completed => "completed", + Self::BodyError => "body_error", + Self::IdleTimeout => "idle_timeout", + } + } +} + +struct EarlyResponseBodyDrainContext { + real_ip: Option, + remote_addr: Option, + method: Method, + uri: Uri, + status_code: u16, + idle_timeout: Duration, +} + +impl EarlyResponseBodyDrainContext { + fn from_request(request: &HttpRequest, idle_timeout: Duration) -> Self { + Self { + real_ip: request.extensions().get::().map(|info| info.real_ip), + remote_addr: request.extensions().get::().map(|addr| addr.0), + method: request.method().clone(), + uri: request.uri().clone(), + status_code: 0, + idle_timeout, + } + } + + fn peer_addr(&self) -> String { + format_peer_addr(self.real_ip, self.remote_addr) + } +} + +fn log_early_response_body_drain(result: EarlyResponseBodyDrainResult, context: &EarlyResponseBodyDrainContext) { + match result { + EarlyResponseBodyDrainResult::Completed => {} + EarlyResponseBodyDrainResult::BodyError => { + let peer_addr = context.peer_addr(); + let uri = redact_sensitive_uri_query(&context.uri); + debug!( + event = EVENT_HTTP_EARLY_RESPONSE_BODY_DRAIN, + component = LOG_COMPONENT_SERVER, + subsystem = LOG_SUBSYSTEM_TRANSPORT, + result = result.as_str(), + peer_addr = %peer_addr, + method = %context.method, + uri = %uri, + status_code = context.status_code, + "HTTP early response body drain closed before EOF" + ); + } + EarlyResponseBodyDrainResult::IdleTimeout => { + let peer_addr = context.peer_addr(); + let uri = redact_sensitive_uri_query(&context.uri); + warn!( + event = EVENT_HTTP_EARLY_RESPONSE_BODY_DRAIN, + component = LOG_COMPONENT_SERVER, + subsystem = LOG_SUBSYSTEM_TRANSPORT, + result = result.as_str(), + peer_addr = %peer_addr, + method = %context.method, + uri = %uri, + status_code = context.status_code, + timeout_secs = context.idle_timeout.as_secs(), + "HTTP early response body drain timed out" + ); + } + } +} + +async fn drain_early_response_body(mut body: B, idle_timeout: Duration) -> EarlyResponseBodyDrainResult +where + B: http_body::Body + Send + Unpin + 'static, + B::Error: Send + 'static, +{ + loop { + let next_frame = std::future::poll_fn(|cx| Pin::new(&mut body).poll_frame(cx)); + let next_frame = if idle_timeout.is_zero() { + next_frame.await + } else { + match tokio::time::timeout(idle_timeout, next_frame).await { + Ok(frame) => frame, + Err(_) => return EarlyResponseBodyDrainResult::IdleTimeout, + } + }; + + match next_frame { + Some(Ok(_)) => {} + Some(Err(_)) => return EarlyResponseBodyDrainResult::BodyError, + None => return EarlyResponseBodyDrainResult::Completed, + } + } +} + +fn spawn_early_response_body_drain(body: B, context: EarlyResponseBodyDrainContext) +where + B: http_body::Body + Send + Unpin + 'static, + B::Error: Send + 'static, +{ + counter!( + METRIC_HTTP_SERVER_EARLY_RESPONSE_BODY_DRAINS_TOTAL, + LABEL_HTTP_BODY_DRAIN_RESULT => "started" + ) + .increment(1); + tokio::spawn(async move { + let result = drain_early_response_body(body, context.idle_timeout).await; + counter!( + METRIC_HTTP_SERVER_EARLY_RESPONSE_BODY_DRAINS_TOTAL, + LABEL_HTTP_BODY_DRAIN_RESULT => result.as_str() + ) + .increment(1); + log_early_response_body_drain(result, &context); + }); +} + +impl Service> for EarlyResponseBodyService +where + S: Service, Response = Response, Error = ServiceError> + + Service>, Response = Response, Error = ServiceError> + + Clone + + Send + + 'static, + >>::Future: Send + 'static, + >>>::Future: Send + 'static, + B: http_body::Body + Send + Unpin + 'static, + B::Error: std::error::Error + Send + Sync + 'static, + ResBody: Send + 'static, + ServiceError: Send + 'static, +{ + type Response = Response; + type Error = ServiceError; + type Future = Pin> + Send>>; + + fn poll_ready(&mut self, cx: &mut Context<'_>) -> Poll> { + match >>::poll_ready(&mut self.inner, cx)? { + Poll::Ready(()) => >>>::poll_ready(&mut self.inner, cx), + Poll::Pending => Poll::Pending, + } + } + + fn call(&mut self, req: HttpRequest) -> Self::Future { + let version = req.version(); + let preserve_on_drop = matches!(version, Version::HTTP_10 | Version::HTTP_11) && !req.body().is_end_stream(); + let mut inner = self.inner.clone(); + if !preserve_on_drop { + return Box::pin(async move { >>::call(&mut inner, req).await }); + } + + let mut drain_context = EarlyResponseBodyDrainContext::from_request(&req, self.idle_timeout); + let state = Arc::new(EarlyResponseBodyState::default()); + let guarded_req = req.map({ + let state = Arc::clone(&state); + move |body| EarlyResponseBody::new(body, state) + }); + + Box::pin(async move { + let result = >>>::call(&mut inner, guarded_req).await; + let Some(abandoned) = state.take_abandoned() else { + return result; + }; + + match result { + Ok(mut response) => { + response.headers_mut().insert(CONNECTION, HeaderValue::from_static("close")); + drain_context.status_code = response.status().as_u16(); + spawn_early_response_body_drain(abandoned, drain_context); + Ok(response) + } + Err(err) => Err(err), + } + }) + } +} + fn trace_on_response(response: &Response, latency: Duration, span: &Span) { span.record("status_code", tracing::field::display(response.status())); let _enter = span.enter(); @@ -573,11 +899,7 @@ fn make_http_trace_span(request: &HttpRequest) -> Span { trace!("No trace context found in request headers, will create root span"); } - let client_info = request.extensions().get::(); - let peer_addr = client_info - .map(|info| info.real_ip.to_string()) - .or_else(|| request.extensions().get::().map(|addr| addr.0.to_string())) - .unwrap_or_else(|| "unknown".to_string()); + let peer_addr = request_peer_addr(request); let span = tracing::info_span!("http-request", request_id = %request_id, @@ -1072,6 +1394,10 @@ pub async fn start_http_server( rustfs_config::ENV_HTTP1_HEADER_READ_TIMEOUT, rustfs_config::DEFAULT_HTTP1_HEADER_READ_TIMEOUT, ); + let http_request_body_read_timeout = Duration::from_secs(rustfs_utils::get_env_u64( + rustfs_config::ENV_HTTP_REQUEST_BODY_READ_TIMEOUT, + rustfs_config::DEFAULT_HTTP_REQUEST_BODY_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); @@ -1268,6 +1594,7 @@ pub async fn start_http_server( rate_limit_layer: api_rate_limit_layer.clone(), server_ctx: Arc::clone(&server_ctx), tls_handshake_timeout: Duration::from_secs(http1_header_read_timeout), + request_body_idle_timeout: http_request_body_read_timeout, }; process_connection(socket, tls_acceptor.clone(), connection_ctx, graceful.watcher(), connection_permit); @@ -1332,6 +1659,9 @@ struct ConnectionContext { /// the existing slow-client bound for the pre-request phase, and is pre-computed with the /// other transport parameters to avoid a per-connection env read. tls_handshake_timeout: Duration, + /// Inter-chunk timeout used while discarding a raw HTTP/1 request body + /// after an S3 operation has already produced its response. + request_body_idle_timeout: Duration, } #[derive(Clone)] @@ -1482,6 +1812,7 @@ fn process_connection( rate_limit_layer, server_ctx, tls_handshake_timeout, + request_body_idle_timeout, } = context; // Build the hybrid service per-connection. @@ -1527,6 +1858,7 @@ fn process_connection( let http_service = SwiftService::new(true, None, s3_service); #[cfg(not(feature = "swift"))] let http_service = s3_service; + let http_service = EarlyResponseBodyService::new(http_service, request_body_idle_timeout); let http_service = InternodeRpcService::new(http_service); let external_service = hybrid(http_service.clone(), rpc_service.clone()); @@ -2079,6 +2411,7 @@ mod tests { use bytes::Bytes; use http::Request as HttpRequest; use http::{HeaderMap, StatusCode}; + use http_body::Frame; use http_body_util::{Empty, Full}; use metrics::with_local_recorder; use metrics_util::debugging::{DebugValue, DebuggingRecorder}; @@ -2086,10 +2419,13 @@ mod tests { use std::collections::HashMap; use std::convert::Infallible; use std::future::Ready; + use std::sync::atomic::{AtomicBool, AtomicUsize}; use std::sync::{Arc, Mutex}; use std::task::{Context, Poll}; use storage::tonic_service::{heal_topology_fingerprint, make_heal_control_server_for_source}; use storage::{Endpoint, EndpointServerPools, Endpoints, PoolEndpoints}; + use tokio::io::{AsyncReadExt, AsyncWriteExt}; + use tokio::sync::{Notify, mpsc}; use tower::{Layer, Service, ServiceBuilder}; type MetricRow = ( @@ -2222,6 +2558,7 @@ mod tests { METRIC_HTTP_SERVER_REQUEST_BODY_SIZE_BYTES, METRIC_HTTP_SERVER_RESPONSE_BODY_BYTES_TOTAL, METRIC_HTTP_SERVER_RESPONSE_BODY_SIZE_BYTES, + METRIC_HTTP_SERVER_EARLY_RESPONSE_BODY_DRAINS_TOTAL, ]; for metric_name in metric_names { @@ -2231,6 +2568,7 @@ mod tests { assert_eq!(LABEL_HTTP_METHOD, "method"); assert_eq!(LABEL_HTTP_STATUS_CLASS, "status_class"); + assert_eq!(LABEL_HTTP_BODY_DRAIN_RESULT, "result"); } #[test] @@ -2268,6 +2606,353 @@ mod tests { assert_eq!(HTTP_STATUS_CLASS_LABELS[HTTP_STATUS_UNKNOWN_INDEX], "unknown"); } + struct TrackedRequestBody { + receiver: mpsc::UnboundedReceiver, + bytes_polled: Arc, + dropped: Arc, + } + + impl http_body::Body for TrackedRequestBody { + type Data = Bytes; + type Error = Infallible; + + fn is_end_stream(&self) -> bool { + self.receiver.is_closed() && self.receiver.is_empty() + } + + fn poll_frame( + mut self: Pin<&mut Self>, + cx: &mut Context<'_>, + ) -> Poll, Self::Error>>> { + match self.receiver.poll_recv(cx) { + Poll::Ready(Some(bytes)) => { + self.bytes_polled.fetch_add(bytes.len(), Ordering::Relaxed); + Poll::Ready(Some(Ok(Frame::data(bytes)))) + } + Poll::Ready(None) => Poll::Ready(None), + Poll::Pending => Poll::Pending, + } + } + } + + impl Drop for TrackedRequestBody { + fn drop(&mut self) { + self.dropped.store(true, Ordering::Release); + } + } + + fn tracked_request_body() -> (mpsc::UnboundedSender, TrackedRequestBody, Arc, Arc) { + let (sender, receiver) = mpsc::unbounded_channel(); + let bytes_polled = Arc::new(AtomicUsize::new(0)); + let dropped = Arc::new(AtomicBool::new(false)); + ( + sender, + TrackedRequestBody { + receiver, + bytes_polled: Arc::clone(&bytes_polled), + dropped: Arc::clone(&dropped), + }, + bytes_polled, + dropped, + ) + } + + #[derive(Clone, Copy)] + struct ConsumeRequestBodyService; + + impl Service> for ConsumeRequestBodyService + where + B: http_body::Body + Send + 'static, + B::Error: Send + 'static, + { + type Response = Response>; + type Error = Infallible; + type Future = Pin> + Send>>; + + fn poll_ready(&mut self, _cx: &mut Context<'_>) -> Poll> { + Poll::Ready(Ok(())) + } + + fn call(&mut self, req: HttpRequest) -> Self::Future { + Box::pin(async move { + let mut body = Box::pin(req.into_body()); + while let Some(frame) = std::future::poll_fn(|cx| body.as_mut().poll_frame(cx)).await { + let _ = frame; + } + Ok(Response::new(Empty::new())) + }) + } + } + + #[derive(Clone, Copy)] + struct RejectWithoutReadingBodyService; + + impl Service> for RejectWithoutReadingBodyService { + type Response = Response>; + type Error = Infallible; + type Future = Ready>; + + fn poll_ready(&mut self, _cx: &mut Context<'_>) -> Poll> { + Poll::Ready(Ok(())) + } + + fn call(&mut self, _req: HttpRequest) -> Self::Future { + std::future::ready(Ok(Response::builder() + .status(StatusCode::SERVICE_UNAVAILABLE) + .body(Empty::new()) + .expect("response"))) + } + } + + struct TransformCountingBody { + inner: B, + polls: Arc, + } + + impl http_body::Body for TransformCountingBody + where + B: http_body::Body + Unpin, + { + type Data = Bytes; + type Error = B::Error; + + fn is_end_stream(&self) -> bool { + self.inner.is_end_stream() + } + + fn poll_frame( + mut self: Pin<&mut Self>, + cx: &mut Context<'_>, + ) -> Poll, Self::Error>>> { + self.polls.fetch_add(1, Ordering::Relaxed); + Pin::new(&mut self.inner).poll_frame(cx) + } + + fn size_hint(&self) -> http_body::SizeHint { + self.inner.size_hint() + } + } + + #[derive(Clone)] + struct TransformThenRejectService { + transformed_polls: Arc, + } + + impl Service> for TransformThenRejectService + where + B: http_body::Body + Unpin, + { + type Response = Response>; + type Error = Infallible; + type Future = Ready>; + + fn poll_ready(&mut self, _cx: &mut Context<'_>) -> Poll> { + Poll::Ready(Ok(())) + } + + fn call(&mut self, req: HttpRequest) -> Self::Future { + let transformed = TransformCountingBody { + inner: req.into_body(), + polls: Arc::clone(&self.transformed_polls), + }; + drop(transformed); + std::future::ready(Ok(Response::builder() + .status(StatusCode::SERVICE_UNAVAILABLE) + .body(Empty::new()) + .expect("response"))) + } + } + + #[tokio::test] + async fn early_response_body_service_leaves_fully_consumed_http1_body_reusable() { + let (sender, body, bytes_polled, dropped) = tracked_request_body(); + sender.send(Bytes::from_static(b"payload")).expect("body receiver"); + drop(sender); + let request = HttpRequest::builder().version(Version::HTTP_11).body(body).expect("request"); + let mut service = EarlyResponseBodyService::new(ConsumeRequestBodyService, Duration::from_secs(1)); + + let response = service.call(request).await.expect("response"); + + assert_eq!(bytes_polled.load(Ordering::Relaxed), 7); + assert!(dropped.load(Ordering::Acquire)); + assert!(response.headers().get(CONNECTION).is_none()); + } + + #[tokio::test] + async fn early_response_body_service_drains_abandoned_http1_body_and_closes_connection() { + let (sender, body, bytes_polled, dropped) = tracked_request_body(); + let request = HttpRequest::builder().version(Version::HTTP_11).body(body).expect("request"); + let mut service = EarlyResponseBodyService::new(RejectWithoutReadingBodyService, Duration::from_secs(1)); + + let response = service.call(request).await.expect("response"); + assert_eq!(response.headers().get(CONNECTION), Some(&HeaderValue::from_static("close"))); + + sender.send(Bytes::from_static(b"payload")).expect("drain receiver"); + drop(sender); + tokio::time::timeout(Duration::from_secs(1), async { + while !dropped.load(Ordering::Acquire) { + tokio::task::yield_now().await; + } + }) + .await + .expect("abandoned body drain should finish"); + assert_eq!(bytes_polled.load(Ordering::Relaxed), 7); + } + + #[tokio::test] + async fn early_response_body_drain_bypasses_downstream_transforms() { + let (sender, body, raw_bytes_polled, dropped) = tracked_request_body(); + let transformed_polls = Arc::new(AtomicUsize::new(0)); + let request = HttpRequest::builder().version(Version::HTTP_11).body(body).expect("request"); + let inner = TransformThenRejectService { + transformed_polls: Arc::clone(&transformed_polls), + }; + let mut service = EarlyResponseBodyService::new(inner, Duration::from_secs(1)); + + let response = service.call(request).await.expect("response"); + assert_eq!(response.headers().get(CONNECTION), Some(&HeaderValue::from_static("close"))); + sender.send(Bytes::from_static(b"payload")).expect("drain receiver"); + drop(sender); + tokio::time::timeout(Duration::from_secs(1), async { + while !dropped.load(Ordering::Acquire) { + tokio::task::yield_now().await; + } + }) + .await + .expect("raw body drain should finish"); + + assert_eq!(raw_bytes_polled.load(Ordering::Relaxed), 7); + assert_eq!(transformed_polls.load(Ordering::Relaxed), 0); + } + + #[tokio::test] + async fn early_response_body_service_keeps_http2_stream_cancellation_semantics() { + let (sender, body, bytes_polled, dropped) = tracked_request_body(); + let request = HttpRequest::builder().version(Version::HTTP_2).body(body).expect("request"); + let mut service = EarlyResponseBodyService::new(RejectWithoutReadingBodyService, Duration::from_secs(1)); + + let response = service.call(request).await.expect("response"); + + assert!(response.headers().get(CONNECTION).is_none()); + assert!(dropped.load(Ordering::Acquire)); + assert!(sender.send(Bytes::from_static(b"payload")).is_err()); + assert_eq!(bytes_polled.load(Ordering::Relaxed), 0); + } + + #[tokio::test(start_paused = true)] + async fn early_response_body_drain_releases_stalled_body_after_idle_timeout() { + let (_sender, body, _bytes_polled, dropped) = tracked_request_body(); + let request = HttpRequest::builder().version(Version::HTTP_11).body(body).expect("request"); + let mut service = EarlyResponseBodyService::new(RejectWithoutReadingBodyService, Duration::from_secs(10)); + + let response = service.call(request).await.expect("response"); + assert_eq!(response.headers().get(CONNECTION), Some(&HeaderValue::from_static("close"))); + tokio::task::yield_now().await; + tokio::time::advance(Duration::from_secs(11)).await; + tokio::task::yield_now().await; + + assert!(dropped.load(Ordering::Acquire)); + } + + async fn assert_http1_early_response_accepts_streaming_body(expect_continue: bool) { + let listener = match TcpListener::bind("127.0.0.1:0").await { + Ok(listener) => listener, + Err(err) if err.kind() == std::io::ErrorKind::PermissionDenied => return, + Err(err) => panic!("test listener should bind: {err}"), + }; + let addr = listener.local_addr().expect("listener address"); + let server = tokio::spawn(async move { + let (socket, _) = listener.accept().await.expect("server accept"); + let service = EarlyResponseBodyService::new(RejectWithoutReadingBodyService, Duration::from_secs(5)); + hyper::server::conn::http1::Builder::new() + .serve_connection(TokioIo::new(socket), TowerToHyperService::new(service)) + .await + }); + + let stream = TcpStream::connect(addr).await.expect("client connect"); + SockRef::from(&stream).set_send_buffer_size(4096).expect("client send buffer"); + let (mut reader, mut writer) = stream.into_split(); + let total_body_len = 2 * 1024 * 1024; + let expect_header = if expect_continue { "Expect: 100-continue\r\n" } else { "" }; + let headers = format!( + "PUT /bucket/object?partNumber=1&uploadId=test HTTP/1.1\r\nHost: {addr}\r\nContent-Length: {total_body_len}\r\n{expect_header}\r\n" + ); + writer.write_all(headers.as_bytes()).await.expect("request headers"); + if expect_continue { + tokio::time::sleep(Duration::from_millis(20)).await; + } + + const FIRST_CHUNK_LEN: usize = 4096; + let first_chunk = [b'a'; FIRST_CHUNK_LEN]; + writer.write_all(&first_chunk).await.expect("first body chunk"); + let continue_upload = Arc::new(Notify::new()); + let writer_gate = Arc::clone(&continue_upload); + let writer_task = tokio::spawn(async move { + writer_gate.notified().await; + let remaining_body = vec![b'b'; total_body_len - FIRST_CHUNK_LEN]; + writer.write_all(&remaining_body).await?; + writer.shutdown().await + }); + + let mut response_bytes = Vec::new(); + let final_response = tokio::time::timeout(Duration::from_secs(5), async { + let mut chunk = [0_u8; 1024]; + loop { + let read = reader.read(&mut chunk).await?; + if read == 0 { + return Err(std::io::Error::new( + std::io::ErrorKind::UnexpectedEof, + "connection closed before the final response", + )); + } + response_bytes.extend_from_slice(&chunk[..read]); + let Some(status_offset) = response_bytes + .windows(b"HTTP/1.1 503".len()) + .position(|window| window == b"HTTP/1.1 503") + else { + continue; + }; + if response_bytes[status_offset..].windows(4).any(|window| window == b"\r\n\r\n") { + return Ok::<(), std::io::Error>(()); + } + } + }) + .await; + continue_upload.notify_one(); + final_response + .expect("early response should arrive before the remaining upload") + .expect("final response should be readable"); + + let response_text = String::from_utf8_lossy(&response_bytes).to_ascii_lowercase(); + assert!(response_text.contains("http/1.1 503")); + assert!(response_text.contains("connection: close")); + if expect_continue { + assert!( + !response_text.contains("100 continue"), + "early rejection must not invite the client to keep uploading" + ); + } + writer_task + .await + .expect("writer task") + .expect("proxy-side upload must not see a broken pipe"); + tokio::time::timeout(Duration::from_secs(5), server) + .await + .expect("server connection should close after the body drain") + .expect("server task") + .expect("HTTP/1 connection should complete cleanly"); + } + + #[tokio::test] + async fn early_response_body_http1_transport_does_not_break_streaming_upload() { + assert_http1_early_response_accepts_streaming_body(false).await; + } + + #[tokio::test] + async fn early_response_body_http1_transport_handles_expect_continue_upload() { + assert_http1_early_response_accepts_streaming_body(true).await; + } + #[test] fn rustfs_s3_config_preserves_compatibility_over_s3s_defaults() { let s3_config = rustfs_s3_config();