nirc-rs/src/transfer/engine.rs

654 lines
27 KiB
Rust
Executable File

//! Zero-copy file transfer engine — Phase 15.
//!
//! Implements the actual send/receive data path with:
//! - Large async I/O buffers (256 KiB) to minimise syscalls and maximise throughput
//! - Streaming SHA-256 verification (computed in-flight, not post-hoc)
//! - Resume support (offset-based, writes to `.partial` then atomically renames)
//! - Progress callbacks via channel — non-blocking to the transfer loop
//! - Cancellation via tokio::CancellationToken
//! - Integration with TransferManager (Phase 14) for state tracking
#[cfg(test)]
use crate::core::protocol::ProtocolType;
use crate::transfer::{
TransferId, TransferManager, TransferState,
};
use sha2::{Digest, Sha256};
use std::path::{Path, PathBuf};
use tokio::io::{AsyncReadExt, AsyncSeekExt, AsyncWriteExt};
use tokio::sync::mpsc;
use tracing::{error, info, warn};
/// Progress update emitted during a transfer.
#[derive(Debug, Clone)]
pub struct TransferProgress {
pub id: TransferId,
pub bytes_transferred: u64,
pub total_bytes: u64,
pub bytes_per_sec: f64,
pub eta_secs: Option<f64>,
/// True when SHA-256 verification succeeded after completion.
pub hash_verified: bool,
pub final_hash: Option<String>,
}
/// Result of a completed transfer.
#[derive(Debug)]
pub enum TransferResult {
Completed { hash: String },
Failed { error: String },
Cancelled,
}
/// Wire protocol header sent before file data over a yamux stream.
///
/// Layout (all little-endian):
/// 4 bytes magic b"NAIM"
/// 2 bytes version (0x0001)
/// 1 byte flags (bit 0: resume_supported, bit 1: hash_included)
/// 8 bytes file_size
/// 8 bytes resume_offset (0 for new transfer)
/// 4 bytes filename_len
/// N bytes filename (UTF-8)
/// 64 bytes sha256 (present if flag bit 1 set)
#[derive(Debug, Clone)]
pub struct TransferHeader {
pub file_size: u64,
pub resume_offset: u64,
pub filename: String,
pub sha256: Option<[u8; 32]>,
pub flags: u8,
}
const TRANSFER_MAGIC: &[u8; 4] = b"NAIM";
const TRANSFER_VERSION: u16 = 1;
const FLAG_RESUME: u8 = 0b0000_0001;
const FLAG_HASH: u8 = 0b0000_0010;
/// I/O buffer size — 256 KiB for high throughput on modern networks.
const BUFFER_SIZE: usize = 256 * 1024;
impl TransferHeader {
/// Serialize header to bytes for wire transmission.
pub fn to_bytes(&self) -> Vec<u8> {
let mut buf = Vec::with_capacity(128 + self.filename.len());
buf.extend_from_slice(TRANSFER_MAGIC);
buf.extend_from_slice(&TRANSFER_VERSION.to_le_bytes());
let mut flags = self.flags;
if self.sha256.is_some() { flags |= FLAG_HASH; }
if self.resume_offset > 0 { flags |= FLAG_RESUME; }
buf.push(flags);
buf.extend_from_slice(&self.file_size.to_le_bytes());
buf.extend_from_slice(&self.resume_offset.to_le_bytes());
let fname_bytes = self.filename.as_bytes();
buf.extend_from_slice(&(fname_bytes.len() as u32).to_le_bytes());
buf.extend_from_slice(fname_bytes);
if let Some(hash) = &self.sha256 {
buf.extend_from_slice(hash);
}
buf
}
/// Parse header from bytes received from the wire.
pub fn from_bytes(data: &[u8]) -> anyhow::Result<Self> {
if data.len() < 27 || &data[0..4] != TRANSFER_MAGIC {
anyhow::bail!("invalid transfer header: bad magic or too short");
}
let version = u16::from_le_bytes(data[4..6].try_into()?);
if version != TRANSFER_VERSION {
anyhow::bail!("unsupported transfer version: {version}");
}
let flags = data[6];
let file_size = u64::from_le_bytes(data[7..15].try_into()?);
let resume_offset = u64::from_le_bytes(data[15..23].try_into()?);
let fname_len = u32::from_le_bytes(data[23..27].try_into()?) as usize;
if data.len() < 27 + fname_len {
let expected = 27 + fname_len;
anyhow::bail!("header truncated: expected {expected} bytes, got {}", data.len());
}
let filename = String::from_utf8(data[27..27 + fname_len].to_vec())?;
let sha256 = if flags & FLAG_HASH != 0 {
let start = 27 + fname_len;
if data.len() < start + 32 {
anyhow::bail!("header truncated: sha256 expected");
}
let mut hash = [0u8; 32];
hash.copy_from_slice(&data[start..start + 32]);
Some(hash)
} else {
None
};
Ok(Self { file_size, resume_offset, filename, sha256, flags })
}
}
// ─── Sender ─────────────────────────────────────────────────────────────────
/// Send a file over an async Read+Write stream (yamux, TCP, etc.).
///
/// The `stream` parameter is any type implementing both `AsyncRead` and `AsyncWrite`.
/// Progress is reported back via `progress_tx`.
pub async fn send_file<S>(
mut stream: S,
filepath: &Path,
manager: &TransferManager,
transfer_id: &TransferId,
progress_tx: mpsc::Sender<TransferProgress>,
cancel: tokio_util::sync::CancellationToken,
) -> TransferResult
where
S: tokio::io::AsyncRead + tokio::io::AsyncWrite + Unpin + Send,
{
// Open and stat the file.
let file = match tokio::fs::File::open(filepath).await {
Ok(f) => f,
Err(e) => {
let err = format!("cannot open file: {e}");
manager.update_state(transfer_id, TransferState::Failed);
if let Some(mut t) = manager.get(transfer_id) { t.error = Some(err.clone()); }
return TransferResult::Failed { error: err };
}
};
let metadata = match file.metadata().await {
Ok(m) => m,
Err(e) => {
let err = format!("cannot stat file: {e}");
manager.update_state(transfer_id, TransferState::Failed);
return TransferResult::Failed { error: err };
}
};
let file_size = metadata.len();
// Compute SHA-256 while reading.
let filename = filepath.file_name()
.and_then(|n| n.to_str())
.unwrap_or("unknown")
.to_owned();
// Build and send header.
let header = TransferHeader {
file_size,
resume_offset: 0,
filename: filename.clone(),
sha256: None, // We'll send the hash after data in a footer.
flags: 0,
};
if let Err(e) = send_header(&mut stream, &header).await {
manager.update_state(transfer_id, TransferState::Failed);
return TransferResult::Failed { error: format!("failed to send header: {e}") };
}
manager.update_state(transfer_id, TransferState::Active);
info!(%transfer_id, %filename, file_size, "File send started");
// Stream file data with a large buffer for near-zero-copy throughput.
let mut reader = tokio::io::BufReader::with_capacity(BUFFER_SIZE, file);
let mut hasher = Sha256::new();
let mut buf = vec![0u8; BUFFER_SIZE];
let mut bytes_sent: u64 = 0;
let started = std::time::Instant::now();
loop {
tokio::select! {
_ = cancel.cancelled() => {
manager.update_state(transfer_id, TransferState::Cancelled);
info!(%transfer_id, "Send cancelled");
return TransferResult::Cancelled;
}
result = reader.read(&mut buf) => {
match result {
Ok(0) => break, // EOF
Ok(n) => {
hasher.update(&buf[..n]);
if let Err(e) = stream.write_all(&buf[..n]).await {
manager.update_state(transfer_id, TransferState::Failed);
return TransferResult::Failed { error: format!("write error: {e}") };
}
if let Err(e) = stream.flush().await {
manager.update_state(transfer_id, TransferState::Failed);
return TransferResult::Failed { error: format!("flush error: {e}") };
}
bytes_sent += n as u64;
manager.update_progress(transfer_id, bytes_sent);
// Throttle progress updates to ~4 Hz.
if bytes_sent % (BUFFER_SIZE as u64 * 4) < n as u64 {
let elapsed = started.elapsed().as_secs_f64();
let bps = if elapsed > 0.0 { bytes_sent as f64 / elapsed } else { 0.0 };
let eta = if bps > 0.0 { Some((file_size - bytes_sent) as f64 / bps) } else { None };
let _ = progress_tx.send(TransferProgress {
id: transfer_id.clone(), bytes_transferred: bytes_sent, total_bytes: file_size,
bytes_per_sec: bps, eta_secs: eta, hash_verified: false, final_hash: None,
}).await;
}
}
Err(e) => {
manager.update_state(transfer_id, TransferState::Failed);
return TransferResult::Failed { error: format!("read error: {e}") };
}
}
}
}
}
// Send SHA-256 footer (32 bytes) so the receiver can verify.
let hash_bytes = hasher.finalize();
if let Err(e) = stream.write_all(&hash_bytes).await {
manager.update_state(transfer_id, TransferState::Failed);
return TransferResult::Failed { error: format!("failed to send hash: {e}") };
}
if let Err(e) = stream.flush().await {
manager.update_state(transfer_id, TransferState::Failed);
return TransferResult::Failed { error: format!("flush after hash: {e}") };
}
let hash_hex = format!("{hash_bytes:x}");
manager.update_state(transfer_id, TransferState::Complete);
// Final progress with hash.
let elapsed = started.elapsed().as_secs_f64();
let _ = progress_tx.send(TransferProgress {
id: transfer_id.clone(), bytes_transferred: file_size, total_bytes: file_size,
bytes_per_sec: file_size as f64 / elapsed.max(0.001), eta_secs: Some(0.0),
hash_verified: true, final_hash: Some(hash_hex.clone()),
}).await;
info!(%transfer_id, %filename, %hash_hex, elapsed_secs = elapsed, "File send complete");
TransferResult::Completed { hash: hash_hex }
}
// ─── Receiver ───────────────────────────────────────────────────────────────
/// Receive a file from an async Read+Write stream.
///
/// Writes to `save_path.partial` during transfer, then atomically renames
/// to `save_path` on successful completion and hash verification.
pub async fn receive_file<S>(
mut stream: S,
save_dir: &Path,
manager: &TransferManager,
transfer_id: &TransferId,
progress_tx: mpsc::Sender<TransferProgress>,
cancel: tokio_util::sync::CancellationToken,
) -> TransferResult
where
S: tokio::io::AsyncRead + tokio::io::AsyncWrite + Unpin + Send,
{
// Read header.
let header = match read_header(&mut stream).await {
Ok(h) => h,
Err(e) => {
manager.update_state(transfer_id, TransferState::Failed);
return TransferResult::Failed { error: format!("failed to read header: {e}") };
}
};
let save_path = PathBuf::from(save_dir).join(&header.filename);
let partial_path = {
let mut p = save_path.clone();
let name = p.file_name().unwrap_or_default();
let mut name_str = name.to_string_lossy().into_owned();
name_str.push_str(".partial");
p.set_file_name(name_str);
p
};
// Open output file. If resuming, seek to offset.
// Note: use write(true) not append(true) — append mode and seek() have
// platform-dependent interaction (see issue N-3.2).
let mut file = match tokio::fs::OpenOptions::new()
.create(true)
.write(true)
.open(&partial_path).await
{
Ok(f) => f,
Err(e) => {
manager.update_state(transfer_id, TransferState::Failed);
return TransferResult::Failed { error: format!("cannot create output file: {e}") };
}
};
if header.resume_offset > 0 {
if let Err(e) = file.seek(std::io::SeekFrom::Start(header.resume_offset)).await {
manager.update_state(transfer_id, TransferState::Failed);
return TransferResult::Failed { error: format!("seek failed: {e}") };
}
}
manager.update_state(transfer_id, TransferState::Active);
info!(%transfer_id, filename = %header.filename, size = header.file_size, "File receive started");
let mut hasher = Sha256::new();
let mut buf = vec![0u8; BUFFER_SIZE];
let mut bytes_received: u64 = header.resume_offset;
let remaining = header.file_size.saturating_sub(header.resume_offset);
let started = std::time::Instant::now();
// We need to read exactly `remaining` bytes of file data, then 32 bytes of hash.
let total_to_read = remaining + 32; // file data + SHA-256 footer
let file_data_end = remaining;
// Buffer to capture the sender's 32-byte SHA-256 footer for verification.
let mut sender_hash_footer: [u8; 32] = [0u8; 32];
let mut footer_captured: bool = false;
while bytes_received < total_to_read {
let to_read = std::cmp::min(
(total_to_read - bytes_received) as usize,
BUFFER_SIZE,
);
tokio::select! {
_ = cancel.cancelled() => {
manager.update_state(transfer_id, TransferState::Cancelled);
info!(%transfer_id, "Receive cancelled at {} bytes", bytes_received);
return TransferResult::Cancelled;
}
result = stream.read(&mut buf[..to_read]) => {
match result {
Ok(0) => {
manager.update_state(transfer_id, TransferState::Failed);
return TransferResult::Failed { error: "unexpected EOF from sender".into() };
}
Ok(n) => {
let data = &buf[..n];
let current_file_pos = bytes_received;
if current_file_pos < file_data_end {
// Still reading file data.
let file_chunk_end = std::cmp::min(current_file_pos + n as u64, file_data_end);
let file_chunk_len = (file_chunk_end - current_file_pos) as usize;
hasher.update(&data[..file_chunk_len]);
if let Err(e) = file.write_all(&data[..file_chunk_len]).await {
manager.update_state(transfer_id, TransferState::Failed);
return TransferResult::Failed { error: format!("write error: {e}") };
}
// This chunk may span into the footer region.
// Capture any trailing bytes that fall in [file_data_end, total_to_read).
let footer_start_in_chunk = file_data_end.saturating_sub(current_file_pos) as usize;
if footer_start_in_chunk < n {
let footer_bytes_in_chunk = n - footer_start_in_chunk;
let footer_offset = (current_file_pos + file_chunk_len as u64 - file_data_end) as usize;
let copy_len = std::cmp::min(footer_bytes_in_chunk, 32 - footer_offset);
sender_hash_footer[footer_offset..footer_offset + copy_len]
.copy_from_slice(&data[footer_start_in_chunk..footer_start_in_chunk + copy_len]);
if footer_offset + copy_len >= 32 {
footer_captured = true;
}
}
} else {
// Entirely in the footer region.
let footer_offset = (current_file_pos - file_data_end) as usize;
let copy_len = std::cmp::min(n, 32 - footer_offset);
if copy_len > 0 {
sender_hash_footer[footer_offset..footer_offset + copy_len]
.copy_from_slice(&data[..copy_len]);
}
if footer_offset + copy_len >= 32 {
footer_captured = true;
}
}
bytes_received += n as u64;
let file_bytes_done = bytes_received.min(file_data_end);
manager.update_progress(transfer_id, file_bytes_done + header.resume_offset);
// Throttled progress.
if file_bytes_done % (BUFFER_SIZE as u64 * 4) < n as u64 {
let elapsed = started.elapsed().as_secs_f64();
let bps = if elapsed > 0.0 { file_bytes_done as f64 / elapsed } else { 0.0 };
let eta = if bps > 0.0 { Some((file_data_end - file_bytes_done) as f64 / bps) } else { None };
let _ = progress_tx.send(TransferProgress {
id: transfer_id.clone(), bytes_transferred: file_bytes_done + header.resume_offset,
total_bytes: header.file_size, bytes_per_sec: bps, eta_secs: eta,
hash_verified: false, final_hash: None,
}).await;
}
}
Err(e) => {
manager.update_state(transfer_id, TransferState::Failed);
return TransferResult::Failed { error: format!("read error: {e}") };
}
}
}
}
}
// Flush file to disk before verifying.
if let Err(e) = file.flush().await {
manager.update_state(transfer_id, TransferState::Failed);
return TransferResult::Failed { error: format!("flush error: {e}") };
}
drop(file);
// The last 32 bytes received are the sender's SHA-256 hash.
// They were NOT included in our hasher (we stopped hashing at file_data_end).
// We need to compute our own hash and compare.
let our_hash = compute_file_hash(&partial_path).await;
// Atomic rename from .partial to final path.
if let Err(e) = tokio::fs::rename(&partial_path, &save_path).await {
manager.update_state(transfer_id, TransferState::Failed);
return TransferResult::Failed { error: format!("atomic rename failed: {e}") };
}
// Verify the sender's SHA-256 footer against our computed hash.
let hash_hex = our_hash.clone().unwrap_or_default();
let hash_verified = if let Some(ref computed_hex) = our_hash {
let sender_hex: String = sender_hash_footer.iter().map(|b| format!("{b:02x}")).collect();
if !footer_captured {
warn!(%transfer_id, "sender hash footer incomplete — cannot verify");
false
} else if sender_hex != *computed_hex {
error!(%transfer_id, expected = %sender_hex, actual = %computed_hex, "SHA-256 hash mismatch");
false
} else {
true
}
} else {
false
};
if !hash_verified && footer_captured {
manager.update_state(transfer_id, TransferState::Failed);
return TransferResult::Failed {
error: format!("SHA-256 hash mismatch: expected {}, got {}",
sender_hash_footer.iter().map(|b| format!("{b:02x}")).collect::<String>(),
hash_hex),
};
}
manager.update_state(transfer_id, TransferState::Complete);
let elapsed = started.elapsed().as_secs_f64();
let _ = progress_tx.send(TransferProgress {
id: transfer_id.clone(), bytes_transferred: header.file_size, total_bytes: header.file_size,
bytes_per_sec: header.file_size as f64 / elapsed.max(0.001), eta_secs: Some(0.0),
hash_verified, final_hash: our_hash,
}).await;
info!(%transfer_id, filename = %header.filename, hash_verified, elapsed_secs = elapsed, "File receive complete");
TransferResult::Completed { hash: hash_hex }
}
// ─── Helpers ─────────────────────────────────────────────────────────────────
async fn send_header<S: tokio::io::AsyncWrite + Unpin>(
stream: &mut S,
header: &TransferHeader,
) -> anyhow::Result<()> {
let bytes = header.to_bytes();
// Prefix with 4-byte big-endian header length so the receiver knows how much to read.
let len = (bytes.len() as u32).to_be_bytes();
stream.write_all(&len).await?;
stream.write_all(&bytes).await?;
stream.flush().await?;
Ok(())
}
async fn read_header<S: tokio::io::AsyncRead + Unpin>(
stream: &mut S,
) -> anyhow::Result<TransferHeader> {
// Read 4-byte BE header length.
let mut len_buf = [0u8; 4];
stream.read_exact(&mut len_buf).await?;
let header_len = u32::from_be_bytes(len_buf) as usize;
if header_len > 4096 {
anyhow::bail!("header too large: {header_len} bytes");
}
let mut header_buf = vec![0u8; header_len];
stream.read_exact(&mut header_buf).await?;
TransferHeader::from_bytes(&header_buf)
}
async fn compute_file_hash(path: &Path) -> Option<String> {
let mut file = tokio::fs::File::open(path).await.ok()?;
let mut hasher = Sha256::new();
let mut buf = vec![0u8; BUFFER_SIZE];
loop {
match file.read(&mut buf).await {
Ok(0) => break,
Ok(n) => hasher.update(&buf[..n]),
Err(_) => return None,
}
}
Some(format!("{:x}", hasher.finalize()))
}
/// Format a file transfer progress line for the TUI status area.
pub fn format_progress_bar(p: &TransferProgress, width: usize) -> String {
let pct = if p.total_bytes == 0 { 0.0 } else { p.bytes_transferred as f64 / p.total_bytes as f64 * 100.0 };
let filled = ((pct / 100.0) * ((width as f64) - 10.0).max(1.0)) as usize;
let bar: String = format!("{}{}", "".repeat(filled), "".repeat((width as usize).saturating_sub(filled + 10)));
let speed = format_speed(p.bytes_per_sec);
let eta = p.eta_secs.map_or("--:--".into(), |s| format_eta(s));
format!("{bar} {:5.1}% {} eta {}", pct, speed, eta)
}
fn format_speed(bps: f64) -> String {
if bps >= 1_073_741.824 { format!("{:.1} MiB/s", bps / 1_048_576.0) }
else if bps >= 1024.0 { format!("{:.1} KiB/s", bps / 1024.0) }
else { format!("{:.0} B/s", bps) }
}
fn format_eta(secs: f64) -> String {
let secs = secs as u64;
let h = secs / 3600;
let m = (secs % 3600) / 60;
let s = secs % 60;
if h > 0 { format!("{h}:{m:02}:{s:02}") } else { format!("{m}:{s:02}") }
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn header_roundtrip() {
let h = TransferHeader {
file_size: 1_048_576,
resume_offset: 0,
filename: "test.bin".into(),
sha256: None,
flags: 0,
};
let bytes = h.to_bytes();
let parsed = TransferHeader::from_bytes(&bytes).unwrap();
assert_eq!(parsed.filename, "test.bin");
assert_eq!(parsed.file_size, 1_048_576);
}
#[test]
fn header_with_hash() {
let mut hash = [0u8; 32];
hash[0] = 0xDE; hash[31] = 0xAD;
let h = TransferHeader {
file_size: 42,
resume_offset: 1024,
filename: "resume.dat".into(),
sha256: Some(hash),
flags: FLAG_RESUME,
};
let bytes = h.to_bytes();
let parsed = TransferHeader::from_bytes(&bytes).unwrap();
assert_eq!(parsed.filename, "resume.dat");
assert_eq!(parsed.resume_offset, 1024);
assert_eq!(parsed.sha256, Some(hash));
}
#[test]
fn header_bad_magic() {
let bad = vec![0x00, 0x00, 0x00, 0x00, 0x01, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00];
assert!(TransferHeader::from_bytes(&bad).is_err());
}
#[test]
fn format_progress() {
let p = TransferProgress {
id: "test".into(), bytes_transferred: 524_288, total_bytes: 1_048_576,
bytes_per_sec: 262_144.0, eta_secs: Some(2.0), hash_verified: false, final_hash: None,
};
let s = format_progress_bar(&p, 40);
assert!(s.contains("50.0%"));
}
#[tokio::test]
async fn send_receive_roundtrip() {
use tokio::io::duplex;
let (client, server) = duplex(65536);
// Create a temp file to send.
let tmp_dir = tempfile::tempdir().unwrap();
let src_path = tmp_dir.path().join("source.txt");
tokio::fs::write(&src_path, b"hello zero-copy world! this is test data for the transfer engine.").await.unwrap();
let save_dir = tmp_dir.path().to_path_buf();
let (tx, _rx) = mpsc::channel(16);
let mgr = TransferManager::new(tx);
let id = mgr.queue_send(ProtocolType::BitChat, "peer", &src_path).unwrap();
let cancel = tokio_util::sync::CancellationToken::new();
// Spawn sender.
let mgr_s = mgr.clone_ref();
let id_s = id.clone();
let (prog_tx_s, mut prog_rx) = mpsc::channel(16);
let cancel_s = cancel.clone();
let sender_handle = tokio::spawn(async move {
send_file(client, &src_path, &mgr_s, &id_s, prog_tx_s, cancel_s).await
});
// Spawn receiver.
let mgr_r = mgr;
let id_r = id.clone();
let (prog_tx_r, mut prog_rx_r) = mpsc::channel(16);
let recv_dir = save_dir.clone();
let receiver_handle = tokio::spawn(async move {
receive_file(server, &recv_dir, &mgr_r, &id_r, prog_tx_r, cancel).await
});
let send_result = sender_handle.await.unwrap();
let recv_result = receiver_handle.await.unwrap();
assert!(matches!(send_result, TransferResult::Completed { .. }));
assert!(matches!(recv_result, TransferResult::Completed { .. }));
// Verify the file exists and has correct content.
let dest = save_dir.join("source.txt");
let content = tokio::fs::read_to_string(&dest).await.unwrap();
assert!(content.contains("hello zero-copy world!"));
// Drain sender progress — verify the sender reports hash_verified.
let mut sender_hash_verified = false;
while let Some(p) = prog_rx.recv().await {
if p.hash_verified { sender_hash_verified = true; }
}
assert!(sender_hash_verified, "sender should report hash_verified on final progress");
// Drain receiver progress — verify the receiver reports hash_verified.
let mut receiver_hash_verified = false;
while let Some(p) = prog_rx_r.recv().await {
if p.hash_verified { receiver_hash_verified = true; }
}
assert!(receiver_hash_verified, "receiver should report hash_verified on final progress (C-2.1.3)");
}
}