use crate::errors::{ParseRangeError, ResumableUploadError, StdResult};
use bytes::Buf;
use core::str::FromStr;
use derive_more::{Debug, Display};
use futures_util::lock::Mutex;
use futures_util::{AsyncRead, AsyncSeek, AsyncSeekExt};
use std::io::Error;
use std::mem::swap;
use std::pin::Pin;
use std::sync::Arc;
use url::Url;
#[derive(Debug)]
pub struct ResumableBody {
#[debug("{:?}", url.as_ref().map(|x| x.as_str()))]
pub(crate) url: Option<Url>,
pub(crate) state: ResumableState,
pub media_body: Option<ResumableMediaUpload<dyn AsyncMediaUploadStream>>,
#[debug(skip)]
pub(crate) save_url: Option<Box<dyn FnOnce(Url)>>,
}
impl ResumableBody {
pub fn new(
stream: Pin<Box<dyn AsyncMediaUploadStream>>,
length: Option<u64>,
mime_type: impl Into<String>,
session_url: Option<Url>,
save_url: Box<dyn FnOnce(Url)>,
) -> Self {
let state = {
if session_url.is_some() {
ResumableState::Resuming
} else {
ResumableState::NotStarted
}
};
let stream = Arc::new(Mutex::new(stream));
Self {
url: session_url,
state,
media_body: Some(ResumableMediaUpload {
mime_type: mime_type.into(),
body: stream,
length,
}),
save_url: Some(save_url),
}
}
pub(crate) fn call_save_url(&mut self, p0: Url) {
if self.save_url.is_some() {
let mut option = None;
swap(&mut self.save_url, &mut option);
option.unwrap()(p0);
}
}
pub(crate) fn is_single_chunk_upload(&self) -> Option<bool> {
match &self.state {
ResumableState::NotStarted => Some(true),
ResumableState::Sending(pending) => {
if pending.ranges.len() == 1 {
let range = &pending.ranges[0];
let total_len = self.media_body.as_ref()?.length?;
let resolved_range = range.resolve_range(Some(total_len)).ok()?;
if resolved_range.start == 0 && resolved_range.end == total_len - 1 {
Some(true)
} else {
Some(false)
}
} else {
None
}
}
ResumableState::Resuming => None,
ResumableState::Done => None,
}
}
}
#[derive(Debug, Clone)]
pub enum ResumableState {
NotStarted,
Sending(PendingChunks),
Resuming,
Done,
}
#[derive(Debug, Clone)]
pub struct PendingChunks {
pub(crate) ranges: Vec<Range>,
}
impl PendingChunks {
pub(crate) fn full(total_size: u64) -> PendingChunks {
Self {
ranges: vec![Range {
start: Some(0),
end: Some(total_size as i64 - 1),
}],
}
}
}
impl PendingChunks {
pub(crate) fn from_ranges(incoming_ranges: String) -> Result<Self, ParseRangeError> {
let ranges: Vec<Range> = incoming_ranges
.split(',')
.filter(|&x| !x.is_empty())
.map(FromStr::from_str)
.collect::<Result<Vec<Range>, ParseRangeError>>()?;
Ok(Self { ranges })
}
}
#[derive(Debug, Clone, Copy, Eq, PartialEq, Hash)]
pub struct Range {
pub(crate) start: Option<u64>,
pub(crate) end: Option<i64>,
}
impl FromStr for Range {
type Err = ParseRangeError;
fn from_str(s: &str) -> Result<Self, Self::Err> {
let start;
let end;
if let Some((start_str, end_str)) = s.split_once('-') {
let mut end_str = end_str.to_string();
if start_str.is_empty() {
start = None;
if !end_str.is_empty() {
end_str = format!("-{end_str}");
}
} else {
start = start_str.parse().ok();
}
if end_str.is_empty() {
end = None;
} else {
end = end_str.parse().ok();
}
Ok(Self { start, end })
} else {
Err(ParseRangeError::StartAndEndEmpty)
}
}
}
#[derive(Debug, Clone, Copy)]
pub struct ResolvedRange {
pub(crate) start: u64,
pub(crate) end: u64,
}
impl Range {
pub fn resolve_range(
&self,
total_len: Option<u64>,
) -> StdResult<ResolvedRange, ResumableUploadError> {
if let Some(total_len) = total_len {
if let Some(&start) = self.start.as_ref() {
let end = self
.end
.as_ref()
.map(|&x| x as u64)
.unwrap_or(total_len - 1);
Ok(ResolvedRange { start, end })
} else if let Some(&original_end) = self.end.as_ref() {
let start;
let end;
if original_end < 0 {
start = total_len - original_end.abs() as u64;
end = total_len - 1;
} else {
start = 0;
end = original_end as u64;
}
Ok(ResolvedRange { start, end })
} else {
Err(ResumableUploadError::ResolveRange)
}
} else {
let start = *self
.start
.as_ref()
.ok_or(ResumableUploadError::ResolveRange)?;
let end = {
let end = *self
.end
.as_ref()
.ok_or(ResumableUploadError::ResolveRange)?;
if end < 0 {
return Err(ResumableUploadError::ResolveRange);
}
end as u64
};
Ok(ResolvedRange { start, end })
}
}
}
#[derive(Debug)]
pub struct ResumableMediaUpload<T: AsyncMediaUploadStream + ?Sized> {
pub mime_type: String,
#[debug(skip)]
pub body: Arc<Mutex<Pin<Box<T>>>>,
pub length: Option<u64>,
}
use crate::media::AsyncMediaUploadStream;
pub use limited_stream::LimitedStream;
pub(crate) mod limited_stream {
use crate::media::AsyncMediaUploadStream;
use core::fmt::Debug;
use futures_util::lock::Mutex;
use futures_util::{AsyncRead, AsyncSeek, AsyncSeekExt};
use std::io::{Error, SeekFrom};
use std::pin::Pin;
use std::sync::Arc;
use std::task::{Context, Poll};
#[derive(Debug)]
pub struct LimitedStream<S: ?Sized> {
limit: u64,
start_position: u64,
inner: Arc<Mutex<Pin<Box<S>>>>,
current_position: u64,
}
impl<S: AsyncMediaUploadStream + ?Sized> LimitedStream<S> {
pub async fn new(inner: Arc<Mutex<Pin<Box<S>>>>, limit: u64) -> Result<Self, Error> {
let start_position = {
let mut inner_guard = inner.lock().await;
inner_guard.stream_position().await?
};
Ok(Self {
limit,
start_position,
inner,
current_position: 0,
})
}
}
impl<S: AsyncMediaUploadStream + ?Sized + Debug> AsyncRead for LimitedStream<S> {
fn poll_read(
mut self: Pin<&mut Self>,
cx: &mut Context<'_>,
buf: &mut [u8],
) -> Poll<Result<usize, Error>> {
dbg!(format!("polling: {:?}", &self));
let current_pos = self.current_position;
if current_pos >= self.limit {
return Poll::Ready(Ok(0));
}
let remaining_limit = self.limit - current_pos;
let read_buf_limit = buf.len().min(remaining_limit as usize);
let my_buf = &mut buf[0..read_buf_limit];
let poll = {
if let Some(mut lock) = self.inner.try_lock() {
let x = lock.as_mut();
x.poll_read(cx, my_buf)
} else {
dbg!("no lock");
Poll::Pending
}
};
match poll {
Poll::Ready(Ok(amount)) => {
self.current_position += amount as u64;
Poll::Ready(Ok(amount))
}
Poll::Ready(Err(err)) => Poll::Ready(Err(err)),
Poll::Pending => Poll::Pending,
}
}
}
impl<S: AsyncMediaUploadStream + ?Sized> AsyncSeek for LimitedStream<S> {
fn poll_seek(
mut self: Pin<&mut Self>,
cx: &mut Context<'_>,
pos: SeekFrom,
) -> Poll<std::io::Result<u64>> {
let limit = (self.start_position + self.limit) as i128;
let start = self.start_position as i128;
let adjusted_pos = match pos {
SeekFrom::Start(start_offset) => {
if start_offset >= self.limit {
SeekFrom::Start(limit as u64)
} else {
SeekFrom::Start(self.start_position + start_offset)
}
}
SeekFrom::End(end_offset) => {
let pos = limit - (end_offset as i128);
if pos >= limit {
SeekFrom::Start(limit as u64)
} else if pos < 0 {
todo!("return some error here")
} else {
SeekFrom::Start(pos as u64)
}
}
SeekFrom::Current(offset) => {
let current_pos = (self.current_position + self.start_position) as i128;
let new_position = current_pos + offset as i128;
if new_position < start {
todo!("return some error here")
} else if new_position > limit {
SeekFrom::Current(
(self.limit as i128 - self.current_position as i128) as i64,
)
} else {
SeekFrom::Current(offset)
}
}
};
let poll = {
let lock = self.inner.try_lock();
if let Some(lock) = lock {
std::pin::pin!(Pin::new(lock)).poll_seek(cx, adjusted_pos)
} else {
Poll::Pending
}
};
match poll {
Poll::Ready(Ok(amount)) => {
self.current_position = amount - self.start_position;
Poll::Ready(Ok(self.current_position))
}
Poll::Ready(Err(err)) => Poll::Ready(Err(err)),
Poll::Pending => Poll::Pending,
}
}
}
#[cfg(test)]
mod tests {
use super::Mutex;
use crate::media::resumable::LimitedStream;
use futures_util::task::noop_waker_ref;
use futures_util::{AsyncReadExt, AsyncSeekExt};
use std::io::Error;
use std::io::{Cursor, SeekFrom};
use std::pin::Pin;
use std::sync::Arc;
use std::task::Poll;
#[derive(Debug)]
struct TestStream {
cursor: Cursor<Vec<u8>>,
}
impl TestStream {
fn new(data: Vec<u8>) -> Self {
Self {
cursor: Cursor::new(data),
}
}
}
impl futures_util::AsyncRead for TestStream {
fn poll_read(
mut self: Pin<&mut Self>,
_cx: &mut std::task::Context<'_>,
buf: &mut [u8],
) -> Poll<Result<usize, Error>> {
Poll::Ready(std::io::Read::read(&mut self.cursor, buf))
}
}
impl futures_util::AsyncSeek for TestStream {
fn poll_seek(
mut self: Pin<&mut Self>,
_cx: &mut std::task::Context<'_>,
pos: std::io::SeekFrom,
) -> Poll<Result<u64, Error>> {
Poll::Ready(std::io::Seek::seek(&mut self.cursor, pos))
}
}
#[tokio::test]
async fn check_bounds_min() {
let data: Vec<u8> = (0..15).collect();
let mut stream = TestStream::new(data);
stream.seek(SeekFrom::Start(5)).await.unwrap();
let mut single = [0u8; 1];
let amount = stream.read(&mut single).await.unwrap();
assert_eq!(amount, 1);
assert_eq!(single[0], 5);
let mutex = Mutex::new(Box::pin(stream));
let mut limited_stream = LimitedStream::new(Arc::new(mutex), 5).await.unwrap();
let mut ten = [0u8; 10];
let amount = limited_stream.read(&mut ten).await.unwrap();
assert_eq!(amount, 5);
assert_eq!(ten, [6, 7, 8, 9, 10, 0, 0, 0, 0, 0]);
let new_position = limited_stream.seek(SeekFrom::Start(0)).await.unwrap();
assert_eq!(new_position, 0);
let new_position = limited_stream.seek(SeekFrom::End(0)).await.unwrap();
assert_eq!(new_position, 5);
let new_position = limited_stream.seek(SeekFrom::End(3)).await.unwrap();
assert_eq!(new_position, 2);
let new_position = limited_stream.seek(SeekFrom::Start(100)).await.unwrap();
assert_eq!(new_position, 5);
let new_position = limited_stream.seek(SeekFrom::Start(0)).await.unwrap();
assert_eq!(new_position, 0);
let mut ten = [0u8; 10];
let amount = limited_stream.read(&mut ten).await.unwrap();
assert_eq!(amount, 5);
assert_eq!(ten, [6, 7, 8, 9, 10, 0, 0, 0, 0, 0]);
let mut ten = [0u8; 10];
let amount = limited_stream.read(&mut ten).await.unwrap();
assert_eq!(amount, 0);
let new_position = limited_stream.seek(SeekFrom::Current(10)).await.unwrap();
assert_eq!(new_position, 5);
let mut ten = [0u8; 10];
let amount = limited_stream.read(&mut ten).await.unwrap();
assert_eq!(amount, 0);
}
}
}