1use crate::codec::compress;
15use crate::codec::crypto::{Handshake, Role, Sealer, HANDSHAKE_MSG_LEN, TAG_LEN};
16use crate::config::Config;
17use crate::error::{Error, Result};
18use crate::io::{FileWriters, WriteHandle};
19use crate::manifest;
20use crate::metrics::{Metrics, Progress, ProgressFn};
21use crate::pool::{BufPool, ObjPool};
22use crate::resume::ResumeState;
23use crate::send::merkle_root;
24use crate::transport::{BoxRecv, Transport};
25use crate::wire::{self, Control, EntryKind, FileEntry, FrameHeader, LocalFileIndex, ResumeEntry};
26use futures_util::stream::{FuturesOrdered, StreamExt};
27use parking_lot::Mutex;
28use std::collections::HashMap;
29use std::path::{Path, PathBuf};
30use std::sync::atomic::{AtomicBool, AtomicU64, Ordering};
31use std::sync::Arc;
32use tokio::io::{AsyncReadExt, AsyncWriteExt};
33use tokio::sync::mpsc;
34
35const CHECKPOINT_INTERVAL: u64 = 256;
40
41const CHECKPOINT_MAX_AGE: std::time::Duration = std::time::Duration::from_secs(5);
46
47const FAREWELL_TIMEOUT: std::time::Duration = std::time::Duration::from_secs(20);
51
52struct RecvFile {
53 entry: FileEntry,
54 dest: PathBuf,
55 preexisting: crate::resume::ChunkBitmap,
58 fresh_part: bool,
66 state: Mutex<ResumeState>,
67 chunk_hashes: Mutex<Vec<[u8; 32]>>,
68 expected_hash: Mutex<Option<[u8; 32]>>,
69 existing: Option<crate::io::ChunkReader>,
73 announced: AtomicBool,
76 finalized: AtomicBool,
77 since_checkpoint: AtomicU64,
78 last_checkpoint: Mutex<std::time::Instant>,
79}
80
81struct RecvShared {
82 cfg: Config,
83 root: PathBuf,
84 files: HashMap<u32, Arc<RecvFile>>,
85 writers: FileWriters,
86 metrics: Metrics,
87 pending: AtomicU64,
89}
90
91pub async fn receive(
93 transport: Arc<dyn Transport>,
94 dest_root: impl AsRef<Path>,
95 cfg: &Config,
96 progress: Option<ProgressFn>,
97) -> Result<Progress> {
98 cfg.validate()?;
99 let root = dest_root.as_ref().to_path_buf();
100 tokio::fs::create_dir_all(&root).await?;
101 let root = tokio::fs::canonicalize(&root).await.unwrap_or(root);
104 let metrics = Metrics::new();
105
106 let (mut ctl_w, mut ctl_r) = transport.accept_bi().await?;
107
108 let hs = Handshake::new(Role::Responder, &cfg.secrecy, cfg.cipher);
109 let mut peer = [0u8; HANDSHAKE_MSG_LEN];
110 tokio::time::timeout(cfg.handshake_timeout, ctl_r.read_exact(&mut peer))
111 .await
112 .map_err(|_| Error::Handshake("timed out waiting for the peer's handshake".into()))?
113 .map_err(map_eof)?;
114 ctl_w.write_all(hs.message()).await?;
115 ctl_w.flush().await?;
116 let crypto = Arc::new(hs.finish(&peer)?);
117
118 let entries = match wire::read_control(
120 &mut ctl_r,
121 cfg.max_frame_bytes,
122 cfg.max_manifest_entries,
123 )
124 .await?
125 {
126 Control::Manifest(e) => e,
127 Control::Abort { reason } => return Err(Error::Closed(reason)),
128 other => return Err(Error::protocol(format!("expected Manifest, got {other:?}"))),
129 };
130 manifest::validate(&entries, cfg)?;
131
132 let mut symlinks = Vec::new();
135 let mut files = HashMap::new();
136 let mut resume_reply = Vec::new();
137 let mut local_index: Vec<LocalFileIndex> = Vec::new();
138 let mut cache = (cfg.delta && cfg.trust_mtime).then(|| crate::index::ChunkIndex::load(&root));
141 #[allow(unused_mut)]
142 let mut seen: std::collections::HashSet<String> = std::collections::HashSet::new();
143 let mut total_bytes = 0u64;
144 let mut file_count = 0u64;
145
146 for entry in &entries {
147 match entry.kind {
148 EntryKind::Directory => {
149 let dest = manifest::safe_join(&root, &entry.path)?;
150 tokio::fs::create_dir_all(&dest).await?;
151 }
152 EntryKind::Symlink => {
153 let (rel, target) = manifest::split_symlink(&entry.path)?;
154 let dest = manifest::safe_join(&root, rel)?;
155 symlinks.push((dest, target.to_string()));
156 }
157 EntryKind::File => {
158 let dest = manifest::safe_join(&root, &entry.path)?;
159 if let Some(parent) = dest.parent() {
160 tokio::fs::create_dir_all(parent).await?;
161 }
162 total_bytes += entry.size;
163 file_count += 1;
164
165 let part = crate::io::part_path_for(&dest);
166 let fresh_part = !part.exists();
167 if fresh_part {
170 ResumeState::load_or_new(&dest, entry.size, entry.chunk_size).clear();
171 }
172
173 if !cfg.resume {
176 ResumeState::load_or_new(&dest, entry.size, entry.chunk_size).clear();
177 }
178 let state = ResumeState::load_or_new(&dest, entry.size, entry.chunk_size);
179 let (existing, local_hashes) = if cfg.delta && dest.exists() {
184 index_existing(
185 &dest,
186 &entry.path,
187 entry.chunk_size,
188 cfg.chunk_hash_budget,
189 cache.as_mut(),
190 )
191 } else {
192 (None, Vec::new())
193 };
194 if !local_hashes.is_empty() {
195 local_index.push(LocalFileIndex {
196 file_id: entry.file_id,
197 hashes: local_hashes,
198 });
199 }
200
201 let preexisting = state.bitmap().clone();
202 if preexisting.count() > 0 {
203 resume_reply.push(ResumeEntry {
204 file_id: entry.file_id,
205 have: preexisting.as_bytes().to_vec(),
206 });
207 metrics.chunk_skipped(0);
208 }
209
210 let n = entry.chunk_count() as usize;
211 files.insert(
212 entry.file_id,
213 Arc::new(RecvFile {
214 entry: entry.clone(),
215 dest,
216 preexisting,
217 existing,
218 fresh_part,
219 state: Mutex::new(state),
220 chunk_hashes: Mutex::new(vec![[0u8; 32]; n]),
221 expected_hash: Mutex::new(None),
222 announced: AtomicBool::new(false),
223 finalized: AtomicBool::new(false),
224 since_checkpoint: AtomicU64::new(0),
225 last_checkpoint: Mutex::new(std::time::Instant::now()),
226 }),
227 );
228 }
229 }
230 }
231
232 metrics.set_totals(file_count, total_bytes);
233 let shared = Arc::new(RecvShared {
234 cfg: cfg.clone(),
235 root: root.clone(),
236 files,
237 writers: FileWriters::new(),
238 metrics: metrics.clone(),
239 pending: AtomicU64::new(file_count),
240 });
241
242 if !local_index.is_empty() {
243 let reusable: usize = local_index.iter().map(|e| e.hashes.len()).sum();
244 tracing::info!(
245 files = local_index.len(),
246 chunks = reusable,
247 "offering existing blocks for reuse"
248 );
249 for batch in split_index(local_index, cfg.max_frame_bytes) {
252 wire::write_control(&mut ctl_w, &Control::LocalIndex(batch)).await?;
253 }
254 }
255 wire::write_control(&mut ctl_w, &Control::ResumeState(resume_reply)).await?;
256
257 for f in shared.files.values() {
259 if f.entry.size == 0 {
260 let h = WriteHandle::open(&f.dest, 0, false)?;
261 h.commit(f.entry.mode, f.entry.mtime, cfg.preserve_metadata)?;
262 f.finalized.store(true, Ordering::Release);
263 shared.pending.fetch_sub(1, Ordering::AcqRel);
264 metrics.file_done();
265 }
266 }
267
268 let stream_count = match wire::read_control(
270 &mut ctl_r,
271 cfg.max_frame_bytes,
272 cfg.max_manifest_entries,
273 )
274 .await?
275 {
276 Control::Start { streams } => streams as usize,
277 Control::Abort { reason } => return Err(Error::Closed(reason)),
278 other => return Err(Error::protocol(format!("expected Start, got {other:?}"))),
279 };
280 if stream_count == 0 || stream_count > 1024 {
281 return Err(Error::protocol(format!(
282 "sender asked for {stream_count} data streams, which is outside the accepted range"
283 )));
284 }
285
286 let cpu = Arc::new(
287 rayon::ThreadPoolBuilder::new()
288 .num_threads(cfg.workers)
289 .thread_name(|i| format!("rst-decode-{i}"))
290 .build()
291 .map_err(|e| Error::Worker(e.to_string()))?,
292 );
293 let pool = BufPool::new(
294 stream_count * cfg.queue_depth * 2 + cfg.workers * 2,
295 cfg.chunk_size + cfg.chunk_size / 8,
296 );
297 let openers: Arc<ObjPool<Sealer>> = ObjPool::new(cfg.workers + stream_count);
298 let progress_task = progress.map(|f| {
299 let m = metrics.clone();
300 tokio::spawn(async move {
301 let mut tick = tokio::time::interval(std::time::Duration::from_millis(500));
302 tick.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Skip);
303 loop {
304 tick.tick().await;
305 f(m.snapshot());
306 }
307 })
308 });
309
310 let (done_tx, mut done_rx) = mpsc::unbounded_channel::<Result<()>>();
314 let ctl_task = {
315 let shared = shared.clone();
316 let done_tx = done_tx.clone();
317 tokio::spawn(async move {
318 let r = run_control(&mut ctl_r, shared).await;
319 let _ = done_tx.send(r);
320 ctl_r
321 })
322 };
323
324 let mut stream_tasks = Vec::with_capacity(stream_count);
325 for _ in 0..stream_count {
326 let src = transport.accept_uni().await?;
327 stream_tasks.push(tokio::spawn(run_stream(
328 src,
329 shared.clone(),
330 openers.clone(),
331 crypto.clone(),
332 cpu.clone(),
333 pool.clone(),
334 )));
335 }
336 drop(done_tx);
337
338 let mut first_err = None;
339 for t in stream_tasks {
340 match t.await {
341 Ok(Ok(())) => {}
342 Ok(Err(e)) => {
343 tracing::error!(error = %e, "data stream failed");
344 first_err.get_or_insert(e);
345 }
346 Err(e) => {
347 first_err.get_or_insert(Error::Worker(e.to_string()));
348 }
349 }
350 }
351
352 let mut ctl_r = ctl_task.await.map_err(|e| Error::Worker(e.to_string()))?;
354 while let Some(r) = done_rx.recv().await {
355 if let Err(e) = r {
356 first_err.get_or_insert(e);
357 }
358 }
359
360 if let Some(e) = first_err {
361 checkpoint_all(&shared);
365 return Err(e);
366 }
367
368 let left = shared.pending.load(Ordering::Acquire);
369 if left != 0 {
370 checkpoint_all(&shared);
373 let names: Vec<_> = shared
374 .files
375 .values()
376 .filter(|f| !f.finalized.load(Ordering::Acquire))
377 .take(5)
378 .map(|f| f.entry.path.clone())
379 .collect();
380 return Err(Error::Protocol(format!(
381 "sender finished with {left} files incomplete (e.g. {names:?})"
382 )));
383 }
384
385 if let Some(mut cache) = cache.take() {
388 for f in shared.files.values() {
389 if !f.finalized.load(Ordering::Acquire) {
390 continue;
391 }
392 seen.insert(f.entry.path.clone());
393 if let Ok(meta) = std::fs::metadata(&f.dest) {
394 let hashes = f.chunk_hashes.lock().clone();
395 if !hashes.is_empty() && hashes.iter().any(|h| *h != [0u8; 32]) {
396 cache.insert(
397 &f.entry.path,
398 meta.len(),
399 crate::index::mtime_of(&meta),
400 f.entry.chunk_size,
401 hashes,
402 );
403 }
404 }
405 }
406 cache.retain(&seen);
407 if let Err(e) = cache.save() {
408 tracing::warn!(error = %e, "could not persist the chunk index");
409 }
410 }
411
412 create_symlinks(&root, &symlinks).await;
413 apply_directory_metadata(&root, &entries, cfg).await;
414
415 wire::write_control(&mut ctl_w, &Control::AllComplete).await?;
423 match tokio::time::timeout(FAREWELL_TIMEOUT, drain_to_eof(&mut ctl_r)).await {
424 Ok(_) => {}
425 Err(_) => tracing::debug!("sender did not close its control stream in time"),
426 }
427 let _ = ctl_w.shutdown().await;
428
429 if let Some(p) = progress_task {
430 p.abort();
431 }
432 let snapshot = metrics.snapshot();
433 tracing::info!(summary = %snapshot, "receive complete");
434 Ok(snapshot)
435}
436
437async fn run_control(ctl_r: &mut BoxRecv, shared: Arc<RecvShared>) -> Result<()> {
438 loop {
439 let msg = match wire::read_control(
440 ctl_r,
441 shared.cfg.max_frame_bytes,
442 shared.cfg.max_manifest_entries,
443 )
444 .await
445 {
446 Ok(m) => m,
447 Err(Error::Closed(_)) => return Ok(()),
449 Err(e) => return Err(e),
450 };
451 match msg {
452 Control::FileComplete { file_id, hash } => {
453 if let Some(f) = shared.files.get(&file_id) {
454 *f.expected_hash.lock() = hash;
455 f.announced.store(true, Ordering::Release);
456 finalize_if_ready(&shared, f)?;
459 }
460 }
461 Control::AllComplete => return Ok(()),
462 Control::Abort { reason } => return Err(Error::Closed(reason)),
463 other => {
464 tracing::debug!(?other, "ignoring unexpected control message");
465 }
466 }
467 }
468}
469
470async fn run_stream(
472 mut src: BoxRecv,
473 shared: Arc<RecvShared>,
474 openers: Arc<ObjPool<Sealer>>,
475 crypto: Arc<crate::codec::crypto::SessionCrypto>,
476 cpu: Arc<rayon::ThreadPool>,
477 pool: Arc<BufPool>,
478) -> Result<()> {
479 let mut inflight: FuturesOrdered<tokio::sync::oneshot::Receiver<Result<()>>> =
480 FuturesOrdered::new();
481 let max_frame = shared.cfg.max_frame_bytes;
482
483 loop {
484 while inflight.len() >= shared.cfg.queue_depth {
487 if let Some(res) = inflight.next().await {
488 res.map_err(|_| Error::Worker("decode worker vanished".into()))??;
489 }
490 }
491
492 let mut head = [0u8; wire::FRAME_HEADER_LEN];
493 match src.read_exact(&mut head).await {
494 Ok(_) => {}
495 Err(e) if e.kind() == std::io::ErrorKind::UnexpectedEof => break,
496 Err(e) => return Err(Error::Io(e)),
497 }
498 let header = FrameHeader::decode(&head, max_frame)?;
499
500 let body_len = header.wire_payload_len();
501 if body_len > max_frame {
502 return Err(Error::FrameTooLarge {
503 got: body_len,
504 limit: max_frame,
505 });
506 }
507 let mut body = pool.take();
508 body.resize(body_len, 0);
509 src.read_exact(&mut body).await.map_err(map_eof)?;
510
511 let (tx, rx) = tokio::sync::oneshot::channel();
512 let shared2 = shared.clone();
513 let openers2 = openers.clone();
514 let crypto2 = crypto.clone();
515 let pool2 = pool.clone();
516 cpu.spawn(move || {
517 let mut opener = openers2.take_or(|| crypto2.opener());
520 let r = decode_and_write(header, head, body, &shared2, &mut opener, &pool2);
521 openers2.put(opener);
522 let _ = tx.send(r);
523 });
524 inflight.push_back(rx);
525 }
526
527 while let Some(res) = inflight.next().await {
528 res.map_err(|_| Error::Worker("decode worker vanished".into()))??;
529 }
530 Ok(())
531}
532
533fn decode_and_write(
535 header: FrameHeader,
536 head: [u8; wire::FRAME_HEADER_LEN],
537 mut body: Vec<u8>,
538 shared: &RecvShared,
539 opener: &mut Sealer,
540 pool: &BufPool,
541) -> Result<()> {
542 let file = shared.files.get(&header.file_id).ok_or_else(|| {
543 Error::protocol(format!(
544 "frame references file_id {} which is not in the manifest",
545 header.file_id
546 ))
547 })?;
548
549 let total_chunks = file.entry.chunk_count();
550 if header.chunk_index >= total_chunks {
551 return Err(Error::protocol(format!(
552 "chunk {} is past the end of file {} ({} chunks)",
553 header.chunk_index, header.file_id, total_chunks
554 )));
555 }
556
557 let chunk_size_u64 = file.entry.chunk_size as u64;
558 let offset = header.chunk_index * chunk_size_u64;
559 let expect_len = chunk_size_u64.min(file.entry.size - offset) as usize;
560
561 if header.reuses_local() {
564 if header.payload_len != 0 || !body.is_empty() {
565 return Err(Error::protocol("reuse frame carries a payload"));
566 }
567 if header.raw_len as usize != expect_len {
568 return Err(Error::protocol(format!(
569 "reuse chunk {} of file {} declares {} bytes, expected {expect_len}",
570 header.chunk_index, header.file_id, header.raw_len
571 )));
572 }
573 pool.put(body);
574 let Some(existing) = file.existing.as_ref() else {
575 return Err(Error::protocol(
576 "peer asked us to reuse a block from a file we never offered",
577 ));
578 };
579
580 let mut buf = pool.take();
581 buf.resize(expect_len, 0);
582 let n = existing.read_at(offset, &mut buf)?;
583 if n != expect_len {
584 return Err(Error::Io(std::io::Error::new(
585 std::io::ErrorKind::UnexpectedEof,
586 format!("existing copy is short at chunk {}", header.chunk_index),
587 )));
588 }
589
590 let handle = shared.writers.get_or_open(header.file_id, || {
591 WriteHandle::open_cloned(&file.dest, file.entry.size, shared.cfg.preallocate)
592 })?;
593 if !handle.matches_at(offset, &buf)? {
596 handle.write_at(offset, &buf)?;
597 }
598
599 if shared.cfg.verify_hashes {
600 let h = *blake3::hash(&buf).as_bytes();
601 let mut hashes = file.chunk_hashes.lock();
602 if let Some(slot) = hashes.get_mut(header.chunk_index as usize) {
603 *slot = h;
604 }
605 }
606 pool.put(buf);
607 return record_chunk(
608 shared,
609 file,
610 &handle,
611 header,
612 ChunkOutcome {
613 raw_len: expect_len as u64,
614 wire_len: wire::FRAME_HEADER_LEN as u64,
615 compressed: false,
616 was_hole: false,
617 was_reuse: true,
618 },
619 );
620 }
621
622 if header.is_zero() {
624 if header.payload_len != 0 || !body.is_empty() {
625 return Err(Error::protocol("zero-chunk frame carries a payload"));
626 }
627 if header.raw_len as usize != expect_len {
628 return Err(Error::protocol(format!(
629 "zero chunk {} of file {} declares {} bytes, expected {expect_len}",
630 header.chunk_index, header.file_id, header.raw_len
631 )));
632 }
633 pool.put(body);
634 let handle = shared.writers.get_or_open(header.file_id, || {
635 WriteHandle::open(&file.dest, file.entry.size, shared.cfg.preallocate)
636 })?;
637 if !file.fresh_part {
641 handle.write_zeros_at(offset, expect_len)?;
642 }
643 if shared.cfg.verify_hashes {
644 let zeros = vec![0u8; expect_len];
645 let h = *blake3::hash(&zeros).as_bytes();
646 let mut hashes = file.chunk_hashes.lock();
647 if let Some(slot) = hashes.get_mut(header.chunk_index as usize) {
648 *slot = h;
649 }
650 }
651 return record_chunk(
652 shared,
653 file,
654 &handle,
655 header,
656 ChunkOutcome {
657 raw_len: expect_len as u64,
658 wire_len: wire::FRAME_HEADER_LEN as u64,
659 compressed: false,
660 was_hole: true,
661 was_reuse: false,
662 },
663 );
664 }
665
666 if header.sealed() {
669 if body.len() < TAG_LEN {
670 return Err(Error::protocol("sealed frame is shorter than its tag"));
671 }
672 let split = body.len() - TAG_LEN;
673 let mut tag = [0u8; TAG_LEN];
674 tag.copy_from_slice(&body[split..]);
675 body.truncate(split);
676 opener.open(
677 header.file_id,
678 header.chunk_index,
679 header.epoch,
680 &head,
681 &mut body,
682 &tag,
683 )?;
684 } else if !opener.is_passthrough() {
685 return Err(Error::protocol(
687 "peer sent an unsealed frame on an encrypted session",
688 ));
689 }
690
691 if header.raw_len as usize != expect_len {
694 return Err(Error::protocol(format!(
695 "chunk {} of file {} declares {} plaintext bytes, expected {expect_len}",
696 header.chunk_index, header.file_id, header.raw_len
697 )));
698 }
699
700 let mut plain = pool.take();
701 compress::decompress_into(header.algorithm, header.raw_len as usize, &body, &mut plain)?;
702 let wire_len = (wire::FRAME_HEADER_LEN + body.len()) as u64;
703 pool.put(body);
704
705 let handle = shared.writers.get_or_open(header.file_id, || {
706 WriteHandle::open(&file.dest, file.entry.size, shared.cfg.preallocate)
707 })?;
708
709 handle.write_at(offset, &plain)?;
710
711 if shared.cfg.verify_hashes {
712 let h = *blake3::hash(&plain).as_bytes();
713 let mut hashes = file.chunk_hashes.lock();
714 if let Some(slot) = hashes.get_mut(header.chunk_index as usize) {
715 *slot = h;
716 }
717 }
718 let raw_len = plain.len() as u64;
719 pool.put(plain);
720
721 record_chunk(
722 shared,
723 file,
724 &handle,
725 header,
726 ChunkOutcome {
727 raw_len,
728 wire_len,
729 compressed: header.algorithm != compress::Algorithm::None,
730 was_hole: false,
731 was_reuse: false,
732 },
733 )
734}
735
736struct ChunkOutcome {
738 raw_len: u64,
739 wire_len: u64,
740 compressed: bool,
741 was_hole: bool,
743 was_reuse: bool,
745}
746
747fn record_chunk(
750 shared: &RecvShared,
751 file: &Arc<RecvFile>,
752 handle: &WriteHandle,
753 header: FrameHeader,
754 outcome: ChunkOutcome,
755) -> Result<()> {
756 let ChunkOutcome {
757 raw_len,
758 wire_len,
759 compressed,
760 was_hole,
761 was_reuse,
762 } = outcome;
763 let complete = {
764 let mut st = file.state.lock();
765 st.record(header.chunk_index);
766 if shared.cfg.resume {
771 let by_count =
772 file.since_checkpoint.fetch_add(1, Ordering::AcqRel) + 1 >= CHECKPOINT_INTERVAL;
773 let by_age = file.last_checkpoint.lock().elapsed() >= CHECKPOINT_MAX_AGE;
774 if by_count || by_age {
775 file.since_checkpoint.store(0, Ordering::Release);
776 *file.last_checkpoint.lock() = std::time::Instant::now();
777 st.checkpoint(handle)?;
778 }
779 }
780 st.is_complete()
781 };
782
783 if was_reuse {
784 shared.metrics.chunk_reused(raw_len, wire_len);
785 } else if was_hole {
786 shared.metrics.chunk_zero(raw_len, wire_len);
787 } else {
788 shared.metrics.chunk_done(raw_len, wire_len, compressed);
789 }
790
791 if complete {
792 finalize_if_ready(shared, file)?;
793 }
794 Ok(())
795}
796
797fn finalize_if_ready(shared: &RecvShared, file: &Arc<RecvFile>) -> Result<()> {
803 if file.finalized.load(Ordering::Acquire) {
804 return Ok(());
805 }
806 if !file.state.lock().is_complete() {
807 return Ok(());
808 }
809 if shared.cfg.verify_hashes && !file.announced.load(Ordering::Acquire) {
810 return Ok(());
812 }
813 if file.finalized.swap(true, Ordering::AcqRel) {
815 return Ok(());
816 }
817
818 let Some(handle) = shared.writers.take(file.entry.file_id) else {
819 let h = WriteHandle::open(&file.dest, file.entry.size, false)?;
821 return commit(shared, file, Arc::new(h));
822 };
823 commit(shared, file, handle)
824}
825
826fn commit(shared: &RecvShared, file: &Arc<RecvFile>, handle: Arc<WriteHandle>) -> Result<()> {
827 if shared.cfg.verify_hashes {
828 if let Some(expected) = *file.expected_hash.lock() {
829 fill_resumed_hashes(file, &handle)?;
832 let actual = merkle_root(&file.chunk_hashes.lock());
833 if actual != expected {
834 return Err(Error::Integrity {
835 path: file.entry.path.clone(),
836 expected: hex(&expected),
837 actual: hex(&actual),
838 });
839 }
840 }
841 }
842
843 {
844 let mut st = file.state.lock();
845 st.checkpoint(&handle)?;
846 }
847
848 handle.commit(
849 file.entry.mode,
850 file.entry.mtime,
851 shared.cfg.preserve_metadata,
852 )?;
853 file.state.lock().clear();
854
855 shared.pending.fetch_sub(1, Ordering::AcqRel);
856 shared.metrics.file_done();
857 tracing::debug!(path = %file.entry.path, "file committed");
858 let _ = &shared.root;
859 Ok(())
860}
861
862fn fill_resumed_hashes(file: &Arc<RecvFile>, handle: &WriteHandle) -> Result<()> {
864 if file.preexisting.count() == 0 {
865 return Ok(());
866 }
867 let chunk_size = file.entry.chunk_size as u64;
868 let mut buf = vec![0u8; chunk_size as usize];
869 let mut hashes = file.chunk_hashes.lock();
870 for i in 0..file.entry.chunk_count() {
871 if !file.preexisting.get(i) {
872 continue;
873 }
874 let offset = i * chunk_size;
875 let len = chunk_size.min(file.entry.size - offset) as usize;
876 let n = handle.read_at(offset, &mut buf[..len])?;
877 if n != len {
878 return Err(Error::Io(std::io::Error::new(
879 std::io::ErrorKind::UnexpectedEof,
880 format!("partial file is short at chunk {i}"),
881 )));
882 }
883 if let Some(slot) = hashes.get_mut(i as usize) {
884 *slot = *blake3::hash(&buf[..len]).as_bytes();
885 }
886 }
887 Ok(())
888}
889
890async fn create_symlinks(root: &Path, links: &[(PathBuf, String)]) {
898 for (dest, target) in links {
899 if !symlink_target_is_contained(root, dest, target) {
900 tracing::warn!(
901 link = %dest.display(),
902 target = %target,
903 "skipping symlink whose target escapes the destination root"
904 );
905 continue;
906 }
907 if let Some(parent) = dest.parent() {
908 let _ = tokio::fs::create_dir_all(parent).await;
909 }
910 let _ = tokio::fs::remove_file(dest).await;
911 #[cfg(unix)]
912 if let Err(e) = tokio::fs::symlink(target, dest).await {
913 tracing::warn!(link = %dest.display(), error = %e, "could not create symlink");
914 }
915 #[cfg(not(unix))]
916 {
917 let _ = (dest, target);
918 tracing::warn!("symlinks are not created on this platform");
919 }
920 }
921}
922
923fn symlink_target_is_contained(root: &Path, link: &Path, target: &str) -> bool {
927 if target.is_empty() {
928 return false;
929 }
930 let t = Path::new(target);
931 if t.is_absolute() {
932 return false;
933 }
934 let Some(parent) = link.parent() else {
935 return false;
936 };
937 let mut resolved = parent.to_path_buf();
938 for comp in t.components() {
939 match comp {
940 std::path::Component::Normal(c) => resolved.push(c),
941 std::path::Component::CurDir => {}
942 std::path::Component::ParentDir => {
943 if !resolved.pop() {
944 return false;
945 }
946 }
947 _ => return false,
948 }
949 }
950 resolved.starts_with(root)
951}
952
953async fn apply_directory_metadata(root: &Path, entries: &[FileEntry], cfg: &Config) {
954 if !cfg.preserve_metadata {
955 return;
956 }
957 let mut dirs: Vec<_> = entries
959 .iter()
960 .filter(|e| e.kind == EntryKind::Directory)
961 .collect();
962 dirs.sort_by_key(|e| std::cmp::Reverse(e.path.matches('/').count()));
963 for e in dirs {
964 let Ok(dest) = manifest::safe_join(root, &e.path) else {
965 continue;
966 };
967 #[cfg(unix)]
968 if e.mode != 0 {
969 use std::os::unix::fs::PermissionsExt;
970 let _ =
971 tokio::fs::set_permissions(&dest, std::fs::Permissions::from_mode(e.mode & 0o7777))
972 .await;
973 }
974 let _ = &dest;
975 }
976}
977
978fn index_existing(
984 dest: &Path,
985 rel: &str,
986 chunk_size: u32,
987 budget: usize,
988 cache: Option<&mut crate::index::ChunkIndex>,
989) -> (Option<crate::io::ChunkReader>, Vec<[u8; 32]>) {
990 if chunk_size == 0 {
991 return (None, Vec::new());
992 }
993 let Ok(reader) = crate::io::ChunkReader::open(dest) else {
994 return (None, Vec::new());
995 };
996 let len = reader.len();
997
998 let meta = std::fs::metadata(dest).ok();
1001 let mtime = meta.as_ref().map(crate::index::mtime_of).unwrap_or(0);
1002 if let Some(cache) = cache {
1003 if let Some(h) = cache.get(rel, len, mtime, chunk_size) {
1004 return (Some(reader), h.to_vec());
1005 }
1006 let hashes = hash_whole_file(&reader, chunk_size, budget);
1007 if !hashes.is_empty() {
1008 cache.insert(rel, len, mtime, chunk_size, hashes.clone());
1009 }
1010 return (Some(reader), hashes);
1011 }
1012 let hashes = hash_whole_file(&reader, chunk_size, budget);
1013 (Some(reader), hashes)
1014}
1015
1016fn hash_whole_file(
1017 reader: &crate::io::ChunkReader,
1018 chunk_size: u32,
1019 budget: usize,
1020) -> Vec<[u8; 32]> {
1021 let len = reader.len();
1022 let chunks = len.div_ceil(chunk_size as u64);
1023 if chunks == 0 || chunks as usize * 32 > budget {
1026 return Vec::new();
1027 }
1028
1029 let mut buf = vec![0u8; chunk_size as usize];
1030 let mut hashes = Vec::with_capacity(chunks as usize);
1031 for i in 0..chunks {
1032 let offset = i * chunk_size as u64;
1033 let want = (chunk_size as u64).min(len - offset) as usize;
1034 match reader.read_at(offset, &mut buf[..want]) {
1035 Ok(n) if n == want => hashes.push(*blake3::hash(&buf[..want]).as_bytes()),
1036 _ => return Vec::new(),
1037 }
1038 }
1039 hashes
1040}
1041
1042fn split_index(mut entries: Vec<LocalFileIndex>, max_frame: usize) -> Vec<Vec<LocalFileIndex>> {
1044 let cap = (max_frame / 2).max(1 << 20);
1045 let mut out = Vec::new();
1046 let mut batch = Vec::new();
1047 let mut size = 0usize;
1048 for e in entries.drain(..) {
1049 let cost = 8 + e.hashes.len() * 32;
1050 if size + cost > cap && !batch.is_empty() {
1051 out.push(std::mem::take(&mut batch));
1052 size = 0;
1053 }
1054 size += cost;
1055 batch.push(e);
1056 }
1057 if !batch.is_empty() {
1058 out.push(batch);
1059 }
1060 out
1061}
1062
1063async fn drain_to_eof(r: &mut BoxRecv) {
1065 let mut scratch = [0u8; 256];
1066 loop {
1067 match r.read(&mut scratch).await {
1068 Ok(0) | Err(_) => return,
1069 Ok(_) => {}
1070 }
1071 }
1072}
1073
1074fn checkpoint_all(shared: &RecvShared) {
1078 for file in shared.files.values() {
1079 if file.finalized.load(Ordering::Acquire) {
1080 continue;
1081 }
1082 let Some(handle) = shared.writers.take(file.entry.file_id) else {
1083 continue;
1084 };
1085 if let Err(e) = file.state.lock().checkpoint(&handle) {
1086 tracing::warn!(path = %file.entry.path, error = %e, "could not persist resume state");
1087 }
1088 }
1089}
1090
1091fn hex(b: &[u8]) -> String {
1092 b.iter().map(|x| format!("{x:02x}")).collect()
1093}
1094
1095fn map_eof(e: std::io::Error) -> Error {
1096 if e.kind() == std::io::ErrorKind::UnexpectedEof {
1097 Error::Closed("stream ended mid-frame".into())
1098 } else {
1099 Error::Io(e)
1100 }
1101}
1102
1103#[cfg(test)]
1104mod tests {
1105 use super::*;
1106
1107 #[test]
1108 fn symlink_containment() {
1109 let root = Path::new("/dest");
1110 let link = Path::new("/dest/sub/link");
1111 assert!(symlink_target_is_contained(root, link, "sibling"));
1112 assert!(symlink_target_is_contained(root, link, "./a/b"));
1113 assert!(symlink_target_is_contained(root, link, "../other"));
1114
1115 assert!(!symlink_target_is_contained(root, link, "/etc"));
1117 assert!(!symlink_target_is_contained(root, link, "../../etc"));
1118 assert!(!symlink_target_is_contained(root, link, "../../../"));
1119 assert!(!symlink_target_is_contained(root, link, ""));
1120 }
1121}