diff --git a/src-tauri/Cargo.lock b/src-tauri/Cargo.lock index 28b73de..858d5db 100644 --- a/src-tauri/Cargo.lock +++ b/src-tauri/Cargo.lock @@ -6953,6 +6953,7 @@ dependencies = [ "serde", "serde_json", "tauri", + "tauri-plugin-deep-link", "thiserror 2.0.18", "tracing", "windows-sys 0.60.2", diff --git a/src-tauri/Cargo.toml b/src-tauri/Cargo.toml index 364791c..68f74bf 100644 --- a/src-tauri/Cargo.toml +++ b/src-tauri/Cargo.toml @@ -39,7 +39,7 @@ sentry = { version = "0.48", default-features = false, features = ["backtrace", fluxa_core = { path = "../../fluxa-core", default-features = false, features = ["desktop"] } fluxa_streaming_engine = { path = "../../fluxa-core/fluxa-streaming-engine" } tauri-plugin-deep-link = "2.4.9" -tauri-plugin-single-instance = "2" +tauri-plugin-single-instance = { version = "2", features = ["deep-link"] } tauri-plugin-updater = "2" tauri-plugin-process = "2" tauri-plugin-dialog = "2" diff --git a/src-tauri/src/lib.rs b/src-tauri/src/lib.rs index 4947ff6..b29f803 100644 --- a/src-tauri/src/lib.rs +++ b/src-tauri/src/lib.rs @@ -52,7 +52,7 @@ use roku::*; use storage::*; use serde_json::{json, Value}; -use std::collections::HashMap; +use std::collections::{HashMap, HashSet}; use std::fs; use std::path::PathBuf; use std::sync::atomic::AtomicBool; @@ -109,6 +109,76 @@ pub struct DesktopState { pub sleep_inhibitor: Mutex, } +#[derive(Clone, serde::Serialize)] +struct OAuthCodePayload { + code: String, + state: Option, +} + +struct PendingOAuthCallbacks { + state: Mutex, +} + +struct OAuthCallbackState { + callbacks: HashMap, + consumed: HashSet, +} + +impl Default for PendingOAuthCallbacks { + fn default() -> Self { + Self { + state: Mutex::new(OAuthCallbackState { + callbacks: HashMap::new(), + consumed: HashSet::new(), + }), + } + } +} + +fn queue_oauth_callback( + app: &tauri::AppHandle, + service: &str, + code: String, + state: Option, +) { + let event = match service { + "trakt" => "trakt-oauth-code", + "anilist" => "anilist-oauth-code", + "simkl" => "simkl-oauth-code", + _ => return, + }; + let payload = OAuthCodePayload { code, state }; + let callback_id = format!("{service}:{}", payload.code); + if let Ok(mut state) = app.state::().state.lock() { + if state.consumed.contains(&callback_id) { + return; + } + state.callbacks.insert(service.to_string(), payload.clone()); + } + let _ = app.emit(event, payload); +} + +#[tauri::command] +fn take_oauth_callback( + service: String, + callbacks: State, +) -> Result, String> { + match service.as_str() { + "trakt" | "anilist" | "simkl" => { + let mut state = callbacks + .state + .lock() + .map_err(|_| "OAuth callback state is unavailable".to_string())?; + let payload = state.callbacks.remove(&service); + if let Some(payload) = &payload { + state.consumed.insert(format!("{service}:{}", payload.code)); + } + Ok(payload) + } + _ => Err("unsupported OAuth service".to_string()), + } +} + impl Default for DesktopState { fn default() -> Self { Self { @@ -649,24 +719,29 @@ pub fn run() { if !arg.starts_with("fluxa://") { continue; } - let query = arg.split('?').nth(1); - let code = query - .and_then(|q| q.split('&').find(|p| p.starts_with("code="))) - .map(|p| p.trim_start_matches("code=").to_string()); - let state = query - .and_then(|q| q.split('&').find(|p| p.starts_with("state="))) - .map(|p| p.trim_start_matches("state=").to_string()); + let url = match tauri::Url::parse(arg) { + Ok(url) => url, + Err(_) => continue, + }; + let code = url + .query_pairs() + .find(|(key, _)| key == "code") + .map(|(_, value)| value.into_owned()); + let state = url + .query_pairs() + .find(|(key, _)| key == "state") + .map(|(_, value)| value.into_owned()); if let Some(code) = code { - let evt = if arg.contains("/trakt") { - "trakt-oauth-code" + let service = if arg.contains("/trakt") { + "trakt" } else if arg.contains("/anilist") { - "anilist-oauth-code" + "anilist" } else if arg.contains("/simkl") { - "simkl-oauth-code" + "simkl" } else { continue; }; - let _ = app.emit(evt, json!({ "code": code, "state": state })); + queue_oauth_callback(app, service, code, state); } } })) @@ -708,6 +783,7 @@ pub fn run() { .plugin(tauri_plugin_dialog::init()) .plugin(tauri_plugin_deep_link::init()) .manage(DesktopState::default()) + .manage(PendingOAuthCallbacks::default()) .manage(discord_presence::DiscordPresenceState::default()) .manage(cast::CastState::default()) .manage(chromecast::ChromecastState::default()) @@ -716,6 +792,9 @@ pub fn run() { .manage(cast_proxy::CastProxyState::default()) .manage(trailer_proxy::TrailerProxyState::default()) .setup(|app| { + #[cfg(any(target_os = "linux", all(debug_assertions, target_os = "windows")))] + app.deep_link().register_all()?; + let data_dir = app .path() .app_data_dir() @@ -763,16 +842,16 @@ pub fn run() { .find(|(k, _)| k == "state") .map(|(_, v)| v.into_owned()); if let Some(code) = code { - let evt = if s.contains("/trakt") { - "trakt-oauth-code" + let service = if s.contains("/trakt") { + "trakt" } else if s.contains("/anilist") { - "anilist-oauth-code" + "anilist" } else if s.contains("/simkl") { - "simkl-oauth-code" + "simkl" } else { continue; }; - let _ = handle.emit(evt, json!({ "code": code, "state": state })); + queue_oauth_callback(&handle, service, code, state); } else { let _ = handle.emit("deep-link-opened", json!({ "url": s })); } @@ -905,6 +984,7 @@ pub fn run() { player_set_episodes, player_clear_episodes, get_oauth_client_id, + take_oauth_callback, nuvio_request, trakt_device_start, trakt_device_poll, diff --git a/src/components/settings/AccountSection.tsx b/src/components/settings/AccountSection.tsx index 48e4a78..85a0308 100644 --- a/src/components/settings/AccountSection.tsx +++ b/src/components/settings/AccountSection.tsx @@ -291,11 +291,12 @@ export function AccountSection({ traktStateRef.current = state; const authUrl = `https://trakt.tv/oauth/authorize?response_type=code&client_id=${traktClientId}&redirect_uri=${encodeURIComponent('fluxa://oauth/trakt')}&state=${state}`; setAuthUrl('trakt', authUrl); - await shellOpen(authUrl); - - const unlisten = await listen('trakt-oauth-code', async (event) => { - unlisten(); - if (event.payload.state !== traktStateRef.current) { + let unlisten: (() => void) | undefined; + const consumeCallback = async () => { + const payload = await invoke('take_oauth_callback', { service: 'trakt' }); + if (!payload) return; + unlisten?.(); + if (payload.state !== traktStateRef.current) { setTraktError(t('settings.oauth_state_mismatch')); setAuthUrl('trakt'); setTraktBusy(false); @@ -304,7 +305,7 @@ export function AccountSection({ traktStateRef.current = null; setAuthUrl('trakt'); try { - const tokenJson = await invoke('trakt_oauth_exchange', { code: event.payload.code }); + const tokenJson = await invoke('trakt_oauth_exchange', { code: payload.code }); const tokens = JSON.parse(tokenJson) as TraktTokenResponse; const updated: UserProfile = { ...activeProfile, traktAccessToken: tokens.access_token, traktRefreshToken: tokens.refresh_token, traktTokenExpiresAt: tokens.created_at + tokens.expires_in }; await saveProfile(updated); @@ -314,7 +315,10 @@ export function AccountSection({ } finally { setTraktBusy(false); } - }); + }; + unlisten = await listen('trakt-oauth-code', () => { void consumeCallback(); }); + await shellOpen(authUrl); + void consumeCallback(); } catch (err) { setTraktError(err instanceof Error ? err.message : String(err)); setAuthUrl('trakt'); @@ -344,11 +348,12 @@ export function AccountSection({ anilistStateRef.current = state; const authUrl = `https://anilist.co/api/v2/oauth/authorize?response_type=code&client_id=${anilistClientId}&redirect_uri=${encodeURIComponent('fluxa://oauth/anilist')}&state=${state}`; setAuthUrl('anilist', authUrl); - await shellOpen(authUrl); - - const unlisten = await listen('anilist-oauth-code', async (event) => { - unlisten(); - if (event.payload.state !== anilistStateRef.current) { + let unlisten: (() => void) | undefined; + const consumeCallback = async () => { + const payload = await invoke('take_oauth_callback', { service: 'anilist' }); + if (!payload) return; + unlisten?.(); + if (payload.state !== anilistStateRef.current) { setAnilistError(t('settings.oauth_state_mismatch')); setAuthUrl('anilist'); setAnilistBusy(false); @@ -357,7 +362,7 @@ export function AccountSection({ anilistStateRef.current = null; setAuthUrl('anilist'); try { - const tokenJson = await invoke('anilist_oauth_exchange', { code: event.payload.code }); + const tokenJson = await invoke('anilist_oauth_exchange', { code: payload.code }); const tokens = JSON.parse(tokenJson) as { access_token: string; refresh_token?: string; expires_in?: number }; const updated: UserProfile = { ...activeProfile, @@ -372,7 +377,10 @@ export function AccountSection({ } finally { setAnilistBusy(false); } - }); + }; + unlisten = await listen('anilist-oauth-code', () => { void consumeCallback(); }); + await shellOpen(authUrl); + void consumeCallback(); } catch (err) { setAnilistError(err instanceof Error ? err.message : String(err)); setAuthUrl('anilist'); @@ -408,11 +416,12 @@ export function AccountSection({ simklStateRef.current = state; const authUrl = `https://simkl.com/oauth/authorize?response_type=code&client_id=${simklClientId}&redirect_uri=${encodeURIComponent('fluxa://oauth/simkl')}&state=${state}`; setAuthUrl('simkl', authUrl); - await shellOpen(authUrl); - - const unlisten = await listen('simkl-oauth-code', async (event) => { - unlisten(); - if (event.payload.state !== simklStateRef.current) { + let unlisten: (() => void) | undefined; + const consumeCallback = async () => { + const payload = await invoke('take_oauth_callback', { service: 'simkl' }); + if (!payload) return; + unlisten?.(); + if (payload.state !== simklStateRef.current) { setSimklError(t('settings.oauth_state_mismatch')); setAuthUrl('simkl'); setSimklBusy(false); @@ -421,7 +430,7 @@ export function AccountSection({ simklStateRef.current = null; setAuthUrl('simkl'); try { - const tokenJson = await invoke('simkl_oauth_exchange', { code: event.payload.code }); + const tokenJson = await invoke('simkl_oauth_exchange', { code: payload.code }); const tokens = JSON.parse(tokenJson) as { access_token: string; refresh_token?: string }; const updated: UserProfile = { ...activeProfile, simklAccessToken: tokens.access_token, simklRefreshToken: tokens.refresh_token }; await saveProfile(updated); @@ -431,7 +440,10 @@ export function AccountSection({ } finally { setSimklBusy(false); } - }); + }; + unlisten = await listen('simkl-oauth-code', () => { void consumeCallback(); }); + await shellOpen(authUrl); + void consumeCallback(); } catch (err) { setSimklError(err instanceof Error ? err.message : String(err)); setAuthUrl('simkl');