Files
Thing/src-tauri/src/download_engine/http_dl.rs
T

668 lines
23 KiB
Rust
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
use futures_util::StreamExt;
use reqwest::Client;
use std::collections::HashMap;
use std::fs::OpenOptions;
use std::io::SeekFrom;
use std::path::Path;
use std::sync::atomic::{AtomicBool, AtomicU64, Ordering};
use std::sync::Arc;
use tokio::io::{AsyncSeekExt, AsyncWriteExt};
use tokio::task::JoinSet;
use super::rate_limit::RateLimiter;
use super::task::{ProbeResult, Segment};
/// 请求超时:连接 + 响应头必须在 30s 内就绪(分段请求若无超时,
/// 服务器挂死时任务将永久卡在 Active,pause→resume 会出现新旧任务并发写同一临时文件)
const REQUEST_TIMEOUT: std::time::Duration = std::time::Duration::from_secs(30);
/// 分块读取停滞超时:30s 内无任何数据视为连接挂死,主动中断(配合取消标志及时退出)
const READ_STALL_TIMEOUT: std::time::Duration = std::time::Duration::from_secs(30);
/// HTTP/HTTPS 下载器
#[derive(Clone)]
pub struct HttpDownloader {
/// 默认客户端:尊重系统代理(reqwest 默认行为,mihomo 开启系统代理时经其转发)
system_client: Client,
/// 直连客户端:强制禁用系统代理(no_proxy)
direct_client: Client,
}
impl HttpDownloader {
pub fn new() -> Self {
let system_client = Client::builder()
.build()
.unwrap_or_else(|_| Client::new());
let direct_client = Client::builder()
// 强制直连:即使系统代理已开启,下载也不经过代理
.no_proxy()
.build()
.unwrap_or_else(|_| Client::new());
Self {
system_client,
direct_client,
}
}
/// 根据 use_proxy 选择客户端
fn client(&self, use_proxy: bool) -> &Client {
if use_proxy {
&self.system_client
} else {
&self.direct_client
}
}
/// 探测下载资源信息(大小、是否支持 Range、文件名)
/// 优先用 GET + Range: bytes=0-0(返回 206 + Content-Range),回退到 HEAD。
/// 代理降级:use_proxy=true 时先走系统代理,失败则回退 no_proxy 直连重试一次
pub async fn probe(
&self,
url: &str,
headers: &HashMap<String, String>,
use_proxy: bool,
) -> Result<ProbeResult, String> {
let first = self.client(use_proxy);
match self.probe_with_client(first, url, headers).await {
Ok(r) => return Ok(r),
Err(e) if use_proxy => {
let direct = self.client(false);
self.probe_with_client(direct, url, headers)
.await
.map_err(|e2| format!("代理探测失败({}),直连重试也失败({}", e, e2))
}
Err(e) => Err(e),
}
}
/// 用指定 client 执行探测(GET Range → 回退 HEAD),供代理降级复用
async fn probe_with_client(
&self,
client: &Client,
url: &str,
headers: &HashMap<String, String>,
) -> Result<ProbeResult, String> {
// 先尝试 Range 请求(能同时判断 Accept-Ranges 和获取大小)
let mut req = client
.get(url)
.header("Range", "bytes=0-0")
.header("User-Agent", "Thing-Download-Engine/1.0");
for (k, v) in headers {
req = req.header(k, v);
}
match tokio::time::timeout(REQUEST_TIMEOUT, req.send())
.await
.map_err(|_| "探测超时(30s 内未收到响应头)".to_string())?
{
Ok(resp) => {
let status = resp.status();
let headers_map = resp.headers().clone();
// 206 Partial Content → 支持 Range
if status.as_u16() == 206 {
let total_size = headers_map
.get("content-range")
.and_then(|v| v.to_str().ok())
.and_then(|v| v.split('/').nth(1))
.and_then(|v| v.trim().parse::<u64>().ok());
let supports_resume = true;
let filename = extract_filename(url, &headers_map);
return Ok(ProbeResult {
total_size,
supports_resume,
filename,
});
}
// 200 OK → 服务器不支持 Range 或忽略 Range 头
// 可能仍有 Content-Length
let total_size = headers_map
.get("content-length")
.and_then(|v| v.to_str().ok())
.and_then(|v| v.trim().parse::<u64>().ok());
// 检查 Accept-Ranges 头
let supports_resume = headers_map
.get("accept-ranges")
.and_then(|v| v.to_str().ok())
.map(|v| v.eq_ignore_ascii_case("bytes"))
.unwrap_or(false);
let filename = extract_filename(url, &headers_map);
Ok(ProbeResult {
total_size,
supports_resume,
filename,
})
}
Err(_) => {
// GET 失败,尝试 HEAD 作为回退
let mut head_req = client.head(url);
for (k, v) in headers {
head_req = head_req.header(k, v);
}
let resp = tokio::time::timeout(REQUEST_TIMEOUT, head_req.send())
.await
.map_err(|_| "探测超时(HEAD 30s 内未收到响应头)".to_string())?
.map_err(|e| format!("探测失败(GET 和 HEAD 均失败): {}", e))?;
let headers_map = resp.headers().clone();
let total_size = headers_map
.get("content-length")
.and_then(|v| v.to_str().ok())
.and_then(|v| v.trim().parse::<u64>().ok());
let supports_resume = headers_map
.get("accept-ranges")
.and_then(|v| v.to_str().ok())
.map(|v| v.eq_ignore_ascii_case("bytes"))
.unwrap_or(false);
let filename = extract_filename(url, &headers_map);
Ok(ProbeResult {
total_size,
supports_resume,
filename,
})
}
}
}
/// 执行多线程 Range 下载
///
/// - `segments`: 分段列表(已根据 probe 结果划分)
/// - `file_path`: 目标文件路径
/// - `cancel`: 取消标志
/// - `progress`: 每个分段的已下载字节(AtomicU64,与 segments 一一对应)
/// - `limiter`: 全局限速器
/// - `use_proxy`: 是否使用系统代理(false=强制直连)
pub async fn download(
&self,
url: &str,
headers: &HashMap<String, String>,
segments: &[Segment],
file_path: &Path,
cancel: Arc<AtomicBool>,
progress: &[Arc<AtomicU64>],
limiter: Arc<RateLimiter>,
use_proxy: bool,
) -> Result<(), String> {
// 代理降级:仅当 use_proxy=true 才有"走代理→失败回退直连"的意义。
// use_proxy=false 直接用直连客户端,无需回退。
// 注意:用户主动暂停/取消(返回"已取消")必须原样透传,不能触发代理回退,
// 否则会把"已取消"包装成"代理失败",导致引擎将其误判为错误而非暂停。
let first = self.client(use_proxy);
match self.download_with_client(first, url, headers, segments, file_path, cancel.clone(), progress, &limiter).await {
Ok(()) => return Ok(()),
Err(e) if use_proxy && e != "已取消" => {
// 回退直连重试(不继承 use_proxy,保证用 no_proxy 客户端)
let direct = self.client(false);
self.download_with_client(direct, url, headers, segments, file_path, cancel, progress, &limiter)
.await
.map_err(|e2| format!("代理下载失败({}),直连重试也失败({}", e, e2))
}
Err(e) => Err(e),
}
}
/// 用指定 client 执行下载(支持单线程与多线程分段),供代理降级复用
async fn download_with_client(
&self,
client: &Client,
url: &str,
headers: &HashMap<String, String>,
segments: &[Segment],
file_path: &Path,
cancel: Arc<AtomicBool>,
progress: &[Arc<AtomicU64>],
limiter: &Arc<RateLimiter>,
) -> Result<(), String> {
let total_size = segments.iter().map(|s| s.len()).sum();
// 预分配文件(若已知大小)
if total_size > 0 {
if let Some(file) = OpenOptions::new()
.create(true)
.write(true)
.open(file_path)
.ok()
{
let _ = file.set_len(total_size);
}
} else {
// 未知大小,确保文件存在
let _ = OpenOptions::new()
.create(true)
.write(true)
.truncate(true)
.open(file_path);
}
if segments.len() <= 1 {
// 单线程下载(不支持 Range 或文件太小)
let seg = &segments[0];
let prog = &progress[0];
self.download_segment(url, headers, seg, file_path, cancel.clone(), prog.clone(), limiter.clone(), client)
.await?;
return Ok(());
}
// 多线程下载:每个分段一个 tokio task
let mut join_set: JoinSet<Result<(), String>> = JoinSet::new();
for seg in segments {
let url = url.to_string();
let headers = headers.clone();
let seg = seg.clone();
let cancel = cancel.clone();
// 找到对应的 progress atomic
let prog = progress
.get(seg.index as usize)
.cloned()
.unwrap_or_default();
let limiter = limiter.clone();
let file_path = file_path.to_path_buf();
let client = client.clone();
join_set.spawn(async move {
download_segment_with_client(
&client,
&url,
&headers,
&seg,
&file_path,
cancel,
prog,
limiter,
)
.await
});
}
// 等待所有分段完成
let mut first_error: Option<String> = None;
while let Some(res) = join_set.join_next().await {
match res {
Ok(Ok(())) => {}
Ok(Err(e)) => {
if first_error.is_none() {
first_error = Some(e);
}
// 一个分段失败,取消其他分段
cancel.store(true, Ordering::SeqCst);
}
Err(e) => {
if first_error.is_none() {
first_error = Some(format!("分段任务异常: {}", e));
}
cancel.store(true, Ordering::SeqCst);
}
}
}
if let Some(e) = first_error {
return Err(e);
}
Ok(())
}
/// 下载单个分段(实例方法,用于单线程场景)
async fn download_segment(
&self,
url: &str,
headers: &HashMap<String, String>,
seg: &Segment,
file_path: &Path,
cancel: Arc<AtomicBool>,
progress: Arc<AtomicU64>,
limiter: Arc<RateLimiter>,
client: &Client,
) -> Result<(), String> {
download_segment_with_client(
client,
url,
headers,
seg,
file_path,
cancel,
progress,
limiter,
)
.await
}
}
/// 下载单个分段的核心逻辑(可被多线程 task 调用)
async fn download_segment_with_client(
client: &Client,
url: &str,
headers: &HashMap<String, String>,
seg: &Segment,
file_path: &Path,
cancel: Arc<AtomicBool>,
progress: Arc<AtomicU64>,
limiter: Arc<RateLimiter>,
) -> Result<(), String> {
// 如果分段已完成,直接返回
if seg.is_done() {
return Ok(());
}
// 计算本次需要下载的范围(从已下载位置继续)
let range_start = seg.start + seg.completed;
let range_end = seg.end;
// start=0, end=0 表示未知大小(不支持 Range 或未探测到大小),不发送 Range 头
let unknown_size = seg.start == 0 && seg.end == 0;
// 打开文件并 seek 到正确位置
let mut file = tokio::fs::OpenOptions::new()
.create(true)
.write(true)
.read(false)
.open(file_path)
.await
.map_err(|e| format!("打开文件失败: {}", e))?;
if !unknown_size {
file.seek(SeekFrom::Start(range_start))
.await
.map_err(|e| format!("文件定位失败: {}", e))?;
}
// 构造请求
let mut req = client
.get(url)
.header("User-Agent", "Thing-Download-Engine/1.0");
if !unknown_size {
req = req.header("Range", format!("bytes={}-{}", range_start, range_end));
}
for (k, v) in headers {
req = req.header(k, v);
}
let resp = tokio::time::timeout(REQUEST_TIMEOUT, req.send())
.await
.map_err(|_| "请求超时(30s 内未收到响应头)".to_string())?
.map_err(|e| format!("请求失败: {}", e))?;
let status = resp.status();
if unknown_size {
// 未知大小:未发送 Range 头,接受任意 2xx
if !status.is_success() {
return Err(format!("服务器返回 HTTP {}", status));
}
} else {
// 已发送 Range 头:必须返回 206。若服务器忽略 Range 返回 200 全文,
// 按 range_start 偏移写入会错位 → 静默损坏文件;此处直接中断。
if status.as_u16() != 206 {
return Err(format!(
"服务器未按分段请求响应(期望 206,实际 {}),已中断以避免文件损坏",
status
));
}
}
// 流式读取并写入文件
let mut stream = resp.bytes_stream();
let mut buf = Vec::with_capacity(64 * 1024);
let mut local_completed = seg.completed;
loop {
if cancel.load(Ordering::SeqCst) {
return Err("已取消".to_string());
}
// 停滞超时:30s 无数据即中断,确保 cancel 标志能及时被感知(配合代际句柄防并发写)
match tokio::time::timeout(READ_STALL_TIMEOUT, stream.next()).await {
Ok(Some(Ok(chunk))) => {
buf.extend_from_slice(&chunk);
// 接收到数据立即更新进度(避免监控周期内进度无变化导致速度显示为 0)
local_completed += chunk.len() as u64;
progress.store(local_completed, Ordering::Relaxed);
// 积累到 64KB 再写入(减少 I/O 次数)
if buf.len() >= 64 * 1024 {
file.write_all(&buf)
.await
.map_err(|e| format!("写入文件失败: {}", e))?;
// 限速
limiter.consume(buf.len() as u64).await;
buf.clear();
}
}
Ok(Some(Err(e))) => {
return Err(format!("读取数据失败: {}", e));
}
Ok(None) => {
// 流结束,写入剩余数据
if !buf.is_empty() {
file.write_all(&buf)
.await
.map_err(|e| format!("写入文件失败: {}", e))?;
// local_completed 和 progress 已在接收 chunk 时更新
limiter.consume(buf.len() as u64).await;
buf.clear();
}
// 校验:已知大小的分段若流提前结束(收到的字节数不足分段长度),
// 说明服务器提前断开或返回不完整内容,不能标记为完成,否则文件会被截断
if !unknown_size && local_completed < seg.len() {
return Err(format!(
"文件不完整:已接收 {} / {} 字节,服务器提前结束连接",
local_completed,
seg.len()
));
}
break;
}
Err(_) => {
return Err("读取超时(30s 无数据,已中断下载)".to_string());
}
}
}
file.flush()
.await
.map_err(|e| format!("刷新文件失败: {}", e))?;
Ok(())
}
/// 从 Content-Disposition 头或 URL 提取文件名
fn extract_filename(url: &str, headers: &reqwest::header::HeaderMap) -> Option<String> {
// 优先从 Content-Disposition 提取
if let Some(cd) = headers.get("content-disposition") {
if let Ok(s) = cd.to_str() {
// filename*=UTF-8''xxx 或 filename="xxx"
if let Some(name) = parse_content_disposition(s) {
return Some(name);
}
}
}
// 从 URL 提取
if let Ok(u) = url::Url::parse(url) {
let path = u.path();
if let Some(name) = path.rsplit('/').next() {
if !name.is_empty() {
return Some(
percent_decode(name),
);
}
}
}
None
}
/// 解析 Content-Disposition 头中的文件名
fn parse_content_disposition(s: &str) -> Option<String> {
// 先尝试 filename*=UTF-8''xxx
if let Some(idx) = s.find("filename*=") {
let rest = &s[idx + "filename*=".len()..];
// 格式: UTF-8''xxx 或 ISO-8859-1''xxx
if let Some(start) = rest.find("''") {
let encoded = &rest[start + 2..];
// 取到分号或结尾
let encoded = encoded.split(';').next().unwrap_or(encoded).trim();
return Some(percent_decode(encoded.trim_matches('"')));
}
}
// 再尝试 filename="xxx" 或 filename=xxx
if let Some(idx) = s.find("filename=") {
let rest = &s[idx + "filename=".len()..];
let rest = rest.trim();
if rest.starts_with('"') {
let end = rest[1..].find('"').unwrap_or(rest.len() - 1);
return Some(rest[1..1 + end].to_string());
} else {
let name = rest.split(';').next().unwrap_or(rest).trim();
if !name.is_empty() {
return Some(name.to_string());
}
}
}
None
}
/// URL 百分号解码
fn percent_decode(s: &str) -> String {
let mut result = String::with_capacity(s.len());
let bytes = s.as_bytes();
let mut i = 0;
while i < bytes.len() {
if bytes[i] == b'%' && i + 2 < bytes.len() {
if let Ok(byte) = u8::from_str_radix(
std::str::from_utf8(&bytes[i + 1..i + 3]).unwrap_or(""),
16,
) {
result.push(byte as char);
i += 3;
continue;
}
}
result.push(bytes[i] as char);
i += 1;
}
result
}
/// 将文件大小划分为 N 个分段
pub fn split_segments(total_size: u64, num_connections: u32) -> Vec<Segment> {
if total_size == 0 || num_connections == 0 {
// 空文件或未指定连接数:单段覆盖整个文件(end 按 total_size 推导,不能硬编码 0
return vec![Segment {
index: 0,
start: 0,
end: total_size.saturating_sub(1),
completed: 0,
}];
}
// 每个分段至少 1MB,否则减少线程数
let min_segment_size = 1024 * 1024;
let actual_connections = {
let max_by_size = (total_size / min_segment_size).max(1) as u32;
num_connections.min(max_by_size)
};
let segment_size = total_size / actual_connections as u64;
let mut segments = Vec::with_capacity(actual_connections as usize);
let mut offset = 0u64;
for i in 0..actual_connections {
let start = offset;
let end = if i == actual_connections - 1 {
total_size - 1
} else {
offset + segment_size - 1
};
segments.push(Segment {
index: i,
start,
end,
completed: 0,
});
offset = end + 1;
}
segments
}
#[cfg(test)]
mod split_segments_tests {
use super::*;
/// 分段必须完整覆盖 [0, total_size),且相互连续无重叠
fn assert_contiguous(segments: &[Segment], total_size: u64) {
assert!(!segments.is_empty());
let mut prev_end: i64 = -1;
for seg in segments {
assert_eq!(seg.start as i64, prev_end + 1, "分段不连续");
assert!(seg.end >= seg.start, "分段 start > end");
prev_end = seg.end as i64;
}
assert_eq!(
segments.last().unwrap().end,
total_size - 1,
"末段未覆盖文件末尾"
);
}
#[test]
fn zero_size_returns_single_segment() {
let segs = split_segments(0, 4);
assert_eq!(segs.len(), 1);
assert_eq!(segs[0].start, 0);
assert_eq!(segs[0].end, 0);
}
#[test]
fn zero_connections_covers_whole_file() {
let segs = split_segments(1024, 0);
assert_eq!(segs.len(), 1);
assert_eq!(segs[0].start, 0);
assert_eq!(segs[0].end, 1023);
}
#[test]
fn divides_evenly_with_contiguous_coverage() {
// 10MB / 4 连接 → 4 段完整覆盖
let total = 10 * 1024 * 1024;
let segs = split_segments(total, 4);
assert_eq!(segs.len(), 4);
assert_contiguous(&segs, total);
}
#[test]
fn clamps_connections_by_min_segment_size() {
// 2MB 文件请求 8 连接 → 受 1MB 最小分段限制,实际 ≤ 2 段
let total = 2 * 1024 * 1024;
let segs = split_segments(total, 8);
assert!(segs.len() <= 2, "连接数未按最小分段收敛: {}", segs.len());
assert_contiguous(&segs, total);
}
#[test]
fn respects_requested_connection_count() {
// 大文件按请求连接数切分
let total = 100 * 1024 * 1024;
let segs = split_segments(total, 3);
assert_eq!(segs.len(), 3);
assert_contiguous(&segs, total);
// 每段大小均匀
for seg in &segs {
let seg_len = seg.end - seg.start + 1;
assert!(
seg_len >= total / 3,
"分段大小不均: {} 段只有 {} 字节",
seg.index,
seg_len
);
}
}
#[test]
fn tiny_file_single_segment() {
// 小于 1MB 的文件始终单段
let total = 100;
let segs = split_segments(total, 4);
assert_eq!(segs.len(), 1);
assert_contiguous(&segs, total);
}
}