use crate::codec::compress::{self, FileHint};
use crate::codec::crypto::{Handshake, Role, Sealer, HANDSHAKE_MSG_LEN};
use crate::config::Config;
use crate::error::{Error, Result};
use crate::io::ChunkReader;
use crate::manifest::{self, Manifest, Source};
use crate::metrics::{Metrics, Progress, ProgressFn};
use crate::pool::{BufPool, ObjPool};
use crate::resume::ChunkBitmap;
use crate::transport::{BoxRecv, BoxSend, Transport};
use crate::wire::{self, Control, EntryKind, FrameHeader};
use futures_util::stream::{FuturesOrdered, StreamExt};
use parking_lot::Mutex;
use std::collections::HashMap;
use std::sync::atomic::{AtomicBool, AtomicU32, AtomicU64, Ordering};
use std::sync::Arc;
use tokio::io::AsyncWriteExt;
use tokio::sync::mpsc;
struct ChunkJob {
file_id: u32,
chunk_index: u64,
offset: u64,
len: usize,
last: bool,
hash_only: bool,
progress: Option<Arc<FileProgress>>,
peer_hash: Option<[u8; 32]>,
cached_hash: Option<[u8; 32]>,
reader: Arc<ChunkReader>,
hint: FileHint,
}
struct EncodedChunk {
frame: Vec<u8>,
file_id: u32,
chunk_index: u64,
raw_len: usize,
compressed: bool,
hash: [u8; 32],
skipped: bool,
reused: bool,
compressor_ran: bool,
}
#[inline]
pub fn is_all_zero(buf: &[u8]) -> bool {
const PREFIX: usize = 64;
if buf.len() >= PREFIX && buf[..PREFIX].iter().any(|&b| b != 0) {
return false;
}
let mut words = buf.chunks_exact(8);
let folded = words
.by_ref()
.map(|w| u64::from_ne_bytes(w.try_into().expect("chunks_exact(8) yields 8 bytes")))
.fold(0u64, |acc, w| acc | w);
folded == 0 && words.remainder().iter().all(|&b| b == 0)
}
struct ChunkCursor {
entry_idx: usize,
chunk_index: u64,
reader: Option<Arc<ChunkReader>>,
hint: FileHint,
}
struct CursorShared {
manifest: Arc<Manifest>,
resume: HashMap<u32, ChunkBitmap>,
state: Mutex<ChunkCursor>,
metrics: Metrics,
outstanding: HashMap<u32, Arc<FileProgress>>,
verify_hashes: bool,
peer_index: HashMap<u32, Vec<[u8; 32]>>,
own_index: HashMap<u32, Vec<[u8; 32]>>,
}
const COMPRESSION_PROBE_CHUNKS: u32 = 8;
const COMPRESSION_WIN_RATE: f32 = 0.25;
struct FileProgress {
remaining: AtomicU64,
chunk_hashes: Mutex<Vec<[u8; 32]>>,
announced: AtomicBool,
probe_attempts: AtomicU32,
probe_wins: AtomicU32,
give_up: AtomicBool,
}
impl FileProgress {
fn should_compress(&self) -> bool {
!self.give_up.load(Ordering::Relaxed)
}
fn record_compression(&self, won: bool) {
if self.give_up.load(Ordering::Relaxed) {
return;
}
if won {
self.probe_wins.fetch_add(1, Ordering::Relaxed);
}
let attempts = self.probe_attempts.fetch_add(1, Ordering::Relaxed) + 1;
if attempts < COMPRESSION_PROBE_CHUNKS {
return;
}
let wins = self.probe_wins.load(Ordering::Relaxed);
if (wins as f32) < attempts as f32 * COMPRESSION_WIN_RATE {
self.give_up.store(true, Ordering::Relaxed);
}
}
}
impl CursorShared {
fn next_job(&self) -> Result<Option<ChunkJob>> {
let mut st = self.state.lock();
loop {
let Some(entry) = self.manifest.entries.get(st.entry_idx) else {
return Ok(None);
};
if entry.kind != EntryKind::File || entry.size == 0 {
st.entry_idx += 1;
st.chunk_index = 0;
st.reader = None;
continue;
}
let total = entry.chunk_count();
if st.chunk_index >= total {
st.entry_idx += 1;
st.chunk_index = 0;
st.reader = None;
continue;
}
if st.reader.is_none() {
let path = &self.manifest.local_paths[st.entry_idx];
let r = ChunkReader::open(path)?;
if r.len() != entry.size {
return Err(Error::Io(std::io::Error::other(format!(
"{} changed size during transfer ({} -> {})",
path.display(),
entry.size,
r.len()
))));
}
st.reader = Some(Arc::new(r));
st.hint = FileHint {
known_incompressible: entry.incompressible,
audio: self.manifest.audio.get(st.entry_idx).copied().flatten(),
chunk_offset: 0,
};
}
let idx = st.chunk_index;
st.chunk_index += 1;
let chunk_size = entry.chunk_size as u64;
let offset = idx * chunk_size;
let len = chunk_size.min(entry.size - offset) as usize;
let last = idx + 1 == total;
let already_there = self
.resume
.get(&entry.file_id)
.is_some_and(|bm| bm.get(idx));
if already_there && !self.verify_hashes {
self.metrics.chunk_skipped(len as u64);
if let Some(fp) = self.outstanding.get(&entry.file_id) {
fp.remaining.fetch_sub(1, Ordering::AcqRel);
}
continue;
}
return Ok(Some(ChunkJob {
file_id: entry.file_id,
chunk_index: idx,
offset,
len,
last,
hash_only: already_there,
progress: self.outstanding.get(&entry.file_id).cloned(),
peer_hash: self
.peer_index
.get(&entry.file_id)
.and_then(|h| h.get(idx as usize))
.copied(),
cached_hash: self
.own_index
.get(&entry.file_id)
.and_then(|h| h.get(idx as usize))
.copied(),
reader: st.reader.clone().expect("reader opened above"),
hint: FileHint {
chunk_offset: offset,
..st.hint
},
}));
}
}
}
fn encode_chunk(
job: &ChunkJob,
cfg: &Config,
sealer: &mut Sealer,
pool: &BufPool,
) -> Result<EncodedChunk> {
if let (Some(mine), Some(theirs)) = (job.cached_hash, job.peer_hash) {
if mine == theirs {
let mut frame = pool.take();
frame.resize(wire::FRAME_HEADER_LEN, 0);
let mut flags = wire::flags::REUSE_LOCAL;
if job.last {
flags |= wire::flags::LAST_CHUNK;
}
let header = FrameHeader {
flags,
algorithm: compress::Algorithm::None,
file_id: job.file_id,
chunk_index: job.chunk_index,
epoch: 0,
raw_len: job.len as u32,
payload_len: 0,
};
let head: &mut [u8; wire::FRAME_HEADER_LEN] = (&mut frame[..])
.try_into()
.expect("frame is exactly a header");
header.encode(head);
return Ok(EncodedChunk {
frame,
file_id: job.file_id,
chunk_index: job.chunk_index,
raw_len: job.len,
compressed: false,
hash: mine,
skipped: false,
reused: true,
compressor_ran: false,
});
}
}
const HDR: usize = wire::FRAME_HEADER_LEN;
let mut frame = pool.take();
frame.resize(HDR + job.len, 0);
let n = job.reader.read_at(job.offset, &mut frame[HDR..])?;
if n != job.len {
return Err(Error::Io(std::io::Error::new(
std::io::ErrorKind::UnexpectedEof,
format!(
"short read at offset {} of file {}: wanted {}, got {n}",
job.offset, job.file_id, job.len
),
)));
}
let hash = *blake3::hash(&frame[HDR..]).as_bytes();
if job.hash_only {
let raw_len = frame.len() - HDR;
pool.put(frame);
return Ok(EncodedChunk {
frame: Vec::new(),
file_id: job.file_id,
chunk_index: job.chunk_index,
raw_len,
compressed: false,
hash,
skipped: true,
reused: false,
compressor_ran: false,
});
}
if job.peer_hash == Some(hash) {
let raw_len = frame.len() - HDR;
frame.truncate(HDR);
let mut flags = wire::flags::REUSE_LOCAL;
if job.last {
flags |= wire::flags::LAST_CHUNK;
}
let header = FrameHeader {
flags,
algorithm: compress::Algorithm::None,
file_id: job.file_id,
chunk_index: job.chunk_index,
epoch: 0,
raw_len: raw_len as u32,
payload_len: 0,
};
let head: &mut [u8; wire::FRAME_HEADER_LEN] = (&mut frame[..])
.try_into()
.expect("frame is exactly a header");
header.encode(head);
return Ok(EncodedChunk {
frame,
file_id: job.file_id,
chunk_index: job.chunk_index,
raw_len,
compressed: false,
hash,
skipped: false,
reused: true,
compressor_ran: false,
});
}
if cfg.sparse && is_all_zero(&frame[HDR..]) {
let raw_len = frame.len() - HDR;
frame.truncate(HDR);
let mut flags = wire::flags::ZERO;
if job.last {
flags |= wire::flags::LAST_CHUNK;
}
let header = FrameHeader {
flags,
algorithm: compress::Algorithm::None,
file_id: job.file_id,
chunk_index: job.chunk_index,
epoch: 0,
raw_len: raw_len as u32,
payload_len: 0,
};
let head: &mut [u8; wire::FRAME_HEADER_LEN] = (&mut frame[..])
.try_into()
.expect("frame is exactly a header");
header.encode(head);
return Ok(EncodedChunk {
frame,
file_id: job.file_id,
chunk_index: job.chunk_index,
raw_len,
compressed: false,
hash,
skipped: false,
reused: false,
compressor_ran: false,
});
}
let mut hint = job.hint;
let probing =
!hint.known_incompressible && job.progress.as_ref().is_some_and(|p| !p.should_compress());
if probing {
hint.known_incompressible = true;
}
let enc =
compress::with_codec(|c| c.compress_in_place(&cfg.compression, hint, &mut frame, HDR))?;
if !hint.known_incompressible {
if let Some(p) = job.progress.as_ref() {
p.record_compression(enc.algorithm != compress::Algorithm::None);
}
}
let payload_len = frame.len() - HDR;
let sealed = !sealer.is_passthrough();
let mut flags = 0u8;
if job.last {
flags |= wire::flags::LAST_CHUNK;
}
if sealed {
flags |= wire::flags::SEALED;
}
let header = FrameHeader {
flags,
algorithm: enc.algorithm,
file_id: job.file_id,
chunk_index: job.chunk_index,
epoch: 0,
raw_len: enc.raw_len as u32,
payload_len: payload_len as u32,
};
let (head, body) = frame.split_at_mut(wire::FRAME_HEADER_LEN);
let head: &mut [u8; wire::FRAME_HEADER_LEN] = head.try_into().expect("split at header length");
header.encode(head);
if sealed {
let tag = sealer.seal(job.file_id, job.chunk_index, 0, head, body)?;
frame.extend_from_slice(&tag);
}
Ok(EncodedChunk {
frame,
file_id: job.file_id,
chunk_index: job.chunk_index,
raw_len: enc.raw_len,
compressed: enc.algorithm != compress::Algorithm::None,
hash,
skipped: false,
reused: false,
compressor_ran: !hint.known_incompressible,
})
}
pub async fn send(
transport: Arc<dyn Transport>,
sources: &[Source],
cfg: &Config,
progress: Option<ProgressFn>,
) -> Result<Progress> {
cfg.validate()?;
let metrics = Metrics::new();
let (mut ctl_w, mut ctl_r) = transport.open_bi().await?;
let hs = Handshake::new(Role::Initiator, &cfg.secrecy, cfg.cipher);
ctl_w.write_all(hs.message()).await?;
ctl_w.flush().await?;
let mut peer = [0u8; HANDSHAKE_MSG_LEN];
tokio::time::timeout(cfg.handshake_timeout, read_exact(&mut ctl_r, &mut peer))
.await
.map_err(|_| Error::Handshake("timed out waiting for the peer's handshake".into()))??;
let crypto = Arc::new(hs.finish(&peer)?);
let manifest = Arc::new(manifest::build(sources, cfg).await?);
metrics.set_totals(manifest.file_count() as u64, manifest.total_bytes());
tracing::info!(
files = manifest.file_count(),
bytes = manifest.total_bytes(),
peer = %transport.peer_label(),
"manifest ready"
);
wire::write_control(&mut ctl_w, &Control::Manifest(manifest.entries.clone())).await?;
let mut peer_index: HashMap<u32, Vec<[u8; 32]>> = HashMap::new();
let resume = loop {
match wire::read_control(&mut ctl_r, cfg.max_frame_bytes, cfg.max_manifest_entries).await? {
Control::LocalIndex(entries) => {
for e in entries {
peer_index.entry(e.file_id).or_default().extend(e.hashes);
}
continue;
}
other => break other,
}
};
let resume = match resume {
Control::ResumeState(entries) => {
let mut map = HashMap::new();
for e in entries {
let Some(me) = manifest.entries.iter().find(|m| m.file_id == e.file_id) else {
continue;
};
let bm = ChunkBitmap::from_bytes(e.have, me.chunk_count())?;
if bm.count() > 0 {
map.insert(e.file_id, bm);
}
}
map
}
Control::Abort { reason } => return Err(Error::Closed(reason)),
other => {
return Err(Error::protocol(format!(
"expected ResumeState, got {other:?}"
)))
}
};
let mut own_index: HashMap<u32, Vec<[u8; 32]>> = HashMap::new();
if cfg.delta && cfg.trust_mtime && !peer_index.is_empty() {
if let Some(root) = index_root(sources) {
let cache = crate::index::ChunkIndex::load(&root);
for (i, e) in manifest.entries.iter().enumerate() {
if e.kind != EntryKind::File || !peer_index.contains_key(&e.file_id) {
continue;
}
let Some(path) = manifest.local_paths.get(i) else {
continue;
};
let Ok(meta) = std::fs::metadata(path) else {
continue;
};
if let Some(h) = cache.get(
&e.path,
meta.len(),
crate::index::mtime_of(&meta),
e.chunk_size,
) {
own_index.insert(e.file_id, h.to_vec());
}
}
}
}
let resumed_chunks: u64 = resume.values().map(|b| b.count()).sum();
if resumed_chunks > 0 {
tracing::info!(
chunks = resumed_chunks,
"receiver already holds chunks; skipping them"
);
}
let mut outstanding = HashMap::new();
for e in &manifest.entries {
if e.kind == EntryKind::File {
let total = e.chunk_count();
outstanding.insert(
e.file_id,
Arc::new(FileProgress {
remaining: AtomicU64::new(total),
chunk_hashes: Mutex::new(vec![[0u8; 32]; total as usize]),
announced: AtomicBool::new(false),
probe_attempts: AtomicU32::new(0),
probe_wins: AtomicU32::new(0),
give_up: AtomicBool::new(false),
}),
);
}
}
let shared = Arc::new(CursorShared {
manifest: manifest.clone(),
resume,
state: Mutex::new(ChunkCursor {
entry_idx: 0,
chunk_index: 0,
reader: None,
hint: FileHint::default(),
}),
metrics: metrics.clone(),
outstanding,
verify_hashes: cfg.verify_hashes,
peer_index,
own_index,
});
let (ctl_tx, mut ctl_rx) = mpsc::unbounded_channel::<Control>();
for e in &manifest.entries {
if e.kind != EntryKind::File {
continue;
}
let fp = &shared.outstanding[&e.file_id];
if fp.remaining.load(Ordering::Acquire) == 0 && !fp.announced.swap(true, Ordering::AcqRel) {
let _ = ctl_tx.send(Control::FileComplete {
file_id: e.file_id,
hash: None,
});
metrics.file_done();
}
}
wire::write_control(
&mut ctl_w,
&Control::Start {
streams: cfg.streams as u32,
},
)
.await?;
let ctl_task = tokio::spawn(async move {
let mut failure = None;
while let Some(msg) = ctl_rx.recv().await {
if let Err(e) = wire::write_control(&mut ctl_w, &msg).await {
failure = Some(e);
break;
}
}
if failure.is_none() {
if let Err(e) = wire::write_control(&mut ctl_w, &Control::AllComplete).await {
failure = Some(e);
}
}
(ctl_w, failure)
});
let cpu = Arc::new(
rayon::ThreadPoolBuilder::new()
.num_threads(cfg.workers)
.thread_name(|i| format!("rst-encode-{i}"))
.build()
.map_err(|e| Error::Worker(e.to_string()))?,
);
let pool = BufPool::new(
cfg.streams * cfg.queue_depth * 2 + cfg.workers * 2,
cfg.chunk_size + cfg.chunk_size / 8,
);
let sealers: Arc<ObjPool<Sealer>> = ObjPool::new(cfg.workers + cfg.streams);
let progress_task = progress.map(|f| spawn_progress(metrics.clone(), f));
let mut stream_tasks = Vec::with_capacity(cfg.streams);
for _ in 0..cfg.streams {
let sink = transport.open_uni().await?;
stream_tasks.push(tokio::spawn(run_stream(
sink,
shared.clone(),
cfg.clone(),
sealers.clone(),
crypto.clone(),
cpu.clone(),
pool.clone(),
metrics.clone(),
ctl_tx.clone(),
)));
}
drop(ctl_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_w, ctl_failure) = ctl_task.await.map_err(|e| Error::Worker(e.to_string()))?;
if let Some(e) = first_err.or(ctl_failure) {
return Err(e);
}
loop {
match wire::read_control(&mut ctl_r, cfg.max_frame_bytes, cfg.max_manifest_entries).await {
Ok(Control::AllComplete) => break,
Ok(Control::Abort { reason }) => return Err(Error::Closed(reason)),
Ok(_) => continue,
Err(Error::Closed(m)) => {
return Err(Error::Closed(format!(
"receiver closed before confirming completion: {m}"
)))
}
Err(e) => return Err(e),
}
}
let _ = ctl_w.shutdown().await;
if let Some(p) = progress_task {
p.abort();
}
if cfg.delta && cfg.trust_mtime {
if let Some(root) = index_root(sources) {
let mut cache = crate::index::ChunkIndex::load(&root);
let mut seen = std::collections::HashSet::new();
for (i, e) in manifest.entries.iter().enumerate() {
if e.kind != EntryKind::File || e.size == 0 {
continue;
}
seen.insert(e.path.clone());
let Some(fp) = shared.outstanding.get(&e.file_id) else {
continue;
};
let hashes = fp.chunk_hashes.lock().clone();
if hashes.is_empty() || hashes.iter().all(|h| *h == [0u8; 32]) {
continue;
}
if let Some(meta) = manifest
.local_paths
.get(i)
.and_then(|p| std::fs::metadata(p).ok())
{
cache.insert(
&e.path,
meta.len(),
crate::index::mtime_of(&meta),
e.chunk_size,
hashes,
);
}
}
cache.retain(&seen);
if let Err(err) = cache.save() {
tracing::warn!(error = %err, "could not persist the source chunk index");
}
}
}
let final_snapshot = metrics.snapshot();
tracing::info!(summary = %final_snapshot, "send complete");
Ok(final_snapshot)
}
#[allow(clippy::too_many_arguments)]
async fn run_stream(
mut sink: BoxSend,
shared: Arc<CursorShared>,
cfg: Config,
sealers: Arc<ObjPool<Sealer>>,
crypto: Arc<crate::codec::crypto::SessionCrypto>,
cpu: Arc<rayon::ThreadPool>,
pool: Arc<BufPool>,
metrics: Metrics,
ctl: mpsc::UnboundedSender<Control>,
) -> Result<()> {
let mut inflight = FuturesOrdered::new();
let mut drained = false;
loop {
while !drained && inflight.len() < cfg.queue_depth {
match shared.next_job()? {
Some(job) => {
let (tx, rx) = tokio::sync::oneshot::channel();
let cfg2 = cfg.clone();
let sealers2 = sealers.clone();
let crypto2 = crypto.clone();
let pool2 = pool.clone();
cpu.spawn(move || {
let mut sealer = sealers2.take_or(|| crypto2.sealer());
let r = encode_chunk(&job, &cfg2, &mut sealer, &pool2);
sealers2.put(sealer);
let _ = tx.send(r);
});
inflight.push_back(rx);
}
None => drained = true,
}
}
let Some(res) = inflight.next().await else {
break;
};
let encoded = res.map_err(|_| Error::Worker("encode worker vanished".into()))??;
if encoded.compressor_ran {
metrics.compressor_ran();
}
if encoded.skipped {
metrics.chunk_skipped(encoded.raw_len as u64);
} else {
let is_hole = encoded.frame.len() == wire::FRAME_HEADER_LEN && !encoded.reused;
sink.write_all(&encoded.frame).await?;
let wire_len = encoded.frame.len() as u64;
pool.put(encoded.frame);
if encoded.reused {
metrics.chunk_reused(encoded.raw_len as u64, wire_len);
} else if is_hole {
metrics.chunk_zero(encoded.raw_len as u64, wire_len);
} else {
metrics.chunk_done(encoded.raw_len as u64, wire_len, encoded.compressed);
}
}
if let Some(fp) = shared.outstanding.get(&encoded.file_id) {
if cfg.verify_hashes {
let mut hashes = fp.chunk_hashes.lock();
if let Some(slot) = hashes.get_mut(encoded.chunk_index as usize) {
*slot = encoded.hash;
}
}
let left = fp.remaining.fetch_sub(1, Ordering::AcqRel) - 1;
if left == 0 && !fp.announced.swap(true, Ordering::AcqRel) {
let hash = if cfg.verify_hashes {
Some(merkle_root(&fp.chunk_hashes.lock()))
} else {
None
};
let _ = ctl.send(Control::FileComplete {
file_id: encoded.file_id,
hash,
});
metrics.file_done();
}
}
}
sink.shutdown().await?;
Ok(())
}
fn index_root(sources: &[Source]) -> Option<std::path::PathBuf> {
let first = sources.first()?;
if first.path.is_dir() {
Some(first.path.clone())
} else {
first.path.parent().map(|p| p.to_path_buf())
}
}
pub fn merkle_root(chunk_hashes: &[[u8; 32]]) -> [u8; 32] {
let mut h = blake3::Hasher::new();
for c in chunk_hashes {
h.update(c);
}
*h.finalize().as_bytes()
}
fn spawn_progress(metrics: Metrics, f: ProgressFn) -> tokio::task::JoinHandle<()> {
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(metrics.snapshot());
}
})
}
async fn read_exact(r: &mut BoxRecv, buf: &mut [u8]) -> Result<()> {
use tokio::io::AsyncReadExt;
r.read_exact(buf).await.map_err(|e| {
if e.kind() == std::io::ErrorKind::UnexpectedEof {
Error::Closed("peer closed during handshake".into())
} else {
Error::Io(e)
}
})?;
Ok(())
}