use bytes::{Bytes, BytesMut};
use futures::stream::{BoxStream, StreamExt};
use std::io;
use std::path::Path;
use tokio::io::{AsyncRead, AsyncReadExt};
const SOURCE_CHUNK_BYTES: usize = 64 * 1024;
pub type PayloadStream = BoxStream<'static, io::Result<Bytes>>;
pub struct PayloadSource {
stream: PayloadStream,
size_bytes: Option<u64>,
}
impl std::fmt::Debug for PayloadSource {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("PayloadSource")
.field("size_bytes", &self.size_bytes)
.finish_non_exhaustive()
}
}
impl PayloadSource {
pub fn stream(stream: PayloadStream) -> Self {
Self {
stream,
size_bytes: None,
}
}
pub fn sized_stream(stream: PayloadStream, size_bytes: u64) -> Self {
Self {
stream,
size_bytes: Some(size_bytes),
}
}
pub fn reader<R>(reader: R) -> Self
where
R: AsyncRead + Send + Unpin + 'static,
{
Self::stream(read_in_chunks(reader))
}
pub async fn open_file(path: impl AsRef<Path>) -> io::Result<Self> {
let file = tokio::fs::File::open(path.as_ref()).await?;
let size_bytes = file.metadata().await?.len();
Ok(Self::sized_stream(read_in_chunks(file), size_bytes))
}
pub fn size_bytes(&self) -> Option<u64> {
self.size_bytes
}
pub fn into_stream(self) -> (PayloadStream, Option<u64>) {
(self.stream, self.size_bytes)
}
}
fn read_in_chunks<R>(reader: R) -> PayloadStream
where
R: AsyncRead + Send + Unpin + 'static,
{
futures::stream::unfold(Some(reader), |reader| async move {
let mut reader = reader?;
let mut buffer = BytesMut::with_capacity(SOURCE_CHUNK_BYTES);
match reader.read_buf(&mut buffer).await {
Ok(0) => None,
Ok(_) => Some((Ok(buffer.freeze()), Some(reader))),
Err(error) => Some((Err(error), None)),
}
})
.boxed()
}
pub(crate) struct PartReader {
stream: PayloadStream,
carry: Option<Bytes>,
part_bytes: usize,
exhausted: bool,
}
impl PartReader {
pub(crate) fn new(stream: PayloadStream, part_bytes: usize) -> Self {
Self {
stream,
carry: None,
part_bytes: part_bytes.max(1),
exhausted: false,
}
}
pub(crate) async fn next_part(&mut self) -> io::Result<Option<Bytes>> {
let mut buffer: Option<BytesMut> = None;
loop {
let filled = buffer.as_ref().map_or(0, BytesMut::len);
if filled >= self.part_bytes {
break;
}
let mut chunk = match self.carry.take() {
Some(chunk) => chunk,
None if self.exhausted => break,
None => match self.stream.next().await {
Some(chunk) => chunk?,
None => {
self.exhausted = true;
break;
}
},
};
let take = (self.part_bytes - filled).min(chunk.len());
let taken = chunk.split_to(take);
if !chunk.is_empty() {
self.carry = Some(chunk);
}
match &mut buffer {
Some(buffer) => buffer.extend_from_slice(&taken),
None if taken.len() == self.part_bytes => return Ok(Some(taken)),
None => {
let mut fresh = BytesMut::with_capacity(self.part_bytes);
fresh.extend_from_slice(&taken);
buffer = Some(fresh);
}
}
}
Ok(buffer
.filter(|buffer| !buffer.is_empty())
.map(BytesMut::freeze))
}
}
#[cfg(test)]
mod tests {
use super::*;
fn source(chunks: Vec<&'static [u8]>) -> PayloadStream {
futures::stream::iter(chunks.into_iter().map(|chunk| Ok(Bytes::from(chunk)))).boxed()
}
async fn cut(reader: &mut PartReader) -> Vec<Vec<u8>> {
let mut parts = Vec::new();
while let Some(part) = reader.next_part().await.expect("cut a part") {
parts.push(part.to_vec());
}
parts
}
#[tokio::test]
async fn parts_are_cut_regardless_of_how_the_source_chunked_them() {
let mut reader = PartReader::new(source(vec![b"abcde", b"fg", b"hijkl"]), 4);
assert_eq!(
cut(&mut reader).await,
vec![b"abcd".to_vec(), b"efgh".to_vec(), b"ijkl".to_vec()]
);
}
#[tokio::test]
async fn the_last_part_is_whatever_is_left() {
let mut reader = PartReader::new(source(vec![b"abcdef"]), 4);
assert_eq!(
cut(&mut reader).await,
vec![b"abcd".to_vec(), b"ef".to_vec()]
);
let mut reader = PartReader::new(source(vec![b"abcd"]), 4);
assert_eq!(cut(&mut reader).await, vec![b"abcd".to_vec()]);
}
#[tokio::test]
async fn an_empty_source_produces_no_parts() {
let mut reader = PartReader::new(source(vec![]), 4);
assert!(reader.next_part().await.expect("cut a part").is_none());
}
#[tokio::test]
async fn a_file_source_knows_its_length() {
let directory = tempfile::tempdir().expect("tempdir");
let path = directory.path().join("payload.bin");
std::fs::write(&path, vec![7u8; 5_000]).expect("write payload");
let source = PayloadSource::open_file(&path).await.expect("open payload");
assert_eq!(source.size_bytes(), Some(5_000));
let (stream, _) = source.into_stream();
let mut reader = PartReader::new(stream, 4_096);
let first = reader.next_part().await.expect("cut").expect("first part");
let second = reader.next_part().await.expect("cut").expect("second part");
assert_eq!(first.len(), 4_096);
assert_eq!(second.len(), 904);
assert!(reader.next_part().await.expect("cut").is_none());
}
#[tokio::test]
async fn a_reader_source_declares_no_length() {
let source = PayloadSource::reader(std::io::Cursor::new(vec![1u8; 100]));
assert_eq!(source.size_bytes(), None);
let (stream, _) = source.into_stream();
let mut reader = PartReader::new(stream, 64);
assert_eq!(
reader.next_part().await.expect("cut").expect("part").len(),
64
);
assert_eq!(
reader.next_part().await.expect("cut").expect("part").len(),
36
);
assert!(reader.next_part().await.expect("cut").is_none());
}
}