Skip to content
Open
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
268 changes: 267 additions & 1 deletion src-tauri/src/ureq.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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 {
Expand Down Expand Up @@ -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;
Expand Down Expand Up @@ -327,6 +332,8 @@ fn download_multi_thread(
cancel_flag: &Arc<AtomicBool>,
pause_flag: &Arc<AtomicBool>,
) -> anyhow::Result<()> {
// 暂停期间不要为新任务发起 HEAD / Range 探测请求。
wait_if_paused(cancel_flag, pause_flag)?;
let head = ureq::head(url)
.set("User-Agent", &user_agent())
.set("Accept", "*/*")
Expand Down Expand Up @@ -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<AtomicBool>, usize)>,
) -> (std::net::SocketAddr, std::thread::JoinHandle<()>) {
fn serve(
mut stream: TcpStream,
body_len: usize,
block: usize,
delay: Duration,
gate: Option<(Arc<AtomicBool>, 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];
Expand Down Expand Up @@ -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();
}
}
Loading