Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
178 changes: 168 additions & 10 deletions src/download.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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;
Expand Down Expand Up @@ -78,15 +77,14 @@ pub fn download_artifact<P: ProviderFactory>(
// 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)?;
Expand Down Expand Up @@ -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<u8>,
destination: Arc<Mutex<Option<PathBuf>>>,
fail_after_write: bool,
}

impl ProviderFactory for TestProviderFactory {
fn get_provider(&self, _provider_type: &str) -> anyhow::Result<Box<dyn Provider>> {
Ok(Box::new(TestProvider {
contents: self.contents.clone(),
destination: Arc::clone(&self.destination),
fail_after_write: self.fail_after_write,
}))
}
}

struct TestProvider {
contents: Vec<u8>,
destination: Arc<Mutex<Option<PathBuf>>>,
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<u8> {
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() {
Expand Down
1 change: 0 additions & 1 deletion src/util/fs_ctx.rs
Original file line number Diff line number Diff line change
Expand Up @@ -127,7 +127,6 @@ pub fn rename<P: AsRef<Path>, Q: AsRef<Path>>(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<P: AsRef<Path>>(path: P) -> io::Result<()> {
fs::remove_file(&path).map_err(|source| wrap1(source, "remove file", path))
}
Expand Down