use crate::codec::compress;
use crate::codec::crypto::{Handshake, Role, Sealer, HANDSHAKE_MSG_LEN, TAG_LEN};
use crate::config::Config;
use crate::error::{Error, Result};
use crate::io::{FileWriters, WriteHandle};
use crate::manifest;
use crate::metrics::{Metrics, Progress, ProgressFn};
use crate::pool::{BufPool, ObjPool};
use crate::resume::ResumeState;
use crate::send::merkle_root;
use crate::transport::{BoxRecv, Transport};
use crate::wire::{self, Control, EntryKind, FileEntry, FrameHeader, LocalFileIndex, ResumeEntry};
use futures_util::stream::{FuturesOrdered, StreamExt};
use parking_lot::Mutex;
use std::collections::HashMap;
use std::path::{Path, PathBuf};
use std::sync::atomic::{AtomicBool, AtomicU64, Ordering};
use std::sync::Arc;
use tokio::io::{AsyncReadExt, AsyncWriteExt};
use tokio::sync::mpsc;
const CHECKPOINT_INTERVAL: u64 = 256;
const CHECKPOINT_MAX_AGE: std::time::Duration = std::time::Duration::from_secs(5);
const FAREWELL_TIMEOUT: std::time::Duration = std::time::Duration::from_secs(20);
struct RecvFile {
entry: FileEntry,
dest: PathBuf,
preexisting: crate::resume::ChunkBitmap,
fresh_part: bool,
state: Mutex<ResumeState>,
chunk_hashes: Mutex<Vec<[u8; 32]>>,
expected_hash: Mutex<Option<[u8; 32]>>,
existing: Option<crate::io::ChunkReader>,
announced: AtomicBool,
finalized: AtomicBool,
since_checkpoint: AtomicU64,
last_checkpoint: Mutex<std::time::Instant>,
}
struct RecvShared {
cfg: Config,
root: PathBuf,
files: HashMap<u32, Arc<RecvFile>>,
writers: FileWriters,
metrics: Metrics,
pending: AtomicU64,
}
pub async fn receive(
transport: Arc<dyn Transport>,
dest_root: impl AsRef<Path>,
cfg: &Config,
progress: Option<ProgressFn>,
) -> Result<Progress> {
cfg.validate()?;
let root = dest_root.as_ref().to_path_buf();
tokio::fs::create_dir_all(&root).await?;
let root = tokio::fs::canonicalize(&root).await.unwrap_or(root);
let metrics = Metrics::new();
let (mut ctl_w, mut ctl_r) = transport.accept_bi().await?;
let hs = Handshake::new(Role::Responder, &cfg.secrecy, cfg.cipher);
let mut peer = [0u8; HANDSHAKE_MSG_LEN];
tokio::time::timeout(cfg.handshake_timeout, ctl_r.read_exact(&mut peer))
.await
.map_err(|_| Error::Handshake("timed out waiting for the peer's handshake".into()))?
.map_err(map_eof)?;
ctl_w.write_all(hs.message()).await?;
ctl_w.flush().await?;
let crypto = Arc::new(hs.finish(&peer)?);
let entries = match wire::read_control(
&mut ctl_r,
cfg.max_frame_bytes,
cfg.max_manifest_entries,
)
.await?
{
Control::Manifest(e) => e,
Control::Abort { reason } => return Err(Error::Closed(reason)),
other => return Err(Error::protocol(format!("expected Manifest, got {other:?}"))),
};
manifest::validate(&entries, cfg)?;
let mut symlinks = Vec::new();
let mut files = HashMap::new();
let mut resume_reply = Vec::new();
let mut local_index: Vec<LocalFileIndex> = Vec::new();
let mut cache = (cfg.delta && cfg.trust_mtime).then(|| crate::index::ChunkIndex::load(&root));
#[allow(unused_mut)]
let mut seen: std::collections::HashSet<String> = std::collections::HashSet::new();
let mut total_bytes = 0u64;
let mut file_count = 0u64;
for entry in &entries {
match entry.kind {
EntryKind::Directory => {
let dest = manifest::safe_join(&root, &entry.path)?;
tokio::fs::create_dir_all(&dest).await?;
}
EntryKind::Symlink => {
let (rel, target) = manifest::split_symlink(&entry.path)?;
let dest = manifest::safe_join(&root, rel)?;
symlinks.push((dest, target.to_string()));
}
EntryKind::File => {
let dest = manifest::safe_join(&root, &entry.path)?;
if let Some(parent) = dest.parent() {
tokio::fs::create_dir_all(parent).await?;
}
total_bytes += entry.size;
file_count += 1;
let part = crate::io::part_path_for(&dest);
let fresh_part = !part.exists();
if fresh_part {
ResumeState::load_or_new(&dest, entry.size, entry.chunk_size).clear();
}
if !cfg.resume {
ResumeState::load_or_new(&dest, entry.size, entry.chunk_size).clear();
}
let state = ResumeState::load_or_new(&dest, entry.size, entry.chunk_size);
let (existing, local_hashes) = if cfg.delta && dest.exists() {
index_existing(
&dest,
&entry.path,
entry.chunk_size,
cfg.chunk_hash_budget,
cache.as_mut(),
)
} else {
(None, Vec::new())
};
if !local_hashes.is_empty() {
local_index.push(LocalFileIndex {
file_id: entry.file_id,
hashes: local_hashes,
});
}
let preexisting = state.bitmap().clone();
if preexisting.count() > 0 {
resume_reply.push(ResumeEntry {
file_id: entry.file_id,
have: preexisting.as_bytes().to_vec(),
});
metrics.chunk_skipped(0);
}
let n = entry.chunk_count() as usize;
files.insert(
entry.file_id,
Arc::new(RecvFile {
entry: entry.clone(),
dest,
preexisting,
existing,
fresh_part,
state: Mutex::new(state),
chunk_hashes: Mutex::new(vec![[0u8; 32]; n]),
expected_hash: Mutex::new(None),
announced: AtomicBool::new(false),
finalized: AtomicBool::new(false),
since_checkpoint: AtomicU64::new(0),
last_checkpoint: Mutex::new(std::time::Instant::now()),
}),
);
}
}
}
metrics.set_totals(file_count, total_bytes);
let shared = Arc::new(RecvShared {
cfg: cfg.clone(),
root: root.clone(),
files,
writers: FileWriters::new(),
metrics: metrics.clone(),
pending: AtomicU64::new(file_count),
});
if !local_index.is_empty() {
let reusable: usize = local_index.iter().map(|e| e.hashes.len()).sum();
tracing::info!(
files = local_index.len(),
chunks = reusable,
"offering existing blocks for reuse"
);
for batch in split_index(local_index, cfg.max_frame_bytes) {
wire::write_control(&mut ctl_w, &Control::LocalIndex(batch)).await?;
}
}
wire::write_control(&mut ctl_w, &Control::ResumeState(resume_reply)).await?;
for f in shared.files.values() {
if f.entry.size == 0 {
let h = WriteHandle::open(&f.dest, 0, false)?;
h.commit(f.entry.mode, f.entry.mtime, cfg.preserve_metadata)?;
f.finalized.store(true, Ordering::Release);
shared.pending.fetch_sub(1, Ordering::AcqRel);
metrics.file_done();
}
}
let stream_count = match wire::read_control(
&mut ctl_r,
cfg.max_frame_bytes,
cfg.max_manifest_entries,
)
.await?
{
Control::Start { streams } => streams as usize,
Control::Abort { reason } => return Err(Error::Closed(reason)),
other => return Err(Error::protocol(format!("expected Start, got {other:?}"))),
};
if stream_count == 0 || stream_count > 1024 {
return Err(Error::protocol(format!(
"sender asked for {stream_count} data streams, which is outside the accepted range"
)));
}
let cpu = Arc::new(
rayon::ThreadPoolBuilder::new()
.num_threads(cfg.workers)
.thread_name(|i| format!("rst-decode-{i}"))
.build()
.map_err(|e| Error::Worker(e.to_string()))?,
);
let pool = BufPool::new(
stream_count * cfg.queue_depth * 2 + cfg.workers * 2,
cfg.chunk_size + cfg.chunk_size / 8,
);
let openers: Arc<ObjPool<Sealer>> = ObjPool::new(cfg.workers + stream_count);
let progress_task = progress.map(|f| {
let m = metrics.clone();
tokio::spawn(async move {
let mut tick = tokio::time::interval(std::time::Duration::from_millis(500));
tick.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Skip);
loop {
tick.tick().await;
f(m.snapshot());
}
})
});
let (done_tx, mut done_rx) = mpsc::unbounded_channel::<Result<()>>();
let ctl_task = {
let shared = shared.clone();
let done_tx = done_tx.clone();
tokio::spawn(async move {
let r = run_control(&mut ctl_r, shared).await;
let _ = done_tx.send(r);
ctl_r
})
};
let mut stream_tasks = Vec::with_capacity(stream_count);
for _ in 0..stream_count {
let src = transport.accept_uni().await?;
stream_tasks.push(tokio::spawn(run_stream(
src,
shared.clone(),
openers.clone(),
crypto.clone(),
cpu.clone(),
pool.clone(),
)));
}
drop(done_tx);
let mut first_err = None;
for t in stream_tasks {
match t.await {
Ok(Ok(())) => {}
Ok(Err(e)) => {
tracing::error!(error = %e, "data stream failed");
first_err.get_or_insert(e);
}
Err(e) => {
first_err.get_or_insert(Error::Worker(e.to_string()));
}
}
}
let mut ctl_r = ctl_task.await.map_err(|e| Error::Worker(e.to_string()))?;
while let Some(r) = done_rx.recv().await {
if let Err(e) = r {
first_err.get_or_insert(e);
}
}
if let Some(e) = first_err {
checkpoint_all(&shared);
return Err(e);
}
let left = shared.pending.load(Ordering::Acquire);
if left != 0 {
checkpoint_all(&shared);
let names: Vec<_> = shared
.files
.values()
.filter(|f| !f.finalized.load(Ordering::Acquire))
.take(5)
.map(|f| f.entry.path.clone())
.collect();
return Err(Error::Protocol(format!(
"sender finished with {left} files incomplete (e.g. {names:?})"
)));
}
if let Some(mut cache) = cache.take() {
for f in shared.files.values() {
if !f.finalized.load(Ordering::Acquire) {
continue;
}
seen.insert(f.entry.path.clone());
if let Ok(meta) = std::fs::metadata(&f.dest) {
let hashes = f.chunk_hashes.lock().clone();
if !hashes.is_empty() && hashes.iter().any(|h| *h != [0u8; 32]) {
cache.insert(
&f.entry.path,
meta.len(),
crate::index::mtime_of(&meta),
f.entry.chunk_size,
hashes,
);
}
}
}
cache.retain(&seen);
if let Err(e) = cache.save() {
tracing::warn!(error = %e, "could not persist the chunk index");
}
}
create_symlinks(&root, &symlinks).await;
apply_directory_metadata(&root, &entries, cfg).await;
wire::write_control(&mut ctl_w, &Control::AllComplete).await?;
match tokio::time::timeout(FAREWELL_TIMEOUT, drain_to_eof(&mut ctl_r)).await {
Ok(_) => {}
Err(_) => tracing::debug!("sender did not close its control stream in time"),
}
let _ = ctl_w.shutdown().await;
if let Some(p) = progress_task {
p.abort();
}
let snapshot = metrics.snapshot();
tracing::info!(summary = %snapshot, "receive complete");
Ok(snapshot)
}
async fn run_control(ctl_r: &mut BoxRecv, shared: Arc<RecvShared>) -> Result<()> {
loop {
let msg = match wire::read_control(
ctl_r,
shared.cfg.max_frame_bytes,
shared.cfg.max_manifest_entries,
)
.await
{
Ok(m) => m,
Err(Error::Closed(_)) => return Ok(()),
Err(e) => return Err(e),
};
match msg {
Control::FileComplete { file_id, hash } => {
if let Some(f) = shared.files.get(&file_id) {
*f.expected_hash.lock() = hash;
f.announced.store(true, Ordering::Release);
finalize_if_ready(&shared, f)?;
}
}
Control::AllComplete => return Ok(()),
Control::Abort { reason } => return Err(Error::Closed(reason)),
other => {
tracing::debug!(?other, "ignoring unexpected control message");
}
}
}
}
async fn run_stream(
mut src: BoxRecv,
shared: Arc<RecvShared>,
openers: Arc<ObjPool<Sealer>>,
crypto: Arc<crate::codec::crypto::SessionCrypto>,
cpu: Arc<rayon::ThreadPool>,
pool: Arc<BufPool>,
) -> Result<()> {
let mut inflight: FuturesOrdered<tokio::sync::oneshot::Receiver<Result<()>>> =
FuturesOrdered::new();
let max_frame = shared.cfg.max_frame_bytes;
loop {
while inflight.len() >= shared.cfg.queue_depth {
if let Some(res) = inflight.next().await {
res.map_err(|_| Error::Worker("decode worker vanished".into()))??;
}
}
let mut head = [0u8; wire::FRAME_HEADER_LEN];
match src.read_exact(&mut head).await {
Ok(_) => {}
Err(e) if e.kind() == std::io::ErrorKind::UnexpectedEof => break,
Err(e) => return Err(Error::Io(e)),
}
let header = FrameHeader::decode(&head, max_frame)?;
let body_len = header.wire_payload_len();
if body_len > max_frame {
return Err(Error::FrameTooLarge {
got: body_len,
limit: max_frame,
});
}
let mut body = pool.take();
body.resize(body_len, 0);
src.read_exact(&mut body).await.map_err(map_eof)?;
let (tx, rx) = tokio::sync::oneshot::channel();
let shared2 = shared.clone();
let openers2 = openers.clone();
let crypto2 = crypto.clone();
let pool2 = pool.clone();
cpu.spawn(move || {
let mut opener = openers2.take_or(|| crypto2.opener());
let r = decode_and_write(header, head, body, &shared2, &mut opener, &pool2);
openers2.put(opener);
let _ = tx.send(r);
});
inflight.push_back(rx);
}
while let Some(res) = inflight.next().await {
res.map_err(|_| Error::Worker("decode worker vanished".into()))??;
}
Ok(())
}
fn decode_and_write(
header: FrameHeader,
head: [u8; wire::FRAME_HEADER_LEN],
mut body: Vec<u8>,
shared: &RecvShared,
opener: &mut Sealer,
pool: &BufPool,
) -> Result<()> {
let file = shared.files.get(&header.file_id).ok_or_else(|| {
Error::protocol(format!(
"frame references file_id {} which is not in the manifest",
header.file_id
))
})?;
let total_chunks = file.entry.chunk_count();
if header.chunk_index >= total_chunks {
return Err(Error::protocol(format!(
"chunk {} is past the end of file {} ({} chunks)",
header.chunk_index, header.file_id, total_chunks
)));
}
let chunk_size_u64 = file.entry.chunk_size as u64;
let offset = header.chunk_index * chunk_size_u64;
let expect_len = chunk_size_u64.min(file.entry.size - offset) as usize;
if header.reuses_local() {
if header.payload_len != 0 || !body.is_empty() {
return Err(Error::protocol("reuse frame carries a payload"));
}
if header.raw_len as usize != expect_len {
return Err(Error::protocol(format!(
"reuse chunk {} of file {} declares {} bytes, expected {expect_len}",
header.chunk_index, header.file_id, header.raw_len
)));
}
pool.put(body);
let Some(existing) = file.existing.as_ref() else {
return Err(Error::protocol(
"peer asked us to reuse a block from a file we never offered",
));
};
let mut buf = pool.take();
buf.resize(expect_len, 0);
let n = existing.read_at(offset, &mut buf)?;
if n != expect_len {
return Err(Error::Io(std::io::Error::new(
std::io::ErrorKind::UnexpectedEof,
format!("existing copy is short at chunk {}", header.chunk_index),
)));
}
let handle = shared.writers.get_or_open(header.file_id, || {
WriteHandle::open_cloned(&file.dest, file.entry.size, shared.cfg.preallocate)
})?;
if !handle.matches_at(offset, &buf)? {
handle.write_at(offset, &buf)?;
}
if shared.cfg.verify_hashes {
let h = *blake3::hash(&buf).as_bytes();
let mut hashes = file.chunk_hashes.lock();
if let Some(slot) = hashes.get_mut(header.chunk_index as usize) {
*slot = h;
}
}
pool.put(buf);
return record_chunk(
shared,
file,
&handle,
header,
ChunkOutcome {
raw_len: expect_len as u64,
wire_len: wire::FRAME_HEADER_LEN as u64,
compressed: false,
was_hole: false,
was_reuse: true,
},
);
}
if header.is_zero() {
if header.payload_len != 0 || !body.is_empty() {
return Err(Error::protocol("zero-chunk frame carries a payload"));
}
if header.raw_len as usize != expect_len {
return Err(Error::protocol(format!(
"zero chunk {} of file {} declares {} bytes, expected {expect_len}",
header.chunk_index, header.file_id, header.raw_len
)));
}
pool.put(body);
let handle = shared.writers.get_or_open(header.file_id, || {
WriteHandle::open(&file.dest, file.entry.size, shared.cfg.preallocate)
})?;
if !file.fresh_part {
handle.write_zeros_at(offset, expect_len)?;
}
if shared.cfg.verify_hashes {
let zeros = vec![0u8; expect_len];
let h = *blake3::hash(&zeros).as_bytes();
let mut hashes = file.chunk_hashes.lock();
if let Some(slot) = hashes.get_mut(header.chunk_index as usize) {
*slot = h;
}
}
return record_chunk(
shared,
file,
&handle,
header,
ChunkOutcome {
raw_len: expect_len as u64,
wire_len: wire::FRAME_HEADER_LEN as u64,
compressed: false,
was_hole: true,
was_reuse: false,
},
);
}
if header.sealed() {
if body.len() < TAG_LEN {
return Err(Error::protocol("sealed frame is shorter than its tag"));
}
let split = body.len() - TAG_LEN;
let mut tag = [0u8; TAG_LEN];
tag.copy_from_slice(&body[split..]);
body.truncate(split);
opener.open(
header.file_id,
header.chunk_index,
header.epoch,
&head,
&mut body,
&tag,
)?;
} else if !opener.is_passthrough() {
return Err(Error::protocol(
"peer sent an unsealed frame on an encrypted session",
));
}
if header.raw_len as usize != expect_len {
return Err(Error::protocol(format!(
"chunk {} of file {} declares {} plaintext bytes, expected {expect_len}",
header.chunk_index, header.file_id, header.raw_len
)));
}
let mut plain = pool.take();
compress::decompress_into(header.algorithm, header.raw_len as usize, &body, &mut plain)?;
let wire_len = (wire::FRAME_HEADER_LEN + body.len()) as u64;
pool.put(body);
let handle = shared.writers.get_or_open(header.file_id, || {
WriteHandle::open(&file.dest, file.entry.size, shared.cfg.preallocate)
})?;
handle.write_at(offset, &plain)?;
if shared.cfg.verify_hashes {
let h = *blake3::hash(&plain).as_bytes();
let mut hashes = file.chunk_hashes.lock();
if let Some(slot) = hashes.get_mut(header.chunk_index as usize) {
*slot = h;
}
}
let raw_len = plain.len() as u64;
pool.put(plain);
record_chunk(
shared,
file,
&handle,
header,
ChunkOutcome {
raw_len,
wire_len,
compressed: header.algorithm != compress::Algorithm::None,
was_hole: false,
was_reuse: false,
},
)
}
struct ChunkOutcome {
raw_len: u64,
wire_len: u64,
compressed: bool,
was_hole: bool,
was_reuse: bool,
}
fn record_chunk(
shared: &RecvShared,
file: &Arc<RecvFile>,
handle: &WriteHandle,
header: FrameHeader,
outcome: ChunkOutcome,
) -> Result<()> {
let ChunkOutcome {
raw_len,
wire_len,
compressed,
was_hole,
was_reuse,
} = outcome;
let complete = {
let mut st = file.state.lock();
st.record(header.chunk_index);
if shared.cfg.resume {
let by_count =
file.since_checkpoint.fetch_add(1, Ordering::AcqRel) + 1 >= CHECKPOINT_INTERVAL;
let by_age = file.last_checkpoint.lock().elapsed() >= CHECKPOINT_MAX_AGE;
if by_count || by_age {
file.since_checkpoint.store(0, Ordering::Release);
*file.last_checkpoint.lock() = std::time::Instant::now();
st.checkpoint(handle)?;
}
}
st.is_complete()
};
if was_reuse {
shared.metrics.chunk_reused(raw_len, wire_len);
} else if was_hole {
shared.metrics.chunk_zero(raw_len, wire_len);
} else {
shared.metrics.chunk_done(raw_len, wire_len, compressed);
}
if complete {
finalize_if_ready(shared, file)?;
}
Ok(())
}
fn finalize_if_ready(shared: &RecvShared, file: &Arc<RecvFile>) -> Result<()> {
if file.finalized.load(Ordering::Acquire) {
return Ok(());
}
if !file.state.lock().is_complete() {
return Ok(());
}
if shared.cfg.verify_hashes && !file.announced.load(Ordering::Acquire) {
return Ok(());
}
if file.finalized.swap(true, Ordering::AcqRel) {
return Ok(());
}
let Some(handle) = shared.writers.take(file.entry.file_id) else {
let h = WriteHandle::open(&file.dest, file.entry.size, false)?;
return commit(shared, file, Arc::new(h));
};
commit(shared, file, handle)
}
fn commit(shared: &RecvShared, file: &Arc<RecvFile>, handle: Arc<WriteHandle>) -> Result<()> {
if shared.cfg.verify_hashes {
if let Some(expected) = *file.expected_hash.lock() {
fill_resumed_hashes(file, &handle)?;
let actual = merkle_root(&file.chunk_hashes.lock());
if actual != expected {
return Err(Error::Integrity {
path: file.entry.path.clone(),
expected: hex(&expected),
actual: hex(&actual),
});
}
}
}
{
let mut st = file.state.lock();
st.checkpoint(&handle)?;
}
handle.commit(
file.entry.mode,
file.entry.mtime,
shared.cfg.preserve_metadata,
)?;
file.state.lock().clear();
shared.pending.fetch_sub(1, Ordering::AcqRel);
shared.metrics.file_done();
tracing::debug!(path = %file.entry.path, "file committed");
let _ = &shared.root;
Ok(())
}
fn fill_resumed_hashes(file: &Arc<RecvFile>, handle: &WriteHandle) -> Result<()> {
if file.preexisting.count() == 0 {
return Ok(());
}
let chunk_size = file.entry.chunk_size as u64;
let mut buf = vec![0u8; chunk_size as usize];
let mut hashes = file.chunk_hashes.lock();
for i in 0..file.entry.chunk_count() {
if !file.preexisting.get(i) {
continue;
}
let offset = i * chunk_size;
let len = chunk_size.min(file.entry.size - offset) as usize;
let n = handle.read_at(offset, &mut buf[..len])?;
if n != len {
return Err(Error::Io(std::io::Error::new(
std::io::ErrorKind::UnexpectedEof,
format!("partial file is short at chunk {i}"),
)));
}
if let Some(slot) = hashes.get_mut(i as usize) {
*slot = *blake3::hash(&buf[..len]).as_bytes();
}
}
Ok(())
}
async fn create_symlinks(root: &Path, links: &[(PathBuf, String)]) {
for (dest, target) in links {
if !symlink_target_is_contained(root, dest, target) {
tracing::warn!(
link = %dest.display(),
target = %target,
"skipping symlink whose target escapes the destination root"
);
continue;
}
if let Some(parent) = dest.parent() {
let _ = tokio::fs::create_dir_all(parent).await;
}
let _ = tokio::fs::remove_file(dest).await;
#[cfg(unix)]
if let Err(e) = tokio::fs::symlink(target, dest).await {
tracing::warn!(link = %dest.display(), error = %e, "could not create symlink");
}
#[cfg(not(unix))]
{
let _ = (dest, target);
tracing::warn!("symlinks are not created on this platform");
}
}
}
fn symlink_target_is_contained(root: &Path, link: &Path, target: &str) -> bool {
if target.is_empty() {
return false;
}
let t = Path::new(target);
if t.is_absolute() {
return false;
}
let Some(parent) = link.parent() else {
return false;
};
let mut resolved = parent.to_path_buf();
for comp in t.components() {
match comp {
std::path::Component::Normal(c) => resolved.push(c),
std::path::Component::CurDir => {}
std::path::Component::ParentDir => {
if !resolved.pop() {
return false;
}
}
_ => return false,
}
}
resolved.starts_with(root)
}
async fn apply_directory_metadata(root: &Path, entries: &[FileEntry], cfg: &Config) {
if !cfg.preserve_metadata {
return;
}
let mut dirs: Vec<_> = entries
.iter()
.filter(|e| e.kind == EntryKind::Directory)
.collect();
dirs.sort_by_key(|e| std::cmp::Reverse(e.path.matches('/').count()));
for e in dirs {
let Ok(dest) = manifest::safe_join(root, &e.path) else {
continue;
};
#[cfg(unix)]
if e.mode != 0 {
use std::os::unix::fs::PermissionsExt;
let _ =
tokio::fs::set_permissions(&dest, std::fs::Permissions::from_mode(e.mode & 0o7777))
.await;
}
let _ = &dest;
}
}
fn index_existing(
dest: &Path,
rel: &str,
chunk_size: u32,
budget: usize,
cache: Option<&mut crate::index::ChunkIndex>,
) -> (Option<crate::io::ChunkReader>, Vec<[u8; 32]>) {
if chunk_size == 0 {
return (None, Vec::new());
}
let Ok(reader) = crate::io::ChunkReader::open(dest) else {
return (None, Vec::new());
};
let len = reader.len();
let meta = std::fs::metadata(dest).ok();
let mtime = meta.as_ref().map(crate::index::mtime_of).unwrap_or(0);
if let Some(cache) = cache {
if let Some(h) = cache.get(rel, len, mtime, chunk_size) {
return (Some(reader), h.to_vec());
}
let hashes = hash_whole_file(&reader, chunk_size, budget);
if !hashes.is_empty() {
cache.insert(rel, len, mtime, chunk_size, hashes.clone());
}
return (Some(reader), hashes);
}
let hashes = hash_whole_file(&reader, chunk_size, budget);
(Some(reader), hashes)
}
fn hash_whole_file(
reader: &crate::io::ChunkReader,
chunk_size: u32,
budget: usize,
) -> Vec<[u8; 32]> {
let len = reader.len();
let chunks = len.div_ceil(chunk_size as u64);
if chunks == 0 || chunks as usize * 32 > budget {
return Vec::new();
}
let mut buf = vec![0u8; chunk_size as usize];
let mut hashes = Vec::with_capacity(chunks as usize);
for i in 0..chunks {
let offset = i * chunk_size as u64;
let want = (chunk_size as u64).min(len - offset) as usize;
match reader.read_at(offset, &mut buf[..want]) {
Ok(n) if n == want => hashes.push(*blake3::hash(&buf[..want]).as_bytes()),
_ => return Vec::new(),
}
}
hashes
}
fn split_index(mut entries: Vec<LocalFileIndex>, max_frame: usize) -> Vec<Vec<LocalFileIndex>> {
let cap = (max_frame / 2).max(1 << 20);
let mut out = Vec::new();
let mut batch = Vec::new();
let mut size = 0usize;
for e in entries.drain(..) {
let cost = 8 + e.hashes.len() * 32;
if size + cost > cap && !batch.is_empty() {
out.push(std::mem::take(&mut batch));
size = 0;
}
size += cost;
batch.push(e);
}
if !batch.is_empty() {
out.push(batch);
}
out
}
async fn drain_to_eof(r: &mut BoxRecv) {
let mut scratch = [0u8; 256];
loop {
match r.read(&mut scratch).await {
Ok(0) | Err(_) => return,
Ok(_) => {}
}
}
}
fn checkpoint_all(shared: &RecvShared) {
for file in shared.files.values() {
if file.finalized.load(Ordering::Acquire) {
continue;
}
let Some(handle) = shared.writers.take(file.entry.file_id) else {
continue;
};
if let Err(e) = file.state.lock().checkpoint(&handle) {
tracing::warn!(path = %file.entry.path, error = %e, "could not persist resume state");
}
}
}
fn hex(b: &[u8]) -> String {
b.iter().map(|x| format!("{x:02x}")).collect()
}
fn map_eof(e: std::io::Error) -> Error {
if e.kind() == std::io::ErrorKind::UnexpectedEof {
Error::Closed("stream ended mid-frame".into())
} else {
Error::Io(e)
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn symlink_containment() {
let root = Path::new("/dest");
let link = Path::new("/dest/sub/link");
assert!(symlink_target_is_contained(root, link, "sibling"));
assert!(symlink_target_is_contained(root, link, "./a/b"));
assert!(symlink_target_is_contained(root, link, "../other"));
assert!(!symlink_target_is_contained(root, link, "/etc"));
assert!(!symlink_target_is_contained(root, link, "../../etc"));
assert!(!symlink_target_is_contained(root, link, "../../../"));
assert!(!symlink_target_is_contained(root, link, ""));
}
}