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, use_proxy: bool, ) -> Result { 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, ) -> Result { // 先尝试 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::().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::().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::().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, segments: &[Segment], file_path: &Path, cancel: Arc, progress: &[Arc], limiter: Arc, 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, segments: &[Segment], file_path: &Path, cancel: Arc, progress: &[Arc], limiter: &Arc, ) -> 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> = 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 = 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, seg: &Segment, file_path: &Path, cancel: Arc, progress: Arc, limiter: Arc, 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, seg: &Segment, file_path: &Path, cancel: Arc, progress: Arc, limiter: Arc, ) -> 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 { // 优先从 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 { // 先尝试 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 { 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); } }