use std::{convert::Infallible, future::ready};
use http::{HeaderMap, header::ACCEPT_ENCODING};
use tower::{ServiceExt, service_fn};
use tower_http::compression::predicate::{NotForContentType, Predicate, SizeAbove};
use crate::{Body, response::Response};
#[derive(Clone, Debug)]
pub struct Compression {
gzip: bool,
brotli: bool,
level: CompressionLevel,
min_size: u64,
}
impl Compression {
#[must_use]
pub fn new() -> Self {
const DEFAULT_MIN_SIZE: u64 = 32;
Self {
gzip: true,
brotli: true,
level: CompressionLevel::default(),
min_size: DEFAULT_MIN_SIZE,
}
}
#[must_use]
pub fn off() -> Self {
Self {
gzip: false,
brotli: false,
..Self::new()
}
}
#[must_use]
pub fn gzip(mut self, enabled: bool) -> Self {
self.gzip = enabled;
self
}
#[must_use]
pub fn brotli(mut self, enabled: bool) -> Self {
self.brotli = enabled;
self
}
#[must_use]
pub fn level(mut self, level: CompressionLevel) -> Self {
self.level = level;
self
}
#[must_use]
pub fn min_size(mut self, bytes: u64) -> Self {
self.min_size = bytes;
self
}
pub(crate) async fn compress(
&self,
request_headers: &HeaderMap,
response: Response,
) -> Response {
if !self.gzip && !self.brotli {
return response;
}
let mut request = http::Request::new(());
for value in request_headers.get_all(ACCEPT_ENCODING) {
request.headers_mut().append(ACCEPT_ENCODING, value.clone());
}
let predicate = SizeAbove::new(self.min_size)
.and(NotForContentType::GRPC)
.and(NotForContentType::IMAGES)
.and(NotForContentType::SSE);
let mut response = Some(response);
let inner = service_fn(move |_: http::Request<()>| {
let response = response.take().expect("one-shot service called once");
ready(Ok::<_, Infallible>(response))
});
let service = tower_http::compression::Compression::new(inner)
.gzip(self.gzip)
.br(self.brotli)
.quality(self.level.into_tower())
.compress_when(predicate);
match service.oneshot(request).await {
Ok(response) => response.map(Body::new),
Err(never) => match never {},
}
}
}
impl Default for Compression {
fn default() -> Self {
Self::new()
}
}
#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
pub enum CompressionLevel {
Fastest,
#[default]
Balanced,
Best,
Precise(i32),
}
impl CompressionLevel {
fn into_tower(self) -> tower_http::CompressionLevel {
const BALANCED_QUALITY: i32 = 4;
match self {
Self::Fastest => tower_http::CompressionLevel::Fastest,
Self::Balanced => tower_http::CompressionLevel::Precise(BALANCED_QUALITY),
Self::Best => tower_http::CompressionLevel::Best,
Self::Precise(quality) => tower_http::CompressionLevel::Precise(quality),
}
}
}
#[cfg(test)]
mod tests {
use std::{borrow::Cow, future::Future};
use http::{
HeaderValue,
header::{CONTENT_ENCODING, VARY},
};
use topcoat_core::context::Cx;
use super::*;
use crate::{
HeaderMap, Method, Path, RouteFn, RouteFuture, RouteHandlerFn, Router, request::Bytes,
response::IntoResponse, to_bytes,
};
fn block_on<F: Future>(future: F) -> F::Output {
tokio::runtime::Builder::new_current_thread()
.build()
.unwrap()
.block_on(future)
}
fn router_with(handler: RouteHandlerFn, compression: Compression) -> Router {
Router::builder()
.route(RouteFn::new(
Method::GET,
Cow::Borrowed(Path::new("/x")),
handler,
))
.compression(compression)
.build()
}
fn send(router: &Router, accept_encoding: Option<&str>) -> (HeaderMap, Bytes) {
let mut request = http::Request::builder().uri("/x");
if let Some(value) = accept_encoding {
request = request.header(ACCEPT_ENCODING, value);
}
let response = block_on(router.handle(request.body(Body::empty()).unwrap()));
let (parts, body) = response.into_parts();
let bytes = block_on(to_bytes(body, usize::MAX)).unwrap();
(parts.headers, bytes)
}
fn long_body() -> String {
"route ".repeat(64)
}
fn long_route(cx: &Cx, _body: Body) -> RouteFuture<'_> {
Box::pin(async move { long_body().into_response(cx) })
}
fn short_route(cx: &Cx, _body: Body) -> RouteFuture<'_> {
Box::pin(async move { "route".into_response(cx) })
}
fn encoded_route(cx: &Cx, _body: Body) -> RouteFuture<'_> {
Box::pin(async move {
let mut response = long_body().into_response(cx)?;
response
.headers_mut()
.insert(CONTENT_ENCODING, HeaderValue::from_static("br"));
Ok(response)
})
}
#[test]
fn compresses_with_the_negotiated_algorithm() {
let router = router_with(long_route, Compression::new());
let (headers, body) = send(&router, Some("gzip"));
assert_eq!(headers.get(CONTENT_ENCODING).unwrap(), "gzip");
assert_eq!(headers.get(VARY).unwrap(), "accept-encoding");
assert!(!body.is_empty());
assert!(body.len() < long_body().len());
let (headers, _) = send(&router, Some("br"));
assert_eq!(headers.get(CONTENT_ENCODING).unwrap(), "br");
}
#[test]
fn passes_through_without_accept_encoding() {
let router = router_with(long_route, Compression::new());
let (headers, body) = send(&router, None);
assert!(!headers.contains_key(CONTENT_ENCODING));
assert_eq!(body, long_body());
}
#[test]
fn off_disables_compression() {
let router = router_with(long_route, Compression::off());
let (headers, body) = send(&router, Some("gzip, br"));
assert!(!headers.contains_key(CONTENT_ENCODING));
assert_eq!(body, long_body());
}
#[test]
fn disabled_algorithms_are_not_offered() {
let router = router_with(long_route, Compression::new().gzip(false));
let (headers, body) = send(&router, Some("gzip"));
assert!(!headers.contains_key(CONTENT_ENCODING));
assert_eq!(body, long_body());
let (headers, _) = send(&router, Some("gzip, br"));
assert_eq!(headers.get(CONTENT_ENCODING).unwrap(), "br");
}
#[test]
fn small_bodies_are_not_compressed() {
let router = router_with(short_route, Compression::new());
let (headers, body) = send(&router, Some("gzip"));
assert!(!headers.contains_key(CONTENT_ENCODING));
assert_eq!(&body[..], b"route");
}
#[test]
fn min_size_lowers_the_compression_threshold() {
let router = router_with(short_route, Compression::new().min_size(0));
let (headers, _) = send(&router, Some("gzip"));
assert_eq!(headers.get(CONTENT_ENCODING).unwrap(), "gzip");
}
#[test]
fn already_encoded_responses_pass_through() {
let router = router_with(encoded_route, Compression::new());
let (headers, body) = send(&router, Some("gzip"));
assert_eq!(headers.get(CONTENT_ENCODING).unwrap(), "br");
assert_eq!(body, long_body());
}
}