feat(transport): Add server side peer cert support (#228)

* feat(transport): Add server side peer cert support
This commit is contained in:
Lucio Franco
2020-01-11 14:17:19 -08:00
committed by GitHub
parent 3be8bc1668
commit af807c3ccd
7 changed files with 59 additions and 13 deletions
+4
View File
@@ -17,6 +17,10 @@ pub struct EchoServer;
#[tonic::async_trait] #[tonic::async_trait]
impl pb::echo_server::Echo for EchoServer { impl pb::echo_server::Echo for EchoServer {
async fn unary_echo(&self, request: Request<EchoRequest>) -> EchoResult<EchoResponse> { async fn unary_echo(&self, request: Request<EchoRequest>) -> EchoResult<EchoResponse> {
if let Some(certs) = request.peer_certs() {
println!("Got {} peer certs!", certs.len());
}
let message = request.into_inner().message; let message = request.into_inner().message;
Ok(Response::new(EchoResponse { message })) Ok(Response::new(EchoResponse { message }))
} }
+2 -2
View File
@@ -33,8 +33,8 @@ transport = [
"tower-load", "tower-load",
"tracing-futures", "tracing-futures",
] ]
tls = ["tokio-rustls"] tls = ["transport", "tokio-rustls"]
tls-roots = ["rustls-native-certs"] tls-roots = ["tls", "rustls-native-certs"]
# [[bench]] # [[bench]]
# name = "bench_main" # name = "bench_main"
+18
View File
@@ -1,7 +1,11 @@
use crate::metadata::MetadataMap; use crate::metadata::MetadataMap;
#[cfg(feature = "transport")]
use crate::transport::Certificate;
use futures_core::Stream; use futures_core::Stream;
use http::Extensions; use http::Extensions;
use std::net::SocketAddr; use std::net::SocketAddr;
#[cfg(feature = "transport")]
use std::sync::Arc;
/// A gRPC request and metadata from an RPC call. /// A gRPC request and metadata from an RPC call.
#[derive(Debug)] #[derive(Debug)]
@@ -14,6 +18,8 @@ pub struct Request<T> {
#[derive(Clone)] #[derive(Clone)]
pub(crate) struct ConnectionInfo { pub(crate) struct ConnectionInfo {
pub(crate) remote_addr: Option<SocketAddr>, pub(crate) remote_addr: Option<SocketAddr>,
#[cfg(feature = "transport")]
pub(crate) peer_certs: Option<Arc<Vec<Certificate>>>,
} }
/// Trait implemented by RPC request types. /// Trait implemented by RPC request types.
@@ -188,6 +194,18 @@ impl<T> Request<T> {
self.get::<ConnectionInfo>()?.remote_addr self.get::<ConnectionInfo>()?.remote_addr
} }
/// Get the peer certificates of the connected client.
///
/// This is used to fetch the certificates from the TLS session
/// and is mostly used for mTLS. This currently only returns
/// `Some` on the server side of the `transport` server with
/// TLS enabled connections.
#[cfg(feature = "transport")]
#[cfg_attr(docsrs, doc(cfg(feature = "transport")))]
pub fn peer_certs(&self) -> Option<Arc<Vec<Certificate>>> {
self.get::<ConnectionInfo>()?.peer_certs.clone()
}
pub(crate) fn get<I: Send + Sync + 'static>(&self) -> Option<&I> { pub(crate) fn get<I: Send + Sync + 'static>(&self) -> Option<&I> {
self.extensions.get::<I>() self.extensions.get::<I>()
} }
+21 -1
View File
@@ -1,7 +1,8 @@
use crate::transport::Certificate;
use hyper::server::conn::AddrStream; use hyper::server::conn::AddrStream;
use std::net::SocketAddr; use std::net::SocketAddr;
#[cfg(feature = "tls")] #[cfg(feature = "tls")]
use tokio_rustls::TlsStream; use tokio_rustls::{rustls::Session, server::TlsStream};
/// Trait that connected IO resources implement. /// Trait that connected IO resources implement.
/// ///
@@ -13,6 +14,11 @@ pub trait Connected {
fn remote_addr(&self) -> Option<SocketAddr> { fn remote_addr(&self) -> Option<SocketAddr> {
None None
} }
/// Return the set of connected peer TLS certificates.
fn peer_certs(&self) -> Option<Vec<Certificate>> {
None
}
} }
impl Connected for AddrStream { impl Connected for AddrStream {
@@ -27,4 +33,18 @@ impl<T: Connected> Connected for TlsStream<T> {
let (inner, _) = self.get_ref(); let (inner, _) = self.get_ref();
inner.remote_addr() inner.remote_addr()
} }
fn peer_certs(&self) -> Option<Vec<Certificate>> {
let (_, session) = self.get_ref();
if let Some(certs) = session.get_peer_certificates() {
let certs = certs
.into_iter()
.map(|c| Certificate::from_pem(c.0))
.collect();
Some(certs)
} else {
None
}
}
} }
+1
View File
@@ -505,6 +505,7 @@ where
fn call(&mut self, io: &ServerIo) -> Self::Future { fn call(&mut self, io: &ServerIo) -> Self::Future {
let conn_info = crate::request::ConnectionInfo { let conn_info = crate::request::ConnectionInfo {
remote_addr: io.remote_addr(), remote_addr: io.remote_addr(),
peer_certs: io.peer_certs().map(Arc::new),
}; };
let interceptor = self.interceptor.clone(); let interceptor = self.interceptor.clone();
+6 -3
View File
@@ -1,4 +1,4 @@
use crate::transport::server::Connected; use crate::transport::{server::Connected, Certificate};
use hyper::client::connect::{Connected as HyperConnected, Connection}; use hyper::client::connect::{Connected as HyperConnected, Connection};
use std::io; use std::io;
use std::net::SocketAddr; use std::net::SocketAddr;
@@ -71,8 +71,11 @@ impl ServerIo {
impl Connected for ServerIo { impl Connected for ServerIo {
fn remote_addr(&self) -> Option<SocketAddr> { fn remote_addr(&self) -> Option<SocketAddr> {
let io = &*self.0; (&*self.0).remote_addr()
io.remote_addr() }
fn peer_certs(&self) -> Option<Vec<Certificate>> {
(&self.0).peer_certs()
} }
} }
+7 -7
View File
@@ -157,17 +157,17 @@ impl TlsAcceptor {
}) })
} }
pub(crate) async fn accept<IO>(&self, io: IO) -> Result<BoxedIo, crate::Error> pub(crate) async fn accept<IO>(
&self,
io: IO,
) -> Result<tokio_rustls::server::TlsStream<IO>, crate::Error>
where where
IO: AsyncRead + AsyncWrite + Connected + Unpin + Send + 'static, IO: AsyncRead + AsyncWrite + Connected + Unpin + Send + 'static,
{ {
let io = { let acceptor = RustlsAcceptor::from(self.inner.clone());
let acceptor = RustlsAcceptor::from(self.inner.clone()); let tls = acceptor.accept(io).await?;
let tls = acceptor.accept(io).await?;
BoxedIo::new(tls)
};
Ok(io) Ok(tls)
} }
} }