1use std::sync::Arc;
11
12use sha2::{Digest, Sha256};
13use tokio::io::{AsyncReadExt, AsyncWriteExt};
14use tokio::sync::{broadcast, mpsc, oneshot};
15use tracing::{error, info, warn};
16
17use crate::network::NetworkProvider;
18use crate::node::Node;
19
20use super::types::{
21 FileOffer, FileTransferEvent, FtMessage, OfferDecision, OfferResponder, OverwritePolicy,
22 TransferDirection, TransferError, TransferProgress,
23};
24use super::MAX_PENDING_OFFERS_PER_PEER;
25
26fn safe_base_name(name: &str) -> Option<String> {
33 let last = name.rsplit(['/', '\\']).next().unwrap_or("").trim();
34 if last.is_empty() || last == "." || last == ".." || last.contains('\0') {
35 return None;
36 }
37 Some(last.to_string())
38}
39
40fn resolve_dest_path(save_path: &str, file_name: &str) -> Result<String, TransferError> {
49 let p = std::path::Path::new(save_path);
50 let treat_as_dir = p.is_dir() || save_path.ends_with('/') || save_path.ends_with('\\');
51 if treat_as_dir {
52 let safe = safe_base_name(file_name).ok_or_else(|| {
53 TransferError::Protocol(format!(
54 "Rejected unsafe file name from peer: {file_name:?}"
55 ))
56 })?;
57 Ok(format!(
58 "{}/{}",
59 save_path.trim_end_matches(['/', '\\']),
60 safe
61 ))
62 } else {
63 Ok(save_path.to_string())
64 }
65}
66
67fn authorize_pull_path(
78 roots: &[std::path::PathBuf],
79 requested: &str,
80 max_size: u64,
81) -> Result<std::path::PathBuf, TransferError> {
82 if roots.is_empty() {
83 return Err(TransferError::Rejected(
84 "pull serving is not enabled on this peer".into(),
85 ));
86 }
87 let canon = std::fs::canonicalize(requested)
88 .map_err(|_| TransferError::Rejected("requested path is not shared".into()))?;
89 if !roots.iter().any(|r| canon.starts_with(r)) {
90 return Err(TransferError::Rejected(
91 "requested path is not shared".into(),
92 ));
93 }
94 let meta = std::fs::metadata(&canon)
95 .map_err(|_| TransferError::Rejected("requested path is not shared".into()))?;
96 if !meta.is_file() {
97 return Err(TransferError::Rejected(
98 "requested path is not shared".into(),
99 ));
100 }
101 if meta.len() > max_size {
102 return Err(TransferError::Rejected(format!(
103 "file size {} exceeds max transfer size {max_size}",
104 meta.len()
105 )));
106 }
107 Ok(canon)
108}
109
110pub fn spawn_receive_handler<N: NetworkProvider + 'static>(
119 node: Arc<Node<N>>,
120 offer_tx: mpsc::Sender<(FileOffer, OfferResponder)>,
121 event_tx: broadcast::Sender<FileTransferEvent>,
122) -> tokio::task::JoinHandle<()> {
123 let cancel = node.tasks.cancel.clone();
124 let tracker = node.tasks.tracker.clone();
125 tracker.spawn(async move {
126 let mut rx = node.subscribe("ft");
127 info!("File transfer receive handler started");
128
129 loop {
130 let recv = tokio::select! {
131 _ = cancel.cancelled() => {
132 info!("FT receive handler: node stopping, exiting");
133 break;
134 }
135 result = rx.recv() => result,
136 };
137 let msg = match recv {
138 Ok(m) => m,
139 Err(broadcast::error::RecvError::Lagged(n)) => {
140 warn!("FT receive handler lagged, missed {n} messages");
141 continue;
142 }
143 Err(broadcast::error::RecvError::Closed) => {
144 info!("FT receive handler: channel closed, exiting");
145 break;
146 }
147 };
148
149 let ft_msg: FtMessage = match serde_json::from_value(msg.payload.clone()) {
150 Ok(m) => m,
151 Err(e) => {
152 warn!(from = msg.from.as_str(), "Bad FT message: {e}");
153 continue;
154 }
155 };
156
157 let node = node.clone();
158 let from = msg.from.clone();
159 let offer_tx = offer_tx.clone();
160 let event_tx = event_tx.clone();
161
162 match ft_msg {
163 FtMessage::Offer {
164 file_name,
165 size,
166 sha256,
167 save_path,
168 token,
169 tcp_port: _,
170 } => {
171 let permit = match node
174 .file_transfer_state
175 .incoming_ops
176 .clone()
177 .try_acquire_owned()
178 {
179 Ok(p) => p,
180 Err(_) => {
181 warn!(
182 from = from.as_str(),
183 "Rejecting offer: incoming-operation limit reached"
184 );
185 spawn_reject(node, from, token, "receiver busy");
186 continue;
187 }
188 };
189 let pending_guard = match PeerPendingGuard::try_acquire(
192 &node.file_transfer_state.pending_offers_per_peer,
193 &from,
194 MAX_PENDING_OFFERS_PER_PEER,
195 ) {
196 Some(g) => g,
197 None => {
198 warn!(
199 from = from.as_str(),
200 "Rejecting offer: per-peer pending-offer limit reached"
201 );
202 spawn_reject(node, from, token, "receiver busy");
203 continue;
204 }
205 };
206 let cancel = node.tasks.cancel.clone();
210 let tracker = node.tasks.tracker.clone();
211 tracker.spawn(async move {
212 let _permit = permit;
213 let _pending = pending_guard;
214 tokio::select! {
215 _ = cancel.cancelled() => {}
216 _ = async {
217 if let Err(e) = handle_incoming_offer(
218 &node, &from, &file_name, size, &sha256, &save_path, &token, &offer_tx,
219 &event_tx,
220 )
221 .await
222 {
223 match &e {
225 TransferError::Rejected(reason) => {
226 info!(
227 from = from.as_str(),
228 file = file_name.as_str(),
229 "File offer rejected: {reason}"
230 );
231 let _ = event_tx.send(FileTransferEvent::Rejected {
232 token,
233 file_name,
234 reason: reason.clone(),
235 });
236 }
237 _ => {
238 error!(
239 from = from.as_str(),
240 file = file_name.as_str(),
241 "Failed to receive file: {e}"
242 );
243 let _ = event_tx.send(FileTransferEvent::Failed {
244 token,
245 direction: TransferDirection::Receive,
246 file_name,
247 reason: e.to_string(),
248 });
249 }
250 }
251 }
252 } => {}
253 }
254 });
255 }
256 FtMessage::PullRequest {
257 path,
258 requester_id: _,
259 token,
260 } => {
261 let permit = match node
264 .file_transfer_state
265 .incoming_ops
266 .clone()
267 .try_acquire_owned()
268 {
269 Ok(p) => p,
270 Err(_) => {
271 warn!(
272 from = from.as_str(),
273 "Rejecting pull: incoming-operation limit reached"
274 );
275 spawn_reject(node, from, token, "receiver busy");
276 continue;
277 }
278 };
279 let cancel = node.tasks.cancel.clone();
282 let tracker = node.tasks.tracker.clone();
283 tracker.spawn(async move {
284 let _permit = permit;
285 tokio::select! {
286 _ = cancel.cancelled() => {}
287 _ = async {
288 if let Err(e) =
289 handle_pull_request(&node, &from, &path, &token, &event_tx).await
290 {
291 error!(
292 from = from.as_str(),
293 path = path.as_str(),
294 "Failed to serve file: {e}"
295 );
296 }
297 } => {}
298 }
299 });
300 }
301 _ => {
302 }
304 }
305 }
306 })
307}
308
309#[allow(clippy::too_many_arguments)]
312async fn handle_incoming_offer<N: NetworkProvider + 'static>(
313 node: &Node<N>,
314 from: &str,
315 file_name: &str,
316 size: u64,
317 sha256: &str,
318 save_path: &str,
319 token: &str,
320 offer_tx: &mpsc::Sender<(FileOffer, OfferResponder)>,
321 event_tx: &broadcast::Sender<FileTransferEvent>,
322) -> Result<(), TransferError> {
323 info!(
324 from = from,
325 file = file_name,
326 size = size,
327 "Received incoming file offer"
328 );
329
330 let max_size = node
334 .file_transfer_state
335 .max_transfer_size
336 .load(std::sync::atomic::Ordering::Relaxed);
337 if size > max_size {
338 return Err(TransferError::Protocol(format!(
339 "Offered file size {size} exceeds max transfer size {max_size}"
340 )));
341 }
342
343 let offer = FileOffer {
347 from_peer: from.to_string(),
348 from_name: from.to_string(), file_name: file_name.to_string(),
350 size,
351 sha256: sha256.to_string(),
352 suggested_path: safe_base_name(save_path).unwrap_or_default(),
353 token: token.to_string(),
354 };
355
356 let _ = event_tx.send(FileTransferEvent::OfferReceived(offer.clone()));
358
359 let (decision_tx, decision_rx) = oneshot::channel::<OfferDecision>();
361 let responder = OfferResponder::new(decision_tx);
362
363 if let Err(e) = offer_tx.try_send((offer, responder)) {
367 let reason = match e {
368 mpsc::error::TrySendError::Full(_) => "receiver busy: offer queue full",
369 mpsc::error::TrySendError::Closed(_) => "receiver has no offer channel",
370 };
371 send_reject(node, from, token, reason).await;
372 return Err(TransferError::Rejected(reason.to_string()));
373 }
374
375 let decision = tokio::time::timeout(tokio::time::Duration::from_secs(60), decision_rx)
377 .await
378 .map_err(|_| TransferError::Timeout)?
379 .map_err(|_| {
380 TransferError::Protocol("Offer responder dropped without decision".to_string())
381 })?;
382
383 match decision {
384 OfferDecision::Accept { save_path: dest } => {
385 accept_and_receive(node, from, file_name, size, sha256, token, &dest, event_tx).await
386 }
387 OfferDecision::Reject { reason } => {
388 let reject = FtMessage::Reject {
390 token: token.to_string(),
391 reason: reason.clone(),
392 };
393 let reject_payload = serde_json::to_value(&reject)
394 .map_err(|e| TransferError::Protocol(format!("Serialize error: {e}")))?;
395 node.send_typed(from, "ft", "reject", &reject_payload)
396 .await
397 .map_err(|e| TransferError::Node(format!("Failed to send REJECT: {e}")))?;
398
399 info!(
400 from = from,
401 file = file_name,
402 reason = reason.as_str(),
403 "Rejected file offer"
404 );
405
406 Err(TransferError::Rejected(reason))
407 }
408 }
409}
410
411#[allow(clippy::too_many_arguments)]
414async fn accept_and_receive<N: NetworkProvider + 'static>(
415 node: &Node<N>,
416 from: &str,
417 file_name: &str,
418 size: u64,
419 sha256: &str,
420 token: &str,
421 save_path: &str,
422 event_tx: &broadcast::Sender<FileTransferEvent>,
423) -> Result<(), TransferError> {
424 let start = std::time::Instant::now();
425
426 let final_path = resolve_dest_path(save_path, file_name)?;
430
431 let policy = node.file_transfer_state.overwrite_policy();
435 if policy == OverwritePolicy::Reject
436 && tokio::fs::try_exists(&final_path).await.unwrap_or(false)
437 {
438 let reason = "destination already exists";
439 send_reject(node, from, token, reason).await;
440 return Err(TransferError::Rejected(format!("{reason}: {final_path}")));
441 }
442
443 if let Some(parent) = std::path::Path::new(&final_path).parent() {
445 tokio::fs::create_dir_all(parent).await?;
446 }
447
448 let mut listener = node
450 .listen_tcp(0)
451 .await
452 .map_err(|e| TransferError::Node(format!("Failed to listen TCP: {e}")))?;
453
454 let accept = FtMessage::Accept {
456 token: token.to_string(),
457 tcp_port: listener.port,
458 };
459 let accept_payload = serde_json::to_value(&accept)
460 .map_err(|e| TransferError::Protocol(format!("Serialize error: {e}")))?;
461 node.send_typed(from, "ft", "accept", &accept_payload)
462 .await
463 .map_err(|e| TransferError::Node(format!("Failed to send ACCEPT: {e}")))?;
464
465 info!(
466 port = listener.port,
467 "Sent ACCEPT, listening for TCP connection"
468 );
469
470 let incoming = tokio::time::timeout(tokio::time::Duration::from_secs(30), listener.accept())
472 .await
473 .map_err(|_| TransferError::Timeout)?
474 .ok_or_else(|| TransferError::Protocol("Listener closed before accepting".to_string()))?;
475
476 let mut stream = incoming.stream;
477
478 let mut size_buf = [0u8; 8];
480 stream.read_exact(&mut size_buf).await?;
481 let file_size = u64::from_be_bytes(size_buf);
482
483 let mut sha_buf = [0u8; 64];
484 stream.read_exact(&mut sha_buf).await?;
485 let received_sha = String::from_utf8_lossy(&sha_buf).to_string();
486
487 if received_sha != sha256 {
489 return Err(TransferError::IntegrityError {
490 expected: sha256.to_string(),
491 actual: received_sha,
492 });
493 }
494
495 if file_size != size {
496 return Err(TransferError::Protocol(format!(
497 "Size mismatch: offer said {size}, stream header says {file_size}"
498 )));
499 }
500
501 let max_size = node
503 .file_transfer_state
504 .max_transfer_size
505 .load(std::sync::atomic::Ordering::Relaxed);
506 if file_size > max_size {
507 return Err(TransferError::Protocol(format!(
508 "File size {file_size} exceeds max transfer size {max_size}"
509 )));
510 }
511
512 let temp_path = format!("{final_path}.{}.truffle-tmp", uuid::Uuid::new_v4());
517 let mut temp_file = tokio::fs::OpenOptions::new()
518 .write(true)
519 .create_new(true)
520 .open(&temp_path)
521 .await?;
522 let mut hasher = Sha256::new();
523 let mut bytes_received: u64 = 0;
524 let progress_start = std::time::Instant::now();
525 let mut last_progress = std::time::Instant::now();
526 let mut buf = vec![0u8; 64 * 1024];
527
528 while bytes_received < file_size {
529 let to_read = ((file_size - bytes_received) as usize).min(buf.len());
530 let n = stream.read(&mut buf[..to_read]).await?;
531 if n == 0 {
532 tokio::fs::remove_file(&temp_path).await.ok();
533 return Err(TransferError::Io(std::io::Error::new(
534 std::io::ErrorKind::UnexpectedEof,
535 format!("Connection closed after {bytes_received}/{file_size} bytes"),
536 )));
537 }
538 hasher.update(&buf[..n]);
539 tokio::io::AsyncWriteExt::write_all(&mut temp_file, &buf[..n]).await?;
540 bytes_received += n as u64;
541
542 if last_progress.elapsed() >= std::time::Duration::from_millis(250) {
544 let elapsed = progress_start.elapsed().as_secs_f64();
545 let speed = if elapsed > 0.0 {
546 bytes_received as f64 / elapsed
547 } else {
548 0.0
549 };
550 let _ = event_tx.send(FileTransferEvent::Progress(TransferProgress {
551 token: token.to_string(),
552 direction: TransferDirection::Receive,
553 file_name: file_name.to_string(),
554 bytes_transferred: bytes_received,
555 total_bytes: file_size,
556 speed_bps: speed,
557 }));
558 last_progress = std::time::Instant::now();
559 }
560 }
561
562 tokio::io::AsyncWriteExt::flush(&mut temp_file).await?;
564
565 let actual_sha = hex::encode(hasher.finalize());
567
568 if actual_sha != sha256 {
569 stream.write_all(&[0x00]).await?;
571 tokio::fs::remove_file(&temp_path).await.ok();
573 return Err(TransferError::IntegrityError {
574 expected: sha256.to_string(),
575 actual: actual_sha,
576 });
577 }
578
579 info!(
583 temp = temp_path.as_str(),
584 final_path = final_path.as_str(),
585 "Moving temp file to final destination"
586 );
587 let placed_path = finalize_received_file(&temp_path, &final_path, policy).await?;
588 info!(final_path = placed_path.as_str(), "File save completed");
589
590 stream.write_all(&[0x01]).await?;
592 tokio::io::AsyncWriteExt::flush(&mut stream).await?;
593
594 tokio::time::sleep(std::time::Duration::from_millis(100)).await;
599
600 let elapsed = start.elapsed().as_secs_f64();
601 info!(
602 file = placed_path.as_str(),
603 bytes = file_size,
604 elapsed_ms = (elapsed * 1000.0) as u64,
605 "File received and verified"
606 );
607
608 let _ = event_tx.send(FileTransferEvent::Completed {
610 token: token.to_string(),
611 direction: TransferDirection::Receive,
612 file_name: file_name.to_string(),
613 bytes_transferred: file_size,
614 sha256: actual_sha,
615 elapsed_secs: elapsed,
616 });
617
618 Ok(())
619}
620
621async fn handle_pull_request<N: NetworkProvider + 'static>(
623 node: &Node<N>,
624 from: &str,
625 path: &str,
626 token: &str,
627 event_tx: &broadcast::Sender<FileTransferEvent>,
628) -> Result<(), TransferError> {
629 info!(from = from, path = path, "Processing PULL_REQUEST");
630
631 let max_size = node
636 .file_transfer_state
637 .max_transfer_size
638 .load(std::sync::atomic::Ordering::Relaxed);
639 let roots = node.file_transfer_state.pull_roots.read().unwrap().clone();
640 let serve_path = match authorize_pull_path(&roots, path, max_size) {
641 Ok(p) => p,
642 Err(e) => {
643 warn!(from = from, path = path, "Denying PULL_REQUEST: {e}");
644 let reject = FtMessage::Reject {
646 token: token.to_string(),
647 reason: "pull denied by peer".to_string(),
648 };
649 if let Ok(payload) = serde_json::to_value(&reject) {
650 if let Err(send_err) = node.send_typed(from, "ft", "reject", &payload).await {
651 warn!(from = from, "Failed to send pull REJECT: {send_err}");
652 }
653 }
654 return Err(e);
655 }
656 };
657
658 let meta = tokio::fs::metadata(&serve_path)
662 .await
663 .map_err(TransferError::Io)?;
664 let size = meta.len();
665 let sha256 = {
666 let mut hasher = Sha256::new();
667 let mut file = tokio::fs::File::open(&serve_path)
668 .await
669 .map_err(TransferError::Io)?;
670 let mut buf = vec![0u8; 64 * 1024];
671 loop {
672 let n = file.read(&mut buf).await.map_err(TransferError::Io)?;
673 if n == 0 {
674 break;
675 }
676 hasher.update(&buf[..n]);
677 }
678 hex::encode(hasher.finalize())
679 };
680
681 let file_name = std::path::Path::new(path)
682 .file_name()
683 .and_then(|n| n.to_str())
684 .unwrap_or("file")
685 .to_string();
686
687 let offer_token = uuid::Uuid::new_v4().to_string();
688
689 let offer = FtMessage::Offer {
691 file_name: file_name.clone(),
692 size,
693 sha256: sha256.clone(),
694 save_path: String::new(),
695 token: offer_token.clone(),
696 tcp_port: 0,
697 };
698 let offer_payload = serde_json::to_value(&offer)
699 .map_err(|e| TransferError::Protocol(format!("Serialize error: {e}")))?;
700
701 let _accept_port = crate::request_reply::send_and_wait(
702 node,
703 from,
704 "ft",
705 "offer",
706 &offer_payload,
707 std::time::Duration::from_secs(30),
708 |msg| {
709 if msg.from != from {
710 return None;
711 }
712 let ft_msg: FtMessage = serde_json::from_value(msg.payload.clone()).ok()?;
713 match ft_msg {
714 FtMessage::Accept {
715 token: ref t,
716 tcp_port,
717 } if *t == offer_token => Some(Ok(tcp_port)),
718 FtMessage::Reject {
719 token: ref t,
720 reason,
721 } if *t == offer_token => Some(Err(TransferError::Rejected(format!(
722 "Peer rejected: {reason}"
723 )))),
724 _ => None,
725 }
726 },
727 )
728 .await
729 .map_err(|e| match e {
730 crate::request_reply::RequestError::Timeout => TransferError::Timeout,
731 crate::request_reply::RequestError::Send(e) => {
732 TransferError::Node(format!("Failed to send OFFER: {e}"))
733 }
734 crate::request_reply::RequestError::ChannelClosed => {
735 TransferError::Protocol("Channel closed".into())
736 }
737 })??;
738
739 let mut stream = node.open_tcp(from, _accept_port).await.map_err(|e| {
741 TransferError::Node(format!("Failed to open TCP to {from}:{_accept_port}: {e}"))
742 })?;
743
744 let start = std::time::Instant::now();
745
746 stream.write_all(&size.to_be_bytes()).await?;
750 stream.write_all(sha256.as_bytes()).await?;
751
752 let mut file = tokio::fs::File::open(&serve_path)
753 .await
754 .map_err(TransferError::Io)?;
755 let mut buf = vec![0u8; 64 * 1024];
756 let mut bytes_sent: u64 = 0;
757 let mut last_progress = std::time::Instant::now();
758
759 loop {
760 let n = file.read(&mut buf).await.map_err(TransferError::Io)?;
761 if n == 0 {
762 break;
763 }
764 stream.write_all(&buf[..n]).await?;
765 bytes_sent += n as u64;
766
767 if last_progress.elapsed() >= std::time::Duration::from_millis(250) {
769 let elapsed = start.elapsed().as_secs_f64();
770 let speed = if elapsed > 0.0 {
771 bytes_sent as f64 / elapsed
772 } else {
773 0.0
774 };
775 let _ = event_tx.send(FileTransferEvent::Progress(TransferProgress {
776 token: offer_token.clone(),
777 direction: TransferDirection::Send,
778 file_name: file_name.clone(),
779 bytes_transferred: bytes_sent,
780 total_bytes: size,
781 speed_bps: speed,
782 }));
783 last_progress = std::time::Instant::now();
784 }
785 }
786
787 stream.flush().await?;
788
789 let mut ack = [0u8; 1];
791 stream.read_exact(&mut ack).await?;
792
793 if ack[0] != 0x01 {
794 return Err(TransferError::IntegrityError {
795 expected: sha256,
796 actual: "peer reported integrity failure".to_string(),
797 });
798 }
799
800 let elapsed = start.elapsed().as_secs_f64();
801 info!(path = path, bytes = size, "File served successfully");
802
803 let _ = event_tx.send(FileTransferEvent::Completed {
805 token: offer_token,
806 direction: TransferDirection::Send,
807 file_name,
808 bytes_transferred: size,
809 sha256,
810 elapsed_secs: elapsed,
811 });
812
813 Ok(())
814}
815
816async fn send_reject<N: NetworkProvider + 'static>(
823 node: &Node<N>,
824 from: &str,
825 token: &str,
826 reason: &str,
827) {
828 let reject = FtMessage::Reject {
829 token: token.to_string(),
830 reason: reason.to_string(),
831 };
832 if let Ok(payload) = serde_json::to_value(&reject) {
833 if let Err(e) = node.send_typed(from, "ft", "reject", &payload).await {
834 warn!(from = from, "Failed to send REJECT: {e}");
835 }
836 }
837}
838
839fn spawn_reject<N: NetworkProvider + 'static>(
842 node: Arc<Node<N>>,
843 from: String,
844 token: String,
845 reason: &'static str,
846) {
847 let tracker = node.tasks.tracker.clone();
848 tracker.spawn(async move {
849 send_reject(&node, &from, &token, reason).await;
850 });
851}
852
853struct PeerPendingGuard {
857 map: Arc<std::sync::Mutex<std::collections::HashMap<String, usize>>>,
858 peer: String,
859}
860
861impl PeerPendingGuard {
862 fn try_acquire(
863 map: &Arc<std::sync::Mutex<std::collections::HashMap<String, usize>>>,
864 peer: &str,
865 max: usize,
866 ) -> Option<Self> {
867 let mut m = map.lock().unwrap();
868 let count = m.entry(peer.to_string()).or_insert(0);
869 if *count >= max {
870 return None;
871 }
872 *count += 1;
873 Some(Self {
874 map: map.clone(),
875 peer: peer.to_string(),
876 })
877 }
878}
879
880impl Drop for PeerPendingGuard {
881 fn drop(&mut self) {
882 let mut m = self.map.lock().unwrap();
883 if let Some(c) = m.get_mut(&self.peer) {
884 *c = c.saturating_sub(1);
885 if *c == 0 {
886 m.remove(&self.peer);
887 }
888 }
889 }
890}
891
892enum PlaceError {
895 Exists,
896 Io(std::io::Error),
897}
898
899async fn place_no_clobber(temp_path: &str, dest: &str) -> Result<(), PlaceError> {
907 match tokio::fs::hard_link(temp_path, dest).await {
908 Ok(()) => {
909 tokio::fs::remove_file(temp_path).await.ok();
910 Ok(())
911 }
912 Err(e) if e.kind() == std::io::ErrorKind::AlreadyExists => Err(PlaceError::Exists),
913 Err(_) => {
914 if tokio::fs::try_exists(dest).await.unwrap_or(false) {
915 return Err(PlaceError::Exists);
916 }
917 tokio::fs::rename(temp_path, dest)
918 .await
919 .map_err(PlaceError::Io)
920 }
921 }
922}
923
924fn dedup_candidate(path: &str, i: u32) -> String {
927 let p = std::path::Path::new(path);
928 let stem = p.file_stem().and_then(|s| s.to_str()).unwrap_or("file");
929 let new_name = match p.extension().and_then(|e| e.to_str()) {
930 Some(ext) => format!("{stem} ({i}).{ext}"),
931 None => format!("{stem} ({i})"),
932 };
933 p.with_file_name(new_name).to_string_lossy().into_owned()
934}
935
936pub(crate) async fn finalize_received_file(
941 temp_path: &str,
942 final_path: &str,
943 policy: OverwritePolicy,
944) -> Result<String, TransferError> {
945 match policy {
946 OverwritePolicy::Replace => {
947 if let Err(rename_err) = tokio::fs::rename(temp_path, final_path).await {
950 info!(err = %rename_err, "Rename failed, trying copy+delete fallback");
951 if let Err(e) = tokio::fs::copy(temp_path, final_path).await {
952 tokio::fs::remove_file(temp_path).await.ok();
953 return Err(TransferError::Io(e));
954 }
955 tokio::fs::remove_file(temp_path).await.ok();
956 }
957 Ok(final_path.to_string())
958 }
959 OverwritePolicy::Reject => match place_no_clobber(temp_path, final_path).await {
960 Ok(()) => Ok(final_path.to_string()),
961 Err(PlaceError::Exists) => {
962 tokio::fs::remove_file(temp_path).await.ok();
963 Err(TransferError::Protocol(format!(
964 "destination already exists: {final_path} (OverwritePolicy::Reject)"
965 )))
966 }
967 Err(PlaceError::Io(e)) => {
968 tokio::fs::remove_file(temp_path).await.ok();
969 Err(TransferError::Io(e))
970 }
971 },
972 OverwritePolicy::Rename => {
973 match place_no_clobber(temp_path, final_path).await {
974 Ok(()) => return Ok(final_path.to_string()),
975 Err(PlaceError::Exists) => {}
976 Err(PlaceError::Io(e)) => {
977 tokio::fs::remove_file(temp_path).await.ok();
978 return Err(TransferError::Io(e));
979 }
980 }
981 for i in 1..=999u32 {
982 let candidate = dedup_candidate(final_path, i);
983 match place_no_clobber(temp_path, &candidate).await {
984 Ok(()) => return Ok(candidate),
985 Err(PlaceError::Exists) => continue,
986 Err(PlaceError::Io(e)) => {
987 tokio::fs::remove_file(temp_path).await.ok();
988 return Err(TransferError::Io(e));
989 }
990 }
991 }
992 tokio::fs::remove_file(temp_path).await.ok();
993 Err(TransferError::Protocol(format!(
994 "no free deduplicated name found for {final_path}"
995 )))
996 }
997 }
998}
999
1000#[cfg(test)]
1001mod tests {
1002 use super::*;
1003
1004 use crate::envelope::codec::JsonCodec;
1006 use crate::envelope::EnvelopeCodec;
1007 use crate::network::*;
1008 use crate::session::PeerRegistry;
1009 use crate::transport::websocket::WebSocketTransport;
1010 use crate::transport::WsConfig;
1011 use std::time::Duration;
1012
1013 #[test]
1014 fn safe_base_name_strips_directories() {
1015 assert_eq!(safe_base_name("file.txt").as_deref(), Some("file.txt"));
1016 assert_eq!(
1017 safe_base_name("../../etc/passwd").as_deref(),
1018 Some("passwd")
1019 );
1020 assert_eq!(safe_base_name("/abs/path/x").as_deref(), Some("x"));
1021 assert_eq!(
1022 safe_base_name("..\\..\\win.exe").as_deref(),
1023 Some("win.exe")
1024 );
1025 }
1026
1027 #[test]
1028 fn safe_base_name_rejects_unsafe() {
1029 assert_eq!(safe_base_name(""), None);
1030 assert_eq!(safe_base_name("."), None);
1031 assert_eq!(safe_base_name(".."), None);
1032 assert_eq!(safe_base_name("../.."), None); assert_eq!(safe_base_name("dir/"), None); }
1035
1036 #[test]
1037 fn resolve_dest_contains_traversal_into_directory() {
1038 let got = resolve_dest_path("/downloads/", "../../../.ssh/authorized_keys").unwrap();
1041 assert_eq!(got, "/downloads/authorized_keys");
1042 assert!(!got.contains(".."));
1043 }
1044
1045 #[test]
1046 fn resolve_dest_rejects_dotdot_filename() {
1047 assert!(resolve_dest_path("/downloads/", "..").is_err());
1048 }
1049
1050 #[test]
1051 fn resolve_dest_passes_through_explicit_file() {
1052 assert_eq!(
1054 resolve_dest_path("/tmp/truffle-explicit-file.bin", "ignored").unwrap(),
1055 "/tmp/truffle-explicit-file.bin"
1056 );
1057 }
1058
1059 #[test]
1062 fn pull_denied_when_no_roots() {
1063 let dir = tempfile::tempdir().unwrap();
1065 let file = dir.path().join("f.txt");
1066 std::fs::write(&file, b"hi").unwrap();
1067 let err = authorize_pull_path(&[], file.to_str().unwrap(), u64::MAX).unwrap_err();
1068 assert!(matches!(err, TransferError::Rejected(_)));
1069 }
1070
1071 #[test]
1072 fn pull_allowlisted_file_authorized() {
1073 let dir = tempfile::tempdir().unwrap();
1075 let root = std::fs::canonicalize(dir.path()).unwrap();
1076 let file = root.join("f.txt");
1077 std::fs::write(&file, b"hi").unwrap();
1078 let got = authorize_pull_path(&[root], file.to_str().unwrap(), u64::MAX).unwrap();
1079 assert_eq!(got, std::fs::canonicalize(&file).unwrap());
1080 }
1081
1082 #[test]
1083 fn pull_outside_root_rejected() {
1084 let dir_a = tempfile::tempdir().unwrap();
1085 let dir_b = tempfile::tempdir().unwrap();
1086 let root_a = std::fs::canonicalize(dir_a.path()).unwrap();
1087 let file_b = dir_b.path().join("secret.txt");
1088 std::fs::write(&file_b, b"secret").unwrap();
1089
1090 let err =
1092 authorize_pull_path(&[root_a.clone()], file_b.to_str().unwrap(), u64::MAX).unwrap_err();
1093 assert!(matches!(err, TransferError::Rejected(_)));
1094
1095 let dir_b_name = dir_b.path().file_name().unwrap().to_str().unwrap();
1098 let dotdot = format!("{}/../{}/secret.txt", dir_a.path().display(), dir_b_name);
1099 let err = authorize_pull_path(&[root_a], &dotdot, u64::MAX).unwrap_err();
1100 assert!(matches!(err, TransferError::Rejected(_)));
1101 }
1102
1103 #[test]
1104 fn pull_prefix_collision_rejected() {
1105 let base = tempfile::tempdir().unwrap();
1108 let shared = base.path().join("shared");
1109 let evil = base.path().join("shared-evil");
1110 std::fs::create_dir(&shared).unwrap();
1111 std::fs::create_dir(&evil).unwrap();
1112 let evil_file = evil.join("f.txt");
1113 std::fs::write(&evil_file, b"x").unwrap();
1114
1115 let root = std::fs::canonicalize(&shared).unwrap();
1116 let err = authorize_pull_path(&[root], evil_file.to_str().unwrap(), u64::MAX).unwrap_err();
1117 assert!(matches!(err, TransferError::Rejected(_)));
1118 }
1119
1120 #[cfg(unix)]
1121 #[test]
1122 fn pull_symlink_escape_rejected() {
1123 let root_dir = tempfile::tempdir().unwrap();
1126 let outside_dir = tempfile::tempdir().unwrap();
1127 let root = std::fs::canonicalize(root_dir.path()).unwrap();
1128
1129 let secret = outside_dir.path().join("secret.txt");
1130 std::fs::write(&secret, b"secret").unwrap();
1131
1132 let link = root.join("link.txt");
1133 std::os::unix::fs::symlink(&secret, &link).unwrap();
1134
1135 let err = authorize_pull_path(&[root], link.to_str().unwrap(), u64::MAX).unwrap_err();
1136 assert!(matches!(err, TransferError::Rejected(_)));
1137 }
1138
1139 #[test]
1140 fn pull_oversize_file_rejected() {
1141 let dir = tempfile::tempdir().unwrap();
1143 let root = std::fs::canonicalize(dir.path()).unwrap();
1144 let file = root.join("big.bin");
1145 std::fs::write(&file, b"1234").unwrap();
1146 let err = authorize_pull_path(&[root], file.to_str().unwrap(), 3).unwrap_err();
1147 assert!(matches!(err, TransferError::Rejected(_)));
1148 }
1149
1150 struct MockNetworkProvider {
1153 identity: NodeIdentity,
1154 local_addr: PeerAddr,
1155 peer_event_tx: tokio::sync::broadcast::Sender<NetworkPeerEvent>,
1156 mock_peers: Arc<tokio::sync::RwLock<Vec<NetworkPeer>>>,
1157 }
1158
1159 impl MockNetworkProvider {
1160 fn new(id: &str) -> Self {
1161 let (peer_event_tx, _) = tokio::sync::broadcast::channel(64);
1162 Self {
1163 identity: NodeIdentity {
1164 app_id: "test".to_string(),
1165 device_id: id.to_string(),
1166 device_name: format!("Test Node {id}"),
1167 tailscale_hostname: format!("truffle-test-{id}"),
1168 tailscale_id: id.to_string(),
1169 dns_name: None,
1170 ip: Some("127.0.0.1".parse().unwrap()),
1171 },
1172 local_addr: PeerAddr {
1173 ip: Some("127.0.0.1".parse().unwrap()),
1174 hostname: format!("truffle-test-{id}"),
1175 dns_name: None,
1176 },
1177 peer_event_tx,
1178 mock_peers: Arc::new(tokio::sync::RwLock::new(Vec::new())),
1179 }
1180 }
1181
1182 fn event_sender(&self) -> tokio::sync::broadcast::Sender<NetworkPeerEvent> {
1183 self.peer_event_tx.clone()
1184 }
1185 }
1186
1187 impl NetworkProvider for MockNetworkProvider {
1188 fn local_identity(&self) -> NodeIdentity {
1189 self.identity.clone()
1190 }
1191 fn local_addr(&self) -> PeerAddr {
1192 self.local_addr.clone()
1193 }
1194 fn peer_events(&self) -> tokio::sync::broadcast::Receiver<NetworkPeerEvent> {
1195 self.peer_event_tx.subscribe()
1196 }
1197
1198 async fn start(&mut self) -> Result<(), NetworkError> {
1199 Ok(())
1200 }
1201 async fn stop(&self) -> Result<(), NetworkError> {
1202 Ok(())
1203 }
1204 async fn peers(&self) -> Vec<NetworkPeer> {
1205 self.mock_peers.read().await.clone()
1206 }
1207 async fn dial_tcp(
1208 &self,
1209 _addr: &str,
1210 _port: u16,
1211 ) -> Result<tokio::net::TcpStream, NetworkError> {
1212 Err(NetworkError::DialFailed("mock".into()))
1213 }
1214 async fn listen_tcp(&self, _port: u16) -> Result<NetworkTcpListener, NetworkError> {
1215 Err(NetworkError::ListenFailed("mock".into()))
1216 }
1217 async fn unlisten_tcp(&self, _port: u16) -> Result<(), NetworkError> {
1218 Ok(())
1219 }
1220 async fn bind_udp(&self, _port: u16) -> Result<NetworkUdpSocket, NetworkError> {
1221 Err(NetworkError::NotRunning)
1222 }
1223 async fn ping(&self, _addr: &str) -> Result<PingResult, NetworkError> {
1224 Ok(PingResult {
1225 latency: Duration::from_millis(1),
1226 connection: "direct".to_string(),
1227 peer_addr: None,
1228 })
1229 }
1230 async fn health(&self) -> HealthInfo {
1231 HealthInfo {
1232 state: "running".to_string(),
1233 healthy: true,
1234 ..Default::default()
1235 }
1236 }
1237 }
1238
1239 fn ws_config(port: u16) -> WsConfig {
1240 WsConfig {
1241 port,
1242 ping_interval: Duration::from_secs(300),
1243 pong_timeout: Duration::from_secs(300),
1244 ..Default::default()
1245 }
1246 }
1247
1248 async fn random_port() -> u16 {
1249 let l = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
1250 l.local_addr().unwrap().port()
1251 }
1252
1253 async fn make_test_node(
1254 id: &str,
1255 ws_port: u16,
1256 ) -> (
1257 Node<MockNetworkProvider>,
1258 tokio::sync::broadcast::Sender<NetworkPeerEvent>,
1259 ) {
1260 let provider = MockNetworkProvider::new(id);
1261 let event_tx = provider.event_sender();
1262 let network = Arc::new(provider);
1263 let ws_transport = Arc::new(WebSocketTransport::new(network.clone(), ws_config(ws_port)));
1264 let session = Arc::new(PeerRegistry::new(network.clone(), ws_transport));
1265 session.start().await;
1266
1267 let codec: Arc<dyn EnvelopeCodec> = Arc::new(JsonCodec);
1268 let node = Node::from_parts(network, session, codec);
1269 (node, event_tx)
1270 }
1271
1272 #[tokio::test]
1280 async fn handle_pull_request_denied_by_default() {
1281 let (node, _event_tx) = make_test_node("node-a", random_port().await).await;
1282
1283 let dir = tempfile::tempdir().unwrap();
1285 let secret = dir.path().join("secret.txt");
1286 std::fs::write(&secret, b"top secret").unwrap();
1287 let secret_path = secret.to_str().unwrap();
1288
1289 let (event_tx, _rx) = tokio::sync::broadcast::channel(16);
1290 let err = handle_pull_request(&node, "peer-x", secret_path, "tok-1", &event_tx)
1291 .await
1292 .unwrap_err();
1293
1294 assert!(
1295 matches!(err, TransferError::Rejected(_)),
1296 "expected Rejected, got {err}"
1297 );
1298 }
1299
1300 #[test]
1303 fn dedup_candidate_naming() {
1304 let got = dedup_candidate("/d/report.pdf", 1);
1307 assert!(got.ends_with("report (1).pdf"), "{got}");
1308 let got = dedup_candidate("/d/Makefile", 2);
1309 assert!(got.ends_with("Makefile (2)"), "{got}");
1310 }
1311
1312 async fn write_file(dir: &std::path::Path, name: &str, content: &[u8]) -> String {
1313 let p = dir.join(name);
1314 tokio::fs::write(&p, content).await.unwrap();
1315 p.to_str().unwrap().to_string()
1316 }
1317
1318 #[tokio::test]
1319 async fn finalize_reject_errors_when_dest_exists() {
1320 let dir = tempfile::tempdir().unwrap();
1321 let dest = write_file(dir.path(), "f.txt", b"old").await;
1322 let temp = write_file(dir.path(), "f.txt.abc.truffle-tmp", b"new").await;
1323
1324 let err = finalize_received_file(&temp, &dest, OverwritePolicy::Reject)
1325 .await
1326 .unwrap_err();
1327 assert!(err.to_string().contains("already exists"), "{err}");
1328 assert_eq!(tokio::fs::read(&dest).await.unwrap(), b"old");
1330 assert!(!tokio::fs::try_exists(&temp).await.unwrap());
1331 }
1332
1333 #[tokio::test]
1334 async fn finalize_reject_places_when_dest_missing() {
1335 let dir = tempfile::tempdir().unwrap();
1336 let temp = write_file(dir.path(), "g.txt.abc.truffle-tmp", b"data").await;
1337 let dest = dir.path().join("g.txt").to_str().unwrap().to_string();
1338
1339 let placed = finalize_received_file(&temp, &dest, OverwritePolicy::Reject)
1340 .await
1341 .unwrap();
1342 assert_eq!(placed, dest);
1343 assert_eq!(tokio::fs::read(&dest).await.unwrap(), b"data");
1344 assert!(!tokio::fs::try_exists(&temp).await.unwrap());
1345 }
1346
1347 #[tokio::test]
1348 async fn finalize_replace_overwrites() {
1349 let dir = tempfile::tempdir().unwrap();
1350 let dest = write_file(dir.path(), "h.txt", b"old").await;
1351 let temp = write_file(dir.path(), "h.txt.abc.truffle-tmp", b"new").await;
1352
1353 let placed = finalize_received_file(&temp, &dest, OverwritePolicy::Replace)
1354 .await
1355 .unwrap();
1356 assert_eq!(placed, dest);
1357 assert_eq!(tokio::fs::read(&dest).await.unwrap(), b"new");
1358 }
1359
1360 #[tokio::test]
1361 async fn finalize_rename_dedupes() {
1362 let dir = tempfile::tempdir().unwrap();
1363 let dest = write_file(dir.path(), "i.txt", b"old").await;
1364 let temp = write_file(dir.path(), "i.txt.abc.truffle-tmp", b"new").await;
1365
1366 let placed = finalize_received_file(&temp, &dest, OverwritePolicy::Rename)
1367 .await
1368 .unwrap();
1369 assert!(placed.ends_with("i (1).txt"), "{placed}");
1370 assert_eq!(tokio::fs::read(&dest).await.unwrap(), b"old");
1371 assert_eq!(tokio::fs::read(&placed).await.unwrap(), b"new");
1372 }
1373
1374 #[test]
1375 fn peer_pending_guard_caps_and_releases() {
1376 let map = Arc::new(std::sync::Mutex::new(std::collections::HashMap::new()));
1377 let g1 = PeerPendingGuard::try_acquire(&map, "p", 2).unwrap();
1378 let _g2 = PeerPendingGuard::try_acquire(&map, "p", 2).unwrap();
1379 assert!(PeerPendingGuard::try_acquire(&map, "p", 2).is_none());
1380 assert!(PeerPendingGuard::try_acquire(&map, "q", 2).is_some());
1382 drop(g1);
1383 assert!(PeerPendingGuard::try_acquire(&map, "p", 2).is_some());
1384 }
1385
1386 #[tokio::test]
1387 async fn offer_rejected_when_queue_full() {
1388 let (node, _peer_events) = make_test_node("node-b", random_port().await).await;
1389
1390 let (offer_tx, _offer_rx) = mpsc::channel(1);
1392 let (dtx, _drx) = tokio::sync::oneshot::channel();
1393 offer_tx
1394 .try_send((
1395 FileOffer {
1396 from_peer: "x".into(),
1397 from_name: "x".into(),
1398 file_name: "a".into(),
1399 size: 1,
1400 sha256: "0".repeat(64),
1401 suggested_path: String::new(),
1402 token: "t0".into(),
1403 },
1404 OfferResponder::new(dtx),
1405 ))
1406 .unwrap();
1407
1408 let (event_tx, _rx) = tokio::sync::broadcast::channel(16);
1409 let err = handle_incoming_offer(
1410 &node,
1411 "peer-x",
1412 "b.txt",
1413 1,
1414 &"0".repeat(64),
1415 "",
1416 "t1",
1417 &offer_tx,
1418 &event_tx,
1419 )
1420 .await
1421 .unwrap_err();
1422 assert!(matches!(err, TransferError::Rejected(_)), "{err}");
1423 }
1424
1425 use proptest::prelude::*;
1428
1429 proptest! {
1430 #[test]
1433 fn safe_base_name_never_escapes(name in ".{0,256}") {
1434 if let Some(safe) = safe_base_name(&name) {
1435 prop_assert!(!safe.contains('/'));
1436 prop_assert!(!safe.contains('\\'));
1437 prop_assert!(!safe.contains('\0'));
1438 prop_assert!(safe != "." && safe != "..");
1439 prop_assert!(!safe.is_empty());
1440 }
1441 }
1442 }
1443
1444 #[tokio::test]
1445 async fn concurrent_finalize_same_dest_reject_single_winner() {
1446 let dir = tempfile::tempdir().unwrap();
1447 let dest = dir.path().join("race.txt").to_str().unwrap().to_string();
1448 let t1 = write_file(dir.path(), "race.txt.aaa.truffle-tmp", b"one").await;
1449 let t2 = write_file(dir.path(), "race.txt.bbb.truffle-tmp", b"two").await;
1450
1451 let (r1, r2) = tokio::join!(
1452 finalize_received_file(&t1, &dest, OverwritePolicy::Reject),
1453 finalize_received_file(&t2, &dest, OverwritePolicy::Reject),
1454 );
1455 assert!(
1456 r1.is_ok() ^ r2.is_ok(),
1457 "exactly one writer should win: {r1:?} / {r2:?}"
1458 );
1459 let winner: &[u8] = if r1.is_ok() { b"one" } else { b"two" };
1460 assert_eq!(tokio::fs::read(&dest).await.unwrap(), winner);
1461 assert!(!tokio::fs::try_exists(&t1).await.unwrap());
1463 assert!(!tokio::fs::try_exists(&t2).await.unwrap());
1464 }
1465
1466 #[tokio::test]
1467 async fn concurrent_finalize_same_dest_rename_both_win() {
1468 let dir = tempfile::tempdir().unwrap();
1469 let dest = dir.path().join("both.txt").to_str().unwrap().to_string();
1470 let t1 = write_file(dir.path(), "both.txt.aaa.truffle-tmp", b"one").await;
1471 let t2 = write_file(dir.path(), "both.txt.bbb.truffle-tmp", b"two").await;
1472
1473 let (r1, r2) = tokio::join!(
1474 finalize_received_file(&t1, &dest, OverwritePolicy::Rename),
1475 finalize_received_file(&t2, &dest, OverwritePolicy::Rename),
1476 );
1477 let (p1, p2) = (r1.unwrap(), r2.unwrap());
1478 assert_ne!(p1, p2, "rename policy must dedupe concurrent writers");
1479 let mut contents = vec![
1480 tokio::fs::read(&p1).await.unwrap(),
1481 tokio::fs::read(&p2).await.unwrap(),
1482 ];
1483 contents.sort();
1484 assert_eq!(contents, vec![b"one".to_vec(), b"two".to_vec()]);
1485 }
1486}