use std::io::Result as IoResult;
use std::pin::Pin;
use std::task::{Context, Poll};
use tokio::io::AsyncWrite;
macro_rules! ok_ready {
($poll: expr) => {
match $poll {
Poll::Pending => return Poll::Pending,
Poll::Ready(Err(e)) => return Poll::Ready(Err(e)),
Poll::Ready(Ok(v)) => v,
}
};
}
pub struct Encoder<W>
where
W: AsyncWrite + Unpin,
{
output: W,
chunks_size: usize,
buffer: Vec<u8>,
flush_after_write: bool,
offset: Option<usize>,
}
const MAX_CHUNK_SIZE: usize = std::u32::MAX as usize;
const MAX_HEADER_SIZE: usize = 6;
impl<W> Encoder<W>
where
W: AsyncWrite + Unpin,
{
pub fn new(output: W) -> Encoder<W> {
Encoder::with_chunks_size(output, 8192)
}
pub fn with_chunks_size(output: W, chunks: usize) -> Encoder<W> {
let chunks_size = chunks.min(MAX_CHUNK_SIZE);
let mut encoder = Encoder {
output,
chunks_size,
buffer: vec![0; MAX_HEADER_SIZE],
flush_after_write: false,
offset: None,
};
encoder.reset_buffer();
encoder
}
pub fn with_flush_after_write(output: W) -> Encoder<W> {
let mut encoder = Encoder {
output,
chunks_size: 8192,
buffer: vec![0; MAX_HEADER_SIZE],
flush_after_write: true,
offset: None,
};
encoder.reset_buffer();
encoder
}
fn reset_buffer(&mut self) {
self.buffer.truncate(MAX_HEADER_SIZE);
self.offset = None;
}
fn is_buffer_empty(&self) -> bool {
self.buffer.len() == MAX_HEADER_SIZE
}
fn buffer_len(&self) -> usize {
self.buffer.len() - MAX_HEADER_SIZE
}
fn send(&mut self, cx: &mut Context) -> Poll<IoResult<()>> {
if let Some(mut offset) = self.offset {
loop {
let wrote =
ok_ready!(Pin::new(&mut self.output).poll_write(cx, &self.buffer[offset..]));
offset += wrote;
self.offset = Some(offset);
if offset >= self.buffer.len() {
self.reset_buffer();
break;
}
}
Poll::Ready(Ok(()))
} else {
if self.is_buffer_empty() {
return Poll::Ready(Ok(()));
}
let prelude = format!("{:x}\r\n", self.buffer_len());
let prelude = prelude.as_bytes();
assert!(
prelude.len() <= MAX_HEADER_SIZE,
"invariant failed: prelude longer than MAX_HEADER_SIZE"
);
let offset = MAX_HEADER_SIZE - prelude.len();
self.buffer[offset..MAX_HEADER_SIZE].clone_from_slice(&prelude);
self.buffer.extend_from_slice(b"\r\n");
self.offset = Some(offset);
self.send(cx)
}
}
fn poll_finish(&mut self, cx: &mut Context) -> Poll<IoResult<()>> {
ok_ready!(self.send(cx));
Pin::new(&mut self.output)
.poll_write(cx, b"0\r\n\r\n" as &[u8])
.map(|r| r.map(|_| ()))
}
pub async fn finish (mut self) -> IoResult<()> {
crate::poll_fn(|cx| self.poll_finish(cx)).await
}
fn priv_poll_write(mut self: Pin<&mut Self>, cx: &mut Context, data: &[u8], data_written: usize) -> Poll<IoResult<usize>> {
let remaining_buffer_space = self.chunks_size - self.buffer_len();
let bytes_to_buffer = std::cmp::min(remaining_buffer_space, data.len());
self.buffer.extend_from_slice(&data[0..bytes_to_buffer]);
let more_to_write: bool = bytes_to_buffer < data.len();
if self.flush_after_write || more_to_write {
ok_ready!(self.send(cx));
}
if more_to_write {
return self.priv_poll_write(cx, &data[bytes_to_buffer..], bytes_to_buffer + data_written);
}
Poll::Ready(Ok(bytes_to_buffer + data_written))
}
}
impl<W> AsyncWrite for Encoder<W>
where
W: AsyncWrite + Unpin,
{
fn poll_write(
self: Pin<&mut Self>,
cx: &mut Context<'_>,
data: &[u8],
) -> Poll<IoResult<usize>> {
self.priv_poll_write(cx, data, 0)
}
fn poll_flush(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<IoResult<()>> {
self.send(cx)
}
fn poll_shutdown(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<IoResult<()>> {
ok_ready!(self.poll_finish(cx));
Pin::new(&mut self.output).poll_shutdown(cx)
}
}
impl<W> Drop for Encoder<W>
where
W: AsyncWrite + Unpin,
{
fn drop(&mut self) {
if !self.is_buffer_empty() {
eprintln!(
"Dropping non-empty Chunked-Transfer encoder. This will cause invalid output."
)
}
}
}
#[cfg(test)]
mod test {
use super::Encoder;
use std::io::Cursor;
use std::io::Write;
use std::str::from_utf8;
use tokio::io::copy;
#[tokio::test]
async fn test() {
let mut source = Cursor::new("hello world".to_string().into_bytes());
let mut dest: Vec<u8> = vec![];
{
let mut encoder = Encoder::with_chunks_size(dest.by_ref(), 5);
copy(&mut source, &mut encoder).await.unwrap();
encoder.finish().await.unwrap();
}
let output = from_utf8(&dest).unwrap();
assert_eq!(output, "5\r\nhello\r\n5\r\n worl\r\n1\r\nd\r\n0\r\n\r\n");
}
#[tokio::test]
async fn flush_after_write() {
let mut source = Cursor::new("hello world".to_string().into_bytes());
let mut dest: Vec<u8> = vec![];
{
let mut encoder = Encoder::with_flush_after_write(dest.by_ref());
copy(&mut source, &mut encoder).await.unwrap();
assert!(encoder.is_buffer_empty());
encoder.finish().await.unwrap();
}
let output = from_utf8(&dest).unwrap();
assert_eq!(output, "b\r\nhello world\r\n0\r\n\r\n");
}
}