diff --git a/Cargo.toml b/Cargo.toml index 05b7f6f..9fd8d2f 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -22,6 +22,7 @@ members = [ "tests/integration_tests", "tests/stream_conflict", "tests/root-crate-path", + "tests/compression", "tonic-web/tests/integration" ] diff --git a/examples/Cargo.toml b/examples/Cargo.toml index 083f211..855f05f 100644 --- a/examples/Cargo.toml +++ b/examples/Cargo.toml @@ -150,6 +150,14 @@ path = "src/hyper_warp_multiplex/client.rs" name = "hyper-warp-multiplex-server" path = "src/hyper_warp_multiplex/server.rs" +[[bin]] +name = "compression-server" +path = "src/compression/server.rs" + +[[bin]] +name = "compression-client" +path = "src/compression/client.rs" + [dependencies] tonic = { path = "../tonic", features = ["tls"] } prost = "0.7" diff --git a/examples/src/compression/client.rs b/examples/src/compression/client.rs new file mode 100644 index 0000000..77ffeeb --- /dev/null +++ b/examples/src/compression/client.rs @@ -0,0 +1,27 @@ +use hello_world::greeter_client::GreeterClient; +use hello_world::HelloRequest; +use tonic::transport::Channel; + +pub mod hello_world { + tonic::include_proto!("helloworld"); +} + +#[tokio::main] +async fn main() -> Result<(), Box> { + let channel = Channel::builder("http://[::1]:50051".parse().unwrap()) + .connect() + .await + .unwrap(); + + let mut client = GreeterClient::new(channel).send_gzip().accept_gzip(); + + let request = tonic::Request::new(HelloRequest { + name: "Tonic".into(), + }); + + let response = client.say_hello(request).await?; + + dbg!(response); + + Ok(()) +} diff --git a/examples/src/compression/server.rs b/examples/src/compression/server.rs new file mode 100644 index 0000000..36f5081 --- /dev/null +++ b/examples/src/compression/server.rs @@ -0,0 +1,40 @@ +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, + ) -> Result, Status> { + println!("Got a request from {:?}", request.remote_addr()); + + let reply = hello_world::HelloReply { + message: format!("Hello {}!", request.into_inner().name), + }; + Ok(Response::new(reply)) + } +} + +#[tokio::main] +async fn main() -> Result<(), Box> { + let addr = "[::1]:50051".parse().unwrap(); + let greeter = MyGreeter::default(); + + println!("GreeterServer listening on {}", addr); + + let service = GreeterServer::new(greeter).send_gzip().accept_gzip(); + + Server::builder().add_service(service).serve(addr).await?; + + Ok(()) +} diff --git a/tests/compression/Cargo.toml b/tests/compression/Cargo.toml new file mode 100644 index 0000000..6646d39 --- /dev/null +++ b/tests/compression/Cargo.toml @@ -0,0 +1,24 @@ +[package] +name = "compression" +version = "0.1.0" +authors = ["Lucio Franco "] +edition = "2018" +publish = false +license = "MIT" + +[dependencies] +tonic = { path = "../../tonic", features = ["compression"] } +prost = "0.7" +tokio = { version = "1.0", features = ["macros", "rt-multi-thread", "net"] } +tower = { version = "0.4", features = [] } +http-body = "0.4" +http = "0.2" +tokio-stream = { version = "0.1.5", features = ["net"] } +tower-http = { version = "0.1", features = ["map-response-body", "map-request-body"] } +bytes = "1" +futures = "0.3" +pin-project = "1.0" +hyper = "0.14" + +[build-dependencies] +tonic-build = { path = "../../tonic-build", features = ["compression"] } diff --git a/tests/compression/build.rs b/tests/compression/build.rs new file mode 100644 index 0000000..a091e94 --- /dev/null +++ b/tests/compression/build.rs @@ -0,0 +1,3 @@ +fn main() { + tonic_build::compile_protos("proto/test.proto").unwrap(); +} diff --git a/tests/compression/proto/test.proto b/tests/compression/proto/test.proto new file mode 100644 index 0000000..325471b --- /dev/null +++ b/tests/compression/proto/test.proto @@ -0,0 +1,19 @@ +syntax = "proto3"; + +package test; + +import "google/protobuf/empty.proto"; + +service Test { + rpc CompressOutputUnary(google.protobuf.Empty) returns (SomeData); + rpc CompressInputUnary(SomeData) returns (google.protobuf.Empty); + rpc CompressOutputServerStream(google.protobuf.Empty) returns (stream SomeData); + rpc CompressInputClientStream(stream SomeData) returns (google.protobuf.Empty); + rpc CompressOutputClientStream(stream SomeData) returns (SomeData); + rpc CompressInputOutputBidirectionalStream(stream SomeData) returns (stream SomeData); +} + +message SomeData { + // include a bunch of data so there actually is something to compress + bytes data = 1; +} diff --git a/tests/compression/src/bidirectional_stream.rs b/tests/compression/src/bidirectional_stream.rs new file mode 100644 index 0000000..53dc833 --- /dev/null +++ b/tests/compression/src/bidirectional_stream.rs @@ -0,0 +1,78 @@ +use super::*; + +#[tokio::test(flavor = "multi_thread")] +async fn client_enabled_server_enabled() { + let (client, server) = tokio::io::duplex(UNCOMPRESSED_MIN_BODY_SIZE * 10); + + let svc = test_server::TestServer::new(Svc::default()) + .accept_gzip() + .send_gzip(); + + let request_bytes_counter = Arc::new(AtomicUsize::new(0)); + let response_bytes_counter = Arc::new(AtomicUsize::new(0)); + + fn assert_right_encoding(req: http::Request) -> http::Request { + assert_eq!(req.headers().get("grpc-encoding").unwrap(), "gzip"); + req + } + + tokio::spawn({ + let request_bytes_counter = request_bytes_counter.clone(); + let response_bytes_counter = response_bytes_counter.clone(); + async move { + Server::builder() + .layer( + ServiceBuilder::new() + .map_request(assert_right_encoding) + .layer(measure_request_body_size_layer( + request_bytes_counter.clone(), + )) + .layer(MapResponseBodyLayer::new(move |body| { + util::CountBytesBody { + inner: body, + counter: response_bytes_counter.clone(), + } + })) + .into_inner(), + ) + .add_service(svc) + .serve_with_incoming(futures::stream::iter(vec![Ok::<_, std::io::Error>( + MockStream(server), + )])) + .await + .unwrap(); + } + }); + + let mut client = test_client::TestClient::new(mock_io_channel(client).await) + .send_gzip() + .accept_gzip(); + + let data = [0_u8; UNCOMPRESSED_MIN_BODY_SIZE].to_vec(); + let stream = futures::stream::iter(vec![SomeData { data: data.clone() }, SomeData { data }]); + let req = Request::new(stream); + + let res = client + .compress_input_output_bidirectional_stream(req) + .await + .unwrap(); + + assert_eq!(res.metadata().get("grpc-encoding").unwrap(), "gzip"); + + let mut stream: Streaming = res.into_inner(); + + stream + .next() + .await + .expect("stream empty") + .expect("item was error"); + + stream + .next() + .await + .expect("stream empty") + .expect("item was error"); + + assert!(request_bytes_counter.load(SeqCst) < UNCOMPRESSED_MIN_BODY_SIZE); + assert!(response_bytes_counter.load(SeqCst) < UNCOMPRESSED_MIN_BODY_SIZE); +} diff --git a/tests/compression/src/client_stream.rs b/tests/compression/src/client_stream.rs new file mode 100644 index 0000000..620f917 --- /dev/null +++ b/tests/compression/src/client_stream.rs @@ -0,0 +1,167 @@ +use super::*; +use http_body::Body as _; + +#[tokio::test(flavor = "multi_thread")] +async fn client_enabled_server_enabled() { + let (client, server) = tokio::io::duplex(UNCOMPRESSED_MIN_BODY_SIZE * 10); + + let svc = test_server::TestServer::new(Svc::default()).accept_gzip(); + + let request_bytes_counter = Arc::new(AtomicUsize::new(0)); + + fn assert_right_encoding(req: http::Request) -> http::Request { + assert_eq!(req.headers().get("grpc-encoding").unwrap(), "gzip"); + req + } + + tokio::spawn({ + let request_bytes_counter = request_bytes_counter.clone(); + async move { + Server::builder() + .layer( + ServiceBuilder::new() + .map_request(assert_right_encoding) + .layer(measure_request_body_size_layer( + request_bytes_counter.clone(), + )) + .into_inner(), + ) + .add_service(svc) + .serve_with_incoming(futures::stream::iter(vec![Ok::<_, std::io::Error>( + MockStream(server), + )])) + .await + .unwrap(); + } + }); + + let mut client = test_client::TestClient::new(mock_io_channel(client).await).send_gzip(); + + let data = [0_u8; UNCOMPRESSED_MIN_BODY_SIZE].to_vec(); + let stream = futures::stream::iter(vec![SomeData { data: data.clone() }, SomeData { data }]); + let req = Request::new(Box::pin(stream)); + + client.compress_input_client_stream(req).await.unwrap(); + + let bytes_sent = request_bytes_counter.load(SeqCst); + assert!(bytes_sent < UNCOMPRESSED_MIN_BODY_SIZE); +} + +#[tokio::test(flavor = "multi_thread")] +async fn client_disabled_server_enabled() { + let (client, server) = tokio::io::duplex(UNCOMPRESSED_MIN_BODY_SIZE * 10); + + let svc = test_server::TestServer::new(Svc::default()).accept_gzip(); + + let request_bytes_counter = Arc::new(AtomicUsize::new(0)); + + fn assert_right_encoding(req: http::Request) -> http::Request { + assert!(req.headers().get("grpc-encoding").is_none()); + req + } + + tokio::spawn({ + let request_bytes_counter = request_bytes_counter.clone(); + async move { + Server::builder() + .layer( + ServiceBuilder::new() + .map_request(assert_right_encoding) + .layer(measure_request_body_size_layer( + request_bytes_counter.clone(), + )) + .into_inner(), + ) + .add_service(svc) + .serve_with_incoming(futures::stream::iter(vec![Ok::<_, std::io::Error>( + MockStream(server), + )])) + .await + .unwrap(); + } + }); + + let mut client = test_client::TestClient::new(mock_io_channel(client).await); + + let data = [0_u8; UNCOMPRESSED_MIN_BODY_SIZE].to_vec(); + let stream = futures::stream::iter(vec![SomeData { data: data.clone() }, SomeData { data }]); + let req = Request::new(Box::pin(stream)); + + client.compress_input_client_stream(req).await.unwrap(); + + let bytes_sent = request_bytes_counter.load(SeqCst); + assert!(bytes_sent > UNCOMPRESSED_MIN_BODY_SIZE); +} + +#[tokio::test(flavor = "multi_thread")] +async fn client_enabled_server_disabled() { + let (client, server) = tokio::io::duplex(UNCOMPRESSED_MIN_BODY_SIZE * 10); + + let svc = test_server::TestServer::new(Svc::default()); + + tokio::spawn(async move { + Server::builder() + .add_service(svc) + .serve_with_incoming(futures::stream::iter(vec![Ok::<_, std::io::Error>( + MockStream(server), + )])) + .await + .unwrap(); + }); + + let mut client = test_client::TestClient::new(mock_io_channel(client).await).send_gzip(); + + let data = [0_u8; UNCOMPRESSED_MIN_BODY_SIZE].to_vec(); + let stream = futures::stream::iter(vec![SomeData { data: data.clone() }, SomeData { data }]); + let req = Request::new(Box::pin(stream)); + + let status = client.compress_input_client_stream(req).await.unwrap_err(); + + assert_eq!(status.code(), tonic::Code::Unimplemented); + assert_eq!( + status.message(), + "Content is compressed with `gzip` which isn't supported" + ); +} + +#[tokio::test(flavor = "multi_thread")] +async fn compressing_response_from_client_stream() { + let (client, server) = tokio::io::duplex(UNCOMPRESSED_MIN_BODY_SIZE * 10); + + let svc = test_server::TestServer::new(Svc::default()).send_gzip(); + + let response_bytes_counter = Arc::new(AtomicUsize::new(0)); + + tokio::spawn({ + let response_bytes_counter = response_bytes_counter.clone(); + async move { + Server::builder() + .layer( + ServiceBuilder::new() + .layer(MapResponseBodyLayer::new(move |body| { + util::CountBytesBody { + inner: body, + counter: response_bytes_counter.clone(), + } + })) + .into_inner(), + ) + .add_service(svc) + .serve_with_incoming(futures::stream::iter(vec![Ok::<_, std::io::Error>( + MockStream(server), + )])) + .await + .unwrap(); + } + }); + + let mut client = test_client::TestClient::new(mock_io_channel(client).await).accept_gzip(); + + let stream = futures::stream::iter(vec![]); + let req = Request::new(Box::pin(stream)); + + let res = client.compress_output_client_stream(req).await.unwrap(); + assert_eq!(res.metadata().get("grpc-encoding").unwrap(), "gzip"); + let bytes_sent = response_bytes_counter.load(SeqCst); + assert!(bytes_sent < UNCOMPRESSED_MIN_BODY_SIZE); +} diff --git a/tests/compression/src/compressing_request.rs b/tests/compression/src/compressing_request.rs new file mode 100644 index 0000000..de41241 --- /dev/null +++ b/tests/compression/src/compressing_request.rs @@ -0,0 +1,89 @@ +use super::*; +use http_body::Body as _; + +#[tokio::test(flavor = "multi_thread")] +async fn client_enabled_server_enabled() { + let (client, server) = tokio::io::duplex(UNCOMPRESSED_MIN_BODY_SIZE * 10); + + let svc = test_server::TestServer::new(Svc::default()).accept_gzip(); + + let request_bytes_counter = Arc::new(AtomicUsize::new(0)); + + fn assert_right_encoding(req: http::Request) -> http::Request { + assert_eq!(req.headers().get("grpc-encoding").unwrap(), "gzip"); + req + } + + tokio::spawn({ + let request_bytes_counter = request_bytes_counter.clone(); + async move { + Server::builder() + .layer( + ServiceBuilder::new() + .layer( + ServiceBuilder::new() + .map_request(assert_right_encoding) + .layer(measure_request_body_size_layer(request_bytes_counter)) + .into_inner(), + ) + .into_inner(), + ) + .add_service(svc) + .serve_with_incoming(futures::stream::iter(vec![Ok::<_, std::io::Error>( + MockStream(server), + )])) + .await + .unwrap(); + } + }); + + let mut client = test_client::TestClient::new(mock_io_channel(client).await).send_gzip(); + + for _ in 0..3 { + client + .compress_input_unary(SomeData { + data: [0_u8; UNCOMPRESSED_MIN_BODY_SIZE].to_vec(), + }) + .await + .unwrap(); + let bytes_sent = request_bytes_counter.load(SeqCst); + assert!(bytes_sent < UNCOMPRESSED_MIN_BODY_SIZE); + } +} + +#[tokio::test(flavor = "multi_thread")] +async fn client_enabled_server_disabled() { + let (client, server) = tokio::io::duplex(UNCOMPRESSED_MIN_BODY_SIZE * 10); + + let svc = test_server::TestServer::new(Svc::default()); + + tokio::spawn(async move { + Server::builder() + .add_service(svc) + .serve_with_incoming(futures::stream::iter(vec![Ok::<_, std::io::Error>( + MockStream(server), + )])) + .await + .unwrap(); + }); + + let mut client = test_client::TestClient::new(mock_io_channel(client).await).send_gzip(); + + let status = client + .compress_input_unary(SomeData { + data: [0_u8; UNCOMPRESSED_MIN_BODY_SIZE].to_vec(), + }) + .await + .unwrap_err(); + + assert_eq!(status.code(), tonic::Code::Unimplemented); + assert_eq!( + status.message(), + "Content is compressed with `gzip` which isn't supported" + ); + + assert_eq!( + status.metadata().get("grpc-accept-encoding").unwrap(), + "identity" + ); +} diff --git a/tests/compression/src/compressing_response.rs b/tests/compression/src/compressing_response.rs new file mode 100644 index 0000000..e60903d --- /dev/null +++ b/tests/compression/src/compressing_response.rs @@ -0,0 +1,360 @@ +use super::*; + +#[tokio::test(flavor = "multi_thread")] +async fn client_enabled_server_enabled() { + let (client, server) = tokio::io::duplex(UNCOMPRESSED_MIN_BODY_SIZE * 10); + + #[derive(Clone, Copy)] + struct AssertCorrectAcceptEncoding(S); + + impl Service> for AssertCorrectAcceptEncoding + where + S: Service>, + { + type Response = S::Response; + type Error = S::Error; + type Future = S::Future; + + fn poll_ready( + &mut self, + cx: &mut std::task::Context<'_>, + ) -> std::task::Poll> { + self.0.poll_ready(cx) + } + + fn call(&mut self, req: http::Request) -> Self::Future { + assert_eq!( + req.headers().get("grpc-accept-encoding").unwrap(), + "gzip,identity" + ); + self.0.call(req) + } + } + + let svc = test_server::TestServer::new(Svc::default()).send_gzip(); + + let response_bytes_counter = Arc::new(AtomicUsize::new(0)); + + tokio::spawn({ + let response_bytes_counter = response_bytes_counter.clone(); + async move { + Server::builder() + .layer( + ServiceBuilder::new() + .layer(layer_fn(AssertCorrectAcceptEncoding)) + .layer(MapResponseBodyLayer::new(move |body| { + util::CountBytesBody { + inner: body, + counter: response_bytes_counter.clone(), + } + })) + .into_inner(), + ) + .add_service(svc) + .serve_with_incoming(futures::stream::iter(vec![Ok::<_, std::io::Error>( + MockStream(server), + )])) + .await + .unwrap(); + } + }); + + let mut client = test_client::TestClient::new(mock_io_channel(client).await).accept_gzip(); + + for _ in 0..3 { + let res = client.compress_output_unary(()).await.unwrap(); + assert_eq!(res.metadata().get("grpc-encoding").unwrap(), "gzip"); + let bytes_sent = response_bytes_counter.load(SeqCst); + assert!(bytes_sent < UNCOMPRESSED_MIN_BODY_SIZE); + } +} + +#[tokio::test(flavor = "multi_thread")] +async fn client_enabled_server_disabled() { + let (client, server) = tokio::io::duplex(UNCOMPRESSED_MIN_BODY_SIZE * 10); + + let svc = test_server::TestServer::new(Svc::default()); + + let response_bytes_counter = Arc::new(AtomicUsize::new(0)); + + tokio::spawn({ + let response_bytes_counter = response_bytes_counter.clone(); + async move { + Server::builder() + // no compression enable on the server so responses should not be compressed + .layer( + ServiceBuilder::new() + .layer(MapResponseBodyLayer::new(move |body| { + util::CountBytesBody { + inner: body, + counter: response_bytes_counter.clone(), + } + })) + .into_inner(), + ) + .add_service(svc) + .serve_with_incoming(futures::stream::iter(vec![Ok::<_, std::io::Error>( + MockStream(server), + )])) + .await + .unwrap(); + } + }); + + let mut client = test_client::TestClient::new(mock_io_channel(client).await).accept_gzip(); + + let res = client.compress_output_unary(()).await.unwrap(); + + assert!(res.metadata().get("grpc-encoding").is_none()); + + let bytes_sent = response_bytes_counter.load(SeqCst); + assert!(bytes_sent > UNCOMPRESSED_MIN_BODY_SIZE); +} + +#[tokio::test(flavor = "multi_thread")] +async fn client_disabled() { + let (client, server) = tokio::io::duplex(UNCOMPRESSED_MIN_BODY_SIZE * 10); + + #[derive(Clone, Copy)] + struct AssertCorrectAcceptEncoding(S); + + impl Service> for AssertCorrectAcceptEncoding + where + S: Service>, + { + type Response = S::Response; + type Error = S::Error; + type Future = S::Future; + + fn poll_ready( + &mut self, + cx: &mut std::task::Context<'_>, + ) -> std::task::Poll> { + self.0.poll_ready(cx) + } + + fn call(&mut self, req: http::Request) -> Self::Future { + assert!(req.headers().get("grpc-accept-encoding").is_none()); + self.0.call(req) + } + } + + let svc = test_server::TestServer::new(Svc::default()).send_gzip(); + + let response_bytes_counter = Arc::new(AtomicUsize::new(0)); + + tokio::spawn({ + let response_bytes_counter = response_bytes_counter.clone(); + async move { + Server::builder() + .layer( + ServiceBuilder::new() + .layer(layer_fn(AssertCorrectAcceptEncoding)) + .layer(MapResponseBodyLayer::new(move |body| { + util::CountBytesBody { + inner: body, + counter: response_bytes_counter.clone(), + } + })) + .into_inner(), + ) + .add_service(svc) + .serve_with_incoming(futures::stream::iter(vec![Ok::<_, std::io::Error>( + MockStream(server), + )])) + .await + .unwrap(); + } + }); + + let mut client = test_client::TestClient::new(mock_io_channel(client).await); + + let res = client.compress_output_unary(()).await.unwrap(); + + assert!(res.metadata().get("grpc-encoding").is_none()); + + let bytes_sent = response_bytes_counter.load(SeqCst); + assert!(bytes_sent > UNCOMPRESSED_MIN_BODY_SIZE); +} + +#[tokio::test(flavor = "multi_thread")] +async fn server_replying_with_unsupported_encoding() { + let (client, server) = tokio::io::duplex(UNCOMPRESSED_MIN_BODY_SIZE * 10); + + let svc = test_server::TestServer::new(Svc::default()).send_gzip(); + + fn add_weird_content_encoding(mut response: http::Response) -> http::Response { + response + .headers_mut() + .insert("grpc-encoding", "br".parse().unwrap()); + response + } + + tokio::spawn(async move { + Server::builder() + .layer( + ServiceBuilder::new() + .map_response(add_weird_content_encoding) + .into_inner(), + ) + .add_service(svc) + .serve_with_incoming(futures::stream::iter(vec![Ok::<_, std::io::Error>( + MockStream(server), + )])) + .await + .unwrap(); + }); + + let mut client = test_client::TestClient::new(mock_io_channel(client).await).accept_gzip(); + let status: Status = client.compress_output_unary(()).await.unwrap_err(); + + assert_eq!(status.code(), tonic::Code::Unimplemented); + assert_eq!( + status.message(), + "Content is compressed with `br` which isn't supported" + ); +} + +#[tokio::test(flavor = "multi_thread")] +async fn disabling_compression_on_single_response() { + let (client, server) = tokio::io::duplex(UNCOMPRESSED_MIN_BODY_SIZE * 10); + + let svc = test_server::TestServer::new(Svc { + disable_compressing_on_response: true, + }) + .send_gzip(); + + let response_bytes_counter = Arc::new(AtomicUsize::new(0)); + + tokio::spawn({ + let response_bytes_counter = response_bytes_counter.clone(); + async move { + Server::builder() + .layer( + ServiceBuilder::new() + .layer(MapResponseBodyLayer::new(move |body| { + util::CountBytesBody { + inner: body, + counter: response_bytes_counter.clone(), + } + })) + .into_inner(), + ) + .add_service(svc) + .serve_with_incoming(futures::stream::iter(vec![Ok::<_, std::io::Error>( + MockStream(server), + )])) + .await + .unwrap(); + } + }); + + let mut client = test_client::TestClient::new(mock_io_channel(client).await).accept_gzip(); + + let res = client.compress_output_unary(()).await.unwrap(); + assert_eq!(res.metadata().get("grpc-encoding").unwrap(), "gzip"); + let bytes_sent = response_bytes_counter.load(SeqCst); + assert!(bytes_sent > UNCOMPRESSED_MIN_BODY_SIZE); +} + +#[tokio::test(flavor = "multi_thread")] +async fn disabling_compression_on_response_but_keeping_compression_on_stream() { + let (client, server) = tokio::io::duplex(UNCOMPRESSED_MIN_BODY_SIZE * 10); + + let svc = test_server::TestServer::new(Svc { + disable_compressing_on_response: true, + }) + .send_gzip(); + + let response_bytes_counter = Arc::new(AtomicUsize::new(0)); + + tokio::spawn({ + let response_bytes_counter = response_bytes_counter.clone(); + async move { + Server::builder() + .layer( + ServiceBuilder::new() + .layer(MapResponseBodyLayer::new(move |body| { + util::CountBytesBody { + inner: body, + counter: response_bytes_counter.clone(), + } + })) + .into_inner(), + ) + .add_service(svc) + .serve_with_incoming(futures::stream::iter(vec![Ok::<_, std::io::Error>( + MockStream(server), + )])) + .await + .unwrap(); + } + }); + + let mut client = test_client::TestClient::new(mock_io_channel(client).await).accept_gzip(); + + let res = client.compress_output_server_stream(()).await.unwrap(); + + assert_eq!(res.metadata().get("grpc-encoding").unwrap(), "gzip"); + + let mut stream: Streaming = res.into_inner(); + + stream + .next() + .await + .expect("stream empty") + .expect("item was error"); + assert!(response_bytes_counter.load(SeqCst) < UNCOMPRESSED_MIN_BODY_SIZE); + + stream + .next() + .await + .expect("stream empty") + .expect("item was error"); + assert!(response_bytes_counter.load(SeqCst) < UNCOMPRESSED_MIN_BODY_SIZE); +} + +#[tokio::test(flavor = "multi_thread")] +async fn disabling_compression_on_response_from_client_stream() { + let (client, server) = tokio::io::duplex(UNCOMPRESSED_MIN_BODY_SIZE * 10); + + let svc = test_server::TestServer::new(Svc { + disable_compressing_on_response: true, + }) + .send_gzip(); + + let response_bytes_counter = Arc::new(AtomicUsize::new(0)); + + tokio::spawn({ + let response_bytes_counter = response_bytes_counter.clone(); + async move { + Server::builder() + .layer( + ServiceBuilder::new() + .layer(MapResponseBodyLayer::new(move |body| { + util::CountBytesBody { + inner: body, + counter: response_bytes_counter.clone(), + } + })) + .into_inner(), + ) + .add_service(svc) + .serve_with_incoming(futures::stream::iter(vec![Ok::<_, std::io::Error>( + MockStream(server), + )])) + .await + .unwrap(); + } + }); + + let mut client = test_client::TestClient::new(mock_io_channel(client).await).accept_gzip(); + + let stream = futures::stream::iter(vec![]); + let req = Request::new(Box::pin(stream)); + + let res = client.compress_output_client_stream(req).await.unwrap(); + assert_eq!(res.metadata().get("grpc-encoding").unwrap(), "gzip"); + let bytes_sent = response_bytes_counter.load(SeqCst); + assert!(bytes_sent > UNCOMPRESSED_MIN_BODY_SIZE); +} diff --git a/tests/compression/src/lib.rs b/tests/compression/src/lib.rs new file mode 100644 index 0000000..38c1203 --- /dev/null +++ b/tests/compression/src/lib.rs @@ -0,0 +1,130 @@ +#![allow(unused_imports)] + +use self::util::*; +use crate::util::{mock_io_channel, MockStream}; +use futures::{Stream, StreamExt}; +use std::convert::TryFrom; +use std::{ + pin::Pin, + sync::{ + atomic::{AtomicUsize, Ordering::SeqCst}, + Arc, + }, +}; +use tokio::net::TcpListener; +use tonic::{ + transport::{Channel, Endpoint, Server, Uri}, + Request, Response, Status, Streaming, +}; +use tower::{layer::layer_fn, service_fn, Service, ServiceBuilder}; +use tower_http::{map_request_body::MapRequestBodyLayer, map_response_body::MapResponseBodyLayer}; + +mod bidirectional_stream; +mod client_stream; +mod compressing_request; +mod compressing_response; +mod server_stream; +mod util; + +tonic::include_proto!("test"); + +#[derive(Debug)] +struct Svc { + disable_compressing_on_response: bool, +} + +impl Default for Svc { + fn default() -> Self { + Self { + disable_compressing_on_response: false, + } + } +} + +const UNCOMPRESSED_MIN_BODY_SIZE: usize = 1024; + +impl Svc { + fn prepare_response(&self, mut res: Response) -> Response { + if self.disable_compressing_on_response { + res.disable_compression(); + } + + res + } +} + +#[tonic::async_trait] +impl test_server::Test for Svc { + async fn compress_output_unary(&self, _req: Request<()>) -> Result, Status> { + let data = [0_u8; UNCOMPRESSED_MIN_BODY_SIZE]; + + Ok(self.prepare_response(Response::new(SomeData { + data: data.to_vec(), + }))) + } + + async fn compress_input_unary(&self, req: Request) -> Result, Status> { + assert_eq!(req.into_inner().data.len(), UNCOMPRESSED_MIN_BODY_SIZE); + Ok(Response::new(())) + } + + type CompressOutputServerStreamStream = + Pin> + Send + Sync + 'static>>; + + async fn compress_output_server_stream( + &self, + _req: Request<()>, + ) -> Result, Status> { + let data = [0_u8; UNCOMPRESSED_MIN_BODY_SIZE].to_vec(); + let stream = futures::stream::repeat(SomeData { data }) + .take(2) + .map(Ok::<_, Status>); + Ok(self.prepare_response(Response::new(Box::pin(stream)))) + } + + async fn compress_input_client_stream( + &self, + req: Request>, + ) -> Result, Status> { + let mut stream = req.into_inner(); + while let Some(item) = stream.next().await { + item.unwrap(); + } + Ok(self.prepare_response(Response::new(()))) + } + + async fn compress_output_client_stream( + &self, + req: Request>, + ) -> Result, Status> { + let mut stream = req.into_inner(); + while let Some(item) = stream.next().await { + item.unwrap(); + } + + let data = [0_u8; UNCOMPRESSED_MIN_BODY_SIZE]; + + Ok(self.prepare_response(Response::new(SomeData { + data: data.to_vec(), + }))) + } + + type CompressInputOutputBidirectionalStreamStream = + Pin> + Send + Sync + 'static>>; + + async fn compress_input_output_bidirectional_stream( + &self, + req: Request>, + ) -> Result, Status> { + let mut stream = req.into_inner(); + while let Some(item) = stream.next().await { + item.unwrap(); + } + + let data = [0_u8; UNCOMPRESSED_MIN_BODY_SIZE].to_vec(); + let stream = futures::stream::repeat(SomeData { data }) + .take(2) + .map(Ok::<_, Status>); + Ok(self.prepare_response(Response::new(Box::pin(stream)))) + } +} diff --git a/tests/compression/src/server_stream.rs b/tests/compression/src/server_stream.rs new file mode 100644 index 0000000..2d302bf --- /dev/null +++ b/tests/compression/src/server_stream.rs @@ -0,0 +1,150 @@ +use super::*; +use tonic::Streaming; + +#[tokio::test(flavor = "multi_thread")] +async fn client_enabled_server_enabled() { + let (client, server) = tokio::io::duplex(UNCOMPRESSED_MIN_BODY_SIZE * 10); + + let svc = test_server::TestServer::new(Svc::default()).send_gzip(); + + let response_bytes_counter = Arc::new(AtomicUsize::new(0)); + + tokio::spawn({ + let response_bytes_counter = response_bytes_counter.clone(); + async move { + Server::builder() + .layer( + ServiceBuilder::new() + .layer(MapResponseBodyLayer::new(move |body| { + util::CountBytesBody { + inner: body, + counter: response_bytes_counter.clone(), + } + })) + .into_inner(), + ) + .add_service(svc) + .serve_with_incoming(futures::stream::iter(vec![Ok::<_, std::io::Error>( + MockStream(server), + )])) + .await + .unwrap(); + } + }); + + let mut client = test_client::TestClient::new(mock_io_channel(client).await).accept_gzip(); + + let res = client.compress_output_server_stream(()).await.unwrap(); + + assert_eq!(res.metadata().get("grpc-encoding").unwrap(), "gzip"); + + let mut stream: Streaming = res.into_inner(); + + stream + .next() + .await + .expect("stream empty") + .expect("item was error"); + assert!(response_bytes_counter.load(SeqCst) < UNCOMPRESSED_MIN_BODY_SIZE); + + stream + .next() + .await + .expect("stream empty") + .expect("item was error"); + assert!(response_bytes_counter.load(SeqCst) < UNCOMPRESSED_MIN_BODY_SIZE); +} + +#[tokio::test(flavor = "multi_thread")] +async fn client_disabled_server_enabled() { + let (client, server) = tokio::io::duplex(UNCOMPRESSED_MIN_BODY_SIZE * 10); + + let svc = test_server::TestServer::new(Svc::default()).send_gzip(); + + let response_bytes_counter = Arc::new(AtomicUsize::new(0)); + + tokio::spawn({ + let response_bytes_counter = response_bytes_counter.clone(); + async move { + Server::builder() + .layer( + ServiceBuilder::new() + .layer(MapResponseBodyLayer::new(move |body| { + util::CountBytesBody { + inner: body, + counter: response_bytes_counter.clone(), + } + })) + .into_inner(), + ) + .add_service(svc) + .serve_with_incoming(futures::stream::iter(vec![Ok::<_, std::io::Error>( + MockStream(server), + )])) + .await + .unwrap(); + } + }); + + let mut client = test_client::TestClient::new(mock_io_channel(client).await); + + let res = client.compress_output_server_stream(()).await.unwrap(); + + assert!(res.metadata().get("grpc-encoding").is_none()); + + let mut stream: Streaming = res.into_inner(); + + stream + .next() + .await + .expect("stream empty") + .expect("item was error"); + assert!(response_bytes_counter.load(SeqCst) > UNCOMPRESSED_MIN_BODY_SIZE); +} + +#[tokio::test(flavor = "multi_thread")] +async fn client_enabled_server_disabled() { + let (client, server) = tokio::io::duplex(UNCOMPRESSED_MIN_BODY_SIZE * 10); + + let svc = test_server::TestServer::new(Svc::default()); + + let response_bytes_counter = Arc::new(AtomicUsize::new(0)); + + tokio::spawn({ + let response_bytes_counter = response_bytes_counter.clone(); + async move { + Server::builder() + .layer( + ServiceBuilder::new() + .layer(MapResponseBodyLayer::new(move |body| { + util::CountBytesBody { + inner: body, + counter: response_bytes_counter.clone(), + } + })) + .into_inner(), + ) + .add_service(svc) + .serve_with_incoming(futures::stream::iter(vec![Ok::<_, std::io::Error>( + MockStream(server), + )])) + .await + .unwrap(); + } + }); + + let mut client = test_client::TestClient::new(mock_io_channel(client).await).accept_gzip(); + + let res = client.compress_output_server_stream(()).await.unwrap(); + + assert!(res.metadata().get("grpc-encoding").is_none()); + + let mut stream: Streaming = res.into_inner(); + + stream + .next() + .await + .expect("stream empty") + .expect("item was error"); + assert!(response_bytes_counter.load(SeqCst) > UNCOMPRESSED_MIN_BODY_SIZE); +} diff --git a/tests/compression/src/util.rs b/tests/compression/src/util.rs new file mode 100644 index 0000000..75df07f --- /dev/null +++ b/tests/compression/src/util.rs @@ -0,0 +1,139 @@ +use super::*; +use bytes::Bytes; +use futures::ready; +use http_body::Body; +use pin_project::pin_project; +use std::{ + pin::Pin, + sync::{ + atomic::{AtomicUsize, Ordering::SeqCst}, + Arc, + }, + task::{Context, Poll}, +}; +use tokio::io::{AsyncRead, AsyncWrite, ReadBuf}; +use tonic::transport::{server::Connected, Channel}; +use tower_http::map_request_body::MapRequestBodyLayer; + +/// A body that tracks how many bytes passes through it +#[pin_project] +pub struct CountBytesBody { + #[pin] + pub inner: B, + pub counter: Arc, +} + +impl Body for CountBytesBody +where + B: Body, +{ + type Data = B::Data; + type Error = B::Error; + + fn poll_data( + self: Pin<&mut Self>, + cx: &mut Context<'_>, + ) -> Poll>> { + let this = self.project(); + let counter: Arc = this.counter.clone(); + match ready!(this.inner.poll_data(cx)) { + Some(Ok(chunk)) => { + println!("response body chunk size = {}", chunk.len()); + counter.fetch_add(chunk.len(), SeqCst); + Poll::Ready(Some(Ok(chunk))) + } + x => Poll::Ready(x), + } + } + + fn poll_trailers( + self: Pin<&mut Self>, + cx: &mut Context<'_>, + ) -> Poll, Self::Error>> { + self.project().inner.poll_trailers(cx) + } + + fn is_end_stream(&self) -> bool { + self.inner.is_end_stream() + } + + fn size_hint(&self) -> http_body::SizeHint { + self.inner.size_hint() + } +} + +#[allow(dead_code)] +pub fn measure_request_body_size_layer( + bytes_sent_counter: Arc, +) -> MapRequestBodyLayer hyper::Body + Clone> { + MapRequestBodyLayer::new(move |mut body: hyper::Body| { + let (mut tx, new_body) = hyper::Body::channel(); + + let bytes_sent_counter = bytes_sent_counter.clone(); + tokio::spawn(async move { + while let Some(chunk) = body.data().await { + let chunk = chunk.unwrap(); + println!("request body chunk size = {}", chunk.len()); + bytes_sent_counter.fetch_add(chunk.len(), SeqCst); + tx.send_data(chunk).await.unwrap(); + } + + if let Some(trailers) = body.trailers().await.unwrap() { + tx.send_trailers(trailers).await.unwrap(); + } + }); + + new_body + }) +} + +#[derive(Debug)] +pub struct MockStream(pub tokio::io::DuplexStream); + +impl Connected for MockStream { + type ConnectInfo = (); + + fn connect_info(&self) -> Self::ConnectInfo {} +} + +impl AsyncRead for MockStream { + fn poll_read( + mut self: Pin<&mut Self>, + cx: &mut Context<'_>, + buf: &mut ReadBuf<'_>, + ) -> Poll> { + Pin::new(&mut self.0).poll_read(cx, buf) + } +} + +impl AsyncWrite for MockStream { + fn poll_write( + mut self: Pin<&mut Self>, + cx: &mut Context<'_>, + buf: &[u8], + ) -> Poll> { + Pin::new(&mut self.0).poll_write(cx, buf) + } + + fn poll_flush(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll> { + Pin::new(&mut self.0).poll_flush(cx) + } + + fn poll_shutdown(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll> { + Pin::new(&mut self.0).poll_shutdown(cx) + } +} + +#[allow(dead_code)] +pub async fn mock_io_channel(client: tokio::io::DuplexStream) -> Channel { + let mut client = Some(client); + + Endpoint::try_from("http://[::]:50051") + .unwrap() + .connect_with_connector(service_fn(move |_: Uri| { + let client = client.take().unwrap(); + async move { Ok::<_, std::io::Error>(MockStream(client)) } + })) + .await + .unwrap() +} diff --git a/tonic-build/Cargo.toml b/tonic-build/Cargo.toml index 08af62c..529cd88 100644 --- a/tonic-build/Cargo.toml +++ b/tonic-build/Cargo.toml @@ -26,6 +26,7 @@ default = ["transport", "rustfmt", "prost"] rustfmt = [] transport = [] prost = ["prost-build"] +compression = [] [package.metadata.docs.rs] all-features = true diff --git a/tonic-build/src/client.rs b/tonic-build/src/client.rs index be7d424..b44c1d0 100644 --- a/tonic-build/src/client.rs +++ b/tonic-build/src/client.rs @@ -20,8 +20,6 @@ pub fn generate( let connect = generate_connect(&service_ident); let service_doc = generate_doc_comments(service.comment()); - let struct_debug = format!("{} {{{{ ... }}}}", &service_ident); - quote! { /// Generated client implementations. pub mod #client_mod { @@ -29,6 +27,7 @@ pub fn generate( use tonic::codegen::*; #service_doc + #[derive(Debug, Clone)] pub struct #service_ident { inner: tonic::client::Grpc, } @@ -59,22 +58,23 @@ pub fn generate( #service_ident::new(InterceptedService::new(inner, interceptor)) } + /// Compress requests with `gzip`. + /// + /// This requires the server to support it otherwise it might respond with an + /// error. + pub fn send_gzip(mut self) -> Self { + self.inner = self.inner.send_gzip(); + self + } + + /// Enable decompressing responses with `gzip`. + pub fn accept_gzip(mut self) -> Self { + self.inner = self.inner.accept_gzip(); + self + } + #methods } - - impl Clone for #service_ident { - fn clone(&self) -> Self { - Self { - inner: self.inner.clone(), - } - } - } - - impl std::fmt::Debug for #service_ident { - fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { - write!(f, #struct_debug) - } - } } } } @@ -153,10 +153,10 @@ fn generate_unary( &mut self, request: impl tonic::IntoRequest<#request>, ) -> Result, tonic::Status> { - self.inner.ready().await.map_err(|e| { - tonic::Status::new(tonic::Code::Unknown, format!("Service was not ready: {}", e.into())) - })?; - let codec = #codec_name::default(); + self.inner.ready().await.map_err(|e| { + tonic::Status::new(tonic::Code::Unknown, format!("Service was not ready: {}", e.into())) + })?; + let codec = #codec_name::default(); let path = http::uri::PathAndQuery::from_static(#path); self.inner.unary(request.into_request(), path, codec).await } @@ -204,7 +204,7 @@ fn generate_client_streaming( pub async fn #ident( &mut self, request: impl tonic::IntoStreamingRequest - ) -> Result, tonic::Status> { + ) -> Result, tonic::Status> where T: std::fmt::Debug { self.inner.ready().await.map_err(|e| { tonic::Status::new(tonic::Code::Unknown, format!("Service was not ready: {}", e.into())) })?; diff --git a/tonic-build/src/server.rs b/tonic-build/src/server.rs index 5179175..f896243 100644 --- a/tonic-build/src/server.rs +++ b/tonic-build/src/server.rs @@ -36,6 +36,32 @@ pub fn generate( ); let transport = generate_transport(&server_service, &server_trait, &path); + let compression_enabled = cfg!(feature = "compression"); + + let compression_config_ty = if compression_enabled { + quote! { EnabledCompressionEncodings } + } else { + quote! { () } + }; + + let configure_compression_methods = if compression_enabled { + quote! { + /// Enable decompressing requests with `gzip`. + pub fn accept_gzip(mut self) -> Self { + self.accept_compression_encodings.enable_gzip(); + self + } + + /// Compress responses with `gzip`, if the client supports it. + pub fn send_gzip(mut self) -> Self { + self.send_compression_encodings.enable_gzip(); + self + } + } + } else { + quote! {} + }; + quote! { /// Generated server implementations. pub mod #server_mod { @@ -48,6 +74,8 @@ pub fn generate( #[derive(Debug)] pub struct #server_service { inner: _Inner, + accept_compression_encodings: #compression_config_ty, + send_compression_encodings: #compression_config_ty, } struct _Inner(Arc); @@ -56,7 +84,11 @@ pub fn generate( pub fn new(inner: T) -> Self { let inner = Arc::new(inner); let inner = _Inner(inner); - Self { inner } + Self { + inner, + accept_compression_encodings: Default::default(), + send_compression_encodings: Default::default(), + } } pub fn with_interceptor(inner: T, interceptor: F) -> InterceptedService @@ -65,6 +97,8 @@ pub fn generate( { InterceptedService::new(Self::new(inner), interceptor) } + + #configure_compression_methods } impl Service> for #server_service @@ -102,7 +136,11 @@ pub fn generate( impl Clone for #server_service { fn clone(&self) -> Self { let inner = self.inner.clone(); - Self { inner } + Self { + inner, + accept_compression_encodings: self.accept_compression_encodings, + send_compression_encodings: self.send_compression_encodings, + } } } @@ -335,13 +373,16 @@ fn generate_unary( } } + let accept_compression_encodings = self.accept_compression_encodings; + let send_compression_encodings = self.send_compression_encodings; let inner = self.inner.clone(); let fut = async move { let inner = inner.0; let method = #service_ident(inner); let codec = #codec_name::default(); - let mut grpc = tonic::server::Grpc::new(codec); + let mut grpc = tonic::server::Grpc::new(codec) + .apply_compression_config(accept_compression_encodings, send_compression_encodings); let res = grpc.unary(method, req).await; Ok(res) @@ -379,19 +420,21 @@ fn generate_server_streaming( let inner = self.0.clone(); let fut = async move { (*inner).#method_ident(request).await - }; Box::pin(fut) } } + let accept_compression_encodings = self.accept_compression_encodings; + let send_compression_encodings = self.send_compression_encodings; let inner = self.inner.clone(); let fut = async move { let inner = inner.0; let method = #service_ident(inner); let codec = #codec_name::default(); - let mut grpc = tonic::server::Grpc::new(codec); + let mut grpc = tonic::server::Grpc::new(codec) + .apply_compression_config(accept_compression_encodings, send_compression_encodings); let res = grpc.server_streaming(method, req).await; Ok(res) @@ -432,13 +475,16 @@ fn generate_client_streaming( } } + let accept_compression_encodings = self.accept_compression_encodings; + let send_compression_encodings = self.send_compression_encodings; let inner = self.inner.clone(); let fut = async move { let inner = inner.0; let method = #service_ident(inner); let codec = #codec_name::default(); - let mut grpc = tonic::server::Grpc::new(codec); + let mut grpc = tonic::server::Grpc::new(codec) + .apply_compression_config(accept_compression_encodings, send_compression_encodings); let res = grpc.client_streaming(method, req).await; Ok(res) @@ -482,13 +528,16 @@ fn generate_streaming( } } + let accept_compression_encodings = self.accept_compression_encodings; + let send_compression_encodings = self.send_compression_encodings; let inner = self.inner.clone(); let fut = async move { let inner = inner.0; let method = #service_ident(inner); let codec = #codec_name::default(); - let mut grpc = tonic::server::Grpc::new(codec); + let mut grpc = tonic::server::Grpc::new(codec) + .apply_compression_config(accept_compression_encodings, send_compression_encodings); let res = grpc.streaming(method, req).await; Ok(res) diff --git a/tonic/Cargo.toml b/tonic/Cargo.toml index a283727..134dc85 100644 --- a/tonic/Cargo.toml +++ b/tonic/Cargo.toml @@ -40,6 +40,7 @@ tls-roots-common = ["tls"] tls-roots = ["tls-roots-common", "rustls-native-certs"] tls-webpki-roots = ["tls-roots-common", "webpki-roots"] prost = ["prost1", "prost-derive"] +compression = ["flate2"] # [[bench]] # name = "bench_main" @@ -82,6 +83,9 @@ tokio-rustls = { version = "0.22", optional = true } rustls-native-certs = { version = "0.5", optional = true } webpki-roots = { version = "0.21.1", optional = true } +# compression +flate2 = { version = "1.0", optional = true } + [dev-dependencies] tokio = { version = "1.0", features = ["rt", "macros"] } static_assertions = "1.0" diff --git a/tonic/benches/decode.rs b/tonic/benches/decode.rs index 41b2496..96f5b49 100644 --- a/tonic/benches/decode.rs +++ b/tonic/benches/decode.rs @@ -22,7 +22,7 @@ macro_rules! bench { b.iter(|| { rt.block_on(async { let decoder = MockDecoder::new($message_size); - let mut stream = Streaming::new_request(decoder, body.clone()); + let mut stream = Streaming::new_request(decoder, body.clone(), None); let mut count = 0; while let Some(msg) = stream.message().await.unwrap() { diff --git a/tonic/src/client/grpc.rs b/tonic/src/client/grpc.rs index 72a803e..c89c4c6 100644 --- a/tonic/src/client/grpc.rs +++ b/tonic/src/client/grpc.rs @@ -1,3 +1,5 @@ +#[cfg(feature = "compression")] +use crate::codec::compression::{CompressionEncoding, EnabledCompressionEncodings}; use crate::{ body::BoxBody, client::GrpcService, @@ -28,12 +30,102 @@ use std::fmt; /// [gRPC protocol definition]: https://github.com/grpc/grpc/blob/master/doc/PROTOCOL-HTTP2.md#requests pub struct Grpc { inner: T, + #[cfg(feature = "compression")] + /// Which compression encodings does the client accept? + accept_compression_encodings: EnabledCompressionEncodings, + #[cfg(feature = "compression")] + /// The compression encoding that will be applied to requests. + send_compression_encodings: Option, } impl Grpc { /// Creates a new gRPC client with the provided [`GrpcService`]. pub fn new(inner: T) -> Self { - Self { inner } + Self { + inner, + #[cfg(feature = "compression")] + send_compression_encodings: None, + #[cfg(feature = "compression")] + accept_compression_encodings: EnabledCompressionEncodings::default(), + } + } + + /// Compress requests with `gzip`. + /// + /// Requires the server to accept `gzip` otherwise it might return an error. + /// + /// # Example + /// + /// The most common way of using this is through a client generated by tonic-build: + /// + /// ```rust + /// use tonic::transport::Channel; + /// # struct TestClient(T); + /// # impl TestClient { + /// # fn new(channel: T) -> Self { Self(channel) } + /// # fn send_gzip(self) -> Self { self } + /// # } + /// + /// # async { + /// let channel = Channel::builder("127.0.0.1:3000".parse().unwrap()) + /// .connect() + /// .await + /// .unwrap(); + /// + /// let client = TestClient::new(channel).send_gzip(); + /// # }; + /// ``` + #[cfg(feature = "compression")] + #[cfg_attr(docsrs, doc(cfg(feature = "compression")))] + pub fn send_gzip(mut self) -> Self { + self.send_compression_encodings = Some(CompressionEncoding::Gzip); + self + } + + #[doc(hidden)] + #[cfg(not(feature = "compression"))] + pub fn send_gzip(self) -> Self { + panic!( + "`send_gzip` called on a client but the `compression` feature is not enabled on tonic" + ); + } + + /// Enable accepting `gzip` compressed responses. + /// + /// Requires the server to also support sending compressed responses. + /// + /// # Example + /// + /// The most common way of using this is through a client generated by tonic-build: + /// + /// ```rust + /// use tonic::transport::Channel; + /// # struct TestClient(T); + /// # impl TestClient { + /// # fn new(channel: T) -> Self { Self(channel) } + /// # fn accept_gzip(self) -> Self { self } + /// # } + /// + /// # async { + /// let channel = Channel::builder("127.0.0.1:3000".parse().unwrap()) + /// .connect() + /// .await + /// .unwrap(); + /// + /// let client = TestClient::new(channel).accept_gzip(); + /// # }; + /// ``` + #[cfg(feature = "compression")] + #[cfg_attr(docsrs, doc(cfg(feature = "compression")))] + pub fn accept_gzip(mut self) -> Self { + self.accept_compression_encodings.enable_gzip(); + self + } + + #[doc(hidden)] + #[cfg(not(feature = "compression"))] + pub fn accept_gzip(self) -> Self { + panic!("`accept_gzip` called on a client but the `compression` feature is not enabled on tonic"); } /// Check if the inner [`GrpcService`] is able to accept a new request. @@ -145,7 +237,14 @@ impl Grpc { let uri = Uri::from_parts(parts).expect("path_and_query only is valid Uri"); let request = request - .map(|s| encode_client(codec.encoder(), s)) + .map(|s| { + encode_client( + codec.encoder(), + s, + #[cfg(feature = "compression")] + self.send_compression_encodings, + ) + }) .map(BoxBody::new); let mut request = request.into_http(uri); @@ -160,12 +259,38 @@ impl Grpc { .headers_mut() .insert(CONTENT_TYPE, HeaderValue::from_static("application/grpc")); + #[cfg(feature = "compression")] + { + if let Some(encoding) = self.send_compression_encodings { + request.headers_mut().insert( + crate::codec::compression::ENCODING_HEADER, + encoding.into_header_value(), + ); + } + + if let Some(header_value) = self + .accept_compression_encodings + .into_accept_encoding_header_value() + { + request.headers_mut().insert( + crate::codec::compression::ACCEPT_ENCODING_HEADER, + header_value, + ); + } + } + let response = self .inner .call(request) .await .map_err(|err| Status::from_error(err.into()))?; + #[cfg(feature = "compression")] + let encoding = CompressionEncoding::from_encoding_header( + response.headers(), + self.accept_compression_encodings, + )?; + let status_code = response.status(); let trailers_only_status = Status::from_header_map(response.headers()); @@ -183,7 +308,13 @@ impl Grpc { let response = response.map(|body| { if expect_additional_trailers { - Streaming::new_response(codec.decoder(), body, status_code) + Streaming::new_response( + codec.decoder(), + body, + status_code, + #[cfg(feature = "compression")] + encoding, + ) } else { Streaming::new_empty(codec.decoder(), body) } @@ -197,12 +328,29 @@ impl Clone for Grpc { fn clone(&self) -> Self { Self { inner: self.inner.clone(), + #[cfg(feature = "compression")] + send_compression_encodings: self.send_compression_encodings, + #[cfg(feature = "compression")] + accept_compression_encodings: self.accept_compression_encodings, } } } impl fmt::Debug for Grpc { fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { - f.debug_struct("Grpc").field("inner", &self.inner).finish() + let mut f = f.debug_struct("Grpc"); + + f.field("inner", &self.inner); + + #[cfg(feature = "compression")] + f.field("compression_encoding", &self.send_compression_encodings); + + #[cfg(feature = "compression")] + f.field( + "accept_compression_encodings", + &self.accept_compression_encodings, + ); + + f.finish() } } diff --git a/tonic/src/codec/compression.rs b/tonic/src/codec/compression.rs new file mode 100644 index 0000000..8f4c279 --- /dev/null +++ b/tonic/src/codec/compression.rs @@ -0,0 +1,189 @@ +use super::encode::BUFFER_SIZE; +use crate::{metadata::MetadataValue, Status}; +use bytes::{Buf, BufMut, BytesMut}; +use flate2::read::{GzDecoder, GzEncoder}; +use std::fmt; + +pub(crate) const ENCODING_HEADER: &str = "grpc-encoding"; +pub(crate) const ACCEPT_ENCODING_HEADER: &str = "grpc-accept-encoding"; + +/// Struct used to configure which encodings are enabled on a server or channel. +#[derive(Debug, Default, Clone, Copy)] +pub struct EnabledCompressionEncodings { + pub(crate) gzip: bool, +} + +impl EnabledCompressionEncodings { + /// Check if `gzip` compression is enabled. + pub fn gzip(self) -> bool { + self.gzip + } + + /// Enable `gzip` compression. + pub fn enable_gzip(&mut self) { + self.gzip = true; + } + + pub(crate) fn into_accept_encoding_header_value(self) -> Option { + let Self { gzip } = self; + if gzip { + Some(http::HeaderValue::from_static("gzip,identity")) + } else { + None + } + } +} + +/// The compression encodings Tonic supports. +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +#[non_exhaustive] +pub enum CompressionEncoding { + #[allow(missing_docs)] + Gzip, +} + +impl CompressionEncoding { + /// Based on the `grpc-accept-encoding` header, pick an encoding to use. + pub(crate) fn from_accept_encoding_header( + map: &http::HeaderMap, + enabled_encodings: EnabledCompressionEncodings, + ) -> Option { + let header_value = map.get(ACCEPT_ENCODING_HEADER)?; + let header_value_str = header_value.to_str().ok()?; + + let EnabledCompressionEncodings { gzip } = enabled_encodings; + + split_by_comma(header_value_str).find_map(|value| match value { + "gzip" if gzip => Some(CompressionEncoding::Gzip), + _ => None, + }) + } + + /// Get the value of `grpc-encoding` header. Returns an error if the encoding isn't supported. + pub(crate) fn from_encoding_header( + map: &http::HeaderMap, + enabled_encodings: EnabledCompressionEncodings, + ) -> Result, Status> { + let header_value = if let Some(value) = map.get(ENCODING_HEADER) { + value + } else { + return Ok(None); + }; + + let header_value_str = if let Ok(value) = header_value.to_str() { + value + } else { + return Ok(None); + }; + + let EnabledCompressionEncodings { gzip } = enabled_encodings; + + match header_value_str { + "gzip" if gzip => Ok(Some(CompressionEncoding::Gzip)), + other => { + let mut status = Status::unimplemented(format!( + "Content is compressed with `{}` which isn't supported", + other + )); + + let header_value = enabled_encodings + .into_accept_encoding_header_value() + .map(MetadataValue::unchecked_from_header_value) + .unwrap_or_else(|| MetadataValue::from_static("identity")); + status + .metadata_mut() + .insert(ACCEPT_ENCODING_HEADER, header_value); + + Err(status) + } + } + } + + pub(crate) fn into_header_value(self) -> http::HeaderValue { + match self { + CompressionEncoding::Gzip => http::HeaderValue::from_static("gzip"), + } + } +} + +impl fmt::Display for CompressionEncoding { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + match self { + CompressionEncoding::Gzip => write!(f, "gzip"), + } + } +} + +fn split_by_comma(s: &str) -> impl Iterator { + s.trim().split(',').map(|s| s.trim()) +} + +/// Compress `len` bytes from `decompressed_buf` into `out_buf`. +pub(crate) fn compress( + encoding: CompressionEncoding, + decompressed_buf: &mut BytesMut, + out_buf: &mut BytesMut, + len: usize, +) -> Result<(), std::io::Error> { + let capacity = ((len / BUFFER_SIZE) + 1) * BUFFER_SIZE; + out_buf.reserve(capacity); + + match encoding { + CompressionEncoding::Gzip => { + let mut gzip_encoder = GzEncoder::new( + &decompressed_buf[0..len], + // FIXME: support customizing the compression level + flate2::Compression::new(6), + ); + let mut out_writer = out_buf.writer(); + + std::io::copy(&mut gzip_encoder, &mut out_writer)?; + } + } + + decompressed_buf.advance(len); + + Ok(()) +} + +/// Decompress `len` bytes from `compressed_buf` into `out_buf`. +pub(crate) fn decompress( + encoding: CompressionEncoding, + compressed_buf: &mut BytesMut, + out_buf: &mut BytesMut, + len: usize, +) -> Result<(), std::io::Error> { + let estimate_decompressed_len = len * 2; + let capacity = ((estimate_decompressed_len / BUFFER_SIZE) + 1) * BUFFER_SIZE; + out_buf.reserve(capacity); + + match encoding { + CompressionEncoding::Gzip => { + let mut gzip_decoder = GzDecoder::new(&compressed_buf[0..len]); + let mut out_writer = out_buf.writer(); + + std::io::copy(&mut gzip_decoder, &mut out_writer)?; + } + } + + compressed_buf.advance(len); + + Ok(()) +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub(crate) enum SingleMessageCompressionOverride { + /// Inherit whatever compression is already configured. If the stream is compressed this + /// message will also be configured. + /// + /// This is the default. + Inherit, + /// Don't compress this message, even if compression is enabled on the stream. + Disable, +} + +impl Default for SingleMessageCompressionOverride { + fn default() -> Self { + Self::Inherit + } +} diff --git a/tonic/src/codec/decode.rs b/tonic/src/codec/decode.rs index 6977efd..6431360 100644 --- a/tonic/src/codec/decode.rs +++ b/tonic/src/codec/decode.rs @@ -1,4 +1,6 @@ -use super::{DecodeBuf, Decoder}; +#[cfg(feature = "compression")] +use super::compression::{decompress, CompressionEncoding}; +use super::{DecodeBuf, Decoder, HEADER_SIZE}; use crate::{body::BoxBody, metadata::MetadataMap, Code, Status}; use bytes::{Buf, BufMut, BytesMut}; use futures_core::Stream; @@ -25,6 +27,10 @@ pub struct Streaming { direction: Direction, buf: BytesMut, trailers: Option, + #[cfg(feature = "compression")] + decompress_buf: BytesMut, + #[cfg(feature = "compression")] + encoding: Option, } impl Unpin for Streaming {} @@ -43,13 +49,24 @@ enum Direction { } impl Streaming { - pub(crate) fn new_response(decoder: D, body: B, status_code: StatusCode) -> Self + pub(crate) fn new_response( + decoder: D, + body: B, + status_code: StatusCode, + #[cfg(feature = "compression")] encoding: Option, + ) -> Self where B: Body + Send + Sync + 'static, B::Error: Into, D: Decoder + Send + Sync + 'static, { - Self::new(decoder, body, Direction::Response(status_code)) + Self::new( + decoder, + body, + Direction::Response(status_code), + #[cfg(feature = "compression")] + encoding, + ) } pub(crate) fn new_empty(decoder: D, body: B) -> Self @@ -58,20 +75,41 @@ impl Streaming { B::Error: Into, D: Decoder + Send + Sync + 'static, { - Self::new(decoder, body, Direction::EmptyResponse) + Self::new( + decoder, + body, + Direction::EmptyResponse, + #[cfg(feature = "compression")] + None, + ) } #[doc(hidden)] - pub fn new_request(decoder: D, body: B) -> Self + pub fn new_request( + decoder: D, + body: B, + #[cfg(feature = "compression")] encoding: Option, + ) -> Self where B: Body + Send + Sync + 'static, B::Error: Into, D: Decoder + Send + Sync + 'static, { - Self::new(decoder, body, Direction::Request) + Self::new( + decoder, + body, + Direction::Request, + #[cfg(feature = "compression")] + encoding, + ) } - fn new(decoder: D, body: B, direction: Direction) -> Self + fn new( + decoder: D, + body: B, + direction: Direction, + #[cfg(feature = "compression")] encoding: Option, + ) -> Self where B: Body + Send + Sync + 'static, B::Error: Into, @@ -87,6 +125,10 @@ impl Streaming { direction, buf: BytesMut::with_capacity(BUFFER_SIZE), trailers: None, + #[cfg(feature = "compression")] + decompress_buf: BytesMut::new(), + #[cfg(feature = "compression")] + encoding, } } } @@ -156,18 +198,21 @@ impl Streaming { fn decode_chunk(&mut self) -> Result, Status> { if let State::ReadHeader = self.state { - if self.buf.remaining() < 5 { + if self.buf.remaining() < HEADER_SIZE { return Ok(None); } let is_compressed = match self.buf.get_u8() { 0 => false, 1 => { - trace!("message compressed, compression not supported yet"); - return Err(Status::new( - Code::Unimplemented, - "Message compressed, compression not supported yet.".to_string(), - )); + if cfg!(feature = "compression") { + true + } else { + return Err(Status::new( + Code::Unimplemented, + "Message compressed, compression support not enabled.".to_string(), + )); + } } f => { trace!("unexpected compression flag"); @@ -191,17 +236,51 @@ impl Streaming { } } - if let State::ReadBody { len, .. } = &self.state { + if let State::ReadBody { len, compression } = &self.state { // if we haven't read enough of the message then return and keep // reading if self.buf.remaining() < *len || self.buf.len() < *len { return Ok(None); } - return match self - .decoder - .decode(&mut DecodeBuf::new(&mut self.buf, *len)) - { + let decoding_result = if *compression { + #[cfg(feature = "compression")] + { + self.decompress_buf.clear(); + + if let Err(err) = decompress( + self.encoding.unwrap_or_else(|| { + unreachable!("message was compressed but `Streaming.encoding` was `None`. This is a bug in Tonic. Please file an issue") + }), + &mut self.buf, + &mut self.decompress_buf, + *len, + ) { + let message = if let Direction::Response(status) = self.direction { + format!( + "Error decompressing: {}, while receiving response with status: {}", + err, status + ) + } else { + format!("Error decompressing: {}, while sending request", err) + }; + return Err(Status::new(Code::Internal, message)); + } + let decompressed_len = self.decompress_buf.len(); + self.decoder.decode(&mut DecodeBuf::new( + &mut self.decompress_buf, + decompressed_len, + )) + } + + #[cfg(not(feature = "compression"))] + unreachable!("should not take this branch if compression is disabled") + } else { + self.decoder + .decode(&mut DecodeBuf::new(&mut self.buf, *len)) + }; + + return match decoding_result { Ok(Some(msg)) => { self.state = State::ReadHeader; Ok(Some(msg)) diff --git a/tonic/src/codec/encode.rs b/tonic/src/codec/encode.rs index a42c195..54bf64c 100644 --- a/tonic/src/codec/encode.rs +++ b/tonic/src/codec/encode.rs @@ -1,4 +1,6 @@ -use super::{EncodeBuf, Encoder}; +#[cfg(feature = "compression")] +use super::compression::{compress, CompressionEncoding, SingleMessageCompressionOverride}; +use super::{EncodeBuf, Encoder, HEADER_SIZE}; use crate::{Code, Status}; use bytes::{BufMut, Bytes, BytesMut}; use futures_core::{Stream, TryStream}; @@ -11,62 +13,124 @@ use std::{ task::{Context, Poll}, }; -const BUFFER_SIZE: usize = 8 * 1024; +pub(super) const BUFFER_SIZE: usize = 8 * 1024; pub(crate) fn encode_server( encoder: T, source: U, + #[cfg(feature = "compression")] compression_encoding: Option, + #[cfg(feature = "compression")] compression_override: SingleMessageCompressionOverride, ) -> EncodeBody>> where T: Encoder + Send + Sync + 'static, T::Item: Send + Sync, U: Stream> + Send + Sync + 'static, { - let stream = encode(encoder, source).into_stream(); + let stream = encode( + encoder, + source, + #[cfg(feature = "compression")] + compression_encoding, + #[cfg(feature = "compression")] + compression_override, + ) + .into_stream(); + EncodeBody::new_server(stream) } pub(crate) fn encode_client( encoder: T, source: U, + #[cfg(feature = "compression")] compression_encoding: Option, ) -> EncodeBody>> where T: Encoder + Send + Sync + 'static, T::Item: Send + Sync, U: Stream + Send + Sync + 'static, { - let stream = encode(encoder, source.map(Ok)).into_stream(); + let stream = encode( + encoder, + source.map(Ok), + #[cfg(feature = "compression")] + compression_encoding, + #[cfg(feature = "compression")] + SingleMessageCompressionOverride::default(), + ) + .into_stream(); EncodeBody::new_client(stream) } -fn encode(mut encoder: T, source: U) -> impl TryStream +fn encode( + mut encoder: T, + source: U, + #[cfg(feature = "compression")] compression_encoding: Option, + #[cfg(feature = "compression")] compression_override: SingleMessageCompressionOverride, +) -> impl TryStream where T: Encoder, U: Stream>, { async_stream::stream! { let mut buf = BytesMut::with_capacity(BUFFER_SIZE); + + #[cfg(feature = "compression")] + let (compression_enabled_for_stream, mut uncompression_buf) = match compression_encoding { + Some(CompressionEncoding::Gzip) => (true, BytesMut::with_capacity(BUFFER_SIZE)), + None => (false, BytesMut::new()), + }; + + #[cfg(feature = "compression")] + let compress_item = compression_enabled_for_stream && compression_override == SingleMessageCompressionOverride::Inherit; + + #[cfg(not(feature = "compression"))] + let compress_item = false; + futures_util::pin_mut!(source); loop { match source.next().await { Some(Ok(item)) => { - buf.reserve(5); + buf.reserve(HEADER_SIZE); unsafe { - buf.advance_mut(5); + buf.advance_mut(HEADER_SIZE); + } + + if compress_item { + #[cfg(feature = "compression")] + { + uncompression_buf.clear(); + + encoder.encode(item, &mut EncodeBuf::new(&mut uncompression_buf)) + .map_err(|err| Status::internal(format!("Error encoding: {}", err)))?; + + let uncompressed_len = uncompression_buf.len(); + + compress( + compression_encoding.unwrap(), + &mut uncompression_buf, + &mut buf, + uncompressed_len, + ).map_err(|err| Status::internal(format!("Error compressing: {}", err)))?; + } + + #[cfg(not(feature = "compression"))] + unreachable!("compression disabled, should not take this branch"); + } else { + encoder.encode(item, &mut EncodeBuf::new(&mut buf)) + .map_err(|err| Status::internal(format!("Error encoding: {}", err)))?; } - encoder.encode(item, &mut EncodeBuf::new(&mut buf)).map_err(drop).unwrap(); // now that we know length, we can write the header - let len = buf.len() - 5; + let len = buf.len() - HEADER_SIZE; assert!(len <= std::u32::MAX as usize); { - let mut buf = &mut buf[..5]; - buf.put_u8(0); // byte must be 0, reserve doesn't auto-zero + let mut buf = &mut buf[..HEADER_SIZE]; + buf.put_u8(compress_item as u8); buf.put_u32(len as u32); } - yield Ok(buf.split_to(len + 5).freeze()); + yield Ok(buf.split_to(len + HEADER_SIZE).freeze()); }, Some(Err(status)) => yield Err(status), None => break, diff --git a/tonic/src/codec/mod.rs b/tonic/src/codec/mod.rs index e100556..d0c9d8b 100644 --- a/tonic/src/codec/mod.rs +++ b/tonic/src/codec/mod.rs @@ -4,20 +4,33 @@ //! and a protobuf codec based on prost. mod buffer; +#[cfg(feature = "compression")] +pub(crate) mod compression; mod decode; mod encode; #[cfg(feature = "prost")] mod prost; +use crate::Status; use std::io; -pub use self::decode::Streaming; pub(crate) use self::encode::{encode_client, encode_server}; + +pub use self::buffer::{DecodeBuf, EncodeBuf}; +#[cfg(feature = "compression")] +#[cfg_attr(docsrs, doc(cfg(feature = "compression")))] +pub use self::compression::{CompressionEncoding, EnabledCompressionEncodings}; +pub use self::decode::Streaming; #[cfg(feature = "prost")] #[cfg_attr(docsrs, doc(cfg(feature = "prost")))] pub use self::prost::ProstCodec; -use crate::Status; -pub use buffer::{DecodeBuf, EncodeBuf}; + +// 5 bytes +const HEADER_SIZE: usize = + // compression flag + std::mem::size_of::() + + // data length + std::mem::size_of::(); /// Trait that knows how to encode and decode gRPC messages. pub trait Codec: Default { diff --git a/tonic/src/codec/prost.rs b/tonic/src/codec/prost.rs index ddfa93e..0db3ee0 100644 --- a/tonic/src/codec/prost.rs +++ b/tonic/src/codec/prost.rs @@ -77,7 +77,10 @@ fn from_decode_error(error: prost1::DecodeError) -> crate::Status { #[cfg(test)] mod tests { - use crate::codec::{encode_server, DecodeBuf, Decoder, EncodeBuf, Encoder, Streaming}; + use crate::codec::compression::SingleMessageCompressionOverride; + use crate::codec::{ + encode_server, DecodeBuf, Decoder, EncodeBuf, Encoder, Streaming, HEADER_SIZE, + }; use crate::Status; use bytes::{Buf, BufMut, BytesMut}; use http_body::Body; @@ -92,7 +95,7 @@ mod tests { let mut buf = BytesMut::new(); - buf.reserve(msg.len() + 5); + buf.reserve(msg.len() + HEADER_SIZE); buf.put_u8(0); buf.put_u32(msg.len() as u32); @@ -100,7 +103,7 @@ mod tests { let body = body::MockBody::new(&buf[..], 10005, 0); - let mut stream = Streaming::new_request(decoder, body); + let mut stream = Streaming::new_request(decoder, body, None); let mut i = 0usize; while let Some(output_msg) = stream.message().await.unwrap() { @@ -119,7 +122,12 @@ mod tests { let messages = std::iter::repeat_with(move || Ok::<_, Status>(msg.clone())).take(10000); let source = futures_util::stream::iter(messages); - let body = encode_server(encoder, source); + let body = encode_server( + encoder, + source, + None, + SingleMessageCompressionOverride::default(), + ); futures_util::pin_mut!(body); @@ -216,6 +224,7 @@ mod tests { } } + #[allow(clippy::drop_ref)] fn poll_trailers( self: Pin<&mut Self>, cx: &mut Context<'_>, diff --git a/tonic/src/codegen.rs b/tonic/src/codegen.rs index dd83f2c..9d3a069 100644 --- a/tonic/src/codegen.rs +++ b/tonic/src/codegen.rs @@ -10,6 +10,8 @@ pub use std::sync::Arc; pub use std::task::{Context, Poll}; pub use tower_service::Service; pub type StdError = Box; +#[cfg(feature = "compression")] +pub use crate::codec::{CompressionEncoding, EnabledCompressionEncodings}; pub use crate::service::interceptor::InterceptedService; pub use http_body::Body; diff --git a/tonic/src/lib.rs b/tonic/src/lib.rs index 4d91643..7fb77a6 100644 --- a/tonic/src/lib.rs +++ b/tonic/src/lib.rs @@ -28,6 +28,9 @@ //! - `tls-webpki-roots`: Add the standard trust roots from the `webpki-roots` crate to //! `rustls`-based gRPC clients. Not enabled by default. //! - `prost`: Enables the [`prost`] based gRPC [`Codec`] implementation. +//! - `compression`: Enables compressing requests, responses, and streams. Note +//! that you must enable the `compression` feature on both `tonic` and +//! `tonic-build` to use it. Depends on [flate2]. Not enabled by default. //! //! # Structure //! @@ -62,6 +65,7 @@ //! [`rustls`]: https://docs.rs/rustls //! [`client`]: client/index.html //! [`transport`]: transport/index.html +//! [flate2]: https://crates.io/crates/flate2 #![recursion_limit = "256"] #![allow(clippy::inconsistent_struct_constructor)] diff --git a/tonic/src/metadata/map.rs b/tonic/src/metadata/map.rs index 461a9fb..197976d 100644 --- a/tonic/src/metadata/map.rs +++ b/tonic/src/metadata/map.rs @@ -200,12 +200,11 @@ pub(crate) const GRPC_TIMEOUT_HEADER: &str = "grpc-timeout"; impl MetadataMap { // Headers reserved by the gRPC protocol. - pub(crate) const GRPC_RESERVED_HEADERS: [&'static str; 7] = [ + pub(crate) const GRPC_RESERVED_HEADERS: [&'static str; 6] = [ "te", "user-agent", "content-type", "grpc-message", - "grpc-encoding", "grpc-message-type", "grpc-status", ]; diff --git a/tonic/src/response.rs b/tonic/src/response.rs index 87f59b4..89fc987 100644 --- a/tonic/src/response.rs +++ b/tonic/src/response.rs @@ -107,6 +107,21 @@ impl Response { pub fn extensions_mut(&mut self) -> &mut Extensions { &mut self.extensions } + + /// Disable compression of the response body. + /// + /// This disables compression of the body of this response, even if compression is enabled on + /// the server. + /// + /// **Note**: This only has effect on responses to unary requests and responses to client to + /// server streams. Response streams (server to client stream and bidirectional streams) will + /// still be compressed according to the configuration of the server. + #[cfg(feature = "compression")] + #[cfg_attr(docsrs, doc(cfg(feature = "compression")))] + pub fn disable_compression(&mut self) { + self.extensions_mut() + .insert(crate::codec::compression::SingleMessageCompressionOverride::Disable); + } } #[cfg(test)] diff --git a/tonic/src/server/grpc.rs b/tonic/src/server/grpc.rs index e640ac8..7978e2b 100644 --- a/tonic/src/server/grpc.rs +++ b/tonic/src/server/grpc.rs @@ -1,3 +1,7 @@ +#[cfg(feature = "compression")] +use crate::codec::compression::{ + CompressionEncoding, EnabledCompressionEncodings, SingleMessageCompressionOverride, +}; use crate::{ body::BoxBody, codec::{encode_server, Codec, Streaming}, @@ -9,6 +13,15 @@ use futures_util::{future, stream, TryStreamExt}; use http_body::Body; use std::fmt; +macro_rules! t { + ($result:expr) => { + match $result { + Ok(value) => value, + Err(status) => return status.to_http(), + } + }; +} + /// A gRPC Server handler. /// /// This will wrap some inner [`Codec`] and provide utilities to handle @@ -20,6 +33,12 @@ use std::fmt; /// implements some [`Body`]. pub struct Grpc { codec: T, + /// Which compression encodings does the server accept for requests? + #[cfg(feature = "compression")] + accept_compression_encodings: EnabledCompressionEncodings, + /// Which compression encodings might the server use for responses. + #[cfg(feature = "compression")] + send_compression_encodings: EnabledCompressionEncodings, } impl Grpc @@ -29,7 +48,121 @@ where { /// Creates a new gRPC server with the provided [`Codec`]. pub fn new(codec: T) -> Self { - Self { codec } + Self { + codec, + #[cfg(feature = "compression")] + accept_compression_encodings: EnabledCompressionEncodings::default(), + #[cfg(feature = "compression")] + send_compression_encodings: EnabledCompressionEncodings::default(), + } + } + + /// Enable accepting `gzip` compressed requests. + /// + /// If a request with an unsupported encoding is received the server will respond with + /// [`Code::UnUnimplemented`](crate::Code). + /// + /// # Example + /// + /// The most common way of using this is through a server generated by tonic-build: + /// + /// ```rust + /// # struct Svc; + /// # struct ExampleServer(T); + /// # impl ExampleServer { + /// # fn new(svc: T) -> Self { Self(svc) } + /// # fn accept_gzip(self) -> Self { self } + /// # } + /// # #[tonic::async_trait] + /// # trait Example {} + /// + /// #[tonic::async_trait] + /// impl Example for Svc { + /// // ... + /// } + /// + /// let service = ExampleServer::new(Svc).accept_gzip(); + /// ``` + #[cfg(feature = "compression")] + #[cfg_attr(docsrs, doc(cfg(feature = "compression")))] + pub fn accept_gzip(mut self) -> Self { + self.accept_compression_encodings.enable_gzip(); + self + } + + #[doc(hidden)] + #[cfg(not(feature = "compression"))] + pub fn accept_gzip(self) -> Self { + panic!("`accept_gzip` called on a server but the `compression` feature is not enabled on tonic"); + } + + /// Enable sending `gzip` compressed responses. + /// + /// Requires the client to also support receiving compressed responses. + /// + /// # Example + /// + /// The most common way of using this is through a server generated by tonic-build: + /// + /// ```rust + /// # struct Svc; + /// # struct ExampleServer(T); + /// # impl ExampleServer { + /// # fn new(svc: T) -> Self { Self(svc) } + /// # fn send_gzip(self) -> Self { self } + /// # } + /// # #[tonic::async_trait] + /// # trait Example {} + /// + /// #[tonic::async_trait] + /// impl Example for Svc { + /// // ... + /// } + /// + /// let service = ExampleServer::new(Svc).send_gzip(); + /// ``` + #[cfg(feature = "compression")] + #[cfg_attr(docsrs, doc(cfg(feature = "compression")))] + pub fn send_gzip(mut self) -> Self { + self.send_compression_encodings.enable_gzip(); + self + } + + #[doc(hidden)] + #[cfg(not(feature = "compression"))] + pub fn send_gzip(self) -> Self { + panic!( + "`send_gzip` called on a server but the `compression` feature is not enabled on tonic" + ); + } + + #[cfg(feature = "compression")] + #[doc(hidden)] + pub fn apply_compression_config( + self, + accept_encodings: EnabledCompressionEncodings, + send_encodings: EnabledCompressionEncodings, + ) -> Self { + let mut this = self; + + let EnabledCompressionEncodings { gzip: accept_gzip } = accept_encodings; + if accept_gzip { + this = this.accept_gzip(); + } + + let EnabledCompressionEncodings { gzip: send_gzip } = send_encodings; + if send_gzip { + this = this.send_gzip(); + } + + this + } + + #[cfg(not(feature = "compression"))] + #[doc(hidden)] + #[allow(unused_variables)] + pub fn apply_compression_config(self, accept_encodings: (), send_encodings: ()) -> Self { + self } /// Handle a single unary gRPC request. @@ -43,13 +176,23 @@ where B: Body + Send + Sync + 'static, B::Error: Into + Send, { + #[cfg(feature = "compression")] + let accept_encoding = CompressionEncoding::from_accept_encoding_header( + req.headers(), + self.send_compression_encodings, + ); + let request = match self.map_request_unary(req).await { Ok(r) => r, Err(status) => { return self - .map_response::>>>(Err( - status, - )); + .map_response::>>>( + Err(status), + #[cfg(feature = "compression")] + accept_encoding, + #[cfg(feature = "compression")] + SingleMessageCompressionOverride::default(), + ); } }; @@ -58,7 +201,16 @@ where .await .map(|r| r.map(|m| stream::once(future::ok(m)))); - self.map_response(response) + #[cfg(feature = "compression")] + let compression_override = compression_override_from_response(&response); + + self.map_response( + response, + #[cfg(feature = "compression")] + accept_encoding, + #[cfg(feature = "compression")] + compression_override, + ) } /// Handle a server side streaming request. @@ -73,16 +225,36 @@ where B: Body + Send + Sync + 'static, B::Error: Into + Send, { + #[cfg(feature = "compression")] + let accept_encoding = CompressionEncoding::from_accept_encoding_header( + req.headers(), + self.send_compression_encodings, + ); + let request = match self.map_request_unary(req).await { Ok(r) => r, Err(status) => { - return self.map_response::(Err(status)); + return self.map_response::( + Err(status), + #[cfg(feature = "compression")] + accept_encoding, + #[cfg(feature = "compression")] + SingleMessageCompressionOverride::default(), + ); } }; let response = service.call(request).await; - self.map_response(response) + self.map_response( + response, + #[cfg(feature = "compression")] + accept_encoding, + // disabling compression of individual stream items must be done on + // the items themselves + #[cfg(feature = "compression")] + SingleMessageCompressionOverride::default(), + ) } /// Handle a client side streaming gRPC request. @@ -96,12 +268,29 @@ where B: Body + Send + Sync + 'static, B::Error: Into + Send + 'static, { - let request = self.map_request_streaming(req); + #[cfg(feature = "compression")] + let accept_encoding = CompressionEncoding::from_accept_encoding_header( + req.headers(), + self.send_compression_encodings, + ); + + let request = t!(self.map_request_streaming(req)); + let response = service .call(request) .await .map(|r| r.map(|m| stream::once(future::ok(m)))); - self.map_response(response) + + #[cfg(feature = "compression")] + let compression_override = compression_override_from_response(&response); + + self.map_response( + response, + #[cfg(feature = "compression")] + accept_encoding, + #[cfg(feature = "compression")] + compression_override, + ) } /// Handle a bi-directional streaming gRPC request. @@ -116,9 +305,23 @@ where B: Body + Send + Sync + 'static, B::Error: Into + Send, { - let request = self.map_request_streaming(req); + #[cfg(feature = "compression")] + let accept_encoding = CompressionEncoding::from_accept_encoding_header( + req.headers(), + self.send_compression_encodings, + ); + + let request = t!(self.map_request_streaming(req)); + let response = service.call(request).await; - self.map_response(response) + + self.map_response( + response, + #[cfg(feature = "compression")] + accept_encoding, + #[cfg(feature = "compression")] + SingleMessageCompressionOverride::default(), + ) } async fn map_request_unary( @@ -129,7 +332,16 @@ where B: Body + Send + Sync + 'static, B::Error: Into + Send, { + #[cfg(feature = "compression")] + let request_compression_encoding = self.request_encoding_if_supported(&request)?; + let (parts, body) = request.into_parts(); + + #[cfg(feature = "compression")] + let stream = + Streaming::new_request(self.codec.decoder(), body, request_compression_encoding); + + #[cfg(not(feature = "compression"))] let stream = Streaming::new_request(self.codec.decoder(), body); futures_util::pin_mut!(stream); @@ -151,42 +363,112 @@ where fn map_request_streaming( &mut self, request: http::Request, - ) -> Request> + ) -> Result>, Status> where B: Body + Send + Sync + 'static, B::Error: Into + Send, { - Request::from_http(request.map(|body| Streaming::new_request(self.codec.decoder(), body))) + #[cfg(feature = "compression")] + let encoding = self.request_encoding_if_supported(&request)?; + + #[cfg(feature = "compression")] + let request = + request.map(|body| Streaming::new_request(self.codec.decoder(), body, encoding)); + + #[cfg(not(feature = "compression"))] + let request = request.map(|body| Streaming::new_request(self.codec.decoder(), body)); + + Ok(Request::from_http(request)) } fn map_response( &mut self, response: Result, Status>, + #[cfg(feature = "compression")] accept_encoding: Option, + #[cfg(feature = "compression")] compression_override: SingleMessageCompressionOverride, ) -> http::Response where B: TryStream + Send + Sync + 'static, { - match response { - Ok(r) => { - let (mut parts, body) = r.into_http().into_parts(); + let response = match response { + Ok(r) => r, + Err(status) => return status.to_http(), + }; - // Set the content type - parts.headers.insert( - http::header::CONTENT_TYPE, - http::header::HeaderValue::from_static("application/grpc"), - ); + let (mut parts, body) = response.into_http().into_parts(); - let body = encode_server(self.codec.encoder(), body.into_stream()); + // Set the content type + parts.headers.insert( + http::header::CONTENT_TYPE, + http::header::HeaderValue::from_static("application/grpc"), + ); - http::Response::from_parts(parts, BoxBody::new(body)) - } - Err(status) => status.to_http(), + #[cfg(feature = "compression")] + if let Some(encoding) = accept_encoding { + // Set the content encoding + parts.headers.insert( + crate::codec::compression::ENCODING_HEADER, + encoding.into_header_value(), + ); } + + let body = encode_server( + self.codec.encoder(), + body.into_stream(), + #[cfg(feature = "compression")] + accept_encoding, + #[cfg(feature = "compression")] + compression_override, + ); + + http::Response::from_parts(parts, BoxBody::new(body)) + } + + #[cfg(feature = "compression")] + fn request_encoding_if_supported( + &self, + request: &http::Request, + ) -> Result, Status> { + CompressionEncoding::from_encoding_header( + request.headers(), + self.accept_compression_encodings, + ) } } impl fmt::Debug for Grpc { fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { - f.debug_struct("Grpc").field("codec", &self.codec).finish() + let mut f = f.debug_struct("Grpc"); + + f.field("codec", &self.codec); + + #[cfg(feature = "compression")] + f.field( + "accept_compression_encodings", + &self.accept_compression_encodings, + ); + + #[cfg(feature = "compression")] + f.field( + "send_compression_encodings", + &self.send_compression_encodings, + ); + + f.finish() } } + +#[cfg(feature = "compression")] +fn compression_override_from_response( + res: &Result, E>, +) -> SingleMessageCompressionOverride { + res.as_ref() + .ok() + .and_then(|response| { + response + .extensions() + .get::() + .copied() + }) + .unwrap_or_default() +}