Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
52 changes: 40 additions & 12 deletions src/proto/streams/recv.rs
Original file line number Diff line number Diff line change
Expand Up @@ -64,11 +64,17 @@ pub(super) struct Recv {
#[derive(Debug)]
pub(super) enum Event {
Headers(peer::PollMessage),
Data(Bytes),
Data(DataEvent),
Trailers(HeaderMap),
InformationalHeaders(peer::PollMessage),
}

#[derive(Debug)]
pub(super) struct DataEvent {
pub(super) payload: Bytes,
pub(super) is_budgeted: bool,
}

#[derive(Debug)]
pub(super) enum RecvHeaderBlockError<T> {
Oversize(T),
Expand Down Expand Up @@ -493,7 +499,12 @@ impl Recv {
}

/// Release any unclaimed capacity for a closed stream.
pub fn release_closed_capacity(&mut self, stream: &mut store::Ptr, task: &mut Option<Waker>) {
pub fn release_closed_capacity(
&mut self,
stream: &mut store::Ptr,
task: &mut Option<Waker>,
counts: &mut Counts,
) {
debug_assert_eq!(stream.ref_count, 0);

if stream.in_flight_recv_data != 0 {
Expand All @@ -507,7 +518,7 @@ impl Recv {
stream.in_flight_recv_data = 0;
}

self.clear_recv_buffer(stream, task);
self.clear_recv_buffer(stream, task, counts);
}

/// Set the "target" connection window size.
Expand Down Expand Up @@ -758,7 +769,11 @@ impl Recv {
return Ok(());
}

let event = Event::Data(frame.into_payload());
let is_budgeted = !frame.is_end_stream();
let event = Event::Data(DataEvent {
payload: frame.into_payload(),
is_budgeted,
});

// Push the frame onto the recv buffer
stream.pending_recv.push_back(&mut self.buffer, event);
Expand Down Expand Up @@ -941,12 +956,20 @@ impl Recv {
stream.notify_push();
}

pub(super) fn clear_recv_buffer(&mut self, stream: &mut Stream, task: &mut Option<Waker>) {
pub(super) fn clear_recv_buffer(
&mut self,
stream: &mut Stream,
task: &mut Option<Waker>,
counts: &mut Counts,
) {
let mut to_release: WindowSize = 0;
while let Some(event) = stream.pending_recv.pop_front(&mut self.buffer) {
if let Event::Data(data) = &event {
if data.is_budgeted {
counts.release_data_frame(data.payload.len());
}
to_release = to_release
.saturating_add(data.len() as WindowSize)
.saturating_add(data.payload.len() as WindowSize)
.min(stream.in_flight_recv_data);
}
}
Expand Down Expand Up @@ -1217,9 +1240,9 @@ impl Recv {
&mut self,
cx: &Context,
stream: &mut Stream,
) -> Poll<Option<Result<Bytes, proto::Error>>> {
) -> Poll<Option<Result<DataEvent, proto::Error>>> {
match stream.pending_recv.pop_front(&mut self.buffer) {
Some(Event::Data(payload)) => Poll::Ready(Some(Ok(payload))),
Some(Event::Data(data)) => Poll::Ready(Some(Ok(data))),
Some(event) => {
// Frame is trailer
stream.pending_recv.push_front(&mut self.buffer, event);
Expand Down Expand Up @@ -1305,14 +1328,19 @@ mod tests {
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.pending_recv.push_back(
&mut recv.buffer,
Event::Data(DataEvent {
payload: data.clone(),
is_budgeted: true,
}),
);
}
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);
let mut counts = Counts::new(peer::Dyn::Server, &config);
recv.clear_recv_buffer(&mut stream, &mut None, &mut counts);

assert!(stream.pending_recv.is_empty());
assert_eq!(stream.in_flight_recv_data, 0);
Expand Down
27 changes: 19 additions & 8 deletions src/proto/streams/streams.rs
Original file line number Diff line number Diff line change
Expand Up @@ -638,9 +638,12 @@ impl Inner {

self.counts.transition(stream, |counts, stream| {
let sz = frame.flow_controlled_len();
let is_end_stream = frame.is_end_stream();
let payload_len = frame.payload().len();
let mut res = actions.recv.recv_data(frame, stream);
if res.is_ok() {
// A stream can receive at most one final DATA frame, so it cannot
// be used to create unbounded framing overhead on that stream.
if res.is_ok() && !is_end_stream {
res = counts.record_data_frame(payload_len).map_err(|_| {
tracing::debug!("too many small DATA frames");
Error::library_go_away_data(Reason::ENHANCE_YOUR_CALM, "too_many_data_frames")
Expand Down Expand Up @@ -1510,11 +1513,19 @@ impl OpaqueStreamRef {

let mut stream = me.store.resolve(self.key);

let poll = me.actions.recv.poll_data(cx, &mut stream);
if let Poll::Ready(Some(Ok(ref payload))) = poll {
me.counts.release_data_frame(payload.len());
}
poll
me.actions
.recv
.poll_data(cx, &mut stream)
.map(|result| match result {
Some(Ok(data)) => {
if data.is_budgeted {
me.counts.release_data_frame(data.payload.len());
}
Some(Ok(data.payload))
}
Some(Err(err)) => Some(Err(err)),
None => None,
})
}

pub fn poll_trailers(&mut self, cx: &Context) -> Poll<Option<Result<HeaderMap, proto::Error>>> {
Expand Down Expand Up @@ -1564,7 +1575,7 @@ impl OpaqueStreamRef {
stream.is_recv = false;
me.actions
.recv
.clear_recv_buffer(&mut stream, &mut me.actions.task);
.clear_recv_buffer(&mut stream, &mut me.actions.task, &mut me.counts);
}

pub fn stream_id(&self) -> StreamId {
Expand Down Expand Up @@ -1659,7 +1670,7 @@ fn drop_stream_ref(inner: &Mutex<Inner>, key: store::Key) {
// it anymore.
actions
.recv
.release_closed_capacity(stream, &mut actions.task);
.release_closed_capacity(stream, &mut actions.task, counts);

// We won't be able to reach our push promises anymore
let mut ppp = stream.pending_push_promises.take();
Expand Down
123 changes: 123 additions & 0 deletions tests/h2-tests/tests/stream_states.rs
Original file line number Diff line number Diff line change
Expand Up @@ -284,6 +284,129 @@ async fn too_many_small_data_frames_sends_goaway() {
join(mock, h2).await;
}

#[tokio::test]
async fn many_small_final_data_frames_do_not_exhaust_budget() {
h2_support::trace_init!();

const NUM_STREAMS: u32 = 200;

let (io, mut srv) = mock::new();
let (done_tx, done_rx) = oneshot::channel();

let mock = async move {
let settings = srv.assert_client_handshake().await;
assert_default_settings!(settings);

for i in 0..NUM_STREAMS {
let stream_id = 1 + i * 2;
srv.recv_frame(
frames::headers(stream_id)
.request("GET", "https://http2.akamai.com/")
.eos(),
)
.await;
}

for i in 0..NUM_STREAMS {
let stream_id = 1 + i * 2;
srv.send_frame(frames::headers(stream_id).response(200))
.await;
srv.send_frame(frames::data(stream_id, "a").eos()).await;
}

done_rx.await.unwrap();
};

let h2 = async move {
let (mut client, h2) = client::handshake(io).await.unwrap();

let requests = async move {
let mut responses = Vec::new();
for _ in 0..NUM_STREAMS {
poll_fn(|cx| client.poll_ready(cx)).await.unwrap();
let request = Request::builder()
.uri("https://http2.akamai.com/")
.body(())
.unwrap();
responses.push(client.send_request(request, true).unwrap().0);
}

// Wait for every response without polling any response body. This
// ensures all final DATA frames can be buffered concurrently.
let mut received = Vec::new();
for response in responses {
received.push(response.await.unwrap());
}
assert_eq!(received.len(), NUM_STREAMS as usize);
done_tx.send(()).unwrap();
};

tokio::spawn(requests);
h2.await.unwrap();
};

join(mock, h2).await;
}

#[tokio::test]
async fn dropping_buffered_data_frames_releases_budget() {
h2_support::trace_init!();

const NUM_STREAMS: u32 = 200;

let (io, mut srv) = mock::new();
let (done_tx, done_rx) = oneshot::channel();

let mock = async move {
let settings = srv.assert_client_handshake().await;
assert_default_settings!(settings);

for i in 0..NUM_STREAMS {
let stream_id = 1 + i * 2;
srv.recv_frame(
frames::headers(stream_id)
.request("GET", "https://http2.akamai.com/")
.eos(),
)
.await;
srv.send_frame(frames::headers(stream_id).response(409))
.await;
srv.send_frame(frames::data(stream_id, "a")).await;
srv.send_frame(frames::data(stream_id, "b").eos()).await;
}

done_rx.await.unwrap();
};

let h2 = async move {
let (mut client, h2) = client::handshake(io).await.unwrap();

let requests = async move {
for _ in 0..NUM_STREAMS {
poll_fn(|cx| client.poll_ready(cx)).await.unwrap();
let request = Request::builder()
.uri("https://http2.akamai.com/")
.body(())
.unwrap();
let response = client.send_request(request, true).unwrap().0.await.unwrap();
assert_eq!(response.status(), StatusCode::CONFLICT);

let mut body = response.into_body();
while body.flow_control().used_capacity() < 2 {
tokio::task::yield_now().await;
}
drop(body);
}
done_tx.send(()).unwrap();
};

tokio::spawn(requests);
h2.await.unwrap();
};

join(mock, h2).await;
}

#[tokio::test]
async fn send_headers_recv_data_single_frame() {
h2_support::trace_init!();
Expand Down