From 44c089fadf44486fc73aa83682c624f425d1c603 Mon Sep 17 00:00:00 2001 From: NimBold Date: Mon, 15 Jun 2026 18:42:03 +0330 Subject: [PATCH] refactor(aria2): modernize architecture with websockets and state mapping - Replace polling loops with global tokio_tungstenite WebSocket listener - Map Download IDs to Aria2 GIDs in AppState - Route all standard downloads to Aria2 daemon, mapping all config options - Refactor path resolution into a DRY helper function - Add missing terabyte multiplier to frontend speed parser --- src-tauri/src/lib.rs | 413 +++++++++++++++++++++---------------------- src/App.tsx | 2 +- 2 files changed, 204 insertions(+), 211 deletions(-) diff --git a/src-tauri/src/lib.rs b/src-tauri/src/lib.rs index 5755eed..e4c8731 100644 --- a/src-tauri/src/lib.rs +++ b/src-tauri/src/lib.rs @@ -357,16 +357,7 @@ async fn open_file(app: tauri::AppHandle, path: String) -> Result<(), String> { println!("open_file called for path: {}", path); use tauri_plugin_opener::OpenerExt; - let mut resolved_dest = std::path::PathBuf::from(&path); - if path.starts_with("~/") { - if let Ok(home) = app.path().home_dir() { - resolved_dest = home.join(&path[2..]); - } - } else if path == "~" { - if let Ok(home) = app.path().home_dir() { - resolved_dest = home; - } - } + let resolved_dest = resolve_path(&path, &app); if !is_safe_path(&resolved_dest, &app) { return Err("Path traversal blocked".to_string()); @@ -380,16 +371,7 @@ async fn show_in_folder(app: tauri::AppHandle, path: String) -> Result<(), Strin println!("show_in_folder called for path: {}", path); use tauri_plugin_opener::OpenerExt; - let mut resolved_dest = std::path::PathBuf::from(&path); - if path.starts_with("~/") { - if let Ok(home) = app.path().home_dir() { - resolved_dest = home.join(&path[2..]); - } - } else if path == "~" { - if let Ok(home) = app.path().home_dir() { - resolved_dest = home; - } - } + let resolved_dest = resolve_path(&path, &app); if !is_safe_path(&resolved_dest, &app) { return Err("Path traversal blocked".to_string()); @@ -468,6 +450,7 @@ pub struct AppState { pub aria2_secret: String, pub media_semaphore: Arc, pub sleep_preventer: Arc>>, + pub aria2_gids: Arc>>, } #[derive(Clone, Serialize, TS)] @@ -480,6 +463,22 @@ pub struct DownloadProgressEvent { size: Option, } + +fn resolve_path(path: &str, app_handle: &tauri::AppHandle) -> std::path::PathBuf { + use tauri::Manager; + let mut resolved = std::path::PathBuf::from(path); + if path.starts_with("~/") { + if let Ok(home) = app_handle.path().home_dir() { + resolved = home.join(&path[2..]); + } + } else if path == "~" { + if let Ok(home) = app_handle.path().home_dir() { + resolved = home; + } + } + resolved +} + fn collect_download_uris(url: &str, mirrors: Option<&str>) -> Vec { let mut uris = Vec::new(); for uri in std::iter::once(url).chain(mirrors.into_iter().flat_map(str::lines)) { @@ -635,97 +634,88 @@ async fn start_download( .unwrap_or("download") .to_string(); - let mut resolved_dest = std::path::PathBuf::from(&destination); - if destination.starts_with("~/") { - if let Ok(home) = app_handle.path().home_dir() { - resolved_dest = home.join(&destination[2..]); - } - } else if destination == "~" { - if let Ok(home) = app_handle.path().home_dir() { - resolved_dest = home; - } - } - + let resolved_dest = resolve_path(&destination, &app_handle); if !is_safe_path(&resolved_dest, &app_handle) { return Err(AppError::Internal("Path traversal blocked".to_string())); } - if connections.unwrap_or(1) > 1 || checksum.is_some() { - println!("Routing multi-part/checksum download to aria2: {}", id); - let mut options = serde_json::Map::new(); - options.insert("dir".to_string(), serde_json::json!(resolved_dest.to_string_lossy().to_string())); - options.insert("out".to_string(), serde_json::json!(safe_filename)); - if let Some(conn) = connections { - options.insert("split".to_string(), serde_json::json!(conn.to_string())); - options.insert("max-connection-per-server".to_string(), serde_json::json!(conn.to_string())); - } - if let Some(speed) = speed_limit { - options.insert("max-download-limit".to_string(), serde_json::json!(speed)); - } - if let Some(user) = username { - options.insert("http-user".to_string(), serde_json::json!(user)); - } - if let Some(pass) = password { - options.insert("http-passwd".to_string(), serde_json::json!(pass)); - } - if let Some(chk) = checksum { - options.insert("checksum".to_string(), serde_json::json!(chk)); - } - if let Some(ua) = user_agent { - options.insert("user-agent".to_string(), serde_json::json!(ua)); - } - if let Some(prox) = proxy { - options.insert("all-proxy".to_string(), serde_json::json!(prox)); - } else { - options.insert("all-proxy".to_string(), serde_json::json!("")); - options.insert("http-proxy".to_string(), serde_json::json!("")); - options.insert("https-proxy".to_string(), serde_json::json!("")); - options.insert("no-proxy".to_string(), serde_json::json!("*")); - } - if let Some(cook) = cookies { - options.insert("header".to_string(), serde_json::json!(format!("Cookie: {}", cook))); - } - - let params = serde_json::json!([ - [url], - options - ]); - - let result = rpc_call(state.aria2_port, &state.aria2_secret, "aria2.addUri", params) - .await - .map_err(|e| AppError::Internal(e))?; - - let gid = result.as_str().unwrap_or("").to_string(); - - let port = state.aria2_port; - let secret = state.aria2_secret.clone(); - let app_handle_clone = app_handle.clone(); - let id_clone = id.clone(); - tauri::async_runtime::spawn(async move { - poll_aria2(app_handle_clone, port, secret, id_clone, gid).await; - }); - - return Ok(()); - } - state - .download_coordinator - .send(download::DownloadCmd::Start(download::DownloadPayload { - id: download_id, - urls: collect_download_uris(&url, mirrors.as_deref()), - output_path: resolved_dest.join(safe_filename), - speed_limit, - username, - password, - headers, - cookies, - user_agent, - max_tries: max_tries.unwrap_or(1).max(1) as u32, - proxy, - })) - .await - .map_err(AppError::Internal)?; + let mt = max_tries.unwrap_or(1).max(1) as u32; - Ok(()) + let mut options = serde_json::Map::new(); + options.insert("dir".to_string(), serde_json::json!(resolved_dest.to_string_lossy().to_string())); + options.insert("out".to_string(), serde_json::json!(safe_filename)); + + let conn = connections.unwrap_or(1); + options.insert("split".to_string(), serde_json::json!(conn.to_string())); + options.insert("max-connection-per-server".to_string(), serde_json::json!(conn.to_string())); + options.insert("max-tries".to_string(), serde_json::json!(mt.to_string())); + + if let Some(speed) = &speed_limit { + options.insert("max-download-limit".to_string(), serde_json::json!(speed)); + } + if let Some(user) = &username { + options.insert("http-user".to_string(), serde_json::json!(user)); + } + if let Some(pass) = &password { + options.insert("http-passwd".to_string(), serde_json::json!(pass)); + } + if let Some(chk) = &checksum { + options.insert("checksum".to_string(), serde_json::json!(chk)); + } + if let Some(ua) = &user_agent { + options.insert("user-agent".to_string(), serde_json::json!(ua)); + } + + let mut header_list = Vec::new(); + if let Some(cook) = &cookies { + header_list.push(format!("Cookie: {}", cook)); + } + if let Some(hdrs) = &headers { + for line in hdrs.lines() { + if !line.trim().is_empty() { + header_list.push(line.trim().to_string()); + } + } + } + if !header_list.is_empty() { + options.insert("header".to_string(), serde_json::json!(header_list)); + } + + if let Some(prox) = &proxy { + options.insert("all-proxy".to_string(), serde_json::json!(prox)); + } + + let uris = collect_download_uris(&url, mirrors.as_deref()); + let params = serde_json::json!([uris, options]); + + match rpc_call(state.aria2_port, &state.aria2_secret, "aria2.addUri", params).await { + Ok(result) => { + let gid = result.as_str().unwrap_or("").to_string(); + state.aria2_gids.write().unwrap().insert(id.clone(), gid); + return Ok(()); + } + Err(e) => { + eprintln!("aria2 failed, falling back to native coordinator: {}", e); + state + .download_coordinator + .send(download::DownloadCmd::Start(download::DownloadPayload { + id: download_id, + urls: collect_download_uris(&url, mirrors.as_deref()), + output_path: resolved_dest.join(safe_filename), + speed_limit, + username, + password, + headers, + cookies, + user_agent, + max_tries: mt, + proxy, + })) + .await + .map_err(AppError::Internal)?; + return Ok(()); + } + } } #[tauri::command] @@ -810,16 +800,7 @@ pub(crate) async fn start_media_download_internal( let ytdlp_path = resource_dir.join("binaries").join(get_binary_name("yt-dlp")); let ffmpeg_path = resource_dir.join("binaries").join(get_binary_name("ffmpeg")); - let mut resolved_dest = std::path::PathBuf::from(&destination); - if destination.starts_with("~/") { - if let Ok(home) = app_handle.path().home_dir() { - resolved_dest = home.join(&destination[2..]); - } - } else if destination == "~" { - if let Ok(home) = app_handle.path().home_dir() { - resolved_dest = home; - } - } + let resolved_dest = resolve_path(&destination, &app_handle); if !is_safe_path(&resolved_dest, &app_handle) { return Err("Path traversal blocked".to_string()); @@ -1015,11 +996,17 @@ pub(crate) async fn start_media_download_internal( #[tauri::command] async fn pause_download(state: tauri::State<'_, AppState>, id: String) -> Result<(), String> { println!("pause_download called for id: {}", id); + + let gid = state.aria2_gids.read().unwrap().get(&id).cloned(); + if let Some(g) = gid { + let _ = rpc_call(state.aria2_port, &state.aria2_secret, "aria2.pause", serde_json::json!([g])).await; + } + if let Ok(download_id) = Uuid::parse_str(&id) { - state + let _ = state .download_coordinator .send(download::DownloadCmd::Pause(download_id)) - .await?; + .await; } state.download_coordinator.pause_media(id).await } @@ -1182,16 +1169,7 @@ fn get_free_space(app_handle: tauri::AppHandle, path: String) -> Result = None; @@ -1246,16 +1224,7 @@ fn delete_keychain_password(id: String) -> Result<(), String> { #[tauri::command] fn check_file_exists(app_handle: tauri::AppHandle, path: String) -> bool { use tauri::Manager; - let mut resolved_dest = std::path::PathBuf::from(&path); - if path.starts_with("~/") { - if let Ok(home) = app_handle.path().home_dir() { - resolved_dest = home.join(&path[2..]); - } - } else if path == "~" { - if let Ok(home) = app_handle.path().home_dir() { - resolved_dest = home; - } - } + let resolved_dest = resolve_path(&path, &app_handle); if !is_safe_path(&resolved_dest, &app_handle) { return false; } @@ -1265,16 +1234,7 @@ fn check_file_exists(app_handle: tauri::AppHandle, path: String) -> bool { #[tauri::command] fn delete_file(app_handle: tauri::AppHandle, path: String) -> Result<(), String> { use tauri::Manager; - let mut resolved_dest = std::path::PathBuf::from(&path); - if path.starts_with("~/") { - if let Ok(home) = app_handle.path().home_dir() { - resolved_dest = home.join(&path[2..]); - } - } else if path == "~" { - if let Ok(home) = app_handle.path().home_dir() { - resolved_dest = home; - } - } + let resolved_dest = resolve_path(&path, &app_handle); if !is_safe_path(&resolved_dest, &app_handle) { return Err("Path traversal blocked".to_string()); } @@ -1442,6 +1402,10 @@ pub fn run() { .plugin(tauri_plugin_deep_link::init()) .manage(Aria2DaemonGuard(std::sync::Mutex::new(None))) .setup(move |app| { + let aria2_gids = Arc::new(RwLock::new(std::collections::HashMap::new())); + let aria2_gids_clone1 = aria2_gids.clone(); + let aria2_gids_clone2 = aria2_gids.clone(); + app.manage(AppState { download_coordinator: download::DownloadCoordinator::spawn(app.handle().clone()), extension_pairing_token, @@ -1450,6 +1414,7 @@ pub fn run() { aria2_secret: aria2_secret.clone(), media_semaphore: Arc::new(tokio::sync::Semaphore::new(3)), sleep_preventer: Arc::new(Mutex::new(None)), + aria2_gids, }); let deep_link_app = app.handle().clone(); app.deep_link().on_open_url(move |event| { @@ -1495,6 +1460,96 @@ pub fn run() { Err(e) => eprintln!("Failed to spawn aria2c daemon: {}", e), } + let app_handle_ws = app.handle().clone(); + let ws_port = aria2_port; + tauri::async_runtime::spawn(async move { + tokio::time::sleep(std::time::Duration::from_secs(1)).await; + let ws_url = format!("ws://127.0.0.1:{}/jsonrpc", ws_port); + if let Ok((ws_stream, _)) = tokio_tungstenite::connect_async(&ws_url).await { + use futures_util::StreamExt; + let (_, mut read) = ws_stream.split(); + while let Some(msg) = read.next().await { + if let Ok(tokio_tungstenite::tungstenite::Message::Text(text)) = msg { + if let Ok(json) = serde_json::from_str::(&text) { + if let Some(method) = json.get("method").and_then(|m| m.as_str()) { + if let Some(params) = json.get("params").and_then(|p| p.as_array()) { + if let Some(event) = params.first().and_then(|p| p.as_object()) { + if let Some(gid) = event.get("gid").and_then(|g| g.as_str()) { + let id = { + let map = aria2_gids_clone1.read().unwrap(); + map.iter().find(|(_, g)| *g == gid).map(|(i, _)| i.clone()) + }; + if let Some(id) = id { + use tauri::Emitter; + match method { + "aria2.onDownloadComplete" => { + let _ = app_handle_ws.emit("download-complete", id); + } + "aria2.onDownloadError" => { + let _ = app_handle_ws.emit("download-failed", id); + } + _ => {} + } + } + } + } + } + } + } + } + } + } + }); + + let app_handle_poll = app.handle().clone(); + let poll_port = aria2_port; + let poll_secret = aria2_secret.clone(); + tauri::async_runtime::spawn(async move { + let mut interval = tokio::time::interval(std::time::Duration::from_millis(1000)); + loop { + interval.tick().await; + let params = serde_json::json!([["gid", "status", "totalLength", "completedLength", "downloadSpeed", "errorMessage"]]); + if let Ok(active_list) = rpc_call(poll_port, &poll_secret, "aria2.tellActive", params).await { + if let Some(active_arr) = active_list.as_array() { + for status_info in active_arr { + let gid = status_info.get("gid").and_then(|s| s.as_str()).unwrap_or(""); + let id = { + let map = aria2_gids_clone2.read().unwrap(); + map.iter().find(|(_, g)| *g == gid).map(|(i, _)| i.clone()) + }; + if let Some(id) = id { + let total = status_info.get("totalLength").and_then(|s| s.as_str()).unwrap_or("0").parse::().unwrap_or(0); + let completed = status_info.get("completedLength").and_then(|s| s.as_str()).unwrap_or("0").parse::().unwrap_or(0); + let speed_bytes = status_info.get("downloadSpeed").and_then(|s| s.as_str()).unwrap_or("0").parse::().unwrap_or(0.0); + + let fraction = if total > 0 { completed as f64 / total as f64 } else { 0.0 }; + let speed = crate::download::format_speed(speed_bytes); + let eta = if speed_bytes > 0.0 && total > completed { + crate::download::format_duration((total - completed) as f64 / speed_bytes) + } else { + "-".to_string() + }; + let size = if total > 0 { + Some(crate::download::format_size(total as f64)) + } else { + None + }; + + use tauri::Emitter; + let _ = app_handle_poll.emit("download-progress", DownloadProgressEvent { + id, + fraction, + speed, + eta, + size, + }); + } + } + } + } + } + }); + let ext_app_handle = app.handle().clone(); tauri::async_runtime::spawn(async move { @@ -1589,65 +1644,3 @@ async fn db_delete_queue(state: tauri::State<'_, crate::db::DbState>, id: String crate::db::delete_queue(&conn, &id).map_err(|e| e.to_string()) } -async fn poll_aria2(app_handle: tauri::AppHandle, port: u16, secret: String, id: String, gid: String) { - use tauri::Emitter; - let mut interval = tokio::time::interval(std::time::Duration::from_millis(1000)); - loop { - interval.tick().await; - let params = serde_json::json!([gid, ["status", "totalLength", "completedLength", "downloadSpeed", "errorMessage"]]); - match rpc_call(port, &secret, "aria2.tellStatus", params).await { - Ok(status_info) => { - let status = status_info.get("status").and_then(|s| s.as_str()).unwrap_or(""); - if status == "active" || status == "waiting" || status == "paused" { - let total_str = status_info.get("totalLength").and_then(|s| s.as_str()).unwrap_or("0"); - let completed_str = status_info.get("completedLength").and_then(|s| s.as_str()).unwrap_or("0"); - let speed_str = status_info.get("downloadSpeed").and_then(|s| s.as_str()).unwrap_or("0"); - - let total = total_str.parse::().unwrap_or(0); - let completed = completed_str.parse::().unwrap_or(0); - let speed_bytes = speed_str.parse::().unwrap_or(0.0); - - let fraction = if total > 0 { completed as f64 / total as f64 } else { 0.0 }; - - let speed = crate::download::format_speed(speed_bytes); - let eta = if speed_bytes > 0.0 && total > completed { - crate::download::format_duration((total - completed) as f64 / speed_bytes) - } else { - "-".to_string() - }; - - let size = if total > 0 { - Some(crate::download::format_size(total as f64)) - } else { - None - }; - - let _ = app_handle.emit("download-progress", serde_json::json!({ - "id": id, - "fraction": fraction, - "speed": speed, - "eta": eta, - "size": size - })); - } else if status == "complete" { - let _ = app_handle.emit("download-complete", id); - break; - } else if status == "error" { - let err = status_info.get("errorMessage").and_then(|s| s.as_str()).unwrap_or("aria2 error").to_string(); - eprintln!("download {} failed: {}", id, err); - let _ = app_handle.emit("download-failed", id.clone()); - break; - } else if status == "removed" { - eprintln!("download {} failed: Removed from aria2", id); - let _ = app_handle.emit("download-failed", id.clone()); - break; - } - } - Err(e) => { - eprintln!("download {} failed: RPC error: {}", id, e); - let _ = app_handle.emit("download-failed", id.clone()); - break; - } - } - } -} diff --git a/src/App.tsx b/src/App.tsx index 6defbf0..8556c78 100644 --- a/src/App.tsx +++ b/src/App.tsx @@ -97,7 +97,7 @@ function App() { if (match) { const amount = Number(match[1]); if (Number.isFinite(amount) && amount > 0) { - const multipliers: Record = { '': 1, k: 1024, m: 1048576, g: 1073741824 }; + const multipliers: Record = { '': 1, k: 1024, m: 1048576, g: 1073741824, t: 1099511627776 }; const bytes = Math.round(amount * multipliers[match[2].toLowerCase()]); formattedLimit = `${bytes}`; }