From 0513b56fafbaad4e5d1875e262eee042d64d6a8a Mon Sep 17 00:00:00 2001 From: David Pedersen Date: Sun, 30 May 2021 01:11:18 +0200 Subject: [PATCH] better readiness handling and less boxing --- Cargo.toml | 7 ++-- src/lib.rs | 118 ++++++++++++++++++++++++++++++++++++++--------------- 2 files changed, 90 insertions(+), 35 deletions(-) diff --git a/Cargo.toml b/Cargo.toml index eafc8c1b..d0333c54 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -7,15 +7,16 @@ edition = "2018" [dependencies] async-trait = "0.1" bytes = "1.0" +futures-util = "0.3" http = "0.2" http-body = "0.4" hyper = "0.14" +pin-project = "1.0" serde = "1.0" -serde_urlencoded = "0.7" serde_json = "1.0" -futures-util = "0.3" -tower = { version = "0.4", features = ["util"] } +serde_urlencoded = "0.7" thiserror = "1.0" +tower = { version = "0.4", features = ["util"] } [dev-dependencies] tokio = { version = "1.6.1", features = ["macros", "rt"] } diff --git a/src/lib.rs b/src/lib.rs index 3538c3bc..c4ee9919 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -31,10 +31,12 @@ use async_trait::async_trait; use bytes::Bytes; use futures_util::{future, ready}; use http::{Method, Request, Response, StatusCode}; +use pin_project::pin_project; use serde::{de::DeserializeOwned, Deserialize, Serialize}; use std::{ future::Future, marker::PhantomData, + pin::Pin, task::{Context, Poll}, }; use tower::{Service, ServiceExt}; @@ -155,21 +157,20 @@ pub enum Error { DeserializeQueryString(#[from] serde_urlencoded::de::Error), } +// TODO(david): make this trait sealed #[async_trait] pub trait Handler { async fn call(self, req: Request) -> Result, Error>; } #[async_trait] -#[allow(non_snake_case)] impl Handler<()> for F where F: Fn(Request) -> Fut + Send + Sync, Fut: Future, Error>> + Send, { async fn call(self, req: Request) -> Result, Error> { - let res = self(req).await?; - Ok(res) + self(req).await } } @@ -253,18 +254,35 @@ where } } -#[async_trait] pub trait FromRequest: Sized { - async fn from_request(req: &mut Request) -> Result; + type Future: Future> + Send; + + fn from_request(req: &mut Request) -> Self::Future; } -#[async_trait] impl FromRequest for Option where T: FromRequest, { - async fn from_request(req: &mut Request) -> Result { - Ok(T::from_request(req).await.ok()) + type Future = OptionFromRequestFuture; + + fn from_request(req: &mut Request) -> Self::Future { + OptionFromRequestFuture(T::from_request(req)) + } +} + +#[pin_project] +pub struct OptionFromRequestFuture(#[pin] F); + +impl Future for OptionFromRequestFuture +where + F: Future>, +{ + type Output = Result, Error>; + + fn poll(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll { + let value = ready!(self.project().0.poll(cx)); + Poll::Ready(Ok(value.ok())) } } @@ -277,15 +295,20 @@ impl Query { } } -#[async_trait] impl FromRequest for Query where - T: DeserializeOwned, + T: DeserializeOwned + Send, { - async fn from_request(req: &mut Request) -> Result { - let query = req.uri().query().ok_or(Error::QueryStringMissing)?; - let value = serde_urlencoded::from_str(query)?; - Ok(Query(value)) + type Future = future::Ready>; + + fn from_request(req: &mut Request) -> Self::Future { + let result = (|| { + let query = req.uri().query().ok_or(Error::QueryStringMissing)?; + let value = serde_urlencoded::from_str(query)?; + Ok(Query(value)) + })(); + + future::ready(result) } } @@ -298,21 +321,24 @@ impl Json { } } -#[async_trait] impl FromRequest for Json where T: DeserializeOwned, { - async fn from_request(req: &mut Request) -> Result { + type Future = future::BoxFuture<'static, Result>; + + fn from_request(req: &mut Request) -> Self::Future { // TODO(david): require the body to have `content-type: application/json` let body = std::mem::take(req.body_mut()); - let bytes = hyper::body::to_bytes(body) - .await - .map_err(Error::ConsumeBody)?; - let value = serde_json::from_slice(&bytes).map_err(Error::DeserializeRequestBody)?; - Ok(Json(value)) + Box::pin(async move { + let bytes = hyper::body::to_bytes(body) + .await + .map_err(Error::ConsumeBody)?; + let value = serde_json::from_slice(&bytes).map_err(Error::DeserializeRequestBody)?; + Ok(Json(value)) + }) } } @@ -331,11 +357,10 @@ impl Service for EmptyRouter { fn call(&mut self, _req: R) -> Self::Future { let mut res = Response::new(Body::empty()); *res.status_mut() = StatusCode::NOT_FOUND; - future::ready(Ok(res)) + future::ok(res) } } -#[derive(Clone)] pub struct Route { handler: H, route_spec: RouteSpec, @@ -344,6 +369,23 @@ pub struct Route { fallback_ready: bool, } +impl Clone for Route +where + H: Clone, + F: Clone, +{ + fn clone(&self) -> Self { + Self { + handler: self.handler.clone(), + fallback: self.fallback.clone(), + route_spec: self.route_spec.clone(), + // important to reset readiness when cloning + handler_ready: false, + fallback_ready: false, + } + } +} + #[derive(Clone)] struct RouteSpec { method: Method, @@ -367,24 +409,36 @@ where type Future = future::Either; fn poll_ready(&mut self, cx: &mut Context<'_>) -> Poll> { - if !self.handler_ready { - ready!(self.handler.poll_ready(cx))?; - self.handler_ready = true; - } + loop { + if !self.handler_ready { + ready!(self.handler.poll_ready(cx))?; + self.handler_ready = true; + } - if !self.fallback_ready { - ready!(self.fallback.poll_ready(cx))?; - self.fallback_ready = true; - } + if !self.fallback_ready { + ready!(self.fallback.poll_ready(cx))?; + self.fallback_ready = true; + } - Poll::Ready(Ok(())) + if self.handler_ready && self.fallback_ready { + return Poll::Ready(Ok(())); + } + } } fn call(&mut self, req: Request) -> Self::Future { if self.route_spec.matches(&req) { + assert!( + self.handler_ready, + "handler not ready. Did you forget to call `poll_ready`?" + ); self.handler_ready = false; future::Either::Left(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)) }