diff --git a/src/proto/streams/recv.rs b/src/proto/streams/recv.rs index e57dc2b3..20a9ce53 100644 --- a/src/proto/streams/recv.rs +++ b/src/proto/streams/recv.rs @@ -506,7 +506,7 @@ impl Recv { self.release_connection_capacity(stream.in_flight_recv_data, task); stream.in_flight_recv_data = 0; - self.clear_recv_buffer(stream); + self.clear_recv_buffer(stream, task); } /// Set the "target" connection window size. @@ -933,9 +933,22 @@ impl Recv { stream.notify_push(); } - pub(super) fn clear_recv_buffer(&mut self, stream: &mut Stream) { - while stream.pending_recv.pop_front(&mut self.buffer).is_some() { - // drop it + pub(super) fn clear_recv_buffer(&mut self, stream: &mut Stream, task: &mut Option) { + let mut to_release: WindowSize = 0; + while let Some(event) = stream.pending_recv.pop_front(&mut self.buffer) { + if let Event::Data(data) = &event { + to_release = to_release + .saturating_add(data.len() as WindowSize) + .min(stream.in_flight_recv_data); + } + } + // Release flow control capacity. Cases: + // * User read data but hasn't released: buf=0, in_flight>0 -> release 0 + // * User released without reading: buf>0, in_flight=0 -> release 0 + // * Normal drop without reading: buf=in_flight -> full release + if to_release > 0 { + stream.in_flight_recv_data -= to_release; + self.release_connection_capacity(to_release, task); } } @@ -1253,6 +1266,52 @@ impl Recv { } } +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn clear_recv_buffer_caps_capacity_before_overflow() { + const FRAME_LEN: usize = 1 << 20; + const FRAME_COUNT: usize = (u32::MAX as usize / FRAME_LEN) + 1; + + let config = Config { + initial_max_send_streams: 0, + local_max_buffer_size: 0, + local_next_stream_id: 2.into(), + local_push_enabled: false, + extended_connect_protocol_enabled: false, + local_reset_duration: Duration::ZERO, + local_reset_max: 0, + remote_reset_max: 0, + remote_init_window_sz: DEFAULT_INITIAL_WINDOW_SIZE, + remote_max_initiated: None, + local_max_error_reset_streams: None, + }; + let mut recv = Recv::new(peer::Dyn::Server, &config); + let mut store = Store::new(); + let mut stream = store.insert( + StreamId::from(1), + Stream::new(StreamId::from(1), 0, DEFAULT_INITIAL_WINDOW_SIZE), + ); + let data = Bytes::from(vec![0; FRAME_LEN]); + + for _ in 0..FRAME_COUNT { + stream + .pending_recv + .push_back(&mut recv.buffer, Event::Data(data.clone())); + } + stream.in_flight_recv_data = DEFAULT_INITIAL_WINDOW_SIZE; + recv.in_flight_data = DEFAULT_INITIAL_WINDOW_SIZE; + + recv.clear_recv_buffer(&mut stream, &mut None); + + assert!(stream.pending_recv.is_empty()); + assert_eq!(stream.in_flight_recv_data, 0); + assert_eq!(recv.in_flight_data, 0); + } +} + // ===== impl Open ===== impl Open { diff --git a/src/proto/streams/streams.rs b/src/proto/streams/streams.rs index 7a6ef66f..ba326d0d 100644 --- a/src/proto/streams/streams.rs +++ b/src/proto/streams/streams.rs @@ -1551,7 +1551,9 @@ impl OpaqueStreamRef { let mut stream = me.store.resolve(self.key); stream.is_recv = false; - me.actions.recv.clear_recv_buffer(&mut stream); + me.actions + .recv + .clear_recv_buffer(&mut stream, &mut me.actions.task); } pub fn stream_id(&self) -> StreamId { diff --git a/tests/h2-tests/tests/flow_control.rs b/tests/h2-tests/tests/flow_control.rs index 5ed61a70..11d1207e 100644 --- a/tests/h2-tests/tests/flow_control.rs +++ b/tests/h2-tests/tests/flow_control.rs @@ -544,6 +544,61 @@ async fn stream_error_release_connection_capacity() { join(srv, client).await; } +#[tokio::test] +async fn recv_stream_drop_releases_only_buffered_connection_capacity() { + h2_support::trace_init!(); + + const FRAME_LEN: usize = 16_384; + const TOTAL_LEN: usize = FRAME_LEN * 2; + + // Exercise all relationships between buffered and in-flight capacity: + // equal, buffered > in-flight, and buffered < in-flight. + for read_frames in 0usize..=1 { + for released_frames in 0usize..=1 { + let expected_used = read_frames.saturating_sub(released_frames) * FRAME_LEN; + let (io, mut peer) = mock::new(); + + let peer = async move { + let _ = peer.assert_server_handshake().await; + peer.send_frame(frames::headers(1).request("POST", "https://example.com/")) + .await; + for _ in 0..2 { + peer.send_frame(frames::data(1, vec![0; FRAME_LEN])).await; + } + + peer.recv_frame(frames::window_update(0, (TOTAL_LEN - expected_used) as u32)) + .await; + if released_frames > 0 { + peer.recv_frame(frames::window_update(1, FRAME_LEN as u32)) + .await; + } + }; + + let server = async move { + let mut server = server::handshake(io).await.unwrap(); + let (request, _respond) = server.next().await.unwrap().unwrap(); + let mut body = request.into_body(); + + for _ in 0..read_frames { + assert_eq!(body.data().await.unwrap().unwrap().len(), FRAME_LEN); + } + + let mut flow = body.flow_control().clone(); + for _ in 0..released_frames { + flow.release_capacity(FRAME_LEN).unwrap(); + } + drop(body); + + assert_eq!(flow.used_capacity(), expected_used); + + let _ = server.next().await; + }; + + join(peer, server).await; + } + } +} + // Regression test for TODO #[tokio::test] async fn padded_data_stream_error_releases_connection_capacity() {