use crate::models::verify::{SignatureError, verify_minisign};
use sha2::{Digest, Sha256};
use std::fs;
use std::io::{self, BufReader, Read, Write};
use std::path::{Path, PathBuf};
#[derive(Debug, thiserror::Error)]
pub enum DownloadError {
#[error("io error on {path}: {source}")]
Io {
path: PathBuf,
#[source]
source: io::Error,
},
#[error("network error fetching {url}: {source}")]
Network {
url: String,
#[source]
source: Box<ureq::Error>,
},
#[error("checksum mismatch for {path}: expected {expected:.16}…, computed {actual:.16}…")]
ChecksumMismatch {
path: PathBuf,
expected: String,
actual: String,
},
#[error("signature invalid for {path}: {source}")]
SignatureInvalid {
path: PathBuf,
#[source]
source: SignatureError,
},
#[error("refusing to fetch model over a non-https URL: {url}")]
InsecureScheme { url: String },
#[error("download for {path} exceeded the {max_bytes}-byte cap")]
TooLarge { path: PathBuf, max_bytes: u64 },
}
pub fn download_with_checksum(
url: &str,
expected_sha256: &str,
dest: &Path,
) -> Result<bool, DownloadError> {
download_with_checksum_and_signature(url, expected_sha256, None, dest)
}
pub fn download_with_checksum_and_signature(
url: &str,
expected_sha256: &str,
signature: Option<&str>,
dest: &Path,
) -> Result<bool, DownloadError> {
download_with_checksum_signature_and_cap(
url,
expected_sha256,
signature,
dest,
DEFAULT_MAX_MODEL_BYTES,
)
}
pub(crate) const DEFAULT_MAX_MODEL_BYTES: u64 = 1024 * 1024 * 1024;
pub(crate) fn max_download_bytes(declared_size: Option<u64>) -> u64 {
match declared_size {
Some(n) if n > 0 => n.saturating_mul(2).min(DEFAULT_MAX_MODEL_BYTES),
_ => DEFAULT_MAX_MODEL_BYTES,
}
}
pub(crate) fn download_with_checksum_signature_and_cap(
url: &str,
expected_sha256: &str,
signature: Option<&str>,
dest: &Path,
max_bytes: u64,
) -> Result<bool, DownloadError> {
if serve_cache_hit(dest, expected_sha256, signature)? {
return Ok(false);
}
require_https(url)?;
let tmp = prepare_partial_path(dest)?;
let prepared = signature
.map(|sig_text| PreparedSignature::new(dest, sig_text))
.transpose()?;
let mut verifier = DownloadVerifier::new(dest, prepared.as_ref())?;
fetch_into_partial(url, &tmp, max_bytes, &mut verifier)?;
verifier.finish(&tmp, dest, expected_sha256)?;
fs::rename(&tmp, dest).map_err(|e| DownloadError::Io {
path: tmp.clone(),
source: e,
})?;
Ok(true)
}
fn serve_cache_hit(
dest: &Path,
expected_sha256: &str,
signature: Option<&str>,
) -> Result<bool, DownloadError> {
if !(dest.exists() && verify_sha256(dest, expected_sha256).is_ok()) {
return Ok(false);
}
if let Some(sig) = signature {
verify_minisign(dest, sig).map_err(|e| DownloadError::SignatureInvalid {
path: dest.to_path_buf(),
source: e,
})?;
}
Ok(true)
}
fn require_https(url: &str) -> Result<(), DownloadError> {
if !url
.get(..8)
.is_some_and(|s| s.eq_ignore_ascii_case("https://"))
{
return Err(DownloadError::InsecureScheme {
url: url.to_owned(),
});
}
Ok(())
}
fn prepare_partial_path(dest: &Path) -> Result<PathBuf, DownloadError> {
if let Some(parent) = dest.parent() {
fs::create_dir_all(parent).map_err(|e| DownloadError::Io {
path: parent.to_path_buf(),
source: e,
})?;
}
let mut tmp = dest.to_path_buf();
let original_name = dest.file_name().and_then(|s| s.to_str()).unwrap_or("model");
tmp.set_file_name(format!(".{original_name}.partial"));
Ok(tmp)
}
struct PreparedSignature {
public_key: minisign_verify::PublicKey,
signature: minisign_verify::Signature,
}
impl PreparedSignature {
fn new(dest: &Path, sig_text: &str) -> Result<Self, DownloadError> {
let public_key =
minisign_verify::PublicKey::from_base64(crate::models::verify::SIGNING_PUBKEY_BASE64)
.map_err(|e| DownloadError::SignatureInvalid {
path: dest.to_path_buf(),
source: SignatureError::BadPublicKey(format!("{e:?}")),
})?;
let signature = minisign_verify::Signature::decode(sig_text).map_err(|e| {
DownloadError::SignatureInvalid {
path: dest.to_path_buf(),
source: SignatureError::BadSignature(format!("{e:?}")),
}
})?;
Ok(Self {
public_key,
signature,
})
}
fn stream_verifier(
&self,
dest: &Path,
) -> Result<minisign_verify::StreamVerifier<'_>, DownloadError> {
self.public_key.verify_stream(&self.signature).map_err(|e| {
DownloadError::SignatureInvalid {
path: dest.to_path_buf(),
source: SignatureError::VerificationFailed(format!("{e:?}")),
}
})
}
}
struct DownloadVerifier<'a> {
hasher: Sha256,
minisign: Option<minisign_verify::StreamVerifier<'a>>,
}
impl<'a> DownloadVerifier<'a> {
fn new(dest: &Path, prepared: Option<&'a PreparedSignature>) -> Result<Self, DownloadError> {
let minisign = match prepared {
Some(p) => Some(p.stream_verifier(dest)?),
None => None,
};
Ok(Self {
hasher: Sha256::new(),
minisign,
})
}
fn update(&mut self, chunk: &[u8]) {
self.hasher.update(chunk);
if let Some(v) = self.minisign.as_mut() {
v.update(chunk);
}
}
fn finish(self, tmp: &Path, dest: &Path, expected_sha256: &str) -> Result<(), DownloadError> {
let actual = format!("{:x}", self.hasher.finalize());
if actual != expected_sha256 {
let _ = fs::remove_file(tmp);
return Err(DownloadError::ChecksumMismatch {
path: dest.to_path_buf(),
expected: expected_sha256.to_owned(),
actual,
});
}
if let Some(mut v) = self.minisign {
v.finalize().map_err(|e| {
let _ = fs::remove_file(tmp);
DownloadError::SignatureInvalid {
path: dest.to_path_buf(),
source: SignatureError::VerificationFailed(format!("{e:?}")),
}
})?;
}
Ok(())
}
}
fn fetch_into_partial(
url: &str,
tmp: &Path,
max_bytes: u64,
verifier: &mut DownloadVerifier<'_>,
) -> Result<(), DownloadError> {
let resp = ureq::get(url).call().map_err(|e| DownloadError::Network {
url: url.to_owned(),
source: Box::new(e),
})?;
let reader = BufReader::new(resp.into_body().into_reader());
let mut file = fs::File::create(tmp).map_err(|e| DownloadError::Io {
path: tmp.to_path_buf(),
source: e,
})?;
write_capped(reader, &mut file, tmp, max_bytes, &mut |chunk| {
verifier.update(chunk);
})?;
file.flush().map_err(|e| DownloadError::Io {
path: tmp.to_path_buf(),
source: e,
})?;
Ok(())
}
fn write_capped<R: Read>(
mut reader: R,
file: &mut fs::File,
tmp: &Path,
max_bytes: u64,
on_chunk: &mut dyn FnMut(&[u8]),
) -> Result<(), DownloadError> {
let mut buf = [0u8; 64 * 1024];
let mut written: u64 = 0;
loop {
let n = reader.read(&mut buf).map_err(|e| DownloadError::Io {
path: tmp.to_path_buf(),
source: e,
})?;
if n == 0 {
break;
}
written += n as u64;
if written > max_bytes {
let _ = fs::remove_file(tmp);
return Err(DownloadError::TooLarge {
path: tmp.to_path_buf(),
max_bytes,
});
}
file.write_all(&buf[..n]).map_err(|e| DownloadError::Io {
path: tmp.to_path_buf(),
source: e,
})?;
on_chunk(&buf[..n]);
}
Ok(())
}
pub fn verify_sha256(path: &Path, expected: &str) -> Result<(), DownloadError> {
let f = fs::File::open(path).map_err(|e| DownloadError::Io {
path: path.to_path_buf(),
source: e,
})?;
let mut reader = BufReader::new(f);
let mut hasher = Sha256::new();
let mut buf = [0u8; 64 * 1024];
loop {
let n = reader.read(&mut buf).map_err(|e| DownloadError::Io {
path: path.to_path_buf(),
source: e,
})?;
if n == 0 {
break;
}
hasher.update(&buf[..n]);
}
let actual = format!("{:x}", hasher.finalize());
if actual == expected {
Ok(())
} else {
Err(DownloadError::ChecksumMismatch {
path: path.to_path_buf(),
expected: expected.to_owned(),
actual,
})
}
}
#[allow(clippy::unwrap_used)]
#[cfg(test)]
#[path = "download_tests.rs"]
mod tests;