use std::{io, pin::Pin};
use async_compression::tokio::bufread::{BrotliEncoder, GzipEncoder, ZlibEncoder, ZstdEncoder};
use bytes::Bytes;
use futures::Stream;
use tokio::io::AsyncReadExt;
use tokio_util::io::{ReaderStream, StreamReader};
use http::header::{CONTENT_ENCODING, CONTENT_LENGTH, HeaderMap};
use crate::{Coding, ContentEncoding};
pub type RequestStream = Pin<Box<dyn Stream<Item = io::Result<Bytes>> + Send>>;
pub async fn encode(headers: &mut HeaderMap, body: &[u8], coding: Coding) -> io::Result<Vec<u8>> {
let compressed = compress_buffer(body, coding.clone()).await?;
declare(headers, coding)?;
Ok(compressed)
}
pub fn encode_stream<S>(
headers: &mut HeaderMap,
body: S,
coding: Coding,
) -> io::Result<RequestStream>
where
S: Stream<Item = io::Result<Bytes>> + Send + 'static,
{
let compressed = compress_stream(body, coding.clone())?;
declare(headers, coding)?;
Ok(compressed)
}
fn declare(headers: &mut HeaderMap, coding: Coding) -> io::Result<()> {
let layered = ContentEncoding::from(&*headers).layer(coding);
let value = layered
.to_header_value()
.ok_or_else(|| io::Error::other(format!("cannot declare {layered:?}")))?;
headers.insert(CONTENT_ENCODING, value);
headers.remove(CONTENT_LENGTH);
Ok(())
}
pub async fn compress_buffer(input: &[u8], coding: Coding) -> io::Result<Vec<u8>> {
let mut output = Vec::new();
match coding {
Coding::Gzip => GzipEncoder::new(input).read_to_end(&mut output).await?,
Coding::Deflate => ZlibEncoder::new(input).read_to_end(&mut output).await?,
Coding::Brotli => BrotliEncoder::new(input).read_to_end(&mut output).await?,
Coding::Zstd => ZstdEncoder::new(input).read_to_end(&mut output).await?,
other => return Err(unsupported(&other)),
};
Ok(output)
}
fn unsupported(coding: &Coding) -> io::Error {
io::Error::new(
io::ErrorKind::Unsupported,
format!("cannot compress in {:?}", coding.token()),
)
}
pub fn compress_stream<S>(input: S, coding: Coding) -> io::Result<RequestStream>
where
S: Stream<Item = io::Result<Bytes>> + Send + 'static,
{
let reader = StreamReader::new(input);
Ok(match coding {
Coding::Gzip => encoder_stream(GzipEncoder::new(reader)),
Coding::Deflate => encoder_stream(ZlibEncoder::new(reader)),
Coding::Brotli => encoder_stream(BrotliEncoder::new(reader)),
Coding::Zstd => encoder_stream(ZstdEncoder::new(reader)),
other => return Err(unsupported(&other)),
})
}
fn encoder_stream<R>(reader: R) -> RequestStream
where
R: tokio::io::AsyncRead + Send + 'static,
{
Box::pin(ReaderStream::new(reader))
}