下载非内核
This commit is contained in:
@@ -0,0 +1,99 @@
|
||||
use std::collections::HashMap;
|
||||
|
||||
use tauri::{AppHandle, State};
|
||||
use tauri_plugin_opener::OpenerExt;
|
||||
|
||||
use super::engine::DownloadEngine;
|
||||
use super::task::{DownloadTask, DownloaderSettings};
|
||||
|
||||
/// 获取所有任务
|
||||
#[tauri::command]
|
||||
pub fn downloader_get_tasks(engine: State<'_, DownloadEngine>) -> Vec<DownloadTask> {
|
||||
engine.get_tasks()
|
||||
}
|
||||
|
||||
/// 添加下载任务
|
||||
#[tauri::command]
|
||||
pub async fn downloader_add_task(
|
||||
engine: State<'_, DownloadEngine>,
|
||||
url: String,
|
||||
filename: Option<String>,
|
||||
dir: Option<String>,
|
||||
headers: Option<HashMap<String, String>>,
|
||||
) -> Result<String, String> {
|
||||
engine.add_task(url, filename, dir, headers.unwrap_or_default()).await
|
||||
}
|
||||
|
||||
/// 暂停任务
|
||||
#[tauri::command]
|
||||
pub fn downloader_pause_task(engine: State<'_, DownloadEngine>, id: String) -> Result<(), String> {
|
||||
engine.pause_task(&id)
|
||||
}
|
||||
|
||||
/// 恢复任务
|
||||
#[tauri::command]
|
||||
pub fn downloader_resume_task(engine: State<'_, DownloadEngine>, id: String) -> Result<(), String> {
|
||||
engine.resume_task(&id)
|
||||
}
|
||||
|
||||
/// 移除任务
|
||||
#[tauri::command]
|
||||
pub fn downloader_remove_task(
|
||||
engine: State<'_, DownloadEngine>,
|
||||
id: String,
|
||||
delete_files: Option<bool>,
|
||||
) -> Result<(), String> {
|
||||
engine.remove_task(&id, delete_files.unwrap_or(false))
|
||||
}
|
||||
|
||||
/// 获取设置
|
||||
#[tauri::command]
|
||||
pub fn downloader_get_settings(engine: State<'_, DownloadEngine>) -> DownloaderSettings {
|
||||
engine.get_settings()
|
||||
}
|
||||
|
||||
/// 保存设置
|
||||
#[tauri::command]
|
||||
pub fn downloader_save_settings(
|
||||
engine: State<'_, DownloadEngine>,
|
||||
settings: DownloaderSettings,
|
||||
) -> Result<(), String> {
|
||||
engine.save_settings(settings);
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// 引擎状态(始终运行中)
|
||||
#[tauri::command]
|
||||
pub fn downloader_status(engine: State<'_, DownloadEngine>) -> serde_json::Value {
|
||||
serde_json::json!({
|
||||
"running": engine.is_started(),
|
||||
})
|
||||
}
|
||||
|
||||
/// 获取扩展服务信息
|
||||
#[tauri::command]
|
||||
pub fn downloader_get_extension_info(engine: State<'_, DownloadEngine>) -> serde_json::Value {
|
||||
let settings = engine.get_settings();
|
||||
serde_json::json!({
|
||||
"url": format!("http://127.0.0.1:{}/", settings.extension_port),
|
||||
"port": settings.extension_port,
|
||||
"secret": settings.extension_secret,
|
||||
"hasSecret": !settings.extension_secret.is_empty(),
|
||||
})
|
||||
}
|
||||
|
||||
/// 用系统资源管理器打开目录
|
||||
#[tauri::command]
|
||||
pub fn downloader_open_dir(app: AppHandle, path: String) -> Result<(), String> {
|
||||
app.opener()
|
||||
.open_path(path, None::<&str>)
|
||||
.map_err(|e| format!("打开目录失败: {}", e))
|
||||
}
|
||||
|
||||
/// 用系统默认浏览器打开 URL
|
||||
#[tauri::command]
|
||||
pub fn downloader_open_url(app: AppHandle, url: String) -> Result<(), String> {
|
||||
app.opener()
|
||||
.open_url(url, None::<&str>)
|
||||
.map_err(|e| format!("打开链接失败: {}", e))
|
||||
}
|
||||
@@ -0,0 +1,623 @@
|
||||
use std::collections::HashMap;
|
||||
use std::path::PathBuf;
|
||||
use std::sync::atomic::{AtomicBool, Ordering};
|
||||
use std::sync::{Arc, Mutex};
|
||||
use std::time::{Duration, Instant};
|
||||
|
||||
use tauri::{AppHandle, Emitter};
|
||||
|
||||
use super::http_dl::{HttpDownloader, split_segments};
|
||||
use super::rate_limit::RateLimiter;
|
||||
use super::storage::{EngineState, Storage};
|
||||
use super::task::{DownloadTask, DownloaderSettings, Segment, TaskStatus};
|
||||
|
||||
/// 进度事件载荷(发给前端 download-progress 事件)
|
||||
#[derive(Debug, Clone, serde::Serialize)]
|
||||
#[serde(rename_all = "camelCase")]
|
||||
pub struct ProgressPayload {
|
||||
pub id: String,
|
||||
pub completed_size: u64,
|
||||
pub total_size: u64,
|
||||
pub speed: u64,
|
||||
pub status: TaskStatus,
|
||||
}
|
||||
|
||||
/// 完成事件载荷(发给前端 download-complete 事件)
|
||||
#[derive(Debug, Clone, serde::Serialize)]
|
||||
#[serde(rename_all = "camelCase")]
|
||||
pub struct CompletePayload {
|
||||
pub id: String,
|
||||
pub filename: String,
|
||||
pub status: TaskStatus,
|
||||
pub error: Option<String>,
|
||||
}
|
||||
|
||||
/// 活跃下载句柄
|
||||
struct TaskHandle {
|
||||
cancel: Arc<AtomicBool>,
|
||||
/// 每个分段的已下载字节(与 segments 一一对应)
|
||||
progress: Vec<Arc<std::sync::atomic::AtomicU64>>,
|
||||
/// 下载任务的 JoinHandle(None = 已完成/已取消)
|
||||
join: Mutex<Option<tauri::async_runtime::JoinHandle<()>>>,
|
||||
}
|
||||
|
||||
/// 下载引擎(进程内运行,通过 Arc 共享状态)
|
||||
#[derive(Clone)]
|
||||
pub struct DownloadEngine {
|
||||
inner: Arc<EngineInner>,
|
||||
}
|
||||
|
||||
struct EngineInner {
|
||||
/// 所有任务状态
|
||||
tasks: Mutex<HashMap<String, DownloadTask>>,
|
||||
/// 活跃下载句柄
|
||||
handles: Mutex<HashMap<String, TaskHandle>>,
|
||||
/// 持久化存储
|
||||
storage: Storage,
|
||||
/// 下载设置
|
||||
settings: Mutex<DownloaderSettings>,
|
||||
/// 全局限速器(bytes/s,0=不限)
|
||||
global_limiter: Arc<RateLimiter>,
|
||||
/// HTTP 下载器
|
||||
http: HttpDownloader,
|
||||
/// Tauri 应用句柄(用于发事件)
|
||||
app_handle: AppHandle,
|
||||
/// 引擎是否已启动
|
||||
started: AtomicBool,
|
||||
/// 持久化节流:上次保存时间
|
||||
last_save: Mutex<Instant>,
|
||||
}
|
||||
|
||||
impl DownloadEngine {
|
||||
/// 创建引擎并从磁盘恢复状态
|
||||
pub fn new(data_dir: PathBuf, app_handle: AppHandle) -> Self {
|
||||
let storage = Storage::new(data_dir);
|
||||
let state = storage.load();
|
||||
|
||||
// 恢复时,将 Active 任务标记为 Paused(避免自动恢复下载)
|
||||
let mut tasks: HashMap<String, DownloadTask> = state
|
||||
.tasks
|
||||
.into_iter()
|
||||
.map(|t| (t.id.clone(), t))
|
||||
.collect();
|
||||
for task in tasks.values_mut() {
|
||||
if task.status == TaskStatus::Active {
|
||||
task.status = TaskStatus::Paused;
|
||||
task.speed = 0;
|
||||
}
|
||||
}
|
||||
|
||||
let settings = state.settings;
|
||||
// 创建全局限速器(KB/s → bytes/s)
|
||||
let global_limit = if settings.global_speed_limit > 0 {
|
||||
settings.global_speed_limit * 1024
|
||||
} else {
|
||||
0
|
||||
};
|
||||
let global_limiter = Arc::new(RateLimiter::new(global_limit));
|
||||
|
||||
let engine = Self {
|
||||
inner: Arc::new(EngineInner {
|
||||
tasks: Mutex::new(tasks),
|
||||
handles: Mutex::new(HashMap::new()),
|
||||
storage,
|
||||
settings: Mutex::new(settings),
|
||||
global_limiter,
|
||||
http: HttpDownloader::new(),
|
||||
app_handle,
|
||||
started: AtomicBool::new(false),
|
||||
last_save: Mutex::new(Instant::now()),
|
||||
}),
|
||||
};
|
||||
|
||||
// 启动后台持久化定时器
|
||||
engine.start_save_timer();
|
||||
engine.inner.started.store(true, Ordering::SeqCst);
|
||||
engine
|
||||
}
|
||||
|
||||
/// 启动后台定时持久化(每 5 秒检查一次)
|
||||
fn start_save_timer(&self) {
|
||||
let engine = self.clone();
|
||||
tauri::async_runtime::spawn(async move {
|
||||
let mut interval = tokio::time::interval(Duration::from_secs(5));
|
||||
loop {
|
||||
interval.tick().await;
|
||||
engine.persist_throttled();
|
||||
}
|
||||
});
|
||||
}
|
||||
|
||||
// ===================== 公开 API =====================
|
||||
|
||||
/// 添加下载任务
|
||||
pub async fn add_task(
|
||||
&self,
|
||||
url: String,
|
||||
filename: Option<String>,
|
||||
dir: Option<String>,
|
||||
headers: HashMap<String, String>,
|
||||
) -> Result<String, String> {
|
||||
// 探测资源信息
|
||||
let probe = self.inner.http.probe(&url, &headers).await;
|
||||
|
||||
let settings = self.inner.settings.lock().unwrap().clone();
|
||||
let task_dir = dir.unwrap_or_else(|| settings.download_dir.clone());
|
||||
|
||||
// 确定文件名
|
||||
let task_filename = filename
|
||||
.or_else(|| probe.as_ref().ok().and_then(|p| p.filename.clone()))
|
||||
.unwrap_or_else(|| {
|
||||
url.split('?')
|
||||
.next()
|
||||
.and_then(|u| u.rsplit('/').next())
|
||||
.filter(|n| !n.is_empty())
|
||||
.map(|n| n.to_string())
|
||||
.unwrap_or_else(|| format!("download_{}", chrono::Utc::now().timestamp()))
|
||||
});
|
||||
|
||||
let id = self.inner.storage.next_task_id();
|
||||
|
||||
// 创建分段
|
||||
let segments = match &probe {
|
||||
Ok(p) if p.supports_resume && p.total_size.map(|s| s > 0).unwrap_or(false) => {
|
||||
split_segments(p.total_size.unwrap(), settings.max_connections)
|
||||
}
|
||||
_ => vec![Segment {
|
||||
index: 0,
|
||||
start: 0,
|
||||
end: 0,
|
||||
completed: 0,
|
||||
}],
|
||||
};
|
||||
|
||||
let total_size = probe.as_ref().ok().and_then(|p| p.total_size).unwrap_or(0);
|
||||
let supports_resume = probe.as_ref().ok().map(|p| p.supports_resume).unwrap_or(false);
|
||||
|
||||
let task = DownloadTask {
|
||||
id: id.clone(),
|
||||
url: url.clone(),
|
||||
filename: task_filename,
|
||||
dir: task_dir,
|
||||
status: TaskStatus::Queued,
|
||||
total_size,
|
||||
completed_size: 0,
|
||||
speed: 0,
|
||||
supports_resume,
|
||||
segments: segments.clone(),
|
||||
error: None,
|
||||
created_at: chrono::Utc::now().timestamp_millis(),
|
||||
headers,
|
||||
};
|
||||
|
||||
// 如果探测失败,标记为 Error
|
||||
let task = match probe {
|
||||
Ok(_) => task,
|
||||
Err(e) => DownloadTask {
|
||||
status: TaskStatus::Error,
|
||||
error: Some(e),
|
||||
..task
|
||||
},
|
||||
};
|
||||
|
||||
{
|
||||
let mut tasks = self.inner.tasks.lock().unwrap();
|
||||
tasks.insert(id.clone(), task);
|
||||
}
|
||||
self.persist_now();
|
||||
|
||||
// 如果探测成功,尝试调度
|
||||
let should_schedule = {
|
||||
let tasks = self.inner.tasks.lock().unwrap();
|
||||
tasks.get(&id).map(|t| t.status == TaskStatus::Queued).unwrap_or(false)
|
||||
};
|
||||
if should_schedule {
|
||||
self.schedule();
|
||||
}
|
||||
|
||||
// 通知前端有新任务加入(扩展通过 HTTP API 添加时,前端需要刷新)
|
||||
let _ = self.inner.app_handle.emit("download-added", serde_json::json!({ "id": id }));
|
||||
|
||||
Ok(id)
|
||||
}
|
||||
|
||||
/// 暂停任务
|
||||
pub fn pause_task(&self, id: &str) -> Result<(), String> {
|
||||
// 1. 设置取消标志 + 读取进度(锁 handles)
|
||||
let progress_values: Vec<u64> = {
|
||||
let handles = self.inner.handles.lock().unwrap();
|
||||
if let Some(handle) = handles.get(id) {
|
||||
handle.cancel.store(true, Ordering::SeqCst);
|
||||
handle.progress.iter().map(|p| p.load(Ordering::Relaxed)).collect()
|
||||
} else {
|
||||
Vec::new()
|
||||
}
|
||||
};
|
||||
|
||||
// 2. 更新任务状态 + 同步进度(锁 tasks,不嵌套锁 handles)
|
||||
{
|
||||
let mut tasks = self.inner.tasks.lock().unwrap();
|
||||
if let Some(task) = tasks.get_mut(id) {
|
||||
if task.status == TaskStatus::Active || task.status == TaskStatus::Queued {
|
||||
task.status = TaskStatus::Paused;
|
||||
task.speed = 0;
|
||||
// 同步分段进度
|
||||
for (i, val) in progress_values.iter().enumerate() {
|
||||
if let Some(seg) = task.segments.get_mut(i) {
|
||||
seg.completed = *val;
|
||||
}
|
||||
}
|
||||
task.recalc_completed();
|
||||
}
|
||||
} else {
|
||||
return Err("任务不存在".to_string());
|
||||
}
|
||||
}
|
||||
self.persist_now();
|
||||
self.schedule();
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// 恢复任务
|
||||
pub fn resume_task(&self, id: &str) -> Result<(), String> {
|
||||
{
|
||||
let mut tasks = self.inner.tasks.lock().unwrap();
|
||||
if let Some(task) = tasks.get_mut(id) {
|
||||
if task.status != TaskStatus::Paused && task.status != TaskStatus::Error {
|
||||
return Err("任务不在可恢复状态".to_string());
|
||||
}
|
||||
// 如果不支持断点续传,从头开始
|
||||
if !task.supports_resume {
|
||||
task.completed_size = 0;
|
||||
for seg in &mut task.segments {
|
||||
seg.completed = 0;
|
||||
}
|
||||
}
|
||||
task.status = TaskStatus::Queued;
|
||||
task.error = None;
|
||||
} else {
|
||||
return Err("任务不存在".to_string());
|
||||
}
|
||||
}
|
||||
self.persist_now();
|
||||
self.schedule();
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// 移除任务
|
||||
pub fn remove_task(&self, id: &str, delete_files: bool) -> Result<(), String> {
|
||||
// 1. 设置取消标志 + abort join handle(锁 handles)
|
||||
{
|
||||
let mut handles = self.inner.handles.lock().unwrap();
|
||||
if let Some(handle) = handles.remove(id) {
|
||||
handle.cancel.store(true, Ordering::SeqCst);
|
||||
if let Ok(mut join) = handle.join.lock() {
|
||||
if let Some(j) = join.take() {
|
||||
j.abort();
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// 2. 从任务列表移除(锁 tasks)
|
||||
let task = {
|
||||
let mut tasks = self.inner.tasks.lock().unwrap();
|
||||
tasks.remove(id)
|
||||
};
|
||||
|
||||
// 3. 可选删除文件(临时文件 + 最终文件)
|
||||
if delete_files {
|
||||
if let Some(task) = &task {
|
||||
let _ = std::fs::remove_file(task.file_path());
|
||||
let _ = std::fs::remove_file(task.temp_file_path());
|
||||
}
|
||||
}
|
||||
|
||||
self.persist_now();
|
||||
self.schedule();
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// 获取所有任务
|
||||
pub fn get_tasks(&self) -> Vec<DownloadTask> {
|
||||
self.inner.tasks.lock().unwrap().values().cloned().collect()
|
||||
}
|
||||
|
||||
/// 获取设置
|
||||
pub fn get_settings(&self) -> DownloaderSettings {
|
||||
self.inner.settings.lock().unwrap().clone()
|
||||
}
|
||||
|
||||
/// 保存设置
|
||||
pub fn save_settings(&self, settings: DownloaderSettings) {
|
||||
// 更新全局限速器
|
||||
let new_limit = if settings.global_speed_limit > 0 {
|
||||
settings.global_speed_limit * 1024
|
||||
} else {
|
||||
0
|
||||
};
|
||||
self.inner.global_limiter.set_limit(new_limit);
|
||||
|
||||
{
|
||||
let mut s = self.inner.settings.lock().unwrap();
|
||||
*s = settings;
|
||||
}
|
||||
self.persist_now();
|
||||
}
|
||||
|
||||
/// 引擎是否已启动
|
||||
pub fn is_started(&self) -> bool {
|
||||
self.inner.started.load(Ordering::SeqCst)
|
||||
}
|
||||
|
||||
/// 退出时清理:停止所有下载、保存状态
|
||||
pub fn cleanup_on_exit(&self) {
|
||||
// 取消所有活跃下载
|
||||
{
|
||||
let handles = self.inner.handles.lock().unwrap();
|
||||
for handle in handles.values() {
|
||||
handle.cancel.store(true, Ordering::SeqCst);
|
||||
}
|
||||
}
|
||||
|
||||
// 将 Active 任务标记为 Paused
|
||||
{
|
||||
let mut tasks = self.inner.tasks.lock().unwrap();
|
||||
for task in tasks.values_mut() {
|
||||
if task.status == TaskStatus::Active {
|
||||
task.status = TaskStatus::Paused;
|
||||
task.speed = 0;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// 等待短暂时间让下载任务退出
|
||||
std::thread::sleep(Duration::from_millis(200));
|
||||
|
||||
// 最终保存
|
||||
self.persist_now();
|
||||
}
|
||||
|
||||
// ===================== 内部调度逻辑 =====================
|
||||
|
||||
/// 调度:如果活跃任务数 < max_concurrent,启动排队任务
|
||||
fn schedule(&self) {
|
||||
let max_concurrent = self.inner.settings.lock().unwrap().max_concurrent as usize;
|
||||
|
||||
let (active_count, queued_ids) = {
|
||||
let tasks = self.inner.tasks.lock().unwrap();
|
||||
let active = tasks.values().filter(|t| t.status == TaskStatus::Active).count();
|
||||
let mut queued: Vec<_> = tasks
|
||||
.values()
|
||||
.filter(|t| t.status == TaskStatus::Queued)
|
||||
.collect();
|
||||
queued.sort_by_key(|t| t.created_at);
|
||||
(active, queued.into_iter().map(|t| t.id.clone()).collect::<Vec<_>>())
|
||||
};
|
||||
|
||||
if active_count >= max_concurrent {
|
||||
return;
|
||||
}
|
||||
|
||||
let slots = max_concurrent.saturating_sub(active_count);
|
||||
for id in queued_ids.into_iter().take(slots) {
|
||||
self.start_download(id);
|
||||
}
|
||||
}
|
||||
|
||||
/// 启动单个下载任务
|
||||
fn start_download(&self, id: String) {
|
||||
let task = {
|
||||
let mut tasks = self.inner.tasks.lock().unwrap();
|
||||
match tasks.get_mut(&id) {
|
||||
Some(task) if task.status == TaskStatus::Queued => {
|
||||
task.status = TaskStatus::Active;
|
||||
task.speed = 0;
|
||||
task.clone()
|
||||
}
|
||||
_ => return,
|
||||
}
|
||||
};
|
||||
|
||||
// 创建取消标志和进度计数器
|
||||
let cancel = Arc::new(AtomicBool::new(false));
|
||||
let progress: Vec<Arc<std::sync::atomic::AtomicU64>> = task
|
||||
.segments
|
||||
.iter()
|
||||
.map(|s| Arc::new(std::sync::atomic::AtomicU64::new(s.completed)))
|
||||
.collect();
|
||||
|
||||
// 存储句柄
|
||||
let handle = TaskHandle {
|
||||
cancel: cancel.clone(),
|
||||
progress: progress.iter().map(|p| p.clone()).collect(),
|
||||
join: Mutex::new(None),
|
||||
};
|
||||
self.inner.handles.lock().unwrap().insert(id.clone(), handle);
|
||||
|
||||
// 生成下载 future
|
||||
let engine = self.clone();
|
||||
let id_clone = id.clone();
|
||||
let cancel_clone = cancel.clone();
|
||||
let progress_clone: Vec<Arc<std::sync::atomic::AtomicU64>> =
|
||||
progress.iter().map(|p| p.clone()).collect();
|
||||
let limiter = self.inner.global_limiter.clone();
|
||||
let url = task.url.clone();
|
||||
let headers = task.headers.clone();
|
||||
let segments = task.segments.clone();
|
||||
let temp_file_path = task.temp_file_path();
|
||||
let final_file_path = task.file_path();
|
||||
let http = self.inner.http.clone();
|
||||
|
||||
let join = tauri::async_runtime::spawn(async move {
|
||||
let result = http
|
||||
.download(&url, &headers, &segments, &temp_file_path, cancel_clone, &progress_clone, limiter)
|
||||
.await;
|
||||
|
||||
// 下载结束,更新任务状态
|
||||
let final_status = match &result {
|
||||
Ok(()) => TaskStatus::Complete,
|
||||
Err(e) if e == "已取消" => TaskStatus::Paused,
|
||||
Err(_) => TaskStatus::Error,
|
||||
};
|
||||
|
||||
// 下载成功后,将临时文件重命名为最终文件名
|
||||
if final_status == TaskStatus::Complete {
|
||||
let _ = tokio::fs::rename(&temp_file_path, &final_file_path).await;
|
||||
}
|
||||
|
||||
// 同步最终进度到任务
|
||||
{
|
||||
let mut tasks = engine.inner.tasks.lock().unwrap();
|
||||
if let Some(task) = tasks.get_mut(&id_clone) {
|
||||
for (i, prog) in progress_clone.iter().enumerate() {
|
||||
if let Some(seg) = task.segments.get_mut(i) {
|
||||
seg.completed = prog.load(Ordering::Relaxed);
|
||||
}
|
||||
}
|
||||
task.recalc_completed();
|
||||
task.speed = 0;
|
||||
task.status = final_status.clone();
|
||||
if let Err(e) = &result {
|
||||
if e != "已取消" {
|
||||
task.error = Some(e.clone());
|
||||
}
|
||||
}
|
||||
if final_status == TaskStatus::Complete {
|
||||
task.completed_size = task.total_size.max(task.completed_size);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// 从活跃句柄中移除
|
||||
engine.inner.handles.lock().unwrap().remove(&id_clone);
|
||||
|
||||
// 发送完成事件
|
||||
let task = engine.inner.tasks.lock().unwrap().get(&id_clone).cloned();
|
||||
if let Some(task) = task {
|
||||
let _ = engine.inner.app_handle.emit(
|
||||
"download-complete",
|
||||
CompletePayload {
|
||||
id: id_clone.clone(),
|
||||
filename: task.filename.clone(),
|
||||
status: task.status.clone(),
|
||||
error: task.error.clone(),
|
||||
},
|
||||
);
|
||||
}
|
||||
|
||||
// 持久化 + 调度下一个
|
||||
engine.persist_now();
|
||||
engine.schedule();
|
||||
});
|
||||
|
||||
// 存储 JoinHandle
|
||||
if let Some(h) = self.inner.handles.lock().unwrap().get_mut(&id) {
|
||||
if let Ok(mut join_guard) = h.join.lock() {
|
||||
*join_guard = Some(join);
|
||||
}
|
||||
}
|
||||
|
||||
// 启动进度监控(每 500ms 更新一次)
|
||||
let engine = self.clone();
|
||||
let id_monitor = id.clone();
|
||||
let progress_monitor = progress;
|
||||
let cancel_monitor = cancel;
|
||||
tauri::async_runtime::spawn(async move {
|
||||
// 初始化为当前已下载量,避免恢复下载时首次计算速度异常
|
||||
let mut last_completed: u64 = progress_monitor
|
||||
.iter()
|
||||
.map(|p| p.load(Ordering::Relaxed))
|
||||
.sum();
|
||||
let mut last_time = Instant::now();
|
||||
let mut interval = tokio::time::interval(Duration::from_millis(500));
|
||||
interval.tick().await; // 跳过第一次立即触发
|
||||
|
||||
loop {
|
||||
interval.tick().await;
|
||||
// 如果任务已不在活跃状态,停止监控
|
||||
let is_active = {
|
||||
let tasks = engine.inner.tasks.lock().unwrap();
|
||||
tasks.get(&id_monitor).map(|t| t.status == TaskStatus::Active).unwrap_or(false)
|
||||
};
|
||||
if !is_active || cancel_monitor.load(Ordering::SeqCst) {
|
||||
break;
|
||||
}
|
||||
|
||||
// 读取进度
|
||||
let completed: u64 = progress_monitor
|
||||
.iter()
|
||||
.map(|p| p.load(Ordering::Relaxed))
|
||||
.sum();
|
||||
|
||||
// 计算速度
|
||||
let now = Instant::now();
|
||||
let elapsed = now.duration_since(last_time).as_secs_f64();
|
||||
let speed = if elapsed > 0.0 && completed >= last_completed {
|
||||
((completed - last_completed) as f64 / elapsed) as u64
|
||||
} else {
|
||||
0
|
||||
};
|
||||
last_completed = completed;
|
||||
last_time = now;
|
||||
|
||||
// 更新任务状态 + 发送进度事件
|
||||
let total_size = {
|
||||
let mut tasks = engine.inner.tasks.lock().unwrap();
|
||||
if let Some(task) = tasks.get_mut(&id_monitor) {
|
||||
task.completed_size = completed;
|
||||
task.speed = speed;
|
||||
for (i, prog) in progress_monitor.iter().enumerate() {
|
||||
if let Some(seg) = task.segments.get_mut(i) {
|
||||
seg.completed = prog.load(Ordering::Relaxed);
|
||||
}
|
||||
}
|
||||
task.total_size
|
||||
} else {
|
||||
break;
|
||||
}
|
||||
};
|
||||
|
||||
let _ = engine.inner.app_handle.emit(
|
||||
"download-progress",
|
||||
ProgressPayload {
|
||||
id: id_monitor.clone(),
|
||||
completed_size: completed,
|
||||
total_size,
|
||||
speed,
|
||||
status: TaskStatus::Active,
|
||||
},
|
||||
);
|
||||
}
|
||||
});
|
||||
}
|
||||
|
||||
// ===================== 持久化 =====================
|
||||
|
||||
/// 节流持久化(至少间隔 3 秒)
|
||||
fn persist_throttled(&self) {
|
||||
let should_save = {
|
||||
let last = self.inner.last_save.lock().unwrap();
|
||||
last.elapsed() >= Duration::from_secs(3)
|
||||
};
|
||||
if should_save {
|
||||
self.persist_now();
|
||||
}
|
||||
}
|
||||
|
||||
/// 立即持久化
|
||||
fn persist_now(&self) {
|
||||
let tasks: Vec<DownloadTask> = {
|
||||
let tasks = self.inner.tasks.lock().unwrap();
|
||||
tasks.values().cloned().collect()
|
||||
};
|
||||
let settings = self.inner.settings.lock().unwrap().clone();
|
||||
let state = EngineState {
|
||||
tasks,
|
||||
settings,
|
||||
next_id: 0, // storage.save 会从 id_counter 读取
|
||||
};
|
||||
self.inner.storage.save(state);
|
||||
*self.inner.last_save.lock().unwrap() = Instant::now();
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -0,0 +1,15 @@
|
||||
pub mod commands;
|
||||
pub mod engine;
|
||||
pub mod http_dl;
|
||||
pub mod rate_limit;
|
||||
pub mod server;
|
||||
pub mod storage;
|
||||
pub mod task;
|
||||
|
||||
pub use commands::{
|
||||
downloader_add_task, downloader_get_extension_info, downloader_get_settings, downloader_get_tasks,
|
||||
downloader_open_dir, downloader_open_url, downloader_pause_task, downloader_remove_task,
|
||||
downloader_resume_task, downloader_save_settings, downloader_status,
|
||||
};
|
||||
pub use engine::DownloadEngine;
|
||||
pub use server::ExtensionServer;
|
||||
@@ -0,0 +1,74 @@
|
||||
use std::sync::atomic::{AtomicU64, Ordering};
|
||||
use std::sync::{Arc, Mutex};
|
||||
use std::time::Duration;
|
||||
|
||||
/// 令牌桶限速器(字节/秒,0=不限)
|
||||
///
|
||||
/// 滑动窗口算法:1 秒窗口内累计消耗字节数,超出限制时 sleep 到下一窗口。
|
||||
pub struct RateLimiter {
|
||||
/// 限速 bytes/s(0=不限),用 Arc<AtomicU64> 以便运行时动态调整
|
||||
limit: Arc<AtomicU64>,
|
||||
/// 当前窗口已消耗字节
|
||||
consumed: AtomicU64,
|
||||
/// 窗口起始时间
|
||||
window_start: Mutex<std::time::Instant>,
|
||||
}
|
||||
|
||||
impl RateLimiter {
|
||||
pub fn new(limit: u64) -> Self {
|
||||
Self {
|
||||
limit: Arc::new(AtomicU64::new(limit)),
|
||||
consumed: AtomicU64::new(0),
|
||||
window_start: Mutex::new(std::time::Instant::now()),
|
||||
}
|
||||
}
|
||||
|
||||
/// 动态设置限速(bytes/s,0=不限)
|
||||
pub fn set_limit(&self, limit: u64) {
|
||||
self.limit.store(limit, Ordering::Relaxed);
|
||||
}
|
||||
|
||||
/// 消费指定字节数,若超出限速则 sleep 等待
|
||||
pub async fn consume(&self, bytes: u64) {
|
||||
let limit = self.limit.load(Ordering::Relaxed);
|
||||
if limit == 0 || bytes == 0 {
|
||||
return;
|
||||
}
|
||||
|
||||
// 尝试在当前窗口消费(用作用域确保 MutexGuard 在 await 前释放)
|
||||
let over_limit = {
|
||||
let now = std::time::Instant::now();
|
||||
let mut start = self.window_start.lock().unwrap();
|
||||
let elapsed = now.duration_since(*start);
|
||||
|
||||
// 窗口过期,重置
|
||||
if elapsed >= Duration::from_secs(1) {
|
||||
self.consumed.store(0, Ordering::Relaxed);
|
||||
*start = now;
|
||||
}
|
||||
|
||||
let current = self.consumed.fetch_add(bytes, Ordering::Relaxed) + bytes;
|
||||
if current > limit {
|
||||
Some(current - limit)
|
||||
} else {
|
||||
None
|
||||
}
|
||||
}; // MutexGuard 在此释放
|
||||
|
||||
// 超出限速,计算需要等待的时间
|
||||
if let Some(over) = over_limit {
|
||||
let wait_ms = (over as f64 / limit as f64 * 1000.0).ceil() as u64;
|
||||
let wait = Duration::from_millis(wait_ms.min(1000));
|
||||
tokio::time::sleep(wait).await;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// 克隆限速器的共享句柄(共享同一限速状态)
|
||||
impl Clone for RateLimiter {
|
||||
fn clone(&self) -> Self {
|
||||
// 注意:clone 的是限速配置,不是状态。新实例独立计数。
|
||||
// 全局限速器应通过引用共享,而非 clone。
|
||||
Self::new(self.limit.load(Ordering::Relaxed))
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,139 @@
|
||||
use std::collections::HashMap;
|
||||
use std::net::SocketAddr;
|
||||
|
||||
use axum::{
|
||||
extract::{Path, State},
|
||||
http::{HeaderMap, StatusCode},
|
||||
response::Json,
|
||||
routing::{get, post},
|
||||
Router,
|
||||
};
|
||||
use serde::{Deserialize, Serialize};
|
||||
|
||||
use super::engine::DownloadEngine;
|
||||
use super::task::DownloadTask;
|
||||
|
||||
/// 扩展 HTTP API 服务器
|
||||
pub struct ExtensionServer;
|
||||
|
||||
#[derive(Serialize)]
|
||||
struct HealthResponse {
|
||||
ok: bool,
|
||||
version: &'static str,
|
||||
}
|
||||
|
||||
#[derive(Deserialize)]
|
||||
struct CreateDownloadRequest {
|
||||
url: String,
|
||||
#[serde(default)]
|
||||
filename: Option<String>,
|
||||
#[serde(default)]
|
||||
dir: Option<String>,
|
||||
#[serde(default)]
|
||||
headers: HashMap<String, String>,
|
||||
}
|
||||
|
||||
#[derive(Serialize)]
|
||||
struct CreateDownloadResponse {
|
||||
id: String,
|
||||
}
|
||||
|
||||
#[derive(Serialize)]
|
||||
struct ErrorResponse {
|
||||
error: String,
|
||||
}
|
||||
|
||||
impl ExtensionServer {
|
||||
/// 启动 HTTP API 服务器(绑定到 127.0.0.1:port)
|
||||
pub async fn start(engine: DownloadEngine, port: u16, secret: String) {
|
||||
let addr: SocketAddr = format!("127.0.0.1:{}", port).parse().expect("无效端口");
|
||||
|
||||
let app = Router::new()
|
||||
.route("/health", get(health))
|
||||
.route("/api/downloads", post(create_download).get(list_downloads))
|
||||
.route("/api/downloads/:id", axum::routing::delete(remove_download))
|
||||
.with_state(AppState { engine, secret });
|
||||
|
||||
let listener = match tokio::net::TcpListener::bind(&addr).await {
|
||||
Ok(l) => l,
|
||||
Err(e) => {
|
||||
eprintln!("[download_engine] 扩展 HTTP 服务启动失败 ({}): {}", addr, e);
|
||||
return;
|
||||
}
|
||||
};
|
||||
|
||||
eprintln!("[download_engine] 扩展 HTTP 服务已启动: http://{}", addr);
|
||||
|
||||
if let Err(e) = axum::serve(listener, app).await {
|
||||
eprintln!("[download_engine] 扩展 HTTP 服务异常: {}", e);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone)]
|
||||
struct AppState {
|
||||
engine: DownloadEngine,
|
||||
secret: String,
|
||||
}
|
||||
|
||||
/// 鉴权检查:如果配置了 secret,校验 Bearer token
|
||||
fn check_auth(headers: &HeaderMap, secret: &str) -> bool {
|
||||
if secret.is_empty() {
|
||||
return true;
|
||||
}
|
||||
if let Some(auth) = headers.get("authorization") {
|
||||
if let Ok(s) = auth.to_str() {
|
||||
if let Some(token) = s.strip_prefix("Bearer ") {
|
||||
return token == secret;
|
||||
}
|
||||
}
|
||||
}
|
||||
false
|
||||
}
|
||||
|
||||
async fn health(State(state): State<AppState>) -> Json<HealthResponse> {
|
||||
let _ = state; // 不需要鉴权
|
||||
Json(HealthResponse {
|
||||
ok: true,
|
||||
version: "thing-download-engine/1.0",
|
||||
})
|
||||
}
|
||||
|
||||
async fn create_download(
|
||||
State(state): State<AppState>,
|
||||
headers: HeaderMap,
|
||||
Json(req): Json<CreateDownloadRequest>,
|
||||
) -> Result<Json<CreateDownloadResponse>, (StatusCode, Json<ErrorResponse>)> {
|
||||
if !check_auth(&headers, &state.secret) {
|
||||
return Err((StatusCode::UNAUTHORIZED, Json(ErrorResponse { error: "未授权".into() })));
|
||||
}
|
||||
|
||||
match state.engine.add_task(req.url, req.filename, req.dir, req.headers).await {
|
||||
Ok(id) => Ok(Json(CreateDownloadResponse { id })),
|
||||
Err(e) => Err((StatusCode::BAD_REQUEST, Json(ErrorResponse { error: e }))),
|
||||
}
|
||||
}
|
||||
|
||||
async fn list_downloads(
|
||||
State(state): State<AppState>,
|
||||
headers: HeaderMap,
|
||||
) -> Result<Json<Vec<DownloadTask>>, (StatusCode, Json<ErrorResponse>)> {
|
||||
if !check_auth(&headers, &state.secret) {
|
||||
return Err((StatusCode::UNAUTHORIZED, Json(ErrorResponse { error: "未授权".into() })));
|
||||
}
|
||||
Ok(Json(state.engine.get_tasks()))
|
||||
}
|
||||
|
||||
async fn remove_download(
|
||||
State(state): State<AppState>,
|
||||
headers: HeaderMap,
|
||||
Path(id): Path<String>,
|
||||
) -> Result<StatusCode, (StatusCode, Json<ErrorResponse>)> {
|
||||
if !check_auth(&headers, &state.secret) {
|
||||
return Err((StatusCode::UNAUTHORIZED, Json(ErrorResponse { error: "未授权".into() })));
|
||||
}
|
||||
match state.engine.remove_task(&id, false) {
|
||||
Ok(()) => Ok(StatusCode::NO_CONTENT),
|
||||
Err(e) => Err((StatusCode::NOT_FOUND, Json(ErrorResponse { error: e }))),
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,82 @@
|
||||
use serde::{Deserialize, Serialize};
|
||||
use std::fs;
|
||||
use std::path::PathBuf;
|
||||
use std::sync::atomic::{AtomicU64, Ordering};
|
||||
|
||||
use super::task::{DownloadTask, DownloaderSettings};
|
||||
|
||||
/// 引擎持久化状态(序列化到 engine_state.json)
|
||||
#[derive(Serialize, Deserialize)]
|
||||
pub struct EngineState {
|
||||
/// 所有任务(active + queued + paused + complete + error)
|
||||
#[serde(default)]
|
||||
pub tasks: Vec<DownloadTask>,
|
||||
/// 下载设置
|
||||
#[serde(default)]
|
||||
pub settings: DownloaderSettings,
|
||||
/// 自增 ID 计数器(不序列化为 hex,存原始数字)
|
||||
#[serde(default)]
|
||||
pub next_id: u64,
|
||||
}
|
||||
|
||||
impl Default for EngineState {
|
||||
fn default() -> Self {
|
||||
Self {
|
||||
tasks: Vec::new(),
|
||||
settings: DownloaderSettings::default(),
|
||||
next_id: 1,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// 存储管理器:负责加载/保存引擎状态到 JSON 文件
|
||||
pub struct Storage {
|
||||
state_path: PathBuf,
|
||||
/// ID 计数器(内存中维护,与 EngineState.next_id 同步)
|
||||
pub id_counter: AtomicU64,
|
||||
}
|
||||
|
||||
impl Storage {
|
||||
pub fn new(data_dir: PathBuf) -> Self {
|
||||
let state_path = data_dir.join("engine_state.json");
|
||||
let existing = Self::load_raw(&state_path);
|
||||
let next_id = existing.as_ref().map(|s| s.next_id).unwrap_or(1);
|
||||
Self {
|
||||
state_path,
|
||||
id_counter: AtomicU64::new(next_id),
|
||||
}
|
||||
}
|
||||
|
||||
/// 生成下一个任务 ID(16 位 hex 字符串)
|
||||
pub fn next_task_id(&self) -> String {
|
||||
let id = self.id_counter.fetch_add(1, Ordering::SeqCst);
|
||||
format!("{:016x}", id)
|
||||
}
|
||||
|
||||
/// 从磁盘加载状态
|
||||
fn load_raw(path: &PathBuf) -> Option<EngineState> {
|
||||
let raw = fs::read_to_string(path).ok()?;
|
||||
serde_json::from_str::<EngineState>(&raw).ok()
|
||||
}
|
||||
|
||||
/// 加载完整状态(若文件不存在返回默认值)
|
||||
pub fn load(&self) -> EngineState {
|
||||
Self::load_raw(&self.state_path).unwrap_or_default()
|
||||
}
|
||||
|
||||
/// 保存状态到磁盘
|
||||
pub fn save(&self, mut state: EngineState) {
|
||||
// 同步 ID 计数器
|
||||
state.next_id = self.id_counter.load(Ordering::SeqCst);
|
||||
match serde_json::to_string_pretty(&state) {
|
||||
Ok(json) => {
|
||||
if let Err(e) = fs::write(&self.state_path, json) {
|
||||
eprintln!("[download_engine] 保存状态失败: {}", e);
|
||||
}
|
||||
}
|
||||
Err(e) => {
|
||||
eprintln!("[download_engine] 序列化状态失败: {}", e);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,177 @@
|
||||
use serde::{Deserialize, Serialize};
|
||||
use std::collections::HashMap;
|
||||
|
||||
/// 任务状态
|
||||
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
|
||||
#[serde(rename_all = "lowercase")]
|
||||
pub enum TaskStatus {
|
||||
/// 排队等待(并发数已满)
|
||||
Queued,
|
||||
/// 下载中
|
||||
Active,
|
||||
/// 已暂停
|
||||
Paused,
|
||||
/// 已完成
|
||||
Complete,
|
||||
/// 错误
|
||||
Error,
|
||||
}
|
||||
|
||||
/// 下载分段(多线程 Range 下载 / 断点续传用)
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
#[serde(rename_all = "camelCase")]
|
||||
pub struct Segment {
|
||||
/// 分段索引
|
||||
pub index: u32,
|
||||
/// 起始字节(含)
|
||||
pub start: u64,
|
||||
/// 结束字节(含)
|
||||
pub end: u64,
|
||||
/// 已下载字节
|
||||
pub completed: u64,
|
||||
}
|
||||
|
||||
impl Segment {
|
||||
/// 该分段总长度
|
||||
pub fn len(&self) -> u64 {
|
||||
self.end.saturating_sub(self.start) + 1
|
||||
}
|
||||
/// 是否已完成
|
||||
pub fn is_done(&self) -> bool {
|
||||
self.completed >= self.len()
|
||||
}
|
||||
}
|
||||
|
||||
/// 下载任务
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
#[serde(rename_all = "camelCase")]
|
||||
pub struct DownloadTask {
|
||||
/// 任务 ID(自增 hex 字符串)
|
||||
pub id: String,
|
||||
/// 下载地址
|
||||
pub url: String,
|
||||
/// 文件名
|
||||
pub filename: String,
|
||||
/// 保存目录(绝对路径)
|
||||
pub dir: String,
|
||||
/// 状态
|
||||
pub status: TaskStatus,
|
||||
/// 文件总大小(字节),0=未知
|
||||
pub total_size: u64,
|
||||
/// 已下载字节
|
||||
pub completed_size: u64,
|
||||
/// 当前下载速度 bytes/s
|
||||
pub speed: u64,
|
||||
/// 服务器是否支持断点续传
|
||||
pub supports_resume: bool,
|
||||
/// 分段信息
|
||||
#[serde(default)]
|
||||
pub segments: Vec<Segment>,
|
||||
/// 错误信息
|
||||
#[serde(default)]
|
||||
pub error: Option<String>,
|
||||
/// 创建时间(Unix 时间戳,毫秒)
|
||||
pub created_at: i64,
|
||||
/// 自定义请求头(Cookie / Referer 等)
|
||||
#[serde(default)]
|
||||
pub headers: HashMap<String, String>,
|
||||
}
|
||||
|
||||
impl DownloadTask {
|
||||
/// 文件完整路径(最终文件名)
|
||||
pub fn file_path(&self) -> std::path::PathBuf {
|
||||
std::path::PathBuf::from(&self.dir).join(&self.filename)
|
||||
}
|
||||
|
||||
/// 下载临时文件路径(下载未完成时使用,完成后重命名为 file_path)
|
||||
pub fn temp_file_path(&self) -> std::path::PathBuf {
|
||||
std::path::PathBuf::from(&self.dir).join(format!("{}.thingdl", self.filename))
|
||||
}
|
||||
|
||||
/// 更新已下载总量(聚合所有分段)
|
||||
pub fn recalc_completed(&mut self) {
|
||||
if self.segments.is_empty() {
|
||||
return;
|
||||
}
|
||||
self.completed_size = self.segments.iter().map(|s| s.completed).sum();
|
||||
}
|
||||
}
|
||||
|
||||
/// 下载设置
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
#[serde(rename_all = "camelCase")]
|
||||
pub struct DownloaderSettings {
|
||||
/// 下载目录
|
||||
#[serde(default = "default_download_dir")]
|
||||
pub download_dir: String,
|
||||
/// 最大同时下载数
|
||||
#[serde(default = "default_max_concurrent")]
|
||||
pub max_concurrent: u32,
|
||||
/// 单任务最大连接数(多线程分段数)
|
||||
#[serde(default = "default_max_connections")]
|
||||
pub max_connections: u32,
|
||||
/// 断点续传
|
||||
#[serde(default = "default_true")]
|
||||
pub continue_download: bool,
|
||||
/// 全局速度限制 KB/s(0=不限)
|
||||
#[serde(default)]
|
||||
pub global_speed_limit: u64,
|
||||
/// 扩展 HTTP API 端口
|
||||
#[serde(default = "default_extension_port")]
|
||||
pub extension_port: u16,
|
||||
/// 扩展认证密钥(空=不认证)
|
||||
#[serde(default)]
|
||||
pub extension_secret: String,
|
||||
/// 删除任务时是否同时删除已下载的文件
|
||||
#[serde(default)]
|
||||
pub delete_files_on_remove: bool,
|
||||
}
|
||||
|
||||
fn default_max_concurrent() -> u32 {
|
||||
5
|
||||
}
|
||||
fn default_max_connections() -> u32 {
|
||||
8
|
||||
}
|
||||
fn default_true() -> bool {
|
||||
true
|
||||
}
|
||||
fn default_extension_port() -> u16 {
|
||||
16800
|
||||
}
|
||||
|
||||
fn default_download_dir() -> String {
|
||||
dirs::download_dir()
|
||||
.or_else(|| {
|
||||
std::env::var("USERPROFILE")
|
||||
.ok()
|
||||
.map(|p| std::path::PathBuf::from(p).join("Downloads"))
|
||||
})
|
||||
.map(|p| p.to_string_lossy().to_string())
|
||||
.unwrap_or_else(|| "downloads".to_string())
|
||||
}
|
||||
|
||||
impl Default for DownloaderSettings {
|
||||
fn default() -> Self {
|
||||
Self {
|
||||
download_dir: default_download_dir(),
|
||||
max_concurrent: default_max_concurrent(),
|
||||
max_connections: default_max_connections(),
|
||||
continue_download: true,
|
||||
global_speed_limit: 0,
|
||||
extension_port: default_extension_port(),
|
||||
extension_secret: String::new(),
|
||||
delete_files_on_remove: false,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// HEAD/Range 探测结果
|
||||
pub struct ProbeResult {
|
||||
/// 文件大小(字节),None=未知
|
||||
pub total_size: Option<u64>,
|
||||
/// 是否支持 Range 请求
|
||||
pub supports_resume: bool,
|
||||
/// 从 Content-Disposition 或 URL 推断的文件名
|
||||
pub filename: Option<String>,
|
||||
}
|
||||
Reference in New Issue
Block a user