use std::path::Path;
use async_trait::async_trait;
use crate::error::{DownloadError, VerifyError};
#[async_trait]
pub trait Sink: Send + Sync {
async fn write_at(&self, offset: u64, bytes: &[u8]) -> Result<(), DownloadError>;
async fn finalize(&self) -> Result<(), DownloadError> {
Ok(())
}
async fn truncate(&self, _len: u64) -> Result<(), DownloadError> {
Err(DownloadError::sink("truncation unsupported by this sink"))
}
fn supports_read_back(&self) -> bool {
false
}
async fn read_at(&self, _offset: u64, _len: u64) -> Result<Vec<u8>, DownloadError> {
Err(DownloadError::sink("read-back unsupported by this sink"))
}
fn staging_path(&self) -> Option<&Path> {
None
}
}
fn read_back_bounds(offset: u64, len: u64) -> Result<(usize, usize), DownloadError> {
let end = offset
.checked_add(len)
.ok_or_else(|| span_too_large(offset, len))?;
let start = usize::try_from(offset).map_err(|_| span_too_large(offset, len))?;
let end = usize::try_from(end).map_err(|_| span_too_large(offset, len))?;
Ok((start, end))
}
fn span_too_large(offset: u64, len: u64) -> DownloadError {
DownloadError::sink(format!(
"read-back span [{offset}, +{len}) does not fit this platform's address space"
))
}
fn try_zeroed_read_buffer(len: u64) -> Result<Vec<u8>, DownloadError> {
let len = usize::try_from(len).map_err(|_| span_too_large(0, len))?;
let mut buf: Vec<u8> = Vec::new();
buf.try_reserve_exact(len).map_err(|e| {
DownloadError::sink(format!(
"cannot allocate a {len}-byte read-back buffer: {e}"
))
})?;
buf.resize(len, 0); Ok(buf)
}
pub(crate) async fn promote_verified(
sink: &dyn Sink,
verified_len: u64,
) -> Result<(), DownloadError> {
sink.truncate(verified_len).await?;
let refuse = |reason: String| Err(DownloadError::Verify(VerifyError::Metadata(reason)));
if !sink.supports_read_back() {
return refuse(format!(
"this sink cannot read back its staging area, so a promotion of {verified_len} verified \
byte(s) cannot be proven to be the verified artifact; refusing to promote"
));
}
if verified_len > 0 && sink.read_at(verified_len - 1, 1).await.is_err() {
return refuse(format!(
"staging area is SHORTER than the verified length {verified_len}; refusing to promote a \
partial artifact as the verified one"
));
}
if sink.read_at(verified_len, 1).await.is_ok() {
return refuse(format!(
"staging area holds bytes past the verified length {verified_len}; refusing to promote an \
artifact that is not the verified one"
));
}
sink.finalize().await
}
#[derive(Debug, Default)]
pub struct InMemorySink {
inner: tokio::sync::Mutex<Inner>,
}
#[derive(Debug, Default)]
struct Inner {
buf: Vec<u8>,
finalized: bool,
}
impl InMemorySink {
pub fn new() -> Self {
InMemorySink::default()
}
pub async fn contents(&self) -> Vec<u8> {
self.inner.lock().await.buf.clone()
}
pub async fn is_finalized(&self) -> bool {
self.inner.lock().await.finalized
}
}
#[async_trait]
impl Sink for InMemorySink {
async fn write_at(&self, offset: u64, bytes: &[u8]) -> Result<(), DownloadError> {
let mut inner = self.inner.lock().await;
let (start, end) = read_back_bounds(offset, bytes.len() as u64)?;
if inner.buf.len() < end {
inner.buf.resize(end, 0);
}
inner.buf[start..end].copy_from_slice(bytes);
Ok(())
}
async fn finalize(&self) -> Result<(), DownloadError> {
self.inner.lock().await.finalized = true;
Ok(())
}
fn supports_read_back(&self) -> bool {
true }
async fn truncate(&self, len: u64) -> Result<(), DownloadError> {
let mut inner = self.inner.lock().await;
if let Ok(len) = usize::try_from(len) {
if inner.buf.len() > len {
inner.buf.truncate(len);
}
}
Ok(())
}
async fn read_at(&self, offset: u64, len: u64) -> Result<Vec<u8>, DownloadError> {
let inner = self.inner.lock().await;
let (start, end) = read_back_bounds(offset, len)?;
if inner.buf.len() < end {
return Err(DownloadError::sink(format!(
"read-back past staged end: want [{start}, {end}), have {}",
inner.buf.len()
)));
}
Ok(inner.buf[start..end].to_vec())
}
}
pub const TMP_SUFFIX: &str = ".download.tmp";
pub const STATE_SUFFIX: &str = ".download.tmp.state";
pub fn staging_path_for(final_path: &Path) -> std::path::PathBuf {
let mut s = final_path.as_os_str().to_owned();
s.push(TMP_SUFFIX);
std::path::PathBuf::from(s)
}
#[derive(Debug)]
pub struct FileSink {
final_path: std::path::PathBuf,
tmp_path: std::path::PathBuf,
file: tokio::sync::Mutex<Option<std::fs::File>>,
}
impl FileSink {
pub fn new(final_path: impl Into<std::path::PathBuf>) -> Self {
let final_path = final_path.into();
let tmp_path = staging_path_for(&final_path);
FileSink {
final_path,
tmp_path,
file: tokio::sync::Mutex::new(None),
}
}
pub fn final_path(&self) -> &Path {
&self.final_path
}
pub fn tmp_path(&self) -> &Path {
&self.tmp_path
}
}
#[async_trait]
impl Sink for FileSink {
async fn write_at(&self, offset: u64, bytes: &[u8]) -> Result<(), DownloadError> {
use std::io::{Seek, SeekFrom, Write};
let mut guard = self.file.lock().await;
if guard.is_none() {
if let Some(parent) = self.tmp_path.parent() {
std::fs::create_dir_all(parent).map_err(DownloadError::sink)?;
}
let f = std::fs::OpenOptions::new()
.read(true)
.write(true)
.create(true)
.truncate(false)
.open(&self.tmp_path)
.map_err(DownloadError::sink)?;
*guard = Some(f);
}
let f = guard.as_mut().expect("file opened above");
f.seek(SeekFrom::Start(offset))
.map_err(DownloadError::sink)?;
f.write_all(bytes).map_err(DownloadError::sink)?;
Ok(())
}
fn supports_read_back(&self) -> bool {
true }
async fn read_at(&self, offset: u64, len: u64) -> Result<Vec<u8>, DownloadError> {
use std::io::{Read, Seek, SeekFrom};
let mut guard = self.file.lock().await;
if guard.is_none() {
let f = std::fs::OpenOptions::new()
.read(true)
.write(true)
.truncate(false)
.open(&self.tmp_path)
.map_err(DownloadError::sink)?;
*guard = Some(f);
}
let f = guard.as_mut().expect("file opened above");
f.seek(SeekFrom::Start(offset))
.map_err(DownloadError::sink)?;
let mut buf = try_zeroed_read_buffer(len)?;
f.read_exact(&mut buf).map_err(|e| {
DownloadError::sink(format!("read-back of {len} bytes at {offset} failed: {e}"))
})?;
Ok(buf)
}
async fn truncate(&self, len: u64) -> Result<(), DownloadError> {
let mut guard = self.file.lock().await;
if guard.is_none() {
if !self.tmp_path.exists() {
return Ok(());
}
let f = std::fs::OpenOptions::new()
.read(true)
.write(true)
.truncate(false)
.open(&self.tmp_path)
.map_err(DownloadError::sink)?;
*guard = Some(f);
}
let f = guard.as_mut().expect("file opened above");
let staged = f.metadata().map_err(DownloadError::sink)?.len();
if staged > len {
f.set_len(len).map_err(DownloadError::sink)?;
}
Ok(())
}
async fn finalize(&self) -> Result<(), DownloadError> {
{
let mut guard = self.file.lock().await;
if let Some(f) = guard.as_mut() {
f.sync_all().map_err(DownloadError::sink)?;
}
*guard = None; }
std::fs::rename(&self.tmp_path, &self.final_path).map_err(DownloadError::sink)?;
Ok(())
}
fn staging_path(&self) -> Option<&Path> {
Some(&self.tmp_path)
}
}
#[cfg(test)]
mod tests {
use super::*;
#[tokio::test]
async fn writes_placed_by_offset_out_of_order() {
let sink = InMemorySink::new();
sink.write_at(3, b"DEF").await.unwrap();
sink.write_at(0, b"ABC").await.unwrap();
assert_eq!(sink.contents().await, b"ABCDEF");
assert!(!sink.is_finalized().await);
sink.finalize().await.unwrap();
assert!(sink.is_finalized().await);
}
#[tokio::test]
async fn overlapping_write_overwrites() {
let sink = InMemorySink::new();
sink.write_at(0, b"ABCDEF").await.unwrap();
sink.write_at(2, b"xy").await.unwrap();
assert_eq!(sink.contents().await, b"ABxyEF");
}
fn temp_dir(tag: &str) -> std::path::PathBuf {
let d = std::env::temp_dir().join(format!(
"dig-download-sink-{tag}-{}-{}",
std::process::id(),
std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap()
.as_nanos()
));
std::fs::create_dir_all(&d).unwrap();
d
}
#[tokio::test]
async fn file_sink_stages_then_atomically_finalizes() {
let dir = temp_dir("finalize");
let final_path = dir.join("resource.dig");
let sink = FileSink::new(&final_path);
sink.write_at(3, b"DEF").await.unwrap();
sink.write_at(0, b"ABC").await.unwrap();
assert!(sink.tmp_path().exists());
assert!(!final_path.exists());
assert_eq!(sink.tmp_path(), staging_path_for(&final_path));
sink.finalize().await.unwrap();
assert!(final_path.exists());
assert!(!sink.tmp_path().exists());
assert_eq!(std::fs::read(&final_path).unwrap(), b"ABCDEF");
let _ = std::fs::remove_dir_all(&dir);
}
#[tokio::test]
async fn file_sink_resume_reattaches_without_truncating() {
let dir = temp_dir("resume");
let final_path = dir.join("resource.dig");
{
let sink = FileSink::new(&final_path);
sink.write_at(3, b"DEF").await.unwrap();
}
assert!(staging_path_for(&final_path).exists());
let sink2 = FileSink::new(&final_path);
sink2.write_at(0, b"ABC").await.unwrap();
sink2.finalize().await.unwrap();
assert_eq!(std::fs::read(&final_path).unwrap(), b"ABCDEF");
let _ = std::fs::remove_dir_all(&dir);
}
#[tokio::test]
async fn an_absurd_read_back_span_is_a_typed_error_not_an_abort() {
let sink = InMemorySink::new();
sink.write_at(0, b"eight!!!").await.unwrap();
let err = sink
.read_at(1, u64::MAX)
.await
.expect_err("an unsatisfiable span is refused");
assert!(matches!(err, DownloadError::Sink(_)), "typed error: {err}");
}
#[tokio::test]
async fn read_back_never_creates_the_staging_file() {
let dir = temp_dir("no-create");
let sink = FileSink::new(dir.join("resource.dig"));
assert!(sink.read_at(0, 4).await.is_err());
assert!(
!sink.tmp_path().exists(),
"a read did not conjure a staging file"
);
let _ = std::fs::remove_dir_all(&dir);
}
#[tokio::test]
async fn truncate_shrinks_and_never_extends() {
let sink = InMemorySink::new();
sink.write_at(0, b"ABCDEF").await.unwrap();
sink.truncate(3).await.unwrap();
assert_eq!(sink.contents().await, b"ABC");
sink.truncate(99).await.unwrap();
assert_eq!(sink.contents().await, b"ABC", "a longer len is a no-op");
}
#[tokio::test]
async fn file_sink_truncate_shrinks_the_staging_file() {
let dir = temp_dir("truncate");
let final_path = dir.join("resource.dig");
let sink = FileSink::new(&final_path);
sink.write_at(0, b"ABCDEF").await.unwrap();
sink.truncate(3).await.unwrap();
sink.finalize().await.unwrap();
assert_eq!(std::fs::read(&final_path).unwrap(), b"ABC");
let _ = std::fs::remove_dir_all(&dir);
}
#[tokio::test]
async fn truncate_never_creates_the_staging_file() {
let dir = temp_dir("truncate-no-create");
let sink = FileSink::new(dir.join("resource.dig"));
sink.truncate(0).await.unwrap();
assert!(!sink.tmp_path().exists());
let _ = std::fs::remove_dir_all(&dir);
}
#[test]
fn staging_path_appends_suffix() {
let p = staging_path_for(Path::new("/data/x.dig"));
assert!(p.to_string_lossy().ends_with(".dig.download.tmp"));
}
#[tokio::test]
async fn promotion_refuses_a_staging_area_that_is_shorter_than_the_verified_length() {
let sink = InMemorySink::new();
sink.write_at(0, b"only ten!!").await.unwrap();
let err = promote_verified(&sink, 40)
.await
.expect_err("a partial artifact is never promoted as the verified one");
assert!(
err.to_string().contains("SHORTER than the verified length"),
"and it names the side it refused, so this cannot silently become the long-side test: {err}"
);
assert!(!sink.is_finalized().await);
}
#[tokio::test]
async fn promotion_refuses_a_staging_area_holding_bytes_past_the_verified_length() {
let sink = InMemorySink::new();
sink.write_at(0, b"eight!!!").await.unwrap();
struct NoOpTruncate(InMemorySink);
#[async_trait]
impl Sink for NoOpTruncate {
async fn write_at(&self, offset: u64, bytes: &[u8]) -> Result<(), DownloadError> {
self.0.write_at(offset, bytes).await
}
async fn truncate(&self, _len: u64) -> Result<(), DownloadError> {
Ok(())
}
fn supports_read_back(&self) -> bool {
true
}
async fn read_at(&self, offset: u64, len: u64) -> Result<Vec<u8>, DownloadError> {
self.0.read_at(offset, len).await
}
async fn finalize(&self) -> Result<(), DownloadError> {
self.0.finalize().await
}
}
let sink = NoOpTruncate(sink);
let err = promote_verified(&sink, 4)
.await
.expect_err("a longer staged artifact is never promoted");
assert!(
err.to_string().contains("past the verified length"),
"names the side it refused: {err}"
);
assert!(!sink.0.is_finalized().await);
}
#[tokio::test]
async fn promotion_refuses_a_sink_that_cannot_read_its_staging_area_back() {
struct WriteOnly;
#[async_trait]
impl Sink for WriteOnly {
async fn write_at(&self, _offset: u64, _bytes: &[u8]) -> Result<(), DownloadError> {
Ok(())
}
async fn truncate(&self, _len: u64) -> Result<(), DownloadError> {
Ok(()) }
}
let err = promote_verified(&WriteOnly, 8)
.await
.expect_err("an unprovable promotion is refused");
assert!(
err.to_string().contains("cannot read back"),
"names the missing capability rather than guessing: {err}"
);
}
}