#![deny(unsafe_code)]
#![deny(missing_docs)]
pub use indicatif::ProgressFinish;
use indicatif::{MultiProgress, ProgressBar, ProgressStyle};
use reqwest::Request;
use tokio::io::AsyncWriteExt;
use tokio::sync::Semaphore;
type Result<T> = std::result::Result<T, Box<dyn std::error::Error + Send + Sync>>;
pub struct RetrieverBuilder {
workers: usize,
show_progress_bar: bool,
pb_style: Option<ProgressStyle>,
pb_finish: ProgressFinish,
}
impl Default for RetrieverBuilder {
fn default() -> Self {
Self {
workers: 10,
show_progress_bar: false,
pb_style: Some(
ProgressStyle::with_template(
"[{elapsed_precise}] [{bar:40.cyan/blue}] {bytes}/{total_bytes} {msg}",
)
.expect("progress bar template should compile")
.progress_chars("=>-"),
),
pb_finish: ProgressFinish::AndLeave,
}
}
}
impl RetrieverBuilder {
pub fn new() -> Self {
Self::default()
}
pub fn show_progress(mut self, show_progress_bar: bool) -> Self {
self.show_progress_bar = show_progress_bar;
self
}
pub fn progress_style(mut self, pb_style: ProgressStyle) -> Self {
self.pb_style = Some(pb_style);
self
}
pub fn with_finish(mut self, pb_finish: ProgressFinish) -> Self {
self.pb_finish = pb_finish;
self
}
pub fn workers(mut self, workers: usize) -> Self {
self.workers = workers;
self
}
pub fn build(self) -> Retriever {
Retriever {
client: reqwest::Client::new(),
job_semaphore: Semaphore::new(self.workers),
mp: if self.show_progress_bar {
Some(MultiProgress::new())
} else {
None
},
pb_style: self.pb_style,
pb_finish: self.pb_finish,
}
}
}
pub struct Retriever {
client: reqwest::Client,
job_semaphore: Semaphore,
mp: Option<MultiProgress>,
pb_style: Option<ProgressStyle>,
pb_finish: ProgressFinish,
}
impl Default for Retriever {
fn default() -> Self {
RetrieverBuilder::new().build()
}
}
impl Retriever {
pub fn with_progress_bar() -> Self {
RetrieverBuilder::new().show_progress(true).build()
}
pub async fn download_file<W>(&self, request: Request, mut writer: W) -> Result<()>
where
W: AsyncWriteExt + Unpin,
{
let _permit = self.job_semaphore.acquire().await?;
let path = String::from(request.url().path());
let mut resp = self.client.execute(request).await?.error_for_status()?;
let mut pb = ProgressBar::hidden();
if let Some(m) = &self.mp {
if let Some(pb_style) = &self.pb_style {
pb = m.add(
ProgressBar::no_length()
.with_style(pb_style.clone())
.with_message(path)
.with_finish(self.pb_finish.clone()),
);
if let Some(total_size) = resp.content_length() {
pb.set_length(total_size);
}
}
}
while let Some(chunk) = resp.chunk().await? {
writer.write_all(chunk.as_ref()).await?;
writer.flush().await?;
pb.inc(chunk.len() as u64);
}
pb.set_length(pb.position());
pb.finish_using_style();
drop(_permit);
Ok(())
}
}
#[cfg(test)]
mod tests {
use super::*;
use http::StatusCode;
use mockito::Matcher;
use reqwest::{Client, Error as ReqwestError};
use tokio::{fs::OpenOptions, io::AsyncReadExt};
#[tokio::test]
async fn download_single() {
let mut server = mockito::Server::new_async().await;
let mock = server
.mock("GET", Matcher::Regex(r"/\d".to_string()))
.with_status(200)
.with_body("hello")
.create();
let retriever = RetrieverBuilder::new()
.show_progress(false)
.workers(1)
.build();
let req = Client::new()
.get(format!("{}/1", server.url()))
.build()
.expect("failed to build request");
let file_path = "/tmp/test";
let file = OpenOptions::new()
.create(true)
.write(true)
.truncate(true)
.open(file_path)
.await
.expect("failed to open file for writing");
let _ = retriever
.download_file(req, file)
.await
.expect("failed to download");
let mut file = OpenOptions::new()
.read(true)
.open(file_path)
.await
.expect("failed to open file for reading");
let mut contents = String::new();
file.read_to_string(&mut contents)
.await
.expect("failed to read file");
assert_eq!(contents, "hello");
mock.assert();
}
#[tokio::test]
async fn download_error_status() {
let mut server = mockito::Server::new_async().await;
let expected_status = StatusCode::NOT_FOUND;
let mock = server
.mock("GET", "/404")
.with_status(expected_status.as_u16().into())
.with_body("not found")
.create();
let retriever = RetrieverBuilder::new()
.show_progress(false)
.workers(1)
.build();
let req = Client::new()
.get(format!("{}/404", server.url()))
.build()
.expect("failed to build request");
let file_path = "/tmp/test_error";
let file = OpenOptions::new()
.create(true)
.write(true)
.truncate(true)
.open(file_path)
.await
.expect("failed to open file for writing");
let err = retriever
.download_file(req, file)
.await
.expect_err("expected error on status");
let reqwest_err = err
.downcast_ref::<ReqwestError>()
.expect("error should be a reqwest::Error");
let status = reqwest_err
.status()
.expect("request error should have a status code");
assert_eq!(status, expected_status);
assert!(reqwest_err.to_string().contains("not found"), "error message should contain response body");
mock.assert();
}
#[tokio::test]
async fn download_multi() {
use std::sync::Arc;
use tokio::task::JoinSet;
let mut server = mockito::Server::new_async().await;
let mock = server
.mock("GET", Matcher::Regex(r"/\d".to_string()))
.with_status(200)
.with_body("hello")
.expect(10)
.create();
let retriever = Arc::new(
RetrieverBuilder::new()
.show_progress(true)
.progress_style(
ProgressStyle::with_template("{bytes}/{total_bytes} {msg}")
.expect("progress bar template should compile"),
)
.with_finish(ProgressFinish::WithMessage("done".into()))
.build(),
);
let mut set = JoinSet::new();
for i in 0..10 {
let ret = Arc::clone(&retriever);
let req = Client::new()
.get(format!("{}/{}", server.url(), i))
.build()
.expect("request should build");
let file = OpenOptions::new()
.create(true)
.write(true)
.truncate(true)
.open(format!("/tmp/test{}", i))
.await
.expect("file should be accessible");
set.spawn(async move { ret.download_file(req, file).await });
}
while let Some(download_result) = set.join_next().await {
assert!(!download_result.is_err());
}
mock.assert();
}
}