Skip to content
Open
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
44 changes: 41 additions & 3 deletions src/body/incoming.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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")
Expand Down Expand Up @@ -54,10 +56,18 @@ enum Kind {
ping: ping::Recorder,
recv: h2::RecvStream,
},
#[cfg(all(feature = "http2", feature = "server"))]
ExpectContinue(Box<ExpectContinue>),
#[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.
Expand Down Expand Up @@ -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(_)) {
Expand Down Expand Up @@ -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),
}
Expand All @@ -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,
}
Expand All @@ -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(),
}
Expand Down
19 changes: 19 additions & 0 deletions src/headers.rs
Original file line number Diff line number Diff line change
Expand Up @@ -67,6 +67,25 @@ pub(super) fn content_length_parse(value: &HeaderValue) -> Option<u64> {
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<u64> {
content_length_parse_all_values(headers.get_all(CONTENT_LENGTH).into_iter())
Expand Down
5 changes: 1 addition & 4 deletions src/proto/h1/role.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
49 changes: 40 additions & 9 deletions src/proto/h2/server.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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};
Expand Down Expand Up @@ -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");
Expand All @@ -308,6 +312,7 @@ where
let fut = H2Stream::new(
service.call(req),
connect_parts,
expect_continue,
respond,
self.date_header,
exec.clone(),
Expand Down Expand Up @@ -382,6 +387,7 @@ pin_project! {
#[pin]
fut: F,
connect_parts: Option<ConnectParts>,
expect_continue: Option<oneshot::Receiver<()>>,
},
Body {
#[pin]
Expand All @@ -403,13 +409,18 @@ where
fn new(
fut: F,
connect_parts: Option<ConnectParts>,
expect_continue: Option<oneshot::Receiver<()>>,
respond: SendResponse<SendBuf<B::Data>>,
date_header: bool,
exec: E,
) -> H2Stream<F, B, E> {
H2Stream {
reply: respond,
state: H2StreamState::Service { fut, connect_parts },
state: H2StreamState::Service {
fut,
connect_parts,
expect_continue,
},
date_header,
exec,
}
Expand Down Expand Up @@ -445,6 +456,7 @@ where
H2StreamStateProj::Service {
fut: h,
connect_parts,
expect_continue,
} => {
let res = match h.poll(cx) {
Poll::Ready(Ok(r)) => r,
Expand All @@ -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)) => {
Expand Down
126 changes: 126 additions & 0 deletions tests/server.rs
Original file line number Diff line number Diff line change
Expand Up @@ -1068,6 +1068,132 @@ async fn expect_continue_waits_for_body_poll() {
child.join().expect("client thread");
}

async fn h2_expect_continue_server<S>(svc: S) -> SendRequest<Bytes>
where
S: hyper::service::HttpService<IncomingBody, ResBody = Empty<Bytes>> + Send + 'static,
S::Future: Send + 'static,
S::Error: Into<Box<dyn std::error::Error + Send + Sync>>,
{
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<Bytes>,
end_of_stream: bool,
) -> (h2::client::ResponseFuture, SendStream<Bytes>) {
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<IncomingBody>| async move {
let body = req.into_body().collect().await?.to_bytes();
assert_eq!(&body[..], b"hello");
Ok::<_, hyper::Error>(Response::new(Empty::<Bytes>::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<IncomingBody>| 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::<Bytes>::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<IncomingBody>| 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::<Bytes>::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<IncomingBody>| 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::<Bytes>::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();
Expand Down