diff --git a/src-tauri/src/lib.rs b/src-tauri/src/lib.rs index 38f9809..7134d0a 100644 --- a/src-tauri/src/lib.rs +++ b/src-tauri/src/lib.rs @@ -2887,7 +2887,7 @@ async fn pause_download( log::info!("pause_download called for id: {}", id); let active_kind = state.queue_manager.active_kind(&id).await; - state.queue_manager.remove_from_pending(&id).await; + let removed_pending = state.queue_manager.remove_from_pending(&id).await; let gid = state.queue_manager.aria2_gid_for_download(&id); if let Some(gid) = gid.as_deref().filter(|gid| !gid.starts_with("native:")) { @@ -2933,6 +2933,10 @@ async fn pause_download( return Ok(()); } + if matches!(active_kind, Some(crate::queue::TaskKind::Aria2)) { + state.queue_manager.cancel_aria2_retries(&id).await; + } + let (tx, rx) = tokio::sync::oneshot::channel(); if matches!(active_kind, Some(crate::queue::TaskKind::Media)) { state @@ -2953,6 +2957,9 @@ async fn pause_download( if !matches!(active_kind, Some(crate::queue::TaskKind::Media)) { state.queue_manager.release_permit(&id).await; } + if removed_pending || matches!(active_kind, Some(crate::queue::TaskKind::Aria2)) { + state.queue_manager.release_registered_id(&id).await; + } use tauri::Emitter; let _ = app_handle.emit( "download-state", diff --git a/src-tauri/src/queue.rs b/src-tauri/src/queue.rs index 6b5364d..2a5b62e 100644 --- a/src-tauri/src/queue.rs +++ b/src-tauri/src/queue.rs @@ -168,6 +168,10 @@ impl QueueManager { self.registered_ids.lock().await.remove(id); } + async fn is_registered(&self, id: &str) -> bool { + self.registered_ids.lock().await.contains(id) + } + /// Enqueue a task. Checks the centralized `registered_ids` for deduplication. pub async fn push(&self, task: QueuedTask) -> Result<(), String> { let id = task.id.clone(); @@ -390,7 +394,30 @@ impl QueueManager { .insert(id.clone(), task.payload.clone()); self.aria2_retry_strikes.lock().await.remove(&id); match self.spawner.add_uri(&id, &task.payload).await { - Ok(gid) => self.remember_gid(id.clone(), gid).await, + Ok(gid) => { + let cancelled = self.aria2_retry_cancelled.lock().await.contains(&id); + if cancelled || !self.is_registered(&id).await { + log::info!( + "aria2 dispatch cancellation [{}]: removing late gid {}", + id, + gid + ); + if !gid.starts_with("native:") { + if let Err(error) = self.spawner.remove_uri(&gid).await { + log::warn!( + "aria2 dispatch cancellation [{}]: failed to remove late gid {}: {}", + id, + gid, + error + ); + } + } + self.clear_aria2_retry_state(&id).await; + self.release_permit(&id).await; + return; + } + self.remember_gid(id.clone(), gid).await; + } Err(error) => { self.clear_aria2_retry_state(&id).await; self.emit_failed(&id, error); @@ -488,7 +515,9 @@ impl QueueManager { PendingOutcome::Error(error) => { if error.to_ascii_lowercase().contains("checksum") { log::warn!("Checksum error detected for {}, cleaning up assets", id); - if let Ok(primary_path) = crate::download_ownership::primary_path_for_id(&self.app_handle, id) { + if let Ok(primary_path) = + crate::download_ownership::primary_path_for_id(&self.app_handle, id) + { if let Some(path) = primary_path.as_deref() { let _ = crate::remove_download_assets(path, &self.app_handle).await; } @@ -957,7 +986,12 @@ async fn probe_bounded_range_support( .redirect(reqwest::redirect::Policy::limited(5)) .timeout(std::time::Duration::from_secs(10)); - if let Some(proxy) = payload.proxy.as_deref().map(str::trim).filter(|value| !value.is_empty()) { + if let Some(proxy) = payload + .proxy + .as_deref() + .map(str::trim) + .filter(|value| !value.is_empty()) + { if proxy.eq_ignore_ascii_case("none") { builder = builder.no_proxy(); } else { @@ -1131,7 +1165,12 @@ impl SidecarSpawner for ProductionSpawner { if !header_list.is_empty() { options.insert("header".to_string(), serde_json::json!(header_list)); } - if let Some(prox) = payload.proxy.as_deref().map(str::trim).filter(|s| !s.is_empty()) { + if let Some(prox) = payload + .proxy + .as_deref() + .map(str::trim) + .filter(|s| !s.is_empty()) + { if prox.eq_ignore_ascii_case("none") { options.insert("all-proxy".to_string(), serde_json::json!("")); } else { diff --git a/src-tauri/tests/queue_manager.rs b/src-tauri/tests/queue_manager.rs index d0fd46f..59ca49a 100644 --- a/src-tauri/tests/queue_manager.rs +++ b/src-tauri/tests/queue_manager.rs @@ -15,6 +15,44 @@ struct CountingSpawner { native_calls: AtomicUsize, } +struct DelayedAria2Spawner { + gid_tx: tokio::sync::Mutex>>, + remove_uri_calls: AtomicUsize, +} + +impl DelayedAria2Spawner { + fn new(gid_tx: tokio::sync::oneshot::Sender<()>) -> Self { + Self { + gid_tx: tokio::sync::Mutex::new(Some(gid_tx)), + remove_uri_calls: AtomicUsize::new(0), + } + } +} + +#[async_trait::async_trait] +impl SidecarSpawner for DelayedAria2Spawner { + async fn add_uri(&self, _id: &str, _payload: &SpawnPayload) -> Result { + let tx = self.gid_tx.lock().await.take().expect("gid release sender"); + let _ = tx.send(()); + tokio::time::sleep(Duration::from_millis(50)).await; + Ok("late-gid".to_string()) + } + + async fn remove_uri(&self, gid: &str) -> Result<(), String> { + assert_eq!(gid, "late-gid"); + self.remove_uri_calls.fetch_add(1, Ordering::SeqCst); + Ok(()) + } + + async fn run_media(&self, _id: &str, _payload: &SpawnPayload) -> Result<(), String> { + unreachable!("media is not used by delayed aria2 tests") + } + + async fn run_native(&self, _id: &str, _payload: &SpawnPayload) -> Result<(), String> { + unreachable!("native is not used by delayed aria2 tests") + } +} + impl CountingSpawner { fn new() -> Self { Self { @@ -476,6 +514,39 @@ async fn aria2_completion_forgets_gid_and_releases_permit() { assert_eq!(mgr.available_permits(), 1); } +#[tokio::test] +async fn late_aria2_gid_after_cancellation_is_removed_without_leaking_permit() { + let app = mock_builder() + .build(mock_context(noop_assets())) + .expect("mock app"); + let (gid_started_tx, gid_started_rx) = tokio::sync::oneshot::channel(); + let spawner = Arc::new(DelayedAria2Spawner::new(gid_started_tx)); + let manager = Arc::new(QueueManager::test_new( + app.handle().clone(), + 1, + spawner.clone(), + )); + manager.push(aria2_task("late")).await.unwrap(); + + let dispatcher = { + let manager = Arc::clone(&manager); + tokio::spawn(async move { manager.run_dispatcher().await }) + }; + + gid_started_rx.await.expect("add_uri should start"); + manager.cancel_aria2_retries("late").await; + manager.release_registered_id("late").await; + manager.release_permit("late").await; + + tokio::time::sleep(Duration::from_millis(100)).await; + + assert!(manager.aria2_gid_for_download("late").is_none()); + assert_eq!(manager.available_permits(), 1); + assert_eq!(spawner.remove_uri_calls.load(Ordering::SeqCst), 1); + + dispatcher.abort(); +} + #[tokio::test] async fn move_up_down_reorders_pending() { use firelink_lib::ipc::QueueDirection; diff --git a/src/store/useDownloadStore.test.ts b/src/store/useDownloadStore.test.ts index 3f92fa8..14cb4ab 100644 --- a/src/store/useDownloadStore.test.ts +++ b/src/store/useDownloadStore.test.ts @@ -118,6 +118,7 @@ describe('useDownloadStore', () => { { id: '2', url: 'http://test2', fileName: 'f2', destination: '/tmp', status: 'queued', category: 'Other', dateAdded: '', queueId: 'MAIN', hasBeenDispatched: false }, ] as any[], backendRegisteredIds: new Set(['1']), // 1 is already registered, so it skips dispatch + pendingOrder: ['1'], }); vi.mocked(ipc.invokeCommand).mockImplementation(async (cmd: string) => { @@ -134,6 +135,28 @@ describe('useDownloadStore', () => { expect((enqueues[0] as any)[1].item.id).toBe('2'); }); + it('repairs stale queued backend registrations before accepting a queue start', async () => { + useDownloadStore.setState({ + downloads: [ + { id: 'stale', url: 'http://test', fileName: 'f', destination: '/tmp', status: 'queued', category: 'Other', dateAdded: '', queueId: 'MAIN', hasBeenDispatched: true }, + ] as any[], + backendRegisteredIds: new Set(['stale']), + pendingOrder: [], + }); + + vi.mocked(ipc.invokeCommand).mockImplementation(async (cmd: string) => { + if (cmd === 'resume_download') return false; + if (cmd === 'get_pending_order') return ['stale']; + return undefined; + }); + + expect(await useDownloadStore.getState().startQueue('MAIN')).toEqual(['stale']); + + const calls = vi.mocked(ipc.invokeCommand).mock.calls; + expect(calls.some(call => call[0] === 'resume_download')).toBe(true); + expect(calls.some(call => call[0] === 'enqueue_download')).toBe(true); + }); + it('does not overwrite a downloading event received while starting a queue', async () => { useDownloadStore.setState({ downloads: [ diff --git a/src/store/useDownloadStore.ts b/src/store/useDownloadStore.ts index 81c70ee..429eace 100644 --- a/src/store/useDownloadStore.ts +++ b/src/store/useDownloadStore.ts @@ -625,12 +625,22 @@ export const useDownloadStore = create((set, get) => ({ const acceptedIds: string[] = []; for (const item of runnable) { + const backendRegistered = get().backendRegisteredIds.has(item.id); + const backendPending = get().pendingOrder.includes(item.id); + + if (item.status === 'queued' && backendRegistered && !backendPending) { + if (await get().resumeDownload(item.id)) { + acceptedIds.push(item.id); + } + continue; + } + if ( item.status === 'ready' || item.status === 'staged' || item.status === 'failed' || !item.hasBeenDispatched || - !get().backendRegisteredIds.has(item.id) + !backendRegistered ) { if (await dispatchItem(item.id)) { const current = get().downloads.find(download => download.id === item.id);