use std::path::Path;
use thiserror::Error;
#[derive(Debug, Error)]
pub enum Error {
#[error("satkit was built without the `download` feature")]
FeatureDisabled,
#[error("File {path} not found and satkit was built without the `download` feature")]
FileNotFoundNoDownload { path: String },
#[error(
"{name} is not present and cannot be downloaded ({reason}). \
Provide it in the data directory (SATKIT_DATA) or install the `satkit-data` bundle; \
sources: {}",
if urls.is_empty() { "(none listed)".to_string() } else { urls.join(", ") }
)]
Offline {
name: String,
reason: &'static str,
urls: Vec<String>,
},
#[error("Path or URL has no valid file name: {path}")]
InvalidFileName { path: String },
#[error("Data manifest is invalid: {reason}")]
ManifestInvalid { reason: String },
#[error(
"{name} from {url}: {what} mismatch (expected {expected}, got {actual}); \
the partial download was discarded"
)]
HashMismatch {
name: String,
url: String,
what: &'static str,
expected: String,
actual: String,
},
#[error("Could not download {name} from any source:\n {}", attempts.join("\n "))]
AllSourcesFailed { name: String, attempts: Vec<String> },
#[error("could not replace {path} (is another process using it?): {source}")]
ReplaceFailed {
path: String,
#[source]
source: std::io::Error,
},
#[error(
"{name} at {path} is corrupt: {what} {} does not match the manifest ({}). \
Delete or replace the file, or allow downloads so it can be re-fetched",
.values.1, .values.0
)]
CorruptFile {
name: String,
path: String,
what: &'static str,
values: Box<(String, String)>,
},
#[error(transparent)]
Io(#[from] std::io::Error),
#[error(transparent)]
Json(#[from] serde_json::Error),
#[cfg(feature = "download")]
#[error(transparent)]
Http(#[from] ureq::Error),
}
pub type Result<T> = std::result::Result<T, Error>;
pub const OFFLINE_ENV: &str = "SATKIT_OFFLINE";
static OFFLINE_OVERRIDE: std::sync::atomic::AtomicU8 = std::sync::atomic::AtomicU8::new(0);
pub fn set_offline(offline: bool) {
OFFLINE_OVERRIDE.store(
if offline { 2 } else { 1 },
std::sync::atomic::Ordering::Relaxed,
);
}
pub fn is_offline() -> bool {
match OFFLINE_OVERRIDE.load(std::sync::atomic::Ordering::Relaxed) {
1 => false,
2 => true,
_ => std::env::var(OFFLINE_ENV)
.map(|v| !(v.is_empty() || v == "0" || v.eq_ignore_ascii_case("false")))
.unwrap_or(false),
}
}
pub fn offline_requested() -> bool {
is_offline()
}
#[cfg(test)]
pub(crate) fn clear_offline_override() {
OFFLINE_OVERRIDE.store(0, std::sync::atomic::Ordering::Relaxed);
}
fn manifest_urls(name: &str) -> Vec<String> {
crate::utils::manifest::embedded()
.entry(name)
.map(|e| e.urls.clone())
.unwrap_or_default()
}
pub(crate) fn offline_error(name: &str, reason: &'static str) -> Error {
Error::Offline {
name: name.to_string(),
reason,
urls: manifest_urls(name),
}
}
#[cfg(feature = "download")]
pub(crate) fn check_online(name: &str) -> Result<()> {
if offline_requested() {
return Err(offline_error(name, "SATKIT_OFFLINE is set"));
}
Ok(())
}
static PART_SEQ: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
#[cfg(feature = "download")]
fn part_path(final_path: &Path) -> std::path::PathBuf {
let seq = PART_SEQ.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
let mut p = final_path.as_os_str().to_owned();
p.push(format!(".part.{}.{seq}", std::process::id()));
std::path::PathBuf::from(p)
}
pub(crate) fn retry_io<F: FnMut() -> std::io::Result<()>>(
attempts: u32,
delay: std::time::Duration,
mut op: F,
) -> std::io::Result<()> {
let mut last = None;
for i in 0..attempts.max(1) {
match op() {
Ok(()) => return Ok(()),
Err(e) if e.kind() == std::io::ErrorKind::NotFound => return Err(e),
Err(e) => {
last = Some(e);
if i + 1 < attempts {
std::thread::sleep(delay);
}
}
}
}
Err(last.unwrap_or_else(|| std::io::Error::other("retry_io: no attempts")))
}
#[cfg(feature = "download")]
fn rename_into_place(part: &Path, final_path: &Path) -> Result<()> {
let r = retry_io(6, std::time::Duration::from_millis(50), || {
std::fs::rename(part, final_path)
});
if let Err(source) = r {
let _ = std::fs::remove_file(part);
return Err(Error::ReplaceFailed {
path: final_path.display().to_string(),
source,
});
}
Ok(())
}
#[cfg(feature = "download")]
fn write_atomic(reader: &mut impl std::io::Read, final_path: &Path) -> Result<()> {
let part = part_path(final_path);
let mut write = || -> Result<()> {
let mut dest = std::fs::File::create(&part)?;
std::io::copy(reader, &mut dest)?;
dest.sync_all()?;
Ok(())
};
match write() {
Ok(()) => rename_into_place(&part, final_path),
Err(e) => {
let _ = std::fs::remove_file(&part);
Err(e)
}
}
}
#[cfg(feature = "download")]
pub(crate) fn write_atomic_verified(
reader: &mut impl std::io::Read,
final_path: &Path,
entry: &crate::utils::manifest::ManifestEntry,
url: &str,
) -> Result<()> {
use sha2::{Digest, Sha256};
let part = part_path(final_path);
let mut write = || -> Result<(u64, String)> {
let mut dest = std::fs::File::create(&part)?;
let mut hasher = Sha256::new();
let mut buf = vec![0u8; 1 << 16];
let mut total: u64 = 0;
loop {
let n = reader.read(&mut buf)?;
if n == 0 {
break;
}
hasher.update(&buf[..n]);
std::io::Write::write_all(&mut dest, &buf[..n])?;
total += n as u64;
}
dest.sync_all()?;
let digest = hasher.finalize();
Ok((total, digest.iter().map(|b| format!("{b:02x}")).collect()))
};
let outcome = write();
let (size, sha) = match outcome {
Ok(v) => v,
Err(e) => {
let _ = std::fs::remove_file(&part);
return Err(e);
}
};
let mismatch = |what: &'static str, expected: String, actual: String| {
let _ = std::fs::remove_file(&part);
Error::HashMismatch {
name: entry.name.clone(),
url: url.to_string(),
what,
expected,
actual,
}
};
if size != entry.size {
return Err(mismatch("size", entry.size.to_string(), size.to_string()));
}
if sha != entry.sha256 {
return Err(mismatch("sha256", entry.sha256.clone(), sha));
}
if final_path.is_file() && entry.verify(final_path).unwrap_or(false) {
let _ = std::fs::remove_file(&part);
let _ = entry.write_verified_marker(final_path);
return Ok(());
}
rename_into_place(&part, final_path)?;
let _ = entry.write_verified_marker(final_path);
Ok(())
}
#[cfg(feature = "download")]
pub fn download_if_not_exist(fname: &Path, seturl: Option<&str>) -> Result<()> {
if fname.is_file() {
return Ok(());
}
let basename =
fname
.file_name()
.and_then(|f| f.to_str())
.ok_or_else(|| Error::InvalidFileName {
path: fname.display().to_string(),
})?;
check_online(basename)?;
if seturl.is_none() {
if let Some(entry) = crate::utils::manifest::embedded().entry(basename) {
let dir = fname.parent().unwrap_or_else(|| Path::new("."));
crate::utils::manifest::fetch_static_file(entry, dir, false)?;
return Ok(());
}
eprintln!(
"Warning: {basename} is not in satkit's data manifest; downloading it unverified"
);
}
let baseurl = seturl.unwrap_or("https://storage.googleapis.com/astrokit-astro-data/");
let url = format!("{}{}", baseurl, basename);
let agent = ureq::Agent::new_with_defaults();
let mut resp = agent.get(url.as_str()).call()?;
write_atomic(&mut resp.body_mut().as_reader(), fname)?;
Ok(())
}
#[cfg(not(feature = "download"))]
pub fn download_if_not_exist(fname: &Path, _seturl: Option<&str>) -> Result<()> {
if fname.is_file() {
return Ok(());
}
let name = fname
.file_name()
.and_then(|f| f.to_str())
.unwrap_or("<unnamed>");
Err(offline_error(
name,
"satkit was built without the `download` feature",
))
}
#[cfg(feature = "download")]
pub fn download_file(url: &str, downloaddir: &Path, overwrite_if_exists: bool) -> Result<bool> {
let fname = std::path::Path::new(url)
.file_name()
.and_then(|f| f.to_str())
.ok_or_else(|| Error::InvalidFileName {
path: url.to_string(),
})?;
let fullpath = downloaddir.join(fname);
if fullpath.exists() && !overwrite_if_exists {
println!("File {} exists; skipping download", fname);
return Ok(false);
}
check_online(fname)?;
let agent = ureq::Agent::new_with_defaults();
let mut resp = agent.get(url).call()?;
println!("Downloading {}", fname);
write_atomic(&mut resp.body_mut().as_reader(), &fullpath)?;
Ok(true)
}
#[cfg(not(feature = "download"))]
pub fn download_file(_url: &str, _downloaddir: &Path, _overwrite_if_exists: bool) -> Result<bool> {
Err(Error::FeatureDisabled)
}
#[cfg(feature = "download")]
pub fn download_file_async(
url: String,
downloaddir: &Path,
overwrite_if_exists: bool,
) -> std::thread::JoinHandle<Result<bool>> {
let dclone = downloaddir.to_path_buf();
let urlclone = url;
let overwriteclone = overwrite_if_exists;
std::thread::spawn(move || download_file(urlclone.as_str(), &dclone, overwriteclone))
}
#[cfg(not(feature = "download"))]
pub fn download_file_async(
_url: String,
_downloaddir: &Path,
_overwrite_if_exists: bool,
) -> std::thread::JoinHandle<Result<bool>> {
std::thread::spawn(|| Err(Error::FeatureDisabled))
}
#[cfg(feature = "download")]
pub fn download_to_string(url: &str) -> Result<String> {
check_online(url)?;
let agent = ureq::Agent::new_with_defaults();
let mut resp = agent.get(url).call()?;
let thestring = std::io::read_to_string(resp.body_mut().as_reader())?;
Ok(thestring)
}
#[cfg(not(feature = "download"))]
pub fn download_to_string(_url: &str) -> Result<String> {
Err(Error::FeatureDisabled)
}
#[cfg(all(test, feature = "download"))]
mod tests {
use super::*;
use std::io::Cursor;
#[test]
fn write_atomic_renames_and_leaves_no_part_file() {
let dir = std::env::temp_dir();
let final_path = dir.join("satkit_write_atomic_test.bin");
let part_path = dir.join("satkit_write_atomic_test.bin.part");
let _ = std::fs::remove_file(&final_path);
let _ = std::fs::remove_file(&part_path);
let data = b"hello satkit";
write_atomic(&mut Cursor::new(&data[..]), &final_path).unwrap();
assert!(
final_path.is_file(),
"final file should exist after success"
);
assert!(!part_path.exists(), "the .part file must be renamed away");
assert!(
!leftover_parts(&final_path),
"no .part.<pid>.<seq> file may remain"
);
assert_eq!(std::fs::read(&final_path).unwrap(), data);
let _ = std::fs::remove_file(&final_path);
}
fn leftover_parts(final_path: &Path) -> bool {
let dir = final_path.parent().unwrap();
let prefix = format!(
"{}.part.",
final_path.file_name().unwrap().to_string_lossy()
);
std::fs::read_dir(dir)
.unwrap()
.flatten()
.any(|e| e.file_name().to_string_lossy().starts_with(&prefix))
}
#[test]
fn part_paths_are_unique_per_call() {
let f = Path::new("/tmp/x.bin");
let a = part_path(f);
let b = part_path(f);
assert_ne!(a, b);
assert!(a
.to_string_lossy()
.contains(&format!(".part.{}.", std::process::id())));
}
#[test]
fn retry_io_retries_transient_failures_then_gives_up() {
let mut calls = 0;
let r = retry_io(5, std::time::Duration::from_millis(1), || {
calls += 1;
if calls < 3 {
Err(std::io::Error::new(
std::io::ErrorKind::PermissionDenied,
"busy",
))
} else {
Ok(())
}
});
assert!(r.is_ok());
assert_eq!(calls, 3, "succeeds on the third attempt");
let mut calls = 0;
let r = retry_io(4, std::time::Duration::from_millis(1), || {
calls += 1;
Err(std::io::Error::new(
std::io::ErrorKind::PermissionDenied,
"busy",
))
});
assert_eq!(r.unwrap_err().kind(), std::io::ErrorKind::PermissionDenied);
assert_eq!(calls, 4, "exhausts every attempt");
let mut calls = 0;
let r = retry_io(4, std::time::Duration::from_millis(1), || {
calls += 1;
Err(std::io::Error::new(std::io::ErrorKind::NotFound, "gone"))
});
assert_eq!(r.unwrap_err().kind(), std::io::ErrorKind::NotFound);
assert_eq!(calls, 1, "a missing source is not retried");
}
#[test]
fn rename_into_place_reports_typed_error() {
let dir = std::env::temp_dir().join(format!("satkit_rename_{}", std::process::id()));
let _ = std::fs::remove_dir_all(&dir);
std::fs::create_dir_all(&dir).unwrap();
let part = dir.join("f.bin.part.1.1");
std::fs::write(&part, b"x").unwrap();
let dest = dir.join("missing-subdir").join("f.bin");
let err = rename_into_place(&part, &dest).unwrap_err();
assert!(matches!(err, Error::ReplaceFailed { .. }), "{err}");
assert!(err.to_string().contains("f.bin"));
assert!(!part.exists(), "temporary file is cleaned up");
let _ = std::fs::remove_dir_all(&dir);
}
#[test]
fn write_atomic_verified_rejects_bad_bytes() {
use crate::utils::manifest::{sha256_hex, ManifestEntry};
let dir = std::env::temp_dir().join(format!("satkit_wav_{}", std::process::id()));
std::fs::create_dir_all(&dir).unwrap();
let final_path = dir.join("f.bin");
let entry = ManifestEntry {
name: "f.bin".into(),
size: 5,
sha256: sha256_hex(b"hello"),
urls: vec!["https://example.invalid/f.bin".into()],
source: String::new(),
license: String::new(),
tier: String::new(),
default: true,
};
let err =
write_atomic_verified(&mut Cursor::new(&b"hellp"[..]), &final_path, &entry, "test")
.unwrap_err();
assert!(
matches!(err, Error::HashMismatch { what: "sha256", .. }),
"{err}"
);
assert!(!final_path.exists() && !part_path(&final_path).exists());
let err = write_atomic_verified(
&mut Cursor::new(&b"hello!"[..]),
&final_path,
&entry,
"test",
)
.unwrap_err();
assert!(
matches!(err, Error::HashMismatch { what: "size", .. }),
"{err}"
);
write_atomic_verified(&mut Cursor::new(&b"hello"[..]), &final_path, &entry, "test")
.unwrap();
assert_eq!(std::fs::read(&final_path).unwrap(), b"hello");
assert!(!part_path(&final_path).exists());
let _ = std::fs::remove_dir_all(&dir);
}
}