diff --git a/tonic-interop/Cargo.toml b/tonic-interop/Cargo.toml index c9a4c24..9c4c86f 100644 --- a/tonic-interop/Cargo.toml +++ b/tonic-interop/Cargo.toml @@ -22,6 +22,7 @@ http = "0.1" futures-core-preview = "=0.3.0-alpha.18" futures-util-preview = "=0.3.0-alpha.18" async-stream = "0.1.1" +tower = "0.3.0-alpha.1a" console = "0.7" structopt = "0.2" diff --git a/tonic-interop/src/bin/server.rs b/tonic-interop/src/bin/server.rs index 606f9ac..83f228c 100644 --- a/tonic-interop/src/bin/server.rs +++ b/tonic-interop/src/bin/server.rs @@ -1,6 +1,9 @@ use structopt::StructOpt; use tonic::Server; use tonic_interop::server; +// TODO: move GrpcService out of client since it can be used for the +// server too. +use tonic::client::GrpcService; #[derive(StructOpt)] struct Opts { @@ -26,6 +29,16 @@ async fn main() -> std::result::Result<(), Box> { builder.tls(ca, key); } + builder.interceptor_fn(|svc, req| { + println!("INBOUND REQUEST={:?}", req); + let call = svc.call(req); + async move { + let res = call.await?; + println!("OUTBOUND RESPONSE={:?}", res); + Ok(res) + } + }); + builder.serve(addr, test_service).await?; Ok(()) diff --git a/tonic/src/transport/server.rs b/tonic/src/transport/server.rs index 6f50e37..36e3995 100644 --- a/tonic/src/transport/server.rs +++ b/tonic/src/transport/server.rs @@ -1,20 +1,29 @@ use super::{ - service::BoxedIo, + service::{BoxedIo, layer_fn}, tls::{Cert, TlsAcceptor}, }; use crate::BoxBody; use futures_core::Stream; -use futures_util::{ready, try_future::MapOk, TryFutureExt, TryStreamExt}; +use futures_util::{ready, try_future::MapErr, TryFutureExt, TryStreamExt}; use http::{Request, Response}; use hyper::server::{accept::Accept, conn}; use hyper::Body; +use std::sync::Arc; use std::{ net::SocketAddr, pin::Pin, task::{Context, Poll}, + fmt, + future::Future }; +use tower::layer::Layer; use tower_make::MakeService; use tower_service::Service; +use tower::layer::util::Stack; +use tower::util::Either; + +type BoxService = tower::util::BoxService, Response, crate::Error>; +type Interceptor = Arc + Send + Sync + 'static>; #[derive(Debug)] pub struct Server {} @@ -25,14 +34,15 @@ impl Server { } } -#[derive(Debug)] +#[derive(Default)] pub struct Builder { tls: Option<(Vec, Vec)>, + interceptor: Option, } impl Builder { fn new() -> Self { - Self { tls: None } + Default::default() } pub fn tls(&mut self, pem: Vec, key: Vec) -> &mut Self { @@ -43,14 +53,29 @@ impl Builder { // pub fn concurrency_limit(&mut self, limit: usize) -> &mut Self { // } + pub fn interceptor_fn(&mut self, f: F) -> &mut Self + where + F: Fn(&mut BoxService, Request) -> Out + Send + Sync + 'static, + Out: Future, crate::Error>> + Send + 'static + { + let f = Arc::new(f); + let interceptor = layer_fn(move |mut s| { + let f = f.clone(); + tower::service_fn(move |req| f(&mut s, req)) + }); + let layer = Stack::new(interceptor, layer_fn(|s| BoxService::new(s))); + self.interceptor = Some(Arc::new(layer)); + self + } + pub async fn serve(self, addr: SocketAddr, svc: M) -> Result<(), super::Error> where M: Service<(), Response = S>, - M::Error: Into + 'static, + M::Error: Into + Send + 'static, M::Future: Send + 'static, S: Service, Response = Response> + Send + 'static, S::Future: Send + 'static, - S::Error: Into, + S::Error: Into + Send, { let tls = if let Some(tls) = self.tls { let cert = Cert { @@ -66,7 +91,10 @@ impl Builder { let incoming = hyper::server::accept::from_stream(incoming(addr, tls)); - let svc = MakeSvc(svc); + let svc = MakeSvc { + inner: svc, + interceptor: self.interceptor.clone(), + }; hyper::Server::builder(incoming) .http2_only(true) @@ -78,6 +106,12 @@ impl Builder { } } +impl fmt::Debug for Builder { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + f.debug_struct("Builder").finish() + } +} + fn incoming( addr: SocketAddr, tls: Option, @@ -128,40 +162,57 @@ struct Svc(S); impl Service> for Svc where S: Service, Response = Response>, + S::Error: Into { type Response = Response; - type Error = S::Error; - type Future = S::Future; + type Error = crate::Error; + type Future = MapErr crate::Error>; - fn poll_ready(&mut self, _cx: &mut Context<'_>) -> Poll> { - Ok(()).into() + fn poll_ready(&mut self, cx: &mut Context<'_>) -> Poll> { + self.0.poll_ready(cx).map_err(Into::into) } fn call(&mut self, req: Request) -> Self::Future { - self.0.call(req) + self.0.call(req).map_err(|e| e.into()) } } -struct MakeSvc(M); +struct MakeSvc { + interceptor: Option, + inner: M, +} impl Service for MakeSvc where M: Service<(), Response = S>, - M::Error: Into, + M::Error: Into + Send, M::Future: Send + 'static, - S: Service, Response = Response>, + S: Service, Response = Response> + Send + 'static, S::Future: Send + 'static, - S::Error: Into, + S::Error: Into + Send, { - type Response = Svc; - type Error = M::Error; - type Future = MapOk Svc>; + type Response = Either, BoxService>; + type Error = crate::Error; + type Future = Pin> + Send + 'static>>; fn poll_ready(&mut self, cx: &mut Context<'_>) -> Poll> { - MakeService::poll_ready(&mut self.0, cx) + MakeService::poll_ready(&mut self.inner, cx).map_err(Into::into) } fn call(&mut self, _: T) -> Self::Future { - self.0.make_service(()).map_ok(|s| Svc(s)) + // self.inner.make_service(()).map_ok(|s| Svc(s)) + let interceptor = self.interceptor.clone(); + // self.inner.make_service(()).map_ok(|s| intercept.layer(BoxService::new(Svc(s)))) + let make = self.inner.make_service(()); + Box::pin(async move { + let svc = make.await.map_err(Into::into)?; + + if let Some(interceptor) = interceptor { + let layered = interceptor.layer(BoxService::new(Svc(svc))); + Ok(Either::B(layered)) + } else { + Ok(Either::A(Svc(svc))) + } + }) } } diff --git a/tonic/src/transport/service/layer.rs b/tonic/src/transport/service/layer.rs index 5fe0402..4a8c11e 100644 --- a/tonic/src/transport/service/layer.rs +++ b/tonic/src/transport/service/layer.rs @@ -41,6 +41,10 @@ impl ServiceBuilderExt for ServiceBuilder { } } +pub(crate) fn layer_fn(f: F) -> LayerFn { + LayerFn(f) +} + #[derive(Clone, Copy, Debug)] pub(crate) struct LayerFn(F); diff --git a/tonic/src/transport/service/mod.rs b/tonic/src/transport/service/mod.rs index f6b6f5d..4a3c8b3 100644 --- a/tonic/src/transport/service/mod.rs +++ b/tonic/src/transport/service/mod.rs @@ -12,3 +12,4 @@ pub(crate) use self::connection::Connection; pub(crate) use self::connector::Connector; pub(crate) use self::discover::ServiceList; pub(crate) use self::io::BoxedIo; +pub(crate) use self::layer::layer_fn;