1use color_eyre::Result;
10use color_eyre::eyre::eyre;
11use std::path::Path;
12use std::sync::Arc;
13
14use crate::unfinished::{Claim, Writer};
15
16#[derive(Debug, Clone)]
20pub struct TempDownload(Arc<Held>);
21
22#[derive(Debug)]
25struct Held {
26 path: tempfile::TempPath,
27 _claim: Option<Claim>,
28 #[cfg(unix)]
32 _lock: Option<std::fs::File>,
33}
34
35#[cfg(unix)]
38pub(crate) fn held_elsewhere(path: &Path) -> bool {
39 use fs2::FileExt;
40 match std::fs::File::open(path) {
41 Ok(file) => file.try_lock_exclusive().is_err(),
42 Err(_) => true,
43 }
44}
45
46impl TempDownload {
47 pub fn create(dir: Option<&Path>, extension: Option<&str>) -> Result<tempfile::NamedTempFile> {
50 let dir = dir
51 .map(Path::to_path_buf)
52 .unwrap_or_else(std::env::temp_dir);
53 let suffix = extension
54 .map(|e| format!(".{e}"))
55 .unwrap_or_else(|| ".tmp".to_string());
56 tempfile::Builder::new()
57 .suffix(&suffix)
58 .tempfile_in(&dir)
59 .map_err(|_| eyre!("Could not create a temporary file."))
60 }
61
62 pub fn keep(file: tempfile::NamedTempFile) -> TempDownload {
64 Self::held(file, None)
65 }
66
67 pub(crate) fn held(file: tempfile::NamedTempFile, claim: Option<Claim>) -> TempDownload {
68 #[cfg(unix)]
69 let lock = file.as_file().try_clone().ok().filter(|lock| {
70 use fs2::FileExt;
71 FileExt::try_lock_shared(lock).is_ok()
72 });
73 TempDownload(Arc::new(Held {
74 path: file.into_temp_path(),
75 _claim: claim,
76 #[cfg(unix)]
77 _lock: lock,
78 }))
79 }
80
81 pub fn path(&self) -> &Path {
82 &self.0.path
83 }
84}
85
86pub const QUEUED_CHUNKS: usize = 4;
88
89const STALL_CHECK: std::time::Duration = std::time::Duration::from_millis(100);
91
92pub type Opened<S> = std::result::Result<(S, Option<u64>), String>;
95
96#[derive(Debug)]
98pub enum StreamError {
99 Open(String),
101 Read(String),
103 Write(color_eyre::Report),
105 Short { expected: u64, got: u64 },
107 Cut,
109}
110
111enum Piece<B> {
113 Opened(Option<u64>),
115 Refused(String),
116 Chunk(B),
117 Failed(String),
118 End,
119}
120
121fn receive<B: AsRef<[u8]>>(
124 stop: &impl Fn() -> bool,
125 mut next: impl FnMut() -> Option<Piece<B>>,
126 mut write: impl FnMut(&[u8]) -> Result<()>,
127) -> std::result::Result<u64, StreamError> {
128 let mut expected = None;
129 let mut written = 0u64;
130 loop {
131 if stop() {
132 return Err(StreamError::Cut);
133 }
134 match next() {
135 Some(Piece::Opened(len)) => expected = len,
136 Some(Piece::Refused(error)) => return Err(StreamError::Open(error)),
137 Some(Piece::Chunk(chunk)) => {
138 let chunk = chunk.as_ref();
139 write(chunk).map_err(StreamError::Write)?;
140 written += chunk.len() as u64;
141 }
142 Some(Piece::Failed(error)) => return Err(StreamError::Read(error)),
143 Some(Piece::End) => break,
144 None => return Err(StreamError::Cut),
145 }
146 }
147 match expected {
148 Some(expected) if expected != written => Err(StreamError::Short {
149 expected,
150 got: written,
151 }),
152 _ => Ok(written),
153 }
154}
155
156fn fill_temp(
164 dir: Option<&Path>,
165 extension: Option<&str>,
166 writer: &Writer,
167 fill: impl FnOnce(&mut dyn FnMut(&[u8]) -> Result<()>) -> std::result::Result<u64, StreamError>,
168) -> std::result::Result<TempDownload, StreamError> {
169 use std::io::Write;
170
171 let Some((mut file, claim)) = writer
172 .create(|| TempDownload::create(dir, extension))
173 .map_err(StreamError::Write)?
174 else {
175 return Err(StreamError::Cut);
176 };
177 let unwritable = |e: std::io::Error| eyre!("Could not write the downloaded file: {e}");
178 let filled = fill(&mut |chunk| file.write_all(chunk).map_err(unwritable))
179 .and_then(|_| file.flush().map_err(|e| StreamError::Write(unwritable(e))));
180 match filled {
181 Ok(_) => Ok(TempDownload::held(file, Some(claim))),
182 Err(e) => {
183 drop(file);
185 drop(claim);
186 Err(e)
187 }
188 }
189}
190
191#[cfg(feature = "cloud")]
201pub fn stream_into<O, S, B, E>(
202 runtime: &tokio::runtime::Handle,
203 open: O,
204 stop: impl Fn() -> bool + Clone + Send + Sync + 'static,
205 write: impl FnMut(&[u8]) -> Result<()>,
206) -> std::result::Result<u64, StreamError>
207where
208 O: std::future::Future<Output = Opened<S>> + Send + 'static,
209 S: futures::Stream<Item = std::result::Result<B, E>> + Send + 'static,
210 B: AsRef<[u8]> + Send + 'static,
211 E: std::fmt::Display,
212{
213 let (tx, mut rx) = tokio::sync::mpsc::channel(QUEUED_CHUNKS);
214 let stopped = stop.clone();
215 runtime.spawn(async move {
219 let gone = || stopped() || tx.is_closed();
222 let (stream, len) = match until_stopped(open, &gone).await {
223 Some(Ok(opened)) => opened,
224 Some(Err(error)) => {
225 let _ = tx.send(Piece::Refused(error)).await;
226 return;
227 }
228 None => return,
229 };
230 if tx.send(Piece::Opened(len)).await.is_err() {
231 return;
232 }
233 let mut stream = std::pin::pin!(stream);
234 loop {
235 let next = futures::StreamExt::next(&mut stream);
236 let piece = match until_stopped(next, &gone).await {
237 None => return,
238 Some(Some(Ok(chunk))) => Piece::Chunk(chunk),
239 Some(Some(Err(error))) => Piece::Failed(error.to_string()),
240 Some(None) => Piece::End,
241 };
242 let last = !matches!(piece, Piece::Chunk(_));
243 if tx.send(piece).await.is_err() || last {
246 return;
247 }
248 }
249 });
250 receive(&stop, || rx.blocking_recv(), write)
251}
252
253#[cfg(feature = "cloud")]
255async fn until_stopped<F: std::future::Future>(
256 future: F,
257 stop: &impl Fn() -> bool,
258) -> Option<F::Output> {
259 let mut future = std::pin::pin!(future);
260 loop {
261 if stop() {
262 return None;
263 }
264 if let Ok(output) = tokio::time::timeout(STALL_CHECK, future.as_mut()).await {
265 return Some(output);
266 }
267 }
268}
269
270#[cfg(feature = "cloud")]
274pub(crate) fn stream_to_temp<O, S, B, E>(
275 runtime: &tokio::runtime::Handle,
276 dir: Option<&Path>,
277 extension: Option<&str>,
278 open: O,
279 writer: &Writer,
280) -> std::result::Result<TempDownload, StreamError>
281where
282 O: std::future::Future<Output = Opened<S>> + Send + 'static,
283 S: futures::Stream<Item = std::result::Result<B, E>> + Send + 'static,
284 B: AsRef<[u8]> + Send + 'static,
285 E: std::fmt::Display,
286{
287 let stop = {
288 let writer = writer.clone();
289 move || writer.stopped()
290 };
291 fill_temp(dir, extension, writer, |write| {
292 stream_into(runtime, open, stop, write)
293 })
294}
295
296const READ_CHUNK: usize = 64 * 1024;
298
299pub fn read_into<R: std::io::Read>(
308 open: impl FnOnce() -> Opened<R> + Send + 'static,
309 stop: impl Fn() -> bool,
310 write: impl FnMut(&[u8]) -> Result<()>,
311) -> std::result::Result<u64, StreamError> {
312 use std::sync::mpsc::RecvTimeoutError;
313
314 let (tx, rx) = std::sync::mpsc::sync_channel(QUEUED_CHUNKS);
315 std::thread::Builder::new()
316 .name("datui-download".to_string())
317 .spawn(move || {
318 let mut reader = match open() {
319 Ok((reader, len)) => {
320 if tx.send(Piece::Opened(len)).is_err() {
321 return;
322 }
323 reader
324 }
325 Err(error) => {
326 let _ = tx.send(Piece::Refused(error));
327 return;
328 }
329 };
330 loop {
331 let mut chunk = vec![0; READ_CHUNK];
332 let piece = match reader.read(&mut chunk) {
333 Ok(0) => Piece::End,
334 Ok(read) => {
335 chunk.truncate(read);
336 Piece::Chunk(chunk)
337 }
338 Err(error) if error.kind() == std::io::ErrorKind::Interrupted => continue,
339 Err(error) => Piece::Failed(error.to_string()),
340 };
341 let last = !matches!(piece, Piece::Chunk(_));
342 if tx.send(piece).is_err() || last {
343 return;
344 }
345 }
346 })
347 .map_err(|e| StreamError::Open(e.to_string()))?;
348 let next = || loop {
349 match rx.recv_timeout(STALL_CHECK) {
350 Ok(piece) => return Some(piece),
351 Err(RecvTimeoutError::Timeout) if !stop() => {}
352 Err(RecvTimeoutError::Timeout) => return None,
353 Err(RecvTimeoutError::Disconnected) => {
355 return Some(Piece::Failed("the download stopped".to_string()));
356 }
357 }
358 };
359 receive(&stop, next, write)
360}
361
362#[cfg(feature = "http")]
370pub(crate) fn read_to_temp<R: std::io::Read>(
371 dir: Option<&Path>,
372 extension: Option<&str>,
373 open: impl FnOnce() -> Opened<R> + Send + 'static,
374 writer: &Writer,
375 limit: Option<u64>,
376) -> std::result::Result<TempDownload, StreamError> {
377 let mut arrived = 0u64;
378 fill_temp(dir, extension, writer, |write| {
379 read_into(
380 open,
381 || writer.stopped(),
382 |chunk| {
383 arrived += chunk.len() as u64;
384 if limit.is_some_and(|limit| arrived > limit) {
385 return Err(color_eyre::Report::new(PastLimit));
386 }
387 write(chunk)
388 },
389 )
390 })
391}
392
393#[cfg(any(feature = "http", feature = "cloud"))]
395#[derive(Debug)]
396pub(crate) struct PastLimit;
397
398#[cfg(any(feature = "http", feature = "cloud"))]
399impl std::fmt::Display for PastLimit {
400 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
401 f.write_str("The download passed the size it may fetch without asking.")
402 }
403}
404
405#[cfg(any(feature = "http", feature = "cloud"))]
406impl std::error::Error for PastLimit {}
407
408pub(crate) fn spool_to_temp<R: std::io::Read>(
412 dir: Option<&Path>,
413 open: impl FnOnce() -> Opened<R> + Send + 'static,
414 writer: &Writer,
415 read: &std::sync::atomic::AtomicU64,
416) -> std::result::Result<TempDownload, StreamError> {
417 use std::sync::atomic::Ordering;
418 fill_temp(dir, None, writer, |write| {
419 read_into(
420 open,
421 || writer.stopped(),
422 |chunk| {
423 write(chunk)?;
424 read.fetch_add(chunk.len() as u64, Ordering::Relaxed);
425 Ok(())
426 },
427 )
428 })
429}
430
431#[cfg(all(test, feature = "cloud"))]
432mod tests {
433 use super::*;
434 use futures::StreamExt;
435 use std::sync::atomic::{AtomicBool, AtomicU64, AtomicUsize, Ordering};
436
437 const CHUNK: usize = 64 * 1024;
438
439 fn runtime() -> tokio::runtime::Runtime {
440 tokio::runtime::Builder::new_multi_thread()
441 .worker_threads(1)
442 .enable_all()
443 .build()
444 .expect("runtime")
445 }
446
447 fn chunk(i: usize, len: usize) -> Vec<u8> {
450 (0..len).map(|j| ((i * 31 + j * 7) % 251) as u8).collect()
451 }
452
453 #[derive(Default)]
456 struct Live {
457 now: AtomicU64,
458 peak: AtomicU64,
459 }
460
461 struct Tracked(Vec<u8>, Arc<Live>);
462
463 impl Tracked {
464 fn new(bytes: Vec<u8>, live: &Arc<Live>) -> Tracked {
465 let now = live.now.fetch_add(bytes.len() as u64, Ordering::SeqCst) + bytes.len() as u64;
466 live.peak.fetch_max(now, Ordering::SeqCst);
467 Tracked(bytes, live.clone())
468 }
469 }
470
471 impl Drop for Tracked {
472 fn drop(&mut self) {
473 self.1.now.fetch_sub(self.0.len() as u64, Ordering::SeqCst);
474 }
475 }
476
477 impl AsRef<[u8]> for Tracked {
478 fn as_ref(&self) -> &[u8] {
479 &self.0
480 }
481 }
482
483 struct DropFlag(Arc<AtomicBool>);
485
486 impl Drop for DropFlag {
487 fn drop(&mut self) {
488 self.0.store(true, Ordering::SeqCst);
489 }
490 }
491
492 async fn opened<S>(stream: S, len: Option<u64>) -> Opened<S> {
493 Ok((stream, len))
494 }
495
496 fn never() -> impl Fn() -> bool + Clone + Send + Sync + 'static {
497 || false
498 }
499
500 fn unstopped() -> Writer {
502 Writer::default()
503 }
504
505 fn files_in(dir: &Path) -> usize {
506 std::fs::read_dir(dir).unwrap().count()
507 }
508
509 #[test]
511 fn a_stream_lands_whole_in_order() {
512 let rt = runtime();
513 let dir = tempfile::tempdir().unwrap();
514 let chunks = (0..40).map(|i| chunk(i, 1000 + i * 13)).collect::<Vec<_>>();
515 let whole = chunks.concat();
516 let stream = futures::stream::iter(chunks.into_iter().map(Ok::<_, String>));
517 let file = stream_to_temp(
518 rt.handle(),
519 Some(dir.path()),
520 Some("csv.gz"),
521 opened(stream, Some(whole.len() as u64)),
522 &unstopped(),
523 )
524 .unwrap();
525 assert_eq!(std::fs::read(file.path()).unwrap(), whole);
526 assert!(file.path().to_string_lossy().ends_with(".csv.gz"));
527 assert_eq!(file.path().parent(), Some(dir.path()));
528 let held = file.clone();
529 drop(file);
530 assert!(held.path().exists(), "a holder keeps it");
531 let path = held.path().to_path_buf();
532 drop(held);
533 assert!(!path.exists(), "the last holder removes it");
534 }
535
536 #[test]
540 fn a_download_is_claimed_for_as_long_as_it_lives() {
541 let rt = runtime();
542 let dir = tempfile::tempdir().unwrap();
543 let unfinished = crate::unfinished::Unfinished::default();
544 let stop = Arc::new(AtomicBool::new(false));
545 let writer = unfinished.writer(stop.clone());
546 let stream = futures::stream::iter(vec![Ok::<_, String>(chunk(0, 10))]);
547 let file = stream_to_temp(
548 rt.handle(),
549 Some(dir.path()),
550 None,
551 opened(stream, None),
552 &writer,
553 )
554 .unwrap();
555 let held = file.clone();
556 drop(file);
557 assert!(unfinished.writing(), "a holder keeps the claim");
558 let path = held.path().to_path_buf();
559 drop(held);
560 assert!(!path.exists());
561 assert!(!unfinished.writing(), "the claim goes with the file");
562
563 stop.store(true, Ordering::SeqCst);
564 let stream = futures::stream::iter(vec![Ok::<_, String>(chunk(0, 10))]);
565 let error = stream_to_temp(
566 rt.handle(),
567 Some(dir.path()),
568 None,
569 opened(stream, None),
570 &writer,
571 )
572 .unwrap_err();
573 assert!(matches!(error, StreamError::Cut), "{error:?}");
574 assert_eq!(files_in(dir.path()), 0);
575 }
576
577 #[test]
581 fn a_slow_writer_holds_the_stream_back() {
582 let rt = runtime();
583 let live = Arc::new(Live::default());
584 let pulled = Arc::new(AtomicUsize::new(0));
585 let chunks = 64;
586 let stream = {
587 let (live, pulled) = (live.clone(), pulled.clone());
588 futures::stream::iter(0..chunks).map(move |i| {
589 pulled.fetch_add(1, Ordering::SeqCst);
590 Ok::<_, String>(Tracked::new(chunk(i, CHUNK), &live))
591 })
592 };
593 let mut written = 0usize;
594 let mut ahead = 0usize;
595 let mut bytes = Vec::new();
596 let total = stream_into(rt.handle(), opened(stream, None), never(), |piece| {
597 std::thread::sleep(std::time::Duration::from_millis(2));
598 written += 1;
599 ahead = ahead.max(pulled.load(Ordering::SeqCst) - written);
600 bytes.extend_from_slice(piece);
601 Ok(())
602 })
603 .unwrap();
604 assert_eq!(total, (chunks * CHUNK) as u64);
605 assert_eq!(
606 bytes,
607 (0..chunks)
608 .flat_map(|i| chunk(i, CHUNK))
609 .collect::<Vec<_>>()
610 );
611 let bound = QUEUED_CHUNKS + 2;
613 assert!(ahead <= bound, "the store ran {ahead} chunks ahead");
614 let peak = live.peak.load(Ordering::SeqCst);
615 assert!(
616 peak <= (bound * CHUNK) as u64,
617 "{peak} bytes held at once, of {} streamed",
618 chunks * CHUNK
619 );
620 assert_eq!(live.now.load(Ordering::SeqCst), 0, "every chunk was let go");
621 }
622
623 #[test]
625 fn a_failure_mid_stream_leaves_no_file() {
626 let rt = runtime();
627 let dir = tempfile::tempdir().unwrap();
628 let stream = futures::stream::iter(vec![
629 Ok(chunk(0, CHUNK)),
630 Ok(chunk(1, CHUNK)),
631 Err("connection reset".to_string()),
632 Ok(chunk(3, CHUNK)),
633 ]);
634 let error = stream_to_temp(
635 rt.handle(),
636 Some(dir.path()),
637 None,
638 opened(stream, None),
639 &unstopped(),
640 )
641 .unwrap_err();
642 assert!(
643 matches!(&error, StreamError::Read(e) if e == "connection reset"),
644 "{error:?}"
645 );
646 assert_eq!(files_in(dir.path()), 0);
647
648 let refused = async {
649 Err::<(futures::stream::Empty<Result<Vec<u8>, String>>, _), _>("403".to_string())
650 };
651 let error =
652 stream_to_temp(rt.handle(), Some(dir.path()), None, refused, &unstopped()).unwrap_err();
653 assert!(
654 matches!(&error, StreamError::Open(e) if e == "403"),
655 "{error:?}"
656 );
657 assert_eq!(files_in(dir.path()), 0);
658
659 let stream = futures::stream::iter(vec![Ok::<_, String>(chunk(0, 10))]);
661 let error = stream_to_temp(
662 rt.handle(),
663 Some(dir.path()),
664 None,
665 opened(stream, Some(20)),
666 &unstopped(),
667 )
668 .unwrap_err();
669 assert!(
670 matches!(
671 error,
672 StreamError::Short {
673 expected: 20,
674 got: 10
675 }
676 ),
677 "{error:?}"
678 );
679 assert_eq!(files_in(dir.path()), 0);
680 }
681
682 #[test]
685 fn a_refused_write_stops_the_stream() {
686 let rt = runtime();
687 let pulled = Arc::new(AtomicUsize::new(0));
688 let stream = {
689 let pulled = pulled.clone();
690 futures::stream::iter(0..1000).map(move |i| {
691 pulled.fetch_add(1, Ordering::SeqCst);
692 Ok::<_, String>(chunk(i, 1024))
693 })
694 };
695 let mut writes = 0;
696 let error = stream_into(rt.handle(), opened(stream, None), never(), |_| {
697 writes += 1;
698 if writes == 3 {
699 return Err(eyre!("No space left on device"));
700 }
701 Ok(())
702 })
703 .unwrap_err();
704 assert!(matches!(&error, StreamError::Write(e) if e.to_string().contains("No space")));
705 std::thread::sleep(std::time::Duration::from_millis(200));
707 let pulled = pulled.load(Ordering::SeqCst);
708 assert!(pulled <= 3 + QUEUED_CHUNKS + 2, "{pulled} chunks read");
709
710 let dropped = Arc::new(AtomicBool::new(false));
713 let stream = {
714 let guard = DropFlag(dropped.clone());
715 futures::stream::iter(vec![Ok::<_, String>(chunk(0, 1024))])
716 .chain(futures::stream::pending())
717 .map(move |chunk| {
718 let _ = &guard;
719 chunk
720 })
721 };
722 let error = stream_into(rt.handle(), opened(stream, None), never(), |_| {
723 Err(eyre!("No space left on device"))
724 })
725 .unwrap_err();
726 assert!(matches!(error, StreamError::Write(_)));
727 let deadline = std::time::Instant::now() + std::time::Duration::from_secs(5);
728 while !dropped.load(Ordering::SeqCst) {
729 assert!(
730 std::time::Instant::now() < deadline,
731 "the request is still open"
732 );
733 std::thread::sleep(std::time::Duration::from_millis(10));
734 }
735
736 let missing = tempfile::tempdir().unwrap().path().join("gone");
738 let asked = Arc::new(AtomicBool::new(false));
739 let open = {
740 let asked = asked.clone();
741 async move {
742 asked.store(true, Ordering::SeqCst);
743 Ok((futures::stream::empty::<Result<Vec<u8>, String>>(), None))
744 }
745 };
746 let error =
747 stream_to_temp(rt.handle(), Some(&missing), None, open, &unstopped()).unwrap_err();
748 assert!(matches!(error, StreamError::Write(_)));
749 assert!(!asked.load(Ordering::SeqCst));
750 }
751
752 #[test]
755 fn a_stop_ends_the_download_and_removes_the_file() {
756 let rt = runtime();
757 let dir = tempfile::tempdir().unwrap();
758 let stop = Arc::new(AtomicBool::new(false));
759 let stopped = crate::unfinished::Unfinished::default().writer(stop.clone());
760
761 let stream = {
762 let stop = stop.clone();
763 futures::stream::iter(0..100).map(move |i| {
764 if i == 3 {
765 stop.store(true, Ordering::SeqCst);
766 }
767 Ok::<_, String>(chunk(i, CHUNK))
768 })
769 };
770 let error = stream_to_temp(
771 rt.handle(),
772 Some(dir.path()),
773 None,
774 opened(stream, None),
775 &stopped,
776 )
777 .unwrap_err();
778 assert!(matches!(error, StreamError::Cut), "{error:?}");
779 assert_eq!(files_in(dir.path()), 0);
780
781 stop.store(false, Ordering::SeqCst);
783 let stream = futures::stream::iter(vec![Ok::<_, String>(chunk(0, CHUNK))])
784 .chain(futures::stream::pending());
785 let stopper = {
786 let stop = stop.clone();
787 std::thread::spawn(move || {
788 std::thread::sleep(std::time::Duration::from_millis(100));
789 stop.store(true, Ordering::SeqCst);
790 })
791 };
792 let began = std::time::Instant::now();
793 let error = stream_to_temp(
794 rt.handle(),
795 Some(dir.path()),
796 None,
797 opened(stream, None),
798 &stopped,
799 )
800 .unwrap_err();
801 stopper.join().unwrap();
802 assert!(matches!(error, StreamError::Cut), "{error:?}");
803 assert!(began.elapsed() < std::time::Duration::from_secs(5));
804 assert_eq!(files_in(dir.path()), 0);
805 }
806
807 #[test]
811 fn a_shutdown_mid_transfer_cuts_it_off() {
812 let rt = runtime();
813 let dir = tempfile::tempdir().unwrap();
814 let stream = futures::stream::iter(vec![Ok::<_, String>(chunk(0, CHUNK))])
815 .chain(futures::stream::pending());
816 let handle = rt.handle().clone();
817 let path = dir.path().to_path_buf();
818 let waiter = std::thread::spawn(move || {
819 stream_to_temp(
820 &handle,
821 Some(&path),
822 None,
823 opened(stream, None),
824 &unstopped(),
825 )
826 });
827 let deadline = std::time::Instant::now() + std::time::Duration::from_secs(30);
831 while !std::fs::read_dir(dir.path())
832 .unwrap()
833 .any(|f| std::fs::metadata(f.unwrap().path()).map_or(0, |m| m.len()) == CHUNK as u64)
834 {
835 assert!(
836 std::time::Instant::now() < deadline,
837 "the first chunk never landed"
838 );
839 std::thread::sleep(std::time::Duration::from_millis(5));
840 }
841 rt.shutdown_background();
842 let error = waiter.join().expect("no panic").unwrap_err();
843 assert!(matches!(error, StreamError::Cut), "{error:?}");
844 assert_eq!(files_in(dir.path()), 0);
845 }
846}
847
848#[cfg(all(test, feature = "http"))]
849mod read_tests {
850 use super::*;
851 use std::io::Read;
852 use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering};
853 use std::time::{Duration, Instant};
854
855 struct Source {
858 chunks: std::vec::IntoIter<Vec<u8>>,
859 release: Option<std::sync::mpsc::Receiver<()>>,
860 reads: Arc<AtomicUsize>,
861 dropped: Arc<AtomicBool>,
862 }
863
864 impl Read for Source {
865 fn read(&mut self, buf: &mut [u8]) -> std::io::Result<usize> {
866 self.reads.fetch_add(1, Ordering::SeqCst);
867 if let Some(chunk) = self.chunks.next() {
868 buf[..chunk.len()].copy_from_slice(&chunk);
869 return Ok(chunk.len());
870 }
871 if let Some(release) = self.release.take() {
872 let _ = release.recv();
873 }
874 Ok(0)
875 }
876 }
877
878 impl Drop for Source {
879 fn drop(&mut self) {
880 self.dropped.store(true, Ordering::SeqCst);
881 }
882 }
883
884 fn source(chunks: Vec<Vec<u8>>, release: Option<std::sync::mpsc::Receiver<()>>) -> Source {
885 Source {
886 chunks: chunks.into_iter(),
887 release,
888 reads: Arc::default(),
889 dropped: Arc::default(),
890 }
891 }
892
893 fn files_in(dir: &Path) -> usize {
894 std::fs::read_dir(dir).unwrap().count()
895 }
896
897 #[track_caller]
898 fn wait_until(what: &str, done: impl Fn() -> bool) {
899 let deadline = Instant::now() + Duration::from_secs(5);
900 while !done() {
901 assert!(Instant::now() < deadline, "{what}");
902 std::thread::sleep(Duration::from_millis(5));
903 }
904 }
905
906 #[test]
908 fn a_reader_lands_whole_in_order() {
909 let dir = tempfile::tempdir().unwrap();
910 let chunks = (0..40u8)
911 .map(|i| vec![i; 1000 + usize::from(i) * 13])
912 .collect::<Vec<_>>();
913 let whole = chunks.concat();
914 let len = whole.len() as u64;
915 let reader = source(chunks.clone(), None);
916 let file = read_to_temp(
917 Some(dir.path()),
918 Some("csv"),
919 move || Ok((reader, Some(len))),
920 &Writer::default(),
921 None,
922 )
923 .unwrap();
924 assert_eq!(std::fs::read(file.path()).unwrap(), whole);
925 assert!(file.path().to_string_lossy().ends_with(".csv"));
926
927 let reader = source(chunks, None);
928 let error = read_to_temp(
929 Some(dir.path()),
930 None,
931 move || Ok((reader, Some(len + 1))),
932 &Writer::default(),
933 None,
934 )
935 .unwrap_err();
936 assert!(matches!(error, StreamError::Short { .. }), "{error:?}");
937 drop(file);
938 assert_eq!(files_in(dir.path()), 0);
939 }
940
941 #[test]
944 fn a_download_past_its_limit_stops_and_leaves_no_file() {
945 let dir = tempfile::tempdir().unwrap();
946 let chunks = (0..8u8).map(|i| vec![i; 1000]).collect::<Vec<_>>();
947 let reader = source(chunks.clone(), None);
948 let error = read_to_temp(
949 Some(dir.path()),
950 None,
951 move || Ok((reader, None)),
952 &Writer::default(),
953 Some(7_999),
954 )
955 .unwrap_err();
956 assert!(
957 matches!(&error, StreamError::Write(report) if report.downcast_ref::<PastLimit>().is_some()),
958 "{error:?}"
959 );
960 assert_eq!(files_in(dir.path()), 0);
961
962 let reader = source(chunks, None);
963 let file = read_to_temp(
964 Some(dir.path()),
965 None,
966 move || Ok((reader, None)),
967 &Writer::default(),
968 Some(8_000),
969 )
970 .unwrap();
971 assert_eq!(std::fs::metadata(file.path()).unwrap().len(), 8_000);
972 }
973
974 #[test]
976 fn a_refusal_or_failed_read_leaves_no_file() {
977 struct Reset(bool);
978 impl Read for Reset {
979 fn read(&mut self, buf: &mut [u8]) -> std::io::Result<usize> {
980 if std::mem::replace(&mut self.0, true) {
981 return Err(std::io::ErrorKind::ConnectionReset.into());
982 }
983 buf[..3].copy_from_slice(b"a,b");
984 Ok(3)
985 }
986 }
987
988 let dir = tempfile::tempdir().unwrap();
989 let error = read_to_temp(
990 Some(dir.path()),
991 None,
992 || Err::<(Reset, _), _>("Server returned 404 Not Found.".to_string()),
993 &Writer::default(),
994 None,
995 )
996 .unwrap_err();
997 assert!(
998 matches!(&error, StreamError::Open(e) if e.contains("404")),
999 "{error:?}"
1000 );
1001 let error = read_to_temp(
1002 Some(dir.path()),
1003 None,
1004 || Ok((Reset(false), None)),
1005 &Writer::default(),
1006 None,
1007 )
1008 .unwrap_err();
1009 assert!(matches!(error, StreamError::Read(_)), "{error:?}");
1010 assert_eq!(files_in(dir.path()), 0);
1011 }
1012
1013 #[test]
1015 fn a_slow_writer_holds_the_reads_back() {
1016 let reader = source((0..64).map(|i| vec![i; 1024]).collect(), None);
1017 let reads = reader.reads.clone();
1018 let mut written = 0usize;
1019 let mut ahead = 0usize;
1020 let total = read_into(
1021 move || Ok((reader, None)),
1022 || false,
1023 |_| {
1024 std::thread::sleep(Duration::from_millis(2));
1025 written += 1;
1026 ahead = ahead.max(reads.load(Ordering::SeqCst) - written);
1027 Ok(())
1028 },
1029 )
1030 .unwrap();
1031 assert_eq!(total, 64 * 1024);
1032 assert!(ahead <= QUEUED_CHUNKS + 2, "read {ahead} chunks ahead");
1034 }
1035
1036 #[test]
1039 fn a_stop_while_the_server_is_silent_ends_it() {
1040 let dir = tempfile::tempdir().unwrap();
1041 let (release, held) = std::sync::mpsc::channel();
1042 let reader = source(vec![vec![1; 1024]], Some(held));
1043 let dropped = reader.dropped.clone();
1044 let stop = Arc::new(AtomicBool::new(false));
1045 let stopper = {
1046 let stop = stop.clone();
1047 let dir = dir.path().to_path_buf();
1048 std::thread::spawn(move || {
1049 let deadline = Instant::now() + Duration::from_secs(5);
1053 let landed = loop {
1054 let sizes = std::fs::read_dir(&dir)
1055 .unwrap()
1056 .filter_map(|f| std::fs::metadata(f.unwrap().path()).ok())
1057 .map(|m| m.len());
1058 if sizes.into_iter().any(|len| len == 1024) {
1059 break true;
1060 }
1061 if Instant::now() > deadline {
1062 break false;
1063 }
1064 std::thread::sleep(Duration::from_millis(5));
1065 };
1066 stop.store(true, Ordering::SeqCst);
1069 (landed, Instant::now())
1070 })
1071 };
1072 let error = read_to_temp(
1073 Some(dir.path()),
1074 None,
1075 move || Ok((reader, None)),
1076 &crate::unfinished::Unfinished::default().writer(stop),
1077 None,
1078 )
1079 .unwrap_err();
1080 let (landed, stopped_at) = stopper.join().unwrap();
1081 assert!(landed, "the first chunk landed");
1082 assert!(matches!(error, StreamError::Cut), "{error:?}");
1083 assert!(stopped_at.elapsed() < Duration::from_secs(2));
1084 assert_eq!(files_in(dir.path()), 0);
1085
1086 assert!(
1087 !dropped.load(Ordering::SeqCst),
1088 "still waiting on the server"
1089 );
1090 release.send(()).unwrap();
1091 wait_until("the reader was let go", || dropped.load(Ordering::SeqCst));
1092 }
1093}