diff --git a/interop/src/server.rs b/interop/src/server.rs index 8c233b230..cb9f5b566 100644 --- a/interop/src/server.rs +++ b/interop/src/server.rs @@ -161,7 +161,7 @@ async fn h3_handle_request(recv_request: RecvRequest) -> Result<()> { Ok(()) } -async fn h3_home(sender: quinn_h3::server::Sender) -> Result<()> { +async fn h3_home(mut sender: quinn_h3::server::Sender) -> Result<()> { let response = Response::builder() .status(StatusCode::OK) .header("server", VERSION) @@ -174,7 +174,7 @@ async fn h3_home(sender: quinn_h3::server::Sender) -> Result<()> { Ok(()) } -async fn h3_payload(sender: quinn_h3::server::Sender, len: usize) -> Result<()> { +async fn h3_payload(mut sender: quinn_h3::server::Sender, len: usize) -> Result<()> { if len > 1_000_000_000 { let response = Response::builder() .status(StatusCode::BAD_REQUEST) diff --git a/quinn-h3/benches/request.rs b/quinn-h3/benches/request.rs index be066a800..18cce7966 100644 --- a/quinn-h3/benches/request.rs +++ b/quinn-h3/benches/request.rs @@ -201,7 +201,7 @@ async fn request_server( select! { _ = &mut stop_recv => break, Some(recv_req) = incoming_req.next() => { - let (_, sender) = recv_req.await.expect("recv_req"); + let (_, mut sender) = recv_req.await.expect("recv_req"); sender.send_response(make_response()) .await .expect("send_response"); diff --git a/quinn-h3/benches/throughput.rs b/quinn-h3/benches/throughput.rs index 29b4dfdba..f8971f0f6 100644 --- a/quinn-h3/benches/throughput.rs +++ b/quinn-h3/benches/throughput.rs @@ -93,7 +93,7 @@ async fn download_server( select! { _ = &mut stop_recv => break, Some(recv_req) = incoming_req.next() => { - let (request, sender) = recv_req.await.expect("recv_req"); + let (request, mut sender) = recv_req.await.expect("recv_req"); let frame_size = request .headers() .get("frame_size") @@ -175,7 +175,7 @@ async fn upload_server( select! { _ = &mut stop_recv => break, Some(recv_req) = incoming_req.next() => { - let (mut req, sender) = recv_req.await.expect("recv_req"); + let (mut req, mut sender) = recv_req.await.expect("recv_req"); while let Some(Ok(_)) = req.body_mut().data().await {} let response = Response::builder() .status(StatusCode::OK) diff --git a/quinn-h3/examples/h3_server.rs b/quinn-h3/examples/h3_server.rs index 27df9ddc7..b488fa54a 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, sender) = recv_request.await?; + let (request, mut sender) = recv_request.await?; info!("received request: {:?}", request); let response = Response::builder() diff --git a/quinn-h3/src/server.rs b/quinn-h3/src/server.rs index e26df6990..db3fe9478 100644 --- a/quinn-h3/src/server.rs +++ b/quinn-h3/src/server.rs @@ -636,8 +636,8 @@ impl Future for RecvRequest { let (header, body) = ready!(self.recv.as_mut().unwrap().poll_unpin(cx))?; let request = self.build_request(header, body)?; let sender = Sender { - send: self.send.take().unwrap(), - conn: self.conn.clone(), + send: self.send.take(), + conn: Some(self.conn.clone()), }; Poll::Ready(Ok((request, sender))) } @@ -659,8 +659,8 @@ impl Future for RecvRequest { /// [`BodyReader`]: ../body/struct.BodyReader.html /// [`cancel()`]: #method.cancel pub struct Sender { - send: SendStream, - conn: ConnectionRef, + send: Option, + conn: Option, } impl Sender { @@ -697,7 +697,7 @@ impl Sender { /// [`Body`]: ../enum.Body.html /// [`BodyWriter`]: ../struct.BodyWriter.html /// [`AsyncWrite`]: https://docs.rs/futures/*/futures/io/trait.AsyncWrite.html - pub fn send_response(self, response: Response) -> SendData + pub fn send_response(&mut self, response: Response) -> SendData where B: HttpBody + Send + 'static, B::Data: Send, @@ -710,7 +710,8 @@ impl Sender { } = response; let header = Header::response(status, headers); - SendData::new(self.send, self.conn, header, body, true) + let (send, conn) = (self.send.take().unwrap(), self.conn.take().unwrap()); + SendData::new(send, conn, header, body, true) } /// Cancel request processing @@ -720,7 +721,10 @@ impl Sender { /// /// Cancelling a request means that some request data have been processed by the application, which /// decided to abandon the response. - pub fn cancel(mut self) { - self.send.reset(ErrorCode::REQUEST_CANCELLED.into()); + pub fn cancel(&mut self) { + self.send + .as_mut() + .unwrap() + .reset(ErrorCode::REQUEST_CANCELLED.into()); } } diff --git a/quinn-h3/src/tests/mod.rs b/quinn-h3/src/tests/mod.rs index 3f2480406..2ca948fe1 100644 --- a/quinn-h3/src/tests/mod.rs +++ b/quinn-h3/src/tests/mod.rs @@ -11,7 +11,7 @@ use helpers::{get, post, timeout_join, Helper}; async fn serve_one(mut incoming: IncomingConnection) -> Result<(), crate::Error> { let mut incoming_req = incoming.next().await.expect("no accept").await?; while let Some(recv_req) = incoming_req.next().await { - let (_, sender) = recv_req.await?; + let (_, mut sender) = recv_req.await?; sender .send_response( Response::builder() @@ -58,7 +58,7 @@ async fn serve_one_request_client_body(mut incoming: IncomingConnection) -> Stri .await .expect("accept"); let recv_req = incoming_req.next().await.expect("wait request"); - let (mut req, sender) = recv_req.await.expect("recv_req"); + let (mut req, mut sender) = recv_req.await.expect("recv_req"); let body = req.body_mut().read_to_end().await.expect("read body"); sender @@ -120,7 +120,7 @@ async fn client_cancel_response() { .expect("accept"); let recv_req = incoming_req.next().await.expect("wait request"); delay_for(Duration::from_millis(25)).await; - let (_, sender) = recv_req.await.expect("recv_req"); + let (_, mut sender) = recv_req.await.expect("recv_req"); sender .send_response( Response::builder() @@ -158,7 +158,7 @@ async fn go_away() { .expect("accept"); let recv_req = incoming_req.next().await.expect("wait request"); incoming_req.go_away(); - let (_, sender) = recv_req.await.expect("recv_req"); + let (_, mut sender) = recv_req.await.expect("recv_req"); sender .send_response( Response::builder() @@ -196,7 +196,7 @@ async fn serve_n_0rtt(mut incoming: IncomingConnection, n: usize) -> Result<(), .map_err(|_| ()) .expect("0rtt failed"); while let Some(recv_req) = incoming_req.next().await { - let (_, sender) = recv_req.await?; + let (_, mut sender) = recv_req.await?; sender .send_response( Response::builder()