1use crate::http::encoding::Encoder;
3use crate::http::header::{ACCEPT_ENCODING, ContentEncoding};
4use crate::service::{Ctx, Middleware, Service};
5use crate::web::{BodyEncoding, State, WebRequest, WebResponse};
6
7#[derive(Debug, Clone)]
8pub struct Compress {
28 enc: ContentEncoding,
29}
30
31impl Compress {
32 pub fn new(encoding: ContentEncoding) -> Self {
37 Compress { enc: encoding }
38 }
39}
40
41impl Default for Compress {
42 fn default() -> Self {
43 Compress::new(ContentEncoding::Auto)
44 }
45}
46
47impl<S, St> Middleware<S, St> for Compress {
48 type Service = CompressMiddleware<S>;
49
50 fn create(&self, _: &St, service: S) -> Self::Service {
51 CompressMiddleware {
52 service,
53 encoding: self.enc,
54 }
55 }
56}
57
58#[derive(Debug)]
59pub struct CompressMiddleware<S> {
60 service: S,
61 encoding: ContentEncoding,
62}
63
64impl<S, St, In> Service<St, WebRequest<In>> for CompressMiddleware<S>
65where
66 S: Service<St, WebRequest<In>, Res = WebResponse>,
67 St: State,
68{
69 type Res = WebResponse;
70 type Error = S::Error;
71
72 crate::forward_ready!(St, service);
73 crate::forward_shutdown!(St, service);
74
75 async fn call(
76 &self,
77 req: WebRequest<In>,
78 ctx: Ctx<'_, Self, St>,
79 ) -> Result<WebResponse, S::Error> {
80 let values = req
82 .headers()
83 .get_all(&ACCEPT_ENCODING)
84 .filter_map(|val| val.to_str().ok());
85 let encoding = AcceptEncoding::parse(values, self.encoding);
86
87 let resp = ctx.call(&self.service, req).await?;
88
89 let enc = if let Some(enc) = resp.response().get_encoding() {
90 enc
91 } else {
92 encoding
93 };
94
95 Ok(resp.map_body(move |head, body| Encoder::response(enc, head, body)))
96 }
97}
98
99const WILDCARD: [ContentEncoding; 3] = [
101 ContentEncoding::Zstd,
102 ContentEncoding::Gzip,
103 ContentEncoding::Deflate,
104];
105
106#[derive(Debug, PartialEq)]
107struct AcceptEncoding {
108 encoding: Option<ContentEncoding>,
110 q: u16,
112}
113
114impl AcceptEncoding {
115 fn new(tag: &str) -> Option<AcceptEncoding> {
116 let mut parts = tag.split(';');
117 let name = parts.next()?.trim();
118 if name.is_empty() {
119 return None;
120 }
121 let encoding = if name == "*" {
122 None
123 } else {
124 Some(ContentEncoding::from(name))
125 };
126
127 let q = parts
129 .filter_map(|param| param.split_once('='))
130 .find(|(name, _)| name.trim().eq_ignore_ascii_case("q"))
131 .map_or(1000, |(_, val)| parse_q(val.trim()).unwrap_or(0));
132 Some(AcceptEncoding { encoding, q })
133 }
134
135 fn parse<'a, I>(values: I, encoding: ContentEncoding) -> ContentEncoding
141 where
142 I: Iterator<Item = &'a str>,
143 {
144 let mut listed: Vec<(ContentEncoding, u16)> = Vec::new();
145 let mut wildcard = None;
146 for item in values.flat_map(|val| val.split(',')) {
147 match AcceptEncoding::new(item) {
148 Some(AcceptEncoding {
149 encoding: Some(enc),
150 q,
151 }) => listed.push((enc, q)),
152 Some(AcceptEncoding { encoding: None, q }) => {
153 wildcard.get_or_insert(q);
154 }
155 None => {}
156 }
157 }
158 let weight = |enc: ContentEncoding| {
159 listed
160 .iter()
161 .find(|(e, _)| *e == enc)
162 .map(|(_, q)| *q)
163 .or(wildcard)
164 .unwrap_or(0)
165 };
166
167 if encoding != ContentEncoding::Auto {
168 return if weight(encoding) > 0 {
169 encoding
170 } else {
171 ContentEncoding::Identity
172 };
173 }
174
175 let mut best: Option<(ContentEncoding, u16)> = None;
176 let candidates = listed.iter().map(|(enc, _)| *enc).chain(WILDCARD);
177 for enc in candidates.filter(|enc| Encoder::can_encode(*enc)) {
178 let q = weight(enc);
179 let better = match best {
180 None => q > 0,
181 Some((b, bq)) => q > bq || (q == bq && enc.quality() > b.quality()),
182 };
183 if better {
184 best = Some((enc, q));
185 }
186 }
187 best.map_or(ContentEncoding::Identity, |(enc, _)| enc)
188 }
189}
190
191fn parse_q(val: &str) -> Option<u16> {
193 let (int, frac) = match val.split_once('.') {
194 Some((int, frac)) => (int, frac),
195 None => (val, ""),
196 };
197 if frac.len() > 3 || !frac.bytes().all(|b| b.is_ascii_digit()) {
198 return None;
199 }
200 let int = match int {
201 "0" => 0,
202 "1" => 1000,
203 _ => return None,
204 };
205 let frac = frac
206 .bytes()
207 .zip([100, 10, 1])
208 .map(|(b, scale)| u16::from(b - b'0') * scale)
209 .sum::<u16>();
210 if int + frac > 1000 { None } else { Some(int + frac) }
211}
212
213#[cfg(test)]
214mod tests {
215 use super::*;
216
217 fn parse(raw: &str, encoding: ContentEncoding) -> ContentEncoding {
218 AcceptEncoding::parse([raw].into_iter(), encoding)
219 }
220
221 #[test]
222 fn test_parse_q() {
223 for (val, q) in [
224 ("0", 0),
225 ("0.", 0),
226 ("0.5", 500),
227 ("0.05", 50),
228 ("0.125", 125),
229 ("0.999", 999),
230 ("1", 1000),
231 ("1.", 1000),
232 ("1.000", 1000),
233 ] {
234 assert_eq!(parse_q(val), Some(q), "{val}");
235 }
236 for val in [
237 "", ".5", "0.1234", "1.001", "1.5", "2", "-0", "0.a", "0.00a", "01", "abc",
238 ] {
239 assert_eq!(parse_q(val), None, "{val}");
240 }
241 }
242
243 #[test]
244 fn test_accept_encoding_entry() {
245 let entry = |tag| AcceptEncoding::new(tag).unwrap();
246 assert_eq!(
247 entry("gzip"),
248 AcceptEncoding {
249 encoding: Some(ContentEncoding::Gzip),
250 q: 1000
251 }
252 );
253 assert_eq!(entry(" gzip ; q=0.5 ").q, 500);
254 assert_eq!(entry("gzip;Q=0.5").q, 500);
255 assert_eq!(entry("gzip;level=1;q=0.25").q, 250);
256 assert_eq!(entry("gzip;q=abc").q, 0);
257 assert_eq!(entry("gzip;0.8").q, 1000);
258 assert_eq!(entry("*;q=0.1").encoding, None);
259 assert!(AcceptEncoding::new(" ").is_none());
260 }
261
262 #[test]
263 fn test_auto_skips_unsupported_encodings() {
264 let auto = ContentEncoding::Auto;
265 assert_eq!(
266 parse("gzip, deflate, br, zstd", auto),
267 ContentEncoding::Zstd
268 );
269 assert_eq!(parse("gzip, deflate, br", auto), ContentEncoding::Gzip);
270 assert_eq!(parse("br, deflate", auto), ContentEncoding::Deflate);
271 assert_eq!(parse("br", auto), ContentEncoding::Identity);
272 assert_eq!(parse("gzip, br", ContentEncoding::Br), ContentEncoding::Br);
273 assert_eq!(parse(" br , gzip ; q=1.0 ", auto), ContentEncoding::Gzip);
274 assert_eq!(parse("", auto), ContentEncoding::Identity);
275 }
276
277 #[test]
278 fn test_auto_honours_q() {
279 let auto = ContentEncoding::Auto;
280 assert_eq!(parse("gzip;q=0.5, zstd;q=0.1", auto), ContentEncoding::Gzip);
281 assert_eq!(
282 parse("deflate;q=0.9, gzip;q=0.8", auto),
283 ContentEncoding::Deflate
284 );
285 assert_eq!(parse("zstd;q=0, gzip", auto), ContentEncoding::Gzip);
286 assert_eq!(parse("gzip;q=0", auto), ContentEncoding::Identity);
287 assert_eq!(parse("gzip;q=abc", auto), ContentEncoding::Identity);
288 assert_eq!(parse("gzip;q=2, deflate", auto), ContentEncoding::Deflate);
289 assert_eq!(parse("gzip, zstd;q=0", auto), ContentEncoding::Gzip);
290 assert_eq!(parse("gzip;q=0, gzip", auto), ContentEncoding::Identity);
292 assert_eq!(
294 parse("deflate;q=0.5, gzip;q=0.5, zstd;q=0.5", auto),
295 ContentEncoding::Zstd
296 );
297 assert_eq!(parse("deflate, gzip", auto), ContentEncoding::Gzip);
298 assert_eq!(parse("x-gzip", auto), ContentEncoding::Gzip);
299 }
300
301 #[test]
302 fn test_auto_wildcard() {
303 let auto = ContentEncoding::Auto;
304 assert_eq!(parse("*", auto), ContentEncoding::Zstd);
305 assert_eq!(parse("*;q=0", auto), ContentEncoding::Identity);
306 assert_eq!(parse("*;q=0.5, zstd;q=0", auto), ContentEncoding::Gzip);
307 assert_eq!(
308 parse("zstd;q=0, gzip;q=0, *", auto),
309 ContentEncoding::Deflate
310 );
311 assert_eq!(parse("*;q=0.5, deflate", auto), ContentEncoding::Deflate);
312 assert_eq!(parse("gzip;q=0.1, *;q=0.5", auto), ContentEncoding::Zstd);
313 assert_eq!(parse("*;q=0, gzip;q=0.1", auto), ContentEncoding::Gzip);
314 assert_eq!(parse("*;q=0.5, *", auto), ContentEncoding::Zstd);
315 assert_eq!(parse("*;q=0, *, zstd;q=0", auto), ContentEncoding::Identity);
316 }
317
318 #[test]
319 fn test_explicit_encoding() {
320 let gzip = ContentEncoding::Gzip;
321 assert_eq!(parse("gzip", gzip), gzip);
322 assert_eq!(parse("zstd", gzip), ContentEncoding::Identity);
323 assert_eq!(parse("gzip;q=0", gzip), ContentEncoding::Identity);
324 assert_eq!(parse("gzip;q=0.001", gzip), gzip);
325 assert_eq!(parse("*", gzip), gzip);
326 assert_eq!(parse("*;q=0", gzip), ContentEncoding::Identity);
327 assert_eq!(parse("*, gzip;q=0", gzip), ContentEncoding::Identity);
328 assert_eq!(parse("*;q=0, gzip", gzip), gzip);
329 }
330
331 #[test]
332 fn test_repeated_headers() {
333 let values = ["zstd;q=0", "gzip;q=0.5, deflate;q=0.1"];
334 assert_eq!(
335 AcceptEncoding::parse(values.into_iter(), ContentEncoding::Auto),
336 ContentEncoding::Gzip
337 );
338 }
339
340 #[crate::rt_test]
341 async fn test_compress_accept_encoding() {
342 use crate::http::header::{CONTENT_ENCODING, HeaderValue};
343 use crate::web::test::{TestRequest, call_service, init_service};
344 use crate::web::{self, App, HttpResponse};
345
346 let srv = init_service(App::new().middleware(Compress::default()).route(
347 "/",
348 web::get().to(async || HttpResponse::Ok().body("a".repeat(1024))),
349 ))
350 .await;
351
352 let req = TestRequest::default()
353 .header(ACCEPT_ENCODING, "gzip")
354 .to_request();
355 let resp = call_service(&srv, req).await;
356 assert_eq!(resp.headers().get(CONTENT_ENCODING).unwrap(), "gzip");
357
358 let req = TestRequest::default()
359 .header(
360 ACCEPT_ENCODING,
361 HeaderValue::from_bytes(b"gzip\xff").unwrap(),
362 )
363 .to_request();
364 let resp = call_service(&srv, req).await;
365 assert!(resp.headers().get(CONTENT_ENCODING).is_none());
366
367 let req = TestRequest::default()
368 .header(ACCEPT_ENCODING, "zstd;q=0")
369 .header(ACCEPT_ENCODING, "gzip;q=0.5")
370 .to_request();
371 let resp = call_service(&srv, req).await;
372 assert_eq!(resp.headers().get(CONTENT_ENCODING).unwrap(), "gzip");
373
374 let req = TestRequest::default()
375 .header(
376 ACCEPT_ENCODING,
377 HeaderValue::from_bytes(b"zstd\xff").unwrap(),
378 )
379 .header(ACCEPT_ENCODING, "deflate")
380 .to_request();
381 let resp = call_service(&srv, req).await;
382 assert_eq!(resp.headers().get(CONTENT_ENCODING).unwrap(), "deflate");
383
384 let req = TestRequest::default()
385 .header(ACCEPT_ENCODING, "gzip;q=0")
386 .to_request();
387 let resp = call_service(&srv, req).await;
388 assert!(resp.headers().get(CONTENT_ENCODING).is_none());
389 }
390}