Add server and basic server side tls

This commit is contained in:
Lucio Franco
2019-09-02 23:02:02 -04:00
parent c30b004475
commit b1cf35e35d
10 changed files with 228 additions and 22 deletions
+3 -4
View File
@@ -1,4 +1,4 @@
use hyper::Server; use tonic::transport::Server;
use tonic::{Request, Response, Status}; use tonic::{Request, Response, Status};
pub mod hello_world { pub mod hello_world {
@@ -34,9 +34,8 @@ async fn main() -> Result<(), Box<dyn std::error::Error>> {
let addr = "[::1]:50051".parse().unwrap(); let addr = "[::1]:50051".parse().unwrap();
let greeter = MyGreeter::default(); let greeter = MyGreeter::default();
Server::bind(&addr) Server::builder()
.http2_only(true) .serve(addr, GreeterServer::new(greeter))
.serve(GreeterServer::new(greeter))
.await?; .await?;
Ok(()) Ok(())
+19 -5
View File
@@ -1,4 +1,5 @@
use hyper::Server; use structopt::StructOpt;
use tonic::transport::Server;
use tonic::{Code, Request, Response, Status}; use tonic::{Code, Request, Response, Status};
pub mod pb { pub mod pb {
@@ -54,17 +55,30 @@ impl TestService {
} }
} }
#[derive(StructOpt)]
struct Opts {
#[structopt(long)]
use_tls: bool,
}
#[tokio::main] #[tokio::main]
async fn main() -> Result<(), Box<dyn std::error::Error>> { async fn main() -> Result<(), Box<dyn std::error::Error>> {
let matches = Opts::from_args();
pretty_env_logger::init(); pretty_env_logger::init();
let addr = "127.0.0.1:10000".parse().unwrap(); let addr = "127.0.0.1:10000".parse().unwrap();
let greeter = TestService::default(); let greeter = TestService::default();
Server::bind(&addr) let mut builder = Server::builder();
.http2_only(true)
.serve(TestServiceServer::new(greeter)) if matches.use_tls {
.await?; 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(()) Ok(())
} }
+5 -1
View File
@@ -13,7 +13,11 @@ impl Endpoint {
Self { Self {
uri, uri,
cert: Some(Cert { ca, domain }), cert: Some(Cert {
ca,
domain,
key: None,
}),
} }
} }
+2
View File
@@ -1,10 +1,12 @@
mod channel; mod channel;
mod endpoint; mod endpoint;
mod server;
mod service; mod service;
mod tls; mod tls;
pub use self::channel::Channel; pub use self::channel::Channel;
pub use self::endpoint::Endpoint; pub use self::endpoint::Endpoint;
pub use self::server::Server;
use std::{error, fmt}; use std::{error, fmt};
+136
View File
@@ -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<u8>, Vec<u8>)>,
}
impl Builder {
fn new() -> Self {
Self { tls: None }
}
pub fn tls(&mut self, pem: Vec<u8>, key: Vec<u8>) -> &mut Self {
self.tls = Some((pem, key));
self
}
// pub fn concurrency_limit(&mut self, limit: usize) -> &mut Self {
// }
pub async fn serve<M, S>(self, addr: SocketAddr, svc: M) -> Result<(), super::Error>
where
M: Service<(), Response = S>,
M::Error: Into<crate::Error> + 'static,
M::Future: Send + 'static,
S: Service<Request<Body>, Response = Response<BoxBody>> + Send + 'static,
S::Future: Send + 'static,
S::Error: Into<crate::Error>,
{
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<TlsAcceptor>,
) -> impl futures_core::Stream<Item = Result<BoxedIo, crate::Error>> {
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>(S);
impl<S> Service<Request<Body>> for Svc<S>
where
S: Service<Request<Body>, Response = Response<BoxBody>>,
{
type Response = Response<BoxBody>;
type Error = S::Error;
type Future = S::Future;
fn poll_ready(&mut self, _cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
Ok(()).into()
}
fn call(&mut self, req: Request<Body>) -> Self::Future {
self.0.call(req)
}
}
pub struct MakeSvc<M>(M);
impl<M, S, T> Service<T> for MakeSvc<M>
where
M: Service<(), Response = S>,
M::Error: Into<crate::Error>,
M::Future: Send + 'static,
S: Service<Request<Body>, Response = Response<BoxBody>>,
S::Future: Send + 'static,
S::Error: Into<crate::Error>,
{
type Response = Svc<S>;
type Error = M::Error;
type Future = MapOk<M::Future, fn(S) -> Svc<S>>;
fn poll_ready(&mut self, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
MakeService::poll_ready(&mut self.0, cx)
}
fn call(&mut self, _: T) -> Self::Future {
self.0.make_service(()).map_ok(|s| Svc(s))
}
}
+4 -4
View File
@@ -1,5 +1,5 @@
use super::io::BoxedIo; use super::io::BoxedIo;
use crate::transport::tls::{Cert, TlsAcceptor}; use crate::transport::tls::{Cert, TlsConnector};
use http::Uri; use http::Uri;
use hyper::client::connect::HttpConnector; use hyper::client::connect::HttpConnector;
use std::future::Future; use std::future::Future;
@@ -12,7 +12,7 @@ type ConnectFuture = <HttpConnector as MakeConnection<Uri>>::Future;
pub struct Connector { pub struct Connector {
http: HttpConnector, http: HttpConnector,
tls: Option<TlsAcceptor>, tls: Option<TlsConnector>,
} }
impl Connector { impl Connector {
@@ -21,7 +21,7 @@ impl Connector {
http.enforce_http(false); http.enforce_http(false);
let tls = if let Some(cert) = cert { let tls = if let Some(cert) = cert {
Some(TlsAcceptor::new(cert)?) Some(TlsConnector::new(cert)?)
} else { } else {
None None
}; };
@@ -51,7 +51,7 @@ impl Service<Uri> for Connector {
async fn connect( async fn connect(
connect: ConnectFuture, connect: ConnectFuture,
tls: Option<TlsAcceptor>, tls: Option<TlsConnector>,
) -> Result<BoxedIo, crate::Error> { ) -> Result<BoxedIo, crate::Error> {
let io = connect.await?; let io = connect.await?;
+5 -2
View File
@@ -3,14 +3,17 @@ use std::pin::Pin;
use std::task::{Context, Poll}; use std::task::{Context, Poll};
use tokio::io::{AsyncRead, AsyncWrite}; 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<T> Io for T where T: AsyncRead + AsyncWrite + Send + Unpin + 'static {} impl<T> Io for T where T: AsyncRead + AsyncWrite + Send + Unpin + 'static {}
pub struct BoxedIo(Pin<Box<dyn Io>>); pub struct BoxedIo(Pin<Box<dyn Io>>);
impl BoxedIo { impl BoxedIo {
pub(super) fn new<I: Io>(io: I) -> Self { pub(in crate::transport) fn new<I: Io>(io: I) -> Self {
BoxedIo(Box::pin(io)) BoxedIo(Box::pin(io))
} }
} }
+1
View File
@@ -9,3 +9,4 @@ pub use self::add_origin::AddOrigin;
pub use self::boxed::BoxService; pub use self::boxed::BoxService;
pub use self::connect::Connection; pub use self::connect::Connection;
pub use self::discover::ServiceList; pub use self::discover::ServiceList;
pub use self::io::BoxedIo;
+19 -1
View File
@@ -1,4 +1,5 @@
// #[cfg(feature = "openssl-1")] // TODO: bring back rustls
// #[cfg(feature = "native-tls")]
// #[cfg(not(feature = "rustls"))] // #[cfg(not(feature = "rustls"))]
// #[path = "rustls.rs"] // #[path = "rustls.rs"]
// mod imp; // mod imp;
@@ -13,9 +14,26 @@ use tokio::net::TcpStream;
#[derive(Debug, Clone)] #[derive(Debug, Clone)]
pub struct Cert { pub struct Cert {
pub(crate) ca: Vec<u8>, pub(crate) ca: Vec<u8>,
pub(crate) key: Option<Vec<u8>>,
pub(crate) domain: String, pub(crate) domain: String,
} }
#[derive(Clone)]
pub struct TlsConnector {
inner: imp::TlsConnector,
}
impl TlsConnector {
pub fn new(cert: Cert) -> Result<Self, crate::Error> {
let inner = imp::TlsConnector::new(cert)?;
Ok(Self { inner })
}
pub async fn connect(&self, io: TcpStream) -> Result<imp::TlsStream, crate::Error> {
self.inner.connect(io).await
}
}
#[derive(Clone)] #[derive(Clone)]
pub struct TlsAcceptor { pub struct TlsAcceptor {
inner: imp::TlsAcceptor, inner: imp::TlsAcceptor,
+34 -5
View File
@@ -1,6 +1,6 @@
use super::Cert; use super::Cert;
use openssl::ssl::{SslConnector, SslMethod}; use openssl::ssl::{SslAcceptor, SslConnector, SslMethod};
use openssl::x509::X509; use openssl::{pkey::PKey, x509::X509};
use std::sync::Arc; use std::sync::Arc;
use tokio::net::TcpStream; use tokio::net::TcpStream;
use tokio_openssl::SslStream; use tokio_openssl::SslStream;
@@ -10,14 +10,14 @@ const ALPN_H2: &[u8] = b"\x02h2";
pub type TlsStream = SslStream<TcpStream>; pub type TlsStream = SslStream<TcpStream>;
#[derive(Clone)] #[derive(Clone)]
pub struct TlsAcceptor { pub struct TlsConnector {
config: SslConnector, config: SslConnector,
domain: Arc<String>, domain: Arc<String>,
} }
impl TlsAcceptor { impl TlsConnector {
pub fn new(cert: Cert) -> Result<Self, crate::Error> { pub fn new(cert: Cert) -> Result<Self, crate::Error> {
let Cert { ca, domain } = cert; let Cert { ca, domain, .. } = cert;
let mut config = SslConnector::builder(SslMethod::tls()).unwrap(); let mut config = SslConnector::builder(SslMethod::tls()).unwrap();
config.set_alpn_protos(ALPN_H2)?; config.set_alpn_protos(ALPN_H2)?;
@@ -40,3 +40,32 @@ impl TlsAcceptor {
Ok(tls) Ok(tls)
} }
} }
#[derive(Clone)]
pub struct TlsAcceptor {
config: SslAcceptor,
}
impl TlsAcceptor {
pub fn new(cert: Cert) -> Result<Self, crate::Error> {
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<TlsStream, crate::Error> {
let config = self.config.clone();
let tls = tokio_openssl::accept(&config, io).await?;
Ok(tls)
}
}