web_faith_encoding/
request.rs1use std::{io, pin::Pin};
10
11use async_compression::tokio::bufread::{BrotliEncoder, GzipEncoder, ZlibEncoder, ZstdEncoder};
12use bytes::Bytes;
13use futures::Stream;
14use tokio::io::AsyncReadExt;
15use tokio_util::io::{ReaderStream, StreamReader};
16
17use http::header::{CONTENT_ENCODING, CONTENT_LENGTH, HeaderMap};
18
19use crate::{Coding, ContentEncoding};
20
21pub type RequestStream = Pin<Box<dyn Stream<Item = io::Result<Bytes>> + Send>>;
23
24pub async fn encode(headers: &mut HeaderMap, body: &[u8], coding: Coding) -> io::Result<Vec<u8>> {
33 let compressed = compress_buffer(body, coding.clone()).await?;
34 declare(headers, coding)?;
35 Ok(compressed)
36}
37
38pub fn encode_stream<S>(
42 headers: &mut HeaderMap,
43 body: S,
44 coding: Coding,
45) -> io::Result<RequestStream>
46where
47 S: Stream<Item = io::Result<Bytes>> + Send + 'static,
48{
49 let compressed = compress_stream(body, coding.clone())?;
50 declare(headers, coding)?;
51 Ok(compressed)
52}
53
54fn declare(headers: &mut HeaderMap, coding: Coding) -> io::Result<()> {
56 let layered = ContentEncoding::from(&*headers).layer(coding);
57 let value = layered
58 .to_header_value()
59 .ok_or_else(|| io::Error::other(format!("cannot declare {layered:?}")))?;
60
61 headers.insert(CONTENT_ENCODING, value);
62 headers.remove(CONTENT_LENGTH);
63 Ok(())
64}
65
66pub async fn compress_buffer(input: &[u8], coding: Coding) -> io::Result<Vec<u8>> {
71 let mut output = Vec::new();
72 match coding {
73 Coding::Gzip => GzipEncoder::new(input).read_to_end(&mut output).await?,
74 Coding::Deflate => ZlibEncoder::new(input).read_to_end(&mut output).await?,
75 Coding::Brotli => BrotliEncoder::new(input).read_to_end(&mut output).await?,
76 Coding::Zstd => ZstdEncoder::new(input).read_to_end(&mut output).await?,
77 other => return Err(unsupported(&other)),
78 };
79 Ok(output)
80}
81
82fn unsupported(coding: &Coding) -> io::Error {
84 io::Error::new(
85 io::ErrorKind::Unsupported,
86 format!("cannot compress in {:?}", coding.token()),
87 )
88}
89
90pub fn compress_stream<S>(input: S, coding: Coding) -> io::Result<RequestStream>
97where
98 S: Stream<Item = io::Result<Bytes>> + Send + 'static,
99{
100 let reader = StreamReader::new(input);
101 Ok(match coding {
102 Coding::Gzip => encoder_stream(GzipEncoder::new(reader)),
103 Coding::Deflate => encoder_stream(ZlibEncoder::new(reader)),
104 Coding::Brotli => encoder_stream(BrotliEncoder::new(reader)),
105 Coding::Zstd => encoder_stream(ZstdEncoder::new(reader)),
106 other => return Err(unsupported(&other)),
107 })
108}
109
110fn encoder_stream<R>(reader: R) -> RequestStream
111where
112 R: tokio::io::AsyncRead + Send + 'static,
113{
114 Box::pin(ReaderStream::new(reader))
115}