gapirs_common/api/client/media/
resumable.rs1use 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, 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 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}