feat(transport): Allow custom IO and UDS example (#184)

Closes #136
This commit is contained in:
Lucio Franco
2019-12-13 17:14:49 -05:00
committed by GitHub
parent 7077d8dfd0
commit b90c340800
12 changed files with 306 additions and 115 deletions
+9 -1
View File
@@ -78,12 +78,20 @@ path = "src/tracing/client.rs"
name = "tracing-server"
path = "src/tracing/server.rs"
[[bin]]
name = "uds-client"
path = "src/uds/client.rs"
[[bin]]
name = "uds-server"
path = "src/uds/server.rs"
[dependencies]
tonic = { path = "../tonic", features = ["tls"] }
bytes = "0.4"
prost = "0.5"
tokio = { version = "0.2", features = ["rt-threaded", "time", "stream", "fs", "macros"] }
tokio = { version = "0.2", features = ["rt-threaded", "time", "stream", "fs", "macros", "uds"] }
futures = { version = "0.3", default-features = false, features = ["alloc"]}
async-stream = "0.2"
http = "0.2"
+39
View File
@@ -0,0 +1,39 @@
#[cfg(unix)]
pub mod hello_world {
tonic::include_proto!("helloworld");
}
use hello_world::{greeter_client::GreeterClient, HelloRequest};
use http::Uri;
use std::convert::TryFrom;
use tokio::net::UnixStream;
use tonic::transport::Endpoint;
use tower::service_fn;
#[tokio::main]
async fn main() -> Result<(), Box<dyn std::error::Error>> {
// We will ignore this uri because uds do not use it
// if your connector does use the uri it will be provided
// as the request to the `MakeConnection`.
let channel = Endpoint::try_from("lttp://[::]:50051")?
.connect_with_connector(service_fn(|_: Uri| {
let path = "/tmp/tonic/helloworld";
// Connect to a Uds socket
UnixStream::connect(path)
}))
.await?;
let mut client = GreeterClient::new(channel);
let request = tonic::Request::new(HelloRequest {
name: "Tonic".into(),
});
let response = client.say_hello(request).await?;
println!("RESPONSE={:?}", response);
Ok(())
}
+48
View File
@@ -0,0 +1,48 @@
use std::path::Path;
use tokio::net::UnixListener;
use tonic::{transport::Server, Request, Response, Status};
pub mod hello_world {
tonic::include_proto!("helloworld");
}
use hello_world::{
greeter_server::{Greeter, GreeterServer},
HelloReply, HelloRequest,
};
#[derive(Default)]
pub struct MyGreeter {}
#[tonic::async_trait]
impl Greeter for MyGreeter {
async fn say_hello(
&self,
request: Request<HelloRequest>,
) -> Result<Response<HelloReply>, Status> {
println!("Got a request: {:?}", request);
let reply = hello_world::HelloReply {
message: format!("Hello {}!", request.into_inner().name).into(),
};
Ok(Response::new(reply))
}
}
#[tokio::main]
async fn main() -> Result<(), Box<dyn std::error::Error>> {
let path = "/tmp/tonic/helloworld";
tokio::fs::create_dir_all(Path::new(path).parent().unwrap()).await?;
let mut uds = UnixListener::bind(path)?;
let greeter = MyGreeter::default();
Server::builder()
.add_service(GreeterServer::new(greeter))
.serve_with_incoming(uds.incoming())
.await?;
Ok(())
}
+31 -1
View File
@@ -1,3 +1,4 @@
use super::super::service;
use super::Channel;
#[cfg(feature = "tls")]
use super::ClientTlsConfig;
@@ -12,6 +13,7 @@ use std::{
sync::Arc,
time::Duration,
};
use tower_make::MakeConnection;
/// Channel builder.
///
@@ -182,7 +184,35 @@ impl Endpoint {
/// Create a channel from this config.
pub async fn connect(&self) -> Result<Channel, Error> {
Channel::connect(self.clone()).await
let mut http = hyper::client::connect::HttpConnector::new();
http.enforce_http(false);
http.set_nodelay(self.tcp_nodelay);
http.set_keepalive(self.tcp_keepalive);
#[cfg(feature = "tls")]
let connector = service::connector(http, self.tls.clone());
#[cfg(not(feature = "tls"))]
let connector = service::connector(http);
Channel::connect(connector, self.clone()).await
}
/// Connect with a custom connector.
pub async fn connect_with_connector<C>(&self, connector: C) -> Result<Channel, Error>
where
C: MakeConnection<Uri> + Send + 'static,
C::Connection: Unpin + Send + 'static,
C::Future: Send + 'static,
crate::Error: From<C::Error> + Send + 'static,
{
#[cfg(feature = "tls")]
let connector = service::connector(connector, self.tls.clone());
#[cfg(not(feature = "tls"))]
let connector = service::connector(connector);
Channel::connect(connector, self.clone()).await
}
}
+10 -2
View File
@@ -15,6 +15,7 @@ use http::{
uri::{InvalidUri, Uri},
Request, Response,
};
use hyper::client::connect::Connection as HyperConnection;
use std::{
fmt,
future::Future,
@@ -22,6 +23,7 @@ use std::{
sync::Arc,
task::{Context, Poll},
};
use tokio::io::{AsyncRead, AsyncWrite};
use tower::{
buffer::{self, Buffer},
discover::Discover,
@@ -121,11 +123,17 @@ impl Channel {
Self::balance(discover, buffer_size, interceptor_headers)
}
pub(crate) async fn connect(endpoint: Endpoint) -> Result<Self, super::Error> {
pub(crate) async fn connect<C>(connector: C, endpoint: Endpoint) -> Result<Self, super::Error>
where
C: Service<Uri> + Send + 'static,
C::Error: Into<crate::Error> + Send,
C::Future: Unpin + Send,
C::Response: AsyncRead + AsyncWrite + HyperConnection + Unpin + Send + 'static,
{
let buffer_size = endpoint.buffer_size.clone().unwrap_or(DEFAULT_BUFFER_SIZE);
let interceptor_headers = endpoint.interceptor_headers.clone();
let svc = Connection::new(endpoint)
let svc = Connection::new(connector, endpoint)
.await
.map_err(|e| super::Error::from_source(super::ErrorKind::Client, e))?;
+75
View File
@@ -0,0 +1,75 @@
use super::Server;
use crate::transport::service::BoxedIo;
use futures_core::Stream;
use futures_util::stream::TryStreamExt;
use hyper::server::{
accept::Accept,
conn::{AddrIncoming, AddrStream},
};
use std::{
net::SocketAddr,
pin::Pin,
task::{Context, Poll},
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<IO, IE>(
incoming: impl Stream<Item = Result<IO, IE>>,
server: Server,
) -> impl Stream<Item = Result<BoxedIo, crate::Error>>
where
IO: AsyncRead + AsyncWrite + Unpin + Send + 'static,
IE: Into<crate::Error>,
{
async_stream::try_stream! {
futures_util::pin_mut!(incoming);
while let Some(stream) = incoming.try_next().await? {
#[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
},
};
yield BoxedIo::new(io);
continue;
}
}
yield BoxedIo::new(stream);
}
}
}
pub(crate) struct TcpIncoming {
inner: AddrIncoming,
}
impl TcpIncoming {
pub(crate) fn new(
addr: SocketAddr,
nodelay: bool,
keepalive: Option<Duration>,
) -> Result<Self, crate::Error> {
let mut inner = AddrIncoming::bind(&addr)?;
inner.set_nodelay(nodelay);
inner.set_keepalive(keepalive);
Ok(TcpIncoming { inner })
}
}
impl Stream for TcpIncoming {
type Item = Result<AddrStream, std::io::Error>;
fn poll_next(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
Pin::new(&mut self.inner).poll_accept(cx)
}
}
+37 -56
View File
@@ -1,5 +1,6 @@
//! Server implementation and builder.
mod incoming;
#[cfg(feature = "tls")]
mod tls;
@@ -9,18 +10,17 @@ pub use tls::ServerTlsConfig;
#[cfg(feature = "tls")]
use super::service::TlsAcceptor;
use super::service::{layer_fn, BoxedIo, Or, Routes, ServiceBuilderExt};
use incoming::TcpIncoming;
use super::service::{layer_fn, Or, Routes, ServiceBuilderExt};
use crate::body::BoxBody;
use futures_core::Stream;
use futures_util::{
future::{self, poll_fn, MapErr},
future::{self, MapErr},
TryFutureExt,
};
use http::{HeaderMap, Request, Response};
use hyper::{
server::{accept::Accept, conn},
Body,
};
use std::time::Duration;
use hyper::{server::accept, Body};
use std::{
fmt,
future::Future,
@@ -28,8 +28,9 @@ use std::{
pin::Pin,
sync::Arc,
task::{Context, Poll},
// time::Duration,
time::Duration,
};
use tokio::io::{AsyncRead, AsyncWrite};
use tower::{
layer::{Layer, Stack},
limit::concurrency::ConcurrencyLimitLayer,
@@ -37,8 +38,6 @@ use tower::{
Service,
ServiceBuilder,
};
#[cfg(feature = "tls")]
use tracing::error;
use tracing_futures::{Instrument, Instrumented};
type BoxService = tower::util::BoxService<Request<Body>, Response<BoxBody>, crate::Error>;
@@ -242,16 +241,19 @@ impl Server {
Router::new(self.clone(), svc)
}
pub(crate) async fn serve_with_shutdown<S, F>(
pub(crate) async fn serve_with_shutdown<S, I, F, IO, IE>(
self,
addr: SocketAddr,
svc: S,
incoming: I,
signal: Option<F>,
) -> Result<(), super::Error>
where
S: Service<Request<Body>, Response = Response<BoxBody>> + Clone + Send + 'static,
S::Future: Send + 'static,
S::Error: Into<crate::Error> + Send,
I: Stream<Item = Result<IO, IE>>,
IO: AsyncRead + AsyncWrite + Unpin + Send + 'static,
IE: Into<crate::Error>,
F: Future<Output = ()>,
{
let interceptor = self.interceptor.clone();
@@ -262,35 +264,8 @@ impl Server {
let max_concurrent_streams = self.max_concurrent_streams;
// let timeout = self.timeout.clone();
let incoming = hyper::server::accept::from_stream::<_, _, crate::Error>(
async_stream::try_stream! {
let mut incoming = conn::AddrIncoming::bind(&addr)?;
incoming.set_nodelay(self.tcp_nodelay);
incoming.set_keepalive(self.tcp_keepalive);
while let Some(stream) = next_accept(&mut incoming).await? {
#[cfg(feature = "tls")]
{
if let Some(tls) = &self.tls {
let io = match tls.connect(stream.into_inner()).await {
Ok(io) => io,
Err(error) => {
error!(message = "Unable to accept incoming connection.", %error);
continue
},
};
yield BoxedIo::new(io);
continue;
}
}
yield BoxedIo::new(stream);
}
},
);
let tcp = incoming::tcp_incoming(incoming, self);
let incoming = accept::from_stream::<_, _, crate::Error>(tcp);
let svc = MakeSvc {
inner: svc,
@@ -384,8 +359,10 @@ where
///
/// [`Server`]: struct.Server.html
pub async fn serve(self, addr: SocketAddr) -> Result<(), super::Error> {
let incoming = TcpIncoming::new(addr, self.server.tcp_nodelay, self.server.tcp_keepalive)
.map_err(map_err)?;
self.server
.serve_with_shutdown::<_, future::Ready<()>>(addr, self.routes, None)
.serve_with_shutdown::<_, _, future::Ready<()>, _, _>(self.routes, incoming, None)
.await
}
@@ -399,8 +376,25 @@ where
addr: SocketAddr,
f: F,
) -> Result<(), super::Error> {
let incoming = TcpIncoming::new(addr, self.server.tcp_nodelay, self.server.tcp_keepalive)
.map_err(map_err)?;
self.server
.serve_with_shutdown(addr, self.routes, Some(f))
.serve_with_shutdown(self.routes, incoming, Some(f))
.await
}
/// Consume this [`Server`] creating a future that will execute the server on
/// the provided incoming stream of `AsyncRead + AsyncWrite`.
///
/// [`Server`]: struct.Server.html
pub async fn serve_with_incoming<I, IO, IE>(self, incoming: I) -> Result<(), super::Error>
where
I: Stream<Item = Result<IO, IE>>,
IO: AsyncRead + AsyncWrite + Unpin + Send + 'static,
IE: Into<crate::Error>,
{
self.server
.serve_with_shutdown::<_, _, future::Ready<()>, _, _>(self.routes, incoming, None)
.await
}
}
@@ -523,16 +517,3 @@ impl Service<Request<Body>> for Unimplemented {
)
}
}
// Implement try_next for `Accept::poll_accept`.
async fn next_accept(
incoming: &mut conn::AddrIncoming,
) -> Result<Option<conn::AddrStream>, crate::Error> {
let res = poll_fn(|cx| Pin::new(&mut *incoming).poll_accept(cx)).await;
if let Some(res) = res {
Ok(Some(res?))
} else {
return Ok(None);
}
}
+11 -12
View File
@@ -1,6 +1,8 @@
use super::{connector, layer::ServiceBuilderExt, reconnect::Reconnect, AddOrigin};
use super::{layer::ServiceBuilderExt, reconnect::Reconnect, AddOrigin};
use crate::{body::BoxBody, transport::Endpoint};
use http::Uri;
use hyper::client::conn::Builder;
use hyper::client::connect::Connection as HyperConnection;
use hyper::client::service::Connect as HyperConnect;
use std::{
fmt,
@@ -8,6 +10,7 @@ use std::{
pin::Pin,
task::{Context, Poll},
};
use tokio::io::{AsyncRead, AsyncWrite};
use tower::{
layer::Layer,
limit::{concurrency::ConcurrencyLimitLayer, rate::RateLimitLayer},
@@ -26,17 +29,13 @@ pub(crate) struct Connection {
}
impl Connection {
pub(crate) async fn new(endpoint: Endpoint) -> Result<Self, crate::Error> {
#[cfg(feature = "tls")]
let connector = connector(endpoint.tls.clone())
.set_keepalive(endpoint.tcp_keepalive)
.set_nodelay(endpoint.tcp_nodelay);
#[cfg(not(feature = "tls"))]
let connector = connector()
.set_keepalive(endpoint.tcp_keepalive)
.set_nodelay(endpoint.tcp_nodelay);
pub(crate) async fn new<C>(connector: C, endpoint: Endpoint) -> Result<Self, crate::Error>
where
C: Service<Uri> + Send + 'static,
C::Error: Into<crate::Error> + Send,
C::Future: Unpin + Send,
C::Response: AsyncRead + AsyncWrite + HyperConnection + Unpin + Send + 'static,
{
let settings = Builder::new()
.http2_initial_stream_window_size(endpoint.init_stream_window_size)
.http2_initial_connection_window_size(endpoint.init_connection_window_size)
+23 -37
View File
@@ -2,64 +2,50 @@ use super::io::BoxedIo;
#[cfg(feature = "tls")]
use super::tls::TlsConnector;
use http::Uri;
use hyper::client::connect::HttpConnector;
use std::future::Future;
use std::pin::Pin;
use std::task::{Context, Poll};
use std::time::Duration;
use tower_make::MakeConnection;
use tower_service::Service;
#[cfg(not(feature = "tls"))]
pub(crate) fn connector() -> Connector {
Connector::new()
pub(crate) fn connector<C>(inner: C) -> Connector<C> {
Connector::new(inner)
}
#[cfg(feature = "tls")]
pub(crate) fn connector(tls: Option<TlsConnector>) -> Connector {
Connector::new(tls)
pub(crate) fn connector<C>(inner: C, tls: Option<TlsConnector>) -> Connector<C> {
Connector::new(inner, tls)
}
pub(crate) struct Connector {
http: HttpConnector,
pub(crate) struct Connector<C> {
inner: C,
#[cfg(feature = "tls")]
tls: Option<TlsConnector>,
#[cfg(not(feature = "tls"))]
#[allow(dead_code)]
tls: Option<()>,
}
impl Connector {
impl<C> Connector<C> {
#[cfg(not(feature = "tls"))]
pub(crate) fn new() -> Self {
Self {
http: Self::http_connector(),
}
pub(crate) fn new(inner: C) -> Self {
Self { inner, tls: None }
}
#[cfg(feature = "tls")]
fn new(tls: Option<TlsConnector>) -> Self {
Self {
http: Self::http_connector(),
tls,
}
}
pub(crate) fn set_nodelay(mut self, enabled: bool) -> Self {
self.http.set_nodelay(enabled);
self
}
pub(crate) fn set_keepalive(mut self, duration: Option<Duration>) -> Self {
self.http.set_keepalive(duration);
self
}
fn http_connector() -> HttpConnector {
let mut http = HttpConnector::new();
http.enforce_http(false);
http
fn new(inner: C, tls: Option<TlsConnector>) -> Self {
Self { inner, tls }
}
}
impl Service<Uri> for Connector {
impl<C> Service<Uri> for Connector<C>
where
C: MakeConnection<Uri>,
C::Connection: Unpin + Send + 'static,
C::Future: Send + 'static,
crate::Error: From<C::Error> + Send + 'static,
{
type Response = BoxedIo;
type Error = crate::Error;
@@ -67,11 +53,11 @@ impl Service<Uri> for Connector {
Pin<Box<dyn Future<Output = Result<Self::Response, Self::Error>> + Send + 'static>>;
fn poll_ready(&mut self, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
MakeConnection::poll_ready(&mut self.http, cx).map_err(Into::into)
MakeConnection::poll_ready(self, cx).map_err(Into::into)
}
fn call(&mut self, uri: Uri) -> Self::Future {
let connect = MakeConnection::make_connection(&mut self.http, uri);
let connect = self.inner.make_connection(uri);
#[cfg(feature = "tls")]
let tls = self.tls.clone();
+5 -1
View File
@@ -49,7 +49,11 @@ impl Discover for ServiceList {
}
if let Some(endpoint) = self.list.pop_front() {
let fut = Connection::new(endpoint);
let mut http = hyper::client::connect::HttpConnector::new();
http.set_nodelay(endpoint.tcp_nodelay);
http.set_keepalive(endpoint.tcp_keepalive);
let fut = Connection::new(http, endpoint);
self.connecting = Some(Box::pin(fut));
} else {
return Poll::Pending;
+9 -2
View File
@@ -1,14 +1,15 @@
use hyper::client::connect::{Connected, Connection};
use std::io;
use std::pin::Pin;
use std::task::{Context, Poll};
use tokio::io::{AsyncRead, AsyncWrite};
pub(in crate::transport) trait Io:
AsyncRead + AsyncWrite + Send + Unpin + 'static
AsyncRead + AsyncWrite + Send + 'static
{
}
impl<T> Io for T where T: AsyncRead + AsyncWrite + Send + Unpin + 'static {}
impl<T> Io for T where T: AsyncRead + AsyncWrite + Send + 'static {}
pub(crate) struct BoxedIo(Pin<Box<dyn Io>>);
@@ -18,6 +19,12 @@ impl BoxedIo {
}
}
impl Connection for BoxedIo {
fn connected(&self) -> Connected {
Connected::new()
}
}
impl AsyncRead for BoxedIo {
fn poll_read(
mut self: Pin<&mut Self>,
+9 -3
View File
@@ -3,7 +3,7 @@ use crate::transport::{Certificate, Identity};
#[cfg(feature = "tls-roots")]
use rustls_native_certs;
use std::{fmt, sync::Arc};
use tokio::net::TcpStream;
use tokio::io::{AsyncRead, AsyncWrite};
#[cfg(feature = "tls")]
use tokio_rustls::{
rustls::{ClientConfig, NoClientAuth, ServerConfig, Session},
@@ -80,7 +80,10 @@ impl TlsConnector {
})
}
pub(crate) async fn connect(&self, io: TcpStream) -> Result<BoxedIo, crate::Error> {
pub(crate) async fn connect<I>(&self, io: I) -> Result<BoxedIo, crate::Error>
where
I: AsyncRead + AsyncWrite + Send + Unpin + 'static,
{
let tls_io = {
let dns = DNSNameRef::try_from_ascii_str(self.domain.as_str())?.to_owned();
@@ -154,7 +157,10 @@ impl TlsAcceptor {
})
}
pub(crate) async fn connect(&self, io: TcpStream) -> Result<BoxedIo, crate::Error> {
pub(crate) async fn accept<IO>(&self, io: IO) -> Result<BoxedIo, crate::Error>
where
IO: AsyncRead + AsyncWrite + Unpin + Send + 'static,
{
let io = {
let acceptor = RustlsAcceptor::from(self.inner.clone());
let tls = acceptor.accept(io).await?;