use crate::error::{Error, Result};
use crate::io::WriteHandle;
use std::path::{Path, PathBuf};
pub const STATE_SUFFIX: &str = ".rst-state";
const STATE_MAGIC: &[u8; 4] = b"RSTS";
const STATE_VERSION: u16 = 1;
const STATE_HEADER_LEN: usize = 4 + 2 + 4 + 8 + 8 + 32;
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct ChunkBitmap {
bits: Vec<u8>,
len: u64,
set_count: u64,
}
impl ChunkBitmap {
pub fn new(len: u64) -> Self {
Self {
bits: vec![0u8; len.div_ceil(8) as usize],
len,
set_count: 0,
}
}
pub fn from_bytes(bits: Vec<u8>, len: u64) -> Result<Self> {
let expect = len.div_ceil(8) as usize;
if bits.len() != expect {
return Err(Error::protocol(format!(
"bitmap is {} bytes, expected {expect} for {len} chunks",
bits.len()
)));
}
if len % 8 != 0 {
if let Some(last) = bits.last() {
let valid = (len % 8) as u32;
if last >> valid != 0 {
return Err(Error::protocol("bitmap has bits set past the chunk count"));
}
}
}
let set_count = bits.iter().map(|b| b.count_ones() as u64).sum();
Ok(Self {
bits,
len,
set_count,
})
}
pub fn as_bytes(&self) -> &[u8] {
&self.bits
}
pub fn len(&self) -> u64 {
self.len
}
pub fn is_empty(&self) -> bool {
self.len == 0
}
#[inline]
pub fn get(&self, i: u64) -> bool {
if i >= self.len {
return false;
}
self.bits[(i / 8) as usize] & (1 << (i % 8)) != 0
}
#[inline]
pub fn set(&mut self, i: u64) -> bool {
if i >= self.len {
return false;
}
let byte = &mut self.bits[(i / 8) as usize];
let mask = 1u8 << (i % 8);
if *byte & mask != 0 {
return false;
}
*byte |= mask;
self.set_count += 1;
true
}
pub fn count(&self) -> u64 {
self.set_count
}
pub fn is_complete(&self) -> bool {
self.set_count == self.len
}
pub fn missing(&self) -> impl Iterator<Item = u64> + '_ {
(0..self.len).filter(move |&i| !self.get(i))
}
pub fn fill(&mut self) {
for b in self.bits.iter_mut() {
*b = 0xFF;
}
if self.len % 8 != 0 {
let valid = (self.len % 8) as u32;
if let Some(last) = self.bits.last_mut() {
*last = (1u8 << valid) - 1;
}
}
self.set_count = self.len;
}
}
pub struct ResumeState {
path: PathBuf,
bitmap: ChunkBitmap,
size: u64,
chunk_size: u32,
dirty: u64,
}
impl ResumeState {
pub fn state_path(final_path: &Path) -> PathBuf {
let mut s = final_path.as_os_str().to_os_string();
s.push(STATE_SUFFIX);
PathBuf::from(s)
}
pub fn load_or_new(final_path: &Path, size: u64, chunk_size: u32) -> Self {
let path = Self::state_path(final_path);
let chunks = if chunk_size == 0 {
0
} else {
size.div_ceil(chunk_size as u64)
};
if let Some(bm) = Self::try_load(&path, size, chunk_size, chunks) {
return Self {
path,
bitmap: bm,
size,
chunk_size,
dirty: 0,
};
}
Self {
path,
bitmap: ChunkBitmap::new(chunks),
size,
chunk_size,
dirty: 0,
}
}
fn try_load(path: &Path, size: u64, chunk_size: u32, chunks: u64) -> Option<ChunkBitmap> {
let raw = std::fs::read(path).ok()?;
if raw.len() < STATE_HEADER_LEN || &raw[0..4] != STATE_MAGIC {
return None;
}
if u16::from_le_bytes(raw[4..6].try_into().ok()?) != STATE_VERSION {
return None;
}
if u32::from_le_bytes(raw[6..10].try_into().ok()?) != chunk_size {
return None;
}
if u64::from_le_bytes(raw[10..18].try_into().ok()?) != size {
return None;
}
let n = u64::from_le_bytes(raw[18..26].try_into().ok()?);
if n != chunks {
return None;
}
let stored_hash = &raw[26..58];
let body = &raw[STATE_HEADER_LEN..];
if blake3::hash(body).as_bytes() != stored_hash {
tracing::warn!(path = %path.display(), "resume state failed its checksum; starting over");
return None;
}
ChunkBitmap::from_bytes(body.to_vec(), chunks).ok()
}
pub fn bitmap(&self) -> &ChunkBitmap {
&self.bitmap
}
pub fn has(&self, chunk: u64) -> bool {
self.bitmap.get(chunk)
}
pub fn record(&mut self, chunk: u64) {
if self.bitmap.set(chunk) {
self.dirty += 1;
}
}
pub fn pending(&self) -> u64 {
self.dirty
}
pub fn is_complete(&self) -> bool {
self.bitmap.is_complete()
}
pub fn checkpoint(&mut self, handle: &WriteHandle) -> Result<()> {
if self.dirty == 0 {
return Ok(());
}
handle.sync()?;
self.persist()?;
self.dirty = 0;
Ok(())
}
fn persist(&self) -> Result<()> {
let body = self.bitmap.as_bytes();
let mut out = Vec::with_capacity(STATE_HEADER_LEN + body.len());
out.extend_from_slice(STATE_MAGIC);
out.extend_from_slice(&STATE_VERSION.to_le_bytes());
out.extend_from_slice(&self.chunk_size.to_le_bytes());
out.extend_from_slice(&self.size.to_le_bytes());
out.extend_from_slice(&self.bitmap.len().to_le_bytes());
out.extend_from_slice(blake3::hash(body).as_bytes());
out.extend_from_slice(body);
let tmp = self.path.with_extension("tmp");
std::fs::write(&tmp, &out)?;
std::fs::rename(&tmp, &self.path)?;
Ok(())
}
pub fn clear(&self) {
let _ = std::fs::remove_file(&self.path);
let _ = std::fs::remove_file(self.path.with_extension("tmp"));
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn bitmap_set_get_count() {
let mut b = ChunkBitmap::new(20);
assert_eq!(b.count(), 0);
assert!(!b.is_complete());
assert!(b.set(0));
assert!(!b.set(0), "setting twice must not double-count");
assert!(b.set(19));
assert!(b.get(0) && b.get(19));
assert!(!b.get(5));
assert_eq!(b.count(), 2);
assert!(!b.set(20), "out of range set is a no-op");
assert!(!b.get(20));
assert_eq!(b.missing().count(), 18);
b.fill();
assert!(b.is_complete());
assert_eq!(b.count(), 20);
assert_eq!(b.missing().count(), 0);
}
#[test]
fn bitmap_roundtrips_and_rejects_bad_input() {
let mut b = ChunkBitmap::new(100);
for i in (0..100).step_by(3) {
b.set(i);
}
let restored = ChunkBitmap::from_bytes(b.as_bytes().to_vec(), 100).unwrap();
assert_eq!(restored, b);
assert_eq!(restored.count(), b.count());
assert!(ChunkBitmap::from_bytes(vec![0u8; 3], 100).is_err());
assert!(ChunkBitmap::from_bytes(vec![0xFF; 13], 100).is_err());
}
#[test]
fn bitmap_scales_to_a_100gb_file() {
let chunks = 100 * 1024u64;
let mut b = ChunkBitmap::new(chunks);
assert_eq!(b.as_bytes().len(), 12_800, "12.8 KB of state for 100 GiB");
for i in 0..chunks {
b.set(i);
}
assert!(b.is_complete());
}
#[test]
fn state_survives_a_reload() {
let tmp = tempfile::tempdir().unwrap();
let dest = tmp.path().join("f.bin");
let h = WriteHandle::open(&dest, 10_000, false).unwrap();
let mut s = ResumeState::load_or_new(&dest, 10_000, 1000);
assert_eq!(s.bitmap().len(), 10);
for c in [0u64, 1, 2, 7] {
h.write_at(c * 1000, &[1u8; 1000]).unwrap();
s.record(c);
}
assert_eq!(s.pending(), 4);
s.checkpoint(&h).unwrap();
assert_eq!(s.pending(), 0);
let s2 = ResumeState::load_or_new(&dest, 10_000, 1000);
assert_eq!(s2.bitmap().count(), 4);
for c in [0u64, 1, 2, 7] {
assert!(s2.has(c));
}
assert!(!s2.has(3));
}
#[test]
fn state_is_discarded_when_the_file_changed() {
let tmp = tempfile::tempdir().unwrap();
let dest = tmp.path().join("f.bin");
let h = WriteHandle::open(&dest, 10_000, false).unwrap();
let mut s = ResumeState::load_or_new(&dest, 10_000, 1000);
s.record(0);
s.checkpoint(&h).unwrap();
let s2 = ResumeState::load_or_new(&dest, 20_000, 1000);
assert_eq!(s2.bitmap().count(), 0);
let s3 = ResumeState::load_or_new(&dest, 10_000, 4096);
assert_eq!(s3.bitmap().count(), 0);
}
#[test]
fn corrupt_state_falls_back_to_a_full_transfer() {
let tmp = tempfile::tempdir().unwrap();
let dest = tmp.path().join("f.bin");
let h = WriteHandle::open(&dest, 10_000, false).unwrap();
let mut s = ResumeState::load_or_new(&dest, 10_000, 1000);
s.record(0);
s.record(1);
s.checkpoint(&h).unwrap();
let p = ResumeState::state_path(&dest);
let mut raw = std::fs::read(&p).unwrap();
let last = raw.len() - 1;
raw[last] ^= 0xFF;
std::fs::write(&p, &raw).unwrap();
let s2 = ResumeState::load_or_new(&dest, 10_000, 1000);
assert_eq!(s2.bitmap().count(), 0, "corruption must not be trusted");
}
#[test]
fn clear_removes_the_sidecar() {
let tmp = tempfile::tempdir().unwrap();
let dest = tmp.path().join("f.bin");
let h = WriteHandle::open(&dest, 1000, false).unwrap();
let mut s = ResumeState::load_or_new(&dest, 1000, 1000);
s.record(0);
s.checkpoint(&h).unwrap();
assert!(ResumeState::state_path(&dest).exists());
s.clear();
assert!(!ResumeState::state_path(&dest).exists());
}
}