From bb9023db024ec4ed29dbbddc0627b1f97afb2238 Mon Sep 17 00:00:00 2001 From: std-external <201375710+std-external@users.noreply.github.com> Date: Sat, 10 Oct 2026 00:28:08 +0800 Subject: [PATCH] =?UTF-8?q?fix(download):=20=E6=9A=82=E5=81=9C=E5=90=8E?= =?UTF-8?q?=E4=B8=8D=E5=86=8D=E8=90=BD=E7=9B=98=E6=AD=A3=E5=9C=A8=E8=AF=BB?= =?UTF-8?q?=E5=8F=96=E7=9A=84=E5=88=86=E5=9D=97?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 下载线程只在每次 read 之前检查暂停标志,而 read 本身无法中断:暂停时 正在读取的那一块读完后仍会被写入文件并上报进度,8 线程分段下载时会 ×8,用户看到的就是「按了暂停还在下载」。 现在读到的一块先留在内存里,等恢复后再落盘,暂停后文件大小与进度会 立刻停住;同时多线程下载在暂停期间不再发起 HEAD / Range 探测请求。 --- src-tauri/src/ureq.rs | 268 +++++++++++++++++++++++++++++++++++++++++- 1 file changed, 267 insertions(+), 1 deletion(-) diff --git a/src-tauri/src/ureq.rs b/src-tauri/src/ureq.rs index e8f84d4..61577ff 100644 --- a/src-tauri/src/ureq.rs +++ b/src-tauri/src/ureq.rs @@ -206,6 +206,9 @@ fn download_single( Ok(n) => n, Err(error) => break Some(error.into()), }; + // 暂停时把刚读到的一块先留在内存里,恢复后再落盘:正在阻塞的 read + // 无法中断,如果这里直接写下去,用户按下暂停后进度和文件还会继续涨。 + wait_if_paused(cancel_flag, pause_flag)?; file.write_all(&buffer[..n])?; downloaded += n as u64; if total_size > 0 { @@ -295,6 +298,8 @@ fn download_range_part( Ok(n) => n, Err(error) => break Some(error.into()), }; + // 与单连接下载一致:暂停后先不落盘,避免进度在暂停后继续增长。 + wait_if_paused(cancel_flag, pause_flag)?; file.write_all(&buffer[..n])?; offset += n as u64; *downloaded_bytes.lock().unwrap() += n as u64; @@ -327,6 +332,8 @@ fn download_multi_thread( cancel_flag: &Arc, pause_flag: &Arc, ) -> anyhow::Result<()> { + // 暂停期间不要为新任务发起 HEAD / Range 探测请求。 + wait_if_paused(cancel_flag, pause_flag)?; let head = ureq::head(url) .set("User-Agent", &user_agent()) .set("Accept", "*/*") @@ -528,10 +535,161 @@ pub fn download_file_with_progress( mod tests { use std::io::{Read, Write}; use std::net::{TcpListener, TcpStream}; - use std::sync::{Arc, Mutex, atomic::AtomicBool}; + use std::sync::{ + Arc, Mutex, + atomic::{AtomicBool, Ordering}, + }; + use std::time::{Duration, Instant}; use super::{download_file_to_path_with_progress, download_single, parse_content_range}; + fn temp_output(tag: &str, port: u16) -> std::path::PathBuf { + std::env::temp_dir().join(format!("celemod-{tag}-{}-{port}.tmp", std::process::id())) + } + + fn file_len(path: &std::path::Path) -> u64 { + std::fs::metadata(path).map(|meta| meta.len()).unwrap_or(0) + } + + fn wait_for_bytes(path: &std::path::Path, target: u64) { + let deadline = Instant::now() + Duration::from_secs(30); + while file_len(path) < target && Instant::now() < deadline { + std::thread::sleep(Duration::from_millis(2)); + } + assert!( + file_len(path) >= target, + "download never reached {target} bytes" + ); + } + + /// 限速流式响应的服务器,用来观察暂停前后到底还写了多少字节。 + fn spawn_streaming_server( + body_len: usize, + block: usize, + delay: Duration, + ) -> (std::net::SocketAddr, std::thread::JoinHandle<()>) { + let listener = TcpListener::bind("127.0.0.1:0").unwrap(); + let address = listener.local_addr().unwrap(); + let handle = std::thread::spawn(move || { + let (mut stream, _) = listener.accept().unwrap(); + let _ = stream.set_nodelay(true); + let _request = read_request(&mut stream); + let _ = stream.write_all( + format!( + "HTTP/1.1 200 OK\r\nContent-Length: {body_len}\r\nAccept-Ranges: none\r\nConnection: close\r\n\r\n" + ) + .as_bytes(), + ); + let chunk = vec![b'x'; block]; + let mut written = 0usize; + while written < body_len { + let length = block.min(body_len - written); + if stream.write_all(&chunk[..length]).is_err() { + break; + } + written += length; + if delay > Duration::ZERO { + std::thread::sleep(delay); + } + } + std::thread::sleep(Duration::from_millis(300)); + }); + (address, handle) + } + + /// 支持 HEAD、单字节 Range 探测和多段 Range 的服务器(多线程下载用)。 + /// `gate` 为 `Some` 时,每个分段先只发 `gate_prefix` 字节,等到 gate 打开再补发, + /// 这样客户端一定停在一次阻塞的 read 上。 + fn spawn_range_server( + body_len: usize, + block: usize, + delay: Duration, + gate: Option<(Arc, usize)>, + ) -> (std::net::SocketAddr, std::thread::JoinHandle<()>) { + fn serve( + mut stream: TcpStream, + body_len: usize, + block: usize, + delay: Duration, + gate: Option<(Arc, usize)>, + ) { + let _ = stream.set_nodelay(true); + let request = read_request(&mut stream); + let lower = request.to_ascii_lowercase(); + if lower.starts_with("head ") { + let _ = stream.write_all( + format!( + "HTTP/1.1 200 OK\r\nContent-Length: {body_len}\r\nAccept-Ranges: bytes\r\nConnection: close\r\n\r\n" + ) + .as_bytes(), + ); + return; + } + let Some(range) = lower + .lines() + .find_map(|line| line.strip_prefix("range: bytes=")) + .map(str::to_string) + else { + return; + }; + let (start, end) = range.split_once('-').unwrap(); + let start: usize = start.trim().parse().unwrap(); + let end: usize = end.trim().parse().unwrap_or(body_len - 1); + let _ = stream.write_all( + format!( + "HTTP/1.1 206 Partial Content\r\nContent-Length: {}\r\nContent-Range: bytes {start}-{end}/{body_len}\r\nConnection: close\r\n\r\n", + end - start + 1 + ) + .as_bytes(), + ); + let chunk = vec![b'y'; block]; + let prefix = gate + .as_ref() + .map(|(_, prefix)| *prefix) + .unwrap_or(usize::MAX); + let mut written = start; + while written <= end { + if written.saturating_sub(start) >= prefix { + let deadline = Instant::now() + Duration::from_secs(30); + while let Some((flag, _)) = gate.as_ref() + && !flag.load(Ordering::Relaxed) + && Instant::now() < deadline + { + std::thread::sleep(Duration::from_millis(10)); + } + } + let length = block.min(end - written + 1); + if stream.write_all(&chunk[..length]).is_err() { + break; + } + written += length; + if delay > Duration::ZERO { + std::thread::sleep(delay); + } + } + } + + let listener = TcpListener::bind("127.0.0.1:0").unwrap(); + let address = listener.local_addr().unwrap(); + let handle = std::thread::spawn(move || { + let mut workers = Vec::new(); + // HEAD + 单字节探测 + 最多 8 个分段连接。 + for _ in 0..10 { + let Ok((stream, _)) = listener.accept() else { + break; + }; + let gate = gate.clone(); + workers.push(std::thread::spawn(move || { + serve(stream, body_len, block, delay, gate); + })); + } + for worker in workers { + let _ = worker.join(); + } + }); + (address, handle) + } + fn read_request(stream: &mut TcpStream) -> String { let mut request = Vec::new(); let mut buffer = [0u8; 1024]; @@ -661,4 +819,112 @@ mod tests { drop(requests); std::fs::remove_file(output_path).unwrap(); } + + /// 按下暂停后不应再有新的字节落盘:正在阻塞的 read 无法中断,所以那一刻 + /// 读到的一块必须留到恢复之后再写,否则「暂停」看起来还在继续下载。 + #[test] + fn pause_stops_writes_until_resumed() { + let body_len = 32 * 1024 * 1024; + let (address, server) = + spawn_streaming_server(body_len, 64 * 1024, Duration::from_millis(2)); + let output = temp_output("pause-writes", address.port()); + let cancel = Arc::new(AtomicBool::new(false)); + let pause = Arc::new(AtomicBool::new(false)); + let worker_pause = Arc::clone(&pause); + let worker_cancel = Arc::clone(&cancel); + let worker_output = output.clone(); + let url = format!("http://{address}/mod.zip"); + let worker = std::thread::spawn(move || { + download_file_to_path_with_progress( + &url, + worker_output.to_string_lossy().as_ref(), + &mut |_| {}, + false, + &worker_cancel, + &worker_pause, + ) + }); + + wait_for_bytes(&output, 4 * 1024 * 1024); + let before = file_len(&output); + pause.store(true, Ordering::Relaxed); + std::thread::sleep(Duration::from_millis(700)); + let after = file_len(&output); + assert_eq!( + before, + after, + "paused download still wrote {} bytes", + after - before + ); + + pause.store(false, Ordering::Relaxed); + worker.join().unwrap().unwrap(); + assert_eq!(file_len(&output), body_len as u64); + std::fs::remove_file(&output).ok(); + server.join().ok(); + } + + /// 多线程分段下载:暂停之后,即使卡住的连接刚好补发了数据也不能再落盘。 + #[test] + fn multi_thread_pause_stops_writes_until_resumed() { + const PREFIX: usize = 64 * 1024; + let body_len = 64 * 1024 * 1024; + let gate = Arc::new(AtomicBool::new(false)); + let (address, server) = spawn_range_server( + body_len, + 64 * 1024, + Duration::ZERO, + Some((Arc::clone(&gate), PREFIX)), + ); + let output = temp_output("multi-pause-writes", address.port()); + let cancel = Arc::new(AtomicBool::new(false)); + let pause = Arc::new(AtomicBool::new(false)); + let observed = Arc::new(Mutex::new(0u64)); + let worker_pause = Arc::clone(&pause); + let worker_cancel = Arc::clone(&cancel); + let worker_output = output.clone(); + let worker_observed = Arc::clone(&observed); + let url = format!("http://{address}/mod.zip"); + let worker = std::thread::spawn(move || { + download_file_to_path_with_progress( + &url, + worker_output.to_string_lossy().as_ref(), + &mut |info| *worker_observed.lock().unwrap() = info.downloaded_bytes, + true, + &worker_cancel, + &worker_pause, + ) + }); + + // 每个分段都只发完前缀,客户端此时停在阻塞的 read 上。 + let deadline = Instant::now() + Duration::from_secs(30); + while *observed.lock().unwrap() < (8 * PREFIX) as u64 && Instant::now() < deadline { + std::thread::sleep(Duration::from_millis(10)); + } + std::thread::sleep(Duration::from_millis(400)); + pause.store(true, Ordering::Relaxed); + std::thread::sleep(Duration::from_millis(300)); + let before = *observed.lock().unwrap(); + assert!( + before >= (8 * PREFIX) as u64, + "expected the eight prefixes ({}) to land, got {before}", + 8 * PREFIX + ); + + gate.store(true, Ordering::Relaxed); + std::thread::sleep(Duration::from_millis(1200)); + let after = *observed.lock().unwrap(); + assert_eq!( + before, + after, + "paused multi-thread download still wrote {} bytes", + after - before + ); + + pause.store(false, Ordering::Relaxed); + worker.join().unwrap().unwrap(); + assert_eq!(*observed.lock().unwrap(), body_len as u64); + std::fs::remove_file(&output).ok(); + server.join().ok(); + } }