diff --git a/src/download.rs b/src/download.rs index e5aef9a..07aabc8 100644 --- a/src/download.rs +++ b/src/download.rs @@ -13,7 +13,6 @@ use std::fs::File; use std::io; use std::io::BufReader; use std::path::Path; -use std::path::PathBuf; use anyhow::Context as _; use rand::distr::Distribution; @@ -78,15 +77,14 @@ pub fn download_artifact( // This must be a sibling to the final artifact_location so that we can // atomically move it into place. let temp_dir_to_mv = fs_ctx::tempdir_in(artifact_parent_dir)?; - let fetch_destination: PathBuf = { - let fetch_destination = fs_ctx::namedtempfile_new_in(artifact_parent_dir) - .context("failed to create fetch temp path")? - .into_temp_path(); - // fetch_destination is dropped after this and is removed from - // disk. This is deliberate since we want a unique path and not - // necessarily the file. - fetch_destination.to_path_buf() - }; + let fetch_destination = fs_ctx::namedtempfile_new_in(artifact_parent_dir) + .context("failed to create fetch temp path")? + .into_temp_path(); + // Providers expect a unique path that they can create themselves. Keep + // the TempPath guard alive so the downloaded file is still removed + // after fetching, verification, and extraction. + fs_ctx::remove_file(&fetch_destination) + .context("failed to remove fetch temp placeholder")?; let provider_type = get_provider_type(provider_config)?; let provider = provider_factory.get_provider(provider_type)?; @@ -330,9 +328,169 @@ pub fn acquire_download_lock_for_artifact( #[cfg(test)] mod tests { + use std::io::Write; + use std::path::PathBuf; + use std::sync::Arc; + use std::sync::Mutex; + + use flate2::Compression; + use flate2::write::GzEncoder; use rand::SeedableRng; use super::*; + use crate::artifact_path::ArtifactPath; + use crate::config::Arg0; + use crate::provider::Provider; + + struct TestProviderFactory { + contents: Vec, + destination: Arc>>, + fail_after_write: bool, + } + + impl ProviderFactory for TestProviderFactory { + fn get_provider(&self, _provider_type: &str) -> anyhow::Result> { + Ok(Box::new(TestProvider { + contents: self.contents.clone(), + destination: Arc::clone(&self.destination), + fail_after_write: self.fail_after_write, + })) + } + } + + struct TestProvider { + contents: Vec, + destination: Arc>>, + fail_after_write: bool, + } + + impl Provider for TestProvider { + fn fetch_artifact( + &self, + _provider_config: &Value, + destination: &Path, + _fetch_lock: &FileLock, + _artifact_entry: &ArtifactEntry, + ) -> anyhow::Result<()> { + fs_ctx::file_create(destination)?.write_all(&self.contents)?; + *self.destination.lock().unwrap() = Some(destination.to_owned()); + if self.fail_after_write { + anyhow::bail!("simulated fetch failure"); + } + Ok(()) + } + } + + fn gzip(contents: &[u8]) -> Vec { + let mut encoder = GzEncoder::new(Vec::new(), Compression::default()); + encoder.write_all(contents).unwrap(); + encoder.finish().unwrap() + } + + fn artifact_entry(contents: &[u8], digest: Digest) -> ArtifactEntry { + ArtifactEntry { + size: contents.len() as u64, + hash: HashAlgorithm::Sha256, + digest, + format: ArtifactFormat::Gz, + path: ArtifactPath::try_from("tool".to_owned()).unwrap(), + providers: vec![serde_json::json!({"type": "test"})], + arg0: Arg0::DotslashFile, + readonly: false, + providers_order: ProvidersOrder::Sequential, + } + } + + fn sha256_digest(contents: &[u8]) -> Digest { + Digest::try_from(format!("{:x}", Sha256::digest(contents))).unwrap() + } + + fn artifact_location(root: &Path) -> ArtifactLocation { + let artifact_directory = root.join("artifact"); + ArtifactLocation { + executable: artifact_directory.join("tool"), + artifact_directory, + lock_path: root.join("artifact.lock"), + arg0: Arg0::DotslashFile, + } + } + + #[test] + fn removes_fetch_destination_after_extracting_archive() { + let extracted_contents = b"downloaded artifact"; + let fetched_contents = gzip(extracted_contents); + let entry = artifact_entry(&fetched_contents, sha256_digest(&fetched_contents)); + let tempdir = tempfile::tempdir().unwrap(); + let location = artifact_location(tempdir.path()); + let destination = Arc::new(Mutex::new(None)); + let provider_factory = TestProviderFactory { + contents: fetched_contents, + destination: Arc::clone(&destination), + fail_after_write: false, + }; + + download_artifact(&entry, &location, &provider_factory).unwrap(); + + assert_eq!( + fs_ctx::read(&location.executable).unwrap(), + extracted_contents + ); + assert!(!destination.lock().unwrap().as_ref().unwrap().exists()); + } + + #[test] + fn removes_fetch_destination_after_verification_failure() { + let fetched_contents = gzip(b"downloaded artifact"); + let entry = artifact_entry(&fetched_contents, sha256_digest(b"different artifact")); + let tempdir = tempfile::tempdir().unwrap(); + let location = artifact_location(tempdir.path()); + let destination = Arc::new(Mutex::new(None)); + let provider_factory = TestProviderFactory { + contents: fetched_contents, + destination: Arc::clone(&destination), + fail_after_write: false, + }; + + download_artifact(&entry, &location, &provider_factory).unwrap_err(); + + assert!(!destination.lock().unwrap().as_ref().unwrap().exists()); + } + + #[test] + fn removes_fetch_destination_after_extraction_failure() { + let fetched_contents = b"not a gzip archive".to_vec(); + let entry = artifact_entry(&fetched_contents, sha256_digest(&fetched_contents)); + let tempdir = tempfile::tempdir().unwrap(); + let location = artifact_location(tempdir.path()); + let destination = Arc::new(Mutex::new(None)); + let provider_factory = TestProviderFactory { + contents: fetched_contents, + destination: Arc::clone(&destination), + fail_after_write: false, + }; + + download_artifact(&entry, &location, &provider_factory).unwrap_err(); + + assert!(!destination.lock().unwrap().as_ref().unwrap().exists()); + } + + #[test] + fn removes_fetch_destination_after_fetch_failure() { + let fetched_contents = gzip(b"partial artifact"); + let entry = artifact_entry(&fetched_contents, sha256_digest(&fetched_contents)); + let tempdir = tempfile::tempdir().unwrap(); + let location = artifact_location(tempdir.path()); + let destination = Arc::new(Mutex::new(None)); + let provider_factory = TestProviderFactory { + contents: fetched_contents, + destination: Arc::clone(&destination), + fail_after_write: true, + }; + + download_artifact(&entry, &location, &provider_factory).unwrap_err(); + + assert!(!destination.lock().unwrap().as_ref().unwrap().exists()); + } #[test] fn providers_in_order_sequential() { diff --git a/src/util/fs_ctx.rs b/src/util/fs_ctx.rs index cdd57f8..b43a96f 100644 --- a/src/util/fs_ctx.rs +++ b/src/util/fs_ctx.rs @@ -127,7 +127,6 @@ pub fn rename, Q: AsRef>(from: P, to: Q) -> io::Result<()> fs::rename(&from, &to).map_err(|source| wrap2(source, "rename from", from, "to", to)) } -#[cfg_attr(not(dotslash_internal), expect(dead_code))] pub fn remove_file>(path: P) -> io::Result<()> { fs::remove_file(&path).map_err(|source| wrap1(source, "remove file", path)) }