From a04c98dd425195fd787d57bc15aae92fcbc3047b Mon Sep 17 00:00:00 2001 From: David Pedersen Date: Sun, 30 May 2021 01:56:52 +0200 Subject: [PATCH] Support any type of response body --- Cargo.toml | 3 +++ src/lib.rs | 74 +++++++++++++++++++++++++++++++++++++++++++++++------- 2 files changed, 68 insertions(+), 9 deletions(-) diff --git a/Cargo.toml b/Cargo.toml index d0333c54..308762c2 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -21,3 +21,6 @@ tower = { version = "0.4", features = ["util"] } [dev-dependencies] tokio = { version = "1.6.1", features = ["macros", "rt"] } serde = { version = "1.0", features = ["derive"] } +tower = { version = "0.4", features = ["util", "make"] } +tower-http = { version = "0.1", features = ["trace"] } +hyper = { version = "0.14", features = ["full"] } diff --git a/src/lib.rs b/src/lib.rs index 6ecb4678..7718ed46 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -31,6 +31,7 @@ use async_trait::async_trait; use bytes::Bytes; use futures_util::{future, ready}; use http::{Method, Request, Response, StatusCode}; +use http_body::{combinators::BoxBody, Body as _}; use pin_project::pin_project; use serde::{de::DeserializeOwned, Deserialize, Serialize}; use std::{ @@ -39,7 +40,7 @@ use std::{ pin::Pin, task::{Context, Poll}, }; -use tower::{Service, ServiceExt}; +use tower::{BoxError, Service, ServiceExt}; pub use hyper::body::Body; @@ -155,6 +156,9 @@ pub enum Error { #[error("failed to deserialize query string")] DeserializeQueryString(#[from] serde_urlencoded::de::Error), + + #[error("failed generating the response body")] + ResponseBody(#[source] BoxError), } // TODO(david): make this trait sealed @@ -408,14 +412,19 @@ impl RouteSpec { } } -impl Service> for Route +impl Service> for Route where - H: Service, Response = Response, Error = Error>, - F: Service, Response = Response, Error = Error>, + H: Service, Response = Response, Error = Error>, + F: Service, Response = Response, Error = Error>, + HB: http_body::Body + Send + Sync + 'static, + HB::Error: Into, + FB: http_body::Body + Send + Sync + 'static, + FB::Error: Into, { - type Response = Response; + type Response = Response>; type Error = Error; - type Future = future::Either; + // type Future = future::BoxFuture<'static, Result>; + type Future = future::Either, BoxResponseBody>; fn poll_ready(&mut self, cx: &mut Context<'_>) -> Poll> { loop { @@ -442,18 +451,38 @@ where "handler not ready. Did you forget to call `poll_ready`?" ); self.handler_ready = false; - future::Either::Left(self.handler.call(req)) + future::Either::Left(BoxResponseBody(self.handler.call(req))) } else { assert!( self.fallback_ready, "fallback not ready. Did you forget to call `poll_ready`?" ); self.fallback_ready = false; - future::Either::Right(self.fallback.call(req)) + // TODO(david): this leads to each route creating one box body, probably not great + future::Either::Right(BoxResponseBody(self.fallback.call(req))) } } } +#[pin_project] +pub struct BoxResponseBody(#[pin] F); + +impl Future for BoxResponseBody +where + F: Future, Error>>, + B: http_body::Body + Send + Sync + 'static, + B::Error: Into, +{ + type Output = Result>, Error>; + + fn poll(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll { + let response: Response = ready!(self.project().0.poll(cx))?; + let response = + response.map(|body| body.map_err(|err| Error::ResponseBody(err.into())).boxed()); + Poll::Ready(Ok(response)) + } +} + impl Service for App where R: Service, @@ -496,6 +525,10 @@ where mod tests { #![allow(warnings)] use super::*; + use hyper::Server; + use std::{fmt, net::SocketAddr}; + use tower::{make::Shared, ServiceBuilder}; + use tower_http::trace::TraceLayer; #[tokio::test] async fn basic() { @@ -517,7 +550,30 @@ mod tests { dbg!(&body); } - async fn body_to_string(res: Response) -> String { + #[allow(dead_code)] + // this should just compile + async fn compatible_with_hyper_and_tower_http() { + let app = app() + .at("/") + .get(root) + .at("/users") + .get(users_index) + .post(users_create); + + let app = ServiceBuilder::new() + .layer(TraceLayer::new_for_http()) + .service(app); + + let addr = SocketAddr::from(([127, 0, 0, 1], 3000)); + let server = Server::bind(&addr).serve(Shared::new(app)); + server.await.unwrap(); + } + + async fn body_to_string(res: Response) -> String + where + B: http_body::Body, + B::Error: fmt::Debug, + { let bytes = hyper::body::to_bytes(res.into_body()).await.unwrap(); String::from_utf8(bytes.to_vec()).unwrap() }