use std::{
error, fs::{File, OpenOptions}, io::{Seek, Write}, path::{Path, PathBuf}, sync::Arc
};
use crate::{bucket::{Bucket, BucketProgress, BucketProgressStream}, models::DownloadStatus};
use reqwest::{self, Client, header};
use tokio::{spawn, sync::{oneshot, watch}};
use futures_util::StreamExt;
#[derive(Debug, Default)]
pub struct DownloadClient {
buckets: Option<Vec<Bucket>>,
url: String,
file_path: String,
error_msg: Option<String>,
cancelled: bool
}
impl DownloadClient {
pub fn init(url: &String, file_path: &String) -> Self {
return Self {
buckets: None,
url: url.clone(),
file_path: file_path.clone(),
error_msg: None,
cancelled: false
};
}
pub async fn begin_download(&mut self) -> Result<(), ()>{
if try_create_file(&self.file_path) == false {
let err_msg = format!("Failed to create file at path {}", self.file_path);
self.error_msg = err_msg.into();
return Err(());
}
match start_download(&self.url, &self.file_path).await {
Ok(b) => {
self.buckets = b.into();
},
Err(e) => {
let err_msg = format!("Failed to start download. {}", e);
self.error_msg = err_msg.into();
return Err(());
},
}
return match self.status() {
DownloadStatus::Failed(_) => Err(()),
_ => Ok(())
};
}
pub fn progress_stream(&self) -> BucketProgressStream {
return match self.buckets.as_ref() {
Some(buckets) => BucketProgressStream::new(buckets),
None => BucketProgressStream::empty(),
};
}
pub fn current_progress(&self) -> impl Iterator<Item = BucketProgress> {
return self.buckets.as_ref().into_iter().flat_map(|buckets| buckets.iter().map(|b| b.bucket_progress()));
}
pub fn bucket_sizes(&mut self) -> Vec<u64> {
let mut sizes: Vec<u64> = vec![];
let buckets = self.buckets.as_mut().unwrap();
for b in buckets {
sizes.push(b.size());
}
return sizes;
}
pub fn status(&mut self) -> DownloadStatus {
if self.error_msg.is_some() {
if self.cancelled == false {
self.cancel();
}
return DownloadStatus::Failed(self.error_msg.as_ref().unwrap().clone());
}
if self.cancelled {
return DownloadStatus::Cancelled;
}
match self.buckets.as_ref() {
Some(buckets) => {
for bucket in buckets {
if !bucket.finished() { return DownloadStatus::InProgress; }
}
return DownloadStatus::Finished;
},
None => return DownloadStatus::NotStarted,
};
}
pub fn cancel(&mut self) {
if let Some(buckets) = self.buckets.as_mut() {
for bucket in buckets {
bucket.cancel();
}
self.delete_unfinished_file();
}
self.cancelled = true;
}
}
impl DownloadClient {
fn delete_unfinished_file(&self) {
println!("Deleting unfinished file...");
let _ = std::fs::remove_file(self.file_path.clone());
println!("Deleted unfinished file.");
}
}
async fn start_download(url: &String, file_path: &String) -> Result<Vec<Bucket>, Box<dyn error::Error>> {
let mut buckets: Vec<Bucket> = vec![];
let client = Arc::new(Client::new());
let head_response = client.head(url).send().await?;
let headers = head_response.headers();
let final_url = head_response.url();
let mut file_name = final_url
.path_segments()
.and_then(|segments| segments.last())
.filter(|s| !s.is_empty())
.unwrap_or("download.dat")
.to_string();
if !file_name.contains('.') { file_name = format!("{}.dat", file_name); }
let directory = dirs::download_dir().unwrap_or_else(|| std::env::current_dir().unwrap_or_default());
let full_file_path = directory.join(file_name);
let content_length: usize = headers.get(header::CONTENT_LENGTH).unwrap().to_str()?.parse::<usize>()?;
let standard_bucket_size: usize = get_standard_bucket_size(content_length, headers.contains_key(header::ACCEPT_RANGES));
let mut bucket_id: u8 = 0;
for start_byte in (0..content_length).step_by(standard_bucket_size) {
let bucket = start_bucket_download(bucket_id, start_byte, standard_bucket_size, content_length, url, &full_file_path, &client).await;
buckets.push(bucket);
bucket_id += 1;
}
return Ok(buckets);
}
async fn start_bucket_download(id: u8, start_byte: usize, standard_bucket_size: usize, content_length: usize, url: &String, file_path: &PathBuf, client: &Arc<Client>) -> Bucket {
let end_byte = (start_byte + standard_bucket_size).min(content_length);
let bucket_size = end_byte - start_byte;
let (w_tx, w_rx) = watch::channel::<u64>(0);
let (ks_tx, ks_rx) = oneshot::channel::<bool>();
spawn(download_range(Arc::clone(client), start_byte, end_byte - 1, url.clone(), file_path.clone(), w_tx, ks_rx));
return Bucket::new(
id,
bucket_size as u64,
w_rx,
ks_tx
);
}
fn try_create_file(file_path: &String) -> bool {
return match File::create(file_path) {
Ok(_) => true,
Err(_) => false,
};
}
fn get_standard_bucket_size(content_length: usize, accepts_ranges: bool) -> usize {
if accepts_ranges == false { return content_length; }
return (content_length / 6) + 1;
}
async fn download_range(client: Arc<Client>, start_byte: usize, end_byte: usize, url: String, file_path: PathBuf, sender: watch::Sender<u64>, mut kill_switch: oneshot::Receiver<bool>) -> Result<(), ()> {
let range = format!("bytes={}-{}", start_byte, end_byte);
if let Ok(response) = client.get(url).header("Range", range).send().await {
let mut file = OpenOptions::new().write(true).open(&file_path).unwrap();
let _ = file.seek(std::io::SeekFrom::Start(start_byte as u64)).unwrap();
let mut stream = response.bytes_stream();
let mut download_offset = 0;
while let Some(item) = stream.next().await {
if let Ok(kill) = kill_switch.try_recv() {
if kill { break; }
}
let bytes = item.unwrap();
let _ = file.write(&bytes).unwrap();
download_offset += bytes.len() as u64;
match sender.send(download_offset) {
Ok(_) => (),
Err(e) => {
println!("error {:?}", e.0);
},
}
}
file.flush().unwrap();
return Ok(());
}
return Err(());
}