ntex 4.0.0-beta.15

Framework for composable network services
//! `Middleware` for compressing response body.
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)]
/// `Middleware` for compressing response body.
///
/// Use `BodyEncoding` trait for overriding response compression.
/// To disable compression set encoding to `ContentEncoding::Identity` value.
///
/// ```rust
/// use ntex::web::{self, middleware, App, HttpResponse};
///
/// fn main() {
///     let app = App::default()
///         .middleware(middleware::Compress::default())
///         .service(
///             web::resource("/test")
///                 .route(web::get().to(async || { HttpResponse::Ok() }))
///                 .route(web::head().to(async || { HttpResponse::MethodNotAllowed() }))
///         );
/// }
/// ```
pub struct Compress {
    enc: ContentEncoding,
}

impl Compress {
    /// Create new `Compress` middleware with the specified encoding.
    ///
    /// Use `Compress::default()` to select the encoding automatically
    /// ([`ContentEncoding::Auto`]).
    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> {
        // negotiate content-encoding
        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 })
    }

    /// Parse a raw Accept-Encoding header value into an ordered list.
    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());
    }
}