feat(codec): compression support (#692)
* Initial compression support * Support configuring compression on `Server` * Minor clean up * Test that compression is actually happening * Clean up some todos * channels compressing requests * Move compression to be on the codecs * Test sending compressed request to server that doesn't support it * Clean up a bit * Compress server streams * Compress client streams * Bidirectional streaming compression * Handle receiving unsupported encoding * Clean up * Add note to future self * Support disabling compression for individual responses * Add docs * Add compression examples * Disable compression behind feature flag * Add some docs * Make flate2 optional dependency * Fix docs wording * Format * Reply with which encodings are supported * Convert tests to use mocked io * Fix lints * Use separate counters * Don't make a long stream * Address review feedback
This commit is contained in:
@@ -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<B>(req: http::Request<B>) -> http::Request<B> {
|
||||
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<SomeData> = 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);
|
||||
}
|
||||
@@ -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<B>(req: http::Request<B>) -> http::Request<B> {
|
||||
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<B>(req: http::Request<B>) -> http::Request<B> {
|
||||
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);
|
||||
}
|
||||
@@ -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<B>(req: http::Request<B>) -> http::Request<B> {
|
||||
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"
|
||||
);
|
||||
}
|
||||
@@ -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>(S);
|
||||
|
||||
impl<S, B> Service<http::Request<B>> for AssertCorrectAcceptEncoding<S>
|
||||
where
|
||||
S: Service<http::Request<B>>,
|
||||
{
|
||||
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<Result<(), Self::Error>> {
|
||||
self.0.poll_ready(cx)
|
||||
}
|
||||
|
||||
fn call(&mut self, req: http::Request<B>) -> 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>(S);
|
||||
|
||||
impl<S, B> Service<http::Request<B>> for AssertCorrectAcceptEncoding<S>
|
||||
where
|
||||
S: Service<http::Request<B>>,
|
||||
{
|
||||
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<Result<(), Self::Error>> {
|
||||
self.0.poll_ready(cx)
|
||||
}
|
||||
|
||||
fn call(&mut self, req: http::Request<B>) -> 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<B>(mut response: http::Response<B>) -> http::Response<B> {
|
||||
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<SomeData> = 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);
|
||||
}
|
||||
@@ -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<B>(&self, mut res: Response<B>) -> Response<B> {
|
||||
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<Response<SomeData>, 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<SomeData>) -> Result<Response<()>, Status> {
|
||||
assert_eq!(req.into_inner().data.len(), UNCOMPRESSED_MIN_BODY_SIZE);
|
||||
Ok(Response::new(()))
|
||||
}
|
||||
|
||||
type CompressOutputServerStreamStream =
|
||||
Pin<Box<dyn Stream<Item = Result<SomeData, Status>> + Send + Sync + 'static>>;
|
||||
|
||||
async fn compress_output_server_stream(
|
||||
&self,
|
||||
_req: Request<()>,
|
||||
) -> Result<Response<Self::CompressOutputServerStreamStream>, 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<Streaming<SomeData>>,
|
||||
) -> Result<Response<()>, 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<Streaming<SomeData>>,
|
||||
) -> Result<Response<SomeData>, 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<Box<dyn Stream<Item = Result<SomeData, Status>> + Send + Sync + 'static>>;
|
||||
|
||||
async fn compress_input_output_bidirectional_stream(
|
||||
&self,
|
||||
req: Request<Streaming<SomeData>>,
|
||||
) -> Result<Response<Self::CompressInputOutputBidirectionalStreamStream>, 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))))
|
||||
}
|
||||
}
|
||||
@@ -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<SomeData> = 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<SomeData> = 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<SomeData> = res.into_inner();
|
||||
|
||||
stream
|
||||
.next()
|
||||
.await
|
||||
.expect("stream empty")
|
||||
.expect("item was error");
|
||||
assert!(response_bytes_counter.load(SeqCst) > UNCOMPRESSED_MIN_BODY_SIZE);
|
||||
}
|
||||
@@ -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<B> {
|
||||
#[pin]
|
||||
pub inner: B,
|
||||
pub counter: Arc<AtomicUsize>,
|
||||
}
|
||||
|
||||
impl<B> Body for CountBytesBody<B>
|
||||
where
|
||||
B: Body<Data = Bytes>,
|
||||
{
|
||||
type Data = B::Data;
|
||||
type Error = B::Error;
|
||||
|
||||
fn poll_data(
|
||||
self: Pin<&mut Self>,
|
||||
cx: &mut Context<'_>,
|
||||
) -> Poll<Option<Result<Self::Data, Self::Error>>> {
|
||||
let this = self.project();
|
||||
let counter: Arc<AtomicUsize> = 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<Result<Option<http::HeaderMap>, 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<AtomicUsize>,
|
||||
) -> MapRequestBodyLayer<impl Fn(hyper::Body) -> 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<std::io::Result<()>> {
|
||||
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<std::io::Result<usize>> {
|
||||
Pin::new(&mut self.0).poll_write(cx, buf)
|
||||
}
|
||||
|
||||
fn poll_flush(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<std::io::Result<()>> {
|
||||
Pin::new(&mut self.0).poll_flush(cx)
|
||||
}
|
||||
|
||||
fn poll_shutdown(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<std::io::Result<()>> {
|
||||
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()
|
||||
}
|
||||
Reference in New Issue
Block a user