use crate::{
io::{api::huggingface, apply_progress_style, create_progress_bar, finish_progress_bar, http, slice_path, ApiResult, ProgressType},
util::constants::{app::DISPLAY_PATH_CHARACTER_COUNT, env::HUGGINGFACE_TOKEN_VARIABLE_NAMES},
};
use ::http::HeaderMap;
use acorn_core::prelude::Vec;
use acorn_host::{
fs::{file_checksum, SafePath},
terminal::Label,
};
use acorn_schema::util::constants::app::DEFAULT_HUGGINGFACE_DOMAIN;
use bon::Builder;
use color_eyre::eyre::eyre;
use fluent_uri::Uri;
use futures::{future::join_all, stream, StreamExt};
use std::{
fs::{create_dir_all, remove_file, rename},
path::{Path, PathBuf},
};
use tracing::info;
#[derive(Builder, Clone, Debug)]
#[builder(start_fn = init, on(String, into))]
pub struct DownloadItem {
pub url: String,
pub path: String,
pub size: Option<u64>,
pub sha: Option<String>,
pub headers: Option<HeaderMap>,
}
#[derive(Clone, Debug)]
pub struct DownloadItems {
pub destination: PathBuf,
pub items: Vec<DownloadItem>,
pub quiet: bool,
pub skip_verify_checksum: bool,
pub concurrency: Option<usize>,
}
#[derive(Clone, Debug)]
pub struct DownloadTask {
pub item: DownloadItem,
pub part: PathBuf,
pub target: PathBuf,
pub quiet: bool,
pub skip_verify_checksum: bool,
}
impl DownloadItem {
const PART_EXTENSION: &'static str = "part";
const DOWNLOAD_ERROR_MESSAGE: &'static str = "Failed to download file";
const HUGGINGFACE_HOST: &'static str = DEFAULT_HUGGINGFACE_DOMAIN;
const HUGGINGFACE_ERROR_MESSAGE: &'static str = "Failed to download Hugging Face file";
const HUGGINGFACE_AUTH_ERROR_MESSAGE: &'static str =
"Hugging Face denied access to this file; verify your token has access to the repository and gated/Xet files";
pub fn local_path(&self) -> PathBuf {
PathBuf::from(&self.path)
}
pub fn target(&self, destination: &Path) -> ApiResult<PathBuf> {
SafePath::new(self.local_path())
.and_then(|path| path.materialize_under(destination))
.map_err(|why| eyre!("Model file path is unsafe: {} — {why}", self.path))
}
pub fn partial_path(&self, destination: &Path) -> ApiResult<PathBuf> {
self.target(destination).map(|target| {
target.with_extension(format!(
"{}{}",
target
.extension()
.and_then(|value| value.to_str())
.map(|value| format!("{value}."))
.unwrap_or_default(),
Self::PART_EXTENSION
))
})
}
pub fn is_huggingface(&self) -> bool {
Uri::parse(self.url.as_str())
.ok()
.and_then(|uri| uri.authority().map(|authority| authority.host().to_string()))
.is_some_and(|host| host == Self::HUGGINGFACE_HOST || host.ends_with(&format!(".{}", Self::HUGGINGFACE_HOST)))
}
pub fn error_message(&self) -> &'static str {
match self.is_huggingface() {
| true => Self::HUGGINGFACE_ERROR_MESSAGE,
| false => Self::DOWNLOAD_ERROR_MESSAGE,
}
}
pub fn auth_error_message(&self) -> Option<String> {
match self.is_huggingface() {
| true if huggingface::has_auth_token() => Some(Self::HUGGINGFACE_AUTH_ERROR_MESSAGE.into()),
| true => Some(format!(
"Model download requires authentication; set {}",
HUGGINGFACE_TOKEN_VARIABLE_NAMES.join(", ")
)),
| false => None,
}
}
fn request_headers_with_huggingface_fallback(&self) -> Option<HeaderMap> {
self.headers
.clone()
.filter(|headers| !headers.is_empty())
.or_else(|| self.is_huggingface().then(huggingface::auth_headers))
}
pub fn is_complete_at(&self, target: &Path) -> bool {
target.exists()
&& self
.size
.is_none_or(|size| target.metadata().map(|metadata| metadata.len() == size).unwrap_or(false))
}
pub fn verify_size(&self, path: &Path) -> ApiResult<()> {
match self.size {
| Some(size) => match path.metadata() {
| Ok(metadata) if metadata.len() == size => Ok(()),
| Ok(metadata) => Err(eyre!(
"Downloaded model size did not match expected size (expected {size}, got {})",
metadata.len()
)),
| Err(why) => Err(eyre!("Failed to inspect downloaded model file {} — {why}", path.display())),
},
| None => Ok(()),
}
}
pub fn verify_checksum(&self, path: &Path, skip_verify_checksum: bool) -> ApiResult<()> {
match (skip_verify_checksum, self.sha.as_ref()) {
| (false, Some(expected)) => match file_checksum(path, None) {
| Ok(actual) if actual.checksum_value.eq_ignore_ascii_case(expected) => Ok(()),
| Ok(actual) => Err(eyre!(
"Model download is incomplete (checksum mismatch for {} (expected {expected}, got {}))",
self.path,
actual.checksum_value
)),
| Err(why) => Err(eyre!(
"Model download is incomplete (failed to compute SHA-256 for {}) — {why}",
self.path
)),
},
| _ => Ok(()),
}
}
}
impl DownloadItems {
/// Creates a new batch of download items.
pub fn new(destination: &Path, items: Vec<DownloadItem>, quiet: bool, skip_verify_checksum: bool) -> Self {
Self {
destination: destination.to_path_buf(),
items,
quiet,
skip_verify_checksum,
concurrency: None,
}
}
/// Downloads all items concurrently (bounded by [`Self::concurrency`] when set) and returns the first error encountered.
pub async fn download(self) -> ApiResult<()> {
let DownloadItems {
destination,
items,
quiet,
skip_verify_checksum,
concurrency,
} = self;
let tasks = items
.into_iter()
.map(|item| DownloadTask::new(item, &destination, quiet, skip_verify_checksum))
.collect::<ApiResult<Vec<_>>>();
match tasks {
| Err(why) => Err(why),
| Ok(tasks) => {
let futures = tasks.into_iter().map(DownloadTask::download).collect::<Vec<_>>();
let results = match concurrency {
| Some(limit) => stream::iter(futures).buffer_unordered(limit.max(1)).collect::<Vec<_>>().await,
| None => join_all(futures).await,
};
results.into_iter().find(|result| result.is_err()).unwrap_or(Ok(()))
}
}
}
}
impl DownloadTask {
/// Creates a new download task with resolved paths.
pub fn new(item: DownloadItem, destination: &Path, quiet: bool, skip_verify_checksum: bool) -> ApiResult<Self> {
item.target(destination).map(|target| Self {
part: target.with_extension(format!(
"{}{}",
target
.extension()
.and_then(|value| value.to_str())
.map(|value| format!("{value}."))
.unwrap_or_default(),
DownloadItem::PART_EXTENSION
)),
target,
item,
quiet,
skip_verify_checksum,
})
}
/// Executes the download, including size and checksum verification.
pub async fn download(self) -> ApiResult<()> {
match self.item.is_complete_at(&self.target) {
| true => {
if !self.quiet {
let target_display = slice_path(&self.target, DISPLAY_PATH_CHARACTER_COUNT);
println!("=> {} Skipping existing model file {target_display}", Label::pass());
}
Ok(())
}
| false => {
let parent_ready = match self.part.parent() {
| Some(parent) => create_dir_all(parent).map_err(|why| eyre!("Failed to create model output directory — {why}")),
| None => Err(eyre!("Model output path has no parent directory")),
};
match parent_ready {
| Ok(_) => match self.resume().await {
| Ok(resume_from) => match self
.clone()
.stream(resume_from)
.await
.and_then(|_| self.item.verify_size(&self.part))
.and_then(|_| self.item.verify_checksum(&self.part, self.skip_verify_checksum))
{
| Ok(_) => rename(&self.part, &self.target)
.map_err(|why| eyre!("Failed to finalize model download {} — {why}", self.target.display())),
| Err(why) => {
if self.part.exists() {
let _ = remove_file(&self.part);
}
Err(why)
}
},
| Err(why) => Err(why),
},
| Err(why) => Err(why),
}
}
}
}
async fn resume(&self) -> ApiResult<Option<u64>> {
match self.part.exists() {
| false => Ok(None),
| true => {
let metadata = self
.part
.metadata()
.map_err(|why| eyre!("Failed to inspect partial download {} — {why}", self.part.display()));
match metadata {
| Ok(data) => {
let part_size = data.len();
match self.item.size {
| Some(expected) if part_size > 0 && part_size < expected => {
let headers = self.item.request_headers_with_huggingface_fallback();
let supports_ranges = http::supports_byte_ranges(&self.item.url, headers).await.unwrap_or(false);
if supports_ranges {
Ok(Some(part_size))
} else {
remove_file(&self.part)
.map_err(|why| eyre!("Failed to remove stale partial download {} — {why}", self.part.display()))
.map(|_| None)
}
}
| _ => remove_file(&self.part)
.map_err(|why| eyre!("Failed to remove stale partial download {} — {why}", self.part.display()))
.map(|_| None),
}
}
| Err(why) => Err(why),
}
}
}
}
async fn stream(self, resume_from: Option<u64>) -> ApiResult<()> {
let Self {
quiet, item, part, target, ..
} = self;
let progress_type = match (quiet, item.size) {
| (true, _) => ProgressType::Silent,
| (false, Some(_)) => ProgressType::Bar,
| (false, None) => ProgressType::Spinner,
};
let target_display = slice_path(&target, DISPLAY_PATH_CHARACTER_COUNT);
if !quiet {
info!("=> {} Downloading {target_display}", Label::run());
}
let progress = create_progress_bar(item.size.and_then(|value| usize::try_from(value).ok()).unwrap_or_default(), progress_type);
match item.size {
| Some(_) => apply_progress_style(&progress, " {spinner:.green} {bytes:>10}/{total_bytes:<10} [{bar:40.green}] {msg}"),
| None => apply_progress_style(&progress, " {spinner:.green} {bytes:>10} {msg}"),
}
progress.set_message(target_display.clone());
let auth_error_message = item.auth_error_message();
let result = http::download_with_progress(
&item.url,
&part,
|downloaded, total| {
if let Some(total) = total {
progress.set_length(total);
}
progress.set_position(downloaded);
},
item.request_headers_with_huggingface_fallback(),
Some(item.error_message()),
auth_error_message.as_deref(),
resume_from,
);
match result.await {
| Ok(_) => {
if !quiet {
finish_progress_bar(&progress, format!("{}Downloaded {target_display}", Label::CHECKMARK));
}
Ok(())
}
| Err(why) => {
progress.finish_and_clear();
Err(why)
}
}
}
}