use std::{cmp, str::FromStr};
use crate::http::encoding::Encoder;
use crate::http::header::{ACCEPT_ENCODING, ContentEncoding};
use crate::service::{Ctx, Middleware, Service};
use crate::web::{BodyEncoding, State, WebRequest, WebResponse};
#[derive(Debug, Clone)]
pub struct Compress {
enc: ContentEncoding,
}
impl Compress {
pub fn new(encoding: ContentEncoding) -> Self {
Compress { enc: encoding }
}
}
impl Default for Compress {
fn default() -> Self {
Compress::new(ContentEncoding::Auto)
}
}
impl<S, St> Middleware<S, St> for Compress {
type Service = CompressMiddleware<S>;
fn create(&self, _: &St, service: S) -> Self::Service {
CompressMiddleware {
service,
encoding: self.enc,
}
}
}
#[derive(Debug)]
pub struct CompressMiddleware<S> {
service: S,
encoding: ContentEncoding,
}
impl<S, St, In> Service<St, WebRequest<In>> for CompressMiddleware<S>
where
S: Service<St, WebRequest<In>, Res = WebResponse>,
St: State,
{
type Res = WebResponse;
type Error = S::Error;
crate::forward_ready!(St, service);
crate::forward_shutdown!(St, service);
async fn call(
&self,
req: WebRequest<In>,
ctx: Ctx<'_, Self, St>,
) -> Result<WebResponse, S::Error> {
let encoding = if let Some(val) = req.headers().get(&ACCEPT_ENCODING) {
if let Ok(enc) = val.to_str() {
AcceptEncoding::parse(enc, self.encoding)
} else {
ContentEncoding::Identity
}
} else {
ContentEncoding::Identity
};
let resp = ctx.call(&self.service, req).await?;
let enc = if let Some(enc) = resp.response().get_encoding() {
enc
} else {
encoding
};
Ok(resp.map_body(move |head, body| Encoder::response(enc, head, body)))
}
}
struct AcceptEncoding {
encoding: ContentEncoding,
quality: f64,
}
impl Eq for AcceptEncoding {}
impl Ord for AcceptEncoding {
#[allow(clippy::comparison_chain)]
fn cmp(&self, other: &AcceptEncoding) -> cmp::Ordering {
if self.quality > other.quality {
cmp::Ordering::Less
} else if self.quality < other.quality {
cmp::Ordering::Greater
} else {
cmp::Ordering::Equal
}
}
}
impl PartialOrd for AcceptEncoding {
fn partial_cmp(&self, other: &AcceptEncoding) -> Option<cmp::Ordering> {
Some(self.cmp(other))
}
}
impl PartialEq for AcceptEncoding {
fn eq(&self, other: &AcceptEncoding) -> bool {
self.quality == other.quality
}
}
impl AcceptEncoding {
fn new(tag: &str) -> Option<AcceptEncoding> {
let mut parts = tag.split(';').map(str::trim);
let encoding = ContentEncoding::from(parts.next()?);
let quality = match parts.next() {
None => encoding.quality(),
Some(q) => f64::from_str(q).unwrap_or(0.0),
};
Some(AcceptEncoding { encoding, quality })
}
fn parse(raw: &str, encoding: ContentEncoding) -> ContentEncoding {
let mut encodings: Vec<_> = raw.split(',').map(AcceptEncoding::new).collect();
encodings.sort();
for enc in encodings.into_iter().flatten() {
if encoding == ContentEncoding::Auto {
if Encoder::can_encode(enc.encoding) {
return enc.encoding;
}
} else if encoding == enc.encoding {
return encoding;
}
}
ContentEncoding::Identity
}
}
#[cfg(test)]
#[allow(clippy::float_cmp)]
mod tests {
use std::cmp::Ordering;
use super::*;
#[test]
fn test_accepting_encodings_equal() {
let accepting_encoding = AcceptEncoding {
encoding: ContentEncoding::Auto,
quality: 0.0,
};
let accepting_encoding2 = AcceptEncoding {
encoding: ContentEncoding::Br,
quality: 0.0,
};
assert!(accepting_encoding == accepting_encoding2);
}
#[test]
fn test_accepting_encodings_not_equal() {
let accepting_encoding = AcceptEncoding {
encoding: ContentEncoding::Auto,
quality: 1.0,
};
let accepting_encoding2 = AcceptEncoding {
encoding: ContentEncoding::Br,
quality: 0.0,
};
assert!(accepting_encoding != accepting_encoding2);
}
#[test]
fn test_accepting_encodings_cmp_order_less() {
let accepting_encoding = AcceptEncoding {
encoding: ContentEncoding::Auto,
quality: 1.0,
};
let accepting_encoding2 = AcceptEncoding {
encoding: ContentEncoding::Br,
quality: 0.0,
};
assert_eq!(accepting_encoding.cmp(&accepting_encoding2), Ordering::Less);
}
#[test]
fn test_accepting_encodings_cmp_order_equal() {
let accepting_encoding = AcceptEncoding {
encoding: ContentEncoding::Auto,
quality: 1.0,
};
assert_eq!(accepting_encoding.cmp(&accepting_encoding), Ordering::Equal);
}
#[test]
fn test_accepting_encoding_from_tag_with_valid_quality() {
let accepting_encoding = AcceptEncoding::new("gzip;0.8").unwrap();
assert_eq!(accepting_encoding.quality, 0.8);
}
#[test]
fn test_auto_skips_unsupported_encodings() {
let auto = ContentEncoding::Auto;
assert_eq!(
AcceptEncoding::parse("gzip, deflate, br, zstd", auto),
ContentEncoding::Gzip
);
assert_eq!(
AcceptEncoding::parse("br, deflate", auto),
ContentEncoding::Deflate
);
assert_eq!(AcceptEncoding::parse("br", auto), ContentEncoding::Identity);
assert_eq!(
AcceptEncoding::parse("gzip, br", ContentEncoding::Br),
ContentEncoding::Br
);
assert_eq!(
AcceptEncoding::parse(" br , gzip ; q=1.0 ", auto),
ContentEncoding::Gzip
);
}
#[test]
fn test_accepting_encoding_from_tag_with_invalid_quality() {
let accepting_encoding = AcceptEncoding::new("gzip;q=abc").unwrap();
assert_eq!(accepting_encoding.quality, 0.0);
}
#[crate::rt_test]
async fn test_compress_accept_encoding() {
use crate::http::header::{CONTENT_ENCODING, HeaderValue};
use crate::web::test::{TestRequest, call_service, init_service};
use crate::web::{self, App, HttpResponse};
let srv = init_service(App::new().middleware(Compress::default()).route(
"/",
web::get().to(async || HttpResponse::Ok().body("a".repeat(1024))),
))
.await;
let req = TestRequest::default()
.header(ACCEPT_ENCODING, "gzip")
.to_request();
let resp = call_service(&srv, req).await;
assert_eq!(resp.headers().get(CONTENT_ENCODING).unwrap(), "gzip");
let req = TestRequest::default()
.header(
ACCEPT_ENCODING,
HeaderValue::from_bytes(b"gzip\xff").unwrap(),
)
.to_request();
let resp = call_service(&srv, req).await;
assert!(resp.headers().get(CONTENT_ENCODING).is_none());
}
}