From b1cf35e35db871d54d50f2edba9df5e9fe4784cc Mon Sep 17 00:00:00 2001 From: Lucio Franco Date: Mon, 2 Sep 2019 23:02:02 -0400 Subject: [PATCH] Add server and basic server side tls --- tonic-examples/src/helloworld/server.rs | 7 +- tonic-interop/src/bin/server.rs | 24 +++- tonic/src/transport/endpoint.rs | 6 +- tonic/src/transport/mod.rs | 2 + tonic/src/transport/server.rs | 136 +++++++++++++++++++++++ tonic/src/transport/service/connector.rs | 8 +- tonic/src/transport/service/io.rs | 7 +- tonic/src/transport/service/mod.rs | 1 + tonic/src/transport/tls/mod.rs | 20 +++- tonic/src/transport/tls/openssl.rs | 39 ++++++- 10 files changed, 228 insertions(+), 22 deletions(-) create mode 100644 tonic/src/transport/server.rs diff --git a/tonic-examples/src/helloworld/server.rs b/tonic-examples/src/helloworld/server.rs index bf389bc..b5b5a1b 100644 --- a/tonic-examples/src/helloworld/server.rs +++ b/tonic-examples/src/helloworld/server.rs @@ -1,4 +1,4 @@ -use hyper::Server; +use tonic::transport::Server; use tonic::{Request, Response, Status}; pub mod hello_world { @@ -34,9 +34,8 @@ async fn main() -> Result<(), Box> { let addr = "[::1]:50051".parse().unwrap(); let greeter = MyGreeter::default(); - Server::bind(&addr) - .http2_only(true) - .serve(GreeterServer::new(greeter)) + Server::builder() + .serve(addr, GreeterServer::new(greeter)) .await?; Ok(()) diff --git a/tonic-interop/src/bin/server.rs b/tonic-interop/src/bin/server.rs index c8ea6ad..c03a3b0 100644 --- a/tonic-interop/src/bin/server.rs +++ b/tonic-interop/src/bin/server.rs @@ -1,4 +1,5 @@ -use hyper::Server; +use structopt::StructOpt; +use tonic::transport::Server; use tonic::{Code, Request, Response, Status}; pub mod pb { @@ -54,17 +55,30 @@ impl TestService { } } +#[derive(StructOpt)] +struct Opts { + #[structopt(long)] + use_tls: bool, +} + #[tokio::main] async fn main() -> Result<(), Box> { + let matches = Opts::from_args(); + pretty_env_logger::init(); let addr = "127.0.0.1:10000".parse().unwrap(); let greeter = TestService::default(); - Server::bind(&addr) - .http2_only(true) - .serve(TestServiceServer::new(greeter)) - .await?; + let mut builder = Server::builder(); + + if matches.use_tls { + let ca = tokio::fs::read("tonic-interop/data/server1.pem").await?; + let key = tokio::fs::read("tonic-interop/data/server1.key").await?; + builder.tls(ca, key); + } + + builder.serve(addr, TestServiceServer::new(greeter)).await?; Ok(()) } diff --git a/tonic/src/transport/endpoint.rs b/tonic/src/transport/endpoint.rs index 74bc0c9..9a0020e 100644 --- a/tonic/src/transport/endpoint.rs +++ b/tonic/src/transport/endpoint.rs @@ -13,7 +13,11 @@ impl Endpoint { Self { uri, - cert: Some(Cert { ca, domain }), + cert: Some(Cert { + ca, + domain, + key: None, + }), } } diff --git a/tonic/src/transport/mod.rs b/tonic/src/transport/mod.rs index 13218d0..5f68d99 100644 --- a/tonic/src/transport/mod.rs +++ b/tonic/src/transport/mod.rs @@ -1,10 +1,12 @@ mod channel; mod endpoint; +mod server; mod service; mod tls; pub use self::channel::Channel; pub use self::endpoint::Endpoint; +pub use self::server::Server; use std::{error, fmt}; diff --git a/tonic/src/transport/server.rs b/tonic/src/transport/server.rs new file mode 100644 index 0000000..27815d7 --- /dev/null +++ b/tonic/src/transport/server.rs @@ -0,0 +1,136 @@ +use super::{ + service::BoxedIo, + tls::{Cert, TlsAcceptor}, +}; +use crate::BoxBody; +use futures_util::{try_future::MapOk, TryFutureExt, TryStreamExt}; +use http::{Request, Response}; +use hyper::server::conn; +use hyper::Body; +use std::net::SocketAddr; +use std::task::{Context, Poll}; +use tower_make::MakeService; +use tower_service::Service; + +pub struct Server {} + +impl Server { + pub fn builder() -> Builder { + Builder::new() + } +} + +pub struct Builder { + tls: Option<(Vec, Vec)>, +} + +impl Builder { + fn new() -> Self { + Self { tls: None } + } + + pub fn tls(&mut self, pem: Vec, key: Vec) -> &mut Self { + self.tls = Some((pem, key)); + self + } + + // pub fn concurrency_limit(&mut self, limit: usize) -> &mut Self { + // } + + pub async fn serve(self, addr: SocketAddr, svc: M) -> Result<(), super::Error> + where + M: Service<(), Response = S>, + M::Error: Into + 'static, + M::Future: Send + 'static, + S: Service, Response = Response> + Send + 'static, + S::Future: Send + 'static, + S::Error: Into, + { + let tcp = conn::AddrIncoming::bind(&addr).unwrap(); + + let tls = if let Some(tls) = self.tls { + let cert = Cert { + ca: tls.0, + key: Some(tls.1), + domain: String::new(), + }; + + Some(TlsAcceptor::new(cert).unwrap()) + } else { + None + }; + + let incoming = incoming(tcp, tls); + + let svc = MakeSvc(svc); + + hyper::Server::builder(incoming) + .http2_only(true) + .serve(svc) + .await + .unwrap(); + + Ok(()) + } +} + +fn incoming( + mut tcp: conn::AddrIncoming, + tls: Option, +) -> impl futures_core::Stream> { + async_stream::try_stream! { + while let Some(stream) = tcp.try_next().await.map_err(Into::into)? { + if let Some(tls) = &tls { + let io = tls.connect(stream.into_inner()).await?; + yield BoxedIo::new(io); + } else { + yield BoxedIo::new(stream); + } + } + } +} + +// TODO: add custom tracing here +#[derive(Debug)] +pub struct Svc(S); + +impl Service> for Svc +where + S: Service, Response = Response>, +{ + type Response = Response; + type Error = S::Error; + type Future = S::Future; + + fn poll_ready(&mut self, _cx: &mut Context<'_>) -> Poll> { + Ok(()).into() + } + + fn call(&mut self, req: Request) -> Self::Future { + self.0.call(req) + } +} + +pub struct MakeSvc(M); + +impl Service for MakeSvc +where + M: Service<(), Response = S>, + M::Error: Into, + M::Future: Send + 'static, + S: Service, Response = Response>, + S::Future: Send + 'static, + S::Error: Into, +{ + type Response = Svc; + type Error = M::Error; + type Future = MapOk Svc>; + + fn poll_ready(&mut self, cx: &mut Context<'_>) -> Poll> { + MakeService::poll_ready(&mut self.0, cx) + } + + fn call(&mut self, _: T) -> Self::Future { + self.0.make_service(()).map_ok(|s| Svc(s)) + } +} diff --git a/tonic/src/transport/service/connector.rs b/tonic/src/transport/service/connector.rs index fd73cef..5135ba1 100644 --- a/tonic/src/transport/service/connector.rs +++ b/tonic/src/transport/service/connector.rs @@ -1,5 +1,5 @@ use super::io::BoxedIo; -use crate::transport::tls::{Cert, TlsAcceptor}; +use crate::transport::tls::{Cert, TlsConnector}; use http::Uri; use hyper::client::connect::HttpConnector; use std::future::Future; @@ -12,7 +12,7 @@ type ConnectFuture = >::Future; pub struct Connector { http: HttpConnector, - tls: Option, + tls: Option, } impl Connector { @@ -21,7 +21,7 @@ impl Connector { http.enforce_http(false); let tls = if let Some(cert) = cert { - Some(TlsAcceptor::new(cert)?) + Some(TlsConnector::new(cert)?) } else { None }; @@ -51,7 +51,7 @@ impl Service for Connector { async fn connect( connect: ConnectFuture, - tls: Option, + tls: Option, ) -> Result { let io = connect.await?; diff --git a/tonic/src/transport/service/io.rs b/tonic/src/transport/service/io.rs index 290aa18..6c8f3aa 100644 --- a/tonic/src/transport/service/io.rs +++ b/tonic/src/transport/service/io.rs @@ -3,14 +3,17 @@ use std::pin::Pin; use std::task::{Context, Poll}; use tokio::io::{AsyncRead, AsyncWrite}; -pub(super) trait Io: AsyncRead + AsyncWrite + Send + Unpin + 'static {} +pub(in crate::transport) trait Io: + AsyncRead + AsyncWrite + Send + Unpin + 'static +{ +} impl Io for T where T: AsyncRead + AsyncWrite + Send + Unpin + 'static {} pub struct BoxedIo(Pin>); impl BoxedIo { - pub(super) fn new(io: I) -> Self { + pub(in crate::transport) fn new(io: I) -> Self { BoxedIo(Box::pin(io)) } } diff --git a/tonic/src/transport/service/mod.rs b/tonic/src/transport/service/mod.rs index f75738d..fccba24 100644 --- a/tonic/src/transport/service/mod.rs +++ b/tonic/src/transport/service/mod.rs @@ -9,3 +9,4 @@ pub use self::add_origin::AddOrigin; pub use self::boxed::BoxService; pub use self::connect::Connection; pub use self::discover::ServiceList; +pub use self::io::BoxedIo; diff --git a/tonic/src/transport/tls/mod.rs b/tonic/src/transport/tls/mod.rs index b47a9ca..454fb06 100644 --- a/tonic/src/transport/tls/mod.rs +++ b/tonic/src/transport/tls/mod.rs @@ -1,4 +1,5 @@ -// #[cfg(feature = "openssl-1")] +// TODO: bring back rustls +// #[cfg(feature = "native-tls")] // #[cfg(not(feature = "rustls"))] // #[path = "rustls.rs"] // mod imp; @@ -13,9 +14,26 @@ use tokio::net::TcpStream; #[derive(Debug, Clone)] pub struct Cert { pub(crate) ca: Vec, + pub(crate) key: Option>, pub(crate) domain: String, } +#[derive(Clone)] +pub struct TlsConnector { + inner: imp::TlsConnector, +} + +impl TlsConnector { + pub fn new(cert: Cert) -> Result { + let inner = imp::TlsConnector::new(cert)?; + Ok(Self { inner }) + } + + pub async fn connect(&self, io: TcpStream) -> Result { + self.inner.connect(io).await + } +} + #[derive(Clone)] pub struct TlsAcceptor { inner: imp::TlsAcceptor, diff --git a/tonic/src/transport/tls/openssl.rs b/tonic/src/transport/tls/openssl.rs index 17bfc1f..a8d9e8c 100644 --- a/tonic/src/transport/tls/openssl.rs +++ b/tonic/src/transport/tls/openssl.rs @@ -1,6 +1,6 @@ use super::Cert; -use openssl::ssl::{SslConnector, SslMethod}; -use openssl::x509::X509; +use openssl::ssl::{SslAcceptor, SslConnector, SslMethod}; +use openssl::{pkey::PKey, x509::X509}; use std::sync::Arc; use tokio::net::TcpStream; use tokio_openssl::SslStream; @@ -10,14 +10,14 @@ const ALPN_H2: &[u8] = b"\x02h2"; pub type TlsStream = SslStream; #[derive(Clone)] -pub struct TlsAcceptor { +pub struct TlsConnector { config: SslConnector, domain: Arc, } -impl TlsAcceptor { +impl TlsConnector { pub fn new(cert: Cert) -> Result { - let Cert { ca, domain } = cert; + let Cert { ca, domain, .. } = cert; let mut config = SslConnector::builder(SslMethod::tls()).unwrap(); config.set_alpn_protos(ALPN_H2)?; @@ -40,3 +40,32 @@ impl TlsAcceptor { Ok(tls) } } + +#[derive(Clone)] +pub struct TlsAcceptor { + config: SslAcceptor, +} + +impl TlsAcceptor { + pub fn new(cert: Cert) -> Result { + let Cert { ca, key, .. } = cert; + + let key = PKey::private_key_from_pem(&key.unwrap()[..])?; + let ca = X509::from_pem(&ca[..])?; + + let mut config = SslAcceptor::mozilla_modern(SslMethod::tls())?; + + config.set_private_key(&key)?; + config.set_certificate(&ca)?; + + Ok(Self { + config: config.build(), + }) + } + + pub async fn connect(&self, io: TcpStream) -> Result { + let config = self.config.clone(); + let tls = tokio_openssl::accept(&config, io).await?; + Ok(tls) + } +}