diff --git a/src-tauri/src/extension_server.rs b/src-tauri/src/extension_server.rs index 32157af..fbf8b1b 100644 --- a/src-tauri/src/extension_server.rs +++ b/src-tauri/src/extension_server.rs @@ -1,7 +1,7 @@ use axum::{ body::Bytes, extract::State, - http::{HeaderMap, StatusCode, Method}, + http::{HeaderMap, Method, StatusCode}, routing::{get, post}, Router, }; @@ -15,12 +15,17 @@ use std::sync::atomic::{AtomicBool, Ordering}; use std::sync::{Arc, Mutex, RwLock}; use std::time::{SystemTime, UNIX_EPOCH}; use tauri::{AppHandle, Emitter, Manager}; +use tauri_plugin_store::StoreExt; +use tokio::sync::watch; use tower_http::cors::{Any, CorsLayer}; use ts_rs::TS; pub const EXTENSION_SERVER_PORT: u16 = 23522; +pub const EXTENSION_SERVER_PORT_RANGE: std::ops::RangeInclusive = + EXTENSION_SERVER_PORT..=23531; const MAX_URL_COUNT: usize = 200; const SIGNATURE_MAX_AGE_MS: u64 = 60_000; +const MAIN_QUEUE_ID: &str = "00000000-0000-0000-0000-000000000001"; type HmacSha256 = Hmac; pub type SharedExtensionToken = Arc>; @@ -44,6 +49,10 @@ struct ExtensionRequest { silent: bool, #[serde(default)] filename: Option, + #[serde(default)] + headers: Option, + #[serde(default)] + cookies: Option, } #[derive(Clone, Serialize, TS)] @@ -53,12 +62,15 @@ pub struct ExtensionDownload { referer: Option, silent: bool, filename: Option, + headers: Option, + cookies: Option, } pub async fn start_server( app_handle: AppHandle, pairing_token: SharedExtensionToken, frontend_ready: SharedFrontendReady, + mut shutdown_rx: watch::Receiver, ) -> Result<(), String> { let state = ServerState { app_handle, @@ -81,19 +93,41 @@ pub async fn start_server( .layer(cors) .with_state(state); - let listener = tokio::net::TcpListener::bind(("127.0.0.1", EXTENSION_SERVER_PORT)) - .await - .map_err(|e| format!("Failed to bind extension server to port {}: {}", EXTENSION_SERVER_PORT, e))?; - - println!("Browser extension server bound to 127.0.0.1:{}", EXTENSION_SERVER_PORT); + let (port, listener) = bind_extension_listener().await?; + + println!("Browser extension server bound to 127.0.0.1:{port}"); axum::serve(listener, app) + .with_graceful_shutdown(async move { + if *shutdown_rx.borrow() { + return; + } + let _ = shutdown_rx.changed().await; + }) .await .map_err(|e| format!("Server error: {}", e))?; Ok(()) } +async fn bind_extension_listener() -> Result<(u16, tokio::net::TcpListener), String> { + let mut errors = Vec::new(); + for port in EXTENSION_SERVER_PORT_RANGE { + match tokio::net::TcpListener::bind(("127.0.0.1", port)).await { + Ok(listener) => return Ok((port, listener)), + Err(error) => { + errors.push(format!("{port}: {error}")); + } + } + } + Err(format!( + "Failed to bind extension server in port range {}-{} ({})", + EXTENSION_SERVER_PORT, + *EXTENSION_SERVER_PORT_RANGE.end(), + errors.join("; ") + )) +} + async fn ping_handler( State(state): State, headers: HeaderMap, @@ -103,12 +137,18 @@ async fn ping_handler( return StatusCode::SERVICE_UNAVAILABLE; } - let signature = match headers.get("x-firelink-signature").and_then(|v| v.to_str().ok()) { + let signature = match headers + .get("x-firelink-signature") + .and_then(|v| v.to_str().ok()) + { Some(v) => v, None => return StatusCode::FORBIDDEN, }; - let timestamp_str = match headers.get("x-firelink-timestamp").and_then(|v| v.to_str().ok()) { + let timestamp_str = match headers + .get("x-firelink-timestamp") + .and_then(|v| v.to_str().ok()) + { Some(v) => v, None => return StatusCode::FORBIDDEN, }; @@ -125,16 +165,18 @@ async fn download_handler( headers: HeaderMap, body: Bytes, ) -> Result { - if !state.frontend_ready.load(Ordering::Acquire) { - return Err(StatusCode::SERVICE_UNAVAILABLE); - } - - let signature = match headers.get("x-firelink-signature").and_then(|v| v.to_str().ok()) { + let signature = match headers + .get("x-firelink-signature") + .and_then(|v| v.to_str().ok()) + { Some(v) => v, None => return Err(StatusCode::FORBIDDEN), }; - let timestamp_str = match headers.get("x-firelink-timestamp").and_then(|v| v.to_str().ok()) { + let timestamp_str = match headers + .get("x-firelink-timestamp") + .and_then(|v| v.to_str().ok()) + { Some(v) => v, None => return Err(StatusCode::FORBIDDEN), }; @@ -172,13 +214,136 @@ async fn download_handler( } } - if state.app_handle.emit("extension-add-download", download).is_err() { + if enqueue_extension_download(&state.app_handle, &download) + .await + .is_err() + { return Err(StatusCode::INTERNAL_SERVER_ERROR); } Ok(StatusCode::OK) } +async fn enqueue_extension_download( + app_handle: &AppHandle, + download: &ExtensionDownload, +) -> Result<(), String> { + let Some(settings) = read_settings(app_handle) else { + return Err("settings unavailable".to_string()); + }; + let state = app_handle.state::(); + let mut created_items = Vec::new(); + let mut tasks = Vec::new(); + + for url in &download.urls { + let filename = download + .filename + .as_deref() + .filter(|_| download.urls.len() == 1) + .map(str::to_string) + .unwrap_or_else(|| filename_from_url(url)); + let category = crate::parity::get_file_category(filename.clone()); + let category_key = format!("{category:?}"); + let destination = settings + .download_directories + .get(&category_key) + .cloned() + .filter(|path| !path.trim().is_empty()) + .unwrap_or_else(|| settings.default_download_path.clone()); + let id = uuid::Uuid::new_v4().to_string(); + let headers = merge_headers(download.referer.as_deref(), download.headers.as_deref()); + let item = crate::ipc::DownloadItem { + id: id.clone(), + url: url.clone(), + file_name: filename.clone(), + status: crate::ipc::DownloadStatus::Queued, + fraction: Some(0.0), + speed: Some("-".to_string()), + eta: Some("-".to_string()), + size: None, + category, + date_added: chrono::Utc::now().to_rfc3339(), + connections: Some(settings.per_server_connections), + speed_limit: (!settings.global_speed_limit.trim().is_empty()) + .then(|| settings.global_speed_limit.clone()), + username: None, + password: None, + headers: headers.clone(), + checksum: None, + cookies: download.cookies.clone(), + mirrors: None, + destination: Some(destination.clone()), + is_media: Some(false), + media_format_selector: None, + queue_id: MAIN_QUEUE_ID.to_string(), + }; + let task = crate::queue::EnqueueItem { + id, + url: url.clone(), + destination, + filename, + connections: Some(settings.per_server_connections), + speed_limit: (!settings.global_speed_limit.trim().is_empty()) + .then(|| settings.global_speed_limit.clone()), + username: None, + password: None, + headers, + checksum: None, + cookies: download.cookies.clone(), + mirrors: None, + user_agent: (!settings.custom_user_agent.trim().is_empty()) + .then(|| settings.custom_user_agent.clone()), + max_tries: Some(settings.max_automatic_retries), + proxy: None, + format_selector: None, + cookie_source: None, + is_media: Some(false), + } + .into_task(); + created_items.push(item); + tasks.push(task); + } + + state.queue_manager.enqueue_many(tasks).await; + let _ = app_handle.emit("extension-downloads-queued", created_items); + Ok(()) +} + +fn read_settings(app_handle: &AppHandle) -> Option { + let store = app_handle.store("store.bin").ok()?; + let settings_value = store.get("settings")?; + let settings_text = settings_value.as_str()?; + serde_json::from_str(settings_text).ok() +} + +fn merge_headers(referer: Option<&str>, headers: Option<&str>) -> Option { + let mut values = Vec::new(); + if let Some(referer) = referer { + values.push(format!("Referer: {referer}")); + } + if let Some(headers) = headers { + values.extend( + headers + .lines() + .map(str::trim) + .filter(|line| !line.is_empty()) + .map(str::to_string), + ); + } + (!values.is_empty()).then(|| values.join("\n")) +} + +fn filename_from_url(raw_url: &str) -> String { + Url::parse(raw_url) + .ok() + .and_then(|url| { + url.path_segments() + .and_then(Iterator::last) + .and_then(|segment| sanitize_filename(segment)) + }) + .unwrap_or_else(|| "download".to_string()) +} + fn normalize_download(payload: ExtensionRequest) -> Option { let mut seen = HashSet::new(); let urls = payload @@ -203,6 +368,8 @@ fn normalize_download(payload: ExtensionRequest) -> Option { referer, silent: payload.silent, filename, + headers: payload.headers.filter(|value| !value.trim().is_empty()), + cookies: payload.cookies.filter(|value| !value.trim().is_empty()), }) } diff --git a/src-tauri/src/lib.rs b/src-tauri/src/lib.rs index f34db78..f213fbe 100644 --- a/src-tauri/src/lib.rs +++ b/src-tauri/src/lib.rs @@ -494,6 +494,7 @@ pub struct AppState { pub download_coordinator: download::DownloadCoordinator, pub extension_pairing_token: extension_server::SharedExtensionToken, pub extension_frontend_ready: extension_server::SharedFrontendReady, + pub extension_server_shutdown: tokio::sync::watch::Sender, pub aria2_port: u16, pub aria2_secret: String, pub media_semaphore: Arc, @@ -1444,6 +1445,7 @@ pub fn run() { let server_pairing_token = extension_pairing_token.clone(); let extension_frontend_ready = Arc::new(AtomicBool::new(false)); let server_frontend_ready = extension_frontend_ready.clone(); + let (extension_server_shutdown_tx, extension_server_shutdown_rx) = tokio::sync::watch::channel(false); let aria2_port = std::net::TcpListener::bind("127.0.0.1:0") .and_then(|listener| listener.local_addr()) @@ -1535,6 +1537,7 @@ pub fn run() { download_coordinator: download::DownloadCoordinator::spawn(app.handle().clone()), extension_pairing_token, extension_frontend_ready, + extension_server_shutdown: extension_server_shutdown_tx.clone(), aria2_port, aria2_secret: aria2_secret.clone(), media_semaphore: Arc::new(tokio::sync::Semaphore::new(3)), @@ -1767,6 +1770,7 @@ pub fn run() { ext_app_handle, server_pairing_token.clone(), server_frontend_ready.clone(), + extension_server_shutdown_rx, ).await { eprintln!("Browser extension server unavailable: {error}"); } @@ -1809,8 +1813,14 @@ pub fn run() { parity::get_system_proxy, parity::get_file_category, parity::check_for_updates, parity::is_supported_media, parity::get_supported_media_domains, parity::create_category_directories ]) - .run(tauri::generate_context!()) - .expect("error while running tauri application"); + .build(tauri::generate_context!()) + .expect("error while building tauri application") + .run(|app_handle, event| { + if let tauri::RunEvent::ExitRequested { .. } = event { + let state = app_handle.state::(); + let _ = state.extension_server_shutdown.send(true); + } + }); } mod extension_server; mod scheduler; diff --git a/src/App.tsx b/src/App.tsx index 0bd60bc..4a83fe2 100644 --- a/src/App.tsx +++ b/src/App.tsx @@ -238,10 +238,27 @@ function App() { const unlistenExtension = listen('extension-add-download', (event) => { useDownloadStore.getState().handleExtensionDownload(event.payload); }); + const unlistenExtensionQueued = listen('extension-downloads-queued', (event) => { + const store = useDownloadStore.getState(); + const incoming = event.payload; + const existing = new Set(store.downloads.map(download => download.id)); + const additions = incoming.filter(download => !existing.has(download.id)); + if (additions.length === 0) return; + useDownloadStore.setState(state => ({ + downloads: [...state.downloads, ...additions], + pendingOrder: [ + ...state.pendingOrder, + ...additions + .filter(download => download.status === 'queued') + .map(download => download.id) + .filter(id => !state.pendingOrder.includes(id)), + ], + })); + }); const unlistenDeepLink = listen('deep-link-add-download', (event) => { useDownloadStore.getState().openAddModalWithUrls(event.payload); }); - Promise.all([unlistenExtension, unlistenDeepLink]) + Promise.all([unlistenExtension, unlistenExtensionQueued, unlistenDeepLink]) .then(() => invoke('set_extension_frontend_ready', { ready: true })) .catch(error => console.error('Failed to activate browser extension integration:', error)); @@ -250,6 +267,7 @@ function App() { unlistenComplete.then(f => f()); unlistenFailed.then(f => f()); unlistenExtension.then(f => f()); + unlistenExtensionQueued.then(f => f()); unlistenDeepLink.then(f => f()); }; }, []); diff --git a/src/bindings/ExtensionDownload.ts b/src/bindings/ExtensionDownload.ts index 2474503..fdd659b 100644 --- a/src/bindings/ExtensionDownload.ts +++ b/src/bindings/ExtensionDownload.ts @@ -1,3 +1,3 @@ // This file was generated by [ts-rs](https://github.com/Aleph-Alpha/ts-rs). Do not edit this file manually. -export type ExtensionDownload = { urls: Array, referer: string | null, silent: boolean, filename: string | null, }; +export type ExtensionDownload = { urls: Array, referer: string | null, silent: boolean, filename: string | null, headers: string | null, cookies: string | null, }; diff --git a/src/ipc.ts b/src/ipc.ts index fcad0e4..dee631b 100644 --- a/src/ipc.ts +++ b/src/ipc.ts @@ -2,6 +2,7 @@ import { invoke as tauriInvoke } from '@tauri-apps/api/core'; import { error as logError } from '@tauri-apps/plugin-log'; import { listen as tauriListen, type Event, type EventCallback, type UnlistenFn } from '@tauri-apps/api/event'; import type { DownloadCategory } from './bindings/DownloadCategory'; +import type { DownloadItem } from './bindings/DownloadItem'; import type { DownloadProgressEvent } from './bindings/DownloadProgressEvent'; import type { DownloadStatus } from './bindings/DownloadStatus'; import type { ExtensionDownload } from './bindings/ExtensionDownload'; @@ -126,6 +127,7 @@ type EventMap = { 'download-complete': string; 'download-failed': string; 'extension-add-download': ExtensionDownload; + 'extension-downloads-queued': DownloadItem[]; 'deep-link-add-download': string; };