diff --git a/src/brotli/lib.rs b/src/brotli/lib.rs index 86e2db54f2b2..e6d668ff39d1 100644 --- a/src/brotli/lib.rs +++ b/src/brotli/lib.rs @@ -101,17 +101,23 @@ impl StreamingDecoder { unsafe { self.brotli.as_mut() } } - /// Consume all of `input`, appending decompressed bytes to `out` - /// (growing in 4096-byte steps). Returns `ShortRead` when more input is - /// required and `is_done` is false. + #[inline] + pub fn is_inflating(&self) -> bool { + matches!(self.state, ReaderState::Inflating) + } + + /// Append decompressed bytes to `out` (growing in 4096-byte steps) until `input` is + /// consumed or `out.len()` reaches `max_output`. Returns the input bytes consumed. + /// Returns `ShortRead` when more input is required and `is_done` is false. pub fn decompress( &mut self, input: &[u8], out: &mut Vec, + max_output: usize, is_done: bool, - ) -> crate::Result<()> { + ) -> crate::Result { if matches!(self.state, ReaderState::End | ReaderState::Error) { - return Ok(()); + return Ok(input.len()); } debug_assert!(out.as_ptr() != input.as_ptr()); @@ -120,12 +126,16 @@ impl StreamingDecoder { self.state, ReaderState::Uninitialized | ReaderState::Inflating ) { + if out.len() >= max_output { + return Ok(total_in); + } if out.try_reserve(4096).is_err() { self.state = ReaderState::Error; return Err(crate::Error::OutOfMemory); } + let budget = max_output - out.len(); let spare = out.spare_capacity_mut(); - let out_len = spare.len(); + let out_len = spare.len().min(budget); let mut next_out: *mut u8 = spare.as_mut_ptr().cast::(); let next_in = &input[total_in..]; @@ -159,7 +169,7 @@ impl StreamingDecoder { match result { c::BrotliDecoderResult::success => { self.state = ReaderState::End; - return Ok(()); + return Ok(input.len()); } c::BrotliDecoderResult::err => { self.state = ReaderState::Error; @@ -192,7 +202,7 @@ impl StreamingDecoder { } } } - Ok(()) + Ok(total_in) } } diff --git a/src/http/Decompressor.rs b/src/http/Decompressor.rs index aca7133da552..fac59577b7d5 100644 --- a/src/http/Decompressor.rs +++ b/src/http/Decompressor.rs @@ -19,6 +19,16 @@ impl Decompressor { // explicit `Drop` is unnecessary. Callers that want a mid-lifecycle reset // assign `*self = Decompressor::None`. + /// Inside a stream: another `decompress_chunk` may produce output with no new input. + pub(crate) fn is_mid_stream(&self) -> bool { + match self { + Decompressor::Zlib(r) => r.is_inflating(), + Decompressor::Brotli(r) => r.is_inflating(), + Decompressor::Zstd(r) => r.is_inflating(), + Decompressor::None => false, + } + } + fn init(&mut self, encoding: Encoding, first_chunk: &[u8]) -> crate::Result<()> { match encoding { Encoding::Gzip | Encoding::Deflate => { @@ -49,27 +59,30 @@ impl Decompressor { } /// Feed one body chunk `buffer` through the decoder, appending the - /// decompressed output to `body_out_str`. Creates the decoder on first - /// call. Returns `ShortRead` when more input is needed and the stream is - /// not yet done. + /// decompressed output to `body_out_str` until it holds `max_output` bytes. Creates the + /// decoder on first call. Returns the input bytes consumed. Returns `ShortRead` when more + /// input is needed and the stream is not yet done. pub(crate) fn decompress_chunk( &mut self, encoding: Encoding, buffer: &[u8], body_out_str: &mut MutableString, + max_output: usize, is_done: bool, - ) -> crate::Result<()> { + ) -> crate::Result { if !encoding.is_compressed() { - return Ok(()); + return Ok(buffer.len()); } if matches!(self, Decompressor::None) { self.init(encoding, buffer)?; } let out = &mut body_out_str.list; match self { - Decompressor::Zlib(reader) => Ok(reader.decompress(buffer, out, is_done)?), - Decompressor::Brotli(reader) => Ok(reader.decompress(buffer, out, is_done)?), - Decompressor::Zstd(reader) => Ok(reader.decompress(buffer, out, is_done)?), + Decompressor::Zlib(reader) => Ok(reader.decompress(buffer, out, max_output, is_done)?), + Decompressor::Brotli(reader) => { + Ok(reader.decompress(buffer, out, max_output, is_done)?) + } + Decompressor::Zstd(reader) => Ok(reader.decompress(buffer, out, max_output, is_done)?), Decompressor::None => { unreachable!("Invalid encoding. This code should not be reachable") } diff --git a/src/http/InternalState.rs b/src/http/InternalState.rs index 96c01181667b..1a083f5e38d2 100644 --- a/src/http/InternalState.rs +++ b/src/http/InternalState.rs @@ -29,6 +29,8 @@ pub struct InternalState<'a> { /// (cap-bounded) after the callback returns. pub(crate) decoded_body: MutableString, pub(crate) compressed_body: MutableString, + /// Prefix of `compressed_body` the decoder has already taken. + compressed_body_consumed: usize, pub(crate) content_length: Option, pub(crate) total_body_received: usize, // Self-borrow into `original_request_body.bytes`; `RawSlice` carries the @@ -80,6 +82,8 @@ pub struct InternalStateFlags { /// `reset()`/`init()` so each redirect/retry hop re-compresses from the /// original uncompressed `original_request_body`. pub(crate) body_compressed: bool, + /// Held input or buffered decoder output remains for `HTTPClient::drain_response_body`. + pub(crate) decompress_output_pending: bool, } impl InternalStateFlags { @@ -95,6 +99,7 @@ impl InternalStateFlags { is_waiting_for_cert_check: false, receive_paused: false, body_compressed: false, + decompress_output_pending: false, } } } @@ -113,6 +118,7 @@ impl Default for InternalState<'_> { stage: Stage::Pending, decoded_body: MutableString::init_empty(), compressed_body: MutableString::init_empty(), + compressed_body_consumed: 0, content_length: None, total_body_received: 0, request_body: bun_ptr::RawSlice::EMPTY, @@ -208,10 +214,19 @@ impl<'a> InternalState<'a> { self.flags.received_last_chunk } + #[inline] + pub(crate) fn has_pending_compressed(&self) -> bool { + self.flags.decompress_output_pending + } + /// True when a socket close during `in_progress` completes the body rather /// than failing it: chunked decoder already in the trailers state, or a /// close-delimited response (no Content-Length, no Transfer-Encoding). pub(crate) fn is_body_complete_on_close(&self) -> bool { + // Every byte arrived; only the decode is outstanding. + if self.flags.decompress_output_pending && self.is_done() { + return true; + } if self.is_chunked_encoding() { return bun_picohttp::phr_decode_chunked_is_in_trailers(&self.chunked_decoder) != 0; } @@ -227,27 +242,30 @@ impl<'a> InternalState<'a> { pub(crate) fn finalize_body_on_eof(&mut self) -> Result<(), Error> { self.flags.received_last_chunk = true; let buffer_snap = core::mem::take(&mut self.get_body_buffer().list); - self.process_body_buffer(buffer_snap, true).map(drop) + self.process_body_buffer(buffer_snap, true, usize::MAX) + .map(drop) } pub(crate) fn decompress_bytes( &mut self, buffer: &[u8], is_final_chunk: bool, - ) -> Result<(), Error> { + max_output: usize, + ) -> Result { // A response that declared a Content-Encoding but sent zero body bytes // (e.g. an empty chunked gzip response) has nothing to decompress. // Running the decompressor anyway makes it report a truncated stream // (ZlibError); Node treats this as an empty body. if buffer.is_empty() && self.total_body_received == 0 { self.compressed_body.reset(); - return Ok(()); + return Ok(0); } // `self.compressed_body.reset()` must run on every exit. scopeguard would // hold &mut self.compressed_body across the body and conflict with &mut self.decompressor, // so each early-return below calls it explicitly. let mut still_needs_to_decompress = true; + let mut consumed = buffer.len(); if bun_core::feature_flags::is_libdeflate_enabled() { // Fast-path: use libdeflate @@ -256,6 +274,7 @@ impl<'a> InternalState<'a> { use bun_libdeflate_sys::libdeflate as bun_libdeflate; if !(is_final_chunk && !self.flags.is_libdeflate_fast_path_disabled + && matches!(self.decompressor, Decompressor::None) && self.encoding.can_use_lib_deflate() && self.is_done()) { @@ -280,6 +299,12 @@ impl<'a> InternalState<'a> { .try_into() .expect("infallible: size matches"), ); + // Under an output budget only `shared_buffer`'s worth may come out in one shot. + if (estimated_size as usize) > deflater.shared_buffer.len() + && max_output != usize::MAX + { + break 'libdeflate; + } // Since this is arbtirary input from the internet, let's set an upper bound of 32 MB for the allocation size. if (estimated_size as usize) > deflater.shared_buffer.len() && estimated_size < 32 * 1024 * 1024 @@ -356,33 +381,40 @@ impl<'a> InternalState<'a> { let min = ((buffer.len() as f64) * 1.5) .ceil() .min(1024.0 * 1024.0 * 2.0); - if let Err(err) = self.decoded_body.grow_by((min as usize).max(32)) { + if let Err(err) = self + .decoded_body + .grow_by((min as usize).max(32).min(max_output)) + { self.compressed_body.reset(); return Err(err.into()); } } let is_done = self.is_done(); - if let Err(err) = self.decompressor.decompress_chunk( + match self.decompressor.decompress_chunk( self.encoding, buffer, &mut self.decoded_body, + max_output, is_done, ) { - if is_done || err != crate::Error::ShortRead { - bun_core::pretty_errorln!( - "Decompression error: {}", - bstr::BStr::new(err.name()), - ); - Output::flush(); - self.compressed_body.reset(); - return Err(err); + Ok(n) => consumed = n, + Err(err) => { + if is_done || err != crate::Error::ShortRead { + bun_core::pretty_errorln!( + "Decompression error: {}", + bstr::BStr::new(err.name()), + ); + Output::flush(); + self.compressed_body.reset(); + return Err(err); + } } } } self.compressed_body.reset(); - Ok(()) + Ok(consumed) } // `buffer` is always the current body buffer's bytes. To avoid aliased &mut/& under @@ -393,6 +425,7 @@ impl<'a> InternalState<'a> { &mut self, mut buffer: Vec, is_final_chunk: bool, + max_output: usize, ) -> Result { if self.flags.is_redirect_pending { // Caller moved the bytes out of the body buffer; put them back so the @@ -403,10 +436,21 @@ impl<'a> InternalState<'a> { match self.encoding { Encoding::Brotli | Encoding::Gzip | Encoding::Deflate | Encoding::Zstd => { - self.decompress_bytes(&buffer, is_final_chunk)?; - // Retain capacity by - // returning the (cleared) allocation to compressed_body instead of dropping it. - buffer.clear(); + let start = self.compressed_body_consumed; + let consumed = + start + self.decompress_bytes(&buffer[start..], is_final_chunk, max_output)?; + let held = buffer.len() - consumed; + // A decoder can hold output with no input left (brotli copy command, zstd flush). + self.flags.decompress_output_pending = self.decoded_body.list.len() >= max_output + && (held != 0 || self.decompressor.is_mid_stream()); + // Shifting only once the taken prefix is the larger part moves each byte once. + if consumed >= held { + buffer.drain(..consumed); + self.compressed_body_consumed = 0; + } else { + self.compressed_body_consumed = consumed; + } + // Retain capacity by returning the allocation to compressed_body. self.compressed_body.list = buffer; } _ => { diff --git a/src/http/Signals.rs b/src/http/Signals.rs index 1d343542a181..868a028bb1f0 100644 --- a/src/http/Signals.rs +++ b/src/http/Signals.rs @@ -84,6 +84,19 @@ impl Signals { .map(bun_ptr::BackRef::from) .is_some_and(|a| a.load(Ordering::Acquire) == BodyReceiveMode::Paused as u8) } + + /// `Flowing` or `Paused`: a consumer takes the body piece by piece. + #[inline] + pub(crate) fn is_demand_driven(self) -> bool { + self.body_receive_mode + .map(bun_ptr::BackRef::from) + .is_some_and(|a| { + matches!( + BodyReceiveMode::from_u8(a.load(Ordering::Acquire)), + BodyReceiveMode::Flowing | BodyReceiveMode::Paused + ) + }) + } } pub struct Store { diff --git a/src/http/lib.rs b/src/http/lib.rs index d95940839ce6..b436832e317a 100644 --- a/src/http/lib.rs +++ b/src/http/lib.rs @@ -2187,6 +2187,11 @@ impl<'a> HTTPClient<'a> { if self.flags.disable_timeout { return; } + // A fully received body that waits on its consumer expects nothing from the socket. + if self.state.has_pending_compressed() && self.state.is_done() { + socket.set_timeout(0); + return; + } bun_core::scoped_log!(fetch, "Timeout {}\n", BStr::new(self.url.href)); // Terminate (mark dead + close) BEFORE failing, matching // `close_and_fail`: `fail()` dispatches the final result, which frees @@ -3911,6 +3916,7 @@ impl<'a> HTTPClient<'a> { self.progress_update::(ctx, socket); return; } + self.maybe_pause_receive(socket); } ResponseStage::BodyChunk => { if !self.state.flags.receive_paused { @@ -3930,6 +3936,7 @@ impl<'a> HTTPClient<'a> { self.progress_update::(ctx, socket); return; } + self.maybe_pause_receive(socket); } ResponseStage::Fail => {} _ => { @@ -4080,6 +4087,33 @@ impl<'a> HTTPClient<'a> { socket.set_timeout(self.effective_idle_timeout_seconds()); } + /// Output budget of one decode pass. h1 only: h2/h3 detach before held input could drain. + #[inline] + fn decompress_output_cap(&self) -> usize { + if self.flags.protocol == Protocol::Http1_1 && self.signals.is_demand_driven() { + signals::BODY_HIGH_WATER_MARK + } else { + usize::MAX + } + } + + /// Decodes what has arrived under the consumer's budget. Returns whether to report bytes. + fn process_received_body(&mut self, is_final_chunk: bool) -> crate::Result { + let max_output = self.decompress_output_cap(); + // Nothing is decoded for a paused consumer (a tunnelled socket keeps reading anyway). + if max_output != usize::MAX + && self.state.encoding.is_compressed() + && self.signals.is_receive_paused() + { + self.state.flags.decompress_output_pending = true; + return Ok(false); + } + // `process_body_buffer` takes `&mut self.state`, so the bytes move out first. + let buffer = core::mem::take(&mut self.state.get_body_buffer().list); + self.state + .process_body_buffer(buffer, is_final_chunk, max_output) + } + fn maybe_pause_receive(&mut self, socket: HttpSocket) { if self.state.flags.receive_paused || self.proxy_tunnel.is_some() @@ -4114,31 +4148,45 @@ impl<'a> HTTPClient<'a> { } pub(crate) fn drain_response_body(&mut self, socket: HttpSocket) { + if self.pump_held_body::(socket) { + let ctx = self.get_ssl_ctx::(); + self.send_progress_update_without_stage_check::(ctx, socket); + } + } + + /// Decodes the next piece of a held body. Returns whether there is an update to send. + fn pump_held_body(&mut self, socket: HttpSocket) -> bool { // Find out if we should not send any update. match self.state.stage { - Stage::Done | Stage::Fail => return, + Stage::Done | Stage::Fail => return false, _ => {} } if self.state.fail.is_some() { // If there's any error at all, do not drain. - return; + return false; } // If there's a pending redirect, then don't bother to send a response body // as that wouldn't make sense and I want to defensively avoid edgecases // from that. if self.state.flags.is_redirect_pending { - return; + return false; } - if self.state.decoded_body.list.is_empty() { - // No update! Don't do anything. - return; + // A consumer that paused again gets another resume when it unpauses. + let pumped = self.state.has_pending_compressed() && !self.signals.is_receive_paused(); + if pumped { + let is_final = self.state.is_done(); + if let Err(err) = self.process_received_body(is_final) { + self.close_and_fail::(err, socket); + return false; + } } - let ctx = self.get_ssl_ctx::(); - self.send_progress_update_without_stage_check::(ctx, socket); + // A pump that ends the body has to say so even with no bytes (a stream trailer alone). + let ended = pumped && self.state.is_done() && !self.state.has_pending_compressed(); + !self.state.decoded_body.list.is_empty() || ended } fn send_progress_update_without_stage_check( @@ -4149,6 +4197,19 @@ impl<'a> HTTPClient<'a> { if self.flags.protocol != Protocol::Http1_1 { return self.send_progress_update_multiplexed(); } + // A loop, not a call back into `drain_response_body`: a consumer that never pauses + // (`BufferAll`, or an S3 error body that is collected whole) takes one pass per turn. + while self.send_one_progress_update::(ctx, socket) + && self.pump_held_body::(socket) + {} + } + + /// Returns whether a held body is left that its consumer will not ask for. + fn send_one_progress_update( + &mut self, + ctx: *mut GenHttpContext, + socket: HttpSocket, + ) -> bool { let callback = self.result_callback; let mut result = self.to_result(); @@ -4264,9 +4325,12 @@ impl<'a> HTTPClient<'a> { self.state.decoded_body = decoded_body; } self.maybe_pause_receive(socket); + // Only a paused consumer asks for the rest. + self.state.has_pending_compressed() && !self.signals.is_receive_paused() } else { result.body_owned = decoded_body.list; callback.run(parent, result); + false } } @@ -4485,7 +4549,8 @@ impl<'a> HTTPClient<'a> { dns_hostname: self.state.dns_hostname.take(), connect_errno: self.state.connect_errno, proxy_connect_response: None, - has_more: self.state.fail.is_none() && !self.state.is_done(), + has_more: self.state.fail.is_none() + && (!self.state.is_done() || self.state.has_pending_compressed()), body_size, certificate_info: None, can_stream: (self.state.request_stage == RequestStage::Body @@ -4507,7 +4572,8 @@ impl<'a> HTTPClient<'a> { proxy_connect_response, // check if we are reporting cert errors, do not have a fail state and we are not done has_more: certificate_info.is_some() - || (self.state.fail.is_none() && !self.state.is_done()), + || (self.state.fail.is_none() + && (!self.state.is_done() || self.state.has_pending_compressed())), body_size, certificate_info, // we can stream the request_body at this stage @@ -4534,6 +4600,8 @@ impl<'a> HTTPClient<'a> { if is_only_buffer && let Some(len) = content_length && incoming_data.len() >= len + // The single-packet path decodes the whole body with no output budget. + && !(self.state.encoding.is_compressed() && self.signals.is_demand_driven()) { self.handle_response_body_from_single_packet(&incoming_data[0..len])?; Ok(true) @@ -4557,7 +4625,8 @@ impl<'a> HTTPClient<'a> { // we can ignore the body data in redirects if !self.state.flags.is_redirect_pending { if self.state.encoding.is_compressed() { - self.state.decompress_bytes(incoming_data, true)?; + self.state + .decompress_bytes(incoming_data, true, usize::MAX)?; } else { self.state .get_body_buffer() @@ -4613,13 +4682,7 @@ impl<'a> HTTPClient<'a> { || self.signals.body_receive_mode.is_some(); if is_done || is_streaming || content_length.is_none() { let is_final_chunk = is_done; - // Move the body buffer's bytes out — process_body_buffer takes `&mut self.state` - // and may mutate `compressed_body` (via decompress_bytes' reset) or `decoded_body`, - // so any `&` into `self.state` held across the call would be aliased UB. - let buffer_snap = core::mem::take(&mut self.state.get_body_buffer().list); - let processed = self - .state - .process_body_buffer(buffer_snap, is_final_chunk)?; + let processed = self.process_received_body(is_final_chunk)?; // We can only use the libdeflate fast path when we are not streaming // If we ever call processBodyBuffer again, it cannot go through the fast path. @@ -4630,6 +4693,7 @@ impl<'a> HTTPClient<'a> { // Close-delimited bodies still need per-packet decompression, but // a non-streaming consumer must not see per-packet progress: the // terminal callback (on close) is the first to carry metadata. + let is_done = is_done && !self.state.has_pending_compressed(); return Ok(is_done || (processed && is_streaming)); } Ok(false) @@ -4640,7 +4704,10 @@ impl<'a> HTTPClient<'a> { incoming_data: &[u8], ) -> crate::Result { let small_len = 16 * 1024usize; - if incoming_data.len() <= small_len && self.state.get_body_buffer().list.is_empty() { + if incoming_data.len() <= small_len + && self.state.get_body_buffer().list.is_empty() + && !(self.state.encoding.is_compressed() && self.signals.is_demand_driven()) + { self.handle_response_body_chunked_encoding_from_single_packet(incoming_data) } else { self.handle_response_body_chunked_encoding_from_multiple_packets(incoming_data) @@ -4703,10 +4770,7 @@ impl<'a> HTTPClient<'a> { { // If we're streaming, we cannot use the libdeflate fast path self.state.flags.is_libdeflate_fast_path_disabled = true; - // Move the - // bytes out so no `&` into self.state aliases the `&mut self.state` call. - let buffer_snap = core::mem::take(&mut self.state.get_body_buffer().list); - return self.state.process_body_buffer(buffer_snap, false); + return self.process_received_body(false); } return Ok(false); @@ -4714,14 +4778,12 @@ impl<'a> HTTPClient<'a> { // Done _ => { self.state.flags.received_last_chunk = true; - // Move the - // bytes out so no `&` into self.state aliases the `&mut self.state` call. - let buffer_snap = core::mem::take(&mut self.state.get_body_buffer().list); - let _ = self.state.process_body_buffer(buffer_snap, true)?; + let processed = self.process_received_body(true)?; self.report_progress(buffer_len); - return Ok(true); + // A held body ends when `drain_response_body` has pumped it dry, not here. + return Ok(processed || !self.state.has_pending_compressed()); } } } @@ -4780,11 +4842,7 @@ impl<'a> HTTPClient<'a> { // If we're streaming, we cannot use the libdeflate fast path self.state.flags.is_libdeflate_fast_path_disabled = true; - // Move - // the bytes out so no `&` into self.state aliases the `&mut self.state` - // taken by process_body_buffer (which mutates compressed_body/decoded_body). - let buffer_snap = core::mem::take(&mut self.state.get_body_buffer().list); - return self.state.process_body_buffer(buffer_snap, false); + return self.process_received_body(false); } Ok(false) diff --git a/src/zlib/lib.rs b/src/zlib/lib.rs index efaa04856cde..4c40c34e7afd 100644 --- a/src/zlib/lib.rs +++ b/src/zlib/lib.rs @@ -938,7 +938,15 @@ impl DeflateEncoder { reserve: usize, flush: FlushValue, ) -> (usize, ReturnCode) { - step(&mut self.strm, input, out, reserve, flush, deflate) + step( + &mut self.strm, + input, + out, + reserve, + usize::MAX, + flush, + deflate, + ) } } @@ -1003,6 +1011,11 @@ impl InflateDecoder { rc } + #[inline] + pub fn is_inflating(&self) -> bool { + matches!(self.state, State::Inflating) + } + /// One `inflate()` call writing into `out`'s spare capacity. Same /// contract as [`DeflateEncoder::step`]. pub fn step( @@ -1012,12 +1025,21 @@ impl InflateDecoder { reserve: usize, flush: FlushValue, ) -> (usize, ReturnCode) { - step(&mut self.strm, input, out, reserve, flush, inflate) + step( + &mut self.strm, + input, + out, + reserve, + usize::MAX, + flush, + inflate, + ) } - /// Consume all of `input`, appending decompressed output to `out` - /// (growing by 4096-byte steps, capped at `max_output_size`). Returns - /// `ShortRead` when more input is required and `is_done` is false. + /// Append decompressed output to `out` (growing by 4096-byte steps, capped at + /// `max_output_size`) until `input` is consumed or `out.len()` reaches `max_output`. + /// Returns the input bytes consumed. Returns `ShortRead` when more input is required and + /// `is_done` is false. /// /// The stream state persists across calls so this can be driven one /// body chunk at a time. @@ -1025,10 +1047,12 @@ impl InflateDecoder { &mut self, mut input: &[u8], out: &mut Vec, + max_output: usize, is_done: bool, - ) -> Result<(), ZlibError> { + ) -> Result { + let input_len = input.len(); if matches!(self.state, State::Error) { - return Ok(()); + return Ok(input_len); } if matches!(self.state, State::End) { // A prior call completed a gzip member at the chunk boundary. @@ -1041,17 +1065,29 @@ impl InflateDecoder { return Err(ZlibError::ZlibError); } } else { - return Ok(()); + return Ok(input_len); } } loop { + if out.len() >= max_output { + return Ok(input_len - input.len()); + } let remaining = self.max_output_size.saturating_sub(out.len()); if remaining == 0 { self.state = State::Error; return Err(ZlibError::ZlibError); } - let reserve = remaining.min(4096); - let (consumed, rc) = self.step(input, out, reserve, FlushValue::NoFlush); + let budget = max_output - out.len(); + let reserve = remaining.min(4096).min(budget); + let (consumed, rc) = step( + &mut self.strm, + input, + out, + reserve, + budget, + FlushValue::NoFlush, + inflate, + ); input = &input[consumed..]; self.state = State::Inflating; if out.len() > self.max_output_size { @@ -1072,7 +1108,7 @@ impl InflateDecoder { } continue; } - return Ok(()); + return Ok(input_len); } ReturnCode::MemError => { self.state = State::Error; @@ -1141,6 +1177,7 @@ fn step( input: &[u8], out: &mut Vec, reserve: usize, + limit: usize, flush: FlushValue, op: unsafe extern "C" fn(*mut zStream_struct, FlushValue) -> ReturnCode, ) -> (usize, ReturnCode) { @@ -1153,7 +1190,7 @@ fn step( strm.avail_in = in_len as uInt; let spare = out.spare_capacity_mut(); - let out_len = spare.len().min(u32::MAX as usize); + let out_len = spare.len().min(limit).min(u32::MAX as usize); strm.next_out = spare.as_mut_ptr().cast::(); strm.avail_out = out_len as uInt; diff --git a/src/zstd/lib.rs b/src/zstd/lib.rs index b8d441a172e8..f4c917dad1f6 100644 --- a/src/zstd/lib.rs +++ b/src/zstd/lib.rs @@ -551,6 +551,8 @@ pub struct StreamingDecoder { /// Decompression-bomb guard: `decompress` errors instead of growing the /// output past this many bytes. Defaults to unbounded. pub(crate) max_output_size: usize, + /// zstd filled its last window and may hold more. `max_output` can end a call there. + output_full: bool, } impl StreamingDecoder { @@ -563,30 +565,35 @@ impl StreamingDecoder { stream, state: State::Uninitialized, max_output_size: usize::MAX, + output_full: false, }) } - /// Consume all of `input`, appending decompressed bytes to `out` - /// (growing in 4096-byte steps). Returns `ShortRead` when more input is - /// required and `is_done` is false. + #[inline] + pub fn is_inflating(&self) -> bool { + matches!(self.state, State::Inflating) + } + + /// Append decompressed bytes to `out` (growing in 4096-byte steps) until `input` is + /// consumed or `out.len()` reaches `max_output`. Returns the input bytes consumed. + /// Returns `ShortRead` when more input is required and `is_done` is false. pub fn decompress( &mut self, input: &[u8], out: &mut Vec, + max_output: usize, is_done: bool, - ) -> core::result::Result<(), ZstdError> { + ) -> core::result::Result { if matches!(self.state, State::End | State::Error) { - return Ok(()); + return Ok(input.len()); } let mut total_in = 0usize; - // zstd may hold decoded bytes it could not fit into the last output - // window. Call it again with no input until it leaves the window short. - let mut output_full = false; while matches!(self.state, State::Uninitialized | State::Inflating) { let next_in = &input[total_in..]; - if next_in.is_empty() && !output_full { + // Call zstd again with no input until it leaves the window short. + if next_in.is_empty() && !self.output_full { if is_done { if self.state == State::Inflating { self.state = State::Error; @@ -594,7 +601,11 @@ impl StreamingDecoder { } self.state = State::End; } - return Ok(()); + return Ok(total_in); + } + + if out.len() >= max_output { + return Ok(total_in); } let remaining_output = self.max_output_size.saturating_sub(out.len()); @@ -607,6 +618,7 @@ impl StreamingDecoder { self.state = State::Error; return Err(ZstdError::OutOfMemory); } + let budget = max_output - out.len(); let spare = out.spare_capacity_mut(); let mut in_buf = c::ZSTD_inBuffer { src: next_in.as_ptr().cast::(), @@ -615,7 +627,7 @@ impl StreamingDecoder { }; let mut out_buf = c::ZSTD_outBuffer { dst: spare.as_mut_ptr().cast::(), - size: spare.len().min(remaining_output), + size: spare.len().min(remaining_output).min(budget), pos: 0, }; @@ -634,20 +646,21 @@ impl StreamingDecoder { let bytes_written = out_buf.pos; let bytes_read = in_buf.pos; - output_full = bytes_written == out_buf.size; + self.output_full = bytes_written == out_buf.size; // SAFETY: zstd wrote exactly `bytes_written` initialized bytes into // the spare capacity starting at the previous len. unsafe { bun_core::vec::commit_spare(out, bytes_written) }; total_in += bytes_read; if rc == 0 { - // Frame complete. + // Frame complete, and fully flushed. self.state = State::Uninitialized; + self.output_full = false; if total_in >= input.len() { if is_done { self.state = State::End; } - return Ok(()); + return Ok(total_in); } // More input available — reinitialize for the next frame. // SAFETY: stream is a valid DStream. @@ -658,7 +671,7 @@ impl StreamingDecoder { self.state = State::Inflating; if bytes_read == next_in.len() { - if output_full { + if self.output_full { continue; } if is_done { @@ -668,7 +681,7 @@ impl StreamingDecoder { return Err(ZstdError::ShortRead); } } - Ok(()) + Ok(total_in) } } diff --git a/test/js/web/fetch/fetch-backpressure.test.ts b/test/js/web/fetch/fetch-backpressure.test.ts index f71611248edf..f697b83759b3 100644 --- a/test/js/web/fetch/fetch-backpressure.test.ts +++ b/test/js/web/fetch/fetch-backpressure.test.ts @@ -9,11 +9,20 @@ import { stat } from "node:fs/promises"; import { createServer } from "node:http"; import { createSecureServer } from "node:http2"; import { createServer as createHttpsServer } from "node:https"; -import { createServer as createTcpServer } from "node:net"; +import { connect, createServer as createTcpServer } from "node:net"; import { join } from "node:path"; import { Readable, Writable } from "node:stream"; import { pipeline } from "node:stream/promises"; -import { gzipSync } from "node:zlib"; +import { createServer as createTlsServer } from "node:tls"; +import { + brotliCompressSync, + createZstdCompress, + deflateRawSync, + deflateSync, + gzipSync, + constants as zlibConstants, + zstdCompressSync, +} from "node:zlib"; const CHUNK = 64 * 1024; const COUNT = 256; // 16 MiB @@ -449,6 +458,7 @@ describe.concurrent("fetch() receive backpressure — Readable.fromWeb bridge", import net from "node:net"; import { Readable } from "node:stream"; import { pipeline } from "node:stream/promises"; +import { createServer as createTlsServer } from "node:tls"; const C = 40, MB = 32; const CHUNK = Buffer.alloc(64 * 1024, 0x41), COUNT = MB * 16, TOTAL = CHUNK.length * COUNT; let peak = process.memoryUsage.rss(); @@ -506,6 +516,343 @@ describe.concurrent("fetch() receive backpressure — Readable.fromWeb bridge", // so withholding WINDOW_UPDATE only takes effect past that. Asserting a tight // RSS bound for h2 needs that window lowered, which is a separate change. +describe.concurrent("fetch() receive backpressure — the decompressor does not run ahead of the reader", () => { + // Pausing the socket bounds COMPRESSED bytes, so a high-ratio body needs its own bound: one + // 512 KB read of 1000:1 input inflates to ~500 MB. Each of these bodies is 256 MB of zeros and + // at most 270 KB on the wire, so one read hands the client all of it. + const DECODED = 256 * 1024 * 1024; + const READS = 17; + // Unbounded, the client holds all of DECODED (+280 to +305 MB in a debug build). Bounded, it + // holds the reader's chunks and one decode pass (+12 to +24 MB). + const PEAK_LIMIT = (isASAN || isDebug ? 96 : 64) * 1024 * 1024; + + type Enc = "gzip" | "deflate" | "br" | "zstd"; + const bombs: Partial> = {}; + function bombFor(enc: Enc) { + return (bombs[enc] ??= (() => { + // A debug build takes seconds to really compress 256 MB, so one compressed MB is repeated: + // gzip members and zstd frames concatenate, and so do raw deflate blocks after a full + // flush (an empty final stored block then ends the stream). brotli has no such shortcut. + const mb = Buffer.alloc(1 << 20); + const repeat = (piece: Buffer) => Buffer.concat(Array(DECODED / mb.length).fill(piece)); + if (enc === "gzip") return repeat(gzipSync(mb, { level: 9 })); + if (enc === "zstd") return repeat(zstdCompressSync(mb)); + if (enc === "deflate") { + const block = deflateRawSync(mb, { level: 9, finishFlush: zlibConstants.Z_FULL_FLUSH }); + return Buffer.concat([repeat(block), Buffer.from([1, 0, 0, 0xff, 0xff])]); + } + return brotliCompressSync(Buffer.alloc(DECODED), { params: { [zlibConstants.BROTLI_PARAM_QUALITY]: 1 } }); + })()); + } + + async function listening(srv: import("node:net").Server) { + const sockets = new Set(); + srv.on("connection", s => sockets.add(s)); + srv.listen(0, "127.0.0.1"); + await once(srv, "listening"); + return { + port: (srv.address() as import("node:net").AddressInfo).port, + [Symbol.asyncDispose]: () => { + for (const s of sockets) s.destroy(); + return new Promise(r => srv.close(() => r())); + }, + }; + } + + // Close-delimited, and the origin never closes: for the client this body does not end. + async function serveBomb(enc: Enc, secure: boolean) { + const bomb = bombFor(enc); + const handler = (s: import("node:net").Socket) => { + s.on("error", () => {}); + s.once("data", () => { + s.write(`HTTP/1.1 200 OK\r\nContent-Encoding: ${enc}\r\nConnection: close\r\n\r\n`); + s.write(bomb); + }); + }; + const server = await listening(secure ? createTlsServer(tls, handler) : createTcpServer(handler)); + return { ...server, url: `${secure ? "https" : "http"}://127.0.0.1:${server.port}/` }; + } + + // A CONNECT proxy that pipes both ways. bun does not pause a tunnelled socket, so the origin's + // bytes keep arriving while the reader is paused. + async function serveConnectProxy() { + const server = await listening( + createTcpServer(client => { + let upstream: import("node:net").Socket | undefined; + client.on("error", () => upstream?.destroy()); + client.on("close", () => upstream?.destroy()); + client.once("data", head => { + const [, target] = head.toString("latin1").split(" "); + const colon = target.lastIndexOf(":"); + upstream = connect(Number(target.slice(colon + 1)), target.slice(0, colon), () => { + client.write("HTTP/1.1 200 Connection Established\r\n\r\n"); + client.pipe(upstream!); + upstream!.pipe(client); + }); + upstream.on("error", () => client.destroy()); + upstream.on("close", () => client.destroy()); + }); + }), + ); + return { ...server, url: `http://127.0.0.1:${server.port}` }; + } + + // Takes one chunk, lets the client's memory settle, takes a few more, and reports the largest + // RSS growth it saw. Nothing here waits on the clock for the decoder: the pulls pace it. + const READ_A_LITTLE = /* js */ ` + const base = process.memoryUsage.rss(); + let peak = 0; + const sample = () => (peak = Math.max(peak, process.memoryUsage.rss() - base)); + const res = await fetch(url, opts); + const reader = res.body.getReader(); + const first = await reader.read(); + for (let last = sample(), stable = 0; stable < 3; ) { + await Bun.sleep(20); + const now = process.memoryUsage.rss() - base; + stable = Math.abs(now - last) < (1 << 20) ? stable + 1 : 0; + last = now; + sample(); + } + let got = first.value.byteLength; + for (let i = 1; i < ${READS}; i++) { + got += (await reader.read()).value.byteLength; + sample(); + } + const zeros = !first.value.some(b => b !== 0); + await reader.cancel(); + process.stdout.write(JSON.stringify({ got, peak, zeros })); + `; + + async function runClient(url: string, opts: object, script: string) { + await using proc = Bun.spawn({ + cmd: [bunExe(), "-e", `const url=${JSON.stringify(url)};const opts=${JSON.stringify(opts)};${script}`], + env: bunEnv, + stdout: "pipe", + stderr: "pipe", + }); + const [stdout, stderr, exitCode] = await Promise.all([proc.stdout.text(), proc.stderr.text(), proc.exited]); + if (!stdout) throw new Error(`client exited ${exitCode}: ${stderr}`); + return { ...JSON.parse(stdout), stderr, exitCode }; + } + + function expectBounded({ + got, + peak, + zeros, + exitCode, + }: { + got: number; + peak: number; + zeros: boolean; + exitCode: number; + }) { + expect({ + got: got > 0, + zeros, + peakUnder: peak < PEAK_LIMIT || { peakMB: peak >> 20, limitMB: PEAK_LIMIT >> 20 }, + }).toEqual({ got: true, zeros: true, peakUnder: true }); + expect(exitCode).toBe(0); + } + + test.each(["gzip", "deflate", "br", "zstd"] as Enc[])( + "%s: a reader that takes a little holds a little", + async enc => { + await using server = await serveBomb(enc, false); + expectBounded(await runClient(server.url, {}, READ_A_LITTLE)); + }, + ); + + test("zstd through a CONNECT proxy: a reader that takes a little holds a little", async () => { + await using server = await serveBomb("zstd", true); + await using proxy = await serveConnectProxy(); + const opts = { proxy: proxy.url, tls: { rejectUnauthorized: false } }; + expectBounded(await runClient(server.url, opts, READ_A_LITTLE)); + }); + + // A live stream: two flushed messages in one packet, then the origin goes quiet with the frame + // open. The budget ends a pass inside the second message's last block, after zstd has taken + // every input byte and while it still holds decoded output. No further packet will come to + // shake that output loose. + test("zstd: output the decoder still holds arrives without another packet", async () => { + const z = createZstdCompress(); + const parts: Buffer[] = []; + z.on("data", d => parts.push(d)); + const message = (n: number) => { + const b = Buffer.alloc(n); + for (let i = 0; i < n; i += 40_009) b[i] = 1 + (i % 251); + return b; + }; + for (const n of [100_000, 200_000]) { + z.write(message(n)); + await new Promise(r => z.flush(zlibConstants.ZSTD_e_flush, () => r())); + } + const wire = Buffer.concat(parts); + await using server = await listening( + createTcpServer(s => { + s.on("error", () => {}); + s.once("data", () => + s.write( + Buffer.concat([ + Buffer.from(`HTTP/1.1 200 OK\r\nContent-Encoding: zstd\r\nTransfer-Encoding: chunked\r\n\r\n`), + Buffer.from(wire.length.toString(16) + "\r\n"), + wire, + Buffer.from("\r\n"), + ]), + ), + ); + }), + ); + const res = await fetch(`http://127.0.0.1:${server.port}/`); + const reader = res.body!.getReader(); + let total = 0; + while (total < 300_000) total += (await reader.read()).value!.byteLength; + await reader.cancel(); + expect(total).toBe(300_000); + }); + + // `await (await fetch(url)).text()` turns the consumer to BufferAll, possibly while the HTTP + // thread is inside a pass it began under the reader's budget. A paused consumer asks for the + // rest of such a pass; one that has just stopped being a reader never does, and without the + // client going on by itself the request hangs. Which thread gets there first is down to + // scheduling, so this repeats the request; a debug build is too slow to repeat it enough. + test.skipIf(isDebug || isASAN)("text() racing a budgeted pass gets the whole body every time", async () => { + const SIZE = 400_000; + const body = gzipSync(Buffer.alloc(SIZE, "abcdefghij"), { level: 1 }); + await using server = await listening( + createTcpServer(s => { + s.on("error", () => {}); + s.on("data", () => { + // The head goes out a tick before the body, so the caller has the Response in hand + // when the body's only packet is decoded. + s.write(`HTTP/1.1 200 OK\r\nContent-Encoding: gzip\r\nTransfer-Encoding: chunked\r\n\r\n`); + setImmediate(() => + s.write( + Buffer.concat([Buffer.from(body.length.toString(16) + "\r\n"), body, Buffer.from("\r\n0\r\n\r\n")]), + ), + ); + }); + }), + ); + const script = /* js */ ` + const TOTAL = 4000, CONCURRENCY = 16; + let started = 0, done = 0, wrong = 0; + async function worker() { + while (started < TOTAL) { + started++; + const text = await (await fetch(url)).text(); + if (text.length !== ${SIZE}) wrong++; + done++; + } + } + await Promise.all(Array.from({ length: CONCURRENCY }, worker)); + process.stdout.write(JSON.stringify({ done, wrong })); + `; + const { done, wrong, exitCode } = await runClient(`http://127.0.0.1:${server.port}/`, {}, script); + expect({ done, wrong }).toEqual({ done: 4000, wrong: 0 }); + expect(exitCode).toBe(0); + }); + + // A bounded decode delivers every byte. The size is not a multiple of anything, so the budget + // ends passes in the middle of blocks, and of the decoders' own flush windows. + describe("a bounded decode still delivers the whole body", () => { + const SIZE = 8 * 1024 * 1024 + 77_777; + // Zero runs broken by short islands: the ratio stays high, the block structure irregular. + const raw = Buffer.alloc(SIZE); + for (let i = 0, x = 12345; i < SIZE; i += 70_001) { + for (let j = 0; j < 257 && i + j < SIZE; j++) raw[i + j] = (x = (x * 1103515245 + 12345) & 0x7fffffff) & 0xff; + } + const digest = md5(raw); + + type Kind = Enc | "br-hq"; + const bodies: { [k: string]: Buffer } = {}; + // br-hq packs the body into a few hundred bytes, so the decoder has taken all of its input + // long before it has produced the budget's worth of output: the decoder holds the rest. + function bodyFor(kind: Kind) { + const q = zlibConstants.BROTLI_PARAM_QUALITY; + return (bodies[kind] ??= + kind === "gzip" + ? gzipSync(raw, { level: 1 }) + : kind === "deflate" + ? deflateSync(raw, { level: 1 }) + : kind === "br" + ? brotliCompressSync(raw, { params: { [q]: 0 } }) + : kind === "br-hq" + ? brotliCompressSync(Buffer.alloc(SIZE), { params: { [q]: 4 } }) + : zstdCompressSync(raw, { level: 1 })); + } + + // Content-Length bodies end with the origin's FIN, which reaches a client that still holds + // most of the body undecoded. Chunked ones stay open. + async function serveBody(kind: Kind, chunked: boolean) { + const body = bodyFor(kind); + const enc = kind === "br-hq" ? "br" : kind; + const server = await listening( + createTcpServer(s => { + s.on("error", () => {}); + s.once("data", () => { + if (chunked) { + s.write(`HTTP/1.1 200 OK\r\nContent-Encoding: ${enc}\r\nTransfer-Encoding: chunked\r\n\r\n`); + s.write(`${body.length.toString(16)}\r\n`); + s.write(body); + s.write("\r\n0\r\n\r\n"); + } else { + s.write( + `HTTP/1.1 200 OK\r\nContent-Encoding: ${enc}\r\nContent-Length: ${body.length}\r\nConnection: close\r\n\r\n`, + ); + s.end(body); + } + }); + }), + ); + return { ...server, url: `http://127.0.0.1:${server.port}/` }; + } + + const STREAM = /* js */ ` + const res = await fetch(url, opts); + const hasher = new Bun.CryptoHasher("md5"); + let total = 0, chunks = 0; + for await (const chunk of res.body) { + total += chunk.byteLength; + chunks++; + hasher.update(chunk); + } + process.stdout.write(JSON.stringify({ total, several: chunks > 1, digest: hasher.digest("hex") })); + `; + const BUFFER = /* js */ ` + const bytes = await (await fetch(url, opts)).bytes(); + const digest = new Bun.CryptoHasher("md5").update(bytes).digest("hex"); + process.stdout.write(JSON.stringify({ total: bytes.byteLength, digest })); + `; + + const cases: [Kind, boolean][] = [ + ["gzip", false], + ["gzip", true], + ["deflate", false], + ["br", false], + ["br-hq", false], + ["zstd", false], + ["zstd", true], + ]; + describe.each(cases)("%s chunked=%p", (kind, chunked) => { + // `several`: the budget split the decode across pulls instead of one SIZE-byte chunk. + test.each([ + ["a streaming reader", STREAM, { several: true }], + ["res.bytes()", BUFFER, {}], + ])("%s", async (_, script, seen) => { + await using server = await serveBody(kind, chunked); + const { stderr, exitCode, ...result } = await runClient(server.url, {}, script); + expect(stderr).toBe(""); + expect(result).toEqual({ + total: SIZE, + digest: kind === "br-hq" ? md5(Buffer.alloc(SIZE)) : digest, + ...seen, + }); + expect(exitCode).toBe(0); + }); + }); + }); +}); + describe.concurrent("fetch() receive backpressure — buffered consumers are not throttled", () => { const cases: [string, (r: Response) => Promise][] = [ ["res.arrayBuffer()", async r => md5(new Uint8Array(await r.arrayBuffer()))],