下载非内核

This commit is contained in:
zhongluofeng
2026-07-22 18:26:26 +08:00
parent 1e31ee8da9
commit f7d3c13f35
24 changed files with 2609 additions and 2532 deletions
+475
View File
@@ -0,0 +1,475 @@
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};
/// HTTP/HTTPS 下载器
#[derive(Clone)]
pub struct HttpDownloader {
client: Client,
}
impl HttpDownloader {
pub fn new() -> Self {
let client = Client::builder()
.build()
.unwrap_or_else(|_| Client::new());
Self { client }
}
/// 探测下载资源信息(大小、是否支持 Range、文件名)
/// 优先用 GET + Range: bytes=0-0(返回 206 + Content-Range),回退到 HEAD
pub async fn probe(
&self,
url: &str,
headers: &HashMap<String, String>,
) -> Result<ProbeResult, String> {
// 先尝试 Range 请求(能同时判断 Accept-Ranges 和获取大小)
let mut req = self
.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 req.send().await {
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 = self.client.head(url);
for (k, v) in headers {
head_req = head_req.header(k, v);
}
let resp = head_req
.send()
.await
.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`: 全局限速器
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>,
) -> 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)
.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 = self.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>,
) -> Result<(), String> {
download_segment_with_client(
&self.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 = req
.send()
.await
.map_err(|e| format!("请求失败: {}", e))?;
let status = resp.status();
if !status.is_success() && status.as_u16() != 206 {
return Err(format!("服务器返回 HTTP {}", 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());
}
match stream.next().await {
Some(Ok(chunk)) => {
buf.extend_from_slice(&chunk);
// 积累到 64KB 再写入(减少 I/O 次数)
if buf.len() >= 64 * 1024 {
file.write_all(&buf)
.await
.map_err(|e| format!("写入文件失败: {}", e))?;
local_completed += buf.len() as u64;
progress.store(local_completed, Ordering::Relaxed);
// 限速
limiter.consume(buf.len() as u64).await;
buf.clear();
}
}
Some(Err(e)) => {
return Err(format!("读取数据失败: {}", e));
}
None => {
// 流结束,写入剩余数据
if !buf.is_empty() {
file.write_all(&buf)
.await
.map_err(|e| format!("写入文件失败: {}", e))?;
local_completed += buf.len() as u64;
progress.store(local_completed, Ordering::Relaxed);
limiter.consume(buf.len() as u64).await;
buf.clear();
}
break;
}
}
}
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 {
return vec![Segment {
index: 0,
start: 0,
end: 0,
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
}