use anyhow::Result;
use async_trait::async_trait;
use bytes::Bytes;
use http::HeaderValue;
use http::StatusCode;
use http::header::ACCEPT_ENCODING;
use http::header::CONTENT_ENCODING;
use http::header::CONTENT_LENGTH;
use http::header::CONTENT_TYPE;
use http::header::VARY;
use http_body_util::BodyExt;
use tako_rs_core::body::TakoBody;
use tako_rs_core::middleware::Next;
use tako_rs_core::plugins::TakoPlugin;
use tako_rs_core::responder::Responder;
use tako_rs_core::router::Router;
use tako_rs_core::types::Request;
use tako_rs_core::types::Response;
use super::brotli_stream::stream_brotli;
use super::config::Config;
use super::deflate_stream::stream_deflate;
use super::encoder::compress_brotli;
use super::encoder::compress_deflate;
use super::encoder::compress_gzip;
#[cfg(feature = "zstd")]
use super::encoder::compress_zstd;
use super::encoding::Encoding;
use super::gzip_stream::stream_gzip;
use super::negotiate::choose_encoding;
#[cfg(feature = "zstd")]
use super::zstd_stream::stream_zstd;
pub enum CompressionResponse<R>
where
R: Responder,
{
Plain(R),
Stream(R),
}
impl<R> Responder for CompressionResponse<R>
where
R: Responder,
{
fn into_response(self) -> Response {
match self {
CompressionResponse::Plain(r) => r.into_response(),
CompressionResponse::Stream(r) => r.into_response(),
}
}
}
#[derive(Clone)]
#[doc(alias = "compression")]
#[doc(alias = "gzip")]
#[doc(alias = "brotli")]
#[doc(alias = "deflate")]
pub struct CompressionPlugin {
pub(crate) cfg: Config,
}
impl Default for CompressionPlugin {
fn default() -> Self {
Self {
cfg: Config::default(),
}
}
}
#[async_trait]
impl TakoPlugin for CompressionPlugin {
fn name(&self) -> &'static str {
"CompressionPlugin"
}
fn setup(&self, router: &Router) -> Result<()> {
let cfg = self.cfg.clone();
router.middleware(move |req, next| {
let cfg = cfg.clone();
let stream = cfg.stream;
async move {
if stream {
CompressionResponse::Stream(
compress_stream_middleware(req, next, cfg)
.await
.into_response(),
)
} else {
CompressionResponse::Plain(compress_middleware(req, next, cfg).await.into_response())
}
}
});
Ok(())
}
}
async fn compress_middleware(req: Request, next: Next, cfg: Config) -> impl Responder {
let accepted = req
.headers()
.get(ACCEPT_ENCODING)
.and_then(|v| v.to_str().ok())
.unwrap_or("")
.to_ascii_lowercase();
let request_is_authenticated = cfg.protect_sensitive && request_carries_credentials(&req);
let mut resp = next.run(req).await;
let chosen = choose_encoding(&accepted, &cfg.enabled);
let status = resp.status();
if !(status.is_success() || status == StatusCode::NOT_MODIFIED) {
return resp.into_response();
}
if resp.headers().contains_key(CONTENT_ENCODING) {
return resp.into_response();
}
if cfg.protect_sensitive
&& (request_is_authenticated || resp.headers().contains_key(http::header::SET_COOKIE))
{
return resp.into_response();
}
if let Some(ct) = resp.headers().get(CONTENT_TYPE) {
let ct = ct.to_str().unwrap_or("");
if !cfg.content_types.matches(ct) {
return resp.into_response();
}
}
ensure_vary_accept_encoding(resp.headers_mut());
let body_bytes = if let Ok(c) = resp.body_mut().collect().await {
c.to_bytes()
} else {
tracing::warn!(
"compression middleware: response body collect() failed; \
returning original status with empty body (no compression)"
);
resp.headers_mut().remove(http::header::CONTENT_ENCODING);
*resp.body_mut() = TakoBody::empty();
return resp.into_response();
};
if body_bytes.len() < cfg.min_size {
*resp.body_mut() = TakoBody::from(body_bytes);
return resp.into_response();
}
if let Some(enc) = chosen {
let compressed = match enc {
Encoding::Gzip => compress_gzip(&body_bytes, cfg.gzip_level).ok(),
Encoding::Brotli => compress_brotli(&body_bytes, cfg.brotli_level).ok(),
Encoding::Deflate => compress_deflate(&body_bytes, cfg.deflate_level).ok(),
#[cfg(feature = "zstd")]
Encoding::Zstd => compress_zstd(&body_bytes, cfg.zstd_level).ok(),
};
if let Some(buf) = compressed {
*resp.body_mut() = TakoBody::from(Bytes::from(buf));
resp
.headers_mut()
.insert(CONTENT_ENCODING, HeaderValue::from_static(enc.as_str()));
resp.headers_mut().remove(CONTENT_LENGTH);
} else {
tracing::warn!(
encoding = enc.as_str(),
"compression failed; serving identity"
);
*resp.body_mut() = TakoBody::from(body_bytes);
resp.headers_mut().remove(CONTENT_ENCODING);
}
} else {
*resp.body_mut() = TakoBody::from(body_bytes);
}
resp.into_response()
}
pub(crate) async fn compress_stream_middleware(
req: Request,
next: Next,
cfg: Config,
) -> impl Responder {
let accepted = req
.headers()
.get(ACCEPT_ENCODING)
.and_then(|v| v.to_str().ok())
.unwrap_or("")
.to_ascii_lowercase();
let request_is_authenticated = cfg.protect_sensitive && request_carries_credentials(&req);
let mut resp = next.run(req).await;
let chosen = choose_encoding(&accepted, &cfg.enabled);
let status = resp.status();
if !(status.is_success() || status == StatusCode::NOT_MODIFIED) {
return resp.into_response();
}
if resp.headers().contains_key(CONTENT_ENCODING) {
return resp.into_response();
}
if cfg.protect_sensitive
&& (request_is_authenticated || resp.headers().contains_key(http::header::SET_COOKIE))
{
return resp.into_response();
}
if let Some(ct) = resp.headers().get(CONTENT_TYPE) {
let ct = ct.to_str().unwrap_or("");
if !cfg.content_types.matches(ct) {
return resp.into_response();
}
}
ensure_vary_accept_encoding(resp.headers_mut());
if let Some(len) = resp
.headers()
.get(CONTENT_LENGTH)
.and_then(|v| v.to_str().ok())
.and_then(|v| v.parse::<usize>().ok())
&& len < cfg.min_size
{
return resp.into_response();
}
if let Some(enc) = chosen {
let body = std::mem::replace(resp.body_mut(), TakoBody::empty());
let new_body = match enc {
Encoding::Gzip => stream_gzip(body, cfg.gzip_level),
Encoding::Brotli => stream_brotli(body, cfg.brotli_level),
Encoding::Deflate => stream_deflate(body, cfg.deflate_level),
#[cfg(feature = "zstd")]
Encoding::Zstd => stream_zstd(body, cfg.zstd_level),
};
*resp.body_mut() = new_body;
resp
.headers_mut()
.insert(CONTENT_ENCODING, HeaderValue::from_static(enc.as_str()));
resp.headers_mut().remove(CONTENT_LENGTH);
}
resp.into_response()
}
fn request_carries_credentials(req: &Request) -> bool {
req.headers().contains_key(http::header::AUTHORIZATION)
|| req
.headers()
.contains_key(http::header::PROXY_AUTHORIZATION)
|| req.headers().contains_key(http::header::COOKIE)
}
fn ensure_vary_accept_encoding(headers: &mut http::HeaderMap) {
let already_present = headers.get_all(VARY).iter().any(|v| {
v.to_str().is_ok_and(|s| {
s.split(',')
.any(|tok| tok.trim().eq_ignore_ascii_case("Accept-Encoding"))
})
});
if !already_present {
headers.append(VARY, HeaderValue::from_static("Accept-Encoding"));
}
}