diff --git a/src/body/incoming.rs b/src/body/incoming.rs index b026a24d6b..21525e8950 100644 --- a/src/body/incoming.rs +++ b/src/body/incoming.rs @@ -3,6 +3,8 @@ use std::pin::Pin; use std::task::{Context, Poll}; use bytes::Bytes; +#[cfg(all(feature = "http2", feature = "server"))] +use futures_channel::oneshot; #[cfg(all( any(feature = "http1", feature = "http2"), any(feature = "client", feature = "server") @@ -54,10 +56,18 @@ enum Kind { ping: ping::Recorder, recv: h2::RecvStream, }, + #[cfg(all(feature = "http2", feature = "server"))] + ExpectContinue(Box), #[cfg(feature = "ffi")] Ffi(crate::ffi::UserBody), } +#[cfg(all(feature = "http2", feature = "server"))] +struct ExpectContinue { + body: Incoming, + tx: oneshot::Sender<()>, +} + /// A sender half created through [`Body::channel()`]. /// /// Useful when wanting to stream chunks from another thread. @@ -127,6 +137,14 @@ impl Incoming { }) } + #[cfg(all(feature = "http2", feature = "server"))] + pub(crate) fn with_expect_continue(self, tx: oneshot::Sender<()>) -> Self { + Incoming::new(Kind::ExpectContinue(Box::new(ExpectContinue { + body: self, + tx, + }))) + } + #[cfg(feature = "ffi")] pub(crate) fn as_ffi_mut(&mut self) -> &mut crate::ffi::UserBody { if !matches!(self.kind, Kind::Ffi(_)) { @@ -226,6 +244,22 @@ impl Body for Incoming { } } + #[cfg(all(feature = "http2", feature = "server"))] + Kind::ExpectContinue(_) => { + let expect = match std::mem::replace(&mut self.kind, Kind::Empty) { + Kind::ExpectContinue(expect) => expect, + _ => unreachable!(), + }; + let ExpectContinue { body, tx } = *expect; + *self = body; + let res = self.as_mut().poll_frame(cx); + // Only ask the client to continue if the body hasn't already arrived. + if res.is_pending() { + let _ = tx.send(()); + } + res + } + #[cfg(feature = "ffi")] Kind::Ffi(body) => body.poll_data(cx), } @@ -238,6 +272,8 @@ impl Body for Incoming { Kind::Chan { content_length, .. } => *content_length == DecodedLength::ZERO, #[cfg(all(feature = "http2", any(feature = "client", feature = "server")))] Kind::H2 { recv: h2, .. } => h2.is_end_stream(), + #[cfg(all(feature = "http2", feature = "server"))] + Kind::ExpectContinue(expect) => expect.body.is_end_stream(), #[cfg(feature = "ffi")] Kind::Ffi(..) => false, } @@ -256,12 +292,14 @@ impl Body for Incoming { } } - match self.kind { + match &self.kind { Kind::Empty => SizeHint::with_exact(0), #[cfg(all(feature = "http1", any(feature = "client", feature = "server")))] - Kind::Chan { content_length, .. } => opt_len(content_length), + Kind::Chan { content_length, .. } => opt_len(*content_length), #[cfg(all(feature = "http2", any(feature = "client", feature = "server")))] - Kind::H2 { content_length, .. } => opt_len(content_length), + Kind::H2 { content_length, .. } => opt_len(*content_length), + #[cfg(all(feature = "http2", feature = "server"))] + Kind::ExpectContinue(expect) => expect.body.size_hint(), #[cfg(feature = "ffi")] Kind::Ffi(..) => SizeHint::default(), } diff --git a/src/headers.rs b/src/headers.rs index caace71f2a..f9b2932462 100644 --- a/src/headers.rs +++ b/src/headers.rs @@ -67,6 +67,25 @@ pub(super) fn content_length_parse(value: &HeaderValue) -> Option { from_digits(value.as_bytes()) } +#[cfg(all(feature = "server", any(feature = "http1", feature = "http2")))] +pub(super) fn expect_continue(value: &HeaderValue) -> bool { + // According to https://datatracker.ietf.org/doc/html/rfc2616#section-14.20 + // Comparison of expectation values is case-insensitive for unquoted tokens + // (including the 100-continue token) + value.as_bytes().eq_ignore_ascii_case(b"100-continue") +} + +// If a message has more than one `Expect` header line, the last one wins, +// same as when parsing HTTP/1 requests. +#[cfg(all(feature = "server", feature = "http2"))] +pub(super) fn expect_last_continue(headers: &HeaderMap) -> bool { + headers + .get_all(http::header::EXPECT) + .iter() + .next_back() + .map_or(false, expect_continue) +} + #[cfg(any(feature = "client", all(feature = "server", feature = "http2")))] pub(super) fn content_length_parse_all(headers: &HeaderMap) -> Option { content_length_parse_all_values(headers.get_all(CONTENT_LENGTH).into_iter()) diff --git a/src/proto/h1/role.rs b/src/proto/h1/role.rs index b83c8ab12b..9b5e3de2f2 100644 --- a/src/proto/h1/role.rs +++ b/src/proto/h1/role.rs @@ -318,10 +318,7 @@ impl Http1Transaction for Server { } } header::EXPECT => { - // According to https://datatracker.ietf.org/doc/html/rfc2616#section-14.20 - // Comparison of expectation values is case-insensitive for unquoted tokens - // (including the 100-continue token) - expect_continue = value.as_bytes().eq_ignore_ascii_case(b"100-continue"); + expect_continue = headers::expect_continue(&value); } header::UPGRADE => { // Upgrades are only allowed with HTTP/1.1 diff --git a/src/proto/h2/server.rs b/src/proto/h2/server.rs index 2098758d9e..241f0af399 100644 --- a/src/proto/h2/server.rs +++ b/src/proto/h2/server.rs @@ -5,10 +5,11 @@ use std::task::{Context, Poll}; use std::time::Duration; use bytes::Bytes; +use futures_channel::oneshot; use futures_core::ready; use h2::server::{Connection, Handshake, SendResponse}; use h2::{Reason, RecvStream}; -use http::{Method, Request}; +use http::{Method, Request, StatusCode}; use pin_project_lite::pin_project; use super::{ping, PipeToSendStream, SendBuf}; @@ -274,14 +275,17 @@ where let is_connect = req.method() == Method::CONNECT; let (mut parts, stream) = req.into_parts(); + let mut expect_continue = None; let (mut req, connect_parts) = if !is_connect { - ( - Request::from_parts( - parts, - IncomingBody::h2(stream, content_length.into(), ping), - ), - None, - ) + let wants_continue = headers::expect_last_continue(&parts.headers) + && !stream.is_end_stream(); + let mut body = IncomingBody::h2(stream, content_length.into(), ping); + if wants_continue { + let (tx, rx) = oneshot::channel(); + body = body.with_expect_continue(tx); + expect_continue = Some(rx); + } + (Request::from_parts(parts, body), None) } else { if content_length.map_or(false, |len| len != 0) { warn!("h2 connect request with non-zero body not supported"); @@ -308,6 +312,7 @@ where let fut = H2Stream::new( service.call(req), connect_parts, + expect_continue, respond, self.date_header, exec.clone(), @@ -382,6 +387,7 @@ pin_project! { #[pin] fut: F, connect_parts: Option, + expect_continue: Option>, }, Body { #[pin] @@ -403,13 +409,18 @@ where fn new( fut: F, connect_parts: Option, + expect_continue: Option>, respond: SendResponse>, date_header: bool, exec: E, ) -> H2Stream { H2Stream { reply: respond, - state: H2StreamState::Service { fut, connect_parts }, + state: H2StreamState::Service { + fut, + connect_parts, + expect_continue, + }, date_header, exec, } @@ -445,6 +456,7 @@ where H2StreamStateProj::Service { fut: h, connect_parts, + expect_continue, } => { let res = match h.poll(cx) { Poll::Ready(Ok(r)) => r, @@ -457,6 +469,25 @@ where debug!("stream received RST_STREAM: {:?}", reason); return Poll::Ready(Err(crate::Error::new_h2(reason.into()))); } + // The service is waiting on an `Expect: 100-continue` body + // before responding: tell the client to send it. + let wanted = expect_continue.as_mut().map(|rx| Pin::new(rx).poll(cx)); + match wanted { + Some(Poll::Ready(Ok(()))) => { + *expect_continue = None; + let mut cont = ::http::Response::new(()); + *cont.status_mut() = StatusCode::CONTINUE; + if let Err(_e) = me.reply.send_informational(cont) { + debug!("send 100-continue error: {}", _e); + } + } + Some(Poll::Ready(Err(_))) => { + // The body was dropped, or didn't have to wait. + // Stop polling the receiver. + *expect_continue = None; + } + Some(Poll::Pending) | None => {} + } return Poll::Pending; } Poll::Ready(Err(e)) => { diff --git a/tests/server.rs b/tests/server.rs index 098ed7a6ee..5583c1c69b 100644 --- a/tests/server.rs +++ b/tests/server.rs @@ -1068,6 +1068,132 @@ async fn expect_continue_waits_for_body_poll() { child.join().expect("client thread"); } +async fn h2_expect_continue_server(svc: S) -> SendRequest +where + S: hyper::service::HttpService> + Send + 'static, + S::Future: Send + 'static, + S::Error: Into>, +{ + let (listener, addr) = setup_tcp_listener(); + tokio::spawn(async move { + let (socket, _) = listener.accept().await.unwrap(); + let _ = http2::Builder::new(TokioExecutor) + .serve_connection(TokioIo::new(socket), svc) + .await; + }); + + let conn = connect_async(addr).await; + let (h2, connection) = h2::client::handshake(conn).await.unwrap(); + tokio::spawn(async move { + let _ = connection.await; + }); + h2.ready().await.unwrap() +} + +fn h2_send_expect_request( + h2: &mut SendRequest, + end_of_stream: bool, +) -> (h2::client::ResponseFuture, SendStream) { + let mut req = Request::post("http://localhost/foo").header("expect", "100-continue"); + if !end_of_stream { + req = req.header("content-length", "5"); + } + h2.send_request(req.body(()).unwrap(), end_of_stream) + .unwrap() +} + +async fn h2_assert_no_informational(response: &mut h2::client::ResponseFuture) { + let info = future::poll_fn(|cx| response.poll_informational(cx)).await; + assert!( + info.is_none(), + "unexpected informational response: {:?}", + info + ); +} + +#[tokio::test] +async fn h2_expect_continue_sends_100_when_body_polled() { + let svc = service_fn(|req: Request| async move { + let body = req.into_body().collect().await?.to_bytes(); + assert_eq!(&body[..], b"hello"); + Ok::<_, hyper::Error>(Response::new(Empty::::new())) + }); + let mut h2 = h2_expect_continue_server(svc).await; + + let (mut response, mut body) = h2_send_expect_request(&mut h2, false); + // The client withholds the body until the server asks for it. + let info = tokio::time::timeout( + Duration::from_secs(1), + future::poll_fn(|cx| response.poll_informational(cx)), + ) + .await + .expect("100 Continue before the timeout") + .expect("an informational response") + .expect("informational response ok"); + assert_eq!(info.status(), StatusCode::CONTINUE); + body.send_data(Bytes::from_static(b"hello"), true).unwrap(); + assert_eq!(response.await.unwrap().status(), StatusCode::OK); +} + +#[tokio::test] +async fn h2_expect_continue_waits_for_body_poll() { + let svc = service_fn(|req: Request| async move { + assert_eq!(req.headers()["expect"], "100-continue"); + // The body is never polled. + drop(req); + // Not responding right away gives the server a chance to send a 100 Continue. + tokio::task::yield_now().await; + Ok::<_, hyper::Error>( + Response::builder() + .status(StatusCode::EXPECTATION_FAILED) + .body(Empty::::new()) + .unwrap(), + ) + }); + let mut h2 = h2_expect_continue_server(svc).await; + + let (mut response, _body) = h2_send_expect_request(&mut h2, false); + h2_assert_no_informational(&mut response).await; + assert_eq!( + response.await.unwrap().status(), + StatusCode::EXPECTATION_FAILED + ); +} + +#[tokio::test] +async fn h2_expect_continue_but_no_body_is_ignored() { + let svc = service_fn(|req: Request| async move { + let body = req.into_body().collect().await?.to_bytes(); + assert!(body.is_empty()); + // Not responding right away gives the server a chance to send a 100 Continue. + tokio::task::yield_now().await; + Ok::<_, hyper::Error>(Response::new(Empty::::new())) + }); + let mut h2 = h2_expect_continue_server(svc).await; + + let (mut response, _body) = h2_send_expect_request(&mut h2, true); + h2_assert_no_informational(&mut response).await; + assert_eq!(response.await.unwrap().status(), StatusCode::OK); +} + +#[tokio::test] +async fn h2_expect_continue_but_body_already_sent_is_ignored() { + let svc = service_fn(|req: Request| async move { + let body = req.into_body().collect().await?.to_bytes(); + assert_eq!(&body[..], b"hello"); + // Not responding right away gives the server a chance to send a 100 Continue. + tokio::task::yield_now().await; + Ok::<_, hyper::Error>(Response::new(Empty::::new())) + }); + let mut h2 = h2_expect_continue_server(svc).await; + + let (mut response, mut body) = h2_send_expect_request(&mut h2, false); + // The client doesn't wait for the server to ask for the body. + body.send_data(Bytes::from_static(b"hello"), true).unwrap(); + h2_assert_no_informational(&mut response).await; + assert_eq!(response.await.unwrap().status(), StatusCode::OK); +} + #[test] fn pipeline_disabled() { let server = serve();