Skip to main content

gapirs_common/api/client/media/
resumable.rs

1use crate::errors::{ParseRangeError, ResumableUploadError, StdResult};
2use bytes::Buf;
3use core::str::FromStr;
4use derive_more::{Debug, Display};
5use futures_util::lock::Mutex;
6use futures_util::{AsyncRead, AsyncSeek, AsyncSeekExt};
7use std::io::Error;
8use std::mem::swap;
9use std::pin::Pin;
10use std::sync::Arc;
11use url::Url;
12
13#[derive(Debug)]
14pub struct ResumableBody {
15    #[debug("{:?}", url.as_ref().map(|x| x.as_str()))]
16    pub(crate) url: Option<Url>,
17    pub(crate) state: ResumableState,
18    pub media_body: Option<ResumableMediaUpload<dyn AsyncMediaUploadStream>>,
19    #[debug(skip)]
20    pub(crate) save_url: Option<Box<dyn FnOnce(Url)>>,
21}
22
23impl ResumableBody {
24    pub fn new(
25        stream: Pin<Box<dyn AsyncMediaUploadStream>>,
26        length: Option<u64>,
27        mime_type: impl Into<String>,
28        session_url: Option<Url>,
29        save_url: Box<dyn FnOnce(Url)>,
30    ) -> Self {
31        let state = {
32            if session_url.is_some() {
33                ResumableState::Resuming
34            } else {
35                ResumableState::NotStarted
36            }
37        };
38        let stream = Arc::new(Mutex::new(stream));
39        Self {
40            url: session_url,
41            state,
42            media_body: Some(ResumableMediaUpload {
43                mime_type: mime_type.into(),
44                body: stream,
45                length,
46            }),
47            save_url: Some(save_url),
48        }
49    }
50    pub(crate) fn call_save_url(&mut self, p0: Url) {
51        if self.save_url.is_some() {
52            let mut option = None;
53            swap(&mut self.save_url, &mut option);
54            option.unwrap()(p0);
55        }
56    }
57    pub(crate) fn is_single_chunk_upload(&self) -> Option<bool> {
58        match &self.state {
59            ResumableState::NotStarted => Some(true),
60            ResumableState::Sending(pending) => {
61                if pending.ranges.len() == 1 {
62                    let range = &pending.ranges[0];
63                    let total_len = self.media_body.as_ref()?.length?;
64                    let resolved_range = range.resolve_range(Some(total_len)).ok()?;
65                    if resolved_range.start == 0 && resolved_range.end == total_len - 1 {
66                        Some(true)
67                    } else {
68                        Some(false)
69                    }
70                } else {
71                    None
72                }
73            }
74            ResumableState::Resuming => None, /*TODO: check if this is actually true?*/
75            ResumableState::Done => None,
76        }
77    }
78}
79
80#[derive(Debug, Clone)]
81pub enum ResumableState {
82    NotStarted,
83    Sending(PendingChunks),
84    Resuming,
85    Done,
86}
87
88#[derive(Debug, Clone)]
89pub struct PendingChunks {
90    pub(crate) ranges: Vec<Range>,
91}
92
93impl PendingChunks {
94    pub(crate) fn full(total_size: u64) -> PendingChunks {
95        Self {
96            ranges: vec![Range {
97                start: Some(0),
98                end: Some(total_size as i64 - 1),
99            }],
100        }
101    }
102}
103
104impl PendingChunks {
105    pub(crate) fn from_ranges(incoming_ranges: String) -> Result<Self, ParseRangeError> {
106        let ranges: Vec<Range> = incoming_ranges
107            .split(',')
108            .filter(|&x| !x.is_empty())
109            .map(FromStr::from_str)
110            .collect::<Result<Vec<Range>, ParseRangeError>>()?;
111
112        Ok(Self { ranges })
113    }
114}
115
116#[derive(Debug, Clone, Copy, Eq, PartialEq, Hash)]
117pub struct Range {
118    pub(crate) start: Option<u64>,
119    pub(crate) end: Option<i64>,
120}
121impl FromStr for Range {
122    type Err = ParseRangeError;
123
124    fn from_str(s: &str) -> Result<Self, Self::Err> {
125        let start;
126        let end;
127        if let Some((start_str, end_str)) = s.split_once('-') {
128            let mut end_str = end_str.to_string();
129            if start_str.is_empty() {
130                start = None;
131                if !end_str.is_empty() {
132                    end_str = format!("-{end_str}");
133                }
134            } else {
135                start = start_str.parse().ok();
136            }
137            if end_str.is_empty() {
138                end = None;
139            } else {
140                end = end_str.parse().ok();
141            }
142            Ok(Self { start, end })
143        } else {
144            Err(ParseRangeError::StartAndEndEmpty)
145        }
146    }
147}
148#[derive(Debug, Clone, Copy)]
149pub struct ResolvedRange {
150    pub(crate) start: u64,
151    pub(crate) end: u64,
152}
153impl Range {
154    pub fn resolve_range(
155        &self,
156        total_len: Option<u64>,
157    ) -> StdResult<ResolvedRange, ResumableUploadError> {
158        if let Some(total_len) = total_len {
159            if let Some(&start) = self.start.as_ref() {
160                let end = self
161                    .end
162                    .as_ref()
163                    .map(|&x| x as u64)
164                    .unwrap_or(total_len - 1);
165                Ok(ResolvedRange { start, end })
166            } else if let Some(&original_end) = self.end.as_ref() {
167                let start;
168                let end;
169                if original_end < 0 {
170                    start = total_len - original_end.abs() as u64;
171                    end = total_len - 1;
172                } else {
173                    start = 0;
174                    end = original_end as u64;
175                }
176                Ok(ResolvedRange { start, end })
177            } else {
178                Err(ResumableUploadError::ResolveRange)
179            }
180        } else {
181            let start = *self
182                .start
183                .as_ref()
184                .ok_or(ResumableUploadError::ResolveRange)?;
185            let end = {
186                let end = *self
187                    .end
188                    .as_ref()
189                    .ok_or(ResumableUploadError::ResolveRange)?;
190
191                if end < 0 {
192                    return Err(ResumableUploadError::ResolveRange);
193                }
194                end as u64
195            };
196            Ok(ResolvedRange { start, end })
197        }
198    }
199}
200
201#[derive(Debug)]
202pub struct ResumableMediaUpload<T: AsyncMediaUploadStream + ?Sized> {
203    pub mime_type: String,
204    #[debug(skip)]
205    pub body: Arc<Mutex<Pin<Box<T>>>>,
206    pub length: Option<u64>,
207}
208use crate::media::AsyncMediaUploadStream;
209pub use limited_stream::LimitedStream;
210
211pub(crate) mod limited_stream {
212    use crate::media::AsyncMediaUploadStream;
213    use core::fmt::Debug;
214    use futures_util::lock::Mutex;
215    use futures_util::{AsyncRead, AsyncSeek, AsyncSeekExt};
216    use std::io::{Error, SeekFrom};
217    use std::pin::Pin;
218    use std::sync::Arc;
219    use std::task::{Context, Poll};
220
221    #[derive(Debug)]
222    pub struct LimitedStream<S: ?Sized> {
223        limit: u64,
224        start_position: u64,
225        inner: Arc<Mutex<Pin<Box<S>>>>,
226        current_position: u64,
227    }
228    impl<S: AsyncMediaUploadStream + ?Sized> LimitedStream<S> {
229        pub async fn new(inner: Arc<Mutex<Pin<Box<S>>>>, limit: u64) -> Result<Self, Error> {
230            let start_position = {
231                let mut inner_guard = inner.lock().await;
232                inner_guard.stream_position().await?
233            };
234            Ok(Self {
235                limit,
236                start_position,
237                inner,
238                current_position: 0,
239            })
240        }
241    }
242    impl<S: AsyncMediaUploadStream + ?Sized + Debug> AsyncRead for LimitedStream<S> {
243        fn poll_read(
244            mut self: Pin<&mut Self>,
245            cx: &mut Context<'_>,
246            buf: &mut [u8],
247        ) -> Poll<Result<usize, Error>> {
248            dbg!(format!("polling: {:?}", &self));
249            let current_pos = self.current_position;
250            if current_pos >= self.limit {
251                return Poll::Ready(Ok(0));
252            }
253            let remaining_limit = self.limit - current_pos;
254            let read_buf_limit = buf.len().min(remaining_limit as usize);
255            let my_buf = &mut buf[0..read_buf_limit];
256
257            let poll = {
258                if let Some(mut lock) = self.inner.try_lock() {
259                    let x = lock.as_mut();
260                    x.poll_read(cx, my_buf)
261                } else {
262                    dbg!("no lock");
263                    // cx.waker().wake_by_ref();//TODO: check, if this is needed
264                    Poll::Pending
265                }
266            };
267            match poll {
268                Poll::Ready(Ok(amount)) => {
269                    self.current_position += amount as u64;
270                    Poll::Ready(Ok(amount))
271                }
272                Poll::Ready(Err(err)) => Poll::Ready(Err(err)),
273                Poll::Pending => Poll::Pending,
274            }
275        }
276    }
277
278    impl<S: AsyncMediaUploadStream + ?Sized> AsyncSeek for LimitedStream<S> {
279        fn poll_seek(
280            mut self: Pin<&mut Self>,
281            cx: &mut Context<'_>,
282            pos: SeekFrom,
283        ) -> Poll<std::io::Result<u64>> {
284            let limit = (self.start_position + self.limit) as i128;
285            let start = self.start_position as i128;
286            let adjusted_pos = match pos {
287                SeekFrom::Start(start_offset) => {
288                    if start_offset >= self.limit {
289                        SeekFrom::Start(limit as u64)
290                    } else {
291                        SeekFrom::Start(self.start_position + start_offset)
292                    }
293                }
294                SeekFrom::End(end_offset) => {
295                    let pos = limit - (end_offset as i128);
296                    if pos >= limit {
297                        SeekFrom::Start(limit as u64)
298                    } else if pos < 0 {
299                        todo!("return some error here")
300                    } else {
301                        SeekFrom::Start(pos as u64)
302                    }
303                }
304                SeekFrom::Current(offset) => {
305                    let current_pos = (self.current_position + self.start_position) as i128;
306                    let new_position = current_pos + offset as i128;
307                    if new_position < start {
308                        todo!("return some error here")
309                    } else if new_position > limit {
310                        SeekFrom::Current(
311                            (self.limit as i128 - self.current_position as i128) as i64,
312                        )
313                    } else {
314                        SeekFrom::Current(offset)
315                    }
316                }
317            };
318
319            let poll = {
320                let lock = self.inner.try_lock();
321                if let Some(lock) = lock {
322                    std::pin::pin!(Pin::new(lock)).poll_seek(cx, adjusted_pos)
323                } else {
324                    Poll::Pending
325                }
326            };
327            match poll {
328                Poll::Ready(Ok(amount)) => {
329                    self.current_position = amount - self.start_position;
330                    Poll::Ready(Ok(self.current_position))
331                }
332                Poll::Ready(Err(err)) => Poll::Ready(Err(err)),
333                Poll::Pending => Poll::Pending,
334            }
335        }
336    }
337
338    #[cfg(test)]
339    mod tests {
340        use super::Mutex;
341        use crate::media::resumable::LimitedStream;
342
343        use futures_util::task::noop_waker_ref;
344        use futures_util::{AsyncReadExt, AsyncSeekExt};
345        use std::io::Error;
346        use std::io::{Cursor, SeekFrom};
347        use std::pin::Pin;
348        use std::sync::Arc;
349        use std::task::Poll;
350
351        #[derive(Debug)]
352        struct TestStream {
353            cursor: Cursor<Vec<u8>>,
354        }
355        impl TestStream {
356            fn new(data: Vec<u8>) -> Self {
357                Self {
358                    cursor: Cursor::new(data),
359                }
360            }
361        }
362        impl futures_util::AsyncRead for TestStream {
363            fn poll_read(
364                mut self: Pin<&mut Self>,
365                _cx: &mut std::task::Context<'_>,
366                buf: &mut [u8],
367            ) -> Poll<Result<usize, Error>> {
368                Poll::Ready(std::io::Read::read(&mut self.cursor, buf))
369            }
370        }
371        impl futures_util::AsyncSeek for TestStream {
372            fn poll_seek(
373                mut self: Pin<&mut Self>,
374                _cx: &mut std::task::Context<'_>,
375                pos: std::io::SeekFrom,
376            ) -> Poll<Result<u64, Error>> {
377                Poll::Ready(std::io::Seek::seek(&mut self.cursor, pos))
378            }
379        }
380
381        #[tokio::test]
382        async fn check_bounds_min() {
383            let data: Vec<u8> = (0..15).collect();
384            let mut stream = TestStream::new(data);
385            stream.seek(SeekFrom::Start(5)).await.unwrap();
386            let mut single = [0u8; 1];
387            let amount = stream.read(&mut single).await.unwrap();
388            assert_eq!(amount, 1);
389            assert_eq!(single[0], 5);
390            let mutex = Mutex::new(Box::pin(stream));
391            let mut limited_stream = LimitedStream::new(Arc::new(mutex), 5).await.unwrap();
392            let mut ten = [0u8; 10];
393            let amount = limited_stream.read(&mut ten).await.unwrap();
394            assert_eq!(amount, 5);
395            assert_eq!(ten, [6, 7, 8, 9, 10, 0, 0, 0, 0, 0]);
396            let new_position = limited_stream.seek(SeekFrom::Start(0)).await.unwrap();
397            assert_eq!(new_position, 0);
398            let new_position = limited_stream.seek(SeekFrom::End(0)).await.unwrap();
399            assert_eq!(new_position, 5);
400            let new_position = limited_stream.seek(SeekFrom::End(3)).await.unwrap();
401            assert_eq!(new_position, 2);
402            let new_position = limited_stream.seek(SeekFrom::Start(100)).await.unwrap();
403            assert_eq!(new_position, 5);
404            let new_position = limited_stream.seek(SeekFrom::Start(0)).await.unwrap();
405            assert_eq!(new_position, 0);
406            let mut ten = [0u8; 10];
407            let amount = limited_stream.read(&mut ten).await.unwrap();
408            assert_eq!(amount, 5);
409            assert_eq!(ten, [6, 7, 8, 9, 10, 0, 0, 0, 0, 0]);
410            let mut ten = [0u8; 10];
411            let amount = limited_stream.read(&mut ten).await.unwrap();
412            assert_eq!(amount, 0);
413
414            let new_position = limited_stream.seek(SeekFrom::Current(10)).await.unwrap();
415            assert_eq!(new_position, 5);
416            let mut ten = [0u8; 10];
417            let amount = limited_stream.read(&mut ten).await.unwrap();
418            assert_eq!(amount, 0);
419        }
420    }
421}