diff --git a/src-tauri/src/lib.rs b/src-tauri/src/lib.rs index ccba712..cd5b583 100644 --- a/src-tauri/src/lib.rs +++ b/src-tauri/src/lib.rs @@ -5687,6 +5687,18 @@ async fn set_concurrent_limit( Ok(()) } +#[tauri::command] +async fn set_download_speed_limit( + state: tauri::State<'_, AppState>, + id: String, + limit: Option, +) -> Result<(), String> { + state + .queue_manager + .set_aria2_download_speed_limit(&id, limit) + .await +} + pub(crate) fn normalize_speed_limit_for_aria2(limit: &str) -> Option { let trimmed = limit.trim(); if trimmed.is_empty() { @@ -9315,7 +9327,7 @@ pub fn run() { authorize_keychain_access, acknowledge_pairing_token_change, check_file_exists, toggle_tray_icon, set_extension_pairing_token, - get_extension_server_port, set_extension_frontend_ready, ack_extension_download, set_concurrent_limit, set_global_speed_limit, remove_download, + get_extension_server_port, set_extension_frontend_ready, ack_extension_download, set_concurrent_limit, set_download_speed_limit, set_global_speed_limit, remove_download, detach_download_for_reconfigure, enqueue_download, enqueue_many, cancel_enqueue_generation, move_in_queue, move_many_in_queue, remove_from_queue, get_pending_order, commands::reveal_in_file_manager, commands::open_downloaded_file, diff --git a/src-tauri/src/queue.rs b/src-tauri/src/queue.rs index a4912b5..cc354ac 100644 --- a/src-tauri/src/queue.rs +++ b/src-tauri/src/queue.rs @@ -165,6 +165,17 @@ pub trait SidecarSpawner: Send + Sync + 'static { Err("aria2 connection refresh is unavailable".to_string()) } + /// Change one active aria2 transfer's runtime download cap. Media + /// runners intentionally keep the default implementation: yt-dlp reads + /// its limit only when the process starts. + async fn set_download_speed_limit( + &self, + _gid: &str, + _limit: Option<&str>, + ) -> Result<(), String> { + Err("live aria2 speed limits are unavailable".to_string()) + } + /// Run a media download to completion. The permit is parked for the full /// duration; release is handled by QueueManager on the runner's exit. async fn run_media( @@ -544,6 +555,68 @@ impl QueueManager { .is_some() } + /// Change an active aria2 transfer's speed cap without replacing its GID + /// or queue permit. The per-download control lock and post-RPC ownership + /// check make a late response harmless if terminal cleanup or a retry + /// transition wins the lifecycle race. + pub async fn set_aria2_download_speed_limit( + &self, + id: &str, + limit: Option, + ) -> Result<(), String> { + let normalized_limit = match limit.as_deref().map(str::trim) { + None | Some("") => None, + Some(raw) => Some( + crate::normalize_speed_limit_for_aria2(raw) + .ok_or_else(|| "invalid download speed limit".to_string())?, + ), + }; + let _control_guard = self.acquire_aria2_control(id).await; + + if !self.is_registered(id).await + || !matches!(self.active_kind(id).await, Some(TaskKind::Aria2)) + { + return Err("download is not an active aria2 transfer".to_string()); + } + let gid = self + .aria2_gid_for_download(id) + .ok_or_else(|| "active aria2 transfer has no gid".to_string())?; + let expected_mapping = self + .aria2_gid_mapping(&gid) + .ok_or_else(|| "active aria2 transfer has no current gid mapping".to_string())?; + if expected_mapping.id != id { + return Err("aria2 gid belongs to another download".to_string()); + } + if !self + .is_aria2_control_epoch_current(id, expected_mapping.epoch) + .await + { + return Err("active aria2 transfer has a stale control epoch".to_string()); + } + + self.spawner + .set_download_speed_limit(&gid, normalized_limit.as_deref()) + .await?; + + let still_current = self.is_registered(id).await + && matches!(self.active_kind(id).await, Some(TaskKind::Aria2)) + && self + .is_aria2_control_epoch_current(id, expected_mapping.epoch) + .await + && self.is_current_aria2_gid_mapping(&gid, &expected_mapping) + && self.aria2_gid_for_download(id).as_deref() == Some(gid.as_str()); + if !still_current { + return Err("download lifecycle changed while setting speed limit".to_string()); + } + + let mut payloads = self.aria2_payloads.lock().await; + let payload = payloads + .get_mut(id) + .ok_or_else(|| "active aria2 transfer payload is unavailable".to_string())?; + payload.speed_limit = normalized_limit; + Ok(()) + } + /// Pop the next task, or None if empty. pub async fn pop_front(&self) -> Option { self.pending.lock().await.pop_front() @@ -1441,22 +1514,38 @@ impl QueueManager { return; } - if !this.active_permits.lock().await.contains_key(&id_for_task) + // Serialize the payload snapshot and addUri with live speed + // changes. Without this guard, a speed update could change the + // old GID and payload just before this worker re-added a new GID + // from a stale clone, silently losing the user's limit. + let control_guard = this.acquire_aria2_control(&id_for_task).await; + let stale_before_add = !this.active_permits.lock().await.contains_key(&id_for_task) || this.is_aria2_retry_cancelled(&id_for_task).await || !this .is_aria2_control_epoch_current(&id_for_task, retry_epoch) .await || !this.is_registered(&id_for_task).await - || this.aria2_gid_for_download(&id_for_task).as_deref() != Some(retry_gid.as_str()) - { + || this.aria2_gid_for_download(&id_for_task).as_deref() != Some(retry_gid.as_str()); + let current_payload = this.aria2_payloads.lock().await.get(&id_for_task).cloned(); + let Some(current_payload) = current_payload else { + drop(control_guard); + this.finish_aria2_retry(&id_for_task, &retry_gid, retry_epoch) + .await; + return; + }; + if stale_before_add { + drop(control_guard); this.finish_aria2_retry(&id_for_task, &retry_gid, retry_epoch) .await; return; } - match this.spawner.add_uri(&id_for_task, &payload).await { + match this + .spawner + .add_uri(&id_for_task, ¤t_payload) + .await + { Ok(new_gid) => { - let control_guard = this.acquire_aria2_control(&id_for_task).await; let stale = this.is_aria2_retry_cancelled(&id_for_task).await || !this .is_aria2_control_epoch_current(&id_for_task, retry_epoch) @@ -1505,7 +1594,6 @@ impl QueueManager { } } Err(retry_error) => { - let control_guard = this.acquire_aria2_control(&id_for_task).await; let stale = this.is_aria2_retry_cancelled(&id_for_task).await || !this .is_aria2_control_epoch_current(&id_for_task, retry_epoch) @@ -2109,6 +2197,30 @@ impl SidecarSpawner for ProductionSpawner { } } + async fn set_download_speed_limit( + &self, + gid: &str, + limit: Option<&str>, + ) -> Result<(), String> { + let state = self.app_handle.state::(); + let limit = limit.unwrap_or("0"); + let result = crate::rpc_call( + state.aria2_port.load(std::sync::atomic::Ordering::Relaxed), + &state.aria2_secret, + "aria2.changeOption", + serde_json::json!([gid, {"max-download-limit": limit}]), + ) + .await + .map_err(|error| format!("aria2 changeOption failed for gid {gid}: {error}"))?; + match result.as_str() { + Some("OK") => Ok(()), + Some(value) => Err(format!( + "aria2.changeOption returned unexpected result {value} for gid {gid}" + )), + None => Err("aria2.changeOption returned a non-string result".to_string()), + } + } + async fn refresh_uri(&self, gid: &str) -> Result { let state = self.app_handle.state::(); let port = state.aria2_port.load(std::sync::atomic::Ordering::Relaxed); diff --git a/src-tauri/tests/queue_manager.rs b/src-tauri/tests/queue_manager.rs index 6032673..ee1f4d4 100644 --- a/src-tauri/tests/queue_manager.rs +++ b/src-tauri/tests/queue_manager.rs @@ -13,6 +13,12 @@ use tokio::time::timeout; struct CountingSpawner { add_uri_calls: AtomicUsize, media_calls: AtomicUsize, + speed_limit_calls: AtomicUsize, + last_speed_limit: std::sync::Mutex>, + add_speed_limits: std::sync::Mutex>>, + block_speed_limit: std::sync::atomic::AtomicBool, + speed_limit_started: tokio::sync::Notify, + speed_limit_release: tokio::sync::Notify, } struct DelayedAria2Spawner { @@ -102,6 +108,12 @@ impl CountingSpawner { Self { add_uri_calls: AtomicUsize::new(0), media_calls: AtomicUsize::new(0), + speed_limit_calls: AtomicUsize::new(0), + last_speed_limit: std::sync::Mutex::new(None), + add_speed_limits: std::sync::Mutex::new(Vec::new()), + block_speed_limit: std::sync::atomic::AtomicBool::new(false), + speed_limit_started: tokio::sync::Notify::new(), + speed_limit_release: tokio::sync::Notify::new(), } } } @@ -133,13 +145,33 @@ impl SidecarSpawner for RefreshOutcomeSpawner { #[async_trait::async_trait] impl firelink_lib::queue::SidecarSpawner for CountingSpawner { - async fn add_uri(&self, _id: &str, _payload: &SpawnPayload) -> Result { + async fn add_uri(&self, _id: &str, payload: &SpawnPayload) -> Result { self.add_uri_calls.fetch_add(1, Ordering::SeqCst); + self.add_speed_limits + .lock() + .unwrap() + .push(payload.speed_limit.clone()); Ok(format!("gid-{}", self.add_uri_calls.load(Ordering::SeqCst))) } async fn remove_uri(&self, _gid: &str) -> Result<(), String> { Ok(()) } + async fn set_download_speed_limit( + &self, + _gid: &str, + limit: Option<&str>, + ) -> Result<(), String> { + self.speed_limit_calls.fetch_add(1, Ordering::SeqCst); + *self.last_speed_limit.lock().unwrap() = limit.map(str::to_string); + if self + .block_speed_limit + .load(std::sync::atomic::Ordering::SeqCst) + { + self.speed_limit_started.notify_one(); + self.speed_limit_release.notified().await; + } + Ok(()) + } async fn run_media(&self, _id: &str, _payload: &SpawnPayload, _generation: u64) -> Result<(), String> { self.media_calls.fetch_add(1, Ordering::SeqCst); Ok(()) @@ -258,6 +290,188 @@ async fn ensure_aria2_permit_does_not_double_acquire() { assert_eq!(mgr.available_permits(), 2); } +#[tokio::test] +async fn live_aria2_speed_limit_updates_the_current_gid_and_payload() { + let (manager, spawner) = make_manager(1); + let manager = Arc::new(manager); + let mut task = aria2_task("speed-limit"); + task.payload.speed_limit = Some("1M".to_string()); + manager.push(task).await.unwrap(); + + let dispatcher = { + let manager = Arc::clone(&manager); + tokio::spawn(async move { manager.run_dispatcher().await }) + }; + timeout(Duration::from_secs(1), async { + loop { + if manager.aria2_gid_for_download("speed-limit").is_some() { + break; + } + tokio::time::sleep(Duration::from_millis(10)).await; + } + }) + .await + .expect("aria2 dispatch should register a gid"); + + manager + .set_aria2_download_speed_limit("speed-limit", Some("512K".to_string())) + .await + .unwrap(); + assert_eq!(spawner.speed_limit_calls.load(Ordering::SeqCst), 1); + assert_eq!( + spawner.last_speed_limit.lock().unwrap().as_deref(), + Some("512K") + ); + assert!(manager.aria2_speed_limited("speed-limit").await); + + manager + .set_aria2_download_speed_limit("speed-limit", None) + .await + .unwrap(); + assert_eq!(spawner.speed_limit_calls.load(Ordering::SeqCst), 2); + assert!(spawner.last_speed_limit.lock().unwrap().is_none()); + assert!(!manager.aria2_speed_limited("speed-limit").await); + + manager + .apply_completion( + "speed-limit", + firelink_lib::queue::PendingOutcome::Complete, + ) + .await; + dispatcher.abort(); +} + +#[tokio::test] +async fn live_aria2_speed_limit_rejects_invalid_and_non_active_requests() { + let (manager, spawner) = make_manager(1); + assert!(manager + .set_aria2_download_speed_limit("missing", Some("not-a-rate".to_string())) + .await + .is_err()); + assert!(manager + .set_aria2_download_speed_limit("missing", Some("512K".to_string())) + .await + .is_err()); + assert_eq!(spawner.speed_limit_calls.load(Ordering::SeqCst), 0); +} + +#[tokio::test] +async fn live_aria2_speed_limit_does_not_update_payload_after_gid_replacement() { + let (manager, spawner) = make_manager(1); + let manager = Arc::new(manager); + let mut task = aria2_task("speed-stale"); + task.payload.speed_limit = Some("1M".to_string()); + manager.push(task).await.unwrap(); + + let dispatcher = { + let manager = Arc::clone(&manager); + tokio::spawn(async move { manager.run_dispatcher().await }) + }; + timeout(Duration::from_secs(1), async { + loop { + if manager.aria2_gid_for_download("speed-stale").is_some() { + break; + } + tokio::time::sleep(Duration::from_millis(10)).await; + } + }) + .await + .expect("aria2 dispatch should register a gid"); + + spawner + .block_speed_limit + .store(true, std::sync::atomic::Ordering::SeqCst); + let started = spawner.speed_limit_started.notified(); + let setter = { + let manager = Arc::clone(&manager); + tokio::spawn(async move { + manager + .set_aria2_download_speed_limit("speed-stale", Some("512K".to_string())) + .await + }) + }; + timeout(Duration::from_secs(1), started) + .await + .expect("speed RPC should start"); + + manager + .remember_gid("speed-stale".to_string(), "gid-replaced".to_string()) + .await; + spawner.speed_limit_release.notify_one(); + assert!(setter.await.unwrap().is_err()); + assert!(manager.aria2_speed_limited("speed-stale").await); + + manager + .apply_completion( + "speed-stale", + firelink_lib::queue::PendingOutcome::Complete, + ) + .await; + dispatcher.abort(); +} + +#[tokio::test] +async fn retry_readds_aria2_with_the_latest_live_speed_limit() { + use firelink_lib::queue::PendingOutcome; + + let (mgr, spawner) = make_manager(1); + let manager = Arc::new(mgr); + let mut task = aria2_task("speed-retry"); + task.payload.max_tries = Some(1); + task.payload.speed_limit = Some("1M".to_string()); + manager.push(task).await.unwrap(); + let dispatcher = { + let manager = Arc::clone(&manager); + tokio::spawn(async move { manager.run_dispatcher().await }) + }; + + timeout(Duration::from_secs(1), async { + loop { + if spawner.add_uri_calls.load(Ordering::SeqCst) >= 1 { + break; + } + tokio::time::sleep(Duration::from_millis(10)).await; + } + }) + .await + .expect("initial aria2 add should run"); + manager + .handle_aria2_event( + "gid-1", + PendingOutcome::Error( + "aria2 error code 1: Failed to receive data, cause: protocol error".to_string(), + ), + ) + .await; + manager + .set_aria2_download_speed_limit("speed-retry", Some("512K".to_string())) + .await + .unwrap(); + + timeout(Duration::from_secs(4), async { + loop { + if spawner.add_uri_calls.load(Ordering::SeqCst) >= 2 { + break; + } + tokio::time::sleep(Duration::from_millis(25)).await; + } + }) + .await + .expect("retry should re-add after backoff"); + assert_eq!( + spawner.add_speed_limits.lock().unwrap().as_slice(), + &[Some("1M".to_string()), Some("512K".to_string())] + ); + + manager + .handle_aria2_event( + "gid-2", + PendingOutcome::Error("permanent failure".to_string()), + ) + .await; + dispatcher.abort(); +} + #[tokio::test] async fn stale_aria2_permit_candidate_cannot_replace_current_permit() { let (mgr, _spawner) = make_manager(2); diff --git a/src/ipc.ts b/src/ipc.ts index d730264..690c902 100644 --- a/src/ipc.ts +++ b/src/ipc.ts @@ -48,6 +48,7 @@ type CommandMap = { perform_system_action: { args: { action: PostQueueAction }; result: void }; ack_schedule_trigger: { args: { action: 'start' | 'stop'; key: string }; result: void }; set_concurrent_limit: { args: { limit: number }; result: void }; + set_download_speed_limit: { args: { id: string; limit: string | null }; result: void }; set_global_speed_limit: { args: { limit: string | null }; result: void }; request_automation_permission: { args: undefined; result: void }; check_automation_permission: { args: undefined; result: void };