use std::fs::File;
use std::io::{Read, Write};
use std::path::{Path, PathBuf};
use std::sync::atomic::{AtomicU64, Ordering};
use crate::error::{IoAction, Result, StryptError};
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
#[non_exhaustive]
pub struct Limits {
pub max_input_bytes: u64,
}
impl Limits {
pub const DEFAULT_MAX_INPUT_BYTES: u64 = 512 * 1024 * 1024;
}
impl Limits {
#[must_use]
pub const fn with_max_input_bytes(max_input_bytes: u64) -> Self {
Self { max_input_bytes }
}
}
impl Default for Limits {
fn default() -> Self {
Self {
max_input_bytes: Self::DEFAULT_MAX_INPUT_BYTES,
}
}
}
pub fn read_bounded(path: &Path, limits: Limits) -> Result<Vec<u8>> {
let file = File::open(path).map_err(|source| StryptError::Io {
action: IoAction::ReadingInput,
source,
})?;
let declared = file
.metadata()
.map_err(|source| StryptError::Io {
action: IoAction::MeasuringInput,
source,
})?
.len();
if declared > limits.max_input_bytes {
return Err(StryptError::InputTooLarge {
limit: limits.max_input_bytes,
actual: Some(declared),
});
}
read_bounded_from(file, limits, Some(declared))
}
fn read_bounded_from<R: Read>(source: R, limits: Limits, hint: Option<u64>) -> Result<Vec<u8>> {
let probe = limits.max_input_bytes.saturating_add(1);
let mut buffer = Vec::new();
if let Some(hint) = hint {
let reserve = hint.min(limits.max_input_bytes);
if let Ok(reserve) = usize::try_from(reserve) {
buffer
.try_reserve_exact(reserve)
.map_err(|_| StryptError::InputTooLarge {
limit: limits.max_input_bytes,
actual: Some(hint),
})?;
}
}
let read = source
.take(probe)
.read_to_end(&mut buffer)
.map_err(|source| StryptError::Io {
action: IoAction::ReadingInput,
source,
})?;
if u64::try_from(read).unwrap_or(u64::MAX) > limits.max_input_bytes {
return Err(StryptError::InputTooLarge {
limit: limits.max_input_bytes,
actual: None,
});
}
Ok(buffer)
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum Overwrite {
Refuse,
Replace,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum Permissions {
OwnerOnly,
Inherit,
}
#[derive(Debug)]
pub struct AtomicWrite {
destination: PathBuf,
temporary: PathBuf,
file: Option<File>,
permissions: Permissions,
}
impl AtomicWrite {
pub fn begin(
destination: &Path,
overwrite: Overwrite,
permissions: Permissions,
) -> Result<Self> {
if overwrite == Overwrite::Refuse && destination.exists() {
return Err(StryptError::Io {
action: IoAction::CreatingTemporary,
source: std::io::Error::new(
std::io::ErrorKind::AlreadyExists,
"destination exists",
),
});
}
let temporary = temporary_path_for(destination);
let file = create_private(&temporary, permissions)?;
Ok(Self {
destination: destination.to_path_buf(),
temporary,
file: Some(file),
permissions,
})
}
pub fn write_all(&mut self, bytes: &[u8]) -> Result<()> {
let Some(file) = self.file.as_mut() else {
return Err(StryptError::Io {
action: IoAction::WritingOutput,
source: std::io::Error::other("write after the file was finished"),
});
};
file.write_all(bytes).map_err(|source| StryptError::Io {
action: IoAction::WritingOutput,
source,
})
}
pub fn commit(mut self) -> Result<()> {
let Some(file) = self.file.take() else {
return Err(StryptError::Io {
action: IoAction::SyncingOutput,
source: std::io::Error::other("already finished"),
});
};
if let Err(source) = file.sync_all() {
self.discard_temporary();
return Err(StryptError::Io {
action: IoAction::SyncingOutput,
source,
});
}
drop(file);
if let Err(source) = std::fs::rename(&self.temporary, &self.destination) {
self.discard_temporary();
return Err(StryptError::Io {
action: IoAction::ReplacingDestination,
source,
});
}
apply_permissions(&self.destination, self.permissions)?;
Ok(())
}
pub fn abort(mut self) {
self.file = None;
self.discard_temporary();
}
fn discard_temporary(&mut self) {
self.file = None;
let _ = std::fs::remove_file(&self.temporary);
}
}
impl Drop for AtomicWrite {
fn drop(&mut self) {
if self.file.is_some() {
self.discard_temporary();
}
}
}
fn temporary_path_for(destination: &Path) -> PathBuf {
static COUNTER: AtomicU64 = AtomicU64::new(0);
let nonce = COUNTER.fetch_add(1, Ordering::Relaxed);
let clock: u32 = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.map_or(0, |d| d.subsec_nanos());
let name = format!(".strypt-{}-{clock}-{nonce}.tmp", std::process::id());
destination.parent().unwrap_or(Path::new(".")).join(name)
}
fn create_private(path: &Path, permissions: Permissions) -> Result<File> {
let mut options = std::fs::OpenOptions::new();
options.write(true).create_new(true);
#[cfg(unix)]
if permissions == Permissions::OwnerOnly {
use std::os::unix::fs::OpenOptionsExt as _;
options.mode(0o600);
}
let file = options.open(path).map_err(|source| StryptError::Io {
action: IoAction::CreatingTemporary,
source,
})?;
let _ = permissions;
Ok(file)
}
fn apply_permissions(path: &Path, permissions: Permissions) -> Result<()> {
#[cfg(unix)]
if permissions == Permissions::OwnerOnly {
use std::os::unix::fs::PermissionsExt as _;
std::fs::set_permissions(path, std::fs::Permissions::from_mode(0o600)).map_err(
|source| StryptError::Io {
action: IoAction::SettingPermissions,
source,
},
)?;
}
let _ = (path, permissions);
Ok(())
}
#[cfg(test)]
mod tests {
#![allow(clippy::unwrap_used)]
use super::*;
struct Scratch(PathBuf);
impl Scratch {
fn new(tag: &str) -> Self {
let path = std::env::temp_dir().join(format!(
"strypt-test-{tag}-{}-{:?}",
std::process::id(),
std::thread::current().id()
));
std::fs::create_dir_all(&path).unwrap();
Self(path)
}
fn join(&self, name: &str) -> PathBuf {
self.0.join(name)
}
}
impl Drop for Scratch {
fn drop(&mut self) {
let _ = std::fs::remove_dir_all(&self.0);
}
}
#[test]
fn a_file_over_the_limit_is_refused_not_truncated() {
let dir = Scratch::new("limit");
let path = dir.join("big.bin");
std::fs::write(&path, vec![0u8; 4096]).unwrap();
let err = read_bounded(
&path,
Limits {
max_input_bytes: 1024,
},
)
.unwrap_err();
assert!(
matches!(
err,
StryptError::InputTooLarge {
limit: 1024,
actual: Some(4096)
}
),
"got {err:?}"
);
}
#[test]
fn a_lying_size_hint_cannot_get_past_the_ceiling() {
let data = vec![0u8; 4096];
let err = read_bounded_from(
data.as_slice(),
Limits {
max_input_bytes: 1024,
},
Some(16),
)
.unwrap_err();
assert!(
matches!(err, StryptError::InputTooLarge { .. }),
"got {err:?}"
);
}
#[test]
fn a_file_exactly_at_the_limit_is_accepted() {
let dir = Scratch::new("exact");
let path = dir.join("exact.bin");
std::fs::write(&path, vec![7u8; 1024]).unwrap();
let got = read_bounded(
&path,
Limits {
max_input_bytes: 1024,
},
)
.unwrap();
assert_eq!(got.len(), 1024);
}
#[test]
fn nothing_appears_at_the_destination_until_commit() {
let dir = Scratch::new("atomic");
let dest = dir.join("out.bin");
let mut w = AtomicWrite::begin(&dest, Overwrite::Refuse, Permissions::OwnerOnly).unwrap();
w.write_all(b"partial").unwrap();
assert!(
!dest.exists(),
"a half-written file must never be visible at the destination path"
);
w.commit().unwrap();
assert_eq!(std::fs::read(&dest).unwrap(), b"partial");
}
#[test]
fn an_abandoned_write_leaves_no_temporary_behind() {
let dir = Scratch::new("abort");
let dest = dir.join("out.bin");
let mut w = AtomicWrite::begin(&dest, Overwrite::Refuse, Permissions::OwnerOnly).unwrap();
w.write_all(b"doomed").unwrap();
w.abort();
assert!(!dest.exists());
let leftovers: Vec<_> = std::fs::read_dir(&dir.0)
.unwrap()
.filter_map(std::result::Result::ok)
.filter(|e| e.file_name().to_string_lossy().starts_with(".strypt-"))
.collect();
assert!(leftovers.is_empty(), "temporary files were left behind");
}
#[test]
fn dropping_a_writer_mid_failure_cleans_up() {
let dir = Scratch::new("drop");
let dest = dir.join("out.bin");
{
let mut w =
AtomicWrite::begin(&dest, Overwrite::Refuse, Permissions::OwnerOnly).unwrap();
w.write_all(b"incomplete").unwrap();
}
assert!(!dest.exists());
let count = std::fs::read_dir(&dir.0).unwrap().count();
assert_eq!(
count, 0,
"the temporary should have been dropped with the writer"
);
}
#[test]
fn an_existing_destination_is_refused_by_default() {
let dir = Scratch::new("refuse");
let dest = dir.join("out.bin");
std::fs::write(&dest, b"the user's file").unwrap();
let err = AtomicWrite::begin(&dest, Overwrite::Refuse, Permissions::OwnerOnly).unwrap_err();
assert!(matches!(err, StryptError::Io { .. }));
assert_eq!(
std::fs::read(&dest).unwrap(),
b"the user's file",
"the existing file must be untouched"
);
}
#[test]
fn replace_is_available_when_the_caller_asks_for_it() {
let dir = Scratch::new("replace");
let dest = dir.join("out.bin");
std::fs::write(&dest, b"old").unwrap();
let mut w = AtomicWrite::begin(&dest, Overwrite::Replace, Permissions::OwnerOnly).unwrap();
w.write_all(b"new").unwrap();
w.commit().unwrap();
assert_eq!(std::fs::read(&dest).unwrap(), b"new");
}
#[cfg(unix)]
#[test]
fn output_is_not_readable_by_anyone_else() {
use std::os::unix::fs::PermissionsExt as _;
let dir = Scratch::new("perms");
let dest = dir.join("out.bin");
let mut w = AtomicWrite::begin(&dest, Overwrite::Refuse, Permissions::OwnerOnly).unwrap();
w.write_all(b"sensitive").unwrap();
w.commit().unwrap();
let mode = std::fs::metadata(&dest).unwrap().permissions().mode() & 0o777;
assert_eq!(mode, 0o600, "ADR-0019: stripped output is owner-only");
}
#[cfg(unix)]
#[test]
fn the_temporary_is_private_while_it_is_being_written() {
use std::os::unix::fs::PermissionsExt as _;
let dir = Scratch::new("temp-perms");
let dest = dir.join("out.bin");
let w = AtomicWrite::begin(&dest, Overwrite::Refuse, Permissions::OwnerOnly).unwrap();
let mode = std::fs::metadata(&w.temporary)
.unwrap()
.permissions()
.mode()
& 0o777;
assert_eq!(mode, 0o600);
w.abort();
}
#[test]
fn the_temporary_sits_beside_the_destination_not_in_tmpdir() {
let dir = Scratch::new("location");
let dest = dir.join("out.bin");
let w = AtomicWrite::begin(&dest, Overwrite::Refuse, Permissions::OwnerOnly).unwrap();
assert_eq!(w.temporary.parent(), dest.parent());
w.abort();
}
}