From 9fd975afa3358234016cacbc2f127340a170ca44 Mon Sep 17 00:00:00 2001 From: stammw Date: Sun, 24 Nov 2019 07:43:50 +0100 Subject: [PATCH] H3: refactor server to use only BodyReader/Writer --- quinn-h3/examples/simple_server.rs | 32 +++--- quinn-h3/src/body.rs | 4 +- quinn-h3/src/server.rs | 156 ++++------------------------- 3 files changed, 33 insertions(+), 159 deletions(-) diff --git a/quinn-h3/examples/simple_server.rs b/quinn-h3/examples/simple_server.rs index c8dca7d59..23f4a414a 100644 --- a/quinn-h3/examples/simple_server.rs +++ b/quinn-h3/examples/simple_server.rs @@ -1,16 +1,15 @@ use std::{net::SocketAddr, path::PathBuf, sync::Arc}; use anyhow::{anyhow, Result}; -use futures::StreamExt; -use http::{Request, Response, StatusCode}; +use futures::{AsyncReadExt, StreamExt}; +use http::{Response, StatusCode}; use structopt::{self, StructOpt}; use quinn::ConnectionDriver as QuicDriver; use quinn_h3::{ self, - body::RecvBody, connection::ConnectionDriver, - server::{Builder as ServerBuilder, IncomingRequest, Sender}, + server::{Builder as ServerBuilder, IncomingRequest, RecvRequest}, }; mod shared; @@ -98,9 +97,8 @@ async fn handle_connection(conn: (QuicDriver, ConnectionDriver, IncomingRequest) tokio::spawn(async move { while let Some(request) = incoming.next().await { - let (req, send) = request.await.expect("recv request failed"); tokio::spawn(async move { - if let Err(e) = handle_request(req, send).await { + if let Err(e) = handle_request(request).await { eprintln!("request error: {}", e) } }); @@ -114,23 +112,18 @@ async fn handle_connection(conn: (QuicDriver, ConnectionDriver, IncomingRequest) Ok(()) } -const INITIAL_CAPACITY: usize = 256; -const MAX_LEN: usize = 256; - -async fn handle_request(request: Request, sender: Sender) -> Result<()> { +async fn handle_request(recv_request: RecvRequest) -> Result<()> { + let (request, mut recv_body, sender) = recv_request.await.expect("recv request failed"); println!("received request: {:?}", request); - let (_, body) = request.into_parts(); - - let (content, trailers) = body - .read_to_end(INITIAL_CAPACITY, MAX_LEN) + 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))?; - if let Some(content) = content { - println!("received body: {}", String::from_utf8_lossy(&content)); - } - if let Some(trailers) = trailers { + println!("received body: {}", String::from_utf8_lossy(&body)); + if let Some(trailers) = recv_body.trailers().await { println!("received trailers: {:?}", trailers); } @@ -141,8 +134,7 @@ async fn handle_request(request: Request, sender: Sender) -> Result<() .expect("failed to build response"); sender - .response(response) - .send() + .send_response(response) .await .map_err(|e| anyhow!("failed to send response: {:?}", e))?; diff --git a/quinn-h3/src/body.rs b/quinn-h3/src/body.rs index 85e7a702e..03f422063 100644 --- a/quinn-h3/src/body.rs +++ b/quinn-h3/src/body.rs @@ -6,7 +6,7 @@ use std::{ task::{Context, Poll}, }; -use bytes::{Bytes, BytesMut}; +use bytes::Bytes; use futures::{ io::{AsyncRead, AsyncWrite}, ready, @@ -27,7 +27,7 @@ use crate::{ ErrorCode, }, streams::Reset, - try_take, Error, + Error, }; pub enum Body { diff --git a/quinn-h3/src/server.rs b/quinn-h3/src/server.rs index 394066ad0..418743f17 100644 --- a/quinn-h3/src/server.rs +++ b/quinn-h3/src/server.rs @@ -12,7 +12,7 @@ use quinn::{EndpointBuilder, EndpointDriver, EndpointError, RecvStream, SendStre use quinn_proto::{Side, StreamId}; use crate::{ - body::{Body, BodyWriter, RecvBody}, + body::{Body, BodyReader, BodyWriter}, connection::{ConnectionDriver, ConnectionRef}, frame::{FrameDecoder, FrameStream, WriteFrame}, headers::{DecodeHeaders, SendHeaders}, @@ -153,22 +153,13 @@ impl RecvRequest { } } - fn build_request( - &self, - headers: Header, - recv: FrameStream, - ) -> Result, Error> { + fn build_request(&self, headers: Header) -> Result, Error> { let (method, uri, headers) = headers.into_request_parts()?; let mut request = Request::builder() .method(method) .uri(uri) .version(http::version::Version::HTTP_3) - .body(RecvBody::new( - recv, - self.conn.clone(), - self.stream_id, - false, - )) + .body(()) .unwrap(); *request.headers_mut() = headers; Ok(request) @@ -184,7 +175,7 @@ impl RecvRequest { } impl Future for RecvRequest { - type Output = Result<(Request, Sender), Error>; + type Output = Result<(Request<()>, BodyReader, Sender), Error>; fn poll(mut self: Pin<&mut Self>, cx: &mut Context) -> Poll { loop { @@ -221,7 +212,8 @@ impl Future for RecvRequest { self.state = RecvRequestState::Finished; let (recv, send) = try_take(&mut self.streams, "Recv request invalid state")?; return Poll::Ready(Ok(( - self.build_request(header, recv)?, + self.build_request(header)?, + BodyReader::new(recv, self.conn.clone(), self.stream_id, false), Sender { send, stream_id: self.stream_id, @@ -244,62 +236,24 @@ pub struct Sender { } impl Sender { - pub fn response(self, response: Response) -> ResponseBuilder { - ResponseBuilder { - response, - sender: self, - trailers: None, - } - } - - pub fn cancel(mut self) { - self.send.reset(ErrorCode::REQUEST_REJECTED.into()); - } -} - -pub struct ResponseBuilder { - sender: Sender, - response: Response, - trailers: Option, -} - -impl ResponseBuilder -where - T: Into, -{ - pub fn trailers(mut self, trailers: HeaderMap) -> Self { - self.trailers = Some(trailers); - self - } - - pub async fn send(self) -> Result<(), Error> { - let Sender { - send, - stream_id, - conn, - } = self.sender; - SendResponse::new(self.response, self.trailers, send, stream_id, conn)?.await?; - Ok(()) - } - - pub async fn stream(self) -> Result { - let Sender { - send, - stream_id, - conn, - } = self.sender; - + pub async fn send_response>( + self, + response: Response, + ) -> Result { let ( response::Parts { status, headers, .. }, body, - ) = self.response.into_parts(); + ) = response.into_parts(); - let trailers = self.trailers; - - let send = - SendHeaders::new(Header::response(status, headers), &conn, send, stream_id)?.await?; + let send = SendHeaders::new( + Header::response(status, headers), + &self.conn, + self.send, + self.stream_id, + )? + .await?; let send = match body.into() { Body::None => send, Body::Buf(payload) => WriteFrame::new(send, DataFrame { payload }).await?, @@ -308,78 +262,6 @@ where } pub fn cancel(mut self) { - let state = mem::replace(&mut self.state, SendResponseState::Finished); - match state { - SendResponseState::SendingHeader(send) => { - send.reset(ErrorCode::REQUEST_CANCELLED); - } - SendResponseState::SendingBody(write) => { - write.reset(ErrorCode::REQUEST_CANCELLED); - } - SendResponseState::SendingTrailers(send) => { - send.reset(ErrorCode::REQUEST_CANCELLED); - } - _ => (), - } - } -} - -impl Future for SendResponse { - type Output = Result<(), Error>; - - fn poll(mut self: Pin<&mut Self>, cx: &mut Context) -> Poll { - loop { - match self.state { - SendResponseState::Finished => panic!("polled after finished"), - SendResponseState::SendingTrailers(ref mut write) => { - ready!(Pin::new(write).poll(cx))?; // drop send - self.state = SendResponseState::Finished; - return Poll::Ready(Ok(())); - } - SendResponseState::SendingHeader(ref mut write) => { - let send = ready!(Pin::new(write).poll(cx))?; - match self.body.take() { - Some(Body::Buf(payload)) => { - self.state = SendResponseState::SendingBody(WriteFrame::new( - send, - DataFrame { payload }, - )); - } - _ => { - self.state = SendResponseState::Finished; - return Poll::Ready(Ok(())); - } - }; - } - SendResponseState::SendingBody(ref mut body) => { - let send = ready!(Pin::new(body).poll(cx))?; - match self.trailer.take() { - None => { - self.state = SendResponseState::Finished; - return Poll::Ready(Ok(())); - } - Some(trailer) => { - self.state = SendResponseState::SendingTrailers(SendHeaders::new( - trailer, - &self.conn, - send, - self.stream_id, - )?); - } - }; - } - } - } - } -} - -impl Drop for SendResponse { - fn drop(&mut self) { - self.conn - .h3 - .lock() - .unwrap() - .inner - .request_finished(self.stream_id); + self.send.reset(ErrorCode::REQUEST_REJECTED.into()); } }