use std::{
fs::File,
io::Write as _,
path::{Path, PathBuf},
};
use sha2::{Digest as _, Sha256};
pub const MAGIC: [u8; 4] = *b"PPP1";
pub const FORMAT_REVISION: u8 = 1;
const PARTITION_DIGEST_LEN: usize = 32;
const INCARNATION_LEN: usize = 32;
const BOUNDARY_LEN: usize = 8;
const CHECKPOINT_DIGEST_LEN: usize = 32;
const CHECKSUM_LEN: usize = 32;
const BOUND_LEN: usize = MAGIC.len()
+ 1 + PARTITION_DIGEST_LEN
+ INCARNATION_LEN
+ BOUNDARY_LEN + CHECKPOINT_DIGEST_LEN
+ 1 + BOUNDARY_LEN; const RECORD_LEN: usize = BOUND_LEN + CHECKSUM_LEN;
pub const FILE_NAME: &str = "prunepass";
const TEMP_FILE_NAME: &str = "prunepass.writing";
const STAGE_INTENT: u8 = 0;
const STAGE_RECEIPT: u8 = 1;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct PrunePass {
partition_digest: [u8; PARTITION_DIGEST_LEN],
incarnation: [u8; INCARNATION_LEN],
requested_boundary: u64,
checkpoint_digest: [u8; CHECKPOINT_DIGEST_LEN],
applied_boundary: Option<u64>,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum PrunePassOutcome {
Absent,
Intent(Box<PrunePass>),
Receipt(Box<PrunePass>),
}
#[derive(Debug, thiserror::Error)]
#[non_exhaustive]
pub enum PruneRecordError {
#[error("prune pass i/o failure: {0}")]
Io(#[from] std::io::Error),
#[error("prune pass is {found} byte(s), not {expected}")]
WrongLength {
expected: usize,
found: usize,
},
#[error("prune pass carries another family's magic")]
WrongMagic,
#[error("prune pass format revision {found} is not {expected}")]
UnsupportedRevision {
expected: u8,
found: u8,
},
#[error("prune pass fails its checksum")]
ChecksumMismatch,
#[error("prune pass was recorded for a different partition")]
WrongPartition,
#[error("prune pass declares an unknown stage {found}")]
UnknownStage {
found: u8,
},
}
impl PrunePass {
#[must_use]
pub fn intent(
partition: &str,
incarnation: [u8; INCARNATION_LEN],
requested_boundary: u64,
checkpoint_digest: [u8; CHECKPOINT_DIGEST_LEN],
) -> Self {
Self {
partition_digest: digest_partition(partition),
incarnation,
requested_boundary,
checkpoint_digest,
applied_boundary: None,
}
}
#[must_use]
pub const fn into_receipt(self, applied_boundary: u64) -> Self {
Self {
applied_boundary: Some(applied_boundary),
..self
}
}
#[must_use]
pub const fn incarnation(&self) -> [u8; INCARNATION_LEN] {
self.incarnation
}
#[must_use]
pub const fn requested_boundary(&self) -> u64 {
self.requested_boundary
}
#[must_use]
pub const fn checkpoint_digest(&self) -> [u8; CHECKPOINT_DIGEST_LEN] {
self.checkpoint_digest
}
#[must_use]
pub const fn applied_boundary(&self) -> Option<u64> {
self.applied_boundary
}
fn to_bytes(self) -> [u8; RECORD_LEN] {
let mut bound = [0_u8; BOUND_LEN];
let mut at = 0;
bound[at..at + MAGIC.len()].copy_from_slice(&MAGIC);
at += MAGIC.len();
bound[at] = FORMAT_REVISION;
at += 1;
bound[at..at + PARTITION_DIGEST_LEN].copy_from_slice(&self.partition_digest);
at += PARTITION_DIGEST_LEN;
bound[at..at + INCARNATION_LEN].copy_from_slice(&self.incarnation);
at += INCARNATION_LEN;
bound[at..at + BOUNDARY_LEN].copy_from_slice(&self.requested_boundary.to_be_bytes());
at += BOUNDARY_LEN;
bound[at..at + CHECKPOINT_DIGEST_LEN].copy_from_slice(&self.checkpoint_digest);
at += CHECKPOINT_DIGEST_LEN;
bound[at] = if self.applied_boundary.is_some() {
STAGE_RECEIPT
} else {
STAGE_INTENT
};
at += 1;
bound[at..at + BOUNDARY_LEN]
.copy_from_slice(&self.applied_boundary.unwrap_or(0).to_be_bytes());
at += BOUNDARY_LEN;
debug_assert_eq!(at, BOUND_LEN);
let checksum = Sha256::digest(bound);
let mut record = [0_u8; RECORD_LEN];
record[..BOUND_LEN].copy_from_slice(&bound);
record[BOUND_LEN..].copy_from_slice(&checksum);
record
}
fn from_bytes(bytes: &[u8], partition: &str) -> Result<Self, PruneRecordError> {
if bytes.len() != RECORD_LEN {
return Err(PruneRecordError::WrongLength {
expected: RECORD_LEN,
found: bytes.len(),
});
}
let (bound, checksum) = bytes.split_at(BOUND_LEN);
let computed = Sha256::digest(bound);
if computed.as_slice() != checksum {
return Err(PruneRecordError::ChecksumMismatch);
}
let mut at = 0;
if bound[at..at + MAGIC.len()] != MAGIC {
return Err(PruneRecordError::WrongMagic);
}
at += MAGIC.len();
let revision = bound[at];
at += 1;
if revision != FORMAT_REVISION {
return Err(PruneRecordError::UnsupportedRevision {
expected: FORMAT_REVISION,
found: revision,
});
}
let partition_digest: [u8; PARTITION_DIGEST_LEN] = bound[at..at + PARTITION_DIGEST_LEN]
.try_into()
.expect("slice is exactly PARTITION_DIGEST_LEN bytes");
at += PARTITION_DIGEST_LEN;
if partition_digest != digest_partition(partition) {
return Err(PruneRecordError::WrongPartition);
}
let incarnation: [u8; INCARNATION_LEN] = bound[at..at + INCARNATION_LEN]
.try_into()
.expect("slice is exactly INCARNATION_LEN bytes");
at += INCARNATION_LEN;
let requested_boundary = read_u64(bound, &mut at);
let checkpoint_digest: [u8; CHECKPOINT_DIGEST_LEN] = bound[at..at + CHECKPOINT_DIGEST_LEN]
.try_into()
.expect("slice is exactly CHECKPOINT_DIGEST_LEN bytes");
at += CHECKPOINT_DIGEST_LEN;
let stage = bound[at];
at += 1;
let applied_boundary_value = read_u64(bound, &mut at);
let applied_boundary = match stage {
STAGE_INTENT => None,
STAGE_RECEIPT => Some(applied_boundary_value),
found => return Err(PruneRecordError::UnknownStage { found }),
};
Ok(Self {
partition_digest,
incarnation,
requested_boundary,
checkpoint_digest,
applied_boundary,
})
}
pub fn write_atomic(&self, dir: &Path) -> Result<(), PruneRecordError> {
std::fs::create_dir_all(dir)?;
let bytes = self.to_bytes();
let temp = dir.join(TEMP_FILE_NAME);
{
let mut file = File::create(&temp)?;
file.write_all(&bytes)?;
file.sync_all()?;
}
std::fs::rename(&temp, dir.join(FILE_NAME))?;
File::open(dir)?.sync_all()?;
Ok(())
}
pub fn read(dir: &Path, partition: &str) -> Result<PrunePassOutcome, PruneRecordError> {
let path: PathBuf = dir.join(FILE_NAME);
let bytes = match std::fs::read(&path) {
Ok(bytes) => bytes,
Err(error) if error.kind() == std::io::ErrorKind::NotFound => {
return Ok(PrunePassOutcome::Absent);
}
Err(error) => return Err(error.into()),
};
let pass = Self::from_bytes(&bytes, partition)?;
Ok(match pass.applied_boundary {
None => PrunePassOutcome::Intent(Box::new(pass)),
Some(_) => PrunePassOutcome::Receipt(Box::new(pass)),
})
}
pub fn invalidate(dir: &Path) -> Result<(), PruneRecordError> {
match std::fs::remove_file(dir.join(FILE_NAME)) {
Ok(()) => {}
Err(error) if error.kind() == std::io::ErrorKind::NotFound => return Ok(()),
Err(error) => return Err(error.into()),
}
File::open(dir)?.sync_all()?;
Ok(())
}
}
fn read_u64(bound: &[u8], at: &mut usize) -> u64 {
let value = u64::from_be_bytes(
bound[*at..*at + 8]
.try_into()
.expect("slice is exactly 8 bytes"),
);
*at += 8;
value
}
fn digest_partition(partition: &str) -> [u8; PARTITION_DIGEST_LEN] {
Sha256::digest(partition.as_bytes()).into()
}
#[cfg(test)]
mod tests {
use super::*;
fn sample_intent(partition: &str) -> PrunePass {
PrunePass::intent(partition, [7_u8; 32], 42, [9_u8; 32])
}
#[test]
fn a_fresh_directory_reads_as_absent_not_an_error() {
let dir = tempfile::tempdir().expect("temp dir");
let outcome = PrunePass::read(dir.path(), "conv-1").expect("read");
assert_eq!(outcome, PrunePassOutcome::Absent);
}
#[test]
fn an_intent_round_trips_as_intent() {
let dir = tempfile::tempdir().expect("temp dir");
let written = sample_intent("conv-1");
written.write_atomic(dir.path()).expect("write");
let outcome = PrunePass::read(dir.path(), "conv-1").expect("read");
assert_eq!(outcome, PrunePassOutcome::Intent(Box::new(written)));
}
#[test]
fn a_receipt_built_from_its_intent_round_trips_as_receipt() {
let dir = tempfile::tempdir().expect("temp dir");
let receipt = sample_intent("conv-1").into_receipt(40);
receipt.write_atomic(dir.path()).expect("write");
let outcome = PrunePass::read(dir.path(), "conv-1").expect("read");
assert_eq!(outcome, PrunePassOutcome::Receipt(Box::new(receipt)));
assert_eq!(receipt.applied_boundary(), Some(40));
}
#[test]
fn a_receipt_overwrites_its_own_intent_rather_than_appending() {
let dir = tempfile::tempdir().expect("temp dir");
let intent = sample_intent("conv-1");
intent.write_atomic(dir.path()).expect("write");
let receipt = intent.into_receipt(41);
receipt.write_atomic(dir.path()).expect("write");
let outcome = PrunePass::read(dir.path(), "conv-1").expect("read");
assert_eq!(outcome, PrunePassOutcome::Receipt(Box::new(receipt)));
assert!(!dir.path().join(TEMP_FILE_NAME).exists());
}
#[test]
fn invalidate_removes_a_present_pass_and_is_a_no_op_when_absent() {
let dir = tempfile::tempdir().expect("temp dir");
PrunePass::invalidate(dir.path()).expect("invalidate an absent pass");
sample_intent("conv-1")
.write_atomic(dir.path())
.expect("write");
PrunePass::invalidate(dir.path()).expect("invalidate a present pass");
assert_eq!(
PrunePass::read(dir.path(), "conv-1").unwrap(),
PrunePassOutcome::Absent
);
}
#[test]
fn a_truncated_record_is_refused_not_read_as_shorter_or_empty() {
let dir = tempfile::tempdir().expect("temp dir");
sample_intent("conv-1")
.write_atomic(dir.path())
.expect("write");
let mut bytes = std::fs::read(dir.path().join(FILE_NAME)).expect("read raw");
bytes.truncate(bytes.len() - 1);
std::fs::write(dir.path().join(FILE_NAME), &bytes).expect("write truncated");
let error = PrunePass::read(dir.path(), "conv-1").expect_err("must refuse");
assert!(matches!(error, PruneRecordError::WrongLength { .. }));
}
#[test]
fn a_flipped_checksum_byte_is_refused() {
let dir = tempfile::tempdir().expect("temp dir");
sample_intent("conv-1")
.write_atomic(dir.path())
.expect("write");
let mut bytes = std::fs::read(dir.path().join(FILE_NAME)).expect("read raw");
let last = bytes.len() - 1;
bytes[last] ^= 0xFF;
std::fs::write(dir.path().join(FILE_NAME), &bytes).expect("write flipped");
let error = PrunePass::read(dir.path(), "conv-1").expect_err("must refuse");
assert!(matches!(error, PruneRecordError::ChecksumMismatch));
}
#[test]
fn another_familys_magic_is_refused() {
let dir = tempfile::tempdir().expect("temp dir");
sample_intent("conv-1")
.write_atomic(dir.path())
.expect("write");
let mut bytes = std::fs::read(dir.path().join(FILE_NAME)).expect("read raw");
bytes[0] ^= 0xFF;
let checksum = Sha256::digest(&bytes[..bytes.len() - CHECKSUM_LEN]);
let new_len = bytes.len() - CHECKSUM_LEN;
bytes[new_len..].copy_from_slice(&checksum);
std::fs::write(dir.path().join(FILE_NAME), &bytes).expect("write foreign magic");
let error = PrunePass::read(dir.path(), "conv-1").expect_err("must refuse");
assert!(matches!(error, PruneRecordError::WrongMagic));
}
#[test]
fn an_unsupported_revision_is_refused() {
let dir = tempfile::tempdir().expect("temp dir");
sample_intent("conv-1")
.write_atomic(dir.path())
.expect("write");
let mut bytes = std::fs::read(dir.path().join(FILE_NAME)).expect("read raw");
bytes[MAGIC.len()] = FORMAT_REVISION + 1;
let checksum = Sha256::digest(&bytes[..bytes.len() - CHECKSUM_LEN]);
let new_len = bytes.len() - CHECKSUM_LEN;
bytes[new_len..].copy_from_slice(&checksum);
std::fs::write(dir.path().join(FILE_NAME), &bytes).expect("write future revision");
let error = PrunePass::read(dir.path(), "conv-1").expect_err("must refuse");
assert!(matches!(
error,
PruneRecordError::UnsupportedRevision {
found: 2,
expected: 1
}
));
}
#[test]
fn a_valid_artifact_for_a_different_partition_is_refused() {
let dir = tempfile::tempdir().expect("temp dir");
sample_intent("conv-1")
.write_atomic(dir.path())
.expect("write");
let error = PrunePass::read(dir.path(), "conv-2").expect_err("must refuse");
assert!(matches!(error, PruneRecordError::WrongPartition));
}
#[test]
fn an_unknown_stage_byte_is_refused() {
let dir = tempfile::tempdir().expect("temp dir");
sample_intent("conv-1")
.write_atomic(dir.path())
.expect("write");
let mut bytes = std::fs::read(dir.path().join(FILE_NAME)).expect("read raw");
let stage_offset = BOUND_LEN - BOUNDARY_LEN - 1;
bytes[stage_offset] = 2;
let checksum = Sha256::digest(&bytes[..bytes.len() - CHECKSUM_LEN]);
let new_len = bytes.len() - CHECKSUM_LEN;
bytes[new_len..].copy_from_slice(&checksum);
std::fs::write(dir.path().join(FILE_NAME), &bytes).expect("write unknown stage");
let error = PrunePass::read(dir.path(), "conv-1").expect_err("must refuse");
assert!(matches!(error, PruneRecordError::UnknownStage { found: 2 }));
}
#[test]
fn accessors_return_exactly_what_was_bound() {
let pass = sample_intent("conv-1");
assert_eq!(pass.incarnation(), [7_u8; 32]);
assert_eq!(pass.requested_boundary(), 42);
assert_eq!(pass.checkpoint_digest(), [9_u8; 32]);
assert_eq!(pass.applied_boundary(), None);
let receipt = pass.into_receipt(40);
assert_eq!(receipt.applied_boundary(), Some(40));
assert_eq!(receipt.requested_boundary(), 42);
assert_eq!(receipt.checkpoint_digest(), [9_u8; 32]);
}
}