use std::fmt;
#[cfg(feature = "http")]
use std::fs::{self, File};
#[cfg(feature = "http")]
use std::io::{self, Read, Write};
use std::path::{Path, PathBuf};
#[cfg(feature = "http")]
use std::time::{Duration, Instant};
use bevy_ecs::message::Message;
use http::header::{HeaderMap, HeaderName};
use http::StatusCode;
use crate::request::RequestId;
use crate::response::BackendError;
pub const DEFAULT_DOWNLOAD_MAX_BYTES: u64 = 256 * 1024 * 1024;
#[cfg(feature = "http")]
const PROGRESS_EVERY: Duration = Duration::from_millis(100);
#[cfg(feature = "http")]
const CHUNK: usize = 64 * 1024;
fn file_name(path: &Path) -> String {
path.file_name().map(|n| n.to_string_lossy().into_owned()).unwrap_or_default()
}
#[cfg(any(feature = "http", feature = "sftp"))]
pub(crate) fn part_path(local: &Path, id: RequestId) -> PathBuf {
static NEXT: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
let mut name = local.file_name().map(std::ffi::OsStr::to_os_string).unwrap_or_default();
name.push(format!(".{}-{}-{}.part", std::process::id(), id.to_string().trim_start_matches('#'), NEXT.fetch_add(1, std::sync::atomic::Ordering::Relaxed)));
local.with_file_name(name)
}
#[cfg(any(feature = "http", feature = "sftp"))]
fn is_stale_part(entry: &str, target: &str, own: u32) -> bool {
let Some(numbers) = entry.strip_prefix(target).and_then(|rest| rest.strip_prefix('.')).and_then(|rest| rest.strip_suffix(".part")) else {
return false;
};
let parts: Vec<&str> = numbers.split('-').collect();
let all_numbers = parts.len() == 3 && parts.iter().all(|p| !p.is_empty() && p.len() <= 20 && p.bytes().all(|b| b.is_ascii_digit()));
all_numbers && parts.first().and_then(|pid| pid.parse::<u32>().ok()).is_some_and(|pid| pid != own)
}
#[cfg(any(feature = "http", feature = "sftp"))]
pub(crate) fn remove_stale_parts(local: &Path) {
let Some(target) = local.file_name().and_then(std::ffi::OsStr::to_str) else { return };
let folder = match local.parent() {
Some(parent) if !parent.as_os_str().is_empty() => parent,
_ => Path::new("."),
};
let Ok(entries) = std::fs::read_dir(folder) else { return };
let own = std::process::id();
for entry in entries.flatten() {
let name = entry.file_name();
if name.to_str().is_some_and(|name| is_stale_part(name, target, own)) && entry.file_type().is_ok_and(|t| t.is_file()) {
let _ = std::fs::remove_file(entry.path());
}
}
}
#[derive(Clone, PartialEq, Eq)]
pub struct HttpDownload {
path: PathBuf,
sha256: Option<String>,
size: Option<u64>,
max_bytes: u64,
progress: bool,
}
impl fmt::Debug for HttpDownload {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("HttpDownload")
.field("file_name", &file_name(&self.path))
.field("sha256", &self.sha256)
.field("size", &self.size)
.field("max_bytes", &self.max_bytes)
.field("progress", &self.progress)
.finish()
}
}
impl HttpDownload {
pub fn to(path: impl Into<PathBuf>) -> Self {
Self { path: path.into(), sha256: None, size: None, max_bytes: DEFAULT_DOWNLOAD_MAX_BYTES, progress: true }
}
pub fn with_sha256(mut self, hex: impl Into<String>) -> Self {
self.sha256 = Some(hex.into().to_ascii_lowercase());
self
}
pub fn with_size(mut self, bytes: u64) -> Self {
self.size = Some(bytes);
self
}
pub fn with_max_bytes(mut self, bytes: u64) -> Self {
self.max_bytes = bytes;
self
}
pub fn with_progress(mut self, report: bool) -> Self {
self.progress = report;
self
}
pub fn path(&self) -> &Path {
&self.path
}
pub fn sha256(&self) -> Option<&str> {
self.sha256.as_deref()
}
pub fn size(&self) -> Option<u64> {
self.size
}
pub fn max_bytes(&self) -> u64 {
self.max_bytes
}
pub fn progress(&self) -> bool {
self.progress
}
pub(crate) fn check(&self) -> Result<(), BackendError> {
let invalid = |why: &str| Err(BackendError::InvalidRequest(format!("download: {why}")));
if self.path.file_name().is_none() {
return invalid("the local path has no file name");
}
if let Some(sha) = &self.sha256 {
if sha.len() != 64 || !sha.bytes().all(|b| b.is_ascii_hexdigit()) {
return invalid("the expected SHA-256 must be 64 hex digits");
}
}
if self.size.is_some_and(|size| size > self.max_bytes) {
return invalid("the expected size is larger than the size limit");
}
Ok(())
}
#[cfg(feature = "http")]
#[cfg_attr(docsrs, doc(cfg(feature = "http")))]
pub fn receive(
&self,
id: RequestId,
body: &mut dyn Read,
content_length: Option<u64>,
progress: &mut dyn FnMut(u64, Option<u64>),
) -> Result<DownloadedFile, BackendError> {
let part = PartFile::create(self, id)?;
part.write_from(body, content_length, progress, &|| false, &|e| BackendError::Network(format!("reading the download failed: {e}")))?.put_in_place()
}
}
#[derive(Clone, PartialEq, Eq)]
#[non_exhaustive]
pub struct DownloadedFile {
pub path: PathBuf,
pub bytes: u64,
pub sha256: String,
pub status: StatusCode,
pub headers: HeaderMap,
}
impl fmt::Debug for DownloadedFile {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
let headers: Vec<&str> = self.headers.keys().map(HeaderName::as_str).collect();
f.debug_struct("DownloadedFile")
.field("file_name", &file_name(&self.path))
.field("bytes", &self.bytes)
.field("sha256", &self.sha256)
.field("status", &self.status)
.field("header_names", &headers)
.finish()
}
}
#[derive(Message, Clone, Debug)]
#[non_exhaustive]
pub struct HttpDownloadResponse {
pub id: RequestId,
pub result: Result<DownloadedFile, BackendError>,
}
#[derive(Message, Clone, Debug, PartialEq, Eq)]
#[non_exhaustive]
pub struct HttpDownloadProgress {
pub id: RequestId,
pub received: u64,
pub total: Option<u64>,
}
#[cfg(feature = "http")]
pub(crate) struct PartFile {
file: Option<File>,
part: PathBuf,
settings: HttpDownload,
done: bool,
}
#[cfg(feature = "http")]
impl Drop for PartFile {
fn drop(&mut self) {
self.file = None;
if !self.done {
let _ = fs::remove_file(&self.part);
}
}
}
#[cfg(feature = "http")]
fn hex(bytes: &[u8]) -> String {
const HEX: &[u8; 16] = b"0123456789abcdef";
let mut out = String::with_capacity(bytes.len() * 2);
for &b in bytes {
out.push(HEX.get(usize::from(b >> 4)).copied().map_or('0', char::from));
out.push(HEX.get(usize::from(b & 0x0f)).copied().map_or('0', char::from));
}
out
}
#[cfg(feature = "http")]
impl PartFile {
pub(crate) fn create(settings: &HttpDownload, id: RequestId) -> Result<Self, BackendError> {
settings.check()?;
if settings.path.is_dir() {
return Err(BackendError::InvalidRequest(format!("download: `{}` is a folder, not a file", file_name(&settings.path))));
}
remove_stale_parts(&settings.path);
let part = part_path(&settings.path, id);
let file = File::options()
.write(true)
.create_new(true)
.open(&part)
.map_err(|e| BackendError::InvalidRequest(format!("download: the part file for `{}` cannot be created: {e}", file_name(&settings.path))))?;
Ok(Self { file: Some(file), part, settings: settings.clone(), done: false })
}
pub(crate) fn write_from(
mut self,
body: &mut dyn Read,
content_length: Option<u64>,
progress: &mut dyn FnMut(u64, Option<u64>),
stop: &dyn Fn() -> bool,
read_error: &dyn Fn(io::Error) -> BackendError,
) -> Result<ReadyFile, BackendError> {
let limit = self.settings.max_bytes;
if let Some(length) = content_length {
if length > limit {
return Err(BackendError::BodyTooLarge { limit });
}
if let Some(size) = self.settings.size.filter(|size| *size != length) {
return Err(BackendError::Network(format!("download: the server announced {length} bytes, the expected size is {size}")));
}
}
let report = self.settings.progress;
let mut hasher = ring::digest::Context::new(&ring::digest::SHA256);
let mut buffer = vec![0u8; CHUNK];
let mut written: u64 = 0;
let mut last: Option<Instant> = None;
let Some(file) = self.file.as_mut() else {
return Err(BackendError::Network("download: the part file is closed".into()));
};
let write_failed = |e: io::Error| BackendError::Network(format!("download: writing the file failed: {e}"));
loop {
if stop() {
return Err(BackendError::Cancelled);
}
let n = match body.read(&mut buffer) {
Ok(0) => break,
Ok(n) => n,
Err(e) if e.kind() == io::ErrorKind::Interrupted => continue,
Err(e) => return Err(read_error(e)),
};
let piece = buffer.get(..n).unwrap_or_default();
written = written.saturating_add(u64::try_from(n).unwrap_or(u64::MAX));
if written > limit {
return Err(BackendError::BodyTooLarge { limit });
}
hasher.update(piece);
file.write_all(piece).map_err(write_failed)?;
if report {
let now = Instant::now();
if last.is_none_or(|at| now.saturating_duration_since(at) >= PROGRESS_EVERY) {
last = Some(now);
progress(written, content_length);
}
}
}
if let Some(length) = content_length.filter(|length| *length != written) {
return Err(BackendError::Network(format!("download: the body ended after {written} of {length} bytes")));
}
if let Some(size) = self.settings.size.filter(|size| *size != written) {
return Err(BackendError::Network(format!("download: the file has {written} bytes, the expected size is {size}")));
}
let sha256 = hex(hasher.finish().as_ref());
if self.settings.sha256.as_deref().is_some_and(|expected| expected != sha256) {
return Err(BackendError::Network("download: the file does not match the expected SHA-256".into()));
}
file.flush().map_err(write_failed)?;
file.sync_all().map_err(write_failed)?;
if stop() {
return Err(BackendError::Cancelled);
}
if report {
progress(written, content_length);
}
self.file = None;
Ok(ReadyFile { part: self, bytes: written, sha256 })
}
}
#[cfg(feature = "http")]
pub(crate) struct ReadyFile {
part: PartFile,
bytes: u64,
sha256: String,
}
#[cfg(feature = "http")]
impl ReadyFile {
pub(crate) fn rename(mut self) -> Result<DownloadedFile, BackendError> {
let target = self.part.settings.path.clone();
fs::rename(&self.part.part, &target)
.map_err(|e| BackendError::Network(format!("download: the file could not be put in place as `{}`: {e}", file_name(&target))))?;
self.part.done = true;
Ok(DownloadedFile { path: target, bytes: self.bytes, sha256: std::mem::take(&mut self.sha256), status: StatusCode::OK, headers: HeaderMap::new() })
}
pub(crate) fn sync_folder(file: &DownloadedFile) {
let folder = match file.path.parent() {
Some(parent) if !parent.as_os_str().is_empty() => parent,
_ => Path::new("."),
};
crate::secret_file::sync_dir(folder);
}
pub(crate) fn put_in_place(self) -> Result<DownloadedFile, BackendError> {
let file = self.rename()?;
Self::sync_folder(&file);
Ok(file)
}
}
#[cfg(all(test, feature = "http"))]
mod tests {
use super::*;
fn dir(name: &str) -> PathBuf {
let dir = std::env::temp_dir().join(format!("bnb-download-unit-{name}-{}", std::process::id()));
let _ = fs::remove_dir_all(&dir);
fs::create_dir_all(&dir).unwrap_or_else(|e| panic!("{e}"));
dir
}
fn leftovers(dir: &Path) -> Vec<String> {
fs::read_dir(dir).map(|it| it.filter_map(Result::ok).map(|e| e.file_name().to_string_lossy().into_owned()).collect()).unwrap_or_default()
}
#[test]
fn checks_and_part_files() {
let dir = dir("checks");
let target = dir.join("file.bin");
let data = vec![7u8; 200_000];
let sha = hex(ring::digest::digest(&ring::digest::SHA256, &data).as_ref());
let id = RequestId::next();
let mut reports = Vec::new();
let ok = HttpDownload::to(&target).with_sha256(sha.to_uppercase()).with_size(200_000).receive(id, &mut &data[..], Some(200_000), &mut |r, t| {
reports.push((r, t));
});
let file = ok.unwrap_or_else(|e| panic!("{e}"));
assert_eq!((file.bytes, file.sha256.as_str()), (200_000, sha.as_str()));
assert_eq!(reports.last(), Some(&(200_000, Some(200_000))));
assert_eq!(fs::read(&target).map(|b| b.len()).unwrap_or(0), 200_000);
assert_eq!(leftovers(&dir), vec!["file.bin".to_string()]);
let wrong = "0".repeat(64);
let cases: Vec<(HttpDownload, Option<u64>)> = vec![
(HttpDownload::to(&target).with_sha256(&wrong), None),
(HttpDownload::to(&target).with_size(5), None),
(HttpDownload::to(&target).with_size(5), Some(200_000)),
(HttpDownload::to(&target).with_max_bytes(1000), None),
(HttpDownload::to(&target).with_max_bytes(1000), Some(200_000)),
(HttpDownload::to(&target), Some(300_000)),
(HttpDownload::to(&target).with_sha256("abc"), None),
(HttpDownload::to(&target).with_size(2000).with_max_bytes(1000), None),
];
for (settings, length) in cases {
let other = vec![1u8; 200_000];
let result = settings.receive(RequestId::next(), &mut &other[..], length, &mut |_, _| {});
assert!(result.is_err(), "{settings:?} {length:?}");
assert_eq!(fs::read(&target).ok(), Some(data.clone()), "{settings:?}");
assert_eq!(leftovers(&dir), vec!["file.bin".to_string()], "{settings:?}");
}
assert!(matches!(
HttpDownload::to(&target).with_max_bytes(1000).receive(RequestId::next(), &mut &data[..], None, &mut |_, _| {}),
Err(BackendError::BodyTooLarge { limit: 1000 })
));
assert!(matches!(
HttpDownload::to(dir.join("missing").join("x.bin")).receive(RequestId::next(), &mut &data[..], None, &mut |_, _| {}),
Err(BackendError::InvalidRequest(_))
));
assert!(matches!(HttpDownload::to("").check(), Err(BackendError::InvalidRequest(_))));
let debug = format!("{:?}", HttpDownload::to(dir.join("secret-folder").join("x.bin")));
assert!(debug.contains("x.bin") && !debug.contains("secret-folder"), "{debug}");
let _ = fs::remove_dir_all(&dir);
}
#[test]
fn a_stop_or_a_read_error_leaves_nothing() {
let dir = dir("stop");
let target = dir.join("file.bin");
let data = vec![3u8; 500_000];
let part = PartFile::create(&HttpDownload::to(&target), RequestId::next()).unwrap_or_else(|e| panic!("{e}"));
assert_eq!(leftovers(&dir).len(), 1, "the part file exists before anything is read");
let result = part.write_from(&mut &data[..], None, &mut |_, _| {}, &|| true, &|e| BackendError::Network(e.to_string()));
assert!(matches!(result, Err(BackendError::Cancelled)));
assert!(leftovers(&dir).is_empty());
struct Broken(usize);
impl Read for Broken {
fn read(&mut self, buf: &mut [u8]) -> io::Result<usize> {
if self.0 == 0 {
return Err(io::Error::new(io::ErrorKind::ConnectionReset, "reset"));
}
self.0 -= 1;
let n = buf.len().min(1000);
Ok(n)
}
}
let result = HttpDownload::to(&target).receive(RequestId::next(), &mut Broken(3), None, &mut |_, _| {});
assert!(matches!(result, Err(BackendError::Network(ref why)) if why.contains("reset")), "{result:?}");
let result = HttpDownload::to(&target).receive(RequestId::next(), &mut Broken(3), Some(10_000), &mut |_, _| {});
assert!(result.is_err());
assert!(leftovers(&dir).is_empty());
let _ = fs::remove_dir_all(&dir);
}
#[test]
fn stale_part_files_of_other_runs_are_removed_and_nothing_else() {
let own = std::process::id();
let other = own.wrapping_add(1);
assert!(is_stale_part(&format!("a.bin.{other}-12-0.part"), "a.bin", own));
for name in [
format!("a.bin.{own}-12-0.part"),
format!("a.bin.{other}-12.part"),
format!("a.bin.{other}-12-0-1.part"),
format!("a.bin.{other}-x-0.part"),
format!("a.bin.{other}--0.part"),
format!("a.bin.{other}-12-0.part.old"),
format!("b.bin.{other}-12-0.part"),
format!("a.bin{other}-12-0.part"),
"a.bin".to_string(),
"a.bin.part".to_string(),
format!("xa.bin.{other}-1-0.part"),
] {
assert!(!is_stale_part(&name, "a.bin", own), "{name}");
}
let dir = dir("stale");
let target = dir.join("a.bin");
let stale = format!("a.bin.{other}-3-7.part");
let keep = [format!("a.bin.{own}-3-7.part"), "a.bin.notes.part".to_string(), "b.bin.1-2-3.part".to_string()];
for name in keep.iter().chain([&stale]) {
fs::write(dir.join(name), b"x").unwrap_or_else(|e| panic!("{e}"));
}
fs::create_dir(dir.join(format!("a.bin.{other}-9-9.part"))).unwrap_or_else(|e| panic!("{e}"));
let data = [1u8; 10];
HttpDownload::to(&target).receive(RequestId::next(), &mut &data[..], None, &mut |_, _| {}).unwrap_or_else(|e| panic!("{e}"));
let mut left = leftovers(&dir);
left.sort();
let mut expected: Vec<String> = keep.to_vec();
expected.push("a.bin".into());
expected.push(format!("a.bin.{other}-9-9.part"));
expected.sort();
assert_eq!(left, expected, "only the other run's part file is gone (a folder of that name stays)");
let folder = dir.join("folder");
fs::create_dir(&folder).unwrap_or_else(|e| panic!("{e}"));
let result = HttpDownload::to(&folder).receive(RequestId::next(), &mut &data[..], None, &mut |_, _| {});
assert!(matches!(&result, Err(BackendError::InvalidRequest(why)) if why.contains("folder")), "{result:?}");
assert_eq!(leftovers(&folder), Vec::<String>::new());
let _ = fs::remove_dir_all(&dir);
}
}