diff --git a/interop/Cargo.toml b/interop/Cargo.toml index 7f7ffc4a2..3563efe45 100644 --- a/interop/Cargo.toml +++ b/interop/Cargo.toml @@ -11,6 +11,7 @@ anyhow = "1.0.22" bytes = "0.5.2" futures = "0.3.1" http = "0.2" +http-body = "0.3" hyper = "0.13" hyper-rustls = "0.20" lazy_static = "1" diff --git a/interop/src/server.rs b/interop/src/server.rs index c2c598d59..71fc791af 100644 --- a/interop/src/server.rs +++ b/interop/src/server.rs @@ -11,9 +11,10 @@ use std::{ use anyhow::{anyhow, bail, Context as _, Result}; use bytes::Bytes; -use futures::{ready, AsyncReadExt, Future, StreamExt, TryFutureExt}; +use futures::{ready, AsyncReadExt, Future, StreamExt, TryFutureExt, future}; use http::{Response, StatusCode}; -use hyper::service::{make_service_fn, service_fn}; +use http_body::Body as _; +use hyper::{body::HttpBody, service::{make_service_fn, service_fn}}; use structopt::{self, StructOpt}; use tokio::net::{TcpListener, TcpStream}; use tokio_rustls::{server::TlsStream, TlsAcceptor}; @@ -137,17 +138,13 @@ async fn h3_handle_connection(connecting: quinn::Connecting) -> Result<()> { } async fn h3_handle_request(recv_request: RecvRequest) -> Result<()> { - let (request, mut recv_body, sender) = recv_request.await?; + let (mut request, sender) = recv_request.await?; println!("received request: {:?}", request); - let mut body = Vec::with_capacity(1024); - recv_body - .read_to_end(&mut body) - .await - .map_err(|e| anyhow!("failed to send response headers: {:?}", e))?; - + let body = request.body_mut().read_to_end().await?; println!("received body: {}", String::from_utf8_lossy(&body)); - if let Some(trailers) = recv_body.trailers().await { + + if let Some(trailers) = request.body_mut().trailers().await? { println!("received trailers: {:?}", trailers); } diff --git a/quinn-h3/examples/h3_client.rs b/quinn-h3/examples/h3_client.rs index be69427b4..c625d85df 100644 --- a/quinn-h3/examples/h3_client.rs +++ b/quinn-h3/examples/h3_client.rs @@ -2,7 +2,6 @@ use std::{fs, io, net::ToSocketAddrs, path::PathBuf}; use structopt::{self, StructOpt}; use anyhow::Result; -use futures::AsyncReadExt; use http::{Request, Uri}; use tracing::{error, info}; use tracing_subscriber::filter::LevelFilter; @@ -56,18 +55,16 @@ async fn main() -> Result<()> { let (send_data, recv_response) = conn.send_request(request); send_data.await?; // Wait for the response - let (response, mut recv_body) = recv_response.await?; + let mut response = recv_response.await?; info!("received response: {:?}", response); // Stream the response body into a vec - let mut body = Vec::with_capacity(1024); - recv_body.read_to_end(&mut body).await?; - + let body = response.body_mut().read_to_end().await?; info!("received body: {}", String::from_utf8_lossy(&body)); // Get the trailers if any - if let Some(trailers) = recv_body.trailers().await { + if let Some(trailers) = response.body_mut().trailers().await? { info!("received trailers: {:?}", trailers); } diff --git a/quinn-h3/examples/h3_server.rs b/quinn-h3/examples/h3_server.rs index 94ff5fae1..2c4d1c3c2 100644 --- a/quinn-h3/examples/h3_server.rs +++ b/quinn-h3/examples/h3_server.rs @@ -81,7 +81,7 @@ async fn main() -> Result<()> { async fn handle_request(recv_request: RecvRequest) -> Result<()> { // Receive the request's headers - let (request, _body_reader, sender) = recv_request.await?; + let (request, sender) = recv_request.await?; info!("received request: {:?}", request); let response = Response::builder() diff --git a/quinn-h3/src/body.rs b/quinn-h3/src/body.rs index 5c2a1e2d7..89b3cbe6a 100644 --- a/quinn-h3/src/body.rs +++ b/quinn-h3/src/body.rs @@ -1,22 +1,24 @@ use std::{ - cmp, + cmp, fmt, + future::Future, io::{self, ErrorKind}, mem, pin::Pin, task::{Context, Poll}, }; -use bytes::{Bytes, BytesMut}; +use bytes::{Buf, Bytes, BytesMut}; use futures::{ + future, io::{AsyncRead, AsyncWrite}, ready, stream::Stream, FutureExt, }; use http::HeaderMap; +use http_body::Body as HttpBody; use quinn::SendStream; use quinn_proto::StreamId; -use std::future::Future; use crate::{ connection::ConnectionRef, @@ -554,3 +556,100 @@ impl HttpBody for SimpleBody { Poll::Ready(Ok(None)) } } + +pub struct RecvBody { + conn: ConnectionRef, + stream_id: StreamId, + recv: FrameStream, + trailers: Option, +} + +impl RecvBody { + pub(crate) fn new(conn: ConnectionRef, stream_id: StreamId, recv: FrameStream) -> Self { + Self { + conn, + stream_id, + recv, + trailers: None, + } + } + + pub async fn read_to_end(&mut self) -> Result { + let mut body = BytesMut::with_capacity(10_240); + + let mut me = self; + let res: Result<(), Error> = future::poll_fn(|cx| { + while let Some(d) = ready!(Pin::new(&mut me).poll_data(cx)) { + body.extend(d?.bytes()); + } + Poll::Ready(Ok(())) + }) + .await; + res?; + + Ok(body.freeze()) + } + + pub async fn trailers(&mut self) -> Result, Error> { + let mut me = self; + Ok(future::poll_fn(|cx| Pin::new(&mut me).poll_trailers(cx)).await?) + } +} + +impl HttpBody for RecvBody { + type Data = bytes::Bytes; + type Error = Error; + + fn poll_data( + mut self: Pin<&mut Self>, + cx: &mut Context<'_>, + ) -> Poll>> { + loop { + return match ready!(Pin::new(&mut self.recv).poll_next(cx)) { + None => Poll::Ready(None), + Some(Ok(HttpFrame::Reserved)) => continue, + Some(Ok(HttpFrame::Data(d))) => Poll::Ready(Some(Ok(d.payload))), + Some(Ok(HttpFrame::Headers(t))) => { + self.trailers = Some(t); + Poll::Ready(None) + } + Some(Err(e)) => { + self.recv.reset(e.code()); + Poll::Ready(Some(Err(e.into()))) + } + Some(Ok(f)) => { + self.recv.reset(ErrorCode::FRAME_UNEXPECTED); + Poll::Ready(Some(Err(Error::Peer(format!( + "Invalid frame type in body: {:?}", + f + ))))) + } + }; + } + } + + fn poll_trailers( + mut self: Pin<&mut Self>, + cx: &mut Context<'_>, + ) -> Poll, Self::Error>> { + if self.trailers.is_none() { + return Poll::Ready(Ok(None)); + } + + let header = { + let mut conn = self.conn.h3.lock().unwrap(); + ready!(conn.poll_decode(cx, self.stream_id, self.trailers.as_ref().unwrap()))? + }; + self.trailers = None; + + Poll::Ready(Ok(Some(header.into_fields()))) + } +} + +impl fmt::Debug for RecvBody { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + f.debug_struct("RecvBody") + .field("stream", &self.stream_id) + .finish() + } +} diff --git a/quinn-h3/src/client.rs b/quinn-h3/src/client.rs index 91efbfa42..055d68b3d 100644 --- a/quinn-h3/src/client.rs +++ b/quinn-h3/src/client.rs @@ -84,7 +84,7 @@ use quinn_proto::{Side, StreamId}; use tracing::trace; use crate::{ - body::BodyReader, + body::RecvBody, connection::{ConnectionDriver, ConnectionRef}, frame::{FrameDecoder, FrameStream}, headers::DecodeHeaders, @@ -681,10 +681,29 @@ impl RecvResponse { .cancel_request(self.stream_id.unwrap()); recv.reset(ErrorCode::REQUEST_CANCELLED); } + + fn build_response( + &self, + header: Header, + recv: FrameStream, + ) -> Result, Error> { + let (status, headers) = header.into_response_parts()?; + let mut response = Response::builder() + .status(status) + .version(http::version::Version::HTTP_3) + .body(RecvBody::new( + self.conn.clone(), + self.stream_id.unwrap(), + recv, + )) + .unwrap(); + *response.headers_mut() = headers; + Ok(response) + } } impl Future for RecvResponse { - type Output = Result<(Response<()>, BodyReader), crate::Error>; + type Output = Result, crate::Error>; fn poll(mut self: Pin<&mut Self>, cx: &mut Context) -> Poll { loop { @@ -737,39 +756,15 @@ impl Future for RecvResponse { } RecvResponseState::Decoding(ref mut decode) => { let headers = ready!(Pin::new(decode).poll(cx))?; - let response = build_response(headers); - match response { - Err(e) => return Poll::Ready(Err(e)), - Ok(r) => { - self.state = RecvResponseState::Finished; - return Poll::Ready(Ok(( - r, - BodyReader::new( - self.recv.take().unwrap(), - self.conn.clone(), - self.stream_id.unwrap(), - true, - ), - ))); - } - } + let recv = self.recv.take().unwrap(); + let response = self.build_response(headers, recv)?; + return Poll::Ready(Ok(response)); } } } } } -fn build_response(header: Header) -> Result, Error> { - let (status, headers) = header.into_response_parts()?; - let mut response = Response::builder() - .status(status) - .version(http::version::Version::HTTP_3) - .body(()) - .unwrap(); - *response.headers_mut() = headers; - Ok(response) -} - #[cfg(test)] impl Connection { pub(crate) fn inner(&self) -> &ConnectionRef { diff --git a/quinn-h3/src/server.rs b/quinn-h3/src/server.rs index 313e51169..3848719d9 100644 --- a/quinn-h3/src/server.rs +++ b/quinn-h3/src/server.rs @@ -115,7 +115,7 @@ use std::{ }; use futures::{ready, Stream}; -use http::{Request, Response}; +use http::{response, Request, Response}; use http_body::Body; use quinn::{ CertificateChain, EndpointBuilder, PrivateKey, RecvStream, SendStream, ZeroRttAccepted, @@ -124,16 +124,12 @@ use quinn_proto::{Side, StreamId}; use rustls::TLSError; use crate::{ - body::BodyReader, + body::RecvBody, connection::{ConnectionDriver, ConnectionRef}, data::SendData, frame::{FrameDecoder, FrameStream}, headers::DecodeHeaders, - proto::{ - frame::HttpFrame, - headers::Header, - ErrorCode, - }, + proto::{frame::HttpFrame, headers::Header, ErrorCode}, streams::Reset, Error, Settings, }; @@ -589,7 +585,7 @@ impl RecvRequest { /// Reject this request with `REQUEST_REJECTED` code. pub fn reject(mut self) { let state = mem::replace(&mut self.state, RecvRequestState::Finished); - if let RecvRequestState::Receiving(recv, mut send) = state { + if let RecvRequestState::Receiving(mut recv, mut send) = state { recv.reset(ErrorCode::REQUEST_REJECTED); send.reset(ErrorCode::REQUEST_REJECTED.into()); } @@ -605,29 +601,39 @@ impl RecvRequest { } } - fn build_request(&self, headers: Header) -> Result, Error> { - let (method, uri, headers) = headers.into_request_parts()?; + fn build_request( + &self, + headers: Header, + recv: FrameStream, + ) -> Result, (Error, FrameStream)> { + let (method, uri, headers) = match headers.into_request_parts() { + Ok(p) => p, + Err(e) => return Err((e.into(), recv)), + }; + if self.is_0rtt && !method.is_idempotent() { + return Err(( + Error::peer(format!( + "Tried an non indempotent method in 0-RTT: {}", + method, + )), + recv, + )); + } + let mut request = Request::builder() .method(method) .uri(uri) .version(http::version::Version::HTTP_3) - .body(()) + .body(RecvBody::new(self.conn.clone(), self.stream_id, recv)) .unwrap(); - if self.is_0rtt && !request.method().is_idempotent() { - return Err(Error::peer(format!( - "Tried an non indempotent method in 0-RTT: {}", - request.method() - ))); - } - *request.headers_mut() = headers; Ok(request) } } impl Future for RecvRequest { - type Output = Result<(Request<()>, BodyReader, Sender), Error>; + type Output = Result<(Request, Sender), Error>; fn poll(mut self: Pin<&mut Self>, cx: &mut Context) -> Poll { loop { @@ -663,21 +669,20 @@ impl Future for RecvRequest { RecvRequestState::Decoding(ref mut decode) => { let header = ready!(Pin::new(decode).poll(cx))?; self.state = RecvRequestState::Finished; - let (mut recv, mut send) = self + let (recv, mut send) = self .streams .take() .ok_or_else(|| Error::internal("Recv request invalid state"))?; - let request = match self.build_request(header) { + let request = match self.build_request(header, recv) { Ok(r) => r, - Err(e) => { + Err((e, mut r)) => { send.reset(ErrorCode::REQUEST_REJECTED.into()); - recv.reset(ErrorCode::REQUEST_REJECTED); + r.reset(ErrorCode::REQUEST_REJECTED); return Poll::Ready(Err(e)); } }; return Poll::Ready(Ok(( request, - BodyReader::new(recv, self.conn.clone(), self.stream_id, false), Sender { send, conn: self.conn.clone(), @@ -776,7 +781,13 @@ impl Sender { B::Error: std::fmt::Debug + Any + Send + Sync, // TODO remove debug { let (response, body) = response.into_parts(); - SendData::new(self.send, self.conn, response, body) + + let response::Parts { + status, headers, .. + } = response; + let header = Header::response(status, headers); + + SendData::new(self.send, self.conn, header, body) } /// Cancel request processing