feat: Add gRPC interceptors (#232)
This change introduces proper gRPC interceptors that are avilable regardless of the transport used. Each codegen service now produces an additional method called `with_interceptor` that accepts a `Interceptor`. All examples have been updated to use this new style and interop has a custom `tower::Service` middleware to echo the headers. There is also a new `interceptor` example that shows basic usage. BREAKING CHANGE: removed `interceptor_fn` and `intercep_headers_fn` from `transport` in favor of using `tonic::Interceptor`.
This commit is contained in:
+1
-3
@@ -2,13 +2,11 @@
|
|||||||
members = [
|
members = [
|
||||||
"tonic",
|
"tonic",
|
||||||
"tonic-build",
|
"tonic-build",
|
||||||
|
|
||||||
# Non-published crates
|
# Non-published crates
|
||||||
"examples",
|
"examples",
|
||||||
"interop",
|
"interop",
|
||||||
|
|
||||||
# Tests
|
# Tests
|
||||||
"tests/included_service",
|
"tests/included_service",
|
||||||
"tests/same_name",
|
"tests/same_name",
|
||||||
"tests/wellknown",
|
"tests/wellknown",
|
||||||
]
|
]
|
||||||
|
|||||||
+11
-8
@@ -86,27 +86,30 @@ path = "src/uds/client.rs"
|
|||||||
name = "uds-server"
|
name = "uds-server"
|
||||||
path = "src/uds/server.rs"
|
path = "src/uds/server.rs"
|
||||||
|
|
||||||
|
[[bin]]
|
||||||
|
name = "interceptor-client"
|
||||||
|
path = "src/interceptor/client.rs"
|
||||||
|
|
||||||
|
[[bin]]
|
||||||
|
name = "interceptor-server"
|
||||||
|
path = "src/interceptor/server.rs"
|
||||||
|
|
||||||
[dependencies]
|
[dependencies]
|
||||||
tonic = { path = "../tonic", features = ["tls"] }
|
tonic = { path = "../tonic", features = ["tls"] }
|
||||||
prost = "0.6"
|
prost = "0.6"
|
||||||
|
|
||||||
tokio = { version = "0.2", features = ["rt-threaded", "time", "stream", "fs", "macros", "uds"] }
|
tokio = { version = "0.2", features = ["rt-threaded", "time", "stream", "fs", "macros", "uds"] }
|
||||||
futures = { version = "0.3", default-features = false, features = ["alloc"]}
|
futures = { version = "0.3", default-features = false, features = ["alloc"] }
|
||||||
async-stream = "0.2"
|
async-stream = "0.2"
|
||||||
http = "0.2"
|
tower = "0.3"
|
||||||
tower = "0.3"
|
|
||||||
|
|
||||||
# Required for routeguide
|
# Required for routeguide
|
||||||
serde = { version = "1.0", features = ["derive"] }
|
serde = { version = "1.0", features = ["derive"] }
|
||||||
serde_json = "1.0"
|
serde_json = "1.0"
|
||||||
rand = "0.7"
|
rand = "0.7"
|
||||||
|
|
||||||
# Tracing
|
# Tracing
|
||||||
tracing = "0.1"
|
tracing = "0.1"
|
||||||
tracing-subscriber = { version = "0.2.0-alpha", features = ["tracing-log"] }
|
tracing-subscriber = { version = "0.2.0-alpha", features = ["tracing-log"] }
|
||||||
tracing-attributes = "0.1"
|
tracing-attributes = "0.1"
|
||||||
tracing-futures = "0.2"
|
tracing-futures = "0.2"
|
||||||
|
|
||||||
# Required for wellknown types
|
# Required for wellknown types
|
||||||
prost-types = "0.6"
|
prost-types = "0.6"
|
||||||
|
|
||||||
|
|||||||
@@ -2,23 +2,19 @@ pub mod pb {
|
|||||||
tonic::include_proto!("grpc.examples.echo");
|
tonic::include_proto!("grpc.examples.echo");
|
||||||
}
|
}
|
||||||
|
|
||||||
use http::header::HeaderValue;
|
|
||||||
use pb::{echo_client::EchoClient, EchoRequest};
|
use pb::{echo_client::EchoClient, EchoRequest};
|
||||||
use tonic::transport::Channel;
|
use tonic::{metadata::MetadataValue, transport::Channel, Request};
|
||||||
|
|
||||||
#[tokio::main]
|
#[tokio::main]
|
||||||
async fn main() -> Result<(), Box<dyn std::error::Error>> {
|
async fn main() -> Result<(), Box<dyn std::error::Error>> {
|
||||||
let channel = Channel::from_static("http://[::1]:50051")
|
let channel = Channel::from_static("http://[::1]:50051").connect().await?;
|
||||||
.intercept_headers(|headers| {
|
|
||||||
headers.insert(
|
|
||||||
"authorization",
|
|
||||||
HeaderValue::from_static("Bearer some-secret-token"),
|
|
||||||
);
|
|
||||||
})
|
|
||||||
.connect()
|
|
||||||
.await?;
|
|
||||||
|
|
||||||
let mut client = EchoClient::new(channel);
|
let token = MetadataValue::from_str("Bearer some-auth-token")?;
|
||||||
|
|
||||||
|
let mut client = EchoClient::with_interceptor(channel, move |mut req: Request<()>| {
|
||||||
|
req.metadata_mut().insert("authorization", token.clone());
|
||||||
|
Ok(req)
|
||||||
|
});
|
||||||
|
|
||||||
let request = tonic::Request::new(EchoRequest {
|
let request = tonic::Request::new(EchoRequest {
|
||||||
message: "hello".into(),
|
message: "hello".into(),
|
||||||
|
|||||||
@@ -5,8 +5,7 @@ pub mod pb {
|
|||||||
use futures::Stream;
|
use futures::Stream;
|
||||||
use pb::{EchoRequest, EchoResponse};
|
use pb::{EchoRequest, EchoResponse};
|
||||||
use std::pin::Pin;
|
use std::pin::Pin;
|
||||||
use tonic::{body::BoxBody, transport::Server, Request, Response, Status, Streaming};
|
use tonic::{metadata::MetadataValue, transport::Server, Request, Response, Status, Streaming};
|
||||||
use tower::Service;
|
|
||||||
|
|
||||||
type EchoResult<T> = Result<Response<T>, Status>;
|
type EchoResult<T> = Result<Response<T>, Status>;
|
||||||
type ResponseStream = Pin<Box<dyn Stream<Item = Result<EchoResponse, Status>> + Send + Sync>>;
|
type ResponseStream = Pin<Box<dyn Stream<Item = Result<EchoResponse, Status>> + Send + Sync>>;
|
||||||
@@ -52,36 +51,18 @@ async fn main() -> Result<(), Box<dyn std::error::Error>> {
|
|||||||
let addr = "[::1]:50051".parse().unwrap();
|
let addr = "[::1]:50051".parse().unwrap();
|
||||||
let server = EchoServer::default();
|
let server = EchoServer::default();
|
||||||
|
|
||||||
Server::builder()
|
let svc = pb::echo_server::EchoServer::with_interceptor(server, check_auth);
|
||||||
.interceptor_fn(move |svc, req| {
|
|
||||||
let auth_header = req.headers().get("authorization").clone();
|
|
||||||
|
|
||||||
let authed = if let Some(auth_header) = auth_header {
|
Server::builder().add_service(svc).serve(addr).await?;
|
||||||
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)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
})
|
|
||||||
.add_service(pb::echo_server::EchoServer::new(server))
|
|
||||||
.serve(addr)
|
|
||||||
.await?;
|
|
||||||
|
|
||||||
Ok(())
|
Ok(())
|
||||||
}
|
}
|
||||||
|
|
||||||
|
fn check_auth(req: Request<()>) -> Result<Request<()>, Status> {
|
||||||
|
let token = MetadataValue::from_str("Bearer some-secret-token").unwrap();
|
||||||
|
|
||||||
|
match req.metadata().get("authorization") {
|
||||||
|
Some(t) if token == t => Ok(req),
|
||||||
|
_ => Err(Status::unauthenticated("No valid auth token")),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|||||||
@@ -3,8 +3,8 @@ pub mod api {
|
|||||||
}
|
}
|
||||||
|
|
||||||
use api::{publisher_client::PublisherClient, ListTopicsRequest};
|
use api::{publisher_client::PublisherClient, ListTopicsRequest};
|
||||||
use http::header::HeaderValue;
|
|
||||||
use tonic::{
|
use tonic::{
|
||||||
|
metadata::MetadataValue,
|
||||||
transport::{Certificate, Channel, ClientTlsConfig},
|
transport::{Certificate, Channel, ClientTlsConfig},
|
||||||
Request,
|
Request,
|
||||||
};
|
};
|
||||||
@@ -23,7 +23,7 @@ async fn main() -> Result<(), Box<dyn std::error::Error>> {
|
|||||||
.ok_or("Expected a project name as the first argument.".to_string())?;
|
.ok_or("Expected a project name as the first argument.".to_string())?;
|
||||||
|
|
||||||
let bearer_token = format!("Bearer {}", token);
|
let bearer_token = format!("Bearer {}", token);
|
||||||
let header_value = HeaderValue::from_str(&bearer_token)?;
|
let header_value = MetadataValue::from_str(&bearer_token)?;
|
||||||
|
|
||||||
let certs = tokio::fs::read("examples/data/gcp/roots.pem").await?;
|
let certs = tokio::fs::read("examples/data/gcp/roots.pem").await?;
|
||||||
|
|
||||||
@@ -32,14 +32,15 @@ async fn main() -> Result<(), Box<dyn std::error::Error>> {
|
|||||||
.domain_name("pubsub.googleapis.com");
|
.domain_name("pubsub.googleapis.com");
|
||||||
|
|
||||||
let channel = Channel::from_static(ENDPOINT)
|
let channel = Channel::from_static(ENDPOINT)
|
||||||
.intercept_headers(move |headers| {
|
|
||||||
headers.insert("authorization", header_value.clone());
|
|
||||||
})
|
|
||||||
.tls_config(tls_config)
|
.tls_config(tls_config)
|
||||||
.connect()
|
.connect()
|
||||||
.await?;
|
.await?;
|
||||||
|
|
||||||
let mut service = PublisherClient::new(channel);
|
let mut service = PublisherClient::with_interceptor(channel, move |mut req: Request<()>| {
|
||||||
|
req.metadata_mut()
|
||||||
|
.insert("authorization", header_value.clone());
|
||||||
|
Ok(req)
|
||||||
|
});
|
||||||
|
|
||||||
let response = service
|
let response = service
|
||||||
.list_topics(Request::new(ListTopicsRequest {
|
.list_topics(Request::new(ListTopicsRequest {
|
||||||
|
|||||||
@@ -0,0 +1,34 @@
|
|||||||
|
use hello_world::greeter_client::GreeterClient;
|
||||||
|
use hello_world::HelloRequest;
|
||||||
|
use tonic::{transport::Endpoint, Request, Status};
|
||||||
|
|
||||||
|
pub mod hello_world {
|
||||||
|
tonic::include_proto!("helloworld");
|
||||||
|
}
|
||||||
|
|
||||||
|
#[tokio::main]
|
||||||
|
async fn main() -> Result<(), Box<dyn std::error::Error>> {
|
||||||
|
let channel = Endpoint::from_static("http://[::1]:50051")
|
||||||
|
.connect()
|
||||||
|
.await?;
|
||||||
|
|
||||||
|
let mut client = GreeterClient::with_interceptor(channel, intercept);
|
||||||
|
|
||||||
|
let request = tonic::Request::new(HelloRequest {
|
||||||
|
name: "Tonic".into(),
|
||||||
|
});
|
||||||
|
|
||||||
|
let response = client.say_hello(request).await?;
|
||||||
|
|
||||||
|
println!("RESPONSE={:?}", response);
|
||||||
|
|
||||||
|
Ok(())
|
||||||
|
}
|
||||||
|
|
||||||
|
/// This function will get called on each outbound request. Returning a
|
||||||
|
/// `Status` here will cancel the request and have that status returned to
|
||||||
|
/// the caller.
|
||||||
|
fn intercept(req: Request<()>) -> Result<Request<()>, Status> {
|
||||||
|
println!("Intercepting request: {:?}", req);
|
||||||
|
Ok(req)
|
||||||
|
}
|
||||||
@@ -0,0 +1,46 @@
|
|||||||
|
use tonic::{transport::Server, Request, Response, Status};
|
||||||
|
|
||||||
|
use hello_world::greeter_server::{Greeter, GreeterServer};
|
||||||
|
use hello_world::{HelloReply, HelloRequest};
|
||||||
|
|
||||||
|
pub mod hello_world {
|
||||||
|
tonic::include_proto!("helloworld");
|
||||||
|
}
|
||||||
|
|
||||||
|
#[derive(Default)]
|
||||||
|
pub struct MyGreeter {}
|
||||||
|
|
||||||
|
#[tonic::async_trait]
|
||||||
|
impl Greeter for MyGreeter {
|
||||||
|
async fn say_hello(
|
||||||
|
&self,
|
||||||
|
request: Request<HelloRequest>,
|
||||||
|
) -> Result<Response<HelloReply>, Status> {
|
||||||
|
let reply = hello_world::HelloReply {
|
||||||
|
message: format!("Hello {}!", request.into_inner().name),
|
||||||
|
};
|
||||||
|
Ok(Response::new(reply))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
#[tokio::main]
|
||||||
|
async fn main() -> Result<(), Box<dyn std::error::Error>> {
|
||||||
|
let addr = "[::1]:50051".parse().unwrap();
|
||||||
|
let greeter = MyGreeter::default();
|
||||||
|
|
||||||
|
let svc = GreeterServer::with_interceptor(greeter, intercept);
|
||||||
|
|
||||||
|
println!("GreeterServer listening on {}", addr);
|
||||||
|
|
||||||
|
Server::builder().add_service(svc).serve(addr).await?;
|
||||||
|
|
||||||
|
Ok(())
|
||||||
|
}
|
||||||
|
|
||||||
|
/// This function will get called on each inbound request, if a `Status`
|
||||||
|
/// is returned, it will cancel the request and return that status to the
|
||||||
|
/// client.
|
||||||
|
fn intercept(req: Request<()>) -> Result<Request<()>, Status> {
|
||||||
|
println!("Intercepting request: {:?}", req);
|
||||||
|
Ok(req)
|
||||||
|
}
|
||||||
@@ -5,11 +5,10 @@ pub mod hello_world {
|
|||||||
}
|
}
|
||||||
|
|
||||||
use hello_world::{greeter_client::GreeterClient, HelloRequest};
|
use hello_world::{greeter_client::GreeterClient, HelloRequest};
|
||||||
use http::Uri;
|
|
||||||
use std::convert::TryFrom;
|
use std::convert::TryFrom;
|
||||||
#[cfg(unix)]
|
#[cfg(unix)]
|
||||||
use tokio::net::UnixStream;
|
use tokio::net::UnixStream;
|
||||||
use tonic::transport::Endpoint;
|
use tonic::transport::{Endpoint, Uri};
|
||||||
use tower::service_fn;
|
use tower::service_fn;
|
||||||
|
|
||||||
#[cfg(unix)]
|
#[cfg(unix)]
|
||||||
|
|||||||
+1
-2
@@ -26,10 +26,9 @@ futures-util = "0.3"
|
|||||||
async-stream = "0.2"
|
async-stream = "0.2"
|
||||||
tower = "0.3"
|
tower = "0.3"
|
||||||
http-body = "0.3"
|
http-body = "0.3"
|
||||||
|
hyper = "0.13"
|
||||||
console = "0.9"
|
console = "0.9"
|
||||||
structopt = "0.3"
|
structopt = "0.3"
|
||||||
|
|
||||||
tracing = "0.1"
|
tracing = "0.1"
|
||||||
tracing-subscriber = "0.2.0-alpha"
|
tracing-subscriber = "0.2.0-alpha"
|
||||||
tracing-log = "0.1.0"
|
tracing-log = "0.1.0"
|
||||||
|
|||||||
@@ -1,10 +1,7 @@
|
|||||||
use http::header::HeaderName;
|
|
||||||
use structopt::StructOpt;
|
use structopt::StructOpt;
|
||||||
use tonic::body::BoxBody;
|
|
||||||
use tonic::client::GrpcService;
|
|
||||||
use tonic::transport::Server;
|
use tonic::transport::Server;
|
||||||
use tonic::transport::{Identity, ServerTlsConfig};
|
use tonic::transport::{Identity, ServerTlsConfig};
|
||||||
use tonic_interop::{server, MergeTrailers};
|
use tonic_interop::server;
|
||||||
|
|
||||||
#[derive(StructOpt)]
|
#[derive(StructOpt)]
|
||||||
struct Opts {
|
struct Opts {
|
||||||
@@ -20,33 +17,7 @@ async fn main() -> std::result::Result<(), Box<dyn std::error::Error>> {
|
|||||||
|
|
||||||
let addr = "127.0.0.1:10000".parse().unwrap();
|
let addr = "127.0.0.1:10000".parse().unwrap();
|
||||||
|
|
||||||
let mut builder = Server::builder().interceptor_fn(|svc, req| {
|
let mut builder = Server::builder();
|
||||||
let echo_header = req
|
|
||||||
.headers()
|
|
||||||
.get("x-grpc-test-echo-initial")
|
|
||||||
.map(Clone::clone);
|
|
||||||
|
|
||||||
let echo_trailer = req
|
|
||||||
.headers()
|
|
||||||
.get("x-grpc-test-echo-trailing-bin")
|
|
||||||
.map(Clone::clone)
|
|
||||||
.map(|v| (HeaderName::from_static("x-grpc-test-echo-trailing-bin"), v));
|
|
||||||
|
|
||||||
let call = svc.call(req);
|
|
||||||
|
|
||||||
async move {
|
|
||||||
let mut res = call.await?;
|
|
||||||
|
|
||||||
if let Some(echo_header) = echo_header {
|
|
||||||
res.headers_mut()
|
|
||||||
.insert("x-grpc-test-echo-initial", echo_header);
|
|
||||||
}
|
|
||||||
|
|
||||||
Ok(res
|
|
||||||
.map(|b| MergeTrailers::new(b, echo_trailer))
|
|
||||||
.map(BoxBody::new))
|
|
||||||
}
|
|
||||||
});
|
|
||||||
|
|
||||||
if matches.use_tls {
|
if matches.use_tls {
|
||||||
let cert = tokio::fs::read("interop/data/server1.pem").await?;
|
let cert = tokio::fs::read("interop/data/server1.pem").await?;
|
||||||
@@ -60,8 +31,11 @@ async fn main() -> std::result::Result<(), Box<dyn std::error::Error>> {
|
|||||||
let unimplemented_service =
|
let unimplemented_service =
|
||||||
server::UnimplementedServiceServer::new(server::UnimplementedService::default());
|
server::UnimplementedServiceServer::new(server::UnimplementedService::default());
|
||||||
|
|
||||||
|
// Wrap this test_service with a service that will echo headers as trailers.
|
||||||
|
let test_service_svc = server::EchoHeadersSvc::new(test_service);
|
||||||
|
|
||||||
builder
|
builder
|
||||||
.add_service(test_service)
|
.add_service(test_service_svc)
|
||||||
.add_service(unimplemented_service)
|
.add_service(unimplemented_service)
|
||||||
.serve(addr)
|
.serve(addr)
|
||||||
.await?;
|
.await?;
|
||||||
|
|||||||
+1
-45
@@ -9,13 +9,7 @@ pub mod pb {
|
|||||||
include!(concat!(env!("OUT_DIR"), "/grpc.testing.rs"));
|
include!(concat!(env!("OUT_DIR"), "/grpc.testing.rs"));
|
||||||
}
|
}
|
||||||
|
|
||||||
use http::header::{HeaderMap, HeaderName, HeaderValue};
|
use std::{default, fmt, iter};
|
||||||
use http_body::Body;
|
|
||||||
use std::{
|
|
||||||
default, fmt, iter,
|
|
||||||
pin::Pin,
|
|
||||||
task::{Context, Poll},
|
|
||||||
};
|
|
||||||
|
|
||||||
pub fn trace_init() {
|
pub fn trace_init() {
|
||||||
let sub = tracing_subscriber::FmtSubscriber::builder()
|
let sub = tracing_subscriber::FmtSubscriber::builder()
|
||||||
@@ -147,41 +141,3 @@ macro_rules! test_assert {
|
|||||||
}
|
}
|
||||||
};
|
};
|
||||||
}
|
}
|
||||||
|
|
||||||
pub struct MergeTrailers<B> {
|
|
||||||
inner: B,
|
|
||||||
trailer: Option<(HeaderName, HeaderValue)>,
|
|
||||||
}
|
|
||||||
|
|
||||||
impl<B> MergeTrailers<B> {
|
|
||||||
pub fn new(inner: B, trailer: Option<(HeaderName, HeaderValue)>) -> Self {
|
|
||||||
Self { inner, trailer }
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
impl<B: Body + Unpin> Body for MergeTrailers<B> {
|
|
||||||
type Data = B::Data;
|
|
||||||
type Error = B::Error;
|
|
||||||
|
|
||||||
fn poll_data(
|
|
||||||
mut self: Pin<&mut Self>,
|
|
||||||
cx: &mut Context<'_>,
|
|
||||||
) -> Poll<Option<Result<Self::Data, Self::Error>>> {
|
|
||||||
Pin::new(&mut self.inner).poll_data(cx)
|
|
||||||
}
|
|
||||||
|
|
||||||
fn poll_trailers(
|
|
||||||
mut self: Pin<&mut Self>,
|
|
||||||
cx: &mut Context<'_>,
|
|
||||||
) -> Poll<Result<Option<HeaderMap>, Self::Error>> {
|
|
||||||
Pin::new(&mut self.inner).poll_trailers(cx).map_ok(|h| {
|
|
||||||
h.map(|mut headers| {
|
|
||||||
if let Some((key, value)) = &self.trailer {
|
|
||||||
headers.insert(key.clone(), value.clone());
|
|
||||||
}
|
|
||||||
|
|
||||||
headers
|
|
||||||
})
|
|
||||||
})
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|||||||
+103
-1
@@ -1,9 +1,14 @@
|
|||||||
use crate::pb::{self, *};
|
use crate::pb::{self, *};
|
||||||
use async_stream::try_stream;
|
use async_stream::try_stream;
|
||||||
use futures_util::{stream, StreamExt, TryStreamExt};
|
use futures_util::{stream, StreamExt, TryStreamExt};
|
||||||
|
use http::header::{HeaderMap, HeaderName, HeaderValue};
|
||||||
|
use http_body::Body;
|
||||||
|
use std::future::Future;
|
||||||
use std::pin::Pin;
|
use std::pin::Pin;
|
||||||
|
use std::task::{Context, Poll};
|
||||||
use std::time::Duration;
|
use std::time::Duration;
|
||||||
use tonic::{Code, Request, Response, Status};
|
use tonic::{body::BoxBody, transport::ServiceName, Code, Request, Response, Status};
|
||||||
|
use tower::Service;
|
||||||
|
|
||||||
pub use pb::test_service_server::TestServiceServer;
|
pub use pb::test_service_server::TestServiceServer;
|
||||||
pub use pb::unimplemented_service_server::UnimplementedServiceServer;
|
pub use pb::unimplemented_service_server::UnimplementedServiceServer;
|
||||||
@@ -159,3 +164,100 @@ impl pb::unimplemented_service_server::UnimplementedService for UnimplementedSer
|
|||||||
Err(Status::unimplemented(""))
|
Err(Status::unimplemented(""))
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
#[derive(Clone, Default)]
|
||||||
|
pub struct EchoHeadersSvc<S> {
|
||||||
|
inner: S,
|
||||||
|
}
|
||||||
|
|
||||||
|
impl<S: ServiceName> ServiceName for EchoHeadersSvc<S> {
|
||||||
|
const NAME: &'static str = S::NAME;
|
||||||
|
}
|
||||||
|
|
||||||
|
impl<S> EchoHeadersSvc<S> {
|
||||||
|
pub fn new(inner: S) -> Self {
|
||||||
|
Self { inner }
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
impl<S> Service<http::Request<hyper::Body>> for EchoHeadersSvc<S>
|
||||||
|
where
|
||||||
|
S: Service<http::Request<hyper::Body>, Response = http::Response<BoxBody>> + Send,
|
||||||
|
S::Future: Send + 'static,
|
||||||
|
{
|
||||||
|
type Response = S::Response;
|
||||||
|
type Error = S::Error;
|
||||||
|
type Future = Pin<
|
||||||
|
Box<dyn Future<Output = std::result::Result<Self::Response, Self::Error>> + Send + 'static>,
|
||||||
|
>;
|
||||||
|
|
||||||
|
fn poll_ready(&mut self, _cx: &mut Context<'_>) -> Poll<std::result::Result<(), Self::Error>> {
|
||||||
|
Ok(()).into()
|
||||||
|
}
|
||||||
|
|
||||||
|
fn call(&mut self, req: http::Request<hyper::Body>) -> Self::Future {
|
||||||
|
let echo_header = req
|
||||||
|
.headers()
|
||||||
|
.get("x-grpc-test-echo-initial")
|
||||||
|
.map(Clone::clone);
|
||||||
|
|
||||||
|
let echo_trailer = req
|
||||||
|
.headers()
|
||||||
|
.get("x-grpc-test-echo-trailing-bin")
|
||||||
|
.map(Clone::clone)
|
||||||
|
.map(|v| (HeaderName::from_static("x-grpc-test-echo-trailing-bin"), v));
|
||||||
|
|
||||||
|
let call = self.inner.call(req);
|
||||||
|
|
||||||
|
Box::pin(async move {
|
||||||
|
let mut res = call.await?;
|
||||||
|
|
||||||
|
if let Some(echo_header) = echo_header {
|
||||||
|
res.headers_mut()
|
||||||
|
.insert("x-grpc-test-echo-initial", echo_header);
|
||||||
|
}
|
||||||
|
|
||||||
|
Ok(res
|
||||||
|
.map(|b| MergeTrailers::new(b, echo_trailer))
|
||||||
|
.map(BoxBody::new))
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
pub struct MergeTrailers<B> {
|
||||||
|
inner: B,
|
||||||
|
trailer: Option<(HeaderName, HeaderValue)>,
|
||||||
|
}
|
||||||
|
|
||||||
|
impl<B> MergeTrailers<B> {
|
||||||
|
pub fn new(inner: B, trailer: Option<(HeaderName, HeaderValue)>) -> Self {
|
||||||
|
Self { inner, trailer }
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
impl<B: Body + Unpin> Body for MergeTrailers<B> {
|
||||||
|
type Data = B::Data;
|
||||||
|
type Error = B::Error;
|
||||||
|
|
||||||
|
fn poll_data(
|
||||||
|
mut self: Pin<&mut Self>,
|
||||||
|
cx: &mut Context<'_>,
|
||||||
|
) -> Poll<Option<std::result::Result<Self::Data, Self::Error>>> {
|
||||||
|
Pin::new(&mut self.inner).poll_data(cx)
|
||||||
|
}
|
||||||
|
|
||||||
|
fn poll_trailers(
|
||||||
|
mut self: Pin<&mut Self>,
|
||||||
|
cx: &mut Context<'_>,
|
||||||
|
) -> Poll<std::result::Result<Option<HeaderMap>, Self::Error>> {
|
||||||
|
Pin::new(&mut self.inner).poll_trailers(cx).map_ok(|h| {
|
||||||
|
h.map(|mut headers| {
|
||||||
|
if let Some((key, value)) = &self.trailer {
|
||||||
|
headers.insert(key.clone(), value.clone());
|
||||||
|
}
|
||||||
|
|
||||||
|
headers
|
||||||
|
})
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|||||||
@@ -34,6 +34,11 @@ pub(crate) fn generate(service: &Service, proto: &str) -> TokenStream {
|
|||||||
Self { inner }
|
Self { inner }
|
||||||
}
|
}
|
||||||
|
|
||||||
|
pub fn with_interceptor(inner: T, interceptor: impl Into<tonic::Interceptor>) -> Self {
|
||||||
|
let inner = tonic::client::Grpc::with_interceptor(inner, interceptor);
|
||||||
|
Self { inner }
|
||||||
|
}
|
||||||
|
|
||||||
#methods
|
#methods
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -29,12 +29,21 @@ pub(crate) fn generate(service: &Service, proto_path: &str) -> TokenStream {
|
|||||||
#[derive(Debug)]
|
#[derive(Debug)]
|
||||||
#[doc(hidden)]
|
#[doc(hidden)]
|
||||||
pub struct #server_service<T: #server_trait> {
|
pub struct #server_service<T: #server_trait> {
|
||||||
inner: Arc<T>,
|
inner: _Inner<T>,
|
||||||
}
|
}
|
||||||
|
|
||||||
|
struct _Inner<T>(Arc<T>, Option<tonic::Interceptor>);
|
||||||
|
|
||||||
impl<T: #server_trait> #server_service<T> {
|
impl<T: #server_trait> #server_service<T> {
|
||||||
pub fn new(inner: T) -> Self {
|
pub fn new(inner: T) -> Self {
|
||||||
let inner = Arc::new(inner);
|
let inner = Arc::new(inner);
|
||||||
|
let inner = _Inner(inner, None);
|
||||||
|
Self { inner }
|
||||||
|
}
|
||||||
|
|
||||||
|
pub fn with_interceptor(inner: T, interceptor: impl Into<tonic::Interceptor>) -> Self {
|
||||||
|
let inner = Arc::new(inner);
|
||||||
|
let inner = _Inner(inner, Some(interceptor.into()));
|
||||||
Self { inner }
|
Self { inner }
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -72,6 +81,18 @@ pub(crate) fn generate(service: &Service, proto_path: &str) -> TokenStream {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
impl<T: #server_trait> Clone for _Inner<T> {
|
||||||
|
fn clone(&self) -> Self {
|
||||||
|
Self(self.0.clone(), self.1.clone())
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
impl<T: std::fmt::Debug> std::fmt::Debug for _Inner<T> {
|
||||||
|
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||||
|
write!(f, "{:?}", self.0)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
#transport
|
#transport
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -246,9 +267,17 @@ fn generate_unary(
|
|||||||
|
|
||||||
let inner = self.inner.clone();
|
let inner = self.inner.clone();
|
||||||
let fut = async move {
|
let fut = async move {
|
||||||
|
let interceptor = inner.1.clone();
|
||||||
|
let inner = inner.0;
|
||||||
let method = #service_ident(inner);
|
let method = #service_ident(inner);
|
||||||
let codec = tonic::codec::ProstCodec::default();
|
let codec = tonic::codec::ProstCodec::default();
|
||||||
let mut grpc = tonic::server::Grpc::new(codec);
|
|
||||||
|
let mut grpc = if let Some(interceptor) = interceptor {
|
||||||
|
tonic::server::Grpc::with_interceptor(codec, interceptor)
|
||||||
|
} else {
|
||||||
|
tonic::server::Grpc::new(codec)
|
||||||
|
};
|
||||||
|
|
||||||
let res = grpc.unary(method, req).await;
|
let res = grpc.unary(method, req).await;
|
||||||
Ok(res)
|
Ok(res)
|
||||||
};
|
};
|
||||||
@@ -289,9 +318,17 @@ fn generate_server_streaming(
|
|||||||
|
|
||||||
let inner = self.inner.clone();
|
let inner = self.inner.clone();
|
||||||
let fut = async move {
|
let fut = async move {
|
||||||
|
let interceptor = inner.1;
|
||||||
|
let inner = inner.0;
|
||||||
let method = #service_ident(inner);
|
let method = #service_ident(inner);
|
||||||
let codec = tonic::codec::ProstCodec::default();
|
let codec = tonic::codec::ProstCodec::default();
|
||||||
let mut grpc = tonic::server::Grpc::new(codec);
|
|
||||||
|
let mut grpc = if let Some(interceptor) = interceptor {
|
||||||
|
tonic::server::Grpc::with_interceptor(codec, interceptor)
|
||||||
|
} else {
|
||||||
|
tonic::server::Grpc::new(codec)
|
||||||
|
};
|
||||||
|
|
||||||
let res = grpc.server_streaming(method, req).await;
|
let res = grpc.server_streaming(method, req).await;
|
||||||
Ok(res)
|
Ok(res)
|
||||||
};
|
};
|
||||||
@@ -330,9 +367,17 @@ fn generate_client_streaming(
|
|||||||
|
|
||||||
let inner = self.inner.clone();
|
let inner = self.inner.clone();
|
||||||
let fut = async move {
|
let fut = async move {
|
||||||
|
let interceptor = inner.1;
|
||||||
|
let inner = inner.0;
|
||||||
let method = #service_ident(inner);
|
let method = #service_ident(inner);
|
||||||
let codec = tonic::codec::ProstCodec::default();
|
let codec = tonic::codec::ProstCodec::default();
|
||||||
let mut grpc = tonic::server::Grpc::new(codec);
|
|
||||||
|
let mut grpc = if let Some(interceptor) = interceptor {
|
||||||
|
tonic::server::Grpc::with_interceptor(codec, interceptor)
|
||||||
|
} else {
|
||||||
|
tonic::server::Grpc::new(codec)
|
||||||
|
};
|
||||||
|
|
||||||
let res = grpc.client_streaming(method, req).await;
|
let res = grpc.client_streaming(method, req).await;
|
||||||
Ok(res)
|
Ok(res)
|
||||||
};
|
};
|
||||||
@@ -373,9 +418,17 @@ fn generate_streaming(
|
|||||||
|
|
||||||
let inner = self.inner.clone();
|
let inner = self.inner.clone();
|
||||||
let fut = async move {
|
let fut = async move {
|
||||||
|
let interceptor = inner.1;
|
||||||
|
let inner = inner.0;
|
||||||
let method = #service_ident(inner);
|
let method = #service_ident(inner);
|
||||||
let codec = tonic::codec::ProstCodec::default();
|
let codec = tonic::codec::ProstCodec::default();
|
||||||
let mut grpc = tonic::server::Grpc::new(codec);
|
|
||||||
|
let mut grpc = if let Some(interceptor) = interceptor {
|
||||||
|
tonic::server::Grpc::with_interceptor(codec, interceptor)
|
||||||
|
} else {
|
||||||
|
tonic::server::Grpc::new(codec)
|
||||||
|
};
|
||||||
|
|
||||||
let res = grpc.streaming(method, req).await;
|
let res = grpc.streaming(method, req).await;
|
||||||
Ok(res)
|
Ok(res)
|
||||||
};
|
};
|
||||||
|
|||||||
@@ -2,6 +2,7 @@ use crate::{
|
|||||||
body::{Body, BoxBody},
|
body::{Body, BoxBody},
|
||||||
client::GrpcService,
|
client::GrpcService,
|
||||||
codec::{encode_client, Codec, Streaming},
|
codec::{encode_client, Codec, Streaming},
|
||||||
|
interceptor::Interceptor,
|
||||||
Code, Request, Response, Status,
|
Code, Request, Response, Status,
|
||||||
};
|
};
|
||||||
use futures_core::Stream;
|
use futures_core::Stream;
|
||||||
@@ -28,12 +29,25 @@ use std::fmt;
|
|||||||
/// [gRPC protocol definition]: https://github.com/grpc/grpc/blob/master/doc/PROTOCOL-HTTP2.md#requests
|
/// [gRPC protocol definition]: https://github.com/grpc/grpc/blob/master/doc/PROTOCOL-HTTP2.md#requests
|
||||||
pub struct Grpc<T> {
|
pub struct Grpc<T> {
|
||||||
inner: T,
|
inner: T,
|
||||||
|
interceptor: Option<Interceptor>,
|
||||||
}
|
}
|
||||||
|
|
||||||
impl<T> Grpc<T> {
|
impl<T> Grpc<T> {
|
||||||
/// Creates a new gRPC client with the provided [`GrpcService`].
|
/// Creates a new gRPC client with the provided [`GrpcService`].
|
||||||
pub fn new(inner: T) -> Self {
|
pub fn new(inner: T) -> Self {
|
||||||
Self { inner }
|
Self {
|
||||||
|
inner,
|
||||||
|
interceptor: None,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Creates a new gRPC client with the provided [`GrpcService`] and will apply
|
||||||
|
/// the provided interceptor on each request.
|
||||||
|
pub fn with_interceptor(inner: T, interceptor: impl Into<Interceptor>) -> Self {
|
||||||
|
Self {
|
||||||
|
inner,
|
||||||
|
interceptor: Some(interceptor.into()),
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
/// Check if the inner [`GrpcService`] is able to accept a new request.
|
/// Check if the inner [`GrpcService`] is able to accept a new request.
|
||||||
@@ -134,6 +148,12 @@ impl<T> Grpc<T> {
|
|||||||
M1: Send + Sync + 'static,
|
M1: Send + Sync + 'static,
|
||||||
M2: Send + Sync + 'static,
|
M2: Send + Sync + 'static,
|
||||||
{
|
{
|
||||||
|
let request = if let Some(interceptor) = &self.interceptor {
|
||||||
|
interceptor.call(request)?
|
||||||
|
} else {
|
||||||
|
request
|
||||||
|
};
|
||||||
|
|
||||||
let mut parts = Parts::default();
|
let mut parts = Parts::default();
|
||||||
parts.path_and_query = Some(path);
|
parts.path_and_query = Some(path);
|
||||||
|
|
||||||
@@ -192,6 +212,7 @@ impl<T: Clone> Clone for Grpc<T> {
|
|||||||
fn clone(&self) -> Self {
|
fn clone(&self) -> Self {
|
||||||
Self {
|
Self {
|
||||||
inner: self.inner.clone(),
|
inner: self.inner.clone(),
|
||||||
|
interceptor: self.interceptor.clone(),
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -0,0 +1,54 @@
|
|||||||
|
use crate::{Request, Status};
|
||||||
|
use std::{fmt, sync::Arc};
|
||||||
|
|
||||||
|
/// Represents a gRPC interceptor.
|
||||||
|
///
|
||||||
|
/// gRPC interceptors are similar to middleware but have much less
|
||||||
|
/// flexibility. This interceptor allows you to do two main things,
|
||||||
|
/// one is to add/remove/check items in the `MetadataMap` of each
|
||||||
|
/// request. Two, cancel a request with any `Status`.
|
||||||
|
///
|
||||||
|
/// An interceptor can be used on both the server and client side through
|
||||||
|
/// the `tonic-build` crate's generated structs.
|
||||||
|
///
|
||||||
|
/// These interceptors do not allow you to modify the `Message` of the request
|
||||||
|
/// but allow you to check for metadata. If you would like to apply middleware like
|
||||||
|
/// features to the body of the request, going through the `tower` abstraction is recommended.
|
||||||
|
#[derive(Clone)]
|
||||||
|
pub struct Interceptor {
|
||||||
|
f: Arc<dyn Fn(Request<()>) -> Result<Request<()>, Status> + Send + Sync + 'static>,
|
||||||
|
}
|
||||||
|
|
||||||
|
impl Interceptor {
|
||||||
|
/// Create a new `Interceptor` from the provided function.
|
||||||
|
pub fn new(
|
||||||
|
f: impl Fn(Request<()>) -> Result<Request<()>, Status> + Send + Sync + 'static,
|
||||||
|
) -> Self {
|
||||||
|
Interceptor { f: Arc::new(f) }
|
||||||
|
}
|
||||||
|
|
||||||
|
pub(crate) fn call<T>(&self, req: Request<T>) -> Result<Request<T>, Status> {
|
||||||
|
let (metadata, ext, message) = req.into_parts();
|
||||||
|
|
||||||
|
let temp_req = Request::from_parts(metadata, ext, ());
|
||||||
|
|
||||||
|
let (metadata, ext, _) = (self.f)(temp_req)?.into_parts();
|
||||||
|
|
||||||
|
Ok(Request::from_parts(metadata, ext, message))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
impl<F> From<F> for Interceptor
|
||||||
|
where
|
||||||
|
F: Fn(Request<()>) -> Result<Request<()>, Status> + Send + Sync + 'static,
|
||||||
|
{
|
||||||
|
fn from(f: F) -> Self {
|
||||||
|
Interceptor::new(f)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
impl fmt::Debug for Interceptor {
|
||||||
|
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
|
||||||
|
f.debug_struct("Interceptor").finish()
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -86,6 +86,7 @@ pub mod server;
|
|||||||
#[cfg_attr(docsrs, doc(cfg(feature = "transport")))]
|
#[cfg_attr(docsrs, doc(cfg(feature = "transport")))]
|
||||||
pub mod transport;
|
pub mod transport;
|
||||||
|
|
||||||
|
mod interceptor;
|
||||||
mod macros;
|
mod macros;
|
||||||
mod request;
|
mod request;
|
||||||
mod response;
|
mod response;
|
||||||
@@ -98,6 +99,7 @@ pub use async_trait::async_trait;
|
|||||||
|
|
||||||
#[doc(inline)]
|
#[doc(inline)]
|
||||||
pub use codec::Streaming;
|
pub use codec::Streaming;
|
||||||
|
pub use interceptor::Interceptor;
|
||||||
pub use request::{IntoRequest, IntoStreamingRequest, Request};
|
pub use request::{IntoRequest, IntoStreamingRequest, Request};
|
||||||
pub use response::Response;
|
pub use response::Response;
|
||||||
pub use status::{Code, Status};
|
pub use status::{Code, Status};
|
||||||
|
|||||||
@@ -145,6 +145,18 @@ impl<T> Request<T> {
|
|||||||
self.message
|
self.message
|
||||||
}
|
}
|
||||||
|
|
||||||
|
pub(crate) fn into_parts(self) -> (MetadataMap, Extensions, T) {
|
||||||
|
(self.metadata, self.extensions, self.message)
|
||||||
|
}
|
||||||
|
|
||||||
|
pub(crate) fn from_parts(metadata: MetadataMap, extensions: Extensions, message: T) -> Self {
|
||||||
|
Self {
|
||||||
|
metadata,
|
||||||
|
extensions,
|
||||||
|
message,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
pub(crate) fn from_http_parts(parts: http::request::Parts, message: T) -> Self {
|
pub(crate) fn from_http_parts(parts: http::request::Parts, message: T) -> Self {
|
||||||
Request {
|
Request {
|
||||||
metadata: MetadataMap::from_headers(parts.headers),
|
metadata: MetadataMap::from_headers(parts.headers),
|
||||||
|
|||||||
+56
-10
@@ -1,6 +1,7 @@
|
|||||||
use crate::{
|
use crate::{
|
||||||
body::BoxBody,
|
body::BoxBody,
|
||||||
codec::{encode_server, Codec, Streaming},
|
codec::{encode_server, Codec, Streaming},
|
||||||
|
interceptor::Interceptor,
|
||||||
server::{ClientStreamingService, ServerStreamingService, StreamingService, UnaryService},
|
server::{ClientStreamingService, ServerStreamingService, StreamingService, UnaryService},
|
||||||
Code, Request, Response, Status,
|
Code, Request, Response, Status,
|
||||||
};
|
};
|
||||||
@@ -9,6 +10,16 @@ use futures_util::{future, stream, TryStreamExt};
|
|||||||
use http_body::Body;
|
use http_body::Body;
|
||||||
use std::fmt;
|
use std::fmt;
|
||||||
|
|
||||||
|
// A try! type macro for intercepting requests
|
||||||
|
macro_rules! t {
|
||||||
|
($expr : expr) => {
|
||||||
|
match $expr {
|
||||||
|
Ok(request) => request,
|
||||||
|
Err(res) => return res,
|
||||||
|
}
|
||||||
|
};
|
||||||
|
}
|
||||||
|
|
||||||
/// A gRPC Server handler.
|
/// A gRPC Server handler.
|
||||||
///
|
///
|
||||||
/// This will wrap some inner [`Codec`] and provide utilities to handle
|
/// This will wrap some inner [`Codec`] and provide utilities to handle
|
||||||
@@ -20,6 +31,7 @@ use std::fmt;
|
|||||||
/// implements some [`Body`].
|
/// implements some [`Body`].
|
||||||
pub struct Grpc<T> {
|
pub struct Grpc<T> {
|
||||||
codec: T,
|
codec: T,
|
||||||
|
interceptor: Option<Interceptor>,
|
||||||
}
|
}
|
||||||
|
|
||||||
impl<T> Grpc<T>
|
impl<T> Grpc<T>
|
||||||
@@ -27,9 +39,21 @@ where
|
|||||||
T: Codec,
|
T: Codec,
|
||||||
T::Encode: Sync,
|
T::Encode: Sync,
|
||||||
{
|
{
|
||||||
/// Creates a new gRPC client with the provided [`Codec`].
|
/// Creates a new gRPC server with the provided [`Codec`].
|
||||||
pub fn new(codec: T) -> Self {
|
pub fn new(codec: T) -> Self {
|
||||||
Self { codec }
|
Self {
|
||||||
|
codec,
|
||||||
|
interceptor: None,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Creates a new gRPC server with the provided [`Codec`] and will apply the provided
|
||||||
|
/// interceptor on each inbound request.
|
||||||
|
pub fn with_interceptor(codec: T, interceptor: impl Into<Interceptor>) -> Self {
|
||||||
|
Self {
|
||||||
|
codec,
|
||||||
|
interceptor: Some(interceptor.into()),
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
/// Handle a single unary gRPC request.
|
/// Handle a single unary gRPC request.
|
||||||
@@ -53,6 +77,8 @@ where
|
|||||||
}
|
}
|
||||||
};
|
};
|
||||||
|
|
||||||
|
let request = t!(self.intercept_request(request));
|
||||||
|
|
||||||
let response = service
|
let response = service
|
||||||
.call(request)
|
.call(request)
|
||||||
.await
|
.await
|
||||||
@@ -80,6 +106,8 @@ where
|
|||||||
}
|
}
|
||||||
};
|
};
|
||||||
|
|
||||||
|
let request = t!(self.intercept_request(request));
|
||||||
|
|
||||||
let response = service.call(request).await;
|
let response = service.call(request).await;
|
||||||
|
|
||||||
self.map_response(response)
|
self.map_response(response)
|
||||||
@@ -97,6 +125,7 @@ where
|
|||||||
B::Error: Into<crate::Error> + Send + 'static,
|
B::Error: Into<crate::Error> + Send + 'static,
|
||||||
{
|
{
|
||||||
let request = self.map_request_streaming(req);
|
let request = self.map_request_streaming(req);
|
||||||
|
let request = t!(self.intercept_request(request));
|
||||||
let response = service
|
let response = service
|
||||||
.call(request)
|
.call(request)
|
||||||
.await
|
.await
|
||||||
@@ -117,6 +146,7 @@ where
|
|||||||
B::Error: Into<crate::Error> + Send,
|
B::Error: Into<crate::Error> + Send,
|
||||||
{
|
{
|
||||||
let request = self.map_request_streaming(req);
|
let request = self.map_request_streaming(req);
|
||||||
|
let request = t!(self.intercept_request(request));
|
||||||
let response = service.call(request).await;
|
let response = service.call(request).await;
|
||||||
self.map_response(response)
|
self.map_response(response)
|
||||||
}
|
}
|
||||||
@@ -180,18 +210,34 @@ where
|
|||||||
|
|
||||||
http::Response::from_parts(parts, BoxBody::new(body))
|
http::Response::from_parts(parts, BoxBody::new(body))
|
||||||
}
|
}
|
||||||
Err(status) => {
|
Err(status) => Self::map_status(status),
|
||||||
let (mut parts, _body) = Response::new(()).into_http().into_parts();
|
}
|
||||||
|
}
|
||||||
|
|
||||||
parts.headers.insert(
|
fn map_status(status: Status) -> http::Response<BoxBody> {
|
||||||
http::header::CONTENT_TYPE,
|
let (mut parts, _body) = Response::new(()).into_http().into_parts();
|
||||||
http::header::HeaderValue::from_static("application/grpc"),
|
|
||||||
);
|
|
||||||
|
|
||||||
status.add_header(&mut parts.headers).unwrap();
|
parts.headers.insert(
|
||||||
|
http::header::CONTENT_TYPE,
|
||||||
|
http::header::HeaderValue::from_static("application/grpc"),
|
||||||
|
);
|
||||||
|
|
||||||
http::Response::from_parts(parts, BoxBody::empty())
|
status.add_header(&mut parts.headers).unwrap();
|
||||||
|
|
||||||
|
http::Response::from_parts(parts, BoxBody::empty())
|
||||||
|
}
|
||||||
|
|
||||||
|
fn intercept_request<A>(&self, req: Request<A>) -> Result<Request<A>, http::Response<BoxBody>> {
|
||||||
|
if let Some(interceptor) = &self.interceptor {
|
||||||
|
match interceptor.call(req) {
|
||||||
|
Ok(req) => Ok(req),
|
||||||
|
Err(status) => {
|
||||||
|
let res = Self::map_status(status);
|
||||||
|
return Err(res);
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
} else {
|
||||||
|
Ok(req)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -10,7 +10,6 @@ use http::uri::{InvalidUri, Uri};
|
|||||||
use std::{
|
use std::{
|
||||||
convert::{TryFrom, TryInto},
|
convert::{TryFrom, TryInto},
|
||||||
fmt,
|
fmt,
|
||||||
sync::Arc,
|
|
||||||
time::Duration,
|
time::Duration,
|
||||||
};
|
};
|
||||||
use tower_make::MakeConnection;
|
use tower_make::MakeConnection;
|
||||||
@@ -27,8 +26,6 @@ pub struct Endpoint {
|
|||||||
#[cfg(feature = "tls")]
|
#[cfg(feature = "tls")]
|
||||||
pub(crate) tls: Option<TlsConnector>,
|
pub(crate) tls: Option<TlsConnector>,
|
||||||
pub(crate) buffer_size: Option<usize>,
|
pub(crate) buffer_size: Option<usize>,
|
||||||
pub(crate) interceptor_headers:
|
|
||||||
Option<Arc<dyn Fn(&mut http::HeaderMap) + Send + Sync + 'static>>,
|
|
||||||
pub(crate) init_stream_window_size: Option<u32>,
|
pub(crate) init_stream_window_size: Option<u32>,
|
||||||
pub(crate) init_connection_window_size: Option<u32>,
|
pub(crate) init_connection_window_size: Option<u32>,
|
||||||
pub(crate) tcp_keepalive: Option<Duration>,
|
pub(crate) tcp_keepalive: Option<Duration>,
|
||||||
@@ -152,29 +149,6 @@ impl Endpoint {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
/// Intercept outbound HTTP Request headers;
|
|
||||||
///
|
|
||||||
/// # Example
|
|
||||||
///
|
|
||||||
/// ```
|
|
||||||
/// # use tonic::transport::Endpoint;
|
|
||||||
/// # use std::time::Duration;
|
|
||||||
/// # let mut builder = Endpoint::from_static("https://example.com");
|
|
||||||
/// builder.intercept_headers(|headers| {
|
|
||||||
/// // Do something with headers
|
|
||||||
/// headers.insert("hello", "world".parse().unwrap());
|
|
||||||
/// });
|
|
||||||
/// ```
|
|
||||||
pub fn intercept_headers<F>(self, f: F) -> Self
|
|
||||||
where
|
|
||||||
F: Fn(&mut http::HeaderMap) + Send + Sync + 'static,
|
|
||||||
{
|
|
||||||
Endpoint {
|
|
||||||
interceptor_headers: Some(Arc::new(f)),
|
|
||||||
..self
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
/// Configures TLS for the endpoint.
|
/// Configures TLS for the endpoint.
|
||||||
#[cfg(feature = "tls")]
|
#[cfg(feature = "tls")]
|
||||||
#[cfg_attr(docsrs, doc(cfg(feature = "tls")))]
|
#[cfg_attr(docsrs, doc(cfg(feature = "tls")))]
|
||||||
@@ -237,7 +211,6 @@ impl From<Uri> for Endpoint {
|
|||||||
#[cfg(feature = "tls")]
|
#[cfg(feature = "tls")]
|
||||||
tls: None,
|
tls: None,
|
||||||
buffer_size: None,
|
buffer_size: None,
|
||||||
interceptor_headers: None,
|
|
||||||
init_stream_window_size: None,
|
init_stream_window_size: None,
|
||||||
init_connection_window_size: None,
|
init_connection_window_size: None,
|
||||||
tcp_keepalive: None,
|
tcp_keepalive: None,
|
||||||
|
|||||||
@@ -21,7 +21,6 @@ use std::{
|
|||||||
fmt,
|
fmt,
|
||||||
future::Future,
|
future::Future,
|
||||||
pin::Pin,
|
pin::Pin,
|
||||||
sync::Arc,
|
|
||||||
task::{Context, Poll},
|
task::{Context, Poll},
|
||||||
};
|
};
|
||||||
use tokio::io::{AsyncRead, AsyncWrite};
|
use tokio::io::{AsyncRead, AsyncWrite};
|
||||||
@@ -63,7 +62,6 @@ 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.
|
||||||
@@ -114,14 +112,9 @@ 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, interceptor_headers)
|
Self::balance(discover, buffer_size)
|
||||||
}
|
}
|
||||||
|
|
||||||
pub(crate) async fn connect<C>(connector: C, endpoint: Endpoint) -> Result<Self, super::Error>
|
pub(crate) async fn connect<C>(connector: C, endpoint: Endpoint) -> Result<Self, super::Error>
|
||||||
@@ -132,7 +125,6 @@ impl Channel {
|
|||||||
C::Response: AsyncRead + AsyncWrite + HyperConnection + Unpin + Send + 'static,
|
C::Response: AsyncRead + AsyncWrite + HyperConnection + Unpin + Send + 'static,
|
||||||
{
|
{
|
||||||
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(connector, endpoint)
|
let svc = Connection::new(connector, endpoint)
|
||||||
.await
|
.await
|
||||||
@@ -140,17 +132,10 @@ impl Channel {
|
|||||||
|
|
||||||
let svc = Buffer::new(Either::A(svc), buffer_size);
|
let svc = Buffer::new(Either::A(svc), buffer_size);
|
||||||
|
|
||||||
Ok(Channel {
|
Ok(Channel { svc })
|
||||||
svc,
|
|
||||||
interceptor_headers,
|
|
||||||
})
|
|
||||||
}
|
}
|
||||||
|
|
||||||
pub(crate) fn balance<D>(
|
pub(crate) fn balance<D>(discover: D, buffer_size: usize) -> Self
|
||||||
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>,
|
||||||
@@ -161,10 +146,7 @@ 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 {
|
Channel { svc }
|
||||||
svc,
|
|
||||||
interceptor_headers,
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -177,11 +159,7 @@ impl GrpcService<BoxBody> for Channel {
|
|||||||
GrpcService::poll_ready(&mut self.svc, cx).map_err(|e| super::Error::from_source(e))
|
GrpcService::poll_ready(&mut self.svc, cx).map_err(|e| super::Error::from_source(e))
|
||||||
}
|
}
|
||||||
|
|
||||||
fn call(&mut self, mut request: Request<BoxBody>) -> Self::Future {
|
fn call(&mut self, 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 }
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -12,7 +12,6 @@
|
|||||||
//! - Timeouts
|
//! - Timeouts
|
||||||
//! - Concurrency Limits
|
//! - Concurrency Limits
|
||||||
//! - Rate limiting
|
//! - Rate limiting
|
||||||
//! - gRPC Interceptors
|
|
||||||
//!
|
//!
|
||||||
//! # Examples
|
//! # Examples
|
||||||
//!
|
//!
|
||||||
@@ -77,10 +76,6 @@
|
|||||||
//! .tls_config(ServerTlsConfig::with_rustls()
|
//! .tls_config(ServerTlsConfig::with_rustls()
|
||||||
//! .identity(Identity::from_pem(&cert, &key)))
|
//! .identity(Identity::from_pem(&cert, &key)))
|
||||||
//! .concurrency_limit_per_connection(256)
|
//! .concurrency_limit_per_connection(256)
|
||||||
//! .interceptor_fn(|svc, req| {
|
|
||||||
//! println!("Request: {:?}", req);
|
|
||||||
//! svc.call(req)
|
|
||||||
//! })
|
|
||||||
//! .add_service(my_svc)
|
//! .add_service(my_svc)
|
||||||
//! .serve(addr)
|
//! .serve(addr)
|
||||||
//! .await?;
|
//! .await?;
|
||||||
@@ -104,7 +99,7 @@ pub use self::error::Error;
|
|||||||
#[doc(inline)]
|
#[doc(inline)]
|
||||||
pub use self::server::{Server, ServiceName};
|
pub use self::server::{Server, ServiceName};
|
||||||
pub use self::tls::{Certificate, Identity};
|
pub use self::tls::{Certificate, Identity};
|
||||||
pub use hyper::Body;
|
pub use hyper::{Body, Uri};
|
||||||
|
|
||||||
#[cfg(feature = "tls")]
|
#[cfg(feature = "tls")]
|
||||||
#[cfg_attr(docsrs, doc(cfg(feature = "tls")))]
|
#[cfg_attr(docsrs, doc(cfg(feature = "tls")))]
|
||||||
|
|||||||
@@ -15,7 +15,7 @@ use super::service::TlsAcceptor;
|
|||||||
|
|
||||||
use incoming::TcpIncoming;
|
use incoming::TcpIncoming;
|
||||||
|
|
||||||
use super::service::{layer_fn, Or, Routes, ServerIo, ServiceBuilderExt};
|
use super::service::{Or, Routes, ServerIo, ServiceBuilderExt};
|
||||||
use crate::{body::BoxBody, request::ConnectionInfo};
|
use crate::{body::BoxBody, request::ConnectionInfo};
|
||||||
use futures_core::Stream;
|
use futures_core::Stream;
|
||||||
use futures_util::{
|
use futures_util::{
|
||||||
@@ -35,15 +35,11 @@ use std::{
|
|||||||
};
|
};
|
||||||
use tokio::io::{AsyncRead, AsyncWrite};
|
use tokio::io::{AsyncRead, AsyncWrite};
|
||||||
use tower::{
|
use tower::{
|
||||||
layer::{Layer, Stack},
|
limit::concurrency::ConcurrencyLimitLayer, timeout::TimeoutLayer, Service, ServiceBuilder,
|
||||||
limit::concurrency::ConcurrencyLimitLayer,
|
|
||||||
timeout::TimeoutLayer,
|
|
||||||
Service, ServiceBuilder,
|
|
||||||
};
|
};
|
||||||
use tracing_futures::{Instrument, Instrumented};
|
use tracing_futures::{Instrument, Instrumented};
|
||||||
|
|
||||||
type BoxService = tower::util::BoxService<Request<Body>, Response<BoxBody>, crate::Error>;
|
type BoxService = tower::util::BoxService<Request<Body>, Response<BoxBody>, crate::Error>;
|
||||||
type Interceptor = Arc<dyn Layer<BoxService, Service = BoxService> + Send + Sync + 'static>;
|
|
||||||
type TraceInterceptor = Arc<dyn Fn(&HeaderMap) -> tracing::Span + Send + Sync + 'static>;
|
type TraceInterceptor = Arc<dyn Fn(&HeaderMap) -> tracing::Span + Send + Sync + 'static>;
|
||||||
|
|
||||||
/// A default batteries included `transport` server.
|
/// A default batteries included `transport` server.
|
||||||
@@ -56,7 +52,6 @@ type TraceInterceptor = Arc<dyn Fn(&HeaderMap) -> tracing::Span + Send + Sync +
|
|||||||
/// wanting to create a more complex and/or specific implementation.
|
/// wanting to create a more complex and/or specific implementation.
|
||||||
#[derive(Default, Clone)]
|
#[derive(Default, Clone)]
|
||||||
pub struct Server {
|
pub struct Server {
|
||||||
interceptor: Option<Interceptor>,
|
|
||||||
trace_interceptor: Option<TraceInterceptor>,
|
trace_interceptor: Option<TraceInterceptor>,
|
||||||
concurrency_limit: Option<usize>,
|
concurrency_limit: Option<usize>,
|
||||||
timeout: Option<Duration>,
|
timeout: Option<Duration>,
|
||||||
@@ -198,35 +193,6 @@ impl Server {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
/// Intercept the execution of gRPC methods.
|
|
||||||
///
|
|
||||||
/// ```
|
|
||||||
/// # use tonic::transport::Server;
|
|
||||||
/// # use tower_service::Service;
|
|
||||||
/// # let mut builder = Server::builder();
|
|
||||||
/// builder.interceptor_fn(|svc, req| {
|
|
||||||
/// println!("request={:?}", req);
|
|
||||||
/// svc.call(req)
|
|
||||||
/// });
|
|
||||||
/// ```
|
|
||||||
pub fn interceptor_fn<F, Out>(self, f: F) -> Self
|
|
||||||
where
|
|
||||||
F: Fn(&mut BoxService, Request<Body>) -> Out + Send + Sync + 'static,
|
|
||||||
Out: Future<Output = Result<Response<BoxBody>, crate::Error>> + Send + 'static,
|
|
||||||
{
|
|
||||||
let f = Arc::new(f);
|
|
||||||
let interceptor = layer_fn(move |mut s| {
|
|
||||||
let f = f.clone();
|
|
||||||
tower::service_fn(move |req| f(&mut s, req))
|
|
||||||
});
|
|
||||||
let layer = Stack::new(interceptor, layer_fn(BoxService::new));
|
|
||||||
|
|
||||||
Server {
|
|
||||||
interceptor: Some(Arc::new(layer)),
|
|
||||||
..self
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
/// Intercept inbound headers and add a [`tracing::Span`] to each response future.
|
/// Intercept inbound headers and add a [`tracing::Span`] to each response future.
|
||||||
pub fn trace_fn<F>(self, f: F) -> Self
|
pub fn trace_fn<F>(self, f: F) -> Self
|
||||||
where
|
where
|
||||||
@@ -270,7 +236,6 @@ impl Server {
|
|||||||
IE: Into<crate::Error>,
|
IE: Into<crate::Error>,
|
||||||
F: Future<Output = ()>,
|
F: Future<Output = ()>,
|
||||||
{
|
{
|
||||||
let interceptor = self.interceptor.clone();
|
|
||||||
let span = self.trace_interceptor.clone();
|
let span = self.trace_interceptor.clone();
|
||||||
let concurrency_limit = self.concurrency_limit;
|
let concurrency_limit = self.concurrency_limit;
|
||||||
let init_connection_window_size = self.init_connection_window_size;
|
let init_connection_window_size = self.init_connection_window_size;
|
||||||
@@ -283,7 +248,6 @@ impl Server {
|
|||||||
|
|
||||||
let svc = MakeSvc {
|
let svc = MakeSvc {
|
||||||
inner: svc,
|
inner: svc,
|
||||||
interceptor,
|
|
||||||
concurrency_limit,
|
concurrency_limit,
|
||||||
timeout,
|
timeout,
|
||||||
span,
|
span,
|
||||||
@@ -480,7 +444,6 @@ impl<S> fmt::Debug for Svc<S> {
|
|||||||
}
|
}
|
||||||
|
|
||||||
struct MakeSvc<S> {
|
struct MakeSvc<S> {
|
||||||
interceptor: Option<Interceptor>,
|
|
||||||
concurrency_limit: Option<usize>,
|
concurrency_limit: Option<usize>,
|
||||||
timeout: Option<Duration>,
|
timeout: Option<Duration>,
|
||||||
inner: S,
|
inner: S,
|
||||||
@@ -508,7 +471,6 @@ where
|
|||||||
peer_certs: io.peer_certs().map(Arc::new),
|
peer_certs: io.peer_certs().map(Arc::new),
|
||||||
};
|
};
|
||||||
|
|
||||||
let interceptor = self.interceptor.clone();
|
|
||||||
let svc = self.inner.clone();
|
let svc = self.inner.clone();
|
||||||
let concurrency_limit = self.concurrency_limit;
|
let concurrency_limit = self.concurrency_limit;
|
||||||
let timeout = self.timeout.clone();
|
let timeout = self.timeout.clone();
|
||||||
@@ -520,20 +482,11 @@ where
|
|||||||
.optional_layer(timeout.map(TimeoutLayer::new))
|
.optional_layer(timeout.map(TimeoutLayer::new))
|
||||||
.service(svc);
|
.service(svc);
|
||||||
|
|
||||||
let svc = if let Some(interceptor) = interceptor {
|
let svc = BoxService::new(Svc {
|
||||||
let layered = interceptor.layer(BoxService::new(Svc {
|
inner: svc,
|
||||||
inner: svc,
|
span,
|
||||||
span,
|
conn_info,
|
||||||
conn_info,
|
});
|
||||||
}));
|
|
||||||
BoxService::new(layered)
|
|
||||||
} else {
|
|
||||||
BoxService::new(Svc {
|
|
||||||
inner: svc,
|
|
||||||
span,
|
|
||||||
conn_info,
|
|
||||||
})
|
|
||||||
};
|
|
||||||
|
|
||||||
Ok(svc)
|
Ok(svc)
|
||||||
})
|
})
|
||||||
|
|||||||
@@ -42,6 +42,8 @@ impl<L> ServiceBuilderExt<L> for ServiceBuilder<L> {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// TODO: figure out why this is causing a warning even though its used in optional_layer_fn
|
||||||
|
#[allow(dead_code)]
|
||||||
pub(crate) fn layer_fn<F>(f: F) -> LayerFn<F> {
|
pub(crate) fn layer_fn<F>(f: F) -> LayerFn<F> {
|
||||||
LayerFn(f)
|
LayerFn(f)
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -14,7 +14,7 @@ pub(crate) use self::connection::Connection;
|
|||||||
pub(crate) use self::connector::connector;
|
pub(crate) use self::connector::connector;
|
||||||
pub(crate) use self::discover::ServiceList;
|
pub(crate) use self::discover::ServiceList;
|
||||||
pub(crate) use self::io::ServerIo;
|
pub(crate) use self::io::ServerIo;
|
||||||
pub(crate) use self::layer::{layer_fn, ServiceBuilderExt};
|
pub(crate) use self::layer::ServiceBuilderExt;
|
||||||
pub(crate) use self::router::{Or, Routes};
|
pub(crate) use self::router::{Or, Routes};
|
||||||
#[cfg(feature = "tls")]
|
#[cfg(feature = "tls")]
|
||||||
pub(crate) use self::tls::{TlsAcceptor, TlsConnector};
|
pub(crate) use self::tls::{TlsAcceptor, TlsConnector};
|
||||||
|
|||||||
Reference in New Issue
Block a user