From 04a8c0c82a4007f48c3bf3539a3f2312746fedd1 Mon Sep 17 00:00:00 2001 From: Lucio Franco Date: Wed, 1 Apr 2020 13:02:19 -0400 Subject: [PATCH] fix(transport): Handle tls accepting on task (#320) Signed-off-by: Lucio Franco --- tonic/src/transport/server/conn.rs | 29 ++++--- tonic/src/transport/server/incoming.rs | 113 +++++++++++++++++++++++-- tonic/src/transport/server/mod.rs | 3 + tonic/src/transport/service/tls.rs | 14 +-- 4 files changed, 132 insertions(+), 27 deletions(-) diff --git a/tonic/src/transport/server/conn.rs b/tonic/src/transport/server/conn.rs index 1f4900a..d7d4588 100644 --- a/tonic/src/transport/server/conn.rs +++ b/tonic/src/transport/server/conn.rs @@ -1,9 +1,11 @@ +#[cfg(feature = "tls")] +use super::TlsStream; use crate::transport::Certificate; use hyper::server::conn::AddrStream; use std::net::SocketAddr; use tokio::net::TcpStream; #[cfg(feature = "tls")] -use tokio_rustls::{rustls::Session, server::TlsStream}; +use tokio_rustls::rustls::Session; /// Trait that connected IO resources implement. /// @@ -37,19 +39,24 @@ impl Connected for TcpStream { #[cfg(feature = "tls")] impl Connected for TlsStream { fn remote_addr(&self) -> Option { - let (inner, _) = self.get_ref(); - inner.remote_addr() + if let Some((inner, _)) = self.get_ref() { + inner.remote_addr() + } else { + None + } } fn peer_certs(&self) -> Option> { - 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) + if let Some((_, 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 + } } else { None } diff --git a/tonic/src/transport/server/incoming.rs b/tonic/src/transport/server/incoming.rs index 5232d6b..e33c77e 100644 --- a/tonic/src/transport/server/incoming.rs +++ b/tonic/src/transport/server/incoming.rs @@ -13,8 +13,6 @@ use std::{ time::Duration, }; use tokio::io::{AsyncRead, AsyncWrite}; -#[cfg(feature = "tls")] -use tracing::error; #[cfg_attr(not(feature = "tls"), allow(unused_variables))] pub(crate) fn tcp_incoming( @@ -32,13 +30,7 @@ where #[cfg(feature = "tls")] { if let Some(tls) = &server.tls { - let io = match tls.accept(stream).await { - Ok(io) => io, - Err(error) => { - error!(message = "Unable to accept incoming connection.", %error); - continue - }, - }; + let io = tls.accept(stream); yield ServerIo::new(io); continue; } @@ -73,3 +65,106 @@ impl Stream for TcpIncoming { Pin::new(&mut self.inner).poll_accept(cx) } } + +// tokio_rustls::server::TlsStream doesn't expose constructor methods, +// so we have to TlsAcceptor::accept and handshake to have access to it +// TlsStream implements AsyncRead/AsyncWrite handshaking tokio_rustls::Accept first +#[cfg(feature = "tls")] +pub(crate) struct TlsStream { + state: State, +} + +#[cfg(feature = "tls")] +enum State { + Handshaking(tokio_rustls::Accept), + Streaming(tokio_rustls::server::TlsStream), +} + +#[cfg(feature = "tls")] +impl TlsStream { + pub(crate) fn new(accept: tokio_rustls::Accept) -> Self { + TlsStream { + state: State::Handshaking(accept), + } + } + + pub(crate) fn get_ref(&self) -> Option<(&IO, &tokio_rustls::rustls::ServerSession)> { + if let State::Streaming(tls) = &self.state { + Some(tls.get_ref()) + } else { + None + } + } +} + +#[cfg(feature = "tls")] +impl AsyncRead for TlsStream +where + IO: AsyncRead + AsyncWrite + Connected + Unpin + Send + 'static, +{ + fn poll_read( + self: Pin<&mut Self>, + cx: &mut Context<'_>, + buf: &mut [u8], + ) -> Poll> { + use std::future::Future; + + let pin = self.get_mut(); + match pin.state { + State::Handshaking(ref mut accept) => { + match futures_core::ready!(Pin::new(accept).poll(cx)) { + Ok(mut stream) => { + let result = Pin::new(&mut stream).poll_read(cx, buf); + pin.state = State::Streaming(stream); + result + } + Err(err) => Poll::Ready(Err(err)), + } + } + State::Streaming(ref mut stream) => Pin::new(stream).poll_read(cx, buf), + } + } +} + +#[cfg(feature = "tls")] +impl AsyncWrite for TlsStream +where + IO: AsyncRead + AsyncWrite + Connected + Unpin + Send + 'static, +{ + fn poll_write( + self: Pin<&mut Self>, + cx: &mut Context<'_>, + buf: &[u8], + ) -> Poll> { + use std::future::Future; + + let pin = self.get_mut(); + match pin.state { + State::Handshaking(ref mut accept) => { + match futures_core::ready!(Pin::new(accept).poll(cx)) { + Ok(mut stream) => { + let result = Pin::new(&mut stream).poll_write(cx, buf); + pin.state = State::Streaming(stream); + result + } + Err(err) => Poll::Ready(Err(err)), + } + } + State::Streaming(ref mut stream) => Pin::new(stream).poll_write(cx, buf), + } + } + + fn poll_flush(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll> { + match self.state { + State::Handshaking(_) => Poll::Ready(Ok(())), + State::Streaming(ref mut stream) => Pin::new(stream).poll_flush(cx), + } + } + + fn poll_shutdown(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll> { + match self.state { + State::Handshaking(_) => Poll::Ready(Ok(())), + State::Streaming(ref mut stream) => Pin::new(stream).poll_shutdown(cx), + } + } +} diff --git a/tonic/src/transport/server/mod.rs b/tonic/src/transport/server/mod.rs index 4ff5aa0..ce84b7d 100644 --- a/tonic/src/transport/server/mod.rs +++ b/tonic/src/transport/server/mod.rs @@ -15,6 +15,9 @@ use super::service::TlsAcceptor; use incoming::TcpIncoming; +#[cfg(feature = "tls")] +pub(crate) use incoming::TlsStream; + use super::service::{Or, Routes, ServerIo, ServiceBuilderExt}; use crate::{body::BoxBody, request::ConnectionInfo}; use futures_core::Stream; diff --git a/tonic/src/transport/service/tls.rs b/tonic/src/transport/service/tls.rs index 0de574d..eef0f9a 100644 --- a/tonic/src/transport/service/tls.rs +++ b/tonic/src/transport/service/tls.rs @@ -1,5 +1,8 @@ use super::io::BoxedIo; -use crate::transport::{server::Connected, Certificate, Identity}; +use crate::transport::{ + server::{Connected, TlsStream}, + Certificate, Identity, +}; #[cfg(feature = "tls-roots")] use rustls_native_certs; use std::{fmt, sync::Arc}; @@ -157,17 +160,14 @@ impl TlsAcceptor { }) } - pub(crate) async fn accept( - &self, - io: IO, - ) -> Result, crate::Error> + pub(crate) fn accept(&self, io: IO) -> TlsStream where IO: AsyncRead + AsyncWrite + Connected + Unpin + Send + 'static, { let acceptor = RustlsAcceptor::from(self.inner.clone()); - let tls = acceptor.accept(io).await?; + let accept = acceptor.accept(io); - Ok(tls) + TlsStream::new(accept) } }