Add auth and lb example

This commit is contained in:
Lucio Franco
2019-09-29 18:33:28 -04:00
parent 772ac84409
commit 2c2b7d7904
9 changed files with 332 additions and 6 deletions
+17
View File
@@ -20,6 +20,22 @@ path = "src/routeguide/server.rs"
name = "routeguide-client" name = "routeguide-client"
path = "src/routeguide/client.rs" path = "src/routeguide/client.rs"
[[bin]]
name = "authentication-client"
path = "src/authentication/client.rs"
[[bin]]
name = "authentication-server"
path = "src/authentication/server.rs"
[[bin]]
name = "load-balance-client"
path = "src/load_balance/client.rs"
[[bin]]
name = "load-balance-server"
path = "src/load_balance/server.rs"
[dependencies] [dependencies]
tonic = { path = "../tonic" } tonic = { path = "../tonic" }
futures-preview = { version = "=0.3.0-alpha.18", default-features = false, features = ["alloc"]} futures-preview = { version = "=0.3.0-alpha.18", default-features = false, features = ["alloc"]}
@@ -31,6 +47,7 @@ serde_json = "1.0"
serde = { version = "1.0", features = ["derive"] } serde = { version = "1.0", features = ["derive"] }
async-stream = "0.1.1" async-stream = "0.1.1"
http = "0.1" http = "0.1"
tower = "0.3.0-alpha.1a"
[build-dependencies] [build-dependencies]
tonic-build = { path = "../tonic-build" } tonic-build = { path = "../tonic-build" }
+1
View File
@@ -1,4 +1,5 @@
fn main() { fn main() {
tonic_build::compile_protos("proto/helloworld/helloworld.proto").unwrap(); tonic_build::compile_protos("proto/helloworld/helloworld.proto").unwrap();
tonic_build::compile_protos("proto/routeguide/route_guide.proto").unwrap(); tonic_build::compile_protos("proto/routeguide/route_guide.proto").unwrap();
tonic_build::compile_protos("proto/echo/echo.proto").unwrap();
} }
+43
View File
@@ -0,0 +1,43 @@
/*
*
* Copyright 2018 gRPC authors.
*
* Licensed under the Apache License, Version 2.0 (the "License");
* you may not use this file except in compliance with the License.
* You may obtain a copy of the License at
*
* http://www.apache.org/licenses/LICENSE-2.0
*
* Unless required by applicable law or agreed to in writing, software
* distributed under the License is distributed on an "AS IS" BASIS,
* WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
* See the License for the specific language governing permissions and
* limitations under the License.
*
*/
syntax = "proto3";
package grpc.examples.echo;
// EchoRequest is the request for echo.
message EchoRequest {
string message = 1;
}
// EchoResponse is the response for echo.
message EchoResponse {
string message = 1;
}
// Echo is the echo service.
service Echo {
// UnaryEcho is unary echo.
rpc UnaryEcho(EchoRequest) returns (EchoResponse) {}
// ServerStreamingEcho is server side streaming.
rpc ServerStreamingEcho(EchoRequest) returns (stream EchoResponse) {}
// ClientStreamingEcho is client side streaming.
rpc ClientStreamingEcho(stream EchoRequest) returns (EchoResponse) {}
// BidirectionalStreamingEcho is bidi streaming.
rpc BidirectionalStreamingEcho(stream EchoRequest) returns (stream EchoResponse) {}
}
@@ -0,0 +1,31 @@
pub mod pb {
include!(concat!(env!("OUT_DIR"), "/grpc.examples.echo.rs"));
}
use http::header::HeaderValue;
use pb::{client::EchoClient, EchoRequest};
use tonic::transport::Channel;
#[tokio::main]
async fn main() -> Result<(), Box<dyn std::error::Error>> {
let channel = Channel::from_static("http://[::1]:50051")
.intercept_headers(|headers| {
headers.insert(
"authorization",
HeaderValue::from_static("Bearer some-secret-token"),
);
})
.channel();
let mut client = EchoClient::new(channel);
let request = tonic::Request::new(EchoRequest {
message: "hello".into(),
});
let response = client.unary_echo(request).await?;
println!("RESPONSE={:?}", response);
Ok(())
}
@@ -0,0 +1,86 @@
pub mod pb {
include!(concat!(env!("OUT_DIR"), "/grpc.examples.echo.rs"));
}
use pb::{EchoRequest, EchoResponse};
use std::collections::VecDeque;
use tonic::{body::BoxBody, transport::Server, Request, Response, Status, Streaming};
use tower::Service;
type EchoResult<T> = Result<Response<T>, Status>;
type Stream = VecDeque<Result<EchoResponse, Status>>;
#[derive(Default)]
pub struct EchoServer;
#[tonic::async_trait]
impl pb::server::Echo for EchoServer {
async fn unary_echo(&self, request: Request<EchoRequest>) -> EchoResult<EchoResponse> {
let message = request.into_inner().message;
Ok(Response::new(EchoResponse { message }))
}
type ServerStreamingEchoStream = Stream;
async fn server_streaming_echo(
&self,
_: Request<EchoRequest>,
) -> EchoResult<Self::ServerStreamingEchoStream> {
Err(Status::unimplemented("not implemented"))
}
async fn client_streaming_echo(
&self,
_: Request<Streaming<EchoRequest>>,
) -> EchoResult<EchoResponse> {
Err(Status::unimplemented("not implemented"))
}
type BidirectionalStreamingEchoStream = Stream;
async fn bidirectional_streaming_echo(
&self,
_: Request<Streaming<EchoRequest>>,
) -> EchoResult<Self::BidirectionalStreamingEchoStream> {
Err(Status::unimplemented("not implemented"))
}
}
#[tokio::main]
async fn main() -> Result<(), Box<dyn std::error::Error>> {
let addr = "[::1]:50051".parse().unwrap();
let server = EchoServer::default();
Server::builder()
.interceptor_fn(move |svc, req| {
let auth_header = req.headers().get("authorization").clone();
let authed = if let Some(auth_header) = auth_header {
auth_header == "Bearer some-secret-token"
} else {
false
};
let fut = svc.call(req);
async move {
if authed {
fut.await
} else {
// Cancel the inner future since we never await it
// the IO never gets registered.
drop(fut);
let res = http::Response::builder()
.header("grpc-status", "16")
.body(BoxBody::empty())
.unwrap();
Ok(res)
}
}
})
.clone()
.serve(addr, pb::server::EchoServer::new(server))
.await?;
Ok(())
}
+29
View File
@@ -0,0 +1,29 @@
pub mod pb {
include!(concat!(env!("OUT_DIR"), "/grpc.examples.echo.rs"));
}
use pb::{client::EchoClient, EchoRequest};
use tonic::transport::Channel;
#[tokio::main]
async fn main() -> Result<(), Box<dyn std::error::Error>> {
let endpoints = ["http://[::1]:50051", "http://[::1]:50052"]
.into_iter()
.map(|a| Channel::from_static(a));
let channel = Channel::balance_list(endpoints);
let mut client = EchoClient::new(channel);
for _ in 0..12 {
let request = tonic::Request::new(EchoRequest {
message: "hello".into(),
});
let response = client.unary_echo(request).await?;
println!("RESPONSE={:?}", response);
}
Ok(())
}
+77
View File
@@ -0,0 +1,77 @@
pub mod pb {
include!(concat!(env!("OUT_DIR"), "/grpc.examples.echo.rs"));
}
use pb::{EchoRequest, EchoResponse};
use std::{collections::VecDeque, net::SocketAddr};
use tokio::sync::mpsc;
use tonic::{transport::Server, Request, Response, Status, Streaming};
type EchoResult<T> = Result<Response<T>, Status>;
type Stream = VecDeque<Result<EchoResponse, Status>>;
#[derive(Debug)]
pub struct EchoServer {
addr: SocketAddr,
}
#[tonic::async_trait]
impl pb::server::Echo for EchoServer {
async fn unary_echo(&self, request: Request<EchoRequest>) -> EchoResult<EchoResponse> {
let message = format!("{} (from {})", request.into_inner().message, self.addr);
Ok(Response::new(EchoResponse { message }))
}
type ServerStreamingEchoStream = Stream;
async fn server_streaming_echo(
&self,
_: Request<EchoRequest>,
) -> EchoResult<Self::ServerStreamingEchoStream> {
Err(Status::unimplemented("not implemented"))
}
async fn client_streaming_echo(
&self,
_: Request<Streaming<EchoRequest>>,
) -> EchoResult<EchoResponse> {
Err(Status::unimplemented("not implemented"))
}
type BidirectionalStreamingEchoStream = Stream;
async fn bidirectional_streaming_echo(
&self,
_: Request<Streaming<EchoRequest>>,
) -> EchoResult<Self::BidirectionalStreamingEchoStream> {
Err(Status::unimplemented("not implemented"))
}
}
#[tokio::main]
async fn main() -> Result<(), Box<dyn std::error::Error>> {
let addrs = ["[::1]:50051", "[::1]:50052"];
let (tx, mut rx) = mpsc::unbounded_channel();
for addr in &addrs {
let addr = addr.parse()?;
let mut tx = tx.clone();
let server = EchoServer { addr };
let serve = Server::builder().serve(addr, pb::server::EchoServer::new(server));
tokio::spawn(async move {
if let Err(e) = serve.await {
eprintln!("Error = {:?}", e);
}
tx.try_send(()).unwrap();
});
}
rx.recv().await;
Ok(())
}
+27 -5
View File
@@ -14,6 +14,7 @@ use std::{
fmt, fmt,
future::Future, future::Future,
pin::Pin, pin::Pin,
sync::Arc,
task::{Context, Poll}, task::{Context, Poll},
}; };
use tower::{ use tower::{
@@ -35,6 +36,7 @@ const DEFAULT_BUFFER_SIZE: usize = 1024;
#[derive(Clone)] #[derive(Clone)]
pub struct Channel { pub struct Channel {
svc: Buffer<Svc, Request<BoxBody>>, svc: Buffer<Svc, Request<BoxBody>>,
interceptor_headers: Option<Arc<dyn Fn(&mut http::HeaderMap) + Send + Sync + 'static>>,
} }
/// A future that resolves to an HTTP response. /// A future that resolves to an HTTP response.
@@ -85,22 +87,35 @@ impl Channel {
.and_then(|e| e.buffer_size) .and_then(|e| e.buffer_size)
.unwrap_or(DEFAULT_BUFFER_SIZE); .unwrap_or(DEFAULT_BUFFER_SIZE);
let interceptor_headers = list
.iter()
.next()
.and_then(|e| e.interceptor_headers.clone());
let discover = ServiceList::new(list); let discover = ServiceList::new(list);
Self::balance(discover, buffer_size) Self::balance(discover, buffer_size, interceptor_headers)
} }
pub(crate) fn connect(endpoint: Endpoint) -> Self { pub(crate) fn connect(endpoint: Endpoint) -> Self {
let buffer_size = endpoint.buffer_size.clone().unwrap_or(DEFAULT_BUFFER_SIZE); 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(endpoint);
let svc = Buffer::new(Either::A(svc), buffer_size); let svc = Buffer::new(Either::A(svc), buffer_size);
Channel { svc } Channel {
svc,
interceptor_headers,
}
} }
pub(crate) fn balance<D>(discover: D, buffer_size: usize) -> Self pub(crate) fn balance<D>(
discover: D,
buffer_size: usize,
interceptor_headers: Option<Arc<dyn Fn(&mut http::HeaderMap) + Send + Sync + 'static>>,
) -> Self
where where
D: Discover<Service = Connection> + Unpin + Send + 'static, D: Discover<Service = Connection> + Unpin + Send + 'static,
D::Error: Into<crate::Error>, D::Error: Into<crate::Error>,
@@ -111,7 +126,10 @@ impl Channel {
let svc = BoxService::new(svc); let svc = BoxService::new(svc);
let svc = Buffer::new(Either::B(svc), buffer_size); let svc = Buffer::new(Either::B(svc), buffer_size);
Channel { svc } Channel {
svc,
interceptor_headers,
}
} }
} }
@@ -125,7 +143,11 @@ impl GrpcService<BoxBody> for Channel {
.map_err(|e| super::Error::from_source(super::ErrorKind::Client, e)) .map_err(|e| super::Error::from_source(super::ErrorKind::Client, e))
} }
fn call(&mut self, request: Request<BoxBody>) -> Self::Future { fn call(&mut self, mut request: Request<BoxBody>) -> Self::Future {
if let Some(interceptor) = self.interceptor_headers.clone() {
interceptor(request.headers_mut());
}
let inner = GrpcService::call(&mut self.svc, request); let inner = GrpcService::call(&mut self.svc, request);
ResponseFuture { inner } ResponseFuture { inner }
} }
+21 -1
View File
@@ -5,13 +5,15 @@ use bytes::Bytes;
use http::uri::{InvalidUriBytes, Uri}; use http::uri::{InvalidUriBytes, Uri};
use std::{ use std::{
convert::{TryFrom, TryInto}, convert::{TryFrom, TryInto},
fmt,
sync::Arc,
time::Duration, time::Duration,
}; };
/// Channel builder. /// Channel builder.
/// ///
/// This struct is used to build and configure HTTP/2 channels. /// This struct is used to build and configure HTTP/2 channels.
#[derive(Debug, Clone)] #[derive(Clone)]
pub struct Endpoint { pub struct Endpoint {
pub(super) uri: Uri, pub(super) uri: Uri,
pub(super) timeout: Option<Duration>, pub(super) timeout: Option<Duration>,
@@ -20,6 +22,8 @@ pub struct Endpoint {
#[cfg(feature = "tls")] #[cfg(feature = "tls")]
pub(super) tls: Option<TlsConnector>, pub(super) tls: Option<TlsConnector>,
pub(super) buffer_size: Option<usize>, pub(super) buffer_size: Option<usize>,
pub(super) interceptor_headers:
Option<Arc<dyn Fn(&mut http::HeaderMap) + Send + Sync + 'static>>,
} }
impl Endpoint { impl Endpoint {
@@ -141,6 +145,15 @@ impl Endpoint {
self self
} }
/// Intercept outbound HTTP Request headers;
pub fn intercept_headers<F>(&mut self, f: F) -> &mut Self
where
F: Fn(&mut http::HeaderMap) + Send + Sync + 'static,
{
self.interceptor_headers = Some(Arc::new(f));
self
}
/// Create a channel from this config. /// Create a channel from this config.
pub fn channel(&self) -> Channel { pub fn channel(&self) -> Channel {
Channel::connect(self.clone()) Channel::connect(self.clone())
@@ -157,6 +170,7 @@ impl From<Uri> for Endpoint {
#[cfg(feature = "tls")] #[cfg(feature = "tls")]
tls: None, tls: None,
buffer_size: None, buffer_size: None,
interceptor_headers: None,
} }
} }
} }
@@ -195,3 +209,9 @@ impl std::fmt::Display for Never {
} }
impl std::error::Error for Never {} impl std::error::Error for Never {}
impl fmt::Debug for Endpoint {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("Endpoint").finish()
}
}