From 5a7832898bfed47571b246856a4a5ee78577e6a4 Mon Sep 17 00:00:00 2001 From: Bob Lee Date: Wed, 22 Jul 2026 11:33:13 +0800 Subject: [PATCH 1/3] fix: harden pages, remote sync, and subscription auth --- .github/workflows/ci.yml | 13 + Cargo.toml | 1 + src/apps/cli/src/account.rs | 73 +- src/apps/cli/src/account_sync.rs | 48 +- src/apps/desktop/src/api/pages_api.rs | 36 +- .../desktop/src/api/remote_connect_api.rs | 254 +++- .../src/api/remote_workspace_policy.rs | 5 + src/apps/desktop/src/lib.rs | 2 + src/apps/relay-server/src/lib.rs | 4 +- src/apps/relay-server/tests/library_compat.rs | 5 + src/crates/adapters/ai-adapters/Cargo.toml | 2 + .../src/subscription_auth/antigravity.rs | 119 +- .../src/subscription_auth/codex.rs | 42 +- .../ai-adapters/src/subscription_auth/jwt.rs | 2 +- .../ai-adapters/src/subscription_auth/mod.rs | 608 ++++++++- .../oauth_callback_locales.json | 26 + .../src/subscription_auth/oauth_server.rs | 198 ++- .../src/subscription_auth/opencode.rs | 41 +- .../src/subscription_auth/store.rs | 978 +++++++++++++- .../tools/implementations/page_deploy_tool.rs | 120 +- .../implementations/page_publish_tool.rs | 262 +++- .../agentic/tools/pipeline/tool_pipeline.rs | 109 +- .../service/remote_connect/settings_sync.rs | 117 +- .../services/page-function-runtime/src/lib.rs | 131 +- src/crates/services/relay-service/Cargo.toml | 2 +- src/crates/services/relay-service/src/db.rs | 179 ++- src/crates/services/relay-service/src/lib.rs | 4 + .../services/relay-service/src/page_data.rs | 716 +++++++++- .../relay-service/src/page_execution.rs | 260 ++++ .../services/relay-service/src/routes/api.rs | 10 + .../relay-service/src/routes/pages.rs | 1169 ++++++++++++++++- .../src/remote_connect.rs | 11 +- .../src/remote_connect/page_upload.rs | 195 ++- src/mobile-web/src/i18n/I18nProvider.tsx | 17 +- src/mobile-web/src/i18n/messages.ts | 3 + src/mobile-web/src/pages/DevicesPage.tsx | 8 +- .../styles/components/language-toggle.scss | 5 +- .../src/styles/components/pairing.scss | 17 +- src/mobile-web/src/styles/global.scss | 2 + .../src/app/components/NavPanel/MainNav.tsx | 13 +- .../src/app/components/NavPanel/NavPanel.scss | 46 + .../RemoteConnectDialog/AccountPanel.scss | 63 +- .../RemoteConnectDialog/AccountPanel.tsx | 235 +++- .../RemoteConnectDialog.scss | 8 +- .../RemoteConnectDialog.tsx | 93 +- .../RemoteConnectDisclaimer.scss | 35 + .../RemoteConnectDisclaimer.tsx | 35 +- .../src/app/components/SceneBar/types.ts | 1 + src/web-ui/src/app/scenes/SceneViewport.tsx | 3 + .../src/app/scenes/pages/PagesScene.scss | 270 ++++ .../src/app/scenes/pages/PagesScene.test.tsx | 165 +++ .../src/app/scenes/pages/PagesScene.tsx | 674 ++++++++++ src/web-ui/src/app/scenes/registry.ts | 9 + .../modern/PermissionRequestPanel.test.tsx | 35 + .../modern/PermissionRequestPanel.tsx | 58 +- .../src/flow_chat/hooks/useTypewriter.test.ts | 24 + .../src/flow_chat/hooks/useTypewriter.ts | 122 +- .../tool-cards/FileOperationToolCard.tsx | 4 +- .../tool-cards/PageDeployToolDisplay.tsx | 23 +- .../tool-cards/PagePublishToolDisplay.tsx | 27 +- .../account/accountSyncStore.test.ts | 58 + .../account/accountSyncStore.ts | 93 +- src/web-ui/src/infrastructure/api/index.ts | 5 +- .../infrastructure/api/service-api/AIApi.ts | 4 +- .../api/service-api/PageAPI.test.ts | 38 + .../infrastructure/api/service-api/PageAPI.ts | 106 ++ .../api/service-api/RemoteConnectAPI.ts | 2 + .../config/components/AIModelConfig.scss | 108 +- .../config/components/AIModelConfig.tsx | 488 ++++++- .../subscriptionLoginCoordinator.test.ts | 74 ++ .../subscriptionLoginCoordinator.ts | 92 ++ .../i18n/presets/namespaceRegistry.ts | 1 + src/web-ui/src/locales/en-US/common.json | 13 +- src/web-ui/src/locales/en-US/flow-chat.json | 14 + .../src/locales/en-US/scenes/pages.json | 92 ++ .../src/locales/en-US/settings/ai-model.json | 24 +- src/web-ui/src/locales/zh-CN/common.json | 13 +- src/web-ui/src/locales/zh-CN/flow-chat.json | 14 + .../src/locales/zh-CN/scenes/pages.json | 92 ++ .../src/locales/zh-CN/settings/ai-model.json | 24 +- src/web-ui/src/locales/zh-TW/common.json | 13 +- src/web-ui/src/locales/zh-TW/flow-chat.json | 14 + .../src/locales/zh-TW/scenes/pages.json | 92 ++ .../src/locales/zh-TW/settings/ai-model.json | 24 +- .../src/shared/utils/textSelection.test.ts | 49 +- src/web-ui/src/shared/utils/textSelection.ts | 39 +- 86 files changed, 8609 insertions(+), 714 deletions(-) create mode 100644 src/crates/adapters/ai-adapters/src/subscription_auth/oauth_callback_locales.json create mode 100644 src/crates/services/relay-service/src/page_execution.rs create mode 100644 src/web-ui/src/app/scenes/pages/PagesScene.scss create mode 100644 src/web-ui/src/app/scenes/pages/PagesScene.test.tsx create mode 100644 src/web-ui/src/app/scenes/pages/PagesScene.tsx create mode 100644 src/web-ui/src/infrastructure/account/accountSyncStore.test.ts create mode 100644 src/web-ui/src/infrastructure/api/service-api/PageAPI.test.ts create mode 100644 src/web-ui/src/infrastructure/api/service-api/PageAPI.ts create mode 100644 src/web-ui/src/infrastructure/config/components/subscriptionLoginCoordinator.test.ts create mode 100644 src/web-ui/src/infrastructure/config/components/subscriptionLoginCoordinator.ts create mode 100644 src/web-ui/src/locales/en-US/scenes/pages.json create mode 100644 src/web-ui/src/locales/zh-CN/scenes/pages.json create mode 100644 src/web-ui/src/locales/zh-TW/scenes/pages.json diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index 3ed6b77bc5..8a5ec82759 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -144,6 +144,19 @@ jobs: - name: Run core and desktop Rust tests run: cargo test --locked -p bitfun-core -p bitfun-desktop + # These crates own platform-sensitive behavior that is not exercised by + # testing bitfun-core/bitfun-desktop alone. Keep their focused contract + # suites in the OS matrix so Linux success cannot hide Windows/macOS + # regressions in worker interruption, relay storage, or OAuth handling. + - name: Run Page Functions runtime tests + run: cargo test --locked -p bitfun-page-function-runtime + + - name: Run Relay service tests + run: cargo test --locked -p bitfun-relay-service + + - name: Run subscription authentication tests + run: cargo test --locked -p bitfun-ai-adapters --features subscription-auth subscription_auth + # ── Frontend: build ──────────────────────────────────────────────── frontend-build: name: Frontend Build diff --git a/Cargo.toml b/Cargo.toml index 9a9a7b0044..23bc31a43f 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -122,6 +122,7 @@ glob = "0.3" ignore = "0.4" notify = "8.2" dirs = "6.0" +keyring = "4.1.5" dark-light = "1.1" dunce = "1" filetime = "0.2" diff --git a/src/apps/cli/src/account.rs b/src/apps/cli/src/account.rs index c672b29a74..dfa8353f77 100644 --- a/src/apps/cli/src/account.rs +++ b/src/apps/cli/src/account.rs @@ -10,7 +10,7 @@ //! The master key lives in memory only and is lost when the CLI exits. use std::sync::{ - atomic::{AtomicBool, Ordering}, + atomic::{AtomicBool, AtomicU64, AtomicUsize, Ordering}, Arc, OnceLock, }; @@ -27,6 +27,71 @@ static ACCOUNT_SESSION: OnceLock>>> = OnceLock /// The relay URL associated with the current account session. static ACCOUNT_RELAY_URL: OnceLock>>> = OnceLock::new(); +static ACCOUNT_CONTEXT_GENERATION: AtomicU64 = AtomicU64::new(1); +static ACCOUNT_CONTEXT_TRANSITIONS: AtomicUsize = AtomicUsize::new(0); +static ACCOUNT_SYNC_LOCK: tokio::sync::Mutex<()> = tokio::sync::Mutex::const_new(()); + +pub(crate) fn account_context_generation() -> u64 { + ACCOUNT_CONTEXT_GENERATION.load(Ordering::Acquire) +} + +pub(crate) fn account_context_is_current(generation: u64) -> bool { + ACCOUNT_CONTEXT_TRANSITIONS.load(Ordering::Acquire) == 0 + && account_context_generation() == generation +} + +struct AccountContextTransitionPermit; + +impl AccountContextTransitionPermit { + fn begin() -> Self { + ACCOUNT_CONTEXT_TRANSITIONS.fetch_add(1, Ordering::AcqRel); + ACCOUNT_CONTEXT_GENERATION.fetch_add(1, Ordering::AcqRel); + Self + } +} + +impl Drop for AccountContextTransitionPermit { + fn drop(&mut self) { + // Reject work queued during the transition before exposing the newly + // installed (or cleared) account context. + ACCOUNT_CONTEXT_GENERATION.fetch_add(1, Ordering::AcqRel); + ACCOUNT_CONTEXT_TRANSITIONS.fetch_sub(1, Ordering::AcqRel); + } +} + +struct AccountContextTransitionGuard { + sync_guard: Option>, + transition: Option, +} + +impl Drop for AccountContextTransitionGuard { + fn drop(&mut self) { + drop(self.sync_guard.take()); + drop(self.transition.take()); + } +} + +pub(crate) async fn lock_account_sync( + generation: u64, +) -> Result> { + let guard = ACCOUNT_SYNC_LOCK.lock().await; + if !account_context_is_current(generation) { + return Err(anyhow!("account sync cancelled")); + } + Ok(guard) +} + +async fn invalidate_and_wait_for_account_sync() -> AccountContextTransitionGuard { + // Create the permit before awaiting so cancellation cannot leak a + // half-started transition or reopen discovery too early. + let transition = AccountContextTransitionPermit::begin(); + let sync_guard = ACCOUNT_SYNC_LOCK.lock().await; + bitfun_core::service::remote_connect::settings_sync::wait_for_sync_operations_idle().await; + AccountContextTransitionGuard { + sync_guard: Some(sync_guard), + transition: Some(transition), + } +} /// The background device-routing relay client. Holding this keeps the WS /// connection alive (the internal read/write tasks own the socket). Dropping it @@ -79,6 +144,7 @@ pub(crate) async fn is_logged_in() -> bool { /// Attempt to restore a persisted session from disk. Called at startup. /// Returns `Some(user_id)` if a session was restored. pub(crate) async fn try_restore_session() -> Option { + let _sync_guard = invalidate_and_wait_for_account_sync().await; match session_store::load_session_detailed() { Ok(Some(loaded)) => { let user_id = loaded.user_id.clone(); @@ -142,6 +208,7 @@ pub(crate) async fn login_with_credentials( username: &str, password: &str, ) -> Result { + let _sync_guard = invalidate_and_wait_for_account_sync().await; let relay_url = relay_url.trim(); let username = username.trim(); if relay_url.is_empty() { @@ -227,6 +294,8 @@ pub(crate) async fn login_with_credentials( /// Persist the in-memory session after the user accepts the sync choice, then /// start device routing (same as a first login with no cloud settings). pub(crate) async fn finalize_login_after_sync_choice() -> Result<()> { + let generation = account_context_generation(); + let _sync_guard = lock_account_sync(generation).await?; let device = current_device_identity()?; let (session, relay_url) = read_account_context().await?; session_store::save_session_with_device( @@ -340,6 +409,7 @@ async fn is_current_routing_client(client: &Arc) -> bool { /// Log out: tear down routing, revoke the token (best-effort), clear state. pub(crate) async fn logout() -> Result<()> { + let _sync_guard = invalidate_and_wait_for_account_sync().await; stop_device_routing().await; // Take the always-on daemon down with the account: the token is revoked // below, so leaving the daemon connected would keep this device online @@ -393,6 +463,7 @@ async fn handle_relay_event( } RelayEvent::AuthError { message } => { tracing::warn!("Device routing auth error: {message}"); + let _sync_guard = invalidate_and_wait_for_account_sync().await; TOKEN_EXPIRED.store(true, Ordering::Relaxed); // Keep CLI/daemon semantics aligned with Desktop: a relay-rejected // token is no longer a usable local login and must not be restored diff --git a/src/apps/cli/src/account_sync.rs b/src/apps/cli/src/account_sync.rs index b42e30a1c2..3d18130026 100644 --- a/src/apps/cli/src/account_sync.rs +++ b/src/apps/cli/src/account_sync.rs @@ -15,7 +15,9 @@ use bitfun_core::service::config::get_global_config_service; use bitfun_core::service::remote_connect::settings_sync; use bitfun_core::service::remote_connect::{sync_state, AccountClient}; -use crate::account::read_account_context; +use crate::account::{ + account_context_generation, account_context_is_current, lock_account_sync, read_account_context, +}; const UPLOAD_CONCURRENCY_CHUNK: usize = 5; @@ -27,8 +29,19 @@ const UPLOAD_CONCURRENCY_CHUNK: usize = 5; pub(crate) fn start_settings_sync_loop() { let hooks = settings_sync::SettingsSyncHooks { account_context: Some(Arc::new(|| { - Box::pin(async { read_account_context().await }) + Box::pin(async { + let generation = account_context_generation(); + if !account_context_is_current(generation) { + return Err(anyhow!("account context is transitioning")); + } + let (account, relay_url) = read_account_context().await?; + if !account_context_is_current(generation) { + return Err(anyhow!("account context changed while reading")); + } + Ok((account, relay_url, generation)) + }) })), + is_account_context_current: Some(Arc::new(account_context_is_current)), on_settings_applied: Some(Arc::new(|| { crate::peer_host::notify_controllers_settings_changed(); })), @@ -56,6 +69,10 @@ pub(crate) async fn push_settings_after_local_change() { if read_account_context().await.is_err() { crate::account::try_restore_session().await; } + let generation = account_context_generation(); + let Ok(_sync_guard) = lock_account_sync(generation).await else { + return; + }; let Ok((account, relay_url)) = read_account_context().await else { return; }; @@ -208,6 +225,8 @@ pub(crate) async fn run_auto_sync( is_first_login: bool, workspace_path: &Path, ) -> Result { + let generation = account_context_generation(); + let _sync_guard = lock_account_sync(generation).await?; set_progress(|p| { *p = SyncProgress { status: SyncStatus::Syncing, @@ -368,6 +387,8 @@ pub(crate) async fn run_auto_sync( } let _ = sync_state::save(&acct_session.user_id, &sync_state_local); + ensure_session_backup_complete(upload_total, exported)?; + tracing::info!("Auto-sync: settings={settings_synced} exported={exported} imported=0"); emit_progress("done", 100, Some(exported), Some(0), None).await; @@ -378,6 +399,15 @@ pub(crate) async fn run_auto_sync( }) } +fn ensure_session_backup_complete(total: usize, uploaded: usize) -> Result<()> { + if uploaded == total { + return Ok(()); + } + Err(anyhow!( + "session backup incomplete: uploaded {uploaded} of {total}; retry will resume remaining sessions" + )) +} + pub(crate) fn sync_phase_label(progress: &SyncProgress) -> String { match progress.phase.as_str() { "uploading_settings" => "Uploading settings…".into(), @@ -398,3 +428,17 @@ pub(crate) fn sync_phase_label(progress: &SyncProgress) -> String { other => other.to_string(), } } + +#[cfg(test)] +mod tests { + use super::ensure_session_backup_complete; + + #[test] + fn partial_session_backup_is_not_reported_as_success() { + assert!(ensure_session_backup_complete(4, 4).is_ok()); + assert!(ensure_session_backup_complete(4, 1) + .unwrap_err() + .to_string() + .contains("uploaded 1 of 4")); + } +} diff --git a/src/apps/desktop/src/api/pages_api.rs b/src/apps/desktop/src/api/pages_api.rs index 98d4a6f135..cfa25d690d 100644 --- a/src/apps/desktop/src/api/pages_api.rs +++ b/src/apps/desktop/src/api/pages_api.rs @@ -1,10 +1,11 @@ //! BitFun Page Tauri commands (Save Version → Deploy). use bitfun_services_integrations::remote_connect::{ - delete_page_version_on_relay, deploy_page_version_on_relay, list_page_versions_from_relay, - list_pages_from_relay, publish_page_to_relay, save_page_version_to_relay, - unpublish_page_from_relay, update_page_on_relay, PageInfo, PagePublishResult, - PageSaveVersionResult, PageVersionInfo, + create_page_open_link_on_relay, delete_page_from_relay, delete_page_version_on_relay, + deploy_page_version_on_relay, list_page_versions_from_relay, list_pages_from_relay, + publish_page_to_relay, save_page_version_to_relay, unpublish_page_from_relay, + update_page_on_relay, PageInfo, PageOpenLink, PagePublishResult, PageSaveVersionResult, + PageVersionInfo, }; use serde::Deserialize; @@ -43,6 +44,12 @@ pub struct PageDeleteVersionRequest { pub version_id: String, } +#[derive(Debug, Deserialize)] +pub struct PageOpenRequest { + pub slug: String, + pub version_id: Option, +} + /// Save a new immutable version (does not change production). #[tauri::command] pub async fn page_save_version( @@ -94,6 +101,19 @@ pub async fn page_list_versions(request: PageSlugRequest) -> Result Result { + let (session, relay_url) = read_account_context().await?; + create_page_open_link_on_relay( + &relay_url, + &session.token, + &request.slug, + request.version_id.as_deref(), + ) + .await + .map_err(|e| e.to_string()) +} + #[tauri::command] pub async fn page_deploy(request: PageDeployRequest) -> Result { let (session, relay_url) = read_account_context().await?; @@ -141,3 +161,11 @@ pub async fn page_unpublish(request: PageSlugRequest) -> Result<(), String> { .await .map_err(|e| e.to_string()) } + +#[tauri::command] +pub async fn page_delete(request: PageSlugRequest) -> Result<(), String> { + let (session, relay_url) = read_account_context().await?; + delete_page_from_relay(&relay_url, &session.token, &request.slug) + .await + .map_err(|e| e.to_string()) +} diff --git a/src/apps/desktop/src/api/remote_connect_api.rs b/src/apps/desktop/src/api/remote_connect_api.rs index 92cbc241e8..63b40e699c 100644 --- a/src/apps/desktop/src/api/remote_connect_api.rs +++ b/src/apps/desktop/src/api/remote_connect_api.rs @@ -24,6 +24,7 @@ use regex::Regex; use serde::{Deserialize, Serialize}; use std::collections::HashMap; use std::path::PathBuf; +use std::sync::atomic::{AtomicU64, AtomicUsize, Ordering}; use std::sync::{Arc, OnceLock}; use tauri::{AppHandle, Emitter, State}; use tokio::sync::RwLock; @@ -39,6 +40,88 @@ static ACCOUNT_SESSION: OnceLock>>> = OnceLock /// and device-routing calls). static ACCOUNT_RELAY_URL: OnceLock>>> = OnceLock::new(); +/// Serializes explicit login-time syncs and lets logout/new login invalidate +/// the active operation before account state is changed. +static ACCOUNT_AUTO_SYNC_LOCK: tokio::sync::Mutex<()> = tokio::sync::Mutex::const_new(()); +static ACTIVE_ACCOUNT_AUTO_SYNC_OPERATION_ID: AtomicU64 = AtomicU64::new(0); +static ACCOUNT_CONTEXT_GENERATION: AtomicU64 = AtomicU64::new(1); +static ACCOUNT_CONTEXT_TRANSITIONS: AtomicUsize = AtomicUsize::new(0); + +fn account_context_generation() -> u64 { + ACCOUNT_CONTEXT_GENERATION.load(Ordering::Acquire) +} + +fn account_context_is_current(generation: u64) -> bool { + ACCOUNT_CONTEXT_TRANSITIONS.load(Ordering::Acquire) == 0 + && account_context_generation() == generation +} + +struct AccountContextTransitionPermit; + +impl AccountContextTransitionPermit { + fn begin() -> Self { + ACCOUNT_CONTEXT_TRANSITIONS.fetch_add(1, Ordering::AcqRel); + ACCOUNT_CONTEXT_GENERATION.fetch_add(1, Ordering::AcqRel); + ACTIVE_ACCOUNT_AUTO_SYNC_OPERATION_ID.store(0, Ordering::Release); + Self + } +} + +impl Drop for AccountContextTransitionPermit { + fn drop(&mut self) { + // Invalidate work queued during the transition before making the newly + // installed (or cleared) account context discoverable. + ACCOUNT_CONTEXT_GENERATION.fetch_add(1, Ordering::AcqRel); + ACCOUNT_CONTEXT_TRANSITIONS.fetch_sub(1, Ordering::AcqRel); + } +} + +struct AccountContextTransitionGuard { + sync_guard: Option>, + transition: Option, +} + +impl Drop for AccountContextTransitionGuard { + fn drop(&mut self) { + // Release the operation lock before reopening context discovery. A + // queued operation that wins this handoff still fails on the gate. + drop(self.sync_guard.take()); + drop(self.transition.take()); + } +} + +async fn lock_account_sync( + generation: u64, +) -> Result, String> { + let guard = ACCOUNT_AUTO_SYNC_LOCK.lock().await; + if !account_context_is_current(generation) { + return Err("account sync cancelled".to_string()); + } + Ok(guard) +} + +fn ensure_account_auto_sync_current(operation_id: u64) -> Result<(), String> { + if operation_id != 0 + && ACTIVE_ACCOUNT_AUTO_SYNC_OPERATION_ID.load(Ordering::Acquire) == operation_id + { + Ok(()) + } else { + Err("account sync cancelled".to_string()) + } +} + +async fn cancel_and_wait_for_account_auto_sync() -> AccountContextTransitionGuard { + // The permit exists before either await, so cancellation cannot leave a + // half-started transition permanently open or accidentally discoverable. + let transition = AccountContextTransitionPermit::begin(); + let sync_guard = ACCOUNT_AUTO_SYNC_LOCK.lock().await; + bitfun_core::service::remote_connect::settings_sync::wait_for_sync_operations_idle().await; + AccountContextTransitionGuard { + sync_guard: Some(sync_guard), + transition: Some(transition), + } +} + /// Global handle to the DialogScheduler, set during app startup. Used by the /// device-routing background task to execute commands received from peer /// devices (ExecuteOnDevice). @@ -99,6 +182,7 @@ fn emit_settings_applied() { /// Emit granular auto-sync progress for the account login / devices UI. fn emit_sync_progress( + operation_id: u64, phase: &str, percent: u8, current: Option, @@ -108,6 +192,7 @@ fn emit_sync_progress( emit_account_event( "account://sync-progress", serde_json::json!({ + "operation_id": operation_id, "phase": phase, "percent": percent.min(100), "current": current, @@ -1396,6 +1481,8 @@ async fn persist_account_session(device_id: Option<&str>) -> Result<(), String> /// overwrite UI must `account_logout` instead of leaving a memory-only session. #[tauri::command] pub async fn account_finalize_login() -> Result<(), String> { + let generation = account_context_generation(); + let _sync_guard = lock_account_sync(generation).await?; let device = current_device_identity()?; persist_account_session(Some(device.device_id.as_str())).await?; PENDING_SYNC_CHOICE.store(false, std::sync::atomic::Ordering::Relaxed); @@ -1418,6 +1505,9 @@ pub async fn account_finalize_login() -> Result<(), String> { #[tauri::command] pub async fn account_login(request: AccountAuthRequest) -> Result { + // A new login must never overlap a detached sync that cloned the prior + // account context. Hold the guard until the new session is installed. + let _sync_guard = cancel_and_wait_for_account_auto_sync().await; let device = current_device_identity()?; let client = AccountClient::new(); let session = client @@ -1508,6 +1598,10 @@ pub async fn account_status() -> Result { /// `revoke_relay_token` is false after deleting this device because the relay /// deletion already revoked the current token along with the device row. async fn clear_account_login(revoke_relay_token: bool) { + // Invalidate the active operation first, then wait for it to observe the + // cancellation and release its guard. This ensures no settings apply or + // progress event can happen after logout completes. + let _sync_guard = cancel_and_wait_for_account_auto_sync().await; // Disconnect device routing before clearing the session. if let Some(service) = get_service_holder().read().await.as_ref() { service.stop_device_connection().await; @@ -1568,6 +1662,8 @@ pub struct OnlineDeviceInfo { /// that logs presence updates; device messages are forwarded to the RemoteConnectService. #[tauri::command] pub async fn account_connect_devices() -> Result, String> { + let account_generation = account_context_generation(); + let sync_guard = lock_account_sync(account_generation).await?; let (session, relay_url) = read_account_context().await?; let identity = current_device_identity()?; let device_name = identity.device_name.clone(); @@ -1608,6 +1704,9 @@ pub async fn account_connect_devices() -> Result, String> Ok(result) => result, Err(e) => { let msg = format!("{e}"); + // Token invalidation re-enters the account transition path, so the + // current account-operation lease must be released first. + drop(sync_guard); drop(holder); if error_indicates_expired_token(&msg) { invalidate_local_account_session(&msg).await; @@ -1632,6 +1731,9 @@ pub async fn account_connect_devices() -> Result, String> tokio::spawn(async move { use bitfun_core::service::remote_connect::relay_client::RelayEvent; while let Some(event) = event_rx.recv().await { + if !account_context_is_current(account_generation) { + break; + } match event { RelayEvent::AuthOk { user_id, device_id } => { // Should not normally arrive — start_device_connection consumes AuthOk. @@ -1671,8 +1773,8 @@ pub async fn account_connect_devices() -> Result, String> emit_device_presence(&pairs); // Another device came online — pull cloud settings if needed. if devices.len() > 1 { - tokio::spawn(async { - pull_and_reconcile().await; + tokio::spawn(async move { + pull_and_reconcile(account_generation).await; }); } } @@ -1756,7 +1858,9 @@ pub async fn account_connect_devices() -> Result, String> session={session_id} bytes={}", session_data.len() ); - match import_session_bundle(&session_data).await { + match import_session_bundle(&session_data, account_generation) + .await + { Ok(()) => { log::info!("Session {session_id} imported from device {source_device_id}"); } @@ -1952,6 +2056,8 @@ pub async fn account_fetch_synced_sessions() -> Result, Strin /// Delete a synced session blob from the relay. #[tauri::command] pub async fn account_delete_synced_session(session_id: String) -> Result<(), String> { + let generation = account_context_generation(); + let _sync_guard = lock_account_sync(generation).await?; let (session, relay_url) = read_account_context().await?; AccountClient::new() .delete_session(&relay_url, &session, &session_id) @@ -1966,6 +2072,8 @@ pub async fn account_delete_synced_session(session_id: String) -> Result<(), Str /// Upload settings blob (encrypted client-side with the master key). #[tauri::command] pub async fn account_sync_settings(settings_json: String) -> Result<(), String> { + let generation = account_context_generation(); + let _sync_guard = lock_account_sync(generation).await?; let (session, relay_url) = read_account_context().await?; bitfun_core::service::remote_connect::settings_sync::upload_settings_payload( &session, @@ -2049,6 +2157,8 @@ pub async fn account_export_local_session( app_state: State<'_, crate::api::app_state::AppState>, path_manager: State<'_, Arc>, ) -> Result<(), String> { + let generation = account_context_generation(); + let _sync_guard = lock_account_sync(generation).await?; let (acct_session, relay_url) = read_account_context().await?; let storage_path = @@ -2109,6 +2219,8 @@ pub async fn account_export_all_sessions( app_state: State<'_, crate::api::app_state::AppState>, path_manager: State<'_, Arc>, ) -> Result { + let generation = account_context_generation(); + let _sync_guard = lock_account_sync(generation).await?; let (acct_session, relay_url) = read_account_context().await?; let storage_path = @@ -2195,6 +2307,8 @@ pub async fn account_import_remote_sessions( app_state: State<'_, crate::api::app_state::AppState>, path_manager: State<'_, Arc>, ) -> Result, String> { + let generation = account_context_generation(); + let _sync_guard = lock_account_sync(generation).await?; let (acct_session, relay_url) = read_account_context().await?; let storage_path = @@ -2266,6 +2380,7 @@ pub async fn account_fetch_session_turns( app_state: State<'_, crate::api::app_state::AppState>, path_manager: State<'_, Arc>, ) -> Result { + let generation = account_context_generation(); // Soft-skip before any disk IO so accidental callers cannot fail-closed // Peer hydrate on metadata load errors. History comes from the peer host. if crate::api::peer_host_invoke::is_peer_controller_active() { @@ -2300,7 +2415,10 @@ pub async fn account_fetch_session_turns( return Ok(false); } - // Fetch the full bundle from the relay (which includes turns). + // Fetch the full bundle from the relay (which includes turns). Keep the + // account lease through the local commit so an account switch cannot write + // a stale account's history after it completes. + let _sync_guard = lock_account_sync(generation).await?; let (acct_session, relay_url) = read_account_context().await?; let fetched = AccountClient::new() .fetch_session(&relay_url, &acct_session, &session_id) @@ -2528,23 +2646,42 @@ pub async fn account_auto_sync( is_first_login: bool, workspace_path: String, config_json: String, + sync_operation_id: u64, app_state: State<'_, crate::api::app_state::AppState>, path_manager: State<'_, Arc>, ) -> Result { - account_auto_sync_inner( + if sync_operation_id == 0 { + return Err("sync operation id must be non-zero".to_string()); + } + // Capture the account generation before queueing. A logout or replacement + // login that wins the lock invalidates this call instead of letting a stale + // request start against the newly installed account. + let generation = account_context_generation(); + let _sync_guard = lock_account_sync(generation).await?; + ACTIVE_ACCOUNT_AUTO_SYNC_OPERATION_ID.store(sync_operation_id, Ordering::Release); + let result = account_auto_sync_inner( is_first_login, workspace_path, config_json, + sync_operation_id, app_state, path_manager, ) - .await + .await; + let _ = ACTIVE_ACCOUNT_AUTO_SYNC_OPERATION_ID.compare_exchange( + sync_operation_id, + 0, + Ordering::AcqRel, + Ordering::Acquire, + ); + result } async fn account_auto_sync_inner( is_first_login: bool, workspace_path: String, config_json: String, + sync_operation_id: u64, app_state: State<'_, crate::api::app_state::AppState>, path_manager: State<'_, Arc>, ) -> Result { @@ -2558,27 +2695,37 @@ async fn account_auto_sync_inner( sessions_imported: 0, }); } + ensure_account_auto_sync_current(sync_operation_id)?; let (acct_session, relay_url) = read_account_context().await?; let client = AccountClient::new(); use bitfun_core::service::remote_connect::settings_sync; // 1. Settings sync let settings_synced = if is_first_login { - emit_sync_progress("uploading_settings", 5, None, None, None); + emit_sync_progress(sync_operation_id, "uploading_settings", 5, None, None, None); settings_sync::upload_settings_payload(&acct_session, &relay_url, &config_json) .await .map_err(|e| format!("upload settings: {e}"))?; + ensure_account_auto_sync_current(sync_operation_id)?; log::info!("First login: uploaded local settings to cloud"); - emit_sync_progress("settings_done", 15, None, None, None); + emit_sync_progress(sync_operation_id, "settings_done", 15, None, None, None); true } else { - emit_sync_progress("downloading_settings", 5, None, None, None); + emit_sync_progress( + sync_operation_id, + "downloading_settings", + 5, + None, + None, + None, + ); let cloud = client .fetch_settings_with_version(&relay_url, &acct_session) .await .map_err(|e| format!("fetch settings: {e}"))?; + ensure_account_auto_sync_current(sync_operation_id)?; if let Some(blob) = cloud { - emit_sync_progress("applying_settings", 10, None, None, None); + emit_sync_progress(sync_operation_id, "applying_settings", 10, None, None, None); // Explicit user choice — always apply, even when the cursor says // this device already has this version. Applies into the global // config service, invalidates the AI client cache, reloads, and @@ -2586,21 +2733,23 @@ async fn account_auto_sync_inner( settings_sync::apply_settings_blob(&acct_session, &blob, true) .await .map_err(|e| format!("apply cloud config: {e}"))?; + ensure_account_auto_sync_current(sync_operation_id)?; log::info!( "Applied cloud settings to local device (version={})", blob.version ); - emit_sync_progress("settings_done", 15, None, None, None); + emit_sync_progress(sync_operation_id, "settings_done", 15, None, None, None); true } else { - emit_sync_progress("settings_done", 15, None, None, None); + emit_sync_progress(sync_operation_id, "settings_done", 15, None, None, None); false } }; // 2. Session sync: upload local sessions only (backup). Do NOT import cloud // sessions into local disk — Remote peer mode reads the peer's live disk. - emit_sync_progress("listing_sessions", 18, None, None, None); + ensure_account_auto_sync_current(sync_operation_id)?; + emit_sync_progress(sync_operation_id, "listing_sessions", 18, None, None, None); let storage_path = desktop_effective_session_storage_path(&app_state, &workspace_path, None, None).await; let manager = PersistenceManager::new(path_manager.inner().clone()) @@ -2613,6 +2762,7 @@ async fn account_auto_sync_inner( let export_candidates = local_sessions.len(); emit_sync_progress( + sync_operation_id, "exporting_sessions", 20, Some(0), @@ -2623,6 +2773,7 @@ async fn account_auto_sync_inner( let mut sync_state_local = sync_state::load(&acct_session.user_id); let mut pending_uploads: Vec<(String, String, String)> = Vec::new(); for meta in local_sessions.iter() { + ensure_account_auto_sync_current(sync_operation_id)?; let turns = manager .load_session_turns(&storage_path, &meta.session_id) .await @@ -2650,7 +2801,14 @@ async fn account_auto_sync_inner( } let upload_total = pending_uploads.len(); - emit_sync_progress("exporting_sessions", 20, Some(0), Some(upload_total), None); + emit_sync_progress( + sync_operation_id, + "exporting_sessions", + 20, + Some(0), + Some(upload_total), + None, + ); let completed = std::sync::Arc::new(std::sync::atomic::AtomicUsize::new(0)); let uploaded: Vec<(String, String, i64)> = stream::iter(pending_uploads) @@ -2660,6 +2818,9 @@ async fn account_auto_sync_inner( let acct_session = acct_session.clone(); let completed = completed.clone(); async move { + if ensure_account_auto_sync_current(sync_operation_id).is_err() { + return None; + } let result = client .upload_session(&relay_url, &acct_session, &session_id, &bundle_json) .await; @@ -2669,7 +2830,11 @@ async fn account_auto_sync_inner( } else { 20 + ((75 * done) / upload_total) as u8 }; + if ensure_account_auto_sync_current(sync_operation_id).is_err() { + return None; + } emit_sync_progress( + sync_operation_id, "exporting_sessions", percent.min(95), Some(done), @@ -2690,6 +2855,8 @@ async fn account_auto_sync_inner( .collect() .await; + ensure_account_auto_sync_current(sync_operation_id)?; + let exported = uploaded.len(); let mut max_uploaded_version = sync_state_local.last_session_since; for (session_id, hash, version) in uploaded { @@ -2703,8 +2870,17 @@ async fn account_auto_sync_inner( } let _ = sync_state::save(&acct_session.user_id, &sync_state_local); + ensure_session_backup_complete(upload_total, exported)?; + log::info!("Auto-sync: settings={settings_synced} exported={exported} imported=0"); - emit_sync_progress("done", 100, Some(exported), Some(0), None); + emit_sync_progress( + sync_operation_id, + "done", + 100, + Some(exported), + Some(0), + None, + ); Ok(AutoSyncResult { settings_synced, sessions_exported: exported, @@ -2712,6 +2888,15 @@ async fn account_auto_sync_inner( }) } +fn ensure_session_backup_complete(total: usize, uploaded: usize) -> Result<(), String> { + if uploaded == total { + return Ok(()); + } + Err(format!( + "session backup incomplete: uploaded {uploaded} of {total}; retry will resume remaining sessions" + )) +} + // ── Auto-sync: debounced upload on session changes ───────────────────────── // // Settings sync (debounced push + 30s pull) is owned by the shared engine in @@ -2750,8 +2935,20 @@ fn start_settings_sync_engine() { use bitfun_core::service::remote_connect::settings_sync; let hooks = settings_sync::SettingsSyncHooks { account_context: Some(std::sync::Arc::new(|| { - Box::pin(async { read_account_context().await.map_err(anyhow::Error::msg) }) + Box::pin(async { + let generation = account_context_generation(); + if !account_context_is_current(generation) { + return Err(anyhow::anyhow!("account context is transitioning")); + } + let (account, relay_url) = + read_account_context().await.map_err(anyhow::Error::msg)?; + if !account_context_is_current(generation) { + return Err(anyhow::anyhow!("account context changed while reading")); + } + Ok((account, relay_url, generation)) + }) })), + is_account_context_current: Some(std::sync::Arc::new(account_context_is_current)), should_pause: Some(std::sync::Arc::new(|| { crate::api::peer_host_invoke::is_peer_controller_active() })), @@ -2864,6 +3061,10 @@ async fn execute_debounced_sync( log::debug!("Debounced sync skipped while peer controller mode is active"); return; } + let generation = account_context_generation(); + let Ok(_sync_guard) = lock_account_sync(generation).await else { + return; + }; // Need to be logged in let (acct_session, relay_url) = match read_account_context().await { Ok(ctx) => ctx, @@ -3100,7 +3301,14 @@ async fn execute_local_remote_command( /// Import a SessionBundle JSON into local storage. Tries all workspace session /// directories and writes to the first one found (or creates one if none exist). -async fn import_session_bundle(bundle_json: &str) -> anyhow::Result<()> { +async fn import_session_bundle(bundle_json: &str, account_generation: u64) -> anyhow::Result<()> { + let _sync_guard = lock_account_sync(account_generation) + .await + .map_err(anyhow::Error::msg)?; + // A queued event from a disconnected account must not write into a new + // account's local session view even if its encrypted payload was already + // received before the socket closed. + read_account_context().await.map_err(anyhow::Error::msg)?; let bundle: SessionBundle = serde_json::from_str(bundle_json)?; let path_manager = std::sync::Arc::new(bitfun_core::infrastructure::PathManager::new()?); @@ -3149,11 +3357,14 @@ async fn import_session_bundle(bundle_json: &str) -> anyhow::Result<()> { /// One-shot cloud settings pull, triggered when another same-account device /// comes online. The periodic pull lives in the shared settings sync engine. -async fn pull_and_reconcile() { +async fn pull_and_reconcile(account_generation: u64) { if crate::api::peer_host_invoke::is_peer_controller_active() { log::debug!("Pull: skip while peer controller mode is active"); return; } + let Ok(_sync_guard) = lock_account_sync(account_generation).await else { + return; + }; let Ok((acct_session, relay_url)) = read_account_context().await else { return; }; @@ -3170,6 +3381,13 @@ async fn pull_and_reconcile() { mod sync_state_tests { use super::*; + #[test] + fn partial_session_backup_is_not_reported_as_success() { + assert!(ensure_session_backup_complete(3, 3).is_ok()); + let error = ensure_session_backup_complete(3, 2).unwrap_err(); + assert!(error.contains("uploaded 2 of 3")); + } + #[test] fn content_hash_is_stable() { let a = sync_state::content_hash(r#"{"session_id":"x"}"#); diff --git a/src/apps/desktop/src/api/remote_workspace_policy.rs b/src/apps/desktop/src/api/remote_workspace_policy.rs index 77aeca4918..f4f0ef2c77 100644 --- a/src/apps/desktop/src/api/remote_workspace_policy.rs +++ b/src/apps/desktop/src/api/remote_workspace_policy.rs @@ -1069,6 +1069,11 @@ pub const REMOTE_WORKSPACE_COMMAND_POLICIES: &[(&str, RemoteWorkspacePolicy)] = ), ("open_remote_workspace", RemoteWorkspacePolicy::RemoteRouted), ("open_workspace", RemoteWorkspacePolicy::LegacyUnaudited), + ( + "page_create_open_link", + RemoteWorkspacePolicy::WorkspaceAgnostic, + ), + ("page_delete", RemoteWorkspacePolicy::WorkspaceAgnostic), ( "page_delete_version", RemoteWorkspacePolicy::WorkspaceAgnostic, diff --git a/src/apps/desktop/src/lib.rs b/src/apps/desktop/src/lib.rs index d91353fe36..621bcd986c 100644 --- a/src/apps/desktop/src/lib.rs +++ b/src/apps/desktop/src/lib.rs @@ -1353,10 +1353,12 @@ pub async fn run() { api::pages_api::page_save_version, api::pages_api::page_list, api::pages_api::page_list_versions, + api::pages_api::page_create_open_link, api::pages_api::page_deploy, api::pages_api::page_delete_version, api::pages_api::page_update, api::pages_api::page_unpublish, + api::pages_api::page_delete, api::peer_host_invoke::peer_host_invoke_complete, api::peer_host_invoke::peer_control_attach, api::peer_host_invoke::peer_control_detach, diff --git a/src/apps/relay-server/src/lib.rs b/src/apps/relay-server/src/lib.rs index e56575d77c..f0f52bb352 100644 --- a/src/apps/relay-server/src/lib.rs +++ b/src/apps/relay-server/src/lib.rs @@ -4,8 +4,8 @@ //! on that crate directly; this facade preserves the existing import paths. pub use bitfun_relay_service::{ - admin, db, relay, routes, AppState, DiskAssetStore, MemoryAssetStore, ResponsePayload, - RoomManager, WebAssetStore, + admin, db, page_execution, relay, routes, AppState, DiskAssetStore, MemoryAssetStore, + ResponsePayload, RoomManager, WebAssetStore, }; /// Builds the shared relay router using this compatibility host's version. diff --git a/src/apps/relay-server/tests/library_compat.rs b/src/apps/relay-server/tests/library_compat.rs index 48aa0309de..8833354036 100644 --- a/src/apps/relay-server/tests/library_compat.rs +++ b/src/apps/relay-server/tests/library_compat.rs @@ -29,6 +29,11 @@ fn legacy_library_path_exposes_supported_relay_api() { asset_store: Arc::new(MemoryAssetStore::new()), db: None, page_data: None, + page_access_manager: Arc::new(routes::pages::PageAccessManager::new()), + page_upload_manager: Arc::new(routes::pages::PageUploadManager::new()), + page_execution_guard: Arc::new( + bitfun_relay_server::page_execution::PageExecutionGuard::new(), + ), login_rate_limiter: Arc::new(routes::auth::LoginRateLimiter::new()), device_manager: relay::DeviceManager::new(), cors_allow_origins: Arc::new(Vec::new()), diff --git a/src/crates/adapters/ai-adapters/Cargo.toml b/src/crates/adapters/ai-adapters/Cargo.toml index 8f2730765e..7903e0073d 100644 --- a/src/crates/adapters/ai-adapters/Cargo.toml +++ b/src/crates/adapters/ai-adapters/Cargo.toml @@ -18,6 +18,7 @@ bitfun-core-types = { path = "../../contracts/core-types" } bitfun-services-core = { path = "../../services/services-core", default-features = false, optional = true } chrono = { workspace = true } dirs = { workspace = true, optional = true } +keyring = { workspace = true, optional = true } eventsource-stream = { workspace = true } futures = { workspace = true } log = { workspace = true } @@ -36,6 +37,7 @@ subscription-auth = [ "dep:base64", "dep:bitfun-services-core", "dep:dirs", + "dep:keyring", "dep:sha2", "dep:uuid", ] diff --git a/src/crates/adapters/ai-adapters/src/subscription_auth/antigravity.rs b/src/crates/adapters/ai-adapters/src/subscription_auth/antigravity.rs index 30558815e0..81571705f8 100644 --- a/src/crates/adapters/ai-adapters/src/subscription_auth/antigravity.rs +++ b/src/crates/adapters/ai-adapters/src/subscription_auth/antigravity.rs @@ -1,8 +1,9 @@ //! Antigravity (Google) subscription login and credential resolution. //! -//! Browser PKCE login against Google OAuth on the fixed loopback port `51121`, -//! then Bearer access to the Cloud Code Assist (`cloudcode-pa`) endpoint using -//! the `gemini-code-assist` request format. Constants mirror +//! Browser PKCE login against Google OAuth on a loopback listener (preferring +//! port `51121` and falling back to an ephemeral port), then Bearer access to +//! the Cloud Code Assist (`cloudcode-pa`) endpoint using the +//! `gemini-code-assist` request format. Constants mirror //! `opencode-antigravity-auth`. use super::store::{self, StoredCredential}; @@ -15,6 +16,7 @@ use tokio_util::sync::CancellationToken; const CLIENT_ID: &str = "1071006060591-tmhssin2h21lcre235vtolojh4g403ep.apps.googleusercontent.com"; const CALLBACK_PATH: &str = "/oauth-callback"; const CALLBACK_PORT: u16 = 51121; +const CALLBACK_PORTS: &[u16] = &[CALLBACK_PORT, 0]; const AUTHORIZE_URL: &str = "https://accounts.google.com/o/oauth2/v2/auth"; const TOKEN_URL: &str = "https://oauth2.googleapis.com/token"; const CODE_ASSIST_BASE_URL: &str = "https://cloudcode-pa.googleapis.com"; @@ -26,8 +28,8 @@ const DEFAULT_MODEL: &str = "gemini-3-pro-high"; const REFRESH_LEEWAY_MS: i64 = 5 * 60 * 1000; const STORE_KEY: &str = "antigravity"; -fn redirect_uri() -> String { - oauth_server::loopback_redirect_uri(CALLBACK_PORT, CALLBACK_PATH) +fn redirect_uri(port: u16) -> String { + oauth_server::loopback_redirect_uri(port, CALLBACK_PATH) } const SCOPES: &[&str] = &[ @@ -49,18 +51,27 @@ fn client_secret() -> String { /// Returns the platform-specific User-Agent and Client-Metadata platform token. fn platform_tokens() -> (String, &'static str) { - if cfg!(target_os = "windows") { - ("windows/amd64".to_string(), "WINDOWS") - } else if cfg!(target_os = "macos") { - let arch = if cfg!(target_arch = "aarch64") { - "darwin/arm64" - } else { - "darwin/amd64" - }; - (arch.to_string(), "MACOS") - } else { - ("linux/amd64".to_string(), "LINUX") - } + platform_tokens_for(std::env::consts::OS, std::env::consts::ARCH) +} + +fn platform_tokens_for(os: &str, arch: &str) -> (String, &'static str) { + let user_agent_os = match os { + "windows" => "windows", + "macos" => "darwin", + _ => "linux", + }; + let metadata_os = match os { + "windows" => "WINDOWS", + "macos" => "MACOS", + _ => "LINUX", + }; + let user_agent_arch = match arch { + "x86_64" => "amd64", + "aarch64" => "arm64", + "x86" => "386", + other => other, + }; + (format!("{user_agent_os}/{user_agent_arch}"), metadata_os) } #[derive(Debug, Deserialize)] @@ -181,9 +192,6 @@ fn metadata_from( } async fn persist_tokens(tokens: TokenResponse) -> Result<()> { - let _guard = super::store_lock(super::SubscriptionProvider::Antigravity) - .lock() - .await; let access = tokens .access_token .clone() @@ -194,9 +202,8 @@ async fn persist_tokens(tokens: TokenResponse) -> Result<()> { .ok_or_else(|| anyhow!("antigravity token response missing refresh_token"))?; let expires = now_ms() + tokens.expires_in.unwrap_or(3600) * 1000; let metadata = metadata_from(&tokens, None); - let mut store = store::load().await.unwrap_or_default(); - store.insert( - STORE_KEY.to_string(), + store::upsert( + STORE_KEY, StoredCredential::Oauth { refresh, access, @@ -204,8 +211,8 @@ async fn persist_tokens(tokens: TokenResponse) -> Result<()> { account_id: None, metadata, }, - ); - store::save(&store).await?; + ) + .await?; log::info!("antigravity subscription tokens saved"); Ok(()) } @@ -214,25 +221,27 @@ async fn persist_tokens(tokens: TokenResponse) -> Result<()> { pub(crate) async fn begin_login(cancel: CancellationToken) -> Result { let pkce = Pkce::generate(); let state = pkce::random_state(); - let redirect_uri = redirect_uri(); + let (listener, callback_port) = oauth_server::bind_loopback_ports(CALLBACK_PORTS).await?; + let redirect_uri = redirect_uri(callback_port); let authorization_url = build_authorize_url(&pkce, &state, &redirect_uri); - let listener = oauth_server::bind_loopback(CALLBACK_PORT).await?; let verifier = pkce.verifier.clone(); let runner = async move { - tokio::select! { - _ = cancel.cancelled() => Err(anyhow!("login cancelled")), - result = async { + super::authorize_then_persist( + super::SubscriptionProvider::Antigravity, + cancel, + async { let params = oauth_server::wait_for_callback(listener, CALLBACK_PATH, &state).await?; let code = params .get("code") .cloned() .ok_or_else(|| anyhow!("antigravity callback missing code"))?; - let tokens = exchange_code(&code, &verifier, &redirect_uri).await?; - persist_tokens(tokens).await - } => result, - } + exchange_code(&code, &verifier, &redirect_uri).await + }, + persist_tokens, + ) + .await }; Ok(StartedLogin { @@ -249,10 +258,8 @@ async fn ensure_fresh() -> Result<(String, i64)> { let _guard = super::store_lock(super::SubscriptionProvider::Antigravity) .lock() .await; - let mut store = store::load().await.unwrap_or_default(); - let entry = store - .get(STORE_KEY) - .cloned() + let entry = store::load_entry(STORE_KEY) + .await? .ok_or_else(|| anyhow!("Antigravity is not connected; sign in first"))?; let StoredCredential::Oauth { refresh: refresh_token, @@ -277,8 +284,8 @@ async fn ensure_fresh() -> Result<(String, i64)> { let new_refresh = refreshed.refresh_token.clone().unwrap_or(refresh_token); let new_expires = now_ms() + refreshed.expires_in.unwrap_or(3600) * 1000; let new_metadata = metadata_from(&refreshed, metadata); - store.insert( - STORE_KEY.to_string(), + store::upsert( + STORE_KEY, StoredCredential::Oauth { refresh: new_refresh, access: new_access.clone(), @@ -286,8 +293,8 @@ async fn ensure_fresh() -> Result<(String, i64)> { account_id, metadata: new_metadata, }, - ); - store::save(&store).await?; + ) + .await?; log::info!("antigravity subscription tokens refreshed"); Ok((new_access, new_expires)) } @@ -326,12 +333,12 @@ pub(crate) fn suggested() -> (&'static str, &'static str, &'static str) { #[cfg(test)] mod tests { - use super::{build_authorize_url, redirect_uri}; + use super::{build_authorize_url, platform_tokens_for, redirect_uri}; use crate::subscription_auth::pkce::Pkce; #[test] fn uses_registered_localhost_redirect_uri() { - let redirect_uri = redirect_uri(); + let redirect_uri = redirect_uri(super::CALLBACK_PORT); assert_eq!(redirect_uri, "http://localhost:51121/oauth-callback"); let authorize_url = build_authorize_url(&Pkce::generate(), "state", &redirect_uri); @@ -339,4 +346,28 @@ mod tests { authorize_url.contains("redirect_uri=http%3A%2F%2Flocalhost%3A51121%2Foauth-callback") ); } + + #[test] + fn reports_real_architecture_on_all_desktop_platforms() { + assert_eq!( + platform_tokens_for("windows", "x86_64"), + ("windows/amd64".to_string(), "WINDOWS") + ); + assert_eq!( + platform_tokens_for("windows", "aarch64"), + ("windows/arm64".to_string(), "WINDOWS") + ); + assert_eq!( + platform_tokens_for("macos", "aarch64"), + ("darwin/arm64".to_string(), "MACOS") + ); + assert_eq!( + platform_tokens_for("linux", "x86_64"), + ("linux/amd64".to_string(), "LINUX") + ); + assert_eq!( + platform_tokens_for("linux", "aarch64"), + ("linux/arm64".to_string(), "LINUX") + ); + } } diff --git a/src/crates/adapters/ai-adapters/src/subscription_auth/codex.rs b/src/crates/adapters/ai-adapters/src/subscription_auth/codex.rs index d92b95abfb..d2a361a492 100644 --- a/src/crates/adapters/ai-adapters/src/subscription_auth/codex.rs +++ b/src/crates/adapters/ai-adapters/src/subscription_auth/codex.rs @@ -139,9 +139,6 @@ fn metadata_from(tokens: &TokenResponse) -> Option { } async fn persist_tokens(tokens: TokenResponse) -> Result<()> { - let _guard = super::store_lock(super::SubscriptionProvider::Codex) - .lock() - .await; let access = tokens .access_token .clone() @@ -153,9 +150,8 @@ async fn persist_tokens(tokens: TokenResponse) -> Result<()> { let expires = now_ms() + tokens.expires_in.unwrap_or(3600) * 1000; let account_id = account_id_from(&tokens); let metadata = metadata_from(&tokens); - let mut store = store::load().await.unwrap_or_default(); - store.insert( - STORE_KEY.to_string(), + store::upsert( + STORE_KEY, StoredCredential::Oauth { refresh, access, @@ -163,8 +159,8 @@ async fn persist_tokens(tokens: TokenResponse) -> Result<()> { account_id, metadata, }, - ); - store::save(&store).await?; + ) + .await?; log::info!("codex subscription tokens saved"); Ok(()) } @@ -179,19 +175,21 @@ pub(crate) async fn begin_login(cancel: CancellationToken) -> Result Err(anyhow!("login cancelled")), - result = async { + super::authorize_then_persist( + super::SubscriptionProvider::Codex, + cancel, + async { let params = oauth_server::wait_for_callback(listener, CALLBACK_PATH, &state).await?; let code = params .get("code") .cloned() .ok_or_else(|| anyhow!("codex callback missing code"))?; - let tokens = exchange_code(&code, &verifier, &redirect_uri).await?; - persist_tokens(tokens).await - } => result, - } + exchange_code(&code, &verifier, &redirect_uri).await + }, + persist_tokens, + ) + .await }; Ok(StartedLogin { @@ -208,10 +206,8 @@ async fn ensure_fresh() -> Result<(String, Option, i64)> { let _guard = super::store_lock(super::SubscriptionProvider::Codex) .lock() .await; - let mut store = store::load().await.unwrap_or_default(); - let entry = store - .get(STORE_KEY) - .cloned() + let entry = store::load_entry(STORE_KEY) + .await? .ok_or_else(|| anyhow!("Codex is not connected; sign in first"))?; let StoredCredential::Oauth { refresh: refresh_token, @@ -237,8 +233,8 @@ async fn ensure_fresh() -> Result<(String, Option, i64)> { let new_expires = now_ms() + refreshed.expires_in.unwrap_or(3600) * 1000; let new_account_id = account_id_from(&refreshed).or(account_id); let new_metadata = metadata_from(&refreshed).or(metadata); - store.insert( - STORE_KEY.to_string(), + store::upsert( + STORE_KEY, StoredCredential::Oauth { refresh: new_refresh, access: new_access.clone(), @@ -246,8 +242,8 @@ async fn ensure_fresh() -> Result<(String, Option, i64)> { account_id: new_account_id.clone(), metadata: new_metadata, }, - ); - store::save(&store).await?; + ) + .await?; log::info!("codex subscription tokens refreshed"); Ok((new_access, new_account_id, new_expires)) } diff --git a/src/crates/adapters/ai-adapters/src/subscription_auth/jwt.rs b/src/crates/adapters/ai-adapters/src/subscription_auth/jwt.rs index 99606ae0a9..b994b992a4 100644 --- a/src/crates/adapters/ai-adapters/src/subscription_auth/jwt.rs +++ b/src/crates/adapters/ai-adapters/src/subscription_auth/jwt.rs @@ -50,7 +50,7 @@ pub(crate) fn chatgpt_account_id(token: &str) -> Option { #[cfg(test)] mod tests { use super::*; - use base64::{engine::general_purpose::URL_SAFE_NO_PAD, Engine as _}; + use base64::engine::general_purpose::URL_SAFE_NO_PAD; fn make_token(payload: serde_json::Value) -> String { let header = URL_SAFE_NO_PAD.encode(b"{\"alg\":\"none\"}"); diff --git a/src/crates/adapters/ai-adapters/src/subscription_auth/mod.rs b/src/crates/adapters/ai-adapters/src/subscription_auth/mod.rs index fb58603f0b..d163f2fbe0 100644 --- a/src/crates/adapters/ai-adapters/src/subscription_auth/mod.rs +++ b/src/crates/adapters/ai-adapters/src/subscription_auth/mod.rs @@ -2,8 +2,9 @@ //! //! Lets BitFun sign in to another product's subscription (Codex/ChatGPT, //! Antigravity/Google, OpenCode Zen) with an OpenCode-style in-app OAuth flow, -//! and use the resulting tokens to authenticate AI requests. Tokens are stored -//! locally in `subscription_auth.json` (mode 0600) and refreshed on resolve. +//! and use the resulting tokens to authenticate AI requests. Secret material +//! is stored in the operating-system credential vault; the local JSON file +//! contains non-secret account metadata only. //! //! There is no upgrade path for the previous Codex/Gemini CLI disk-scan import. @@ -88,6 +89,14 @@ pub struct SubscriptionAccount { /// Unix seconds when the current credential expires (for UI display). pub expires_at: Option, pub connected: bool, + /// The account was known previously, but its secret is absent from the + /// system credential vault. The UI should ask the user to sign in again. + #[serde(default)] + pub reauthentication_required: bool, + /// The system credential vault is currently locked or unavailable. Unlike + /// a missing entry, this is retryable and should not request re-login. + #[serde(default)] + pub vault_unavailable: bool, pub suggested_format: String, pub suggested_base_url: String, pub suggested_model: String, @@ -194,9 +203,44 @@ pub(crate) fn store_lock(provider: SubscriptionProvider) -> &'static tokio::sync } } +/// Runs the externally cancellable authorization/polling phase, then commits +/// the resulting credential without cancellation. Dropping a credential-vault +/// write can leave an orphan secret because blocking platform keyring calls +/// continue running after their Rust future is dropped. +pub(crate) async fn authorize_then_persist( + provider: SubscriptionProvider, + cancel: CancellationToken, + authorize: Authorize, + persist: Persist, +) -> Result<()> +where + Authorize: std::future::Future>, + Persist: FnOnce(T) -> PersistFuture, + PersistFuture: std::future::Future>, +{ + let credential = tokio::select! { + _ = cancel.cancelled() => return Err(anyhow!("login cancelled")), + result = tokio::time::timeout(LOGIN_TIMEOUT, authorize) => match result { + Ok(result) => result?, + Err(_) => return Err(anyhow!("Login timed out")), + }, + }; + // Logout/re-login cancels the generation before waiting on this same + // provider lock. Whichever side reaches the lock boundary first wins: + // an already-started commit finishes before logout deletes it, while a + // cancelled commit waiting on the lock is discarded before writing. + let _guard = store_lock(provider).lock().await; + if cancel.is_cancelled() { + return Err(anyhow!("login cancelled")); + } + persist(credential).await +} + fn build_account( provider: SubscriptionProvider, entry: Option<&StoredCredential>, + reauthentication_required: bool, + vault_unavailable: bool, ) -> SubscriptionAccount { let (format, base_url, model) = provider.suggested(); let (connected, account, expires_at) = match entry { @@ -223,6 +267,8 @@ fn build_account( account, expires_at, connected, + reauthentication_required, + vault_unavailable, suggested_format: format.to_string(), suggested_base_url: base_url.to_string(), suggested_model: model.to_string(), @@ -230,73 +276,149 @@ fn build_account( } async fn account_snapshot(provider: SubscriptionProvider) -> SubscriptionAccount { - let store = store::load().await.unwrap_or_default(); - build_account(provider, store.get(provider.key())) + let state = store::load_with_state().await.unwrap_or_else(|error| { + log::warn!("load subscription credential state failed: {error:#}"); + store::LoadState { + credentials: store::Store::new(), + requires_reauthentication: std::collections::HashSet::new(), + vault_unavailable: std::collections::HashSet::new(), + } + }); + build_account( + provider, + state.credentials.get(provider.key()), + state.requires_reauthentication.contains(provider.key()), + state.vault_unavailable.contains(provider.key()), + ) } /// Lists all providers with their current connection state. pub async fn list_accounts() -> Vec { - let store = store::load().await.unwrap_or_default(); + let state = store::load_with_state().await.unwrap_or_else(|error| { + log::warn!("load subscription credential state failed: {error:#}"); + store::LoadState { + credentials: store::Store::new(), + requires_reauthentication: std::collections::HashSet::new(), + vault_unavailable: std::collections::HashSet::new(), + } + }); SubscriptionProvider::ALL .iter() - .map(|provider| build_account(*provider, store.get(provider.key()))) + .map(|provider| { + build_account( + *provider, + state.credentials.get(provider.key()), + state.requires_reauthentication.contains(provider.key()), + state.vault_unavailable.contains(provider.key()), + ) + }) .collect() } /// Starts a login session, cancelling any existing pending session for the /// same provider. Returns immediately with the authorization URL / user code. pub async fn start_login(provider: SubscriptionProvider) -> Result { - if let Some(previous) = { + let cancel = CancellationToken::new(); + let generation = next_generation(); + let previous = { let mut map = sessions() .lock() .map_err(|_| anyhow!("subscription login session lock poisoned"))?; - map.remove(&provider) - } { + map.insert( + provider, + SessionState { + status: LoginStatus::Pending, + authorization_url: None, + user_code: None, + instructions: None, + error: None, + account: None, + cancel: cancel.clone(), + generation, + }, + ) + }; + if let Some(previous) = previous { previous.cancel.cancel(); } - let cancel = CancellationToken::new(); - let started = match provider { - SubscriptionProvider::Codex => codex::begin_login(cancel.clone()).await, - SubscriptionProvider::Antigravity => antigravity::begin_login(cancel.clone()).await, - SubscriptionProvider::Opencode => opencode::begin_login(cancel.clone()).await, - }?; + // The placeholder above makes cancellation visible even while a provider + // is still binding its callback listener or requesting a device code. + let begin = async { + match provider { + SubscriptionProvider::Codex => codex::begin_login(cancel.clone()).await, + SubscriptionProvider::Antigravity => antigravity::begin_login(cancel.clone()).await, + SubscriptionProvider::Opencode => opencode::begin_login(cancel.clone()).await, + } + }; + let started_result = tokio::select! { + _ = cancel.cancelled() => Err(anyhow!("login cancelled")), + result = begin => result, + }; + let started = match started_result { + Ok(started) if !cancel.is_cancelled() => started, + Ok(_) => return Err(anyhow!("login cancelled")), + Err(error) => { + if let Ok(mut map) = sessions().lock() { + if let Some(state) = map + .get_mut(&provider) + .filter(|state| state.generation == generation) + { + state.status = if cancel.is_cancelled() { + LoginStatus::Cancelled + } else { + LoginStatus::Failed + }; + state.error = Some(format!("{error:#}")); + } + } + return Err(error); + } + }; let authorization_url = started.authorization_url.clone(); // Desktop opener rejects relative URLs ("Not allowed to open url /..."). // Every provider must return an absolute http(s) authorization URL. if !(authorization_url.starts_with("https://") || authorization_url.starts_with("http://")) { cancel.cancel(); + if let Ok(mut map) = sessions().lock() { + if let Some(state) = map + .get_mut(&provider) + .filter(|state| state.generation == generation) + { + state.status = LoginStatus::Failed; + state.error = Some( + "Subscription login returned a non-absolute authorization URL".to_string(), + ); + } + } return Err(anyhow!( "subscription login returned a non-absolute authorization URL: {authorization_url}" )); } let user_code = started.user_code.clone(); let instructions = started.instructions.clone(); - let generation = next_generation(); - { let mut map = sessions() .lock() .map_err(|_| anyhow!("subscription login session lock poisoned"))?; - map.insert( - provider, - SessionState { - status: LoginStatus::Pending, - authorization_url: Some(authorization_url.clone()), - user_code: user_code.clone(), - instructions: Some(instructions.clone()), - error: None, - account: None, - cancel: cancel.clone(), - generation, - }, - ); + let Some(state) = map + .get_mut(&provider) + .filter(|state| state.generation == generation && !state.cancel.is_cancelled()) + else { + cancel.cancel(); + return Err(anyhow!("login cancelled")); + }; + state.authorization_url = Some(authorization_url.clone()); + state.user_code = user_code.clone(); + state.instructions = Some(instructions.clone()); } let runner = started.runner; tokio::spawn(async move { - let outcome = tokio::time::timeout(LOGIN_TIMEOUT, runner).await; + // Authorization timeout lives inside `authorize_then_persist`; once + // persistence begins it must not be dropped by a surrounding timeout. + let outcome: Result, tokio::time::error::Elapsed> = Ok(runner.await); finalize_session(provider, generation, &cancel, outcome).await; }); @@ -350,8 +472,21 @@ async fn finalize_session( } }; + update_session_if_current(provider, generation, status, error, account); +} + +fn update_session_if_current( + provider: SubscriptionProvider, + generation: u64, + status: LoginStatus, + error: Option, + account: Option, +) { if let Ok(mut map) = sessions().lock() { - if let Some(state) = map.get_mut(&provider) { + if let Some(state) = map + .get_mut(&provider) + .filter(|state| state.generation == generation) + { state.status = status; state.error = error; if account.is_some() { @@ -395,6 +530,11 @@ pub async fn cancel_login(provider: SubscriptionProvider) { state.error = Some("Login cancelled".to_string()); } } + // Act as a completion barrier for the commit phase. If cancellation wins + // the provider lock, the runner observes the cancelled token and skips its + // write. If persistence already owns the lock, let that atomic commit + // finish before reporting cancellation back to the UI. + let _guard = store_lock(provider).lock().await; } /// Removes the stored credential for a provider. @@ -407,9 +547,7 @@ pub async fn logout(provider: SubscriptionProvider) -> Result<()> { } } let _guard = store_lock(provider).lock().await; - let mut store = store::load().await.unwrap_or_default(); - store.remove(provider.key()); - store::save(&store).await?; + store::remove(provider.key()).await?; drop(_guard); log::info!("subscription provider {} logged out", provider.key()); Ok(()) @@ -486,6 +624,11 @@ mod tests { ); store::save(&store).await.unwrap(); + let metadata_file = std::fs::read_to_string(store_path_override_for_assertion()).unwrap(); + assert!(!metadata_file.contains("refresh-token")); + assert!(!metadata_file.contains("access-token")); + assert!(metadata_file.contains("user@example.com")); + let loaded = store::load().await.unwrap(); let entry = loaded.get("codex").expect("codex entry present"); match entry { @@ -506,6 +649,209 @@ mod tests { assert!(codex.connected); assert_eq!(codex.account.as_deref(), Some("user@example.com")); assert_eq!(codex.expires_at, Some(1_800_000_000)); + assert!(!codex.reauthentication_required); + } + + fn store_path_override_for_assertion() -> std::path::PathBuf { + super::store::store_path_for_test_assertion() + } + + #[tokio::test] + async fn legacy_plaintext_store_is_migrated_and_scrubbed() { + let _guard = test_lock() + .lock() + .unwrap_or_else(|poisoned| poisoned.into_inner()); + let path = temp_store_path(); + store::set_store_path_for_test(path.clone()); + let legacy = serde_json::json!({ + "opencode": { + "type": "oauth", + "refresh": "legacy-refresh-secret", + "access": "legacy-access-secret", + "expires": 1_900_000_000_000_i64, + "metadata": { "email": "legacy@example.com" } + } + }); + std::fs::write(&path, serde_json::to_vec_pretty(&legacy).unwrap()).unwrap(); + + let loaded = store::load().await.unwrap(); + assert!(loaded.contains_key("opencode")); + + let migrated = std::fs::read_to_string(&path).unwrap(); + assert!(migrated.contains("\"version\": 2")); + assert!(migrated.contains("legacy@example.com")); + assert!(!migrated.contains("legacy-refresh-secret")); + assert!(!migrated.contains("legacy-access-secret")); + } + + #[tokio::test] + async fn concurrent_provider_upserts_preserve_both_metadata_entries() { + let _guard = test_lock() + .lock() + .unwrap_or_else(|poisoned| poisoned.into_inner()); + let path = temp_store_path(); + store::set_store_path_for_test(path.clone()); + + let codex = store::upsert( + "codex", + StoredCredential::Oauth { + refresh: "codex-refresh".to_string(), + access: "codex-access".to_string(), + expires: 1_900_000_000_000, + account_id: Some("codex-account".to_string()), + metadata: None, + }, + ); + let opencode = store::upsert( + "opencode", + StoredCredential::Oauth { + refresh: "opencode-refresh".to_string(), + access: "opencode-access".to_string(), + expires: 1_900_000_000_000, + account_id: None, + metadata: Some(serde_json::json!({ "email": "zen@example.com" })), + }, + ); + let (codex_result, opencode_result) = tokio::join!(codex, opencode); + codex_result.unwrap(); + opencode_result.unwrap(); + + let loaded = store::load().await.unwrap(); + assert!(loaded.contains_key("codex")); + assert!(loaded.contains_key("opencode")); + let metadata = std::fs::read_to_string(path).unwrap(); + assert!(metadata.contains("\"codex\"")); + assert!(metadata.contains("\"opencode\"")); + assert!(!metadata.contains("codex-access")); + assert!(!metadata.contains("opencode-access")); + } + + #[tokio::test] + async fn repeated_upsert_replaces_existing_metadata_file() { + let _guard = test_lock() + .lock() + .unwrap_or_else(|poisoned| poisoned.into_inner()); + let path = temp_store_path(); + store::set_store_path_for_test(path.clone()); + + store::upsert( + "codex", + StoredCredential::Oauth { + refresh: "old-refresh".to_string(), + access: "old-access".to_string(), + expires: 1_800_000_000_000, + account_id: None, + metadata: None, + }, + ) + .await + .unwrap(); + store::upsert( + "codex", + StoredCredential::Oauth { + refresh: "new-refresh".to_string(), + access: "new-access".to_string(), + expires: 1_900_000_000_000, + account_id: Some("updated-account".to_string()), + metadata: None, + }, + ) + .await + .unwrap(); + + let loaded = store::load_entry("codex").await.unwrap().unwrap(); + match loaded { + StoredCredential::Oauth { + refresh, + access, + expires, + account_id, + .. + } => { + assert_eq!(refresh, "new-refresh"); + assert_eq!(access, "new-access"); + assert_eq!(expires, 1_900_000_000_000); + assert_eq!(account_id.as_deref(), Some("updated-account")); + } + _ => panic!("expected oauth credential"), + } + let metadata = std::fs::read_to_string(path).unwrap(); + assert!(!metadata.contains("old-access")); + assert!(!metadata.contains("new-access")); + assert!(metadata.contains("updated-account")); + } + + #[tokio::test] + async fn long_tokens_are_split_below_the_windows_vault_limit() { + let _guard = test_lock() + .lock() + .unwrap_or_else(|poisoned| poisoned.into_inner()); + let path = temp_store_path(); + store::set_store_path_for_test(path.clone()); + let refresh = "r".repeat(5_000); + let access = "a".repeat(9_000); + + store::upsert( + "codex", + StoredCredential::Oauth { + refresh: refresh.clone(), + access: access.clone(), + expires: 1_900_000_000_000, + account_id: None, + metadata: None, + }, + ) + .await + .unwrap(); + + let entries = store::test_vault_entries_for_assertion(); + assert!(entries.len() > 2, "long tokens must use multiple entries"); + assert!(entries.keys().all(|name| name != "codex")); + assert!(entries.values().all(|part| part.len() <= 2_048)); + let loaded = store::load_entry("codex").await.unwrap().unwrap(); + match loaded { + StoredCredential::Oauth { + refresh: loaded_refresh, + access: loaded_access, + .. + } => { + assert_eq!(loaded_refresh, refresh); + assert_eq!(loaded_access, access); + } + _ => panic!("expected oauth credential"), + } + let metadata = std::fs::read_to_string(path).unwrap(); + assert!(!metadata.contains(&refresh)); + assert!(!metadata.contains(&access)); + } + + #[tokio::test] + async fn unavailable_vault_is_retryable_not_missing_credential() { + let _guard = test_lock() + .lock() + .unwrap_or_else(|poisoned| poisoned.into_inner()); + store::set_store_path_for_test(temp_store_path()); + store::upsert( + "opencode", + StoredCredential::Api { + key: "sk-present-but-locked".to_string(), + metadata: None, + }, + ) + .await + .unwrap(); + + store::set_test_vault_unavailable(true); + let state = store::load_with_state().await.unwrap(); + assert!(state.credentials.get("opencode").is_none()); + assert!(!state.requires_reauthentication.contains("opencode")); + assert!(state.vault_unavailable.contains("opencode")); + let error = store::load_entry("opencode").await.unwrap_err(); + assert!(error.to_string().contains("locked or unavailable")); + store::set_test_vault_unavailable(false); + + let restored = store::load_entry("opencode").await.unwrap(); + assert!(restored.is_some()); } #[tokio::test] @@ -529,8 +875,41 @@ mod tests { assert!(loaded.get("opencode").is_none()); } + #[tokio::test] + async fn failed_logout_metadata_commit_preserves_usable_credential() { + let _guard = test_lock() + .lock() + .unwrap_or_else(|poisoned| poisoned.into_inner()); + store::set_store_path_for_test(temp_store_path()); + store::upsert( + "opencode", + StoredCredential::Api { + key: "sk-still-usable".to_string(), + metadata: None, + }, + ) + .await + .unwrap(); + let entries_before = store::test_vault_entries_for_assertion(); + + store::set_test_metadata_write_failure(true); + let error = store::remove("opencode").await.unwrap_err(); + assert!(error.to_string().contains("injected")); + store::set_test_metadata_write_failure(false); + + assert_eq!(store::test_vault_entries_for_assertion(), entries_before); + let loaded = store::load_entry("opencode").await.unwrap().unwrap(); + match loaded { + StoredCredential::Api { key, .. } => assert_eq!(key, "sk-still-usable"), + _ => panic!("expected api credential"), + } + } + #[tokio::test] async fn finalize_ignores_superseded_session() { + let _guard = test_lock() + .lock() + .unwrap_or_else(|poisoned| poisoned.into_inner()); let provider = SubscriptionProvider::Codex; let stale_generation = next_generation(); { @@ -568,4 +947,163 @@ mod tests { }; assert_eq!(status, Some(LoginStatus::Pending)); } + + #[tokio::test] + async fn cancellation_does_not_drop_started_credential_persistence() { + let cancel = CancellationToken::new(); + let (persist_started_tx, persist_started_rx) = tokio::sync::oneshot::channel(); + let (allow_persist_tx, allow_persist_rx) = tokio::sync::oneshot::channel(); + + let task = tokio::spawn(authorize_then_persist( + SubscriptionProvider::Codex, + cancel.clone(), + async { Ok::<_, anyhow::Error>("authorized-token") }, + move |token| async move { + assert_eq!(token, "authorized-token"); + persist_started_tx.send(()).unwrap(); + allow_persist_rx.await.unwrap(); + Ok(()) + }, + )); + + persist_started_rx.await.unwrap(); + cancel.cancel(); + tokio::task::yield_now().await; + assert!(!task.is_finished()); + + allow_persist_tx.send(()).unwrap(); + assert!(task.await.unwrap().is_ok()); + } + + #[tokio::test] + async fn cancellation_before_authorization_skips_persistence() { + let cancel = CancellationToken::new(); + cancel.cancel(); + let persisted = std::sync::Arc::new(std::sync::atomic::AtomicBool::new(false)); + let persisted_for_task = persisted.clone(); + + let result = authorize_then_persist( + SubscriptionProvider::Opencode, + cancel, + std::future::pending::>(), + move |_| async move { + persisted_for_task.store(true, Ordering::SeqCst); + Ok(()) + }, + ) + .await; + + assert!(result.is_err()); + assert!(!persisted.load(Ordering::SeqCst)); + } + + #[tokio::test] + async fn cancelled_commit_waiting_for_provider_lock_does_not_persist() { + let provider = SubscriptionProvider::Antigravity; + let store_guard = store_lock(provider).lock().await; + let cancel = CancellationToken::new(); + let (authorized_tx, authorized_rx) = tokio::sync::oneshot::channel(); + let persisted = std::sync::Arc::new(std::sync::atomic::AtomicBool::new(false)); + let persisted_for_task = persisted.clone(); + + let task = tokio::spawn(authorize_then_persist( + provider, + cancel.clone(), + async move { + authorized_tx.send(()).unwrap(); + Ok::<_, anyhow::Error>("authorized-token") + }, + move |_| async move { + persisted_for_task.store(true, Ordering::SeqCst); + Ok(()) + }, + )); + + authorized_rx.await.unwrap(); + cancel.cancel(); + drop(store_guard); + + assert!(task.await.unwrap().is_err()); + assert!(!persisted.load(Ordering::SeqCst)); + } + + #[tokio::test] + async fn cancel_command_waits_for_the_commit_boundary() { + let provider = SubscriptionProvider::Opencode; + let store_guard = store_lock(provider).lock().await; + let cancel = CancellationToken::new(); + let generation = next_generation(); + { + let mut map = sessions().lock().unwrap(); + map.insert( + provider, + SessionState { + status: LoginStatus::Pending, + authorization_url: None, + user_code: None, + instructions: None, + error: None, + account: None, + cancel: cancel.clone(), + generation, + }, + ); + } + + let task = tokio::spawn(cancel_login(provider)); + tokio::task::yield_now().await; + assert!(cancel.is_cancelled()); + assert!(!task.is_finished()); + + drop(store_guard); + task.await.unwrap(); + let status = sessions() + .lock() + .unwrap() + .remove(&provider) + .map(|state| state.status); + assert_eq!(status, Some(LoginStatus::Cancelled)); + } + + #[test] + fn final_state_update_rechecks_generation_after_async_work() { + let _guard = test_lock() + .lock() + .unwrap_or_else(|poisoned| poisoned.into_inner()); + let provider = SubscriptionProvider::Codex; + let old_generation = next_generation(); + let new_generation = next_generation(); + { + let mut map = sessions().lock().unwrap(); + map.insert( + provider, + SessionState { + status: LoginStatus::Pending, + authorization_url: None, + user_code: None, + instructions: None, + error: None, + account: None, + cancel: CancellationToken::new(), + generation: new_generation, + }, + ); + } + + update_session_if_current( + provider, + old_generation, + LoginStatus::Authorized, + None, + None, + ); + + let status = { + let mut map = sessions().lock().unwrap(); + let status = map.get(&provider).map(|state| state.status); + map.remove(&provider); + status + }; + assert_eq!(status, Some(LoginStatus::Pending)); + } } diff --git a/src/crates/adapters/ai-adapters/src/subscription_auth/oauth_callback_locales.json b/src/crates/adapters/ai-adapters/src/subscription_auth/oauth_callback_locales.json new file mode 100644 index 0000000000..0432668850 --- /dev/null +++ b/src/crates/adapters/ai-adapters/src/subscription_auth/oauth_callback_locales.json @@ -0,0 +1,26 @@ +{ + "en": { + "success_title": "Sign-in complete", + "success_message": "You are now signed in. You can close this window and return to BitFun.", + "error_title": "Sign-in failed", + "bad_request": "Bad request", + "missing_code": "Missing authorization code", + "invalid_state": "Invalid authorization state" + }, + "zh-CN": { + "success_title": "登录完成", + "success_message": "登录已完成。你可以关闭此窗口并返回 BitFun。", + "error_title": "登录失败", + "bad_request": "请求无效", + "missing_code": "缺少授权码", + "invalid_state": "授权状态无效" + }, + "zh-TW": { + "success_title": "登入完成", + "success_message": "登入已完成。你可以關閉此視窗並返回 BitFun。", + "error_title": "登入失敗", + "bad_request": "請求無效", + "missing_code": "缺少授權碼", + "invalid_state": "授權狀態無效" + } +} diff --git a/src/crates/adapters/ai-adapters/src/subscription_auth/oauth_server.rs b/src/crates/adapters/ai-adapters/src/subscription_auth/oauth_server.rs index 71b5ba2de2..420dea56ed 100644 --- a/src/crates/adapters/ai-adapters/src/subscription_auth/oauth_server.rs +++ b/src/crates/adapters/ai-adapters/src/subscription_auth/oauth_server.rs @@ -6,7 +6,9 @@ //! query parameters. use anyhow::{anyhow, Context, Result}; +use serde::Deserialize; use std::collections::HashMap; +use std::sync::OnceLock; use tokio::io::{AsyncReadExt, AsyncWriteExt}; use tokio::net::TcpListener; @@ -30,18 +32,8 @@ pub(crate) fn loopback_redirect_uri(port: u16, path: &str) -> String { format!("http://{LOOPBACK_REDIRECT_HOST}:{port}{path}") } -/// Binds the OAuth callback listener on [`LOOPBACK_BIND_HOST`]. -pub(crate) async fn bind_loopback(port: u16) -> Result { - TcpListener::bind((LOOPBACK_BIND_HOST, port)) - .await - .with_context(|| { - format!( - "bind OAuth callback on {LOOPBACK_BIND_HOST}:{port} (is another app using this port?)" - ) - }) -} - -/// Binds the first available registered callback port. +/// Binds the first available provider-supported callback port. A final `0` +/// entry requests an ephemeral port for desktop OAuth providers that permit it. /// /// Fallback is attempted only when a preferred port is already in use. The /// returned port must be used to construct both the authorize and token @@ -60,7 +52,7 @@ pub(crate) async fn bind_loopback_ports(ports: &[u16]) -> Result<(TcpListener, u .port(); if index > 0 { log::warn!( - "OAuth callback port {preferred_port} is unavailable; using registered fallback port {actual_port}" + "OAuth callback port {preferred_port} is unavailable; using fallback port {actual_port}" ); } return Ok((listener, actual_port)); @@ -100,8 +92,14 @@ pub(crate) async fn wait_for_callback( } }; let request = String::from_utf8_lossy(&buf[..n]); + let locale = preferred_locale(&request); let Some(request_line) = request.lines().next() else { - write_response(&mut stream, 400, &error_page("Bad request")).await; + write_response( + &mut stream, + 400, + &error_page(&callback_messages(locale).bad_request, locale), + ) + .await; continue; }; let target = request_line @@ -119,26 +117,39 @@ pub(crate) async fn wait_for_callback( } let params = parse_query(query); + // Ignore unsolicited loopback requests instead of letting a local + // process/browser probe terminate the real OAuth session. Validate + // state before accepting provider errors for the same reason. + match params.get("state") { + Some(state) if state == expected_state => {} + _ => { + write_response( + &mut stream, + 400, + &error_page(&callback_messages(locale).invalid_state, locale), + ) + .await; + continue; + } + } if let Some(error) = params.get("error") { let message = params .get("error_description") .cloned() .unwrap_or_else(|| error.clone()); - write_response(&mut stream, 200, &error_page(&message)).await; + write_response(&mut stream, 200, &error_page(&message, locale)).await; return Err(anyhow!("authorization failed: {message}")); } if params.get("code").map(String::is_empty).unwrap_or(true) { - write_response(&mut stream, 400, &error_page("Missing authorization code")).await; + write_response( + &mut stream, + 400, + &error_page(&callback_messages(locale).missing_code, locale), + ) + .await; return Err(anyhow!("authorization callback missing code")); } - match params.get("state") { - Some(state) if state == expected_state => {} - _ => { - write_response(&mut stream, 400, &error_page("Invalid state")).await; - return Err(anyhow!("authorization state mismatch")); - } - } - write_response(&mut stream, 200, &success_page()).await; + write_response(&mut stream, 200, &success_page(locale)).await; return Ok(params); } } @@ -181,25 +192,75 @@ async fn write_response(stream: &mut tokio::net::TcpStream, status: u16, body: & let _ = stream.flush().await; } -fn success_page() -> String { - result_page( - "Sign-in complete", - "You are now signed in. You can close this window and return to BitFun.", - ) +#[derive(Debug, Deserialize)] +struct CallbackMessages { + success_title: String, + success_message: String, + error_title: String, + bad_request: String, + missing_code: String, + invalid_state: String, +} + +fn callback_locales() -> &'static HashMap { + static LOCALES: OnceLock> = OnceLock::new(); + LOCALES.get_or_init(|| { + serde_json::from_str(include_str!("oauth_callback_locales.json")) + .expect("embedded OAuth callback locales are valid JSON") + }) +} + +fn callback_messages(locale: &str) -> &'static CallbackMessages { + callback_locales() + .get(locale) + .or_else(|| callback_locales().get("en")) + .expect("OAuth callback English locale is embedded") +} + +fn preferred_locale(request: &str) -> &'static str { + for line in request.lines() { + let Some((name, value)) = line.split_once(':') else { + continue; + }; + if !name.eq_ignore_ascii_case("accept-language") { + continue; + } + let language = value + .split(',') + .next() + .unwrap_or_default() + .trim() + .to_ascii_lowercase(); + if language.starts_with("zh-tw") || language.starts_with("zh-hk") { + return "zh-TW"; + } + if language.starts_with("zh") { + return "zh-CN"; + } + break; + } + "en" +} + +fn success_page(locale: &str) -> String { + let messages = callback_messages(locale); + result_page(locale, &messages.success_title, &messages.success_message) } -fn error_page(message: &str) -> String { - result_page("Sign-in failed", message) +fn error_page(message: &str, locale: &str) -> String { + result_page(locale, &callback_messages(locale).error_title, message) } -fn result_page(title: &str, message: &str) -> String { +fn result_page(language: &str, title: &str, message: &str) -> String { let message = escape_html(message); format!( - "{title}\ -

{title}

{message}

" ) } @@ -217,9 +278,11 @@ fn escape_html(text: &str) -> String { #[cfg(test)] mod tests { use super::{ - bind_loopback_ports, escape_html, loopback_redirect_uri, LOOPBACK_BIND_HOST, + bind_loopback_ports, callback_messages, escape_html, loopback_redirect_uri, + preferred_locale, success_page, wait_for_callback, LOOPBACK_BIND_HOST, LOOPBACK_REDIRECT_HOST, }; + use tokio::io::{AsyncReadExt, AsyncWriteExt}; #[test] fn escapes_html_injection() { @@ -243,6 +306,67 @@ mod tests { assert_eq!(LOOPBACK_REDIRECT_HOST, "localhost"); } + #[tokio::test] + async fn invalid_state_does_not_terminate_the_real_callback_session() { + let listener = tokio::net::TcpListener::bind((LOOPBACK_BIND_HOST, 0)) + .await + .unwrap(); + let address = listener.local_addr().unwrap(); + let waiter = tokio::spawn(async move { + wait_for_callback(listener, "/auth/callback", "expected-state").await + }); + + let mut invalid = tokio::net::TcpStream::connect(address).await.unwrap(); + invalid + .write_all( + b"GET /auth/callback?error=denied&state=attacker-state HTTP/1.1\r\nHost: localhost\r\n\r\n", + ) + .await + .unwrap(); + let mut invalid_response = Vec::new(); + invalid.read_to_end(&mut invalid_response).await.unwrap(); + assert!(String::from_utf8_lossy(&invalid_response).contains("400 Bad Request")); + assert!(!waiter.is_finished()); + + let mut valid = tokio::net::TcpStream::connect(address).await.unwrap(); + valid + .write_all( + b"GET /auth/callback?code=real-code&state=expected-state HTTP/1.1\r\nHost: localhost\r\n\r\n", + ) + .await + .unwrap(); + let mut valid_response = Vec::new(); + valid.read_to_end(&mut valid_response).await.unwrap(); + assert!(String::from_utf8_lossy(&valid_response).contains("200 OK")); + + let params = waiter.await.unwrap().unwrap(); + assert_eq!(params.get("code").map(String::as_str), Some("real-code")); + } + + #[test] + fn callback_page_uses_browser_language_and_color_scheme() { + assert_eq!( + preferred_locale("GET / HTTP/1.1\r\nAccept-Language: zh-CN,zh;q=0.9\r\n"), + "zh-CN" + ); + assert_eq!( + preferred_locale("GET / HTTP/1.1\r\nAccept-Language: zh-TW,zh;q=0.9\r\n"), + "zh-TW" + ); + assert_eq!( + preferred_locale("GET / HTTP/1.1\r\nAccept-Language: en-US,en;q=0.9\r\n"), + "en" + ); + assert_ne!( + callback_messages("zh-CN").success_title, + callback_messages("en").success_title + ); + let chinese = success_page("zh-CN"); + assert!(chinese.contains("lang=\"zh-CN\"")); + assert!(chinese.contains(&callback_messages("zh-CN").success_title)); + assert!(chinese.contains("prefers-color-scheme:dark")); + } + #[tokio::test] async fn falls_back_when_preferred_callback_port_is_occupied() { let occupied = tokio::net::TcpListener::bind((LOOPBACK_BIND_HOST, 0)) diff --git a/src/crates/adapters/ai-adapters/src/subscription_auth/opencode.rs b/src/crates/adapters/ai-adapters/src/subscription_auth/opencode.rs index 2370e792de..328789e592 100644 --- a/src/crates/adapters/ai-adapters/src/subscription_auth/opencode.rs +++ b/src/crates/adapters/ai-adapters/src/subscription_auth/opencode.rs @@ -178,14 +178,10 @@ async fn fetch_metadata(access: &str) -> serde_json::Value { } async fn persist_tokens(tokens: TokenResponse) -> Result<()> { - let _guard = super::store_lock(super::SubscriptionProvider::Opencode) - .lock() - .await; let expires = now_ms() + tokens.expires_in * 1000; let metadata = fetch_metadata(&tokens.access_token).await; - let mut store = store::load().await.unwrap_or_default(); - store.insert( - STORE_KEY.to_string(), + store::upsert( + STORE_KEY, StoredCredential::Oauth { refresh: tokens.refresh_token, access: tokens.access_token, @@ -193,8 +189,8 @@ async fn persist_tokens(tokens: TokenResponse) -> Result<()> { account_id: None, metadata: Some(metadata), }, - ); - store::save(&store).await?; + ) + .await?; log::info!("opencode subscription tokens saved"); Ok(()) } @@ -246,14 +242,15 @@ pub(crate) async fn begin_login(cancel: CancellationToken) -> Result Err(anyhow!("login cancelled")), - result = async { + super::authorize_then_persist( + super::SubscriptionProvider::Opencode, + cancel, + async { let mut wait = interval; loop { tokio::time::sleep(Duration::from_secs(wait)).await; match poll_once(&device_code).await? { - DevicePoll::Authorized(tokens) => return persist_tokens(tokens).await, + DevicePoll::Authorized(tokens) => return Ok(tokens), DevicePoll::Pending => { wait = interval; } @@ -264,8 +261,10 @@ pub(crate) async fn begin_login(cancel: CancellationToken) -> Result result, - } + }, + persist_tokens, + ) + .await }; Ok(StartedLogin { @@ -281,10 +280,8 @@ async fn ensure_fresh() -> Result { let _guard = super::store_lock(super::SubscriptionProvider::Opencode) .lock() .await; - let mut store = store::load().await.unwrap_or_default(); - let entry = store - .get(STORE_KEY) - .cloned() + let entry = store::load_entry(STORE_KEY) + .await? .ok_or_else(|| anyhow!("OpenCode Zen is not connected; sign in first"))?; match entry { StoredCredential::Api { key, .. } => Ok(key), @@ -300,8 +297,8 @@ async fn ensure_fresh() -> Result { } let refreshed = refresh(&refresh_token).await?; let new_expires = now_ms() + refreshed.expires_in * 1000; - store.insert( - STORE_KEY.to_string(), + store::upsert( + STORE_KEY, StoredCredential::Oauth { refresh: refreshed.refresh_token, access: refreshed.access_token.clone(), @@ -309,8 +306,8 @@ async fn ensure_fresh() -> Result { account_id, metadata, }, - ); - store::save(&store).await?; + ) + .await?; log::info!("opencode subscription tokens refreshed"); Ok(refreshed.access_token) } diff --git a/src/crates/adapters/ai-adapters/src/subscription_auth/store.rs b/src/crates/adapters/ai-adapters/src/subscription_auth/store.rs index 089903086e..22f18eccf9 100644 --- a/src/crates/adapters/ai-adapters/src/subscription_auth/store.rs +++ b/src/crates/adapters/ai-adapters/src/subscription_auth/store.rs @@ -1,22 +1,27 @@ -//! On-disk persistence for subscription auth credentials. +//! Persistence for subscription-account credentials. //! -//! Tokens live in a single JSON document keyed by provider id -//! (`codex` / `antigravity` / `opencode`). The file is written with mode -//! `0600` on Unix so other local users cannot read the stored tokens. +//! OAuth tokens and API keys are stored in the operating-system credential +//! vault (macOS Keychain, Windows Credential Manager, or Linux Secret +//! Service). The JSON file contains only non-secret account metadata and +//! references used to discover the corresponding vault entries. //! //! Path: `{dirs::config_dir()}/bitfun/data/subscription_auth.json`. use anyhow::{anyhow, Context, Result}; use serde::{Deserialize, Serialize}; -use std::collections::HashMap; -use std::path::PathBuf; -use std::sync::{OnceLock, RwLock}; - -/// A single stored credential for one provider. -/// -/// `oauth` credentials keep the refresh/access token pair plus the millisecond -/// epoch expiry; `api` credentials keep a static API key. Both may carry an -/// opaque `metadata` object (email, org info, etc.). +use std::collections::{HashMap, HashSet}; +use std::path::{Path, PathBuf}; +use std::sync::{Mutex, OnceLock, RwLock}; + +const STORE_VERSION: u8 = 2; +const KEYRING_SERVICE: &str = "openbitfun.bitfun.subscription-auth.v1"; +// Windows Credential Manager limits a generic credential blob to 2560 bytes. +// Leave headroom for platform-store implementations and split every logical +// secret so a long JWT or refresh token remains portable across all hosts. +const SECRET_CHUNK_BYTES: usize = 2_048; + +/// A single credential assembled in memory after its secret material has been +/// read from the platform credential vault. #[derive(Debug, Clone, Serialize, Deserialize)] #[serde(tag = "type", rename_all = "snake_case")] pub enum StoredCredential { @@ -37,27 +42,340 @@ pub enum StoredCredential { }, } -/// Provider id -> stored credential. +/// Provider id -> in-memory credential. pub type Store = HashMap; +#[derive(Debug, Clone, Serialize, Deserialize)] +#[serde(tag = "type", rename_all = "snake_case")] +enum CredentialMetadata { + Oauth { + expires: i64, + #[serde(default, skip_serializing_if = "Option::is_none")] + account_id: Option, + #[serde(default, skip_serializing_if = "Option::is_none")] + metadata: Option, + #[serde(default, skip_serializing_if = "std::ops::Not::not")] + needs_reauthentication: bool, + /// Unique namespace for this committed set of vault chunks. `None` + /// denotes the legacy single-password entry keyed by provider. + #[serde(default, skip_serializing_if = "Option::is_none")] + secret_set_id: Option, + #[serde(default, skip_serializing_if = "is_zero")] + refresh_parts: u32, + #[serde(default, skip_serializing_if = "is_zero")] + access_parts: u32, + }, + Api { + #[serde(default, skip_serializing_if = "Option::is_none")] + metadata: Option, + #[serde(default, skip_serializing_if = "std::ops::Not::not")] + needs_reauthentication: bool, + #[serde(default, skip_serializing_if = "Option::is_none")] + secret_set_id: Option, + #[serde(default, skip_serializing_if = "is_zero")] + key_parts: u32, + }, +} + +fn is_zero(value: &u32) -> bool { + *value == 0 +} + +fn secret_part_count(secret: &str) -> u32 { + // Even an empty value has one explicit part, distinguishing a committed + // empty refresh token from a missing vault entry. + secret.len().max(1).div_ceil(SECRET_CHUNK_BYTES) as u32 +} + +impl CredentialMetadata { + fn from_credential(credential: &StoredCredential) -> Self { + match credential { + StoredCredential::Oauth { + refresh, + access, + expires, + account_id, + metadata, + .. + } => Self::Oauth { + expires: *expires, + account_id: account_id.clone(), + metadata: metadata.clone(), + needs_reauthentication: false, + secret_set_id: Some(uuid::Uuid::new_v4().simple().to_string()), + refresh_parts: secret_part_count(refresh), + access_parts: secret_part_count(access), + }, + StoredCredential::Api { key, metadata } => Self::Api { + metadata: metadata.clone(), + needs_reauthentication: false, + secret_set_id: Some(uuid::Uuid::new_v4().simple().to_string()), + key_parts: secret_part_count(key), + }, + } + } + + fn requiring_reauthentication(credential: &StoredCredential) -> Self { + let mut metadata = Self::from_credential(credential); + match &mut metadata { + Self::Oauth { + needs_reauthentication, + secret_set_id, + refresh_parts, + access_parts, + .. + } => { + *needs_reauthentication = true; + *secret_set_id = None; + *refresh_parts = 0; + *access_parts = 0; + } + Self::Api { + needs_reauthentication, + secret_set_id, + key_parts, + .. + } => { + *needs_reauthentication = true; + *secret_set_id = None; + *key_parts = 0; + } + } + metadata + } + + fn needs_reauthentication(&self) -> bool { + match self { + Self::Oauth { + needs_reauthentication, + .. + } + | Self::Api { + needs_reauthentication, + .. + } => *needs_reauthentication, + } + } + + fn combine(&self, secret: SecretMaterial) -> Option { + match (self, secret) { + ( + Self::Oauth { + expires, + account_id, + metadata, + .. + }, + SecretMaterial::Oauth { refresh, access }, + ) => Some(StoredCredential::Oauth { + refresh, + access, + expires: *expires, + account_id: account_id.clone(), + metadata: metadata.clone(), + }), + (Self::Api { metadata, .. }, SecretMaterial::Api { key }) => { + Some(StoredCredential::Api { + key, + metadata: metadata.clone(), + }) + } + _ => None, + } + } + + fn vault_entries(&self, provider: &str) -> Vec { + match self { + Self::Oauth { + secret_set_id: Some(set_id), + refresh_parts, + access_parts, + .. + } => secret_entry_names(provider, set_id, "refresh", *refresh_parts) + .chain(secret_entry_names( + provider, + set_id, + "access", + *access_parts, + )) + .collect(), + Self::Api { + secret_set_id: Some(set_id), + key_parts, + .. + } => secret_entry_names(provider, set_id, "api-key", *key_parts).collect(), + // The old representation used one password entry named exactly + // after the provider. + _ => vec![provider.to_string()], + } + } +} + +fn secret_entry_names<'a>( + provider: &'a str, + set_id: &'a str, + field: &'a str, + parts: u32, +) -> impl Iterator + 'a { + (0..parts).map(move |index| format!("{provider}/v2/{set_id}/{field}/{index}")) +} + +#[derive(Debug, Clone, Serialize, Deserialize)] +#[serde(tag = "type", rename_all = "snake_case")] +enum SecretMaterial { + Oauth { refresh: String, access: String }, + Api { key: String }, +} + +impl From<&StoredCredential> for SecretMaterial { + fn from(value: &StoredCredential) -> Self { + match value { + StoredCredential::Oauth { + refresh, access, .. + } => Self::Oauth { + refresh: refresh.clone(), + access: access.clone(), + }, + StoredCredential::Api { key, .. } => Self::Api { key: key.clone() }, + } + } +} + +#[derive(Debug, Default, Serialize, Deserialize)] +struct SecureStoreFile { + version: u8, + #[serde(default)] + accounts: HashMap, +} + +/// Result used by account discovery so a missing/locked vault entry is visible +/// to the UI instead of silently looking like a never-configured account. +pub(crate) struct LoadState { + pub credentials: Store, + pub requires_reauthentication: HashSet, + /// Metadata and secret entries exist, but the OS vault is currently + /// locked/unavailable. This is retryable and must not be shown as lost. + pub vault_unavailable: HashSet, +} + fn store_path_override() -> &'static RwLock> { static OVERRIDE: OnceLock>> = OnceLock::new(); OVERRIDE.get_or_init(|| RwLock::new(None)) } -/// Overrides the store path for tests. Pass a temp-dir file path. +/// Test-only secret material, keyed by the overridden metadata path. Tests +/// must never read from or write to a developer's real system credential vault. +fn test_secrets() -> &'static Mutex>>> { + static SECRETS: OnceLock>>>> = OnceLock::new(); + SECRETS.get_or_init(|| Mutex::new(HashMap::new())) +} + +#[cfg(test)] +fn unavailable_test_vaults() -> &'static Mutex> { + static PATHS: OnceLock>> = OnceLock::new(); + PATHS.get_or_init(|| Mutex::new(HashSet::new())) +} + +#[cfg(test)] +fn failing_metadata_writes() -> &'static Mutex> { + static PATHS: OnceLock>> = OnceLock::new(); + PATHS.get_or_init(|| Mutex::new(HashSet::new())) +} + +fn native_keyring_lock() -> &'static Mutex<()> { + static LOCK: OnceLock> = OnceLock::new(); + LOCK.get_or_init(|| Mutex::new(())) +} + +fn store_operation_lock() -> &'static tokio::sync::Mutex<()> { + static LOCK: tokio::sync::Mutex<()> = tokio::sync::Mutex::const_new(()); + &LOCK +} + +/// Overrides the metadata path for tests. The override also switches secret +/// persistence to the process-local test vault above. pub fn set_store_path_for_test(path: PathBuf) { if let Ok(mut guard) = store_path_override().write() { *guard = Some(path); } } -fn store_path() -> Result { - if let Ok(guard) = store_path_override().read() { - if let Some(path) = guard.as_ref() { - return Ok(path.clone()); +#[cfg(test)] +pub(crate) fn store_path_for_test_assertion() -> PathBuf { + overridden_store_path().expect("subscription test store path is configured") +} + +#[cfg(test)] +pub(crate) fn test_vault_entries_for_assertion() -> HashMap> { + let path = store_path_for_test_assertion(); + test_secrets() + .lock() + .ok() + .and_then(|vault| vault.get(&path).cloned()) + .unwrap_or_default() +} + +#[cfg(test)] +pub(crate) fn set_test_vault_unavailable(unavailable: bool) { + let path = store_path_for_test_assertion(); + if let Ok(mut paths) = unavailable_test_vaults().lock() { + if unavailable { + paths.insert(path); + } else { + paths.remove(&path); } } +} + +#[cfg(test)] +pub(crate) fn set_test_metadata_write_failure(fail: bool) { + let path = store_path_for_test_assertion(); + if let Ok(mut paths) = failing_metadata_writes().lock() { + if fail { + paths.insert(path); + } else { + paths.remove(&path); + } + } +} + +#[cfg(test)] +fn test_vault_is_unavailable(path: &Path) -> bool { + unavailable_test_vaults() + .lock() + .map(|paths| paths.contains(path)) + .unwrap_or(true) +} + +#[cfg(not(test))] +fn test_vault_is_unavailable(_path: &Path) -> bool { + false +} + +#[cfg(test)] +fn metadata_write_should_fail(path: &Path) -> bool { + failing_metadata_writes() + .lock() + .map(|paths| paths.contains(path)) + .unwrap_or(true) +} + +#[cfg(not(test))] +fn metadata_write_should_fail(_path: &Path) -> bool { + false +} + +fn overridden_store_path() -> Option { + store_path_override() + .read() + .ok() + .and_then(|guard| guard.clone()) +} + +fn store_path() -> Result { + if let Some(path) = overridden_store_path() { + return Ok(path); + } let base = dirs::config_dir().ok_or_else(|| anyhow!("system config directory unavailable"))?; Ok(base .join("bitfun") @@ -65,39 +383,547 @@ fn store_path() -> Result { .join("subscription_auth.json")) } -/// Loads the full credential store. Returns an empty store when the file does -/// not exist yet. -pub async fn load() -> Result { - let path = store_path()?; +async fn read_bytes(path: &Path) -> Result>> { + #[cfg(windows)] + restore_windows_backup_if_needed(path).await?; if !path.exists() { - return Ok(Store::new()); + return Ok(None); } - let bytes = tokio::fs::read(&path) + let bytes = tokio::fs::read(path) .await - .with_context(|| format!("read subscription auth store at {}", path.display()))?; - if bytes.is_empty() { - return Ok(Store::new()); + .with_context(|| format!("read subscription auth metadata at {}", path.display()))?; + Ok((!bytes.is_empty()).then_some(bytes)) +} + +/// Recover the old metadata index if the process stopped after rotating the +/// destination but before moving the new temp file into place. +#[cfg(windows)] +async fn restore_windows_backup_if_needed(path: &Path) -> Result<()> { + if path.exists() { + return Ok(()); + } + let backup = path.with_extension("bak"); + if !backup.exists() { + return Ok(()); + } + tokio::fs::rename(&backup, path).await.with_context(|| { + format!( + "restore interrupted subscription auth metadata {} -> {}", + backup.display(), + path.display() + ) + }) +} + +fn parse_secure_file(bytes: &[u8], path: &Path) -> Result { + let file: SecureStoreFile = serde_json::from_slice(bytes) + .with_context(|| format!("parse subscription auth metadata at {}", path.display()))?; + if file.version != STORE_VERSION { + return Err(anyhow!( + "unsupported subscription auth metadata version {} at {}", + file.version, + path.display() + )); } - serde_json::from_slice(&bytes) - .with_context(|| format!("parse subscription auth store at {}", path.display())) + Ok(file) } -/// Persists the full credential store with restrictive permissions. -pub async fn save(store: &Store) -> Result<()> { - let path = store_path()?; +async fn read_secure_file(path: &Path) -> Result { + let Some(bytes) = read_bytes(path).await? else { + return Ok(SecureStoreFile { + version: STORE_VERSION, + accounts: HashMap::new(), + }); + }; + parse_secure_file(&bytes, path) +} + +async fn get_secret_bytes(entry_name: &str) -> Result>> { + if let Some(path) = overridden_store_path() { + if test_vault_is_unavailable(&path) { + return Err(anyhow!("subscription test vault unavailable")); + } + return test_secrets() + .lock() + .map_err(|_| anyhow!("subscription test vault lock poisoned")) + .map(|vault| { + vault + .get(&path) + .and_then(|items| items.get(entry_name)) + .cloned() + }); + } + + let entry_name = entry_name.to_string(); + tokio::task::spawn_blocking(move || { + let _guard = native_keyring_lock() + .lock() + .map_err(|_| "subscription keyring lock poisoned".to_string())?; + let entry = keyring::Entry::new(KEYRING_SERVICE, &entry_name) + .map_err(|err| format!("open system credential entry: {err}"))?; + match entry.get_secret() { + Ok(secret) => Ok(Some(secret)), + Err(keyring::Error::NoEntry) => Ok(None), + Err(err) => Err(format!("read system credential entry: {err}")), + } + }) + .await + .context("join system credential read task")? + .map_err(anyhow::Error::msg) +} + +/// Reads the v1 combined JSON entry. It was written through the password API, +/// which uses a platform-specific text encoding on Windows, so it cannot be +/// safely read through `get_secret` there. +async fn get_legacy_password(provider: &str) -> Result> { + if let Some(path) = overridden_store_path() { + if test_vault_is_unavailable(&path) { + return Err(anyhow!("subscription test vault unavailable")); + } + return test_secrets() + .lock() + .map_err(|_| anyhow!("subscription test vault lock poisoned")) + .map(|vault| { + vault + .get(&path) + .and_then(|items| items.get(provider)) + .and_then(|bytes| String::from_utf8(bytes.clone()).ok()) + }); + } + + let provider = provider.to_string(); + tokio::task::spawn_blocking(move || { + let _guard = native_keyring_lock() + .lock() + .map_err(|_| "subscription keyring lock poisoned".to_string())?; + let entry = keyring::Entry::new(KEYRING_SERVICE, &provider) + .map_err(|err| format!("open system credential entry: {err}"))?; + match entry.get_password() { + Ok(secret) => Ok(Some(secret)), + Err(keyring::Error::NoEntry) => Ok(None), + Err(err) => Err(format!("read legacy system credential entry: {err}")), + } + }) + .await + .context("join legacy system credential read task")? + .map_err(anyhow::Error::msg) +} + +async fn set_secret_bytes(entry_name: &str, secret: Vec) -> Result<()> { + if secret.len() > SECRET_CHUNK_BYTES { + return Err(anyhow!( + "subscription credential chunk exceeds portable size limit: {} bytes", + secret.len() + )); + } + if let Some(path) = overridden_store_path() { + if test_vault_is_unavailable(&path) { + return Err(anyhow!("subscription test vault unavailable")); + } + let mut vault = test_secrets() + .lock() + .map_err(|_| anyhow!("subscription test vault lock poisoned"))?; + vault + .entry(path) + .or_default() + .insert(entry_name.to_string(), secret); + return Ok(()); + } + + let entry_name = entry_name.to_string(); + tokio::task::spawn_blocking(move || { + let _guard = native_keyring_lock() + .lock() + .map_err(|_| "subscription keyring lock poisoned".to_string())?; + let entry = keyring::Entry::new(KEYRING_SERVICE, &entry_name) + .map_err(|err| format!("open system credential entry: {err}"))?; + entry + .set_secret(&secret) + .map_err(|err| format!("write system credential entry: {err}")) + }) + .await + .context("join system credential write task")? + .map_err(anyhow::Error::msg) +} + +async fn delete_secret_entry(entry_name: &str) -> Result<()> { + if let Some(path) = overridden_store_path() { + if test_vault_is_unavailable(&path) { + return Err(anyhow!("subscription test vault unavailable")); + } + if let Ok(mut vault) = test_secrets().lock() { + if let Some(items) = vault.get_mut(&path) { + items.remove(entry_name); + } + } + return Ok(()); + } + + let entry_name = entry_name.to_string(); + tokio::task::spawn_blocking(move || { + let _guard = native_keyring_lock() + .lock() + .map_err(|_| "subscription keyring lock poisoned".to_string())?; + let entry = keyring::Entry::new(KEYRING_SERVICE, &entry_name) + .map_err(|err| format!("open system credential entry: {err}"))?; + match entry.delete_credential() { + Ok(()) | Err(keyring::Error::NoEntry) => Ok(()), + Err(err) => Err(format!("delete system credential entry: {err}")), + } + }) + .await + .context("join system credential delete task")? + .map_err(anyhow::Error::msg) +} + +fn secret_chunks(secret: &str) -> Vec> { + if secret.is_empty() { + return vec![Vec::new()]; + } + secret + .as_bytes() + .chunks(SECRET_CHUNK_BYTES) + .map(<[u8]>::to_vec) + .collect() +} + +async fn read_chunked_field( + provider: &str, + set_id: &str, + field: &str, + parts: u32, +) -> Result> { + if parts == 0 { + return Ok(None); + } + let mut bytes = Vec::new(); + for entry_name in secret_entry_names(provider, set_id, field, parts) { + let Some(mut part) = get_secret_bytes(&entry_name).await? else { + return Ok(None); + }; + bytes.append(&mut part); + } + match String::from_utf8(bytes) { + Ok(secret) => Ok(Some(secret)), + Err(error) => { + log::warn!( + "subscription credential vault chunks are invalid for provider {provider} field {field}: {error}" + ); + Ok(None) + } + } +} + +async fn read_secret_material( + provider: &str, + metadata: &CredentialMetadata, +) -> Result> { + match metadata { + CredentialMetadata::Oauth { + secret_set_id: Some(set_id), + refresh_parts, + access_parts, + .. + } => { + let Some(refresh) = + read_chunked_field(provider, set_id, "refresh", *refresh_parts).await? + else { + return Ok(None); + }; + let Some(access) = + read_chunked_field(provider, set_id, "access", *access_parts).await? + else { + return Ok(None); + }; + Ok(Some(SecretMaterial::Oauth { refresh, access })) + } + CredentialMetadata::Api { + secret_set_id: Some(set_id), + key_parts, + .. + } => Ok(read_chunked_field(provider, set_id, "api-key", *key_parts) + .await? + .map(|key| SecretMaterial::Api { key })), + // Backward-compatible read of the original combined JSON password. + _ => { + let Some(secret) = get_legacy_password(provider).await? else { + return Ok(None); + }; + match serde_json::from_str(&secret) { + Ok(material) => Ok(Some(material)), + Err(error) => { + log::warn!( + "legacy subscription credential vault entry is invalid for provider {provider}: {error}" + ); + Ok(None) + } + } + } + } +} + +async fn write_secret_material( + provider: &str, + metadata: &CredentialMetadata, + credential: &StoredCredential, +) -> Result<()> { + let (set_id, fields): (&str, Vec<(&str, &str)>) = match (metadata, credential) { + ( + CredentialMetadata::Oauth { + secret_set_id: Some(set_id), + .. + }, + StoredCredential::Oauth { + refresh, access, .. + }, + ) => (set_id, vec![("refresh", refresh), ("access", access)]), + ( + CredentialMetadata::Api { + secret_set_id: Some(set_id), + .. + }, + StoredCredential::Api { key, .. }, + ) => (set_id, vec![("api-key", key)]), + _ => return Err(anyhow!("subscription credential metadata type mismatch")), + }; + + let mut written: Vec = Vec::new(); + for (field, value) in fields { + for (index, chunk) in secret_chunks(value).into_iter().enumerate() { + let entry_name = format!("{provider}/v2/{set_id}/{field}/{index}"); + if let Err(error) = set_secret_bytes(&entry_name, chunk).await { + for previous in written { + let _ = delete_secret_entry(&previous).await; + } + return Err(error); + } + written.push(entry_name); + } + } + Ok(()) +} + +async fn delete_secret_material(provider: &str, metadata: &CredentialMetadata) -> Result<()> { + let mut first_error = None; + for entry_name in metadata.vault_entries(provider) { + if let Err(error) = delete_secret_entry(&entry_name).await { + if first_error.is_none() { + first_error = Some(error); + } + } + } + if let Some(error) = first_error { + Err(error) + } else { + Ok(()) + } +} + +async fn write_secure_file(path: &Path, file: &SecureStoreFile) -> Result<()> { + if metadata_write_should_fail(path) { + return Err(anyhow!("injected subscription metadata write failure")); + } if let Some(parent) = path.parent() { tokio::fs::create_dir_all(parent) .await - .with_context(|| format!("create subscription auth store dir {}", parent.display()))?; + .with_context(|| format!("create subscription auth directory {}", parent.display()))?; + } + let bytes = serde_json::to_vec_pretty(file)?; + write_atomic(path, &bytes).await +} + +async fn migrate_legacy_store(path: &Path, legacy: Store) -> Result { + let mut secure = SecureStoreFile { + version: STORE_VERSION, + accounts: HashMap::new(), + }; + let mut written: Vec<(String, CredentialMetadata)> = Vec::new(); + for (provider, credential) in &legacy { + let metadata = CredentialMetadata::from_credential(credential); + if let Err(error) = write_secret_material(provider, &metadata, credential).await { + for (previous_provider, previous_metadata) in written { + let _ = delete_secret_material(&previous_provider, &previous_metadata).await; + } + // Never keep plaintext tokens after the security migration has + // been attempted. Preserve account labels/expiry only and make the + // required one-time sign-in explicit to the UI. + secure.accounts = legacy + .iter() + .map(|(key, value)| { + ( + key.clone(), + CredentialMetadata::requiring_reauthentication(value), + ) + }) + .collect(); + write_secure_file(path, &secure).await?; + log::warn!( + "subscription credential vault migration failed; plaintext credentials were removed and reauthentication is required: {error:#}" + ); + return Ok(LoadState { + credentials: Store::new(), + requires_reauthentication: secure.accounts.keys().cloned().collect(), + vault_unavailable: HashSet::new(), + }); + } + secure.accounts.insert(provider.clone(), metadata.clone()); + written.push((provider.clone(), metadata)); + } + if let Err(error) = write_secure_file(path, &secure).await { + for (provider, metadata) in written { + let _ = delete_secret_material(&provider, &metadata).await; + } + return Err(error); + } + log::info!("subscription credentials migrated to the system credential vault"); + Ok(LoadState { + credentials: legacy, + requires_reauthentication: HashSet::new(), + vault_unavailable: HashSet::new(), + }) +} + +/// Loads credentials plus vault availability state. Legacy plaintext files are +/// migrated in place and immediately rewritten without secret fields. +async fn load_with_state_unlocked() -> Result { + let path = store_path()?; + let Some(bytes) = read_bytes(&path).await? else { + return Ok(LoadState { + credentials: Store::new(), + requires_reauthentication: HashSet::new(), + vault_unavailable: HashSet::new(), + }); + }; + + let secure = match parse_secure_file(&bytes, &path) { + Ok(file) => file, + Err(secure_error) => match serde_json::from_slice::(&bytes) { + Ok(legacy) => return migrate_legacy_store(&path, legacy).await, + Err(_) => return Err(secure_error), + }, + }; + + let mut credentials = Store::new(); + let mut requires_reauthentication = HashSet::new(); + let mut vault_unavailable = HashSet::new(); + for (provider, metadata) in secure.accounts { + if metadata.needs_reauthentication() { + requires_reauthentication.insert(provider); + continue; + } + let material = match read_secret_material(&provider, &metadata).await { + Ok(Some(material)) => material, + Ok(None) => { + requires_reauthentication.insert(provider); + continue; + } + Err(error) => { + log::warn!( + "subscription credential vault is unavailable for provider {provider}: {error:#}" + ); + vault_unavailable.insert(provider); + continue; + } + }; + if let Some(credential) = metadata.combine(material) { + credentials.insert(provider, credential); + } else { + requires_reauthentication.insert(provider); + } } - let bytes = serde_json::to_vec_pretty(store)?; - write_atomic(&path, &bytes).await + Ok(LoadState { + credentials, + requires_reauthentication, + vault_unavailable, + }) } -/// Writes `bytes` to `path` atomically (temp file + rename) so a crash -/// mid-write cannot corrupt the token store. On Unix the temp file is created -/// with mode `0600` up front, so tokens are never briefly world-readable. -async fn write_atomic(path: &std::path::Path, bytes: &[u8]) -> Result<()> { +/// Serializes discovery with migrations and metadata mutations so callers +/// never observe or race a partially rewritten credential index. +pub(crate) async fn load_with_state() -> Result { + let _guard = store_operation_lock().lock().await; + load_with_state_unlocked().await +} + +/// Loads all credentials that are currently available from the system vault. +pub async fn load() -> Result { + Ok(load_with_state().await?.credentials) +} + +/// Loads one provider credential without exposing its secret in the metadata +/// file. `None` means the provider needs a new sign-in. +pub async fn load_entry(provider: &str) -> Result> { + let mut state = load_with_state().await?; + if state.vault_unavailable.contains(provider) { + return Err(anyhow!( + "system credential vault is locked or unavailable; unlock it and retry" + )); + } + Ok(state.credentials.remove(provider)) +} + +/// Inserts or replaces a provider credential. Secret material is committed to +/// the platform vault before the non-secret metadata advertises the entry. +pub async fn upsert(provider: &str, credential: StoredCredential) -> Result<()> { + let _guard = store_operation_lock().lock().await; + let path = store_path()?; + // Trigger one-time migration before modifying an older file. + let _ = load_with_state_unlocked().await?; + let mut file = read_secure_file(&path).await?; + let previous = file.accounts.get(provider).cloned(); + let metadata = CredentialMetadata::from_credential(&credential); + write_secret_material(provider, &metadata, &credential).await?; + file.accounts.insert(provider.to_string(), metadata.clone()); + if let Err(error) = write_secure_file(&path, &file).await { + let _ = delete_secret_material(provider, &metadata).await; + return Err(error); + } + if let Some(previous) = previous { + if let Err(error) = delete_secret_material(provider, &previous).await { + log::warn!( + "remove superseded subscription credential chunks failed for provider {provider}: {error:#}" + ); + } + } + Ok(()) +} + +/// Removes one provider from both the native vault and metadata index. +pub async fn remove(provider: &str) -> Result<()> { + let _guard = store_operation_lock().lock().await; + let path = store_path()?; + let _ = load_with_state_unlocked().await?; + let mut file = read_secure_file(&path).await?; + let previous = file.accounts.remove(provider); + // Commit the discovery-index removal first. If this write fails, the old + // metadata and vault chunks remain a usable pair. A crash after this point + // can leave only unreachable vault chunks, never a visible broken account. + write_secure_file(&path, &file).await?; + let cleanup = match previous.as_ref() { + Some(metadata) => delete_secret_material(provider, metadata).await, + None => delete_secret_entry(provider).await, + }; + if let Err(error) = cleanup { + log::warn!( + "remove unreachable subscription credential chunks failed for provider {provider}: {error:#}" + ); + } + Ok(()) +} + +/// Persists all supplied credentials. Kept for compatibility with focused +/// tests; production refresh/login paths should call [`upsert`] for one +/// provider so concurrent providers cannot overwrite each other's tokens. +pub async fn save(store: &Store) -> Result<()> { + for (provider, credential) in store { + upsert(provider, credential.clone()).await?; + } + Ok(()) +} + +/// Writes `bytes` atomically (temp file + rename). Although the v2 file is +/// non-secret, restrictive Unix permissions protect account metadata too. +async fn write_atomic(path: &Path, bytes: &[u8]) -> Result<()> { use tokio::io::AsyncWriteExt; let tmp = path.with_extension("tmp"); @@ -126,12 +952,78 @@ async fn write_atomic(path: &std::path::Path, bytes: &[u8]) -> Result<()> { .await .with_context(|| format!("sync subscription auth temp file {}", tmp.display()))?; } - tokio::fs::rename(&tmp, path).await.with_context(|| { + replace_metadata_file(&tmp, path).await +} + +#[cfg(not(windows))] +async fn replace_metadata_file(tmp: &Path, path: &Path) -> Result<()> { + tokio::fs::rename(tmp, path).await.with_context(|| { format!( - "rename subscription auth store {} -> {}", + "rename subscription auth metadata {} -> {}", tmp.display(), path.display() ) - })?; + }) +} + +/// Windows does not reliably replace an existing destination with `rename`. +/// Rotate the prior metadata file to a backup first and restore it if moving +/// the newly-synced temp file into place fails. +#[cfg(windows)] +async fn replace_metadata_file(tmp: &Path, path: &Path) -> Result<()> { + let backup = path.with_extension("bak"); + match tokio::fs::remove_file(&backup).await { + Ok(()) => {} + Err(error) if error.kind() == std::io::ErrorKind::NotFound => {} + Err(error) => { + return Err(error).with_context(|| { + format!( + "remove stale subscription auth metadata backup {}", + backup.display() + ) + }); + } + } + + let had_existing = match tokio::fs::rename(path, &backup).await { + Ok(()) => true, + Err(error) if error.kind() == std::io::ErrorKind::NotFound => false, + Err(error) => { + return Err(error).with_context(|| { + format!( + "rotate subscription auth metadata {} -> {}", + path.display(), + backup.display() + ) + }); + } + }; + + if let Err(error) = tokio::fs::rename(tmp, path).await { + if had_existing { + if let Err(restore_error) = tokio::fs::rename(&backup, path).await { + return Err(anyhow!( + "replace subscription auth metadata failed: {error}; restoring {} also failed: {restore_error}", + backup.display() + )); + } + } + return Err(error).with_context(|| { + format!( + "rename subscription auth metadata {} -> {}", + tmp.display(), + path.display() + ) + }); + } + + if had_existing { + tokio::fs::remove_file(&backup).await.with_context(|| { + format!( + "remove subscription auth metadata backup {}", + backup.display() + ) + })?; + } Ok(()) } diff --git a/src/crates/assembly/core/src/agentic/tools/implementations/page_deploy_tool.rs b/src/crates/assembly/core/src/agentic/tools/implementations/page_deploy_tool.rs index d0312db44d..e03a07de45 100644 --- a/src/crates/assembly/core/src/agentic/tools/implementations/page_deploy_tool.rs +++ b/src/crates/assembly/core/src/agentic/tools/implementations/page_deploy_tool.rs @@ -31,9 +31,9 @@ impl Tool for PageDeployTool { Ok( r#"Switch the production pointer of an existing BitFun Page to a previously saved version_id (rollback or promote a prior version). -Requires a logged-in BitFun account. This tool is only available after account login. To create or update page content and publish, use PagePublish instead — do not ask the user for a version_id they do not have, and do not mention a Page management scene. +Requires a logged-in BitFun account. This tool is only available after account login. To create or update page content and publish, use PagePublish instead. Existing versions can also be reviewed from the Pages scene. -Input: slug (page path id), version_id (immutable saved version from a prior PagePublish). Returns absolute `url` plus url_path / deployed_version_id. +Input: slug (page path id), version_id (immutable saved version from a prior PagePublish). Returns absolute `url` plus url_path / deployed_version_id. Public links can be shared directly. Private and relay links must be opened or copied through the Pages scene/tool card so the browser receives a scoped one-time access handoff. When telling the user the link: paste the full absolute URL and put a trailing space after it (before any punctuation or newline). @@ -70,10 +70,43 @@ Preview a version at /p/{username}/{slug}/@v/{version_id}."# fn permission_intents( &self, - _input: &Value, + input: &Value, _context: &ToolUseContext, ) -> BitFunResult> { - Ok(Vec::new()) + let slug = input + .get("slug") + .and_then(Value::as_str) + .map(str::trim) + .filter(|value| !value.is_empty()) + .unwrap_or(""); + let version_id = input + .get("version_id") + .and_then(Value::as_str) + .map(str::trim) + .filter(|value| !value.is_empty()) + .unwrap_or(""); + let resource = format!("page:{slug}; production-version={version_id}"); + let mut intent = PermissionIntent::new("page_deploy", vec![resource]); + intent.save_resources.clear(); + intent.display_metadata.insert( + "permissionScope".to_string(), + Value::String("account".to_string()), + ); + intent + .display_metadata + .insert("requiresFreshApproval".to_string(), Value::Bool(true)); + intent.display_metadata.insert( + "pageOperation".to_string(), + Value::String("deploy".to_string()), + ); + intent + .display_metadata + .insert("pageSlug".to_string(), Value::String(slug.to_string())); + intent.display_metadata.insert( + "pageVersion".to_string(), + Value::String(version_id.to_string()), + ); + Ok(vec![intent]) } async fn is_available_in_context(&self, _context: Option<&ToolUseContext>) -> bool { @@ -116,15 +149,11 @@ Preview a version at /p/{username}/{slug}/@v/{version_id}."# .filter(|s| !s.is_empty()) .or_else(|| result.get("url_path").and_then(|v| v.as_str())) .unwrap_or(""); - // Trailing space after URL keeps chat linkifiers from eating the next char. - let assistant = if url.is_empty() { - format!("Deployed BitFun Page '{slug}' version '{version_id}' to production.") - } else { - format!( - "Deployed BitFun Page '{slug}' version '{version_id}'. Production URL: {url} \n\ - Share this full absolute URL with the user, and keep a trailing space after the URL." - ) - }; + let visibility = result + .get("visibility") + .and_then(Value::as_str) + .unwrap_or("private"); + let assistant = deploy_result_for_assistant(&slug, &version_id, visibility, url); Ok(vec![ToolResult::Result { data: result, @@ -134,6 +163,28 @@ Preview a version at /p/{username}/{slug}/@v/{version_id}."# } } +fn deploy_result_for_assistant( + slug: &str, + version_id: &str, + visibility: &str, + url: &str, +) -> String { + if visibility != "public" { + return format!( + "Deployed BitFun Page '{slug}' version '{version_id}' with {visibility} visibility. Open or copy it from the Pages scene/tool card so BitFun can create a scoped browser-access link; do not share the raw Page URL." + ); + } + // Trailing space after URL keeps chat linkifiers from eating the next char. + if url.is_empty() { + format!("Deployed public BitFun Page '{slug}' version '{version_id}' to production.") + } else { + format!( + "Deployed public BitFun Page '{slug}' version '{version_id}'. Production URL: {url} \n\ + Share this full absolute URL with the user, and keep a trailing space after the URL." + ) + } +} + #[cfg(test)] mod tests { use super::*; @@ -180,4 +231,47 @@ mod tests { .expect_err("should reject without login"); assert!(err.to_string().contains("logged-in")); } + + #[test] + fn permission_intent_names_page_and_target_version() { + let tool = PageDeployTool::new(); + let intents = tool + .permission_intents( + &json!({ "slug": "demo", "version_id": "v2026" }), + &empty_context(), + ) + .expect("permission intent"); + + assert_eq!(intents.len(), 1); + assert_eq!(intents[0].action, "page_deploy"); + assert_eq!( + intents[0].resources, + vec!["page:demo; production-version=v2026"] + ); + assert!(intents[0].save_resources.is_empty()); + assert_eq!( + intents[0].display_metadata.get("permissionScope"), + Some(&Value::String("account".to_string())) + ); + assert_eq!( + intents[0].display_metadata.get("requiresFreshApproval"), + Some(&Value::Bool(true)) + ); + assert_eq!( + intents[0].display_metadata.get("pageOperation"), + Some(&Value::String("deploy".to_string())) + ); + } + + #[test] + fn private_deploy_result_never_advertises_the_raw_url() { + let message = deploy_result_for_assistant( + "demo", + "v1", + "private", + "https://relay.example/p/alice/demo", + ); + assert!(!message.contains("https://")); + assert!(message.contains("scoped browser-access link")); + } } diff --git a/src/crates/assembly/core/src/agentic/tools/implementations/page_publish_tool.rs b/src/crates/assembly/core/src/agentic/tools/implementations/page_publish_tool.rs index ea5293ce2e..77e0875b69 100644 --- a/src/crates/assembly/core/src/agentic/tools/implementations/page_publish_tool.rs +++ b/src/crates/assembly/core/src/agentic/tools/implementations/page_publish_tool.rs @@ -1,10 +1,10 @@ //! PagePublish tool — create/update a BitFun Page from inline files or a directory, -//! then optionally deploy to production (default: deploy). +//! then optionally deploy to production. use std::collections::HashMap; use crate::agentic::tools::account_login_capability::account_login_available; -use crate::agentic::tools::framework::{Tool, ToolResult, ToolUseContext}; +use crate::agentic::tools::framework::{PermissionIntent, Tool, ToolResult, ToolUseContext}; use crate::agentic::tools::page_publish_host::{invoke_page_publish, PagePublishHostRequest}; use crate::util::errors::{BitFunError, BitFunResult}; use async_trait::async_trait; @@ -32,26 +32,26 @@ impl Tool for PagePublishTool { async fn description(&self) -> BitFunResult { Ok( - r#"Publish a BitFun Page to the account relay: upload content, freeze an immutable version, and deploy to production by default. + r#"Publish a BitFun Page to the account relay: upload content, freeze an immutable version, and optionally deploy it to production. -Requires a logged-in BitFun account. This tool is only available after account login. There is no separate Page management scene — create and ship pages from the conversation with this tool. +Requires a logged-in BitFun account. This tool is only available after account login. Published Pages can be reviewed and managed later from the Pages scene. When you produce self-contained publishable web content (landing page, docs site, or a Page with server/worker.js) and the user is logged in, proactively ask whether they want it published to BitFun Page (suggest a slug and visibility). If they already said publish/deploy/上线, proceed with permission confirmation. IMPORTANT — content source: - Prefer `files` (inline path→UTF-8 content). For agent-authored pages, pass HTML/JS directly in `files` and call PagePublish. Do NOT Write/Edit page files into the user workspace just to publish, and do NOT create folders like bitfun-page/ unless the user explicitly asked to keep a local copy. -- Use `directory` only when the user already has (or explicitly wants) page sources on disk in the workspace. +- Use `directory` only when the user already has (or explicitly wants) page sources on disk in a local workspace. Remote workspaces must use `files` so BitFun never mistakes a remote path for a local path. Input: - slug (required): page path id (lowercase letters, digits, hyphens) -- visibility: private | relay | public (default public) +- visibility: private | relay | public (default private) - title?, note? -- deploy: boolean (default true). false = save version only; still returns absolute preview_url +- deploy: boolean (default false). false = save version only; still returns absolute preview_url - Exactly one of: - files: object map of relative path → UTF-8 file content (default/preferred). Must include index.html and/or server/worker.js - directory: existing local workspace path (only when user wants on-disk sources) -Returns version_id, absolute `url` / `preview_url` (plus relative paths), deployed_version_id when deployed. +Returns version_id, absolute `url` / `preview_url` (plus relative paths), deployed_version_id when deployed. Public links can be shared directly. Private and relay links must be opened or copied through the Pages scene/tool card so the browser receives a scoped one-time access handoff; never share their raw URL as if it were independently accessible. When telling the user the link: paste the full absolute URL and put a trailing space after it (before any punctuation or newline), so chat linkifiers do not swallow the next character. @@ -77,7 +77,7 @@ Use PageDeploy only to switch an already-saved version_id (rollback / promote a "visibility": { "type": "string", "enum": ["private", "relay", "public"], - "description": "Page visibility. Defaults to public." + "description": "Page visibility. Defaults to private. Use public only when the user explicitly intends to share the page publicly." }, "title": { "type": "string", @@ -89,7 +89,7 @@ Use PageDeploy only to switch an already-saved version_id (rollback / promote a }, "deploy": { "type": "boolean", - "description": "Deploy to production after saving. Defaults to true." + "description": "Deploy to production after saving. Defaults to false; set true only when the user explicitly asked to publish or deploy." }, "files": { "type": "object", @@ -98,7 +98,7 @@ Use PageDeploy only to switch an already-saved version_id (rollback / promote a }, "directory": { "type": "string", - "description": "Only when page sources already exist (or user asked to keep them) on disk. Mutually exclusive with files." + "description": "Only when page sources already exist (or user asked to keep them) on disk in a local workspace. Remote workspaces must use files. Mutually exclusive with files." } } }) @@ -108,6 +108,60 @@ Use PageDeploy only to switch an already-saved version_id (rollback / promote a false } + fn permission_intents( + &self, + input: &Value, + _context: &ToolUseContext, + ) -> BitFunResult> { + let slug = input + .get("slug") + .and_then(Value::as_str) + .map(str::trim) + .filter(|value| !value.is_empty()) + .unwrap_or(""); + let visibility = input + .get("visibility") + .and_then(Value::as_str) + .map(str::trim) + .filter(|value| !value.is_empty()) + .unwrap_or("private"); + let deploy = input + .get("deploy") + .and_then(Value::as_bool) + .unwrap_or(false); + let resource = format!( + "page:{slug}; visibility={visibility}; deploy={}", + if deploy { + "production" + } else { + "saved-version-only" + } + ); + let mut intent = PermissionIntent::new("page_publish", vec![resource]); + // Publishing decisions should stay per-call: a remembered wildcard grant could + // otherwise hide a later change from private preview to public production. + intent.save_resources.clear(); + intent.display_metadata.insert( + "permissionScope".to_string(), + Value::String("account".to_string()), + ); + intent + .display_metadata + .insert("requiresFreshApproval".to_string(), Value::Bool(true)); + intent.display_metadata.insert( + "pageOperation".to_string(), + Value::String(if deploy { "publish" } else { "save" }.to_string()), + ); + intent + .display_metadata + .insert("pageSlug".to_string(), Value::String(slug.to_string())); + intent.display_metadata.insert( + "pageVisibility".to_string(), + Value::String(visibility.to_string()), + ); + Ok(vec![intent]) + } + async fn is_available_in_context(&self, _context: Option<&ToolUseContext>) -> bool { account_login_available() } @@ -115,7 +169,7 @@ Use PageDeploy only to switch an already-saved version_id (rollback / promote a async fn call_impl( &self, input: &Value, - _context: &ToolUseContext, + context: &ToolUseContext, ) -> BitFunResult> { if !account_login_available() { return Err(BitFunError::tool( @@ -136,7 +190,7 @@ Use PageDeploy only to switch an already-saved version_id (rollback / promote a .and_then(|v| v.as_str()) .map(str::trim) .filter(|s| !s.is_empty()) - .unwrap_or("public") + .unwrap_or("private") .to_string(); let title = input @@ -155,7 +209,7 @@ Use PageDeploy only to switch an already-saved version_id (rollback / promote a let deploy = input .get("deploy") .and_then(|v| v.as_bool()) - .unwrap_or(true); + .unwrap_or(false); let directory = input .get("directory") @@ -164,6 +218,19 @@ Use PageDeploy only to switch an already-saved version_id (rollback / promote a .filter(|s| !s.is_empty()) .map(str::to_string); + if directory.is_some() && context.workspace.is_none() { + return Err(BitFunError::tool( + "PagePublish directory requires a local workspace; use inline files when no workspace is open" + .to_string(), + )); + } + if directory.is_some() && context.is_remote() { + return Err(BitFunError::tool( + "PagePublish cannot read a remote workspace directory through the local desktop host; use inline files" + .to_string(), + )); + } + let files = parse_files_map(input.get("files"))?; match (&directory, &files) { @@ -182,7 +249,7 @@ Use PageDeploy only to switch an already-saved version_id (rollback / promote a let result = invoke_page_publish(PagePublishHostRequest { slug: slug.clone(), - visibility, + visibility: visibility.clone(), title, note, deploy, @@ -213,24 +280,8 @@ Use PageDeploy only to switch an already-saved version_id (rollback / promote a .and_then(|v| v.as_bool()) .unwrap_or(false); - // Trailing space after URL keeps chat linkifiers from eating the next char. - let assistant = if deployed { - if url.is_empty() { - format!("Published BitFun Page '{slug}' version '{version_id}' to production.") - } else { - format!( - "Published BitFun Page '{slug}' version '{version_id}'. Production URL: {url} \n\ - Share this full absolute URL with the user, and keep a trailing space after the URL." - ) - } - } else if preview.is_empty() { - format!("Saved BitFun Page '{slug}' version '{version_id}' (not deployed).") - } else { - format!( - "Saved BitFun Page '{slug}' version '{version_id}' (not deployed). Preview URL: {preview} \n\ - Share this full absolute URL with the user, and keep a trailing space after the URL." - ) - }; + let assistant = + publish_result_for_assistant(&slug, version_id, &visibility, deployed, url, preview); Ok(vec![ToolResult::Result { data: result, @@ -240,6 +291,45 @@ Use PageDeploy only to switch an already-saved version_id (rollback / promote a } } +fn publish_result_for_assistant( + slug: &str, + version_id: &str, + visibility: &str, + deployed: bool, + url: &str, + preview: &str, +) -> String { + if visibility != "public" { + let state = if deployed { + "published to production" + } else { + "saved without changing production" + }; + return format!( + "BitFun Page '{slug}' version '{version_id}' was {state} with {visibility} visibility. Open or copy it from the Pages scene/tool card so BitFun can create a scoped browser-access link; do not share the raw Page URL." + ); + } + + // Trailing space after URL keeps chat linkifiers from eating the next char. + if deployed { + if url.is_empty() { + format!("Published public BitFun Page '{slug}' version '{version_id}' to production.") + } else { + format!( + "Published public BitFun Page '{slug}' version '{version_id}'. Production URL: {url} \n\ + Share this full absolute URL with the user, and keep a trailing space after the URL." + ) + } + } else if preview.is_empty() { + format!("Saved public BitFun Page '{slug}' version '{version_id}' (not deployed).") + } else { + format!( + "Saved public BitFun Page '{slug}' version '{version_id}' (not deployed). Preview URL: {preview} \n\ + Share this full absolute URL with the user, and keep a trailing space after the URL." + ) + } +} + fn parse_files_map(value: Option<&Value>) -> BitFunResult>> { let Some(value) = value else { return Ok(None); @@ -326,4 +416,110 @@ mod tests { assert!(err.to_string().contains("directory or files")); set_account_login_available(false); } + + #[test] + fn permission_intent_exposes_safe_defaults_and_publish_scope() { + let tool = PagePublishTool::new(); + let intents = tool + .permission_intents( + &json!({ + "slug": "release-notes", + "files": { "index.html": "" } + }), + &empty_context(), + ) + .expect("permission intent"); + + assert_eq!(intents.len(), 1); + assert_eq!(intents[0].action, "page_publish"); + assert_eq!( + intents[0].resources, + vec!["page:release-notes; visibility=private; deploy=saved-version-only"] + ); + assert!(intents[0].save_resources.is_empty()); + assert_eq!( + intents[0].display_metadata.get("permissionScope"), + Some(&Value::String("account".to_string())) + ); + assert_eq!( + intents[0].display_metadata.get("requiresFreshApproval"), + Some(&Value::Bool(true)) + ); + assert_eq!( + intents[0].display_metadata.get("pageOperation"), + Some(&Value::String("save".to_string())) + ); + + let public_deploy = tool + .permission_intents( + &json!({ + "slug": "launch", + "visibility": "public", + "deploy": true, + "files": { "index.html": "" } + }), + &empty_context(), + ) + .expect("public deploy permission intent"); + assert_eq!( + public_deploy[0].resources, + vec!["page:launch; visibility=public; deploy=production"] + ); + assert_eq!( + public_deploy[0].display_metadata.get("pageOperation"), + Some(&Value::String("publish".to_string())) + ); + } + + #[tokio::test] + async fn directory_source_requires_a_local_workspace() { + let _guard = LOGIN_GATE.lock().unwrap(); + let tool = PagePublishTool::new(); + set_account_login_available(true); + + let no_workspace_error = tool + .call_impl( + &json!({ "slug": "demo", "directory": "page" }), + &empty_context(), + ) + .await + .expect_err("directory without workspace should be rejected"); + assert!(no_workspace_error.to_string().contains("local workspace")); + + let mut remote_context = empty_context(); + remote_context.workspace = Some(crate::agentic::WorkspaceBinding::new_remote( + None, + std::path::PathBuf::from("/srv/page"), + "connection-1".to_string(), + "Remote".to_string(), + crate::service::remote_ssh::workspace_state::WorkspaceSessionIdentity { + hostname: "remote.example".to_string(), + logical_workspace_path: "/srv/page".to_string(), + remote_connection_id: Some("connection-1".to_string()), + }, + )); + let remote_error = tool + .call_impl( + &json!({ "slug": "demo", "directory": "." }), + &remote_context, + ) + .await + .expect_err("remote directory should be rejected"); + assert!(remote_error.to_string().contains("remote workspace")); + set_account_login_available(false); + } + + #[test] + fn private_result_never_advertises_the_raw_url() { + let message = publish_result_for_assistant( + "demo", + "v1", + "private", + true, + "https://relay.example/p/alice/demo", + "https://relay.example/p/alice/demo/@v/v1", + ); + assert!(!message.contains("https://")); + assert!(message.contains("scoped browser-access link")); + } } diff --git a/src/crates/assembly/core/src/agentic/tools/pipeline/tool_pipeline.rs b/src/crates/assembly/core/src/agentic/tools/pipeline/tool_pipeline.rs index dadee63a78..ea8eb24dab 100644 --- a/src/crates/assembly/core/src/agentic/tools/pipeline/tool_pipeline.rs +++ b/src/crates/assembly/core/src/agentic/tools/pipeline/tool_pipeline.rs @@ -500,6 +500,40 @@ fn permission_project_path(context: &ToolUseContext) -> BitFunResult { .to_string()) } +const ACCOUNT_PERMISSION_SCOPE: &str = "account"; +const ACCOUNT_PERMISSION_PROJECT_ID: &str = "__bitfun_account_actions__"; +const ACCOUNT_PERMISSION_PROJECT_PATH: &str = "BitFun account"; + +fn permission_scope( + context: &ToolUseContext, + intents: &[PermissionIntent], +) -> BitFunResult<(String, String)> { + if context.workspace.is_some() { + return Ok(( + permission_project_id(context)?, + permission_project_path(context)?, + )); + } + + let account_scoped = intents.iter().all(|intent| { + intent + .display_metadata + .get("permissionScope") + .and_then(serde_json::Value::as_str) + == Some(ACCOUNT_PERMISSION_SCOPE) + }); + if account_scoped { + return Ok(( + ACCOUNT_PERMISSION_PROJECT_ID.to_string(), + ACCOUNT_PERMISSION_PROJECT_PATH.to_string(), + )); + } + + Err(BitFunError::validation( + "A workspace is required for file permissions".to_string(), + )) +} + fn permission_resource_case_sensitivity( context: &ToolUseContext, ) -> PermissionResourceCaseSensitivity { @@ -564,10 +598,21 @@ fn permission_intent_effect( } } - if intent.resources.is_empty() { + let effect = if intent.resources.is_empty() { PermissionEffect::Ask } else { aggregate + }; + if effect != PermissionEffect::Deny + && intent + .display_metadata + .get("requiresFreshApproval") + .and_then(serde_json::Value::as_bool) + .unwrap_or(false) + { + PermissionEffect::Ask + } else { + effect } } @@ -623,8 +668,7 @@ impl ToolPipeline { return Ok(PermissionPlanDraft::Allowed); } - let project_id = permission_project_id(&context)?; - let project_path = permission_project_path(&context)?; + let (project_id, project_path) = permission_scope(&context, &intents)?; let permission_rules = task.options.permission_rules.clone(); let case_sensitivity = permission_resource_case_sensitivity(&context); let round_id = task.context.round_id.clone(); @@ -2209,6 +2253,65 @@ mod tests { ); } + #[test] + fn account_scoped_fresh_approval_works_without_a_workspace_and_ignores_allow_rules() { + let mut intent = PermissionIntent::new( + "page_publish", + vec!["page:demo; visibility=private; deploy=saved-version-only".to_string()], + ); + intent.display_metadata.insert( + "permissionScope".to_string(), + json!(ACCOUNT_PERMISSION_SCOPE), + ); + intent + .display_metadata + .insert("requiresFreshApproval".to_string(), json!(true)); + let context = ToolUseContext::for_tool_listing(None, None); + assert_eq!( + permission_scope(&context, &[intent.clone()]).expect("account scope"), + ( + ACCOUNT_PERMISSION_PROJECT_ID.to_string(), + ACCOUNT_PERMISSION_PROJECT_PATH.to_string(), + ) + ); + + let allow = vec![PermissionRule::new( + "page_publish", + "*", + PermissionEffect::Allow, + )]; + assert_eq!( + permission_intent_effect( + &intent, + &allow, + &[], + PermissionResourceCaseSensitivity::Sensitive, + ), + PermissionEffect::Ask + ); + let deny = vec![PermissionRule::new( + "page_publish", + "*", + PermissionEffect::Deny, + )]; + assert_eq!( + permission_intent_effect( + &intent, + &deny, + &[], + PermissionResourceCaseSensitivity::Sensitive, + ), + PermissionEffect::Deny + ); + } + + #[test] + fn ordinary_permission_intents_still_require_a_workspace() { + let context = ToolUseContext::for_tool_listing(None, None); + let intent = PermissionIntent::new("edit", vec!["src/main.rs".to_string()]); + assert!(permission_scope(&context, &[intent]).is_err()); + } + struct StaticTestTool { name: String, response: serde_json::Value, diff --git a/src/crates/assembly/core/src/service/remote_connect/settings_sync.rs b/src/crates/assembly/core/src/service/remote_connect/settings_sync.rs index c111b9c72a..22450c9669 100644 --- a/src/crates/assembly/core/src/service/remote_connect/settings_sync.rs +++ b/src/crates/assembly/core/src/service/remote_connect/settings_sync.rs @@ -26,7 +26,7 @@ use std::time::Duration; use anyhow::{anyhow, Result}; use log::{debug, warn}; -use tokio::sync::mpsc; +use tokio::sync::{mpsc, Notify}; use bitfun_services_integrations::remote_connect::account::{ error_indicates_expired_token, AccountClient, AccountSession, SettingsBlob, @@ -40,7 +40,7 @@ pub const SETTINGS_PUSH_DEBOUNCE: Duration = Duration::from_secs(5); /// Account context needed for every relay call: the session (token + /// master_key) and the relay base URL. -pub type AccountContext = (AccountSession, String); +pub type AccountContext = (AccountSession, String, u64); type AccountContextFn = dyn Fn() -> std::pin::Pin> + Send>> + Send @@ -54,6 +54,10 @@ pub struct SettingsSyncHooks { /// logged out. Required for the background loop; one-shot helpers take /// the context explicitly. pub account_context: Option>, + /// Confirms that a context generation captured before an async relay call + /// still belongs to the active account. Hosts bump the generation before + /// logout or replacement login. + pub is_account_context_current: Option bool + Send + Sync>>, /// When true, push and pull are paused (Desktop: Peer controller mode). pub should_pause: Option bool + Send + Sync>>, /// Fired after cloud settings were applied to the local config. @@ -73,6 +77,7 @@ static PUSH_TX: OnceLock> = OnceLock::new(); /// user-chosen direction. A counter (not a flag) so concurrent ops do not /// clear each other's in-flight state on completion. static SYNC_OPS_IN_FLIGHT: AtomicUsize = AtomicUsize::new(0); +static SYNC_OPS_IDLE: Notify = Notify::const_new(); /// RAII guard that marks a sync upload/apply as in flight. struct SyncOpGuard; @@ -84,13 +89,29 @@ impl SyncOpGuard { } impl Drop for SyncOpGuard { fn drop(&mut self) { - SYNC_OPS_IN_FLIGHT.fetch_sub(1, Ordering::SeqCst); + if SYNC_OPS_IN_FLIGHT.fetch_sub(1, Ordering::SeqCst) == 1 { + SYNC_OPS_IDLE.notify_waiters(); + } + } +} + +/// Wait until every settings upload/apply critical section has completed. +/// Hosts call this after invalidating their account generation and before +/// completing logout or installing a replacement account. +pub async fn wait_for_sync_operations_idle() { + loop { + let notified = SYNC_OPS_IDLE.notified(); + if SYNC_OPS_IN_FLIGHT.load(Ordering::SeqCst) == 0 { + return; + } + notified.await; } } fn hooks() -> &'static SettingsSyncHooks { static DEFAULT: SettingsSyncHooks = SettingsSyncHooks { account_context: None, + is_account_context_current: None, should_pause: None, on_settings_applied: None, on_settings_pushed: None, @@ -99,6 +120,14 @@ fn hooks() -> &'static SettingsSyncHooks { HOOKS.get().unwrap_or(&DEFAULT) } +fn is_account_context_current(generation: u64) -> bool { + hooks() + .is_account_context_current + .as_ref() + .map(|check| check(generation)) + .unwrap_or(true) +} + fn should_pause() -> bool { hooks().should_pause.as_ref().map(|f| f()).unwrap_or(false) } @@ -183,11 +212,29 @@ pub async fn upload_settings_payload( relay_url: &str, payload: &str, ) -> Result { + upload_settings_payload_for_generation(account, relay_url, payload, None).await +} + +async fn upload_settings_payload_for_generation( + account: &AccountSession, + relay_url: &str, + payload: &str, + generation: Option, +) -> Result { + if generation.is_some_and(|value| !is_account_context_current(value)) { + return Err(anyhow!("account context changed before settings upload")); + } let _op = SyncOpGuard::begin(); + if generation.is_some_and(|value| !is_account_context_current(value)) { + return Err(anyhow!("account context changed before settings upload")); + } let client = AccountClient::new(); // The version returned here is the exact one stored on the relay — // recording it keeps the next pull from re-applying our own upload. let version = client.upload_settings(relay_url, account, payload).await?; + if generation.is_some_and(|value| !is_account_context_current(value)) { + return Err(anyhow!("account context changed during settings upload")); + } let hash = settings_content_hash(payload).unwrap_or_default(); record_settings_cursor(&account.user_id, version, hash); debug!("Settings sync: uploaded settings (version={version})"); @@ -198,6 +245,14 @@ pub async fn upload_settings_payload( /// Export the current config and upload it when the content differs from the /// last uploaded/applied blob. Returns `true` when an upload happened. pub async fn push_settings_now(account: &AccountSession, relay_url: &str) -> Result { + push_settings_now_for_generation(account, relay_url, None).await +} + +async fn push_settings_now_for_generation( + account: &AccountSession, + relay_url: &str, + generation: Option, +) -> Result { let config_service = crate::service::config::get_global_config_service() .await .map_err(|e| anyhow!("config service: {e}"))?; @@ -207,6 +262,10 @@ pub async fn push_settings_now(account: &AccountSession, relay_url: &str) -> Res .map_err(|e| anyhow!("export config: {e}"))?; let payload = serde_json::to_string(&exported).map_err(|e| anyhow!("serialize config: {e}"))?; + if generation.is_some_and(|value| !is_account_context_current(value)) { + return Err(anyhow!("account context changed while exporting settings")); + } + let hash = settings_content_hash(&payload)?; let known = sync_state::load_settings_cursor(&account.user_id); if known.hash == hash && known.version != 0 { @@ -214,7 +273,7 @@ pub async fn push_settings_now(account: &AccountSession, relay_url: &str) -> Res return Ok(false); } - upload_settings_payload(account, relay_url, &payload).await?; + upload_settings_payload_for_generation(account, relay_url, &payload, generation).await?; Ok(true) } @@ -226,6 +285,15 @@ pub async fn apply_settings_blob( account: &AccountSession, blob: &SettingsBlob, force: bool, +) -> Result { + apply_settings_blob_for_generation(account, blob, force, None).await +} + +async fn apply_settings_blob_for_generation( + account: &AccountSession, + blob: &SettingsBlob, + force: bool, + generation: Option, ) -> Result { if !force { let known = sync_state::load_settings_cursor(&account.user_id); @@ -233,7 +301,13 @@ pub async fn apply_settings_blob( return Ok(false); } } + if generation.is_some_and(|value| !is_account_context_current(value)) { + return Err(anyhow!("account context changed before settings apply")); + } let _op = SyncOpGuard::begin(); + if generation.is_some_and(|value| !is_account_context_current(value)) { + return Err(anyhow!("account context changed before settings apply")); + } let inner_config = inner_config_value(&blob.plaintext)?; let config_service = crate::service::config::get_global_config_service() @@ -273,6 +347,14 @@ pub async fn apply_settings_blob( /// when new settings were applied; `Ok(false)` also when no cloud settings /// exist yet. pub async fn pull_and_apply_settings(account: &AccountSession, relay_url: &str) -> Result { + pull_and_apply_settings_for_generation(account, relay_url, None).await +} + +async fn pull_and_apply_settings_for_generation( + account: &AccountSession, + relay_url: &str, + generation: Option, +) -> Result { let client = AccountClient::new(); let Some(blob) = client .fetch_settings_with_version(relay_url, account) @@ -280,7 +362,10 @@ pub async fn pull_and_apply_settings(account: &AccountSession, relay_url: &str) else { return Ok(false); }; - apply_settings_blob(account, &blob, false).await + if generation.is_some_and(|value| !is_account_context_current(value)) { + return Err(anyhow!("account context changed during settings pull")); + } + apply_settings_blob_for_generation(account, &blob, false, generation).await } async fn account_context() -> Result { @@ -296,11 +381,17 @@ async fn push_from_loop() { debug!("Settings sync: push paused by host app"); return; } - let (account, relay_url) = match account_context().await { + let (account, relay_url, generation) = match account_context().await { Ok(ctx) => ctx, Err(_) => return, // logged out — silently skip }; - if let Err(e) = push_settings_now(&account, &relay_url).await { + if !is_account_context_current(generation) { + return; + } + if let Err(e) = push_settings_now_for_generation(&account, &relay_url, Some(generation)).await { + if !is_account_context_current(generation) { + return; + } note_relay_error(&e, "push"); } } @@ -314,11 +405,19 @@ async fn pull_from_loop() { debug!("Settings sync: pull skipped while an upload/apply is in flight"); return; } - let (account, relay_url) = match account_context().await { + let (account, relay_url, generation) = match account_context().await { Ok(ctx) => ctx, Err(_) => return, // logged out — silently skip }; - if let Err(e) = pull_and_apply_settings(&account, &relay_url).await { + if !is_account_context_current(generation) { + return; + } + if let Err(e) = + pull_and_apply_settings_for_generation(&account, &relay_url, Some(generation)).await + { + if !is_account_context_current(generation) { + return; + } note_relay_error(&e, "pull"); } } diff --git a/src/crates/services/page-function-runtime/src/lib.rs b/src/crates/services/page-function-runtime/src/lib.rs index 4101f019bb..2eab200978 100644 --- a/src/crates/services/page-function-runtime/src/lib.rs +++ b/src/crates/services/page-function-runtime/src/lib.rs @@ -7,6 +7,7 @@ //! host bindings are synchronous. use std::collections::HashMap; +use std::sync::atomic::{AtomicBool, Ordering}; use std::sync::{Arc, Mutex}; use std::time::{Duration, Instant}; @@ -97,19 +98,34 @@ pub fn run_fetch( timeout: Duration, ) -> Result { let started = Instant::now(); + let interrupted = Arc::new(AtomicBool::new(false)); let runtime = Runtime::new().map_err(|e| PageFunctionError::Init(e.to_string()))?; runtime.set_memory_limit(16 * 1024 * 1024); runtime.set_max_stack_size(256 * 1024); + let interrupt_flag = Arc::clone(&interrupted); + runtime.set_interrupt_handler(Some(Box::new(move || { + let timed_out = started.elapsed() >= timeout; + if timed_out { + interrupt_flag.store(true, Ordering::Relaxed); + } + timed_out + }))); let context = Context::full(&runtime).map_err(|e| PageFunctionError::Init(e.to_string()))?; let json_out: String = context.with(|ctx| { - if started.elapsed() > timeout { + if started.elapsed() >= timeout { return Err(PageFunctionError::Timeout(timeout)); } - ctx.eval::<(), _>(worker_source) - .map_err(|e| PageFunctionError::Eval(format!("{e}")))?; + ctx.eval::<(), _>(worker_source).map_err(|e| { + execution_error( + &interrupted, + started, + timeout, + PageFunctionError::Eval(format!("{e}")), + ) + })?; // Ensure fetch exists before wrapping. let globals = ctx.globals(); @@ -150,7 +166,14 @@ pub fn run_fetch( }; "#, ) - .map_err(|e| PageFunctionError::Eval(format!("wrap fetch: {e}")))?; + .map_err(|e| { + execution_error( + &interrupted, + started, + timeout, + PageFunctionError::Eval(format!("wrap fetch: {e}")), + ) + })?; let invoke: Function = globals .get("__bitfun_invoke") @@ -160,17 +183,29 @@ pub fn run_fetch( let req_obj = request_to_object(&ctx, request)?; let maybe: MaybePromise = invoke .call((req_obj, env)) - .map_err(|e| PageFunctionError::Handler(format!("{e}")))?; + .map_err(|e| { + execution_error( + &interrupted, + started, + timeout, + PageFunctionError::Handler(format!("{e}")), + ) + })?; // Drive QuickJS microtasks until the (maybe) promise settles, respecting timeout. let out = loop { - if started.elapsed() > timeout { + if started.elapsed() >= timeout { return Err(PageFunctionError::Timeout(timeout)); } match maybe.result::() { Some(Ok(s)) => break s, Some(Err(e)) => { - return Err(PageFunctionError::Handler(format!("async fetch failed: {e}"))); + return Err(execution_error( + &interrupted, + started, + timeout, + PageFunctionError::Handler(format!("async fetch failed: {e}")), + )); } None => { if !ctx.execute_pending_job() { @@ -190,6 +225,19 @@ pub fn run_fetch( .map_err(|e| PageFunctionError::Handler(format!("invalid fetch response JSON: {e}"))) } +fn execution_error( + interrupted: &AtomicBool, + started: Instant, + timeout: Duration, + fallback: PageFunctionError, +) -> PageFunctionError { + if interrupted.load(Ordering::Relaxed) || started.elapsed() >= timeout { + PageFunctionError::Timeout(timeout) + } else { + fallback + } +} + fn build_env_object<'js>( ctx: &Ctx<'js>, host: Arc, @@ -209,16 +257,22 @@ fn build_env_object<'js>( KV: { get: function(k) { var r = JSON.parse(hostCall("kv_get", String(k), "", "")); + if (r.error) throw new Error(r.error); return r.v; }, put: function(k, v) { - JSON.parse(hostCall("kv_put", String(k), String(v), "")); + var r = JSON.parse(hostCall("kv_put", String(k), String(v), "")); + if (!r.ok) throw new Error(r.error || "KV put failed"); }, delete: function(k) { - return JSON.parse(hostCall("kv_delete", String(k), "", "")).ok; + var r = JSON.parse(hostCall("kv_delete", String(k), "", "")); + if (r.error) throw new Error(r.error); + return r.ok; }, list: function() { - return JSON.parse(hostCall("kv_list", "", "", "")).keys; + var r = JSON.parse(hostCall("kv_list", "", "", "")); + if (r.error) throw new Error(r.error); + return r.keys; } }, DB: { @@ -231,14 +285,19 @@ fn build_env_object<'js>( }, BLOBS: { put: function(id, contentType, dataB64) { - return JSON.parse(hostCall("blob_put", String(id), String(contentType||"application/octet-stream"), String(dataB64))).ok; + var r = JSON.parse(hostCall("blob_put", String(id), String(contentType||"application/octet-stream"), String(dataB64))); + if (!r.ok) throw new Error(r.error || "Blob put failed"); + return true; }, get: function(id) { var r = JSON.parse(hostCall("blob_get", String(id), "", "")); + if (r.error) throw new Error(r.error); return r.found ? { contentType: r.contentType, data: r.data } : null; }, delete: function(id) { - return JSON.parse(hostCall("blob_delete", String(id), "", "")).ok; + var r = JSON.parse(hostCall("blob_delete", String(id), "", "")); + if (r.error) throw new Error(r.error); + return r.ok; } }, ASSETS: { @@ -273,14 +332,14 @@ fn host_call(host: &dyn PageHost, op: &str, a: &str, b: &str, c: &str) -> String }, "kv_delete" => match host.kv_delete(a) { Ok(ok) => format!(r#"{{"ok":{ok}}}"#), - Err(_) => r#"{"ok":false}"#.into(), + Err(e) => format!(r#"{{"ok":false,"error":{}}}"#, json_str(&e)), }, "kv_list" => match host.kv_list() { Ok(keys) => format!( r#"{{"keys":{}}}"#, serde_json::to_string(&keys).unwrap_or_else(|_| "[]".into()) ), - Err(_) => r#"{"keys":[]}"#.into(), + Err(e) => format!(r#"{{"keys":[],"error":{}}}"#, json_str(&e)), }, "db_execute" => host .db_execute(a, b) @@ -303,7 +362,7 @@ fn host_call(host: &dyn PageHost, op: &str, a: &str, b: &str, c: &str) -> String }, "blob_delete" => match host.blob_delete(a) { Ok(ok) => format!(r#"{{"ok":{ok}}}"#), - Err(_) => r#"{"ok":false}"#.into(), + Err(e) => format!(r#"{{"ok":false,"error":{}}}"#, json_str(&e)), }, "assets_fetch" => match host.assets_get(a) { Ok(Some((ct, bytes))) => { @@ -497,4 +556,46 @@ mod tests { Some("application/json") ); } + + #[test] + fn top_level_infinite_loop_is_interrupted() { + let timeout = Duration::from_millis(25); + let started = Instant::now(); + let err = run_fetch( + "while (true) {}", + &test_request(), + Arc::new(MemoryPageHost::default()), + timeout, + ) + .unwrap_err(); + + assert!(matches!(err, PageFunctionError::Timeout(value) if value == timeout)); + assert!(started.elapsed() < Duration::from_secs(1)); + } + + #[test] + fn fetch_handler_infinite_loop_is_interrupted() { + let timeout = Duration::from_millis(25); + let started = Instant::now(); + let err = run_fetch( + "function fetch() { while (true) {} }", + &test_request(), + Arc::new(MemoryPageHost::default()), + timeout, + ) + .unwrap_err(); + + assert!(matches!(err, PageFunctionError::Timeout(value) if value == timeout)); + assert!(started.elapsed() < Duration::from_secs(1)); + } + + fn test_request() -> FetchRequest { + FetchRequest { + method: "GET".into(), + url: "https://example/p/u/s/".into(), + path: "/".into(), + headers: HashMap::new(), + body: None, + } + } } diff --git a/src/crates/services/relay-service/Cargo.toml b/src/crates/services/relay-service/Cargo.toml index 5134e90459..3d6b29fbee 100644 --- a/src/crates/services/relay-service/Cargo.toml +++ b/src/crates/services/relay-service/Cargo.toml @@ -34,7 +34,7 @@ libsqlite3-sys = { version = "0.30", features = ["bundled"] } argon2 = "0.5" aes-gcm = "0.10" bitfun-page-function-runtime = { path = "../page-function-runtime" } -rusqlite = { version = "0.32", features = ["bundled"] } +rusqlite = { version = "0.32", features = ["bundled", "extra_check", "hooks", "limits"] } [dev-dependencies] tower = { version = "0.5", features = ["util"] } diff --git a/src/crates/services/relay-service/src/db.rs b/src/crates/services/relay-service/src/db.rs index ef1b44ac9e..180d6960d8 100644 --- a/src/crates/services/relay-service/src/db.rs +++ b/src/crates/services/relay-service/src/db.rs @@ -13,6 +13,13 @@ use std::time::Duration; pub type DbPool = Pool; +pub const MAX_PAGE_KV_KEY_BYTES: usize = 256; +pub const MAX_PAGE_KV_VALUE_BYTES: usize = 64 * 1024; +pub const MAX_PAGE_KV_ENTRIES: i64 = 1_024; +pub const MAX_USER_KV_ENTRIES: i64 = 10_000; +pub const MAX_PAGE_KV_BYTES: i64 = 5 * 1024 * 1024; +pub const MAX_USER_KV_BYTES: i64 = 25 * 1024 * 1024; + const SCHEMA: &str = r#" CREATE TABLE IF NOT EXISTS users ( user_id TEXT PRIMARY KEY, @@ -1153,32 +1160,57 @@ impl PageRow { Ok(()) } + pub async fn clear_deployed_version(pool: &DbPool, user_id: &str, slug: &str) -> Result { + let now = Utc::now().timestamp(); + let result = sqlx::query( + "UPDATE pages SET deployed_version_id = NULL, updated_at = ? \ + WHERE user_id = ? AND slug = ?", + ) + .bind(now) + .bind(user_id) + .bind(slug) + .execute(pool) + .await + .map_err(|e| anyhow!("clear deployed version: {e}"))?; + Ok(result.rows_affected() > 0) + } + pub async fn delete(pool: &DbPool, user_id: &str, slug: &str) -> Result { + let mut tx = pool + .begin() + .await + .map_err(|e| anyhow!("begin page delete transaction: {e}"))?; sqlx::query("DELETE FROM page_kv WHERE user_id = ? AND slug = ?") .bind(user_id) .bind(slug) - .execute(pool) + .execute(&mut *tx) .await .map_err(|e| anyhow!("delete page_kv: {e}"))?; sqlx::query("DELETE FROM page_blobs WHERE user_id = ? AND slug = ?") .bind(user_id) .bind(slug) - .execute(pool) + .execute(&mut *tx) .await .map_err(|e| anyhow!("delete page_blobs: {e}"))?; sqlx::query("DELETE FROM page_versions WHERE user_id = ? AND slug = ?") .bind(user_id) .bind(slug) - .execute(pool) + .execute(&mut *tx) .await .map_err(|e| anyhow!("delete page_versions: {e}"))?; let result = sqlx::query("DELETE FROM pages WHERE user_id = ? AND slug = ?") .bind(user_id) .bind(slug) - .execute(pool) + .execute(&mut *tx) .await .map_err(|e| anyhow!("delete page: {e}"))?; - Ok(result.rows_affected() > 0) + if result.rows_affected() == 0 { + return Ok(false); + } + tx.commit() + .await + .map_err(|e| anyhow!("commit page delete transaction: {e}"))?; + Ok(true) } /// Resolve a page by public URL components `(username, slug)`. @@ -1311,12 +1343,24 @@ impl PageVersionRow { pub mod page_kv { use super::*; + fn validate_key(key: &str) -> Result<()> { + if key.is_empty() || key.len() > MAX_PAGE_KV_KEY_BYTES || key.chars().any(char::is_control) + { + return Err(anyhow!( + "page KV key must be non-empty, control-free, and at most {} bytes", + MAX_PAGE_KV_KEY_BYTES + )); + } + Ok(()) + } + pub async fn get( pool: &DbPool, user_id: &str, slug: &str, key: &str, ) -> Result> { + validate_key(key)?; let row: Option<(String,)> = sqlx::query_as("SELECT value FROM page_kv WHERE user_id = ? AND slug = ? AND key = ?") .bind(user_id) @@ -1335,6 +1379,53 @@ pub mod page_kv { key: &str, value: &str, ) -> Result<()> { + validate_key(key)?; + if value.len() > MAX_PAGE_KV_VALUE_BYTES { + return Err(anyhow!( + "page KV value exceeds the {} byte operation limit", + MAX_PAGE_KV_VALUE_BYTES + )); + } + + let mut tx = pool + .begin() + .await + .map_err(|e| anyhow!("page_kv begin quota transaction: {e}"))?; + let old_bytes: Option<(i64,)> = sqlx::query_as( + "SELECT length(CAST(key AS BLOB)) + length(CAST(value AS BLOB)) \ + FROM page_kv WHERE user_id = ? AND slug = ? AND key = ?", + ) + .bind(user_id) + .bind(slug) + .bind(key) + .fetch_optional(&mut *tx) + .await + .map_err(|e| anyhow!("page_kv read existing size: {e}"))?; + let page_usage: (i64, i64) = sqlx::query_as( + "SELECT COUNT(*), COALESCE(SUM(length(CAST(key AS BLOB)) + \ + length(CAST(value AS BLOB))), 0) FROM page_kv WHERE user_id = ? AND slug = ?", + ) + .bind(user_id) + .bind(slug) + .fetch_one(&mut *tx) + .await + .map_err(|e| anyhow!("page_kv read page quota: {e}"))?; + let user_usage: (i64, i64) = sqlx::query_as( + "SELECT COUNT(*), COALESCE(SUM(length(CAST(key AS BLOB)) + \ + length(CAST(value AS BLOB))), 0) FROM page_kv WHERE user_id = ?", + ) + .bind(user_id) + .fetch_one(&mut *tx) + .await + .map_err(|e| anyhow!("page_kv read account quota: {e}"))?; + enforce_quota( + page_usage, + user_usage, + old_bytes.map_or(0, |row| row.0), + key.len().saturating_add(value.len()) as i64, + old_bytes.is_none(), + )?; + let now = Utc::now().timestamp(); sqlx::query( "INSERT INTO page_kv (user_id, slug, key, value, updated_at) VALUES (?, ?, ?, ?, ?) \ @@ -1345,13 +1436,17 @@ pub mod page_kv { .bind(key) .bind(value) .bind(now) - .execute(pool) + .execute(&mut *tx) .await .map_err(|e| anyhow!("page_kv put: {e}"))?; + tx.commit() + .await + .map_err(|e| anyhow!("page_kv commit: {e}"))?; Ok(()) } pub async fn delete(pool: &DbPool, user_id: &str, slug: &str, key: &str) -> Result { + validate_key(key)?; let result = sqlx::query("DELETE FROM page_kv WHERE user_id = ? AND slug = ? AND key = ?") .bind(user_id) .bind(slug) @@ -1372,6 +1467,78 @@ pub mod page_kv { .map_err(|e| anyhow!("page_kv list: {e}"))?; Ok(rows.into_iter().map(|r| r.0).collect()) } + + fn enforce_quota( + page_usage: (i64, i64), + user_usage: (i64, i64), + replaced_bytes: i64, + added_bytes: i64, + is_new: bool, + ) -> Result<()> { + let added_entries = i64::from(is_new); + if page_usage.0.saturating_add(added_entries) > MAX_PAGE_KV_ENTRIES { + return Err(anyhow!("page KV entry quota exceeded")); + } + if user_usage.0.saturating_add(added_entries) > MAX_USER_KV_ENTRIES { + return Err(anyhow!("account KV entry quota exceeded")); + } + if page_usage + .1 + .saturating_sub(replaced_bytes) + .saturating_add(added_bytes) + > MAX_PAGE_KV_BYTES + { + return Err(anyhow!("page KV byte quota exceeded")); + } + if user_usage + .1 + .saturating_sub(replaced_bytes) + .saturating_add(added_bytes) + > MAX_USER_KV_BYTES + { + return Err(anyhow!("account KV byte quota exceeded")); + } + Ok(()) + } + + #[cfg(test)] + mod quota_tests { + use super::*; + + #[test] + fn quota_projection_handles_insert_and_overwrite() { + assert!(enforce_quota( + (MAX_PAGE_KV_ENTRIES, 100), + (MAX_PAGE_KV_ENTRIES, 100), + 10, + 10, + false, + ) + .is_ok()); + assert!(enforce_quota( + (MAX_PAGE_KV_ENTRIES, 100), + (MAX_PAGE_KV_ENTRIES, 100), + 0, + 1, + true, + ) + .is_err()); + assert!( + enforce_quota((1, MAX_PAGE_KV_BYTES), (1, MAX_PAGE_KV_BYTES), 1, 2, false,) + .is_err() + ); + assert!(enforce_quota((1, 1), (1, MAX_USER_KV_BYTES), 0, 1, false,).is_err()); + } + + #[tokio::test] + async fn put_rejects_oversized_single_values_before_writing() { + let pool = connect(":memory:").await.unwrap(); + let value = "x".repeat(MAX_PAGE_KV_VALUE_BYTES + 1); + let err = put(&pool, "u1", "site", "key", &value).await.unwrap_err(); + assert!(err.to_string().contains("operation limit")); + assert!(get(&pool, "u1", "site", "key").await.unwrap().is_none()); + } + } } /// Legacy single-room asset key (pre-versioning). Used only for one-time migration. diff --git a/src/crates/services/relay-service/src/lib.rs b/src/crates/services/relay-service/src/lib.rs index eb3aaedd2d..a3e25dea7a 100644 --- a/src/crates/services/relay-service/src/lib.rs +++ b/src/crates/services/relay-service/src/lib.rs @@ -12,6 +12,7 @@ pub mod admin; pub mod db; pub mod page_data; +pub mod page_execution; pub mod relay; pub mod routes; @@ -1050,6 +1051,9 @@ pub fn build_relay_router_with_page_data_and_origins( asset_store, db, page_data, + page_access_manager: Arc::new(routes::pages::PageAccessManager::new()), + page_upload_manager: Arc::new(routes::pages::PageUploadManager::new()), + page_execution_guard: Arc::new(crate::page_execution::PageExecutionGuard::new()), login_rate_limiter: std::sync::Arc::new(crate::routes::auth::LoginRateLimiter::new()), device_manager: crate::relay::DeviceManager::new(), cors_allow_origins: Arc::new(cors_allow_origins.clone()), diff --git a/src/crates/services/relay-service/src/page_data.rs b/src/crates/services/relay-service/src/page_data.rs index 4db0c52452..195137e3ab 100644 --- a/src/crates/services/relay-service/src/page_data.rs +++ b/src/crates/services/relay-service/src/page_data.rs @@ -2,27 +2,52 @@ //! Survives version deploy/rollback; separate from immutable version assets. use std::path::{Path, PathBuf}; -use std::sync::Arc; +use std::sync::{Arc, Mutex}; +use std::time::{Duration, Instant}; use anyhow::{anyhow, Result}; use base64::{engine::general_purpose::STANDARD as B64, Engine}; use bitfun_page_function_runtime::{PageHost, PageMeta}; -use chrono::Utc; +use dashmap::DashMap; +use rusqlite::hooks::{AuthAction, AuthContext, Authorization}; +use rusqlite::limits::Limit; use tokio::runtime::Handle; use crate::db::{page_kv, DbPool}; +pub const MAX_BLOB_ID_BYTES: usize = 128; +pub const MAX_BLOB_BYTES: usize = 4 * 1024 * 1024; +pub const MAX_BLOB_FILES_PER_PAGE: u64 = 2_048; +pub const MAX_BLOB_FILES_PER_USER: u64 = 10_000; +pub const MAX_MUTABLE_BYTES_PER_PAGE: u64 = 64 * 1024 * 1024; +pub const MAX_MUTABLE_BYTES_PER_USER: u64 = 256 * 1024 * 1024; +pub const MAX_PAGE_DB_BYTES: u64 = 20 * 1024 * 1024; +pub const MAX_DB_SQL_BYTES: usize = 64 * 1024; +pub const MAX_DB_PARAMS_BYTES: usize = 256 * 1024; +pub const MAX_DB_PARAMS: usize = 256; +pub const MAX_DB_QUERY_ROWS: usize = 1_000; +pub const MAX_DB_QUERY_BYTES: usize = 2 * 1024 * 1024; +const MAX_DB_VALUE_BYTES: i32 = 2 * 1024 * 1024; +#[cfg(not(test))] +const MAX_DB_OPERATION_TIME: Duration = Duration::from_secs(1); +#[cfg(test)] +const MAX_DB_OPERATION_TIME: Duration = Duration::from_millis(50); + /// Root directory for page-data (`{base}/{user_id}/{slug}/...`). #[derive(Clone)] pub struct PageDataStore { base_dir: PathBuf, + user_mutation_locks: Arc>>>, } impl PageDataStore { pub fn new(base_dir: impl Into) -> Self { let base_dir = base_dir.into(); let _ = std::fs::create_dir_all(&base_dir); - Self { base_dir } + Self { + base_dir, + user_mutation_locks: Arc::new(DashMap::new()), + } } pub fn base_dir(&self) -> &Path { @@ -42,18 +67,41 @@ impl PageDataStore { } pub fn cleanup_page(&self, user_id: &str, slug: &str) { - let dir = self.page_dir(user_id, slug); - if dir.exists() { + let lock = self.user_mutation_lock(user_id); + let Ok(_guard) = lock.lock() else { + return; + }; + if let Ok(Some(dir)) = self.existing_page_dir(user_id, slug) { let _ = std::fs::remove_dir_all(&dir); } } fn ensure_page_dir(&self, user_id: &str, slug: &str) -> Result { + ensure_directory(&self.base_dir)?; + ensure_directory(&self.base_dir.join(user_id))?; let dir = self.page_dir(user_id, slug); - std::fs::create_dir_all(&dir).map_err(|e| anyhow!("create page-data dir: {e}"))?; + ensure_directory(&dir)?; Ok(dir) } + fn existing_page_dir(&self, user_id: &str, slug: &str) -> Result> { + if existing_directory(&self.base_dir)?.is_none() + || existing_directory(&self.base_dir.join(user_id))?.is_none() + { + return Ok(None); + } + existing_directory(&self.page_dir(user_id, slug)) + } + + fn user_mutation_lock(&self, user_id: &str) -> Arc> { + Arc::clone( + self.user_mutation_locks + .entry(user_id.to_string()) + .or_insert_with(|| Arc::new(Mutex::new(()))) + .value(), + ) + } + pub fn blob_put( &self, user_id: &str, @@ -62,17 +110,44 @@ impl PageDataStore { content_type: &str, data: &[u8], ) -> Result<()> { - if blob_id.contains("..") || blob_id.contains('/') || blob_id.contains('\\') { - return Err(anyhow!("invalid blob id")); + validate_blob_id(blob_id)?; + validate_content_type(content_type)?; + if data.len() > MAX_BLOB_BYTES { + return Err(anyhow!( + "blob exceeds the {} byte operation limit", + MAX_BLOB_BYTES + )); } + let lock = self.user_mutation_lock(user_id); + let _guard = lock + .lock() + .map_err(|_| anyhow!("page-data mutation lock poisoned"))?; + self.ensure_page_dir(user_id, slug)?; let blobs = self.blobs_dir(user_id, slug); - std::fs::create_dir_all(&blobs).map_err(|e| anyhow!("create blobs dir: {e}"))?; + ensure_directory(&blobs)?; + ensure_directory(&blobs.join(".metadata"))?; let path = blobs.join(blob_id); + let meta = blob_meta_path(&blobs, blob_id); + + if !path.exists() { + let page_blob_files = direct_file_count(&blobs)?; + let user_blob_files = user_blob_file_count(&self.base_dir.join(user_id))?; + if page_blob_files >= MAX_BLOB_FILES_PER_PAGE { + return Err(anyhow!("page blob file quota exceeded")); + } + if user_blob_files >= MAX_BLOB_FILES_PER_USER { + return Err(anyhow!("account blob file quota exceeded")); + } + } + + let replaced_bytes = file_len(&path).saturating_add(file_len(&meta)); + let added_bytes = (data.len() as u64).saturating_add(content_type.len() as u64); + let page_bytes = directory_size(&self.page_dir(user_id, slug))?; + let user_bytes = directory_size(&self.base_dir.join(user_id))?; + enforce_storage_quota(page_bytes, user_bytes, replaced_bytes, added_bytes)?; + std::fs::write(&path, data).map_err(|e| anyhow!("write blob: {e}"))?; - let meta = blobs.join(format!("{blob_id}.meta")); std::fs::write(&meta, content_type).map_err(|e| anyhow!("write blob meta: {e}"))?; - let _ = content_type; - let _ = Utc::now(); Ok(()) } @@ -82,33 +157,59 @@ impl PageDataStore { slug: &str, blob_id: &str, ) -> Result)>> { - if blob_id.contains("..") || blob_id.contains('/') || blob_id.contains('\\') { - return Err(anyhow!("invalid blob id")); + validate_blob_id(blob_id)?; + let lock = self.user_mutation_lock(user_id); + let _guard = lock + .lock() + .map_err(|_| anyhow!("page-data mutation lock poisoned"))?; + if self.existing_page_dir(user_id, slug)?.is_none() { + return Ok(None); } - let path = self.blobs_dir(user_id, slug).join(blob_id); - if !path.is_file() { + let Some(blobs) = existing_directory(&self.blobs_dir(user_id, slug))? else { return Ok(None); + }; + let path = blobs.join(blob_id); + let metadata = match std::fs::symlink_metadata(&path) { + Ok(metadata) => metadata, + Err(error) if error.kind() == std::io::ErrorKind::NotFound => return Ok(None), + Err(error) => return Err(anyhow!("read blob metadata: {error}")), + }; + if metadata.file_type().is_symlink() { + return Err(anyhow!("blob path must not be a symlink")); + } + if !metadata.is_file() { + return Ok(None); + } + let size = metadata.len(); + if size > MAX_BLOB_BYTES as u64 { + return Err(anyhow!("stored blob exceeds the read limit")); } let data = std::fs::read(&path).map_err(|e| anyhow!("read blob: {e}"))?; - let meta = self - .blobs_dir(user_id, slug) - .join(format!("{blob_id}.meta")); - let content_type = std::fs::read_to_string(&meta) - .unwrap_or_else(|_| "application/octet-stream".to_string()); + let meta = blob_meta_path(&blobs, blob_id); + let content_type = read_small_metadata(&meta) + .or_else(|| read_small_metadata(&blobs.join(format!("{blob_id}.meta")))) + .unwrap_or_else(|| "application/octet-stream".to_string()); Ok(Some((content_type, data))) } pub fn blob_delete(&self, user_id: &str, slug: &str, blob_id: &str) -> Result { - if blob_id.contains("..") || blob_id.contains('/') || blob_id.contains('\\') { - return Err(anyhow!("invalid blob id")); + validate_blob_id(blob_id)?; + let lock = self.user_mutation_lock(user_id); + let _guard = lock + .lock() + .map_err(|_| anyhow!("page-data mutation lock poisoned"))?; + if self.existing_page_dir(user_id, slug)?.is_none() { + return Ok(false); } - let path = self.blobs_dir(user_id, slug).join(blob_id); - let meta = self - .blobs_dir(user_id, slug) - .join(format!("{blob_id}.meta")); + let Some(blobs) = existing_directory(&self.blobs_dir(user_id, slug))? else { + return Ok(false); + }; + let path = blobs.join(blob_id); + let meta = blob_meta_path(&blobs, blob_id); let existed = path.exists(); let _ = std::fs::remove_file(&path); let _ = std::fs::remove_file(&meta); + let _ = std::fs::remove_file(blobs.join(format!("{blob_id}.meta"))); Ok(existed) } @@ -119,15 +220,36 @@ impl PageDataStore { sql: &str, params_json: &str, ) -> Result { + validate_db_input(sql, params_json)?; + let params = parse_db_params(params_json)?; + let lock = self.user_mutation_lock(user_id); + let _guard = lock + .lock() + .map_err(|_| anyhow!("page-data mutation lock poisoned"))?; self.ensure_page_dir(user_id, slug)?; let path = self.db_path(user_id, slug); - let conn = rusqlite::Connection::open(&path).map_err(|e| anyhow!("open page db: {e}"))?; - let params: Vec = - serde_json::from_str(params_json).unwrap_or_else(|_| Vec::new()); - let mut stmt = conn.prepare(sql).map_err(|e| anyhow!("prepare: {e}"))?; - let changes = stmt - .execute(rusqlite::params_from_iter(params.iter().map(json_to_sql))) - .map_err(|e| anyhow!("execute: {e}"))?; + ensure_regular_file_or_missing(&path)?; + let mut conn = + rusqlite::Connection::open(&path).map_err(|e| anyhow!("open page db: {e}"))?; + configure_connection(&conn); + configure_database_quota(self, &conn, user_id, slug, &path)?; + + let tx = conn + .transaction() + .map_err(|e| anyhow!("begin transaction: {e}"))?; + tx.authorizer(Some(authorize_execute)); + let result = tx.execute( + sql, + rusqlite::params_from_iter(params.iter().map(json_to_sql)), + ); + tx.authorizer(None::) -> Authorization>); + let changes = result.map_err(|e| anyhow!("execute: {e}"))?; + let logical_bytes = database_logical_bytes(&tx)?; + if logical_bytes > MAX_PAGE_DB_BYTES { + return Err(anyhow!("page database quota exceeded")); + } + tx.commit().map_err(|e| anyhow!("commit: {e}"))?; + enforce_current_storage_quota(self, user_id, slug)?; Ok(serde_json::json!({ "ok": true, "changes": changes }).to_string()) } @@ -138,21 +260,38 @@ impl PageDataStore { sql: &str, params_json: &str, ) -> Result { + validate_db_input(sql, params_json)?; + let params = parse_db_params(params_json)?; + let lock = self.user_mutation_lock(user_id); + let _guard = lock + .lock() + .map_err(|_| anyhow!("page-data mutation lock poisoned"))?; self.ensure_page_dir(user_id, slug)?; let path = self.db_path(user_id, slug); + ensure_regular_file_or_missing(&path)?; let conn = rusqlite::Connection::open(&path).map_err(|e| anyhow!("open page db: {e}"))?; - let params: Vec = - serde_json::from_str(params_json).unwrap_or_else(|_| Vec::new()); + configure_connection(&conn); + conn.authorizer(Some(authorize_query)); let mut stmt = conn.prepare(sql).map_err(|e| anyhow!("prepare: {e}"))?; + if !stmt.readonly() { + return Err(anyhow!("DB.query only accepts read-only statements")); + } let col_count = stmt.column_count(); let col_names: Vec = (0..col_count) .map(|i| stmt.column_name(i).unwrap_or("?").to_string()) .collect(); let mut rows_out = Vec::new(); + let mut response_bytes = 0usize; let mut rows = stmt .query(rusqlite::params_from_iter(params.iter().map(json_to_sql))) .map_err(|e| anyhow!("query: {e}"))?; while let Some(row) = rows.next().map_err(|e| anyhow!("row: {e}"))? { + if rows_out.len() >= MAX_DB_QUERY_ROWS { + return Err(anyhow!( + "query exceeds the {} row result limit", + MAX_DB_QUERY_ROWS + )); + } let mut obj = serde_json::Map::new(); for (i, name) in col_names.iter().enumerate() { let val = match row.get_ref(i).map_err(|e| anyhow!("get: {e}"))? { @@ -166,12 +305,368 @@ impl PageDataStore { }; obj.insert(name.clone(), val); } - rows_out.push(serde_json::Value::Object(obj)); + let value = serde_json::Value::Object(obj); + let row_bytes = serde_json::to_vec(&value) + .map_err(|e| anyhow!("serialize query row: {e}"))? + .len(); + if response_bytes.saturating_add(row_bytes) > MAX_DB_QUERY_BYTES { + return Err(anyhow!( + "query exceeds the {} byte result limit", + MAX_DB_QUERY_BYTES + )); + } + response_bytes = response_bytes.saturating_add(row_bytes); + rows_out.push(value); } Ok(serde_json::json!({ "ok": true, "rows": rows_out }).to_string()) } } +fn ensure_directory(path: &Path) -> Result<()> { + match std::fs::symlink_metadata(path) { + Ok(metadata) if metadata.file_type().is_symlink() => Err(anyhow!( + "page-data directory must not be a symlink: {}", + path.display() + )), + Ok(metadata) if metadata.is_dir() => Ok(()), + Ok(_) => Err(anyhow!( + "page-data directory path is not a directory: {}", + path.display() + )), + Err(error) if error.kind() == std::io::ErrorKind::NotFound => std::fs::create_dir(path) + .map_err(|e| anyhow!("create page-data directory {}: {e}", path.display())), + Err(error) => Err(anyhow!( + "inspect page-data directory {}: {error}", + path.display() + )), + } +} + +fn existing_directory(path: &Path) -> Result> { + match std::fs::symlink_metadata(path) { + Ok(metadata) if metadata.file_type().is_symlink() => Err(anyhow!( + "page-data directory must not be a symlink: {}", + path.display() + )), + Ok(metadata) if metadata.is_dir() => Ok(Some(path.to_path_buf())), + Ok(_) => Err(anyhow!( + "page-data directory path is not a directory: {}", + path.display() + )), + Err(error) if error.kind() == std::io::ErrorKind::NotFound => Ok(None), + Err(error) => Err(anyhow!( + "inspect page-data directory {}: {error}", + path.display() + )), + } +} + +fn ensure_regular_file_or_missing(path: &Path) -> Result<()> { + match std::fs::symlink_metadata(path) { + Ok(metadata) if metadata.file_type().is_symlink() => Err(anyhow!( + "page-data file must not be a symlink: {}", + path.display() + )), + Ok(metadata) if metadata.is_file() => Ok(()), + Ok(_) => Err(anyhow!( + "page-data file path is not a regular file: {}", + path.display() + )), + Err(error) if error.kind() == std::io::ErrorKind::NotFound => Ok(()), + Err(error) => Err(anyhow!( + "inspect page-data file {}: {error}", + path.display() + )), + } +} + +fn validate_blob_id(blob_id: &str) -> Result<()> { + if blob_id.is_empty() + || blob_id.len() > MAX_BLOB_ID_BYTES + || matches!(blob_id, "." | ".." | ".metadata") + || blob_id.contains('/') + || blob_id.contains('\\') + || blob_id.chars().any(char::is_control) + { + return Err(anyhow!("invalid blob id")); + } + Ok(()) +} + +fn validate_content_type(content_type: &str) -> Result<()> { + if content_type.is_empty() + || content_type.len() > 256 + || content_type.chars().any(char::is_control) + { + return Err(anyhow!("invalid blob content type")); + } + Ok(()) +} + +fn blob_meta_path(blobs_dir: &Path, blob_id: &str) -> PathBuf { + blobs_dir.join(".metadata").join(blob_id) +} + +fn file_len(path: &Path) -> u64 { + path.metadata().map_or(0, |metadata| metadata.len()) +} + +fn read_small_metadata(path: &Path) -> Option { + if file_len(path) > 256 { + return None; + } + std::fs::read_to_string(path).ok() +} + +fn directory_size(path: &Path) -> Result { + if !path.exists() { + return Ok(0); + } + let metadata = std::fs::symlink_metadata(path) + .map_err(|e| anyhow!("inspect page-data path {}: {e}", path.display()))?; + if metadata.file_type().is_symlink() { + return Err(anyhow!( + "page-data path must not contain symlinks: {}", + path.display() + )); + } + if metadata.is_file() { + return Ok(metadata.len()); + } + if !metadata.is_dir() { + return Ok(0); + } + + let mut total = 0u64; + for entry in std::fs::read_dir(path) + .map_err(|e| anyhow!("read page-data directory {}: {e}", path.display()))? + { + let entry = entry.map_err(|e| anyhow!("read page-data entry: {e}"))?; + total = total + .checked_add(directory_size(&entry.path())?) + .ok_or_else(|| anyhow!("page-data size overflow"))?; + } + Ok(total) +} + +fn direct_file_count(path: &Path) -> Result { + if !path.exists() { + return Ok(0); + } + let mut count = 0u64; + for entry in std::fs::read_dir(path) + .map_err(|e| anyhow!("read blob directory {}: {e}", path.display()))? + { + let entry = entry.map_err(|e| anyhow!("read blob entry: {e}"))?; + let metadata = std::fs::symlink_metadata(entry.path()) + .map_err(|e| anyhow!("inspect blob entry: {e}"))?; + if metadata.file_type().is_symlink() { + return Err(anyhow!("blob directory must not contain symlinks")); + } + if metadata.is_file() { + count = count + .checked_add(1) + .ok_or_else(|| anyhow!("blob file count overflow"))?; + } + } + Ok(count) +} + +fn user_blob_file_count(user_dir: &Path) -> Result { + if !user_dir.exists() { + return Ok(0); + } + let mut count = 0u64; + for page in + std::fs::read_dir(user_dir).map_err(|e| anyhow!("read account page-data directory: {e}"))? + { + let page = page.map_err(|e| anyhow!("read account page-data entry: {e}"))?; + let metadata = std::fs::symlink_metadata(page.path()) + .map_err(|e| anyhow!("inspect account page-data entry: {e}"))?; + if metadata.file_type().is_symlink() { + return Err(anyhow!("page-data directory must not contain symlinks")); + } + if metadata.is_dir() { + count = count + .checked_add(direct_file_count(&page.path().join("blobs"))?) + .ok_or_else(|| anyhow!("account blob file count overflow"))?; + } + } + Ok(count) +} + +fn enforce_storage_quota( + page_bytes: u64, + user_bytes: u64, + replaced_bytes: u64, + added_bytes: u64, +) -> Result<()> { + let projected_page = page_bytes + .saturating_sub(replaced_bytes) + .checked_add(added_bytes) + .ok_or_else(|| anyhow!("page-data size overflow"))?; + if projected_page > MAX_MUTABLE_BYTES_PER_PAGE { + return Err(anyhow!( + "page mutable storage exceeds the {} byte quota", + MAX_MUTABLE_BYTES_PER_PAGE + )); + } + let projected_user = user_bytes + .saturating_sub(replaced_bytes) + .checked_add(added_bytes) + .ok_or_else(|| anyhow!("account page-data size overflow"))?; + if projected_user > MAX_MUTABLE_BYTES_PER_USER { + return Err(anyhow!( + "account mutable storage exceeds the {} byte quota", + MAX_MUTABLE_BYTES_PER_USER + )); + } + Ok(()) +} + +fn enforce_current_storage_quota(store: &PageDataStore, user_id: &str, slug: &str) -> Result<()> { + enforce_storage_quota( + directory_size(&store.page_dir(user_id, slug))?, + directory_size(&store.base_dir.join(user_id))?, + 0, + 0, + ) +} + +fn validate_db_input(sql: &str, params_json: &str) -> Result<()> { + if sql.trim().is_empty() || sql.len() > MAX_DB_SQL_BYTES { + return Err(anyhow!( + "SQL must be non-empty and at most {} bytes", + MAX_DB_SQL_BYTES + )); + } + if params_json.len() > MAX_DB_PARAMS_BYTES { + return Err(anyhow!( + "SQL parameters exceed the {} byte limit", + MAX_DB_PARAMS_BYTES + )); + } + Ok(()) +} + +fn parse_db_params(params_json: &str) -> Result> { + let params: Vec = + serde_json::from_str(params_json).map_err(|e| anyhow!("invalid SQL parameters: {e}"))?; + if params.len() > MAX_DB_PARAMS { + return Err(anyhow!( + "SQL parameters exceed the {} item limit", + MAX_DB_PARAMS + )); + } + Ok(params) +} + +fn configure_connection(conn: &rusqlite::Connection) { + conn.set_limit(Limit::SQLITE_LIMIT_LENGTH, MAX_DB_VALUE_BYTES); + conn.set_limit(Limit::SQLITE_LIMIT_SQL_LENGTH, MAX_DB_SQL_BYTES as i32); + conn.set_limit(Limit::SQLITE_LIMIT_COLUMN, 128); + conn.set_limit(Limit::SQLITE_LIMIT_EXPR_DEPTH, 64); + conn.set_limit(Limit::SQLITE_LIMIT_COMPOUND_SELECT, 16); + conn.set_limit(Limit::SQLITE_LIMIT_FUNCTION_ARG, 32); + conn.set_limit(Limit::SQLITE_LIMIT_ATTACHED, 0); + conn.set_limit(Limit::SQLITE_LIMIT_VARIABLE_NUMBER, MAX_DB_PARAMS as i32); + conn.set_limit(Limit::SQLITE_LIMIT_TRIGGER_DEPTH, 8); + conn.set_limit(Limit::SQLITE_LIMIT_WORKER_THREADS, 0); + let started = Instant::now(); + conn.progress_handler( + 1_000, + Some(move || started.elapsed() >= MAX_DB_OPERATION_TIME), + ); + let _ = conn.busy_timeout(Duration::from_millis(250)); +} + +fn configure_database_quota( + store: &PageDataStore, + conn: &rusqlite::Connection, + user_id: &str, + slug: &str, + db_path: &Path, +) -> Result<()> { + let page_bytes = directory_size(&store.page_dir(user_id, slug))?; + let user_bytes = directory_size(&store.base_dir.join(user_id))?; + enforce_storage_quota(page_bytes, user_bytes, 0, 0)?; + + let current_db_bytes = file_len(db_path); + if current_db_bytes > MAX_PAGE_DB_BYTES { + return Err(anyhow!("page database already exceeds its quota")); + } + let remaining_page = MAX_MUTABLE_BYTES_PER_PAGE.saturating_sub(page_bytes); + let remaining_user = MAX_MUTABLE_BYTES_PER_USER.saturating_sub(user_bytes); + let allowed_db_bytes = + MAX_PAGE_DB_BYTES.min(current_db_bytes.saturating_add(remaining_page.min(remaining_user))); + let page_size: u64 = conn + .query_row("PRAGMA page_size", [], |row| row.get(0)) + .map_err(|e| anyhow!("read database page size: {e}"))?; + let max_pages = (allowed_db_bytes / page_size).max(1); + let _: u64 = conn + .query_row(&format!("PRAGMA max_page_count = {max_pages}"), [], |row| { + row.get(0) + }) + .map_err(|e| anyhow!("set database quota: {e}"))?; + Ok(()) +} + +fn database_logical_bytes(conn: &rusqlite::Connection) -> Result { + let page_count: u64 = conn + .query_row("PRAGMA page_count", [], |row| row.get(0)) + .map_err(|e| anyhow!("read database page count: {e}"))?; + let page_size: u64 = conn + .query_row("PRAGMA page_size", [], |row| row.get(0)) + .map_err(|e| anyhow!("read database page size: {e}"))?; + Ok(page_count.saturating_mul(page_size)) +} + +fn authorize_execute(context: AuthContext<'_>) -> Authorization { + match context.action { + AuthAction::Attach { .. } + | AuthAction::Detach { .. } + | AuthAction::Pragma { .. } + | AuthAction::Transaction { .. } + | AuthAction::Savepoint { .. } + | AuthAction::CreateTempIndex { .. } + | AuthAction::CreateTempTable { .. } + | AuthAction::CreateTempTrigger { .. } + | AuthAction::CreateTempView { .. } + | AuthAction::DropTempIndex { .. } + | AuthAction::DropTempTable { .. } + | AuthAction::DropTempTrigger { .. } + | AuthAction::DropTempView { .. } + | AuthAction::CreateVtable { .. } + | AuthAction::DropVtable { .. } + | AuthAction::Analyze { .. } + | AuthAction::Reindex { .. } + | AuthAction::Unknown { .. } => Authorization::Deny, + AuthAction::Function { function_name } if forbidden_db_function(function_name) => { + Authorization::Deny + } + _ => Authorization::Allow, + } +} + +fn authorize_query(context: AuthContext<'_>) -> Authorization { + match context.action { + AuthAction::Read { .. } | AuthAction::Select | AuthAction::Recursive => { + Authorization::Allow + } + AuthAction::Function { function_name } if !forbidden_db_function(function_name) => { + Authorization::Allow + } + _ => Authorization::Deny, + } +} + +fn forbidden_db_function(function_name: &str) -> bool { + matches!( + function_name.to_ascii_lowercase().as_str(), + "load_extension" | "readfile" | "writefile" | "edit" | "shell" + ) +} + fn json_to_sql(v: &serde_json::Value) -> rusqlite::types::Value { match v { serde_json::Value::Null => rusqlite::types::Value::Null, @@ -232,6 +727,10 @@ impl PageHost for RelayPageHost { } fn kv_put(&self, key: &str, value: &str) -> Result<(), String> { + let lock = self.page_data.user_mutation_lock(&self.user_id); + let _guard = lock + .lock() + .map_err(|_| "page-data mutation lock poisoned".to_string())?; let db = Arc::clone(&self.db); let user_id = self.user_id.clone(); let slug = self.slug.clone(); @@ -268,6 +767,12 @@ impl PageHost for RelayPageHost { } fn blob_put(&self, blob_id: &str, content_type: &str, data_b64: &str) -> Result<(), String> { + const MAX_ENCODED_BLOB_BYTES: usize = MAX_BLOB_BYTES.div_ceil(3) * 4; + if data_b64.len() > MAX_ENCODED_BLOB_BYTES { + return Err(format!( + "encoded blob exceeds the {MAX_ENCODED_BLOB_BYTES} byte operation limit" + )); + } let data = B64.decode(data_b64).map_err(|e| e.to_string())?; self.page_data .blob_put(&self.user_id, &self.slug, blob_id, content_type, &data) @@ -323,3 +828,140 @@ fn mime_from_path(p: &str) -> &'static str { pub fn default_page_data_dir(room_web_dir: &Path) -> PathBuf { room_web_dir.join("page-data") } + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn storage_projection_enforces_page_and_account_quotas() { + assert!(enforce_storage_quota( + MAX_MUTABLE_BYTES_PER_PAGE, + MAX_MUTABLE_BYTES_PER_PAGE, + 10, + 10, + ) + .is_ok()); + assert!(enforce_storage_quota( + MAX_MUTABLE_BYTES_PER_PAGE, + MAX_MUTABLE_BYTES_PER_PAGE, + 0, + 1, + ) + .is_err()); + assert!(enforce_storage_quota(0, MAX_MUTABLE_BYTES_PER_USER, 0, 1,).is_err()); + } + + #[test] + fn blob_limits_are_enforced_and_metadata_names_do_not_collide() { + let temp = tempfile::tempdir().unwrap(); + let store = PageDataStore::new(temp.path()); + let oversized = vec![0u8; MAX_BLOB_BYTES + 1]; + assert!(store + .blob_put( + "u1", + "site", + "too-big", + "application/octet-stream", + &oversized + ) + .unwrap_err() + .to_string() + .contains("operation limit")); + + store + .blob_put("u1", "site", "item", "text/plain", b"first") + .unwrap(); + store + .blob_put( + "u1", + "site", + "item.meta", + "application/octet-stream", + b"second", + ) + .unwrap(); + assert_eq!( + store.blob_get("u1", "site", "item").unwrap().unwrap(), + ("text/plain".to_string(), b"first".to_vec()) + ); + assert_eq!( + store.blob_get("u1", "site", "item.meta").unwrap().unwrap(), + ("application/octet-stream".to_string(), b"second".to_vec()) + ); + } + + #[test] + fn database_denies_file_attachment_and_caps_query_rows() { + let temp = tempfile::tempdir().unwrap(); + let store = PageDataStore::new(temp.path().join("page-data")); + let attached = temp.path().join("outside.sqlite"); + let error = store + .db_execute( + "u1", + "site", + &format!("ATTACH DATABASE '{}' AS outside", attached.display()), + "[]", + ) + .unwrap_err(); + assert!( + error.to_string().contains("not authorized") + || error.to_string().contains("too many attached") + ); + assert!(!attached.exists()); + + let error = store + .db_query( + "u1", + "site", + "WITH RECURSIVE items(n) AS (VALUES(1) UNION ALL SELECT n + 1 FROM items WHERE n <= 1000) SELECT n FROM items", + "[]", + ) + .unwrap_err(); + assert!(error.to_string().contains("row result limit")); + } + + #[test] + fn database_execute_accepts_normal_schema_and_data_changes() { + let temp = tempfile::tempdir().unwrap(); + let store = PageDataStore::new(temp.path()); + store + .db_execute( + "u1", + "site", + "CREATE TABLE notes (id INTEGER PRIMARY KEY, body TEXT)", + "[]", + ) + .unwrap(); + store + .db_execute( + "u1", + "site", + "INSERT INTO notes (body) VALUES (?)", + r#"["hello"]"#, + ) + .unwrap(); + let output = store + .db_query("u1", "site", "SELECT body FROM notes", "[]") + .unwrap(); + let json: serde_json::Value = serde_json::from_str(&output).unwrap(); + assert_eq!(json["rows"][0]["body"], "hello"); + } + + #[test] + fn database_long_running_query_is_interrupted() { + let temp = tempfile::tempdir().unwrap(); + let store = PageDataStore::new(temp.path()); + let started = Instant::now(); + let error = store + .db_query( + "u1", + "site", + "WITH RECURSIVE items(n) AS (VALUES(1) UNION ALL SELECT n + 1 FROM items WHERE n < 1000000000) SELECT sum(n) FROM items", + "[]", + ) + .unwrap_err(); + assert!(error.to_string().contains("interrupted")); + assert!(started.elapsed() < Duration::from_secs(1)); + } +} diff --git a/src/crates/services/relay-service/src/page_execution.rs b/src/crates/services/relay-service/src/page_execution.rs new file mode 100644 index 0000000000..98ddcc32da --- /dev/null +++ b/src/crates/services/relay-service/src/page_execution.rs @@ -0,0 +1,260 @@ +//! Admission control for public Page Function execution. + +use std::collections::{HashMap, VecDeque}; +use std::sync::{Arc, Mutex}; +use std::time::{Duration, Instant}; + +use dashmap::DashMap; +use tokio::sync::{OwnedSemaphorePermit, Semaphore}; + +pub const MAX_PAGE_FUNCTION_REQUEST_BODY_BYTES: usize = 1024 * 1024; +const GLOBAL_CONCURRENCY: usize = 64; +const USER_CONCURRENCY: usize = 16; +const PAGE_CONCURRENCY: usize = 8; +const USER_REQUESTS_PER_WINDOW: usize = 3_000; +const PAGE_REQUESTS_PER_WINDOW: usize = 600; +const RATE_WINDOW: Duration = Duration::from_secs(60); +const MAX_TRACKED_IDENTITIES: usize = 10_000; + +#[derive(Debug, Clone, Copy, Eq, PartialEq)] +pub enum PageExecutionRejection { + Busy, + RateLimited, +} + +#[derive(Clone, Copy)] +struct Limits { + global_concurrency: usize, + user_concurrency: usize, + page_concurrency: usize, + user_requests_per_window: usize, + page_requests_per_window: usize, + rate_window: Duration, +} + +impl Default for Limits { + fn default() -> Self { + Self { + global_concurrency: GLOBAL_CONCURRENCY, + user_concurrency: USER_CONCURRENCY, + page_concurrency: PAGE_CONCURRENCY, + user_requests_per_window: USER_REQUESTS_PER_WINDOW, + page_requests_per_window: PAGE_REQUESTS_PER_WINDOW, + rate_window: RATE_WINDOW, + } + } +} + +/// Process-local guard against one account or page exhausting blocking workers. +pub struct PageExecutionGuard { + limits: Limits, + global: Arc, + users: DashMap>, + pages: DashMap>, + rate_windows: Mutex>>, +} + +impl Default for PageExecutionGuard { + fn default() -> Self { + Self::new() + } +} + +impl PageExecutionGuard { + pub fn new() -> Self { + Self::with_limits(Limits::default()) + } + + fn with_limits(limits: Limits) -> Self { + Self { + global: Arc::new(Semaphore::new(limits.global_concurrency)), + users: DashMap::new(), + pages: DashMap::new(), + rate_windows: Mutex::new(HashMap::new()), + limits, + } + } + + pub fn try_acquire( + &self, + user_id: &str, + slug: &str, + ) -> Result { + self.check_rate(user_id, slug)?; + + let global = Arc::clone(&self.global) + .try_acquire_owned() + .map_err(|_| PageExecutionRejection::Busy)?; + let user = self.user_semaphore(user_id); + let user = user + .try_acquire_owned() + .map_err(|_| PageExecutionRejection::Busy)?; + let page = self.page_semaphore(user_id, slug); + let page = page + .try_acquire_owned() + .map_err(|_| PageExecutionRejection::Busy)?; + + Ok(PageExecutionPermit { + _global: global, + _user: user, + _page: page, + }) + } + + fn user_semaphore(&self, user_id: &str) -> Arc { + self.prune_idle_semaphores(); + Arc::clone( + self.users + .entry(user_id.to_string()) + .or_insert_with(|| Arc::new(Semaphore::new(self.limits.user_concurrency))) + .value(), + ) + } + + fn page_semaphore(&self, user_id: &str, slug: &str) -> Arc { + self.prune_idle_semaphores(); + let key = format!("{user_id}\0{slug}"); + Arc::clone( + self.pages + .entry(key) + .or_insert_with(|| Arc::new(Semaphore::new(self.limits.page_concurrency))) + .value(), + ) + } + + fn prune_idle_semaphores(&self) { + if self.users.len() > MAX_TRACKED_IDENTITIES { + self.users + .retain(|_, semaphore| Arc::strong_count(semaphore) > 1); + } + if self.pages.len() > MAX_TRACKED_IDENTITIES { + self.pages + .retain(|_, semaphore| Arc::strong_count(semaphore) > 1); + } + } + + fn check_rate(&self, user_id: &str, slug: &str) -> Result<(), PageExecutionRejection> { + let now = Instant::now(); + let cutoff = now.checked_sub(self.limits.rate_window).unwrap_or(now); + let mut windows = self + .rate_windows + .lock() + .map_err(|_| PageExecutionRejection::Busy)?; + if windows.len() > MAX_TRACKED_IDENTITIES * 2 { + windows.retain(|_, entries| { + entries.retain(|timestamp| *timestamp > cutoff); + !entries.is_empty() + }); + } + + record_request( + &mut windows, + format!("user\0{user_id}"), + cutoff, + now, + self.limits.user_requests_per_window, + )?; + if let Err(error) = record_request( + &mut windows, + format!("page\0{user_id}\0{slug}"), + cutoff, + now, + self.limits.page_requests_per_window, + ) { + if let Some(user_window) = windows.get_mut(&format!("user\0{user_id}")) { + if user_window.back() == Some(&now) { + user_window.pop_back(); + } + } + return Err(error); + } + Ok(()) + } +} + +fn record_request( + windows: &mut HashMap>, + key: String, + cutoff: Instant, + now: Instant, + limit: usize, +) -> Result<(), PageExecutionRejection> { + let entries = windows.entry(key).or_default(); + while entries + .front() + .is_some_and(|timestamp| *timestamp <= cutoff) + { + entries.pop_front(); + } + if entries.len() >= limit { + return Err(PageExecutionRejection::RateLimited); + } + entries.push_back(now); + Ok(()) +} + +#[derive(Debug)] +pub struct PageExecutionPermit { + _global: OwnedSemaphorePermit, + _user: OwnedSemaphorePermit, + _page: OwnedSemaphorePermit, +} + +#[cfg(test)] +mod tests { + use super::*; + + fn test_guard() -> PageExecutionGuard { + PageExecutionGuard::with_limits(Limits { + global_concurrency: 3, + user_concurrency: 2, + page_concurrency: 1, + user_requests_per_window: 100, + page_requests_per_window: 100, + rate_window: Duration::from_secs(60), + }) + } + + #[test] + fn page_concurrency_is_isolated_and_permits_recover() { + let guard = test_guard(); + let first = guard.try_acquire("user", "one").unwrap(); + assert_eq!( + guard.try_acquire("user", "one").unwrap_err(), + PageExecutionRejection::Busy + ); + let second_page = guard.try_acquire("user", "two").unwrap(); + assert_eq!( + guard.try_acquire("user", "three").unwrap_err(), + PageExecutionRejection::Busy + ); + drop((first, second_page)); + assert!(guard.try_acquire("user", "one").is_ok()); + } + + #[test] + fn page_and_account_rate_windows_are_enforced() { + let guard = PageExecutionGuard::with_limits(Limits { + global_concurrency: 3, + user_concurrency: 2, + page_concurrency: 1, + user_requests_per_window: 4, + page_requests_per_window: 2, + rate_window: Duration::from_secs(60), + }); + drop(guard.try_acquire("user", "one").unwrap()); + drop(guard.try_acquire("user", "one").unwrap()); + assert_eq!( + guard.try_acquire("user", "one").unwrap_err(), + PageExecutionRejection::RateLimited + ); + + drop(guard.try_acquire("user", "two").unwrap()); + drop(guard.try_acquire("user", "two").unwrap()); + assert_eq!( + guard.try_acquire("user", "three").unwrap_err(), + PageExecutionRejection::RateLimited + ); + assert!(guard.try_acquire("other", "one").is_ok()); + } +} diff --git a/src/crates/services/relay-service/src/routes/api.rs b/src/crates/services/relay-service/src/routes/api.rs index 6388ee4055..0c1d8d2e35 100644 --- a/src/crates/services/relay-service/src/routes/api.rs +++ b/src/crates/services/relay-service/src/routes/api.rs @@ -65,6 +65,13 @@ pub struct AppState { pub db: Option>, /// Optional per-page mutable data root (KV/SQLite/blobs). Required for Page Functions data plane. pub page_data: Option, + /// Short-lived browser handoff tickets and page-scoped grants. These stay + /// process-local so no account credential is persisted or exposed in URLs. + pub page_access_manager: Arc, + /// Manifest-bound Page draft upload sessions and per-Page serialization. + pub page_upload_manager: Arc, + /// Global/account/page admission control for public Page Function execution. + pub page_execution_guard: Arc, /// Per-IP rate limiter for auth endpoints (brute-force protection). pub login_rate_limiter: Arc, /// Per-user online device registry for account-based device routing. @@ -622,6 +629,9 @@ mod tests { asset_store: Arc::new(MemoryAssetStore::new()), db: None, page_data: None, + page_access_manager: Arc::new(crate::routes::pages::PageAccessManager::new()), + page_upload_manager: Arc::new(crate::routes::pages::PageUploadManager::new()), + page_execution_guard: Arc::new(crate::page_execution::PageExecutionGuard::new()), login_rate_limiter: Arc::new(crate::routes::auth::LoginRateLimiter::new()), device_manager: crate::relay::DeviceManager::new(), cors_allow_origins: Arc::new(Vec::new()), diff --git a/src/crates/services/relay-service/src/routes/pages.rs b/src/crates/services/relay-service/src/routes/pages.rs index 5209905feb..83e43f6981 100644 --- a/src/crates/services/relay-service/src/routes/pages.rs +++ b/src/crates/services/relay-service/src/routes/pages.rs @@ -5,18 +5,21 @@ //! Production: `/p/{user}/{slug}/...` (deployed version only). use axum::extract::{DefaultBodyLimit, Path, State}; -use axum::http::{header, HeaderMap, Method, StatusCode}; -use axum::response::IntoResponse; +use axum::http::{header, HeaderMap, HeaderValue, Method, StatusCode}; +use axum::response::{IntoResponse, Redirect, Response}; use axum::routing::{get, post}; use axum::{Json, Router}; use base64::{engine::general_purpose::STANDARD as B64, Engine}; use bitfun_page_function_runtime::{ - run_fetch, FetchRequest, PageMeta, DEFAULT_TIMEOUT, WORKER_ENTRY_PATH, + run_fetch, FetchRequest, PageFunctionError, PageMeta, DEFAULT_TIMEOUT, WORKER_ENTRY_PATH, }; +use dashmap::DashMap; use serde::{Deserialize, Serialize}; use sha2::{Digest, Sha256}; use std::collections::HashMap; +use std::hash::{DefaultHasher, Hash, Hasher}; use std::sync::Arc; +use std::time::{Duration, Instant}; use crate::db::{ new_page_version_id, page_draft_asset_key, page_legacy_asset_key, page_version_asset_key, @@ -31,7 +34,207 @@ pub const MAX_PAGES_PER_USER: i64 = 50; pub const MAX_VERSIONS_PER_PAGE: i64 = 30; pub const MAX_PAGE_BYTES: u64 = 100 * 1024 * 1024; pub const MAX_FILE_BYTES: u64 = 10 * 1024 * 1024; +pub const MAX_PAGE_FILES: usize = 4096; pub const PAGE_UPLOAD_BODY_LIMIT: usize = 12 * 1024 * 1024; +const PAGE_OPEN_TICKET_TTL: Duration = Duration::from_secs(60); +const PAGE_BROWSER_GRANT_TTL: Duration = Duration::from_secs(10 * 60); +const MAX_PAGE_OPEN_TICKETS: usize = 4096; +const MAX_PAGE_BROWSER_GRANTS: usize = 8192; +const PAGE_ACCESS_COOKIE: &str = "bitfun_page_access"; +const PAGE_UPLOAD_SESSION_TTL: Duration = Duration::from_secs(15 * 60); +const MAX_PAGE_UPLOAD_SESSIONS: usize = 4096; +const PAGE_UPLOAD_LOCK_SHARDS: usize = 64; + +#[derive(Clone)] +struct PageAccessScope { + user_id: String, + username: String, + slug: String, + version_id: Option, + expires_at: Instant, +} + +/// Process-local, time-bounded grants used to hand an authenticated Page URL +/// from the desktop client to the user's external browser without putting the +/// account bearer token in a URL. +#[derive(Default)] +pub struct PageAccessManager { + open_tickets: DashMap, + browser_grants: DashMap, +} + +#[derive(Clone)] +struct PageUploadSession { + upload_id: String, + draft_key: String, + manifest: HashMap, + finalized: bool, + expires_at: Instant, +} + +/// Tracks the one active, manifest-bound upload session for each Page. Drafts +/// use upload-specific asset namespaces, so superseded/concurrent uploads can +/// never mix files before an immutable version is frozen. +pub struct PageUploadManager { + sessions: DashMap, + locks: Vec>, +} + +impl Default for PageUploadManager { + fn default() -> Self { + Self::new() + } +} + +impl PageUploadManager { + pub fn new() -> Self { + Self { + sessions: DashMap::new(), + locks: (0..PAGE_UPLOAD_LOCK_SHARDS) + .map(|_| tokio::sync::Mutex::new(())) + .collect(), + } + } + + fn lock_for(&self, page_key: &str) -> &tokio::sync::Mutex<()> { + let mut hasher = DefaultHasher::new(); + page_key.hash(&mut hasher); + &self.locks[hasher.finish() as usize % self.locks.len()] + } + + fn prune_expired(&self, now: Instant) -> Vec { + let mut expired_drafts = Vec::new(); + self.sessions.retain(|_, session| { + let keep = session.expires_at > now; + if !keep { + expired_drafts.push(session.draft_key.clone()); + } + keep + }); + expired_drafts + } +} + +impl PageAccessManager { + pub fn new() -> Self { + Self::default() + } + + fn prune_expired(&self, now: Instant) { + self.open_tickets.retain(|_, scope| scope.expires_at > now); + self.browser_grants + .retain(|_, scope| scope.expires_at > now); + } + + fn issue_open_ticket( + &self, + user_id: String, + username: String, + slug: String, + version_id: Option, + ) -> Result { + let now = Instant::now(); + self.prune_expired(now); + if self.open_tickets.len() >= MAX_PAGE_OPEN_TICKETS { + return Err(StatusCode::TOO_MANY_REQUESTS); + } + let ticket = random_page_access_token(); + self.open_tickets.insert( + ticket.clone(), + PageAccessScope { + user_id, + username, + slug, + version_id, + expires_at: now + PAGE_OPEN_TICKET_TTL, + }, + ); + Ok(ticket) + } + + fn exchange_ticket( + &self, + ticket: &str, + ) -> Result, StatusCode> { + let now = Instant::now(); + self.prune_expired(now); + if self.browser_grants.len() >= MAX_PAGE_BROWSER_GRANTS { + return Err(StatusCode::TOO_MANY_REQUESTS); + } + let Some((_, mut scope)) = self.open_tickets.remove(ticket) else { + return Ok(None); + }; + if scope.expires_at <= now { + return Ok(None); + } + let grant = random_page_access_token(); + scope.expires_at = now + PAGE_BROWSER_GRANT_TTL; + self.browser_grants.insert(grant.clone(), scope.clone()); + Ok(Some((grant, scope))) + } + + fn authorizes_page( + &self, + headers: &HeaderMap, + user_id: &str, + slug: &str, + version_id: Option<&str>, + ) -> bool { + let now = Instant::now(); + let mut authorized = false; + for cookie_header in headers.get_all(header::COOKIE) { + let Ok(value) = cookie_header.to_str() else { + continue; + }; + for cookie in value.split(';') { + let Some((name, token)) = cookie.trim().split_once('=') else { + continue; + }; + if name != PAGE_ACCESS_COOKIE { + continue; + } + let Some(scope) = self.browser_grants.get(token) else { + continue; + }; + if scope.expires_at > now + && scope.user_id == user_id + && scope.slug == slug + && scope.version_id.as_deref() == version_id + { + authorized = true; + break; + } + } + if authorized { + break; + } + } + if !authorized { + self.prune_expired(now); + } + authorized + } +} + +fn random_page_access_token() -> String { + format!( + "{}{}", + uuid::Uuid::new_v4().simple(), + uuid::Uuid::new_v4().simple() + ) +} + +fn is_valid_page_upload_id(upload_id: &str) -> bool { + upload_id.len() == 32 && upload_id.bytes().all(|byte| byte.is_ascii_hexdigit()) +} + +fn page_upload_session_key(user_id: &str, slug: &str) -> String { + format!("{user_id}\0{slug}") +} + +fn page_upload_draft_key(user_id: &str, slug: &str, upload_id: &str) -> String { + format!("pages/{user_id}/{slug}/draft/{upload_id}") +} fn is_valid_slug(slug: &str) -> bool { let bytes = slug.as_bytes(); @@ -73,18 +276,26 @@ pub fn pages_router() -> Router { get(list_versions).post(freeze_version), ) .route("/api/pages/{slug}/deploy", post(deploy_version)) + .route("/api/pages/{slug}/unpublish", post(unpublish_page)) .route( "/api/pages/{slug}/versions/{version_id}", axum::routing::delete(delete_version), ) .route( "/api/pages/{slug}", - axum::routing::patch(update_page).delete(delete_page), + post(create_open_ticket) + .patch(update_page) + .delete(delete_page), ) + .route("/api/page-open/{ticket}", get(exchange_open_ticket)) // Preview routes (more specific first). .route( "/p/{username}/{slug}/@v/{version_id}", - get(serve_preview_root).post(serve_preview_root), + get(serve_preview_root) + .post(serve_preview_root) + .layer(DefaultBodyLimit::max( + crate::page_execution::MAX_PAGE_FUNCTION_REQUEST_BODY_BYTES, + )), ) .route( "/p/{username}/{slug}/@v/{version_id}/{*path}", @@ -92,11 +303,18 @@ pub fn pages_router() -> Router { .post(serve_preview_path) .put(serve_preview_path) .delete(serve_preview_path) - .patch(serve_preview_path), + .patch(serve_preview_path) + .layer(DefaultBodyLimit::max( + crate::page_execution::MAX_PAGE_FUNCTION_REQUEST_BODY_BYTES, + )), ) .route( "/p/{username}/{slug}", - get(serve_prod_root).post(serve_prod_root), + get(serve_prod_root) + .post(serve_prod_root) + .layer(DefaultBodyLimit::max( + crate::page_execution::MAX_PAGE_FUNCTION_REQUEST_BODY_BYTES, + )), ) .route( "/p/{username}/{slug}/{*path}", @@ -104,7 +322,10 @@ pub fn pages_router() -> Router { .post(serve_prod_path) .put(serve_prod_path) .delete(serve_prod_path) - .patch(serve_prod_path), + .patch(serve_prod_path) + .layer(DefaultBodyLimit::max( + crate::page_execution::MAX_PAGE_FUNCTION_REQUEST_BODY_BYTES, + )), ) } @@ -120,11 +341,13 @@ pub struct FileManifestEntry { #[derive(Deserialize)] pub struct CheckPageFilesRequest { pub slug: String, + pub upload_id: String, pub files: Vec, } #[derive(Serialize)] pub struct CheckPageFilesResponse { + pub upload_id: String, pub needed: Vec, pub existing_count: usize, pub total_count: usize, @@ -139,6 +362,7 @@ pub struct UploadFileEntry { #[derive(Deserialize)] pub struct UploadPageFilesRequest { pub slug: String, + pub upload_id: String, #[serde(default)] pub title: String, pub visibility: String, @@ -149,6 +373,7 @@ pub struct UploadPageFilesRequest { #[derive(Deserialize)] pub struct FreezeVersionRequest { + pub upload_id: String, #[serde(default)] pub title: String, #[serde(default)] @@ -193,8 +418,128 @@ pub struct UpdatePageRequest { pub title: Option, } +#[derive(Deserialize)] +pub struct PageOpenRequest { + #[serde(default)] + pub version_id: Option, +} + +#[derive(Serialize)] +pub struct PageOpenResponse { + pub open_url_path: String, + pub expires_in_seconds: u64, +} + // ── Management ────────────────────────────────────────────────────────── +async fn create_open_ticket( + State(state): State, + headers: HeaderMap, + Path(slug): Path, + Json(body): Json, +) -> Result, StatusCode> { + let auth = validate_auth(&state, &headers).await?; + let db = require_db(&state)?; + if !is_valid_slug(&slug) { + return Err(StatusCode::BAD_REQUEST); + } + let page = PageRow::get(db, &auth.user_id, &slug) + .await + .map_err(|_| StatusCode::INTERNAL_SERVER_ERROR)? + .ok_or(StatusCode::NOT_FOUND)?; + let username = UserRow::find_by_username_for_user_id(db, &auth.user_id) + .await + .map_err(|_| StatusCode::INTERNAL_SERVER_ERROR)? + .ok_or(StatusCode::NOT_FOUND)?; + if !is_safe_page_username(&username) { + return Err(StatusCode::BAD_REQUEST); + } + + let version_id = body + .version_id + .map(|value| value.trim().to_string()) + .filter(|value| !value.is_empty()); + if let Some(version_id) = version_id.as_deref() { + PageVersionRow::get(db, &auth.user_id, &slug, version_id) + .await + .map_err(|_| StatusCode::INTERNAL_SERVER_ERROR)? + .ok_or(StatusCode::NOT_FOUND)?; + } else if page.deployed_version_id.is_none() { + return Err(StatusCode::NOT_FOUND); + } + + let ticket = + state + .page_access_manager + .issue_open_ticket(auth.user_id, username, slug, version_id)?; + Ok(Json(PageOpenResponse { + open_url_path: format!("/api/page-open/{ticket}"), + expires_in_seconds: PAGE_OPEN_TICKET_TTL.as_secs(), + })) +} + +async fn exchange_open_ticket( + State(state): State, + headers: HeaderMap, + Path(ticket): Path, +) -> Result { + if ticket.len() != 64 || !ticket.bytes().all(|byte| byte.is_ascii_hexdigit()) { + return Err(StatusCode::NOT_FOUND); + } + let (grant, scope) = state + .page_access_manager + .exchange_ticket(&ticket)? + .ok_or(StatusCode::NOT_FOUND)?; + let page_path = format!("/p/{}/{}", scope.username, scope.slug); + let target = match scope.version_id { + Some(version_id) => format!("{page_path}/@v/{version_id}"), + None => page_path.clone(), + }; + let secure = request_used_https(&headers); + let cookie = format!( + "{PAGE_ACCESS_COOKIE}={grant}; Path={page_path}; Max-Age={}; HttpOnly; SameSite=Lax{}", + PAGE_BROWSER_GRANT_TTL.as_secs(), + if secure { "; Secure" } else { "" } + ); + let mut response = Redirect::temporary(&target).into_response(); + response.headers_mut().insert( + header::SET_COOKIE, + HeaderValue::from_str(&cookie).map_err(|_| StatusCode::INTERNAL_SERVER_ERROR)?, + ); + response + .headers_mut() + .insert(header::CACHE_CONTROL, HeaderValue::from_static("no-store")); + Ok(response) +} + +fn request_used_https(headers: &HeaderMap) -> bool { + headers + .get("forwarded") + .and_then(|value| value.to_str().ok()) + .is_some_and(|value| { + value + .split(';') + .any(|part| part.trim().eq_ignore_ascii_case("proto=https")) + }) + || headers + .get("x-forwarded-proto") + .and_then(|value| value.to_str().ok()) + .is_some_and(|value| { + value + .split(',') + .next() + .is_some_and(|proto| proto.trim().eq_ignore_ascii_case("https")) + }) +} + +fn is_safe_page_username(username: &str) -> bool { + !username.is_empty() + && username.len() <= 128 + && username + .bytes() + .all(|byte| byte.is_ascii_alphanumeric() || matches!(byte, b'-' | b'_' | b'.')) +} + async fn list_pages( State(state): State, headers: HeaderMap, @@ -230,11 +575,15 @@ async fn check_page_files( ) -> Result, StatusCode> { let auth = validate_auth(&state, &headers).await?; let db = require_db(&state)?; - if !is_valid_slug(&body.slug) { + if !is_valid_slug(&body.slug) || !is_valid_page_upload_id(&body.upload_id) { return Err(StatusCode::BAD_REQUEST); } let mut total_bytes: u64 = 0; + let mut manifest = HashMap::with_capacity(body.files.len()); + if body.files.is_empty() || body.files.len() > MAX_PAGE_FILES { + return Err(StatusCode::BAD_REQUEST); + } for entry in &body.files { if crate::validated_asset_relative_path(&entry.path).is_err() || !crate::is_valid_content_hash(&entry.hash) @@ -245,6 +594,12 @@ async fn check_page_files( return Err(StatusCode::PAYLOAD_TOO_LARGE); } total_bytes = total_bytes.saturating_add(entry.size); + if manifest + .insert(entry.path.clone(), entry.hash.clone()) + .is_some() + { + return Err(StatusCode::BAD_REQUEST); + } } if total_bytes > MAX_PAGE_BYTES { return Err(StatusCode::PAYLOAD_TOO_LARGE); @@ -262,13 +617,50 @@ async fn check_page_files( } } - let draft_key = page_draft_asset_key(&auth.user_id, &body.slug); + let page_key = page_upload_session_key(&auth.user_id, &body.slug); + let _page_guard = state.page_upload_manager.lock_for(&page_key).lock().await; + let now = Instant::now(); + for stale_draft in state.page_upload_manager.prune_expired(now) { + state.asset_store.cleanup_room(&stale_draft); + } + let replacing_existing_page_session = + state.page_upload_manager.sessions.contains_key(&page_key); + if !replacing_existing_page_session + && state.page_upload_manager.sessions.len() >= MAX_PAGE_UPLOAD_SESSIONS + { + return Err(StatusCode::TOO_MANY_REQUESTS); + } + + let draft_key = page_upload_draft_key(&auth.user_id, &body.slug, &body.upload_id); + if let Some((_, previous)) = state.page_upload_manager.sessions.remove(&page_key) { + state.asset_store.cleanup_room(&previous.draft_key); + } + // Reclaim upload-specific drafts left by a prior process crash. The + // in-memory store uses exact namespaces (the active one was removed + // above); the disk store recursively clears this Page's draft subtree. + state + .asset_store + .cleanup_room(&page_draft_asset_key(&auth.user_id, &body.slug)); + // A retry using the same upload id must also start from the submitted + // manifest, never from partial mappings left by the failed attempt. + state.asset_store.cleanup_room(&draft_key); let asset_store = Arc::clone(&state.asset_store); + let response_upload_id = body.upload_id.clone(); let response = tokio::task::spawn_blocking(move || { - process_check_page_files(asset_store, &draft_key, body.files) + process_check_page_files(asset_store, &draft_key, response_upload_id, body.files) }) .await - .map_err(|_| StatusCode::INTERNAL_SERVER_ERROR)?; + .map_err(|_| StatusCode::INTERNAL_SERVER_ERROR)??; + state.page_upload_manager.sessions.insert( + page_key, + PageUploadSession { + upload_id: body.upload_id.clone(), + draft_key: page_upload_draft_key(&auth.user_id, &body.slug, &body.upload_id), + manifest, + finalized: false, + expires_at: now + PAGE_UPLOAD_SESSION_TTL, + }, + ); Ok(Json(response)) } @@ -279,7 +671,7 @@ async fn upload_page_files( ) -> Result, StatusCode> { let auth = validate_auth(&state, &headers).await?; let db = require_db(&state)?; - if !is_valid_slug(&body.slug) { + if !is_valid_slug(&body.slug) || !is_valid_page_upload_id(&body.upload_id) { return Err(StatusCode::BAD_REQUEST); } let visibility = PageVisibility::parse(&body.visibility).ok_or(StatusCode::BAD_REQUEST)?; @@ -296,10 +688,32 @@ async fn upload_page_files( } } + let page_key = page_upload_session_key(&auth.user_id, &body.slug); + let _page_guard = state.page_upload_manager.lock_for(&page_key).lock().await; + let session = state + .page_upload_manager + .sessions + .get(&page_key) + .map(|entry| entry.value().clone()) + .ok_or(StatusCode::CONFLICT)?; + if session.upload_id != body.upload_id { + return Err(StatusCode::CONFLICT); + } + if session.expires_at <= Instant::now() { + state.page_upload_manager.sessions.remove(&page_key); + state.asset_store.cleanup_room(&session.draft_key); + return Err(StatusCode::GONE); + } + for (path, entry) in &body.files { + if session.manifest.get(path) != Some(&entry.hash) { + return Err(StatusCode::CONFLICT); + } + } + let existing = PageRow::get(db, &auth.user_id, &body.slug) .await .map_err(|_| StatusCode::INTERNAL_SERVER_ERROR)?; - if existing.is_none() { + if body.finalize && existing.is_none() { let count = PageRow::count_for_user(db, &auth.user_id) .await .map_err(|_| StatusCode::INTERNAL_SERVER_ERROR)?; @@ -316,11 +730,7 @@ async fn upload_page_files( } else { body.title.clone() }; - PageRow::ensure(db, &auth.user_id, &body.slug, visibility, &title) - .await - .map_err(|_| StatusCode::INTERNAL_SERVER_ERROR)?; - - let draft_key = page_draft_asset_key(&auth.user_id, &body.slug); + let draft_key = session.draft_key.clone(); let asset_store = Arc::clone(&state.asset_store); let files = body.files; let stored = tokio::task::spawn_blocking(move || { @@ -329,10 +739,35 @@ async fn upload_page_files( .await .map_err(|_| StatusCode::INTERNAL_SERVER_ERROR)??; + if body.finalize { + let actual_manifest = state + .asset_store + .list_room_entries(&session.draft_key) + .into_iter() + .collect::>(); + if actual_manifest != session.manifest { + return Err(StatusCode::CONFLICT); + } + PageRow::ensure(db, &auth.user_id, &body.slug, visibility, &title) + .await + .map_err(|_| StatusCode::INTERNAL_SERVER_ERROR)?; + let mut active = state + .page_upload_manager + .sessions + .get_mut(&page_key) + .ok_or(StatusCode::CONFLICT)?; + if active.upload_id != body.upload_id { + return Err(StatusCode::CONFLICT); + } + active.finalized = true; + active.expires_at = Instant::now() + PAGE_UPLOAD_SESSION_TTL; + } + Ok(Json(serde_json::json!({ "status": "ok", "files_stored": stored, "slug": body.slug, + "upload_id": body.upload_id, "draft": true, "finalize": body.finalize, }))) @@ -346,10 +781,27 @@ async fn freeze_version( ) -> Result, StatusCode> { let auth = validate_auth(&state, &headers).await?; let db = require_db(&state)?; - if !is_valid_slug(&slug) { + if !is_valid_slug(&slug) || !is_valid_page_upload_id(&body.upload_id) { return Err(StatusCode::BAD_REQUEST); } + let page_key = page_upload_session_key(&auth.user_id, &slug); + let _page_guard = state.page_upload_manager.lock_for(&page_key).lock().await; + let session = state + .page_upload_manager + .sessions + .get(&page_key) + .map(|entry| entry.value().clone()) + .ok_or(StatusCode::CONFLICT)?; + if session.upload_id != body.upload_id || !session.finalized { + return Err(StatusCode::CONFLICT); + } + if session.expires_at <= Instant::now() { + state.page_upload_manager.sessions.remove(&page_key); + state.asset_store.cleanup_room(&session.draft_key); + return Err(StatusCode::GONE); + } + let page = PageRow::get(db, &auth.user_id, &slug) .await .map_err(|_| StatusCode::INTERNAL_SERVER_ERROR)? @@ -362,7 +814,7 @@ async fn freeze_version( return Err(StatusCode::TOO_MANY_REQUESTS); } - let draft_key = page_draft_asset_key(&auth.user_id, &slug); + let draft_key = session.draft_key.clone(); if !state.asset_store.has_room_files(&draft_key) { return Err(StatusCode::BAD_REQUEST); } @@ -400,7 +852,7 @@ async fn freeze_version( body.title.clone() }; - PageVersionRow::insert( + if PageVersionRow::insert( db, &auth.user_id, &slug, @@ -412,9 +864,14 @@ async fn freeze_version( &body.note, ) .await - .map_err(|_| StatusCode::INTERNAL_SERVER_ERROR)?; + .is_err() + { + state.asset_store.cleanup_room(&version_key); + return Err(StatusCode::INTERNAL_SERVER_ERROR); + } - // Clear draft after successful freeze. + // Consume the exact finalized upload session after a successful freeze. + state.page_upload_manager.sessions.remove(&page_key); state.asset_store.cleanup_room(&draft_key); let username = UserRow::find_by_username_for_user_id(db, &auth.user_id) @@ -520,6 +977,25 @@ async fn deploy_version( Ok(Json(page_to_info(&page, &username))) } +async fn unpublish_page( + State(state): State, + headers: HeaderMap, + Path(slug): Path, +) -> Result { + let auth = validate_auth(&state, &headers).await?; + let db = require_db(&state)?; + if !is_valid_slug(&slug) { + return Err(StatusCode::BAD_REQUEST); + } + let updated = PageRow::clear_deployed_version(db, &auth.user_id, &slug) + .await + .map_err(|_| StatusCode::INTERNAL_SERVER_ERROR)?; + if !updated { + return Err(StatusCode::NOT_FOUND); + } + Ok(StatusCode::NO_CONTENT) +} + async fn delete_version( State(state): State, headers: HeaderMap, @@ -593,9 +1069,23 @@ async fn delete_page( if !is_valid_slug(&slug) { return Err(StatusCode::BAD_REQUEST); } + let page_key = page_upload_session_key(&auth.user_id, &slug); + let _page_guard = state.page_upload_manager.lock_for(&page_key).lock().await; let versions = PageVersionRow::list_for_page(db, &auth.user_id, &slug) .await .map_err(|_| StatusCode::INTERNAL_SERVER_ERROR)?; + // Commit the relational deletion first. Filesystem/object cleanup is + // idempotent, whereas deleting assets before a failed multi-table DB + // mutation could leave a partially present Page with missing content. + let deleted = PageRow::delete(db, &auth.user_id, &slug) + .await + .map_err(|_| StatusCode::INTERNAL_SERVER_ERROR)?; + if !deleted { + return Err(StatusCode::NOT_FOUND); + } + if let Some((_, active_upload)) = state.page_upload_manager.sessions.remove(&page_key) { + state.asset_store.cleanup_room(&active_upload.draft_key); + } for v in &versions { state.asset_store.cleanup_room(&page_version_asset_key( &auth.user_id, @@ -612,12 +1102,6 @@ async fn delete_page( if let Some(store) = &state.page_data { store.cleanup_page(&auth.user_id, &slug); } - let deleted = PageRow::delete(db, &auth.user_id, &slug) - .await - .map_err(|_| StatusCode::INTERNAL_SERVER_ERROR)?; - if !deleted { - return Err(StatusCode::NOT_FOUND); - } Ok(Json(serde_json::json!({ "status": "ok", "slug": slug }))) } @@ -710,7 +1194,7 @@ async fn serve_page( .await .map_err(|_| StatusCode::INTERNAL_SERVER_ERROR)? .ok_or(StatusCode::NOT_FOUND)?; - enforce_visibility(&state, &headers, &page).await?; + enforce_visibility(&state, &headers, &page, version_override).await?; // Resolve version. let version_id = if let Some(v) = version_override { @@ -794,6 +1278,13 @@ async fn serve_with_worker( raw_path: &str, body: axum::body::Bytes, ) -> Result { + if body.len() > crate::page_execution::MAX_PAGE_FUNCTION_REQUEST_BODY_BYTES { + return Err(StatusCode::PAYLOAD_TOO_LARGE); + } + let _execution_permit = state + .page_execution_guard + .try_acquire(&page.user_id, &page.slug) + .map_err(|_| StatusCode::TOO_MANY_REQUESTS)?; let page_data = state.page_data.clone().ok_or(StatusCode::NOT_IMPLEMENTED)?; let db = state.db.clone().ok_or(StatusCode::NOT_IMPLEMENTED)?; let asset_store = Arc::clone(&state.asset_store); @@ -875,7 +1366,11 @@ async fn serve_with_worker( .map_err(|_| StatusCode::INTERNAL_SERVER_ERROR)? .map_err(|e| { tracing::warn!("Page function error: {e}"); - StatusCode::BAD_GATEWAY + if matches!(e, PageFunctionError::Timeout(_)) { + StatusCode::GATEWAY_TIMEOUT + } else { + StatusCode::BAD_GATEWAY + } })?; // If worker returns 404 for document GET, fall back to static assets. @@ -908,7 +1403,9 @@ async fn serve_with_worker( axum::http::HeaderName::try_from(k), axum::http::HeaderValue::try_from(v), ) { - builder = builder.header(name, val); + if should_forward_page_worker_response_header(&name) { + builder = builder.header(name, val); + } } } builder @@ -916,6 +1413,38 @@ async fn serve_with_worker( .map_err(|_| StatusCode::INTERNAL_SERVER_ERROR) } +/// Filter worker-controlled response headers that can mutate origin-wide +/// browser policy/state or connection framing. This is defense in depth: Page +/// documents still share an origin with the relay until hosting moves to a +/// dedicated Page origin. Ordinary representation and CORS headers remain +/// available to Page authors. +fn should_forward_page_worker_response_header(name: &axum::http::HeaderName) -> bool { + !matches!( + name.as_str(), + "accept-ch" + | "alt-svc" + | "clear-site-data" + | "connection" + | "content-length" + | "critical-ch" + | "keep-alive" + | "nel" + | "proxy-authenticate" + | "proxy-authorization" + | "report-to" + | "reporting-endpoints" + | "service-worker-allowed" + | "set-cookie" + | "set-cookie2" + | "strict-transport-security" + | "te" + | "trailer" + | "transfer-encoding" + | "upgrade" + | "x-content-type-options" + ) +} + // ── Helpers ───────────────────────────────────────────────────────────── fn page_to_info(page: &PageRow, username: &str) -> PageInfo { @@ -1005,6 +1534,7 @@ async fn enforce_visibility( state: &AppState, headers: &HeaderMap, page: &PageWithUsername, + version_id: Option<&str>, ) -> Result<(), StatusCode> { let visibility = page .visibility_enum() @@ -1012,19 +1542,37 @@ async fn enforce_visibility( match visibility { PageVisibility::Public => Ok(()), PageVisibility::Relay => { - resolve_viewer(state, headers).await?; - Ok(()) + if resolve_viewer(state, headers).await.is_ok() + || state.page_access_manager.authorizes_page( + headers, + &page.user_id, + &page.slug, + version_id, + ) + { + Ok(()) + } else { + Err(StatusCode::UNAUTHORIZED) + } } PageVisibility::Private => { // Return NOT_FOUND for any auth failure so the existence of a // private page is not revealed to anonymous or foreign viewers. - let viewer = resolve_viewer(state, headers) - .await - .map_err(|_| StatusCode::NOT_FOUND)?; - if viewer.user_id != page.user_id { - return Err(StatusCode::NOT_FOUND); + if let Ok(viewer) = resolve_viewer(state, headers).await { + if viewer.user_id == page.user_id { + return Ok(()); + } + } + if state.page_access_manager.authorizes_page( + headers, + &page.user_id, + &page.slug, + version_id, + ) { + Ok(()) + } else { + Err(StatusCode::NOT_FOUND) } - Ok(()) } } } @@ -1040,8 +1588,9 @@ async fn resolve_viewer(state: &AppState, headers: &HeaderMap) -> Result, asset_key: &str, + upload_id: String, files: Vec, -) -> CheckPageFilesResponse { +) -> Result { let mut needed = Vec::new(); let mut existing_count = 0usize; let total_count = files.len(); @@ -1054,16 +1603,19 @@ fn process_check_page_files( } if asset_store.has_content(&entry.hash) { existing_count += 1; - let _ = asset_store.map_to_room(asset_key, &entry.path, &entry.hash); + asset_store + .map_to_room(asset_key, &entry.path, &entry.hash) + .map_err(|_| StatusCode::INTERNAL_SERVER_ERROR)?; } else { needed.push(entry.path); } } - CheckPageFilesResponse { + Ok(CheckPageFilesResponse { + upload_id, needed, existing_count, total_count, - } + }) } fn process_upload_page_files( @@ -1146,6 +1698,30 @@ mod tests { use axum::http::{Request, StatusCode}; use tower::ServiceExt; + #[test] + fn worker_response_headers_cannot_mutate_origin_wide_browser_state() { + for name in [ + "service-worker-allowed", + "set-cookie", + "clear-site-data", + "strict-transport-security", + "reporting-endpoints", + "transfer-encoding", + ] { + let name = axum::http::HeaderName::from_bytes(name.as_bytes()).unwrap(); + assert!(!should_forward_page_worker_response_header(&name), "{name}"); + } + + for name in [ + "content-type", + "cache-control", + "access-control-allow-origin", + ] { + let name = axum::http::HeaderName::from_bytes(name.as_bytes()).unwrap(); + assert!(should_forward_page_worker_response_header(&name), "{name}"); + } + } + async fn setup_app() -> (axum::Router, String, String) { let pool = connect(":memory:").await.unwrap(); let pool = Arc::new(pool); @@ -1178,6 +1754,100 @@ mod tests { (app, tok_alice.token, tok_bob.token) } + async fn begin_test_upload( + app: &axum::Router, + token: &str, + slug: &str, + files: &[(&str, &[u8])], + ) -> (String, serde_json::Map) { + let upload_id = uuid::Uuid::new_v4().simple().to_string(); + let mut manifest = Vec::new(); + let mut upload_files = serde_json::Map::new(); + for (path, content) in files { + let hash = hex_sha256(content); + manifest.push(serde_json::json!({ + "path": path, + "hash": hash, + "size": content.len(), + })); + upload_files.insert( + (*path).to_string(), + serde_json::json!({ "content": B64.encode(content), "hash": hash }), + ); + } + let check = serde_json::json!({ + "slug": slug, + "upload_id": upload_id, + "files": manifest, + }); + let response = app + .clone() + .oneshot( + Request::builder() + .method("POST") + .uri("/api/pages/check-files") + .header("Authorization", format!("Bearer {token}")) + .header("content-type", "application/json") + .body(axum::body::Body::from(check.to_string())) + .unwrap(), + ) + .await + .unwrap(); + assert_eq!(response.status(), StatusCode::OK); + (upload_id, upload_files) + } + + async fn finish_test_upload( + app: &axum::Router, + token: &str, + slug: &str, + visibility: &str, + upload_id: &str, + files: serde_json::Map, + ) -> String { + let upload = serde_json::json!({ + "slug": slug, + "upload_id": upload_id, + "title": slug, + "visibility": visibility, + "files": files, + "finalize": true, + }); + let response = app + .clone() + .oneshot( + Request::builder() + .method("POST") + .uri("/api/pages/upload-files") + .header("Authorization", format!("Bearer {token}")) + .header("content-type", "application/json") + .body(axum::body::Body::from(upload.to_string())) + .unwrap(), + ) + .await + .unwrap(); + assert_eq!(response.status(), StatusCode::OK); + + let freeze = serde_json::json!({ "upload_id": upload_id, "title": slug }); + let response = app + .clone() + .oneshot( + Request::builder() + .method("POST") + .uri(format!("/api/pages/{slug}/versions")) + .header("Authorization", format!("Bearer {token}")) + .header("content-type", "application/json") + .body(axum::body::Body::from(freeze.to_string())) + .unwrap(), + ) + .await + .unwrap(); + assert_eq!(response.status(), StatusCode::OK); + let body = to_bytes(response.into_body(), usize::MAX).await.unwrap(); + let version: serde_json::Value = serde_json::from_slice(&body).unwrap(); + version["version_id"].as_str().unwrap().to_string() + } + async fn save_and_deploy(app: &axum::Router, token: &str, slug: &str, html: &str) -> String { save_and_deploy_with_visibility(app, token, slug, html, "public").await } @@ -1191,8 +1861,29 @@ mod tests { ) -> String { let hash = hex_sha256(html.as_bytes()); let b64 = B64.encode(html.as_bytes()); + let upload_id = uuid::Uuid::new_v4().simple().to_string(); + let check = serde_json::json!({ + "slug": slug, + "upload_id": upload_id, + "files": [{ "path": "index.html", "hash": hash, "size": html.len() }] + }); + let resp = app + .clone() + .oneshot( + Request::builder() + .method("POST") + .uri("/api/pages/check-files") + .header("Authorization", format!("Bearer {token}")) + .header("content-type", "application/json") + .body(axum::body::Body::from(check.to_string())) + .unwrap(), + ) + .await + .unwrap(); + assert_eq!(resp.status(), StatusCode::OK); let upload = serde_json::json!({ "slug": slug, + "upload_id": upload_id, "title": slug, "visibility": visibility, "finalize": true, @@ -1213,7 +1904,11 @@ mod tests { .unwrap(); assert_eq!(resp.status(), StatusCode::OK); - let freeze = serde_json::json!({ "title": slug, "note": "test" }); + let freeze = serde_json::json!({ + "upload_id": upload_id, + "title": slug, + "note": "test" + }); let resp = app .clone() .oneshot( @@ -1255,8 +1950,33 @@ mod tests { let (app, alice, _) = setup_app().await; let hash = hex_sha256(b"draft"); let b64 = B64.encode(b"draft"); + let upload_id = uuid::Uuid::new_v4().simple().to_string(); + let check = serde_json::json!({ + "slug": "staged", + "upload_id": upload_id, + "files": [{ + "path": "index.html", + "hash": hash, + "size": b"draft".len() + }] + }); + let resp = app + .clone() + .oneshot( + Request::builder() + .method("POST") + .uri("/api/pages/check-files") + .header("Authorization", format!("Bearer {alice}")) + .header("content-type", "application/json") + .body(axum::body::Body::from(check.to_string())) + .unwrap(), + ) + .await + .unwrap(); + assert_eq!(resp.status(), StatusCode::OK); let upload = serde_json::json!({ "slug": "staged", + "upload_id": upload_id, "title": "staged", "visibility": "public", "files": { "index.html": { "content": b64, "hash": hash } } @@ -1313,6 +2033,123 @@ mod tests { assert_eq!(resp.status(), StatusCode::OK); } + #[tokio::test] + async fn republish_manifest_removes_files_not_present_in_the_new_version() { + let (app, alice, _) = setup_app().await; + let first_files = [ + ("index.html", b"first".as_slice()), + ("removed.txt", b"must not survive".as_slice()), + ]; + let (first_upload, first_map) = + begin_test_upload(&app, &alice, "clean-republish", &first_files).await; + let first_version = finish_test_upload( + &app, + &alice, + "clean-republish", + "public", + &first_upload, + first_map, + ) + .await; + assert_eq!( + get_page( + &app, + &format!("/p/alice/clean-republish/@v/{first_version}/removed.txt"), + None, + ) + .await, + StatusCode::OK + ); + + let second_files = [("index.html", b"second".as_slice())]; + let (second_upload, second_map) = + begin_test_upload(&app, &alice, "clean-republish", &second_files).await; + let second_version = finish_test_upload( + &app, + &alice, + "clean-republish", + "public", + &second_upload, + second_map, + ) + .await; + let response = app + .clone() + .oneshot( + Request::builder() + .uri(format!( + "/p/alice/clean-republish/@v/{second_version}/removed.txt" + )) + .body(axum::body::Body::empty()) + .unwrap(), + ) + .await + .unwrap(); + // Static Page routing falls back to the new index.html for unknown + // paths. The old file content must not survive in the new manifest. + assert_eq!(response.status(), StatusCode::OK); + let body = to_bytes(response.into_body(), usize::MAX).await.unwrap(); + assert_eq!(String::from_utf8_lossy(&body), "second"); + } + + #[tokio::test] + async fn superseded_upload_session_cannot_mix_or_freeze_files() { + let (app, alice, _) = setup_app().await; + let first_files = [("index.html", b"first".as_slice())]; + let second_files = [("index.html", b"second".as_slice())]; + let (first_upload, first_map) = + begin_test_upload(&app, &alice, "concurrent", &first_files).await; + let (second_upload, second_map) = + begin_test_upload(&app, &alice, "concurrent", &second_files).await; + + let stale_upload = serde_json::json!({ + "slug": "concurrent", + "upload_id": first_upload, + "visibility": "private", + "files": first_map, + "finalize": true, + }); + let response = app + .clone() + .oneshot( + Request::builder() + .method("POST") + .uri("/api/pages/upload-files") + .header("Authorization", format!("Bearer {alice}")) + .header("content-type", "application/json") + .body(axum::body::Body::from(stale_upload.to_string())) + .unwrap(), + ) + .await + .unwrap(); + assert_eq!(response.status(), StatusCode::CONFLICT); + + let _ = finish_test_upload( + &app, + &alice, + "concurrent", + "private", + &second_upload, + second_map, + ) + .await; + let stale_freeze = app + .oneshot( + Request::builder() + .method("POST") + .uri("/api/pages/concurrent/versions") + .header("Authorization", format!("Bearer {alice}")) + .header("content-type", "application/json") + .body(axum::body::Body::from( + serde_json::json!({ "upload_id": first_upload }).to_string(), + )) + .unwrap(), + ) + .await + .unwrap(); + assert_eq!(stale_freeze.status(), StatusCode::CONFLICT); + } + #[tokio::test] async fn worker_can_use_kv() { let (app, alice, _) = setup_app().await; @@ -1329,19 +2166,14 @@ mod tests { ("index.html", b"static".as_slice()), ("server/worker.js", worker.as_bytes()), ]; - let mut map = serde_json::Map::new(); - for (path, content) in files { - let hash = hex_sha256(content); - map.insert( - path.to_string(), - serde_json::json!({ "content": B64.encode(content), "hash": hash }), - ); - } + let (upload_id, map) = begin_test_upload(&app, &alice, "fn", &files).await; let upload = serde_json::json!({ "slug": "fn", + "upload_id": upload_id, "title": "fn", "visibility": "public", - "files": map + "files": map, + "finalize": true }); let resp = app .clone() @@ -1357,7 +2189,7 @@ mod tests { .await .unwrap(); assert_eq!(resp.status(), StatusCode::OK); - let freeze = serde_json::json!({}); + let freeze = serde_json::json!({ "upload_id": upload_id }); let resp = app .clone() .oneshot( @@ -1464,6 +2296,197 @@ mod tests { assert_eq!(get_page(&app, "/p/alice/pub", None).await, StatusCode::OK); } + #[tokio::test] + async fn one_time_open_ticket_exchanges_for_scoped_http_only_cookie() { + let (app, alice, bob) = setup_app().await; + let private_version = save_and_deploy_with_visibility( + &app, + &alice, + "private-open", + "p", + "private", + ) + .await; + save_and_deploy_with_visibility(&app, &alice, "other-private", "o", "private") + .await; + + let unauthorized = app + .clone() + .oneshot( + Request::builder() + .method("POST") + .uri("/api/pages/private-open") + .header("content-type", "application/json") + .body(axum::body::Body::from("{}")) + .unwrap(), + ) + .await + .unwrap(); + assert_eq!(unauthorized.status(), StatusCode::UNAUTHORIZED); + + let foreign = app + .clone() + .oneshot( + Request::builder() + .method("POST") + .uri("/api/pages/private-open") + .header("Authorization", format!("Bearer {bob}")) + .header("content-type", "application/json") + .body(axum::body::Body::from("{}")) + .unwrap(), + ) + .await + .unwrap(); + assert_eq!(foreign.status(), StatusCode::NOT_FOUND); + + let ticket_response = app + .clone() + .oneshot( + Request::builder() + .method("POST") + .uri("/api/pages/private-open") + .header("Authorization", format!("Bearer {alice}")) + .header("content-type", "application/json") + .body(axum::body::Body::from("{}")) + .unwrap(), + ) + .await + .unwrap(); + assert_eq!(ticket_response.status(), StatusCode::OK); + let body = to_bytes(ticket_response.into_body(), usize::MAX) + .await + .unwrap(); + let json: serde_json::Value = serde_json::from_slice(&body).unwrap(); + let open_path = json["open_url_path"].as_str().unwrap(); + assert!(!open_path.contains(&alice)); + + let exchange = app + .clone() + .oneshot( + Request::builder() + .uri(open_path) + .header("x-forwarded-proto", "https") + .body(axum::body::Body::empty()) + .unwrap(), + ) + .await + .unwrap(); + assert_eq!(exchange.status(), StatusCode::TEMPORARY_REDIRECT); + assert_eq!( + exchange.headers().get(header::LOCATION).unwrap(), + "/p/alice/private-open" + ); + let cookie = exchange + .headers() + .get(header::SET_COOKIE) + .unwrap() + .to_str() + .unwrap() + .to_string(); + assert!(cookie.contains("HttpOnly")); + assert!(cookie.contains("SameSite=Lax")); + assert!(cookie.contains("Secure")); + assert!(cookie.contains("Path=/p/alice/private-open")); + assert!(!cookie.contains(&alice)); + let browser_cookie = cookie.split(';').next().unwrap(); + + let replay = app + .clone() + .oneshot( + Request::builder() + .uri(open_path) + .body(axum::body::Body::empty()) + .unwrap(), + ) + .await + .unwrap(); + assert_eq!(replay.status(), StatusCode::NOT_FOUND); + + let authorized = app + .clone() + .oneshot( + Request::builder() + .uri("/p/alice/private-open") + .header(header::COOKIE, browser_cookie) + .body(axum::body::Body::empty()) + .unwrap(), + ) + .await + .unwrap(); + assert_eq!(authorized.status(), StatusCode::OK); + + // A production ticket is not a blanket grant for immutable preview + // routes, even when the preview belongs to the same Page. + let wrong_version_scope = app + .clone() + .oneshot( + Request::builder() + .uri(format!("/p/alice/private-open/@v/{private_version}")) + .header(header::COOKIE, browser_cookie) + .body(axum::body::Body::empty()) + .unwrap(), + ) + .await + .unwrap(); + assert_eq!(wrong_version_scope.status(), StatusCode::NOT_FOUND); + + let wrong_page = app + .oneshot( + Request::builder() + .uri("/p/alice/other-private") + .header(header::COOKIE, browser_cookie) + .body(axum::body::Body::empty()) + .unwrap(), + ) + .await + .unwrap(); + assert_eq!(wrong_page.status(), StatusCode::NOT_FOUND); + } + + #[tokio::test] + async fn unpublish_stops_production_without_deleting_versions() { + let (app, alice, _) = setup_app().await; + let version_id = save_and_deploy(&app, &alice, "pause-me", "live").await; + + let response = app + .clone() + .oneshot( + Request::builder() + .method("POST") + .uri("/api/pages/pause-me/unpublish") + .header("Authorization", format!("Bearer {alice}")) + .body(axum::body::Body::empty()) + .unwrap(), + ) + .await + .unwrap(); + assert_eq!(response.status(), StatusCode::NO_CONTENT); + assert_eq!( + get_page(&app, "/p/alice/pause-me", None).await, + StatusCode::NOT_FOUND + ); + assert_eq!( + get_page(&app, &format!("/p/alice/pause-me/@v/{version_id}"), None,).await, + StatusCode::OK + ); + + let versions = app + .oneshot( + Request::builder() + .uri("/api/pages/pause-me/versions") + .header("Authorization", format!("Bearer {alice}")) + .body(axum::body::Body::empty()) + .unwrap(), + ) + .await + .unwrap(); + assert_eq!(versions.status(), StatusCode::OK); + let body = to_bytes(versions.into_body(), usize::MAX).await.unwrap(); + let list: serde_json::Value = serde_json::from_slice(&body).unwrap(); + assert_eq!(list.as_array().unwrap().len(), 1); + assert_eq!(list[0]["deployed"], false); + } + #[tokio::test] async fn worker_does_not_receive_credential_headers() { let (app, alice, _) = setup_app().await; @@ -1474,19 +2497,14 @@ mod tests { } "#; let files = [("server/worker.js", worker.as_bytes())]; - let mut map = serde_json::Map::new(); - for (path, content) in files { - let hash = hex_sha256(content); - map.insert( - path.to_string(), - serde_json::json!({ "content": B64.encode(content), "hash": hash }), - ); - } + let (upload_id, map) = begin_test_upload(&app, &alice, "hdr", &files).await; let upload = serde_json::json!({ "slug": "hdr", + "upload_id": upload_id, "title": "hdr", "visibility": "relay", - "files": map + "files": map, + "finalize": true }); let resp = app .clone() @@ -1510,7 +2528,9 @@ mod tests { .uri("/api/pages/hdr/versions") .header("Authorization", format!("Bearer {alice}")) .header("content-type", "application/json") - .body(axum::body::Body::from("{}")) + .body(axum::body::Body::from( + serde_json::json!({ "upload_id": upload_id }).to_string(), + )) .unwrap(), ) .await @@ -1574,4 +2594,21 @@ mod tests { assert_eq!(page["file_count"].as_i64().unwrap(), 1); assert_eq!(page["total_bytes"].as_i64().unwrap(), html.len() as i64); } + + #[tokio::test] + async fn page_request_body_limit_is_applied_before_worker_execution() { + let (app, _, _) = setup_app().await; + let body = vec![0u8; crate::page_execution::MAX_PAGE_FUNCTION_REQUEST_BODY_BYTES + 1]; + let response = app + .oneshot( + Request::builder() + .method("POST") + .uri("/p/alice/missing/api") + .body(axum::body::Body::from(body)) + .unwrap(), + ) + .await + .unwrap(); + assert_eq!(response.status(), StatusCode::PAYLOAD_TOO_LARGE); + } } diff --git a/src/crates/services/services-integrations/src/remote_connect.rs b/src/crates/services/services-integrations/src/remote_connect.rs index 79d467905a..4238749cad 100644 --- a/src/crates/services/services-integrations/src/remote_connect.rs +++ b/src/crates/services/services-integrations/src/remote_connect.rs @@ -50,11 +50,12 @@ pub use ngrok::{ cleanup_all_ngrok, detect_running_ngrok, is_ngrok_available, start_ngrok_tunnel, NgrokTunnel, }; pub use page_upload::{ - delete_page_version_on_relay, deploy_page_version_on_relay, join_relay_url, - list_page_versions_from_relay, list_pages_from_relay, publish_page_content_on_relay, - publish_page_to_relay, save_page_version_from_inline_files, save_page_version_to_relay, - unpublish_page_from_relay, update_page_on_relay, PageContentPublishResult, PageInfo, - PagePublishResult, PageSaveVersionResult, PageVersionInfo, + create_page_open_link_on_relay, delete_page_from_relay, delete_page_version_on_relay, + deploy_page_version_on_relay, join_relay_url, list_page_versions_from_relay, + list_pages_from_relay, publish_page_content_on_relay, publish_page_to_relay, + save_page_version_from_inline_files, save_page_version_to_relay, unpublish_page_from_relay, + update_page_on_relay, PageContentPublishResult, PageInfo, PageOpenLink, PagePublishResult, + PageSaveVersionResult, PageVersionInfo, }; pub use pairing::{PairingChallenge, PairingProtocol, PairingResponse, PairingState, QrPayload}; pub use qr_generator::QrGenerator; diff --git a/src/crates/services/services-integrations/src/remote_connect/page_upload.rs b/src/crates/services/services-integrations/src/remote_connect/page_upload.rs index cdf4f0fa04..b6b84c4b4f 100644 --- a/src/crates/services/services-integrations/src/remote_connect/page_upload.rs +++ b/src/crates/services/services-integrations/src/remote_connect/page_upload.rs @@ -9,6 +9,7 @@ use std::path::Path; const MAX_UPLOAD_BATCH_BASE64_BYTES: usize = 256 * 1024; const MAX_PAGE_BYTES: u64 = 100 * 1024 * 1024; const MAX_FILE_BYTES: u64 = 10 * 1024 * 1024; +const MAX_PAGE_FILES: usize = 4096; #[derive(Debug, Clone, PartialEq, Eq, Serialize)] struct PageUploadManifestEntry { @@ -53,6 +54,18 @@ pub struct PageVersionInfo { pub preview_url_path: String, } +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct PageOpenLink { + pub open_url: String, + pub expires_in_seconds: u64, +} + +#[derive(Debug, Deserialize)] +struct PageOpenLinkRelayResponse { + open_url_path: String, + expires_in_seconds: u64, +} + #[derive(Debug, Clone, Serialize, Deserialize)] pub struct PageSaveVersionResult { pub slug: String, @@ -270,12 +283,17 @@ async fn save_page_version_from_collected_files( size: f.content.len() as u64, }) .collect(); + let upload_id = uuid::Uuid::new_v4().simple().to_string(); let check_url = format!("{relay_base}/api/pages/check-files"); let check_resp = client .post(&check_url) .header("Authorization", &auth) - .json(&serde_json::json!({ "slug": slug, "files": manifest })) + .json(&serde_json::json!({ + "slug": slug, + "upload_id": upload_id, + "files": manifest, + })) .timeout(std::time::Duration::from_secs(30)) .send() .await @@ -289,6 +307,9 @@ async fn save_page_version_from_collected_files( .json() .await .map_err(|e| anyhow!("parse check-files response: {e}"))?; + if check_body["upload_id"].as_str() != Some(upload_id.as_str()) { + return Err(anyhow!("check-files returned a mismatched upload session")); + } let needed: Vec = check_body["needed"] .as_array() .map(|items| { @@ -301,7 +322,7 @@ async fn save_page_version_from_collected_files( if !needed.is_empty() { upload_needed_page_files( - &client, relay_base, &auth, slug, &title, visibility, &all_files, &needed, + &client, relay_base, &auth, slug, &upload_id, &title, visibility, &all_files, &needed, ) .await?; } else { @@ -311,6 +332,7 @@ async fn save_page_version_from_collected_files( &format!("{relay_base}/api/pages/upload-files"), &auth, slug, + &upload_id, &title, visibility, &HashMap::new(), @@ -325,6 +347,7 @@ async fn save_page_version_from_collected_files( .post(&freeze_url) .header("Authorization", &auth) .json(&serde_json::json!({ + "upload_id": upload_id, "title": title, "note": note.unwrap_or(""), })) @@ -418,6 +441,43 @@ pub async fn list_page_versions_from_relay( .map_err(|e| anyhow!("parse list versions: {e}"))?) } +/// Create a short-lived, one-time browser handoff URL for a production Page or +/// a specific immutable preview. The account bearer token is sent only in the +/// authenticated request header and is never embedded in the returned URL. +pub async fn create_page_open_link_on_relay( + relay_url: &str, + token: &str, + slug: &str, + version_id: Option<&str>, +) -> Result { + validate_slug(slug)?; + let client = reqwest::Client::new(); + let url = format!("{}/api/pages/{}", relay_url.trim_end_matches('/'), slug); + let resp = client + .post(&url) + .header("Authorization", format!("Bearer {token}")) + .json(&serde_json::json!({ "version_id": version_id })) + .timeout(std::time::Duration::from_secs(15)) + .send() + .await + .map_err(|e| anyhow!("create page open link failed: {e}"))?; + if !resp.status().is_success() { + let status = resp.status(); + let body = resp.text().await.unwrap_or_default(); + return Err(anyhow!( + "create page open link failed: HTTP {status} — {body}" + )); + } + let result: PageOpenLinkRelayResponse = resp + .json() + .await + .map_err(|e| anyhow!("parse page open link response: {e}"))?; + Ok(PageOpenLink { + open_url: join_relay_url(relay_url, &result.open_url_path), + expires_in_seconds: result.expires_in_seconds, + }) +} + pub async fn deploy_page_version_on_relay( relay_url: &str, token: &str, @@ -521,9 +581,13 @@ pub async fn update_page_on_relay( pub async fn unpublish_page_from_relay(relay_url: &str, token: &str, slug: &str) -> Result<()> { validate_slug(slug)?; let client = reqwest::Client::new(); - let url = format!("{}/api/pages/{}", relay_url.trim_end_matches('/'), slug); + let url = format!( + "{}/api/pages/{}/unpublish", + relay_url.trim_end_matches('/'), + slug + ); let resp = client - .delete(&url) + .post(&url) .header("Authorization", format!("Bearer {token}")) .timeout(std::time::Duration::from_secs(15)) .send() @@ -537,6 +601,25 @@ pub async fn unpublish_page_from_relay(relay_url: &str, token: &str, slug: &str) Ok(()) } +pub async fn delete_page_from_relay(relay_url: &str, token: &str, slug: &str) -> Result<()> { + validate_slug(slug)?; + let client = reqwest::Client::new(); + let url = format!("{}/api/pages/{}", relay_url.trim_end_matches('/'), slug); + let resp = client + .delete(&url) + .header("Authorization", format!("Bearer {token}")) + .timeout(std::time::Duration::from_secs(15)) + .send() + .await + .map_err(|e| anyhow!("delete page failed: {e}"))?; + if !resp.status().is_success() { + let status = resp.status(); + let body = resp.text().await.unwrap_or_default(); + return Err(anyhow!("delete page failed: HTTP {status} — {body}")); + } + Ok(()) +} + fn validate_slug(slug: &str) -> Result<()> { let bytes = slug.as_bytes(); if bytes.is_empty() || bytes.len() > 64 { @@ -571,19 +654,22 @@ fn validate_visibility(v: &str) -> Result<()> { } fn collect_page_files(base: &Path) -> Result> { - if !base.is_dir() { + let base_metadata = std::fs::symlink_metadata(base) + .map_err(|_| anyhow!("page directory does not exist: {}", base.display()))?; + if base_metadata.file_type().is_symlink() || !base_metadata.is_dir() { return Err(anyhow!("page directory does not exist: {}", base.display())); } - let has_index = base.join("index.html").exists(); - let has_worker = base.join("server").join("worker.js").exists(); + let has_index = is_regular_page_source_file(&base.join("index.html"))?; + let has_worker = is_regular_page_source_file(&base.join("server").join("worker.js"))?; if !has_index && !has_worker { return Err(anyhow!( "page directory must contain index.html and/or server/worker.js: {}", base.display() )); } + let canonical_base = std::fs::canonicalize(base)?; let mut all_files = Vec::new(); - collect_files_with_hash(base, base, &mut all_files)?; + collect_files_with_hash(&canonical_base, &canonical_base, &mut all_files)?; Ok(all_files) } @@ -593,6 +679,12 @@ fn collect_inline_page_files(files: &HashMap) -> Result MAX_PAGE_FILES { + return Err(anyhow!( + "page exceeds file count limit ({} > {MAX_PAGE_FILES})", + files.len() + )); + } let has_index = files.contains_key("index.html"); let has_worker = files.contains_key("server/worker.js"); @@ -640,18 +732,35 @@ fn collect_files_with_hash( for entry in std::fs::read_dir(dir)? { let entry = entry?; let path = entry.path(); - if path.is_dir() { - collect_files_with_hash(base, &path, out)?; - } else if path.is_file() { - let rel = path + let file_type = entry.file_type()?; + if file_type.is_symlink() { + return Err(anyhow!( + "symbolic links are not allowed in Page sources: {}", + path.display() + )); + } + let canonical_path = std::fs::canonicalize(&path)?; + if !canonical_path.starts_with(base) { + return Err(anyhow!( + "Page source entry resolves outside its directory: {}", + path.display() + )); + } + if file_type.is_dir() { + collect_files_with_hash(base, &canonical_path, out)?; + } else if file_type.is_file() { + if out.len() >= MAX_PAGE_FILES { + return Err(anyhow!("page exceeds file count limit ({MAX_PAGE_FILES})")); + } + let rel = canonical_path .strip_prefix(base) - .unwrap_or(&path) + .unwrap_or(&canonical_path) .to_string_lossy() .replace('\\', "/"); if rel.contains("..") { continue; } - let content = std::fs::read(&path)?; + let content = std::fs::read(&canonical_path)?; if content.len() as u64 > MAX_FILE_BYTES { return Err(anyhow!( "file exceeds size limit: {rel} ({} > {MAX_FILE_BYTES} bytes)", @@ -666,17 +775,35 @@ fn collect_files_with_hash( content, hash, }); + } else { + return Err(anyhow!( + "unsupported Page source entry type: {}", + path.display() + )); } } Ok(()) } +fn is_regular_page_source_file(path: &Path) -> Result { + match std::fs::symlink_metadata(path) { + Ok(metadata) if metadata.file_type().is_symlink() => Err(anyhow!( + "symbolic links are not allowed in Page sources: {}", + path.display() + )), + Ok(metadata) => Ok(metadata.is_file()), + Err(error) if error.kind() == std::io::ErrorKind::NotFound => Ok(false), + Err(error) => Err(error.into()), + } +} + #[allow(clippy::too_many_arguments)] async fn upload_needed_page_files( client: &reqwest::Client, relay_base: &str, auth: &str, slug: &str, + upload_id: &str, title: &str, visibility: &str, all_files: &[CollectedPageFile], @@ -710,6 +837,7 @@ async fn upload_needed_page_files( &url, auth, slug, + upload_id, title, visibility, ¤t_batch, @@ -730,6 +858,7 @@ async fn upload_needed_page_files( &url, auth, slug, + upload_id, title, visibility, ¤t_batch, @@ -747,6 +876,7 @@ async fn post_upload_batch( url: &str, auth: &str, slug: &str, + upload_id: &str, title: &str, visibility: &str, files: &HashMap, @@ -758,6 +888,7 @@ async fn post_upload_batch( .header("Authorization", auth) .json(&serde_json::json!({ "slug": slug, + "upload_id": upload_id, "title": title, "visibility": visibility, "files": files, @@ -835,6 +966,42 @@ mod tests { assert!(collect_inline_page_files(&empty_entry).is_err()); } + #[cfg(unix)] + #[test] + fn directory_source_rejects_symlinks_to_external_files() { + use std::os::unix::fs::symlink; + + let root = std::env::temp_dir().join(format!( + "bitfun-page-symlink-external-{}", + uuid::Uuid::new_v4() + )); + let page = root.join("page"); + std::fs::create_dir_all(&page).unwrap(); + std::fs::write(page.join("index.html"), b"").unwrap(); + std::fs::write(root.join("secret.txt"), b"secret").unwrap(); + symlink(root.join("secret.txt"), page.join("secret.txt")).unwrap(); + + let error = collect_page_files(&page).expect_err("external symlink must be rejected"); + assert!(error.to_string().contains("symbolic links are not allowed")); + let _ = std::fs::remove_dir_all(root); + } + + #[cfg(unix)] + #[test] + fn directory_source_rejects_symlink_loops() { + use std::os::unix::fs::symlink; + + let root = + std::env::temp_dir().join(format!("bitfun-page-symlink-loop-{}", uuid::Uuid::new_v4())); + std::fs::create_dir_all(&root).unwrap(); + std::fs::write(root.join("index.html"), b"").unwrap(); + symlink(&root, root.join("loop")).unwrap(); + + let error = collect_page_files(&root).expect_err("symlink loop must be rejected"); + assert!(error.to_string().contains("symbolic links are not allowed")); + let _ = std::fs::remove_dir_all(root); + } + #[tokio::test] async fn publish_source_xor_is_enforced() { let err = publish_page_content_on_relay( diff --git a/src/mobile-web/src/i18n/I18nProvider.tsx b/src/mobile-web/src/i18n/I18nProvider.tsx index 6bbb16a78c..25dd58c26b 100644 --- a/src/mobile-web/src/i18n/I18nProvider.tsx +++ b/src/mobile-web/src/i18n/I18nProvider.tsx @@ -17,6 +17,7 @@ interface I18nContextValue { toggleLanguage: () => void; t: (key: string, params?: TranslateParams) => string; formatDate: (date: Date | number, options?: Intl.DateTimeFormatOptions) => string; + formatRelativeTime: (date: Date | number) => string; } const STORAGE_KEY = 'bitfun-mobile-language'; @@ -99,8 +100,22 @@ export const I18nContext = createContext({ toggleLanguage: () => {}, t: (key) => key, formatDate: (date, options) => new Intl.DateTimeFormat(DEFAULT_LANGUAGE, options).format(date), + formatRelativeTime: (date) => formatRelativeTime(DEFAULT_LANGUAGE, date), }); +function formatRelativeTime(language: MobileLanguage, date: Date | number): string { + const timestamp = date instanceof Date ? date.getTime() : date; + const diffSeconds = (timestamp - Date.now()) / 1000; + const absoluteSeconds = Math.abs(diffSeconds); + const formatter = new Intl.RelativeTimeFormat(language, { numeric: 'auto' }); + + if (absoluteSeconds < 60) return formatter.format(Math.round(diffSeconds), 'second'); + if (absoluteSeconds < 3600) return formatter.format(Math.round(diffSeconds / 60), 'minute'); + if (absoluteSeconds < 86_400) return formatter.format(Math.round(diffSeconds / 3600), 'hour'); + if (absoluteSeconds < 2_592_000) return formatter.format(Math.round(diffSeconds / 86_400), 'day'); + return new Intl.DateTimeFormat(language, { dateStyle: 'medium', timeStyle: 'short' }).format(timestamp); +} + export const I18nProvider: React.FC<{ children: React.ReactNode }> = ({ children }) => { const [language, setLanguageState] = useState(detectInitialLanguage); @@ -127,6 +142,7 @@ export const I18nProvider: React.FC<{ children: React.ReactNode }> = ({ children toggleLanguage, t: (key, params) => translate(language, key, params), formatDate: (date, options) => new Intl.DateTimeFormat(language, options).format(date), + formatRelativeTime: (date) => formatRelativeTime(language, date), }), [language, setLanguage, toggleLanguage]); return ( @@ -137,4 +153,3 @@ export const I18nProvider: React.FC<{ children: React.ReactNode }> = ({ children }; export type { MobileLanguage, TranslateParams }; - diff --git a/src/mobile-web/src/i18n/messages.ts b/src/mobile-web/src/i18n/messages.ts index d0a17dbf84..4488fd1bfb 100644 --- a/src/mobile-web/src/i18n/messages.ts +++ b/src/mobile-web/src/i18n/messages.ts @@ -176,6 +176,7 @@ export const messages: Record = { noDevices: 'No devices found', online: 'Online', offline: 'Offline', + lastSeen: 'Last seen {time}', current: 'Current', pairedDesktop: 'Paired', switchFailed: 'Failed to switch device', @@ -364,6 +365,7 @@ export const messages: Record = { noDevices: '未找到设备', online: '在线', offline: '离线', + lastSeen: '上次在线:{time}', current: '当前', pairedDesktop: '已配对', switchFailed: '切换设备失败', @@ -552,6 +554,7 @@ export const messages: Record = { noDevices: '未找到設備', online: '在線', offline: '離線', + lastSeen: '上次在線:{time}', current: '目前', pairedDesktop: '已配對', switchFailed: '切換設備失敗', diff --git a/src/mobile-web/src/pages/DevicesPage.tsx b/src/mobile-web/src/pages/DevicesPage.tsx index b5ac28f310..3325bea0e9 100644 --- a/src/mobile-web/src/pages/DevicesPage.tsx +++ b/src/mobile-web/src/pages/DevicesPage.tsx @@ -57,7 +57,7 @@ const NoIdentityIcon = () => ( ); const DevicesPage: React.FC = ({ client, onBack }) => { - const { t } = useI18n(); + const { t, formatRelativeTime } = useI18n(); const { setControlTarget, resetForDeviceSwitch } = useMobileStore(); const [devices, setDevices] = useState([]); const [identityReady, setIdentityReady] = useState(client.hasDelegatedIdentity); @@ -270,7 +270,11 @@ const DevicesPage: React.FC = ({ client, onBack }) => { - {d.online ? t('devices.online') : t('devices.offline')} + {d.online + ? t('devices.online') + : d.last_seen_at + ? t('devices.lastSeen', { time: formatRelativeTime(d.last_seen_at * 1000) }) + : t('devices.offline')} {d.device_id.slice(0, 8)} diff --git a/src/mobile-web/src/styles/components/language-toggle.scss b/src/mobile-web/src/styles/components/language-toggle.scss index 7b322b0388..c0fae2ec9e 100644 --- a/src/mobile-web/src/styles/components/language-toggle.scss +++ b/src/mobile-web/src/styles/components/language-toggle.scss @@ -2,8 +2,8 @@ display: inline-flex; align-items: center; justify-content: center; - min-width: 32px; - height: 32px; + min-width: 44px; + height: 44px; padding: 0 10px; border: 1px solid var(--border-subtle); border-radius: 999px; @@ -21,4 +21,3 @@ background: var(--element-bg-base); } } - diff --git a/src/mobile-web/src/styles/components/pairing.scss b/src/mobile-web/src/styles/components/pairing.scss index d70ad7143a..327b63399c 100644 --- a/src/mobile-web/src/styles/components/pairing.scss +++ b/src/mobile-web/src/styles/components/pairing.scss @@ -6,16 +6,21 @@ flex-direction: column; align-items: center; justify-content: center; - height: 100%; + min-height: 100%; + min-height: 100dvh; gap: var(--size-gap-6); - padding: var(--size-gap-8); + padding: + max(var(--size-gap-8), env(safe-area-inset-top, 0px)) + max(var(--size-gap-8), env(safe-area-inset-right, 0px)) + max(var(--size-gap-8), env(safe-area-inset-bottom, 0px)) + max(var(--size-gap-8), env(safe-area-inset-left, 0px)); animation: fadeIn var(--motion-slow) motion.$easing-decelerate; } .pairing-page__actions { position: absolute; - top: var(--size-gap-4); - right: var(--size-gap-4); + top: calc(var(--size-gap-4) + env(safe-area-inset-top, 0px)); + right: calc(var(--size-gap-4) + env(safe-area-inset-right, 0px)); display: flex; align-items: center; gap: var(--size-gap-2); @@ -25,8 +30,8 @@ display: inline-flex; align-items: center; justify-content: center; - width: 32px; - height: 32px; + width: 44px; + height: 44px; padding: 0; border: 1px solid var(--border-subtle); border-radius: 50%; diff --git a/src/mobile-web/src/styles/global.scss b/src/mobile-web/src/styles/global.scss index c46f836c7c..be2bac3361 100644 --- a/src/mobile-web/src/styles/global.scss +++ b/src/mobile-web/src/styles/global.scss @@ -9,6 +9,7 @@ html, body, #root { height: 100%; + height: 100dvh; width: 100%; overflow: hidden; } @@ -25,6 +26,7 @@ body { .mobile-app { height: 100%; + height: 100dvh; width: 100%; background: var(--color-bg-primary); color: var(--color-text-primary); diff --git a/src/web-ui/src/app/components/NavPanel/MainNav.tsx b/src/web-ui/src/app/components/NavPanel/MainNav.tsx index b25e71aaf7..82b4822a2d 100644 --- a/src/web-ui/src/app/components/NavPanel/MainNav.tsx +++ b/src/web-ui/src/app/components/NavPanel/MainNav.tsx @@ -13,7 +13,7 @@ import React, { useCallback, useState, useMemo, useEffect, useRef } from 'react'; import { createPortal } from 'react-dom'; -import { Plus, FolderOpen, FolderPlus, History, Check, User, Users, Puzzle, Blocks, ChevronDown, Search } from 'lucide-react'; +import { Plus, FolderOpen, FolderPlus, History, Check, User, Users, Puzzle, Blocks, ChevronDown, Search, PanelsTopLeft } from 'lucide-react'; import { Tooltip } from '@/component-library'; import { useApp } from '../../hooks/useApp'; import { useSceneManager } from '../../hooks/useSceneManager'; @@ -75,6 +75,7 @@ const MainNav: React.FC = ({ const activeTabId = useSceneStore(s => s.activeTabId); const setSelectedAssistantWorkspaceId = useMyAgentStore((s) => s.setSelectedAssistantWorkspaceId); const { t } = useI18n('common'); + const { t: tPages } = useI18n('scenes/pages'); const { currentWorkspace, loading: workspaceLoading, @@ -680,6 +681,16 @@ const MainNav: React.FC = ({ {/* ── Bottom: MiniApp ───────────────────────── */}
+
= ({ const { enterPeerMode } = usePeerDeviceMode(); const syncStatus = useAccountSyncStore((s) => s.status); const syncProgress = useAccountSyncStore((s) => s.progress); + const lastSyncIsFirstLogin = useAccountSyncStore((s) => s.lastSyncIsFirstLogin); const setSyncing = useAccountSyncStore((s) => s.setSyncing); const setSyncDone = useAccountSyncStore((s) => s.setDone); const setSyncFailed = useAccountSyncStore((s) => s.setFailed); + const clearSync = useAccountSyncStore((s) => s.clear); const [username, setUsername] = useState(''); const [password, setPassword] = useState(''); @@ -154,16 +161,19 @@ export const AccountPanel: React.FC = ({ const handleCopyRelayUrl = useCallback(async () => { if (!accountRelayUrl) return; - try { - await navigator.clipboard.writeText(accountRelayUrl); + const copied = await copyTextToClipboard(accountRelayUrl); + if (copied) { setCopiedServerUrl(true); window.setTimeout(() => setCopiedServerUrl(false), 1500); - } catch (e) { - log.warn('copy relay url failed', e); + } else { + warning(t('accountLogin.copyServerFailed')); } - }, [accountRelayUrl]); + }, [accountRelayUrl, t, warning]); const handleSessionExpired = useCallback(async (_error: unknown) => { + // Invalidate detached retries before the logout request yields control. + syncInFlightRef.current = false; + clearSync(); try { await remoteConnectAPI.accountLogout(); } catch (e) { @@ -172,7 +182,7 @@ export const AccountPanel: React.FC = ({ resetState(); setView('login'); setError(t('accountLogin.sessionExpired')); - }, [resetState, t]); + }, [clearSync, resetState, t]); const markRelayUnreachable = useCallback(() => { setDevicesReady(false); @@ -292,12 +302,14 @@ export const AccountPanel: React.FC = ({ useEffect(() => { return () => { if (viewRef.current === 'overwrite') { + syncInFlightRef.current = false; + clearSync(); void remoteConnectAPI.accountLogout().catch((e) => { log.warn('logout on overwrite abandon failed', e); }); } }; - }, []); + }, [clearSync]); useEffect(() => { remoteConnectAPI.getDeviceInfo().then((info) => { @@ -359,7 +371,11 @@ export const AccountPanel: React.FC = ({ } syncInFlightRef.current = true; ensureAccountSyncProgressListener(); - setSyncing(); + setSyncing(isFirstLogin); + const operationId = useAccountSyncStore.getState().operationId; + const isCurrentOperation = () => ( + useAccountSyncStore.getState().operationId === operationId + ); info(t('accountLogin.syncStarted')); // Connect device presence immediately so the device list can populate @@ -373,6 +389,7 @@ export const AccountPanel: React.FC = ({ let configJson = '{}'; if (isFirstLogin) { useAccountSyncStore.getState().applyProgress({ + operation_id: operationId, phase: 'uploading_settings', percent: 2, }); @@ -382,22 +399,32 @@ export const AccountPanel: React.FC = ({ } catch (e) { log.warn('export config failed', e); } + if (!isCurrentOperation()) return; } const wp = workspacePath || '/'; const maxAttempts = 3; let result: Awaited> | null = null; let lastError: unknown = null; for (let attempt = 1; attempt <= maxAttempts; attempt += 1) { + if (!isCurrentOperation()) return; try { - result = await remoteConnectAPI.accountAutoSync(isFirstLogin, wp, configJson); + result = await remoteConnectAPI.accountAutoSync( + isFirstLogin, + wp, + configJson, + operationId, + ); + if (!isCurrentOperation()) return; lastError = null; break; } catch (e) { + if (!isCurrentOperation()) return; lastError = e; log.warn(`Auto-sync attempt ${attempt}/${maxAttempts} failed`, e); if (attempt < maxAttempts) { info(t('accountLogin.syncRetrying', { attempt, max: maxAttempts })); await new Promise((resolve) => setTimeout(resolve, 2000 * attempt)); + if (!isCurrentOperation()) return; } } } @@ -406,28 +433,35 @@ export const AccountPanel: React.FC = ({ ? lastError : new Error(String(lastError ?? 'auto-sync failed')); } + if (!isCurrentOperation()) return; log.info( `Auto-sync done: settings=${result.settings_synced} exported=${result.sessions_exported}`, ); if (result.settings_synced && !isFirstLogin) { + if (!isCurrentOperation()) return; try { await configAPI.reloadConfig(); + if (!isCurrentOperation()) return; configManager.clearCache(); success(t('accountLogin.settingsApplied')); } catch (e) { log.warn('reloadConfig after sync failed', e); } } + if (!isCurrentOperation()) return; setSyncDone(result); success(t('accountLogin.syncDone', { exported: result.sessions_exported, })); } catch (e) { + if (!isCurrentOperation()) return; log.error('Auto-sync failed', e); setSyncFailed(e instanceof Error ? e.message : String(e)); warning(t('accountLogin.syncFailed')); } finally { - syncInFlightRef.current = false; + if (isCurrentOperation()) { + syncInFlightRef.current = false; + } } })(); }, [ @@ -441,6 +475,11 @@ export const AccountPanel: React.FC = ({ workspacePath, ]); + const handleRetrySync = useCallback(() => { + if (syncStatus !== 'failed' || syncInFlightRef.current) return; + startBackgroundSync(lastSyncIsFirstLogin ?? false); + }, [lastSyncIsFirstLogin, startBackgroundSync, syncStatus]); + /** Landing path after a completed login: devices view + background sync. */ const completeLogin = useCallback((relayUrl: string, isFirstLogin: boolean) => { setAccountRelayUrl(relayUrl); @@ -462,7 +501,12 @@ export const AccountPanel: React.FC = ({ completeLogin(server, true); } catch (e: unknown) { setError(e instanceof Error ? e.message : String(e)); - } finally { setLoading(false); } + } finally { + // The account session has its own token after this call; retaining the + // password in React state while the device list is open is unnecessary. + setPassword(''); + setLoading(false); + } }, [completeLogin, success, t]); const handleLogin = useCallback(async () => { @@ -470,10 +514,16 @@ export const AccountPanel: React.FC = ({ const relayUrl = parseRelayServer(authServer); if (!relayUrl) return; const isLoopback = ['localhost', '127.0.0.1', '[::1]', '::1'].includes(relayUrl.hostname); - if (relayUrl.protocol === 'http:' - && !isLoopback - && !window.confirm(t('accountLogin.insecureServerConfirm'))) { - return; + if (relayUrl.protocol === 'http:' && !isLoopback) { + const confirmed = await confirmWarning( + t('accountLogin.insecureServerTitle'), + t('accountLogin.insecureServerConfirm'), + { + confirmText: t('accountLogin.continueInsecure'), + cancelText: t('accountLogin.cancel'), + }, + ); + if (!confirmed) return; } await performLogin(authServer.trim(), username.trim(), password); }, [validate, authServer, username, password, performLogin, t]); @@ -500,9 +550,12 @@ export const AccountPanel: React.FC = ({ } catch (e: unknown) { if (isAccountAuthFailure(e)) { await handleSessionExpired(e); - } else { - setError(e instanceof Error ? e.message : String(e)); + return; } + setError(e instanceof Error ? e.message : String(e)); + // Stop any detached work before accountLogout can yield. + syncInFlightRef.current = false; + clearSync(); try { await remoteConnectAPI.accountLogout(); } catch (logoutErr) { log.warn('logout after finalize failure failed', logoutErr); } @@ -511,7 +564,7 @@ export const AccountPanel: React.FC = ({ } finally { setLoading(false); } - }, [authServer, completeLogin, handleSessionExpired, resetState, success, t, username]); + }, [authServer, clearSync, completeLogin, handleSessionExpired, resetState, success, t, username]); const handleConfirmOverwrite = useCallback(() => { void finalizeAndSync(false); @@ -522,13 +575,17 @@ export const AccountPanel: React.FC = ({ }, [finalizeAndSync]); const handleCancelOverwrite = useCallback(async () => { + syncInFlightRef.current = false; + clearSync(); try { await remoteConnectAPI.accountLogout(); } catch (e) { log.warn('logout failed', e); } resetState(); setView('login'); - }, [resetState]); + }, [clearSync, resetState]); const handleLogout = useCallback(async () => { setLoading(true); + syncInFlightRef.current = false; + clearSync(); try { await remoteConnectAPI.accountLogout(); resetState(); @@ -536,16 +593,36 @@ export const AccountPanel: React.FC = ({ } catch (e: unknown) { setError(e instanceof Error ? e.message : String(e)); } finally { setLoading(false); } - }, [resetState]); + }, [clearSync, resetState]); const handleDeleteDevice = useCallback(async (deviceId: string, deviceName: string) => { const isLocal = localDeviceId === deviceId; const confirmation = isLocal ? t('accountLogin.confirmRemoveCurrentDevice', { name: deviceName }) : t('accountLogin.confirmRemoveDevice', { name: deviceName }); - if (!window.confirm(confirmation)) return; + const confirmed = await confirmDanger( + isLocal + ? t('accountLogin.removeCurrentDevice') + : t('accountLogin.removeDevice'), + confirmation, + { + confirmText: isLocal + ? t('accountLogin.removeCurrentDevice') + : t('accountLogin.removeDevice'), + cancelText: t('accountLogin.cancel'), + }, + ); + if (!confirmed) return; + const previousSyncStatus = syncStatus; + const previousSyncDirection = lastSyncIsFirstLogin; setLoading(true); setError(null); + if (isLocal) { + // A current-device removal is also a logout. Invalidate retries and + // late progress before the backend request yields. + syncInFlightRef.current = false; + clearSync(); + } try { await remoteConnectAPI.accountDeleteDevice(deviceId); if (isLocal) { @@ -560,12 +637,35 @@ export const AccountPanel: React.FC = ({ if (isAccountAuthFailure(e)) { await handleSessionExpired(e); } else { - setError(e instanceof Error ? e.message : String(e)); + const message = e instanceof Error ? e.message : String(e); + setError(message); + if ( + isLocal + && previousSyncDirection !== null + && (previousSyncStatus === 'syncing' || previousSyncStatus === 'failed') + ) { + // Preserve the direction so Retry remains meaningful after a failed + // current-device removal invalidated the previous generation. + setSyncing(previousSyncDirection); + setSyncFailed(message); + } } } finally { setLoading(false); } - }, [handleSessionExpired, localDeviceId, refreshDevices, resetState, success, t]); + }, [ + clearSync, + handleSessionExpired, + lastSyncIsFirstLogin, + localDeviceId, + refreshDevices, + resetState, + setSyncFailed, + setSyncing, + success, + syncStatus, + t, + ]); const selectDevice = useCallback(async (device: AccountDeviceInfo) => { if (!device.online) return; @@ -574,6 +674,9 @@ export const AccountPanel: React.FC = ({ info(t('accountLogin.syncInProgressHint')); return; } + if (syncStatus === 'failed') { + warning(t('accountLogin.syncFailedPeerHint')); + } setLoading(true); setError(null); try { @@ -585,7 +688,7 @@ export const AccountPanel: React.FC = ({ } finally { setLoading(false); } - }, [enterPeerMode, info, localDeviceId, onCloseDialog, success, syncStatus, t]); + }, [enterPeerMode, info, localDeviceId, onCloseDialog, success, syncStatus, t, warning]); return ( <> @@ -731,6 +834,17 @@ export const AccountPanel: React.FC = ({ {syncStatus === 'done' && t('accountLogin.syncDoneShort')} {syncStatus === 'failed' && t('accountLogin.syncFailed')} + {syncStatus === 'failed' && ( + + )} {syncStatus === 'syncing' && ( {t('accountLogin.syncProgressPercent', { percent: syncProgress.percent })} @@ -781,43 +895,48 @@ export const AccountPanel: React.FC = ({ const displayName = d.device_name || t('accountLogin.unknownDevice'); return (
isSelectable && selectDevice(d)} - onKeyDown={(event) => { - if (isSelectable && (event.key === 'Enter' || event.key === ' ')) { - event.preventDefault(); - void selectDevice(d); - } - }} - role={isSelectable ? 'button' : undefined} - tabIndex={isSelectable ? 0 : undefined}> - -
- - {displayName} - {isLocal && {t('accountLogin.thisDevice')}} - - - - {d.device_id.slice(0, 8)} + className={`account-panel__device-card ${isSelectable ? 'selectable' : ''} ${d.online ? '' : 'offline'} ${isLocal ? 'current' : ''} ${syncStatus === 'syncing' && !isLocal ? 'syncing' : ''}`}> +
- {!isLocal && d.online && syncStatus !== 'syncing' && } - {!isLocal && d.online && syncStatus === 'syncing' && ( - - )} - +
); + const handleCopyPairingUrl = useCallback(async () => { + if (!connectionResult?.qr_url) return; + const copied = await copyTextToClipboard(connectionResult.qr_url); + if (copied) { + setQrCopied(true); + window.setTimeout(() => setQrCopied(false), 2000); + } else { + notifyError(t('remoteConnect.copyUrlFailed')); + } + }, [connectionResult?.qr_url, notifyError, t]); + const renderPairingInProgress = () => { if (!connectionResult) return null; return (
{connectionResult.qr_url && ( -
{ - navigator.clipboard.writeText(connectionResult.qr_url!); - setQrCopied(true); - setTimeout(() => setQrCopied(false), 2000); - }} + title={t('remoteConnect.copyUrl')} + aria-label={t('remoteConnect.copyUrl')} + onClick={() => void handleCopyPairingUrl()} > -
+ )} {connectionResult.bot_pairing_code && (
@@ -939,14 +976,14 @@ export const RemoteConnectDialog: React.FC = ({ {isWeixinRasterQrSrc(weixinQrImageUrl) ? ( WeChat QR ) : (
= ({ >
{/* ── Group tabs ── */} -
+
= ({ >
) : ( -
+
{BOT_TABS.map((tab, i) => ( {i > 0 && } + ), + Input: (props: React.InputHTMLAttributes) => , + Select: () =>
, + confirmDanger: vi.fn(), + confirmWarning: vi.fn(), +})); + +vi.mock('@/app/components', () => ({ + GalleryLayout: ({ children }: { children: React.ReactNode }) =>
{children}
, + GalleryPageHeader: ({ title }: { title: React.ReactNode }) =>
{title}
, + GalleryEmpty: ({ message, action, testId }: { message: React.ReactNode; action?: React.ReactNode; testId?: string }) => ( +
{message}{action}
+ ), +})); + +describe('PagesScene initial loading', () => { + let container: HTMLDivElement; + let root: Root; + + beforeEach(() => { + container = document.createElement('div'); + document.body.appendChild(container); + root = createRoot(container); + mocks.accountStatus.mockReset().mockResolvedValue({ logged_in: true, user_id: 'u1' }); + mocks.accountGetCredentialHint.mockReset().mockResolvedValue({ relay_url: 'https://relay.test' }); + mocks.listPages.mockReset().mockRejectedValue(new Error('relay unavailable')); + mocks.createOpenLink.mockReset(); + mocks.update.mockReset(); + }); + + afterEach(() => { + act(() => root.unmount()); + container.remove(); + }); + + it('attempts a failed initial load only once until the user retries', async () => { + await act(async () => { + root.render(); + await Promise.resolve(); + await new Promise((resolve) => setTimeout(resolve, 0)); + }); + await act(async () => { + await Promise.resolve(); + await new Promise((resolve) => setTimeout(resolve, 0)); + }); + + // The failed relay call triggers one bounded auth re-check so an expired + // session can switch to the sign-in state; it must not retry listPages. + expect(mocks.accountStatus).toHaveBeenCalledTimes(2); + expect(mocks.listPages).toHaveBeenCalledTimes(1); + expect(container.querySelector('[data-testid="pages-error"]')).not.toBeNull(); + }); + + it('locks every action on one Page while an operation is pending and exposes title editing', async () => { + mocks.listPages.mockResolvedValue([{ + slug: 'demo', + visibility: 'public', + title: 'Demo', + file_count: 1, + total_bytes: 20, + created_at: 1, + updated_at: 1, + url_path: '/p/alice/demo', + preview_url_path: '/p/alice/demo/@v/v1', + deployed_version_id: 'v1', + }]); + let resolveOpenLink: ((value: { open_url: string; expires_in_seconds: number }) => void) | undefined; + mocks.createOpenLink.mockImplementation(() => new Promise((resolve) => { + resolveOpenLink = resolve; + })); + + await act(async () => { + root.render(); + await Promise.resolve(); + await new Promise((resolve) => setTimeout(resolve, 0)); + }); + + expect(container.querySelector('input[aria-label="titleField.inputAria"]')).not.toBeNull(); + const buttons = [...container.querySelectorAll('button')]; + const open = buttons.find((button) => button.textContent?.includes('actions.openProduction')); + const remove = buttons.find((button) => button.textContent?.includes('actions.deletePage')); + expect(open).toBeDefined(); + expect(remove).toBeDefined(); + + await act(async () => { + open?.click(); + await Promise.resolve(); + }); + expect(remove?.disabled).toBe(true); + + await act(async () => { + resolveOpenLink?.({ open_url: 'https://relay.test/open', expires_in_seconds: 60 }); + await Promise.resolve(); + }); + expect(remove?.disabled).toBe(false); + }); +}); diff --git a/src/web-ui/src/app/scenes/pages/PagesScene.tsx b/src/web-ui/src/app/scenes/pages/PagesScene.tsx new file mode 100644 index 0000000000..49d4ec4037 --- /dev/null +++ b/src/web-ui/src/app/scenes/pages/PagesScene.tsx @@ -0,0 +1,674 @@ +import React, { Suspense, lazy, useCallback, useEffect, useMemo, useRef, useState } from 'react'; +import { + ChevronDown, + ChevronUp, + Copy, + ExternalLink, + FileClock, + PanelsTopLeft, + RefreshCw, + Rocket, + Save, + Trash2, +} from 'lucide-react'; +import { + Button, + Input, + Select, + confirmDanger, + confirmWarning, + type SelectOption, +} from '@/component-library'; +import { GalleryEmpty, GalleryLayout, GalleryPageHeader } from '@/app/components'; +import { + pageAPI, + type PageInfo, + type PageVersionInfo, + type PageVisibility, +} from '@/infrastructure/api/service-api/PageAPI'; +import { remoteConnectAPI } from '@/infrastructure/api/service-api/RemoteConnectAPI'; +import { systemAPI } from '@/infrastructure/api/service-api/SystemAPI'; +import { useI18n } from '@/infrastructure/i18n'; +import { useNotification } from '@/shared/notification-system'; +import { createLogger } from '@/shared/utils/logger'; +import './PagesScene.scss'; + +const log = createLogger('PagesScene'); +const RemoteConnectDialog = lazy(() => import('@/app/components/RemoteConnectDialog')); + +interface PagesSceneProps { + isActive?: boolean; +} + +function errorText(error: unknown): string { + return error instanceof Error ? error.message : String(error); +} + +function replacePage(pages: PageInfo[], updated: PageInfo): PageInfo[] { + return pages.map((page) => (page.slug === updated.slug ? updated : page)); +} + +const PagesScene: React.FC = ({ isActive = true }) => { + const { t, formatDate, formatNumber } = useI18n('scenes/pages'); + const notification = useNotification(); + const attemptedLoadRef = useRef(false); + const [pages, setPages] = useState([]); + const [relayBaseUrl, setRelayBaseUrl] = useState(''); + const [versionsBySlug, setVersionsBySlug] = useState>({}); + const [expandedSlugs, setExpandedSlugs] = useState>(() => new Set()); + const [loading, setLoading] = useState(false); + const [loadError, setLoadError] = useState(''); + const [loginRequired, setLoginRequired] = useState(false); + const [showAccountDialog, setShowAccountDialog] = useState(false); + const [pendingBySlug, setPendingBySlug] = useState>({}); + const busySlugsRef = useRef>(new Set()); + const [titleDrafts, setTitleDrafts] = useState>({}); + + const beginPageAction = useCallback((slug: string, key: string): boolean => { + if (busySlugsRef.current.has(slug)) return false; + busySlugsRef.current.add(slug); + setPendingBySlug((current) => ({ ...current, [slug]: key })); + return true; + }, []); + + const endPageAction = useCallback((slug: string, key: string) => { + busySlugsRef.current.delete(slug); + setPendingBySlug((current) => { + if (current[slug] !== key) return current; + const next = { ...current }; + delete next[slug]; + return next; + }); + }, []); + + const visibilityOptions = useMemo(() => [ + { value: 'private', label: t('visibility.private') }, + { value: 'relay', label: t('visibility.relay') }, + { value: 'public', label: t('visibility.public') }, + ], [t]); + + const visibilityLabel = useCallback((visibility: PageVisibility): string => { + switch (visibility) { + case 'private': return t('visibility.private'); + case 'relay': return t('visibility.relay'); + case 'public': return t('visibility.public'); + } + }, [t]); + + const formatBytes = useCallback((bytes: number): string => { + if (bytes < 1024) return t('bytes.b', { value: formatNumber(bytes) }); + if (bytes < 1024 * 1024) { + return t('bytes.kb', { value: formatNumber(bytes / 1024, { maximumFractionDigits: 1 }) }); + } + return t('bytes.mb', { + value: formatNumber(bytes / (1024 * 1024), { maximumFractionDigits: 1 }), + }); + }, [formatNumber, t]); + + const formatTimestamp = useCallback((seconds: number): string => formatDate( + new Date(seconds * 1000), + { year: 'numeric', month: 'short', day: 'numeric', hour: '2-digit', minute: '2-digit' }, + ), [formatDate]); + + const loadPages = useCallback(async () => { + attemptedLoadRef.current = true; + setLoading(true); + setLoadError(''); + setLoginRequired(false); + try { + const status = await remoteConnectAPI.accountStatus(); + if (!status.logged_in) { + setPages([]); + setLoginRequired(true); + return; + } + const [nextPages, hint] = await Promise.all([ + pageAPI.listPages(), + remoteConnectAPI.accountGetCredentialHint().catch(() => null), + ]); + setPages(nextPages); + setRelayBaseUrl(hint?.relay_url?.replace(/\/$/, '') ?? ''); + } catch (error) { + log.error('Failed to load published Pages', { error }); + const latestStatus = await remoteConnectAPI.accountStatus().catch(() => null); + if (latestStatus && !latestStatus.logged_in) { + setPages([]); + setLoginRequired(true); + return; + } + setLoadError(errorText(error)); + } finally { + setLoading(false); + } + }, []); + + useEffect(() => { + if (isActive && !attemptedLoadRef.current && !loading) { + void loadPages(); + } + }, [isActive, loadPages, loading]); + + const loadVersions = useCallback(async (slug: string) => { + const key = `versions:${slug}`; + if (!beginPageAction(slug, key)) return; + try { + const versions = await pageAPI.listVersions(slug); + setVersionsBySlug((current) => ({ ...current, [slug]: versions })); + } catch (error) { + log.error('Failed to load Page versions', { slug, error }); + notification.error(t('notifications.versionsLoadFailed', { error: errorText(error) })); + throw error; + } finally { + endPageAction(slug, key); + } + }, [beginPageAction, endPageAction, notification, t]); + + const toggleVersions = useCallback(async (slug: string) => { + if (expandedSlugs.has(slug)) { + setExpandedSlugs((current) => { + const next = new Set(current); + next.delete(slug); + return next; + }); + return; + } + if (!versionsBySlug[slug]) { + try { + await loadVersions(slug); + } catch { + return; + } + } + setExpandedSlugs((current) => new Set(current).add(slug)); + }, [expandedSlugs, loadVersions, versionsBySlug]); + + const openPage = useCallback(async (page: PageInfo, versionId?: string) => { + const key = `open:${page.slug}:${versionId ?? 'production'}`; + if (!beginPageAction(page.slug, key)) return; + try { + const link = await pageAPI.createOpenLink(page.slug, versionId); + await systemAPI.openExternal(link.open_url); + } catch (error) { + log.error('Failed to open Page', { slug: page.slug, versionId, error }); + notification.error(t('notifications.openFailed', { error: errorText(error) })); + } finally { + endPageAction(page.slug, key); + } + }, [beginPageAction, endPageAction, notification, t]); + + const copyPageLink = useCallback(async (page: PageInfo, version?: PageVersionInfo) => { + const key = `copy:${page.slug}:${version?.version_id ?? 'production'}`; + if (!beginPageAction(page.slug, key)) return; + try { + let url = ''; + let temporary = false; + const path = version?.preview_url_path ?? page.url_path; + if (page.visibility === 'public' && relayBaseUrl && path) { + url = `${relayBaseUrl}${path.startsWith('/') ? '' : '/'}${path}`; + } else { + const link = await pageAPI.createOpenLink(page.slug, version?.version_id); + url = link.open_url; + temporary = true; + } + await systemAPI.setClipboard(url); + notification.success(temporary + ? t('notifications.temporaryLinkCopied') + : t('notifications.linkCopied')); + } catch (error) { + log.error('Failed to copy Page link', { + slug: page.slug, + versionId: version?.version_id, + error, + }); + notification.error(t('notifications.copyFailed', { error: errorText(error) })); + } finally { + endPageAction(page.slug, key); + } + }, [beginPageAction, endPageAction, notification, relayBaseUrl, t]); + + const changeVisibility = useCallback(async (page: PageInfo, visibility: PageVisibility) => { + if (visibility === page.visibility) return; + const confirmed = await confirmWarning( + t('confirm.visibilityTitle'), + t('confirm.visibilityMessage', { + slug: page.slug, + current: visibilityLabel(page.visibility), + target: visibilityLabel(visibility), + }), + { confirmText: t('actions.changeVisibility') }, + ); + if (!confirmed) return; + const key = `visibility:${page.slug}`; + if (!beginPageAction(page.slug, key)) return; + try { + const updated = await pageAPI.update(page.slug, { visibility }); + setPages((current) => replacePage(current, updated)); + notification.success(t('notifications.visibilityUpdated', { + visibility: visibilityLabel(visibility), + })); + } catch (error) { + log.error('Failed to update Page visibility', { slug: page.slug, visibility, error }); + notification.error(t('notifications.visibilityFailed', { error: errorText(error) })); + } finally { + endPageAction(page.slug, key); + } + }, [beginPageAction, endPageAction, notification, t, visibilityLabel]); + + const saveTitle = useCallback(async (page: PageInfo) => { + const title = (titleDrafts[page.slug] ?? page.title).trim(); + if (!title || title === page.title) return; + const key = `title:${page.slug}`; + if (!beginPageAction(page.slug, key)) return; + try { + const updated = await pageAPI.update(page.slug, { title }); + setPages((current) => replacePage(current, updated)); + setTitleDrafts((current) => ({ ...current, [page.slug]: updated.title })); + notification.success(t('notifications.titleUpdated')); + } catch (error) { + log.error('Failed to update Page title', { slug: page.slug, error }); + notification.error(t('notifications.titleFailed', { error: errorText(error) })); + } finally { + endPageAction(page.slug, key); + } + }, [beginPageAction, endPageAction, notification, t, titleDrafts]); + + const deployVersion = useCallback(async (page: PageInfo, version: PageVersionInfo) => { + if (version.deployed) return; + const confirmed = await confirmWarning( + t('confirm.deployTitle'), + t('confirm.deployMessage', { + slug: page.slug, + current: page.deployed_version_id ?? t('status.notDeployed'), + target: version.version_id, + }), + { confirmText: t('actions.deploy') }, + ); + if (!confirmed) return; + const key = `deploy:${page.slug}:${version.version_id}`; + if (!beginPageAction(page.slug, key)) return; + try { + const updated = await pageAPI.deploy(page.slug, version.version_id); + setPages((current) => replacePage(current, updated)); + setVersionsBySlug((current) => ({ + ...current, + [page.slug]: (current[page.slug] ?? []).map((item) => ({ + ...item, + deployed: item.version_id === version.version_id, + })), + })); + notification.success(t('notifications.deployed', { version: version.version_id })); + } catch (error) { + log.error('Failed to deploy Page version', { + slug: page.slug, + versionId: version.version_id, + error, + }); + notification.error(t('notifications.deployFailed', { error: errorText(error) })); + } finally { + endPageAction(page.slug, key); + } + }, [beginPageAction, endPageAction, notification, t]); + + const unpublishPage = useCallback(async (page: PageInfo) => { + if (!page.deployed_version_id) return; + const confirmed = await confirmWarning( + t('confirm.unpublishTitle'), + t('confirm.unpublishMessage', { + slug: page.slug, + current: page.deployed_version_id, + }), + { confirmText: t('actions.unpublish') }, + ); + if (!confirmed) return; + const key = `unpublish:${page.slug}`; + if (!beginPageAction(page.slug, key)) return; + try { + await pageAPI.unpublish(page.slug); + setPages((current) => current.map((item) => ( + item.slug === page.slug ? { ...item, deployed_version_id: null } : item + ))); + setVersionsBySlug((current) => ({ + ...current, + [page.slug]: (current[page.slug] ?? []).map((item) => ({ ...item, deployed: false })), + })); + notification.success(t('notifications.unpublished')); + } catch (error) { + log.error('Failed to unpublish Page', { slug: page.slug, error }); + notification.error(t('notifications.unpublishFailed', { error: errorText(error) })); + } finally { + endPageAction(page.slug, key); + } + }, [beginPageAction, endPageAction, notification, t]); + + const deleteVersion = useCallback(async (page: PageInfo, version: PageVersionInfo) => { + if (version.deployed) return; + const confirmed = await confirmDanger( + t('confirm.deleteVersionTitle'), + t('confirm.deleteVersionMessage', { version: version.version_id, title: page.title }), + { confirmText: t('actions.deleteVersion') }, + ); + if (!confirmed) return; + const key = `delete-version:${page.slug}:${version.version_id}`; + if (!beginPageAction(page.slug, key)) return; + try { + await pageAPI.deleteVersion(page.slug, version.version_id); + setVersionsBySlug((current) => ({ + ...current, + [page.slug]: (current[page.slug] ?? []) + .filter((item) => item.version_id !== version.version_id), + })); + notification.success(t('notifications.versionDeleted')); + } catch (error) { + log.error('Failed to delete Page version', { + slug: page.slug, + versionId: version.version_id, + error, + }); + notification.error(t('notifications.versionDeleteFailed', { error: errorText(error) })); + } finally { + endPageAction(page.slug, key); + } + }, [beginPageAction, endPageAction, notification, t]); + + const deletePage = useCallback(async (page: PageInfo) => { + const confirmed = await confirmDanger( + t('confirm.deletePageTitle'), + t('confirm.deletePageMessage', { title: page.title, slug: page.slug }), + { confirmText: t('actions.deletePage') }, + ); + if (!confirmed) return; + const key = `delete-page:${page.slug}`; + if (!beginPageAction(page.slug, key)) return; + try { + await pageAPI.deletePage(page.slug); + setPages((current) => current.filter((item) => item.slug !== page.slug)); + setVersionsBySlug((current) => { + const next = { ...current }; + delete next[page.slug]; + return next; + }); + notification.success(t('notifications.pageDeleted')); + } catch (error) { + log.error('Failed to delete Page', { slug: page.slug, error }); + notification.error(t('notifications.pageDeleteFailed', { error: errorText(error) })); + } finally { + endPageAction(page.slug, key); + } + }, [beginPageAction, endPageAction, notification, t]); + + const refreshButton = ( + + ); + + return ( + + + + {loading && pages.length === 0 ? ( + } + message={t('loading')} + testId="pages-loading" + /> + ) : loginRequired ? ( + } + message={<>{t('signInRequired')}{t('signInHint')}} + action={( + + )} + testId="pages-sign-in-required" + /> + ) : loadError && pages.length === 0 ? ( + } + message={<>{t('loadFailed')}{loadError}} + isError + action={} + testId="pages-error" + /> + ) : pages.length === 0 ? ( + } + message={<>{t('empty')}{t('emptyHint')}} + testId="pages-empty" + /> + ) : ( +
+ {pages.map((page) => { + const versions = versionsBySlug[page.slug] ?? []; + const expanded = expandedSlugs.has(page.slug); + const deployed = Boolean(page.deployed_version_id); + const pendingAction = pendingBySlug[page.slug]; + const pageBusy = Boolean(pendingAction); + const titleDraft = titleDrafts[page.slug] ?? page.title; + return ( +
+
+
+

{page.title || page.slug}

+ /{page.slug} +
+ + {deployed ? t('status.deployed') : t('status.savedOnly')} + +
+ +
+ {t('meta.updated', { date: formatTimestamp(page.updated_at) })} + {t('meta.size', { size: formatBytes(page.total_bytes) })} + {t('meta.files', { count: page.file_count })} +
+ +
+ {t('titleField.label')} +
+ setTitleDrafts((current) => ({ + ...current, + [page.slug]: event.currentTarget.value, + }))} + onKeyDown={(event) => { + if (event.key === 'Enter') void saveTitle(page); + }} + aria-label={t('titleField.inputAria', { slug: page.slug })} + /> + +
+
+ +
+ {t('visibility.label')} + , Select: () =>
, - confirmDanger: vi.fn(), + confirmDanger: mocks.confirmDanger, confirmWarning: vi.fn(), })); vi.mock('@/app/components', () => ({ GalleryLayout: ({ children }: { children: React.ReactNode }) =>
{children}
, - GalleryPageHeader: ({ title }: { title: React.ReactNode }) =>
{title}
, + GalleryPageHeader: ({ + title, + actions, + }: { + title: React.ReactNode; + actions?: React.ReactNode; + }) =>
{title}{actions}
, GalleryEmpty: ({ message, action, testId }: { message: React.ReactNode; action?: React.ReactNode; testId?: string }) => (
{message}{action}
), @@ -92,8 +107,13 @@ describe('PagesScene initial loading', () => { mocks.accountStatus.mockReset().mockResolvedValue({ logged_in: true, user_id: 'u1' }); mocks.accountGetCredentialHint.mockReset().mockResolvedValue({ relay_url: 'https://relay.test' }); mocks.listPages.mockReset().mockRejectedValue(new Error('relay unavailable')); + mocks.listVersions.mockReset().mockResolvedValue([]); mocks.createOpenLink.mockReset(); mocks.update.mockReset(); + mocks.deletePage.mockReset().mockResolvedValue(undefined); + mocks.confirmDanger.mockReset().mockResolvedValue(true); + mocks.listen.mockReset().mockImplementation(() => vi.fn()); + mocks.openExternal.mockReset().mockResolvedValue(undefined); }); afterEach(() => { @@ -122,6 +142,7 @@ describe('PagesScene initial loading', () => { it('locks every action on one Page while an operation is pending and exposes title editing', async () => { mocks.listPages.mockResolvedValue([{ slug: 'demo', + generation: 'aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa', visibility: 'public', title: 'Demo', file_count: 1, @@ -162,4 +183,325 @@ describe('PagesScene initial loading', () => { }); expect(remove?.disabled).toBe(false); }); + + it('does not restore a deleted Page from an older refresh response', async () => { + const page = { + slug: 'demo', + generation: 'aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa', + visibility: 'public', + title: 'Demo', + file_count: 1, + total_bytes: 20, + created_at: 1, + updated_at: 1, + url_path: '/p/alice/demo', + preview_url_path: '/p/alice/demo/@v/v1', + deployed_version_id: 'v1', + }; + mocks.listPages.mockResolvedValueOnce([page]); + + await act(async () => { + root.render(); + await Promise.resolve(); + await new Promise((resolve) => setTimeout(resolve, 0)); + }); + + let resolveRefresh: ((pages: typeof page[]) => void) | undefined; + mocks.listPages.mockImplementationOnce(() => new Promise((resolve) => { + resolveRefresh = resolve; + })); + const refresh = [...container.querySelectorAll('button')] + .find((button) => button.textContent === 'actions.refresh'); + await act(async () => { + refresh?.click(); + await Promise.resolve(); + }); + + const remove = [...container.querySelectorAll('button')] + .find((button) => button.textContent?.includes('actions.deletePage')); + await act(async () => { + remove?.click(); + await Promise.resolve(); + await new Promise((resolve) => setTimeout(resolve, 0)); + }); + expect(container.textContent).not.toContain('Demo'); + + await act(async () => { + resolveRefresh?.([page]); + await Promise.resolve(); + await new Promise((resolve) => setTimeout(resolve, 0)); + }); + expect(container.textContent).not.toContain('Demo'); + }); + + it('drops slug caches when the same account recreates a Page with a new generation', async () => { + const oldPage = { + slug: 'recreated', + generation: 'aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa', + visibility: 'public' as const, + title: 'Old Page', + file_count: 1, + total_bytes: 20, + created_at: 1, + updated_at: 1, + url_path: '/p/alice/recreated', + preview_url_path: '/p/alice/recreated/@v/a1', + deployed_version_id: 'a1', + }; + const newPage = { + ...oldPage, + generation: 'bbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbb', + title: 'New Page', + preview_url_path: '/p/alice/recreated/@v/b1', + deployed_version_id: 'b1', + }; + mocks.listPages.mockResolvedValueOnce([oldPage]); + mocks.listVersions + .mockResolvedValueOnce([{ + generation: oldPage.generation, + version_id: 'a1', + title: oldPage.title, + file_count: 1, + total_bytes: 20, + has_worker: false, + note: 'old generation note', + created_at: 1, + deployed: true, + preview_url_path: oldPage.preview_url_path, + }]) + .mockResolvedValueOnce([{ + generation: newPage.generation, + version_id: 'b1', + title: newPage.title, + file_count: 1, + total_bytes: 20, + has_worker: false, + note: 'new generation note', + created_at: 2, + deployed: true, + preview_url_path: newPage.preview_url_path, + }]); + + await act(async () => { + root.render(); + await Promise.resolve(); + await new Promise((resolve) => setTimeout(resolve, 0)); + }); + const oldInput = container.querySelector('input') as HTMLInputElement; + const valueSetter = Object.getOwnPropertyDescriptor( + HTMLInputElement.prototype, + 'value', + )?.set; + await act(async () => { + valueSetter?.call(oldInput, 'old generation draft'); + oldInput.dispatchEvent(new Event('input', { bubbles: true })); + await Promise.resolve(); + }); + expect(oldInput.value).toBe('old generation draft'); + + const oldVersions = [...container.querySelectorAll('button')] + .find((button) => button.textContent?.includes('actions.versions')); + await act(async () => { + oldVersions?.click(); + await Promise.resolve(); + await new Promise((resolve) => setTimeout(resolve, 0)); + }); + expect(container.textContent).toContain('old generation note'); + + mocks.listPages.mockResolvedValueOnce([newPage]); + const refresh = [...container.querySelectorAll('button')] + .find((button) => button.textContent === 'actions.refresh'); + await act(async () => { + refresh?.click(); + await Promise.resolve(); + await new Promise((resolve) => setTimeout(resolve, 0)); + }); + expect(container.textContent).toContain('New Page'); + expect(container.textContent).not.toContain('old generation note'); + expect((container.querySelector('input') as HTMLInputElement).value).toBe('New Page'); + + const newVersions = [...container.querySelectorAll('button')] + .find((button) => button.textContent?.includes('actions.versions')); + await act(async () => { + newVersions?.click(); + await Promise.resolve(); + await new Promise((resolve) => setTimeout(resolve, 0)); + }); + expect(mocks.listVersions).toHaveBeenLastCalledWith( + newPage.slug, + newPage.generation, + ); + expect(container.textContent).toContain('new generation note'); + expect(container.textContent).not.toContain('old generation note'); + }); + + it('clears account-owned state immediately and fences stale same-slug actions', async () => { + const pageA = { + slug: 'shared', + generation: 'aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa', + visibility: 'public' as const, + title: 'Account A Page', + file_count: 1, + total_bytes: 20, + created_at: 1, + updated_at: 1, + url_path: '/p/alice/shared', + preview_url_path: '/p/alice/shared/@v/a1', + deployed_version_id: 'a1', + }; + const pageB = { + ...pageA, + generation: 'bbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbb', + title: 'Account B Page', + url_path: '/p/bob/shared', + preview_url_path: '/p/bob/shared/@v/b1', + deployed_version_id: 'b1', + }; + mocks.listPages.mockResolvedValueOnce([pageA]); + mocks.listVersions.mockResolvedValueOnce([{ + generation: pageA.generation, + version_id: 'a1', + title: pageA.title, + file_count: 1, + total_bytes: 20, + has_worker: false, + note: 'A-only note', + created_at: 1, + deployed: true, + preview_url_path: pageA.preview_url_path, + }]); + + await act(async () => { + root.render(); + await Promise.resolve(); + await new Promise((resolve) => setTimeout(resolve, 0)); + }); + const versions = [...container.querySelectorAll('button')] + .find((button) => button.textContent?.includes('actions.versions')); + await act(async () => { + versions?.click(); + await Promise.resolve(); + await new Promise((resolve) => setTimeout(resolve, 0)); + }); + expect(container.textContent).toContain('A-only note'); + + let resolveConfirmation: ((confirmed: boolean) => void) | undefined; + mocks.confirmDanger.mockImplementationOnce(() => new Promise((resolve) => { + resolveConfirmation = resolve; + })); + const staleDelete = [...container.querySelectorAll('button')] + .find((button) => button.textContent?.includes('actions.deletePage')); + await act(async () => { + staleDelete?.click(); + await Promise.resolve(); + }); + + let resolveBPages: ((pages: typeof pageB[]) => void) | undefined; + mocks.accountStatus.mockResolvedValue({ logged_in: true, user_id: 'u2' }); + mocks.listPages.mockImplementationOnce(() => new Promise((resolve) => { + resolveBPages = resolve; + })); + const loginStateListener = mocks.listen.mock.calls + .find(([event]) => event === 'account://login-state')?.[1] as + ((payload: { logged_in: boolean }) => void) | undefined; + expect(loginStateListener).toBeDefined(); + await act(async () => { + loginStateListener?.({ logged_in: true }); + await Promise.resolve(); + await new Promise((resolve) => setTimeout(resolve, 0)); + }); + + // B ownership is adopted before its list response arrives, so no A-only + // versions, drafts, or cards remain actionable during the gap. + expect(container.textContent).not.toContain('Account A Page'); + expect(container.textContent).not.toContain('A-only note'); + + await act(async () => { + resolveConfirmation?.(true); + await Promise.resolve(); + await new Promise((resolve) => setTimeout(resolve, 0)); + }); + expect(mocks.deletePage).not.toHaveBeenCalled(); + + await act(async () => { + resolveBPages?.([pageB]); + await Promise.resolve(); + await new Promise((resolve) => setTimeout(resolve, 0)); + }); + expect(container.textContent).toContain('Account B Page'); + expect(container.textContent).not.toContain('A-only note'); + expect((container.querySelector('input') as HTMLInputElement | null)?.value) + .toBe('Account B Page'); + }); + + it('invalidates in-flight actions when the same user logs in again', async () => { + const oldPage = { + slug: 'same-user', + generation: 'aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa', + visibility: 'public' as const, + title: 'Old session Page', + file_count: 1, + total_bytes: 20, + created_at: 1, + updated_at: 1, + url_path: '/p/alice/same-user', + preview_url_path: '/p/alice/same-user/@v/a1', + deployed_version_id: 'a1', + }; + const freshPage = { + ...oldPage, + generation: 'bbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbb', + title: 'Fresh session Page', + preview_url_path: '/p/alice/same-user/@v/b1', + deployed_version_id: 'b1', + }; + mocks.listPages.mockResolvedValueOnce([oldPage]); + let resolveStaleOpen: ((value: { open_url: string; expires_in_seconds: number }) => void) + | undefined; + mocks.createOpenLink.mockImplementationOnce(() => new Promise((resolve) => { + resolveStaleOpen = resolve; + })); + + await act(async () => { + root.render(); + await Promise.resolve(); + await new Promise((resolve) => setTimeout(resolve, 0)); + }); + const open = [...container.querySelectorAll('button')] + .find((button) => button.textContent?.includes('actions.openProduction')); + await act(async () => { + open?.click(); + await Promise.resolve(); + await new Promise((resolve) => setTimeout(resolve, 0)); + }); + expect(mocks.createOpenLink).toHaveBeenCalledWith( + oldPage.slug, + oldPage.generation, + undefined, + ); + + mocks.listPages.mockResolvedValueOnce([freshPage]); + const loginStateListener = mocks.listen.mock.calls + .find(([event]) => event === 'account://login-state')?.[1] as + ((payload: { logged_in: boolean }) => void) | undefined; + expect(loginStateListener).toBeDefined(); + await act(async () => { + loginStateListener?.({ logged_in: true }); + await Promise.resolve(); + await new Promise((resolve) => setTimeout(resolve, 0)); + }); + expect(container.textContent).toContain('Fresh session Page'); + expect(container.textContent).not.toContain('Old session Page'); + + await act(async () => { + resolveStaleOpen?.({ + open_url: 'https://relay.test/stale-open', + expires_in_seconds: 60, + }); + await Promise.resolve(); + await new Promise((resolve) => setTimeout(resolve, 0)); + }); + expect(mocks.openExternal).not.toHaveBeenCalled(); + expect(container.textContent).toContain('Fresh session Page'); + }); }); diff --git a/src/web-ui/src/app/scenes/pages/PagesScene.tsx b/src/web-ui/src/app/scenes/pages/PagesScene.tsx index 49d4ec4037..709d92b0a0 100644 --- a/src/web-ui/src/app/scenes/pages/PagesScene.tsx +++ b/src/web-ui/src/app/scenes/pages/PagesScene.tsx @@ -26,6 +26,7 @@ import { type PageVersionInfo, type PageVisibility, } from '@/infrastructure/api/service-api/PageAPI'; +import { api } from '@/infrastructure/api/service-api/ApiClient'; import { remoteConnectAPI } from '@/infrastructure/api/service-api/RemoteConnectAPI'; import { systemAPI } from '@/infrastructure/api/service-api/SystemAPI'; import { useI18n } from '@/infrastructure/i18n'; @@ -40,6 +41,19 @@ interface PagesSceneProps { isActive?: boolean; } +interface PageOwner { + userId: string; + epoch: number; +} + +interface PageActionLease { + slug: string; + key: string; + token: string; + userId: string; + ownerEpoch: number; +} + function errorText(error: unknown): string { return error instanceof Error ? error.message : String(error); } @@ -52,6 +66,12 @@ const PagesScene: React.FC = ({ isActive = true }) => { const { t, formatDate, formatNumber } = useI18n('scenes/pages'); const notification = useNotification(); const attemptedLoadRef = useRef(false); + const pageLoadEpochRef = useRef(0); + const pageOwnerRef = useRef(null); + const pageOwnerEpochCounterRef = useRef(0); + const nextActionTokenRef = useRef(0); + const [pageOwnerEpoch, setPageOwnerEpoch] = useState(0); + const pagesRef = useRef([]); const [pages, setPages] = useState([]); const [relayBaseUrl, setRelayBaseUrl] = useState(''); const [versionsBySlug, setVersionsBySlug] = useState>({}); @@ -61,26 +81,147 @@ const PagesScene: React.FC = ({ isActive = true }) => { const [loginRequired, setLoginRequired] = useState(false); const [showAccountDialog, setShowAccountDialog] = useState(false); const [pendingBySlug, setPendingBySlug] = useState>({}); - const busySlugsRef = useRef>(new Set()); + const busySlugsRef = useRef>(new Map()); const [titleDrafts, setTitleDrafts] = useState>({}); - const beginPageAction = useCallback((slug: string, key: string): boolean => { - if (busySlugsRef.current.has(slug)) return false; - busySlugsRef.current.add(slug); - setPendingBySlug((current) => ({ ...current, [slug]: key })); - return true; + const cancelPendingPageLoad = useCallback(() => { + pageLoadEpochRef.current += 1; + setLoading(false); + }, []); + + const adoptPageOwner = useCallback(( + userId: string | null, + cancelLoads: boolean, + forceNewEpoch = false, + ): number => { + const current = pageOwnerRef.current; + if (!forceNewEpoch && (current?.userId ?? null) === userId) { + return current?.epoch ?? pageOwnerEpochCounterRef.current; + } + const epoch = pageOwnerEpochCounterRef.current + 1; + pageOwnerEpochCounterRef.current = epoch; + pageOwnerRef.current = userId ? { userId, epoch } : null; + setPageOwnerEpoch(epoch); + if (cancelLoads) pageLoadEpochRef.current += 1; + busySlugsRef.current.clear(); + pagesRef.current = []; + setPages([]); + setRelayBaseUrl(''); + setVersionsBySlug({}); + setExpandedSlugs(new Set()); + setTitleDrafts({}); + setPendingBySlug({}); + setLoadError(''); + setLoading(false); + setLoginRequired(userId === null); + return epoch; + }, []); + + const updateOwnedPages = useCallback(( + update: (current: PageInfo[]) => PageInfo[], + ) => { + const next = update(pagesRef.current); + pagesRef.current = next; + setPages(next); + }, []); + + const commitLoadedPages = useCallback((nextPages: PageInfo[]) => { + const previousGenerationBySlug = new Map( + pagesRef.current.map((page) => [page.slug, page.generation]), + ); + const nextGenerationBySlug = new Map( + nextPages.map((page) => [page.slug, page.generation]), + ); + const canRetainSlugState = (slug: string) => { + const previousGeneration = previousGenerationBySlug.get(slug); + return previousGeneration !== undefined + && previousGeneration === nextGenerationBySlug.get(slug); + }; + + for (const slug of busySlugsRef.current.keys()) { + if (!canRetainSlugState(slug)) busySlugsRef.current.delete(slug); + } + setVersionsBySlug((current) => Object.fromEntries( + Object.entries(current).filter(([slug]) => canRetainSlugState(slug)), + )); + setExpandedSlugs((current) => new Set( + [...current].filter((slug) => canRetainSlugState(slug)), + )); + setTitleDrafts((current) => Object.fromEntries( + Object.entries(current).filter(([slug]) => canRetainSlugState(slug)), + )); + setPendingBySlug((current) => Object.fromEntries( + Object.entries(current).filter(([slug]) => canRetainSlugState(slug)), + )); + pagesRef.current = nextPages; + setPages(nextPages); + }, []); + + const isPageActionCurrent = useCallback((lease: PageActionLease): boolean => { + const owner = pageOwnerRef.current; + return owner?.userId === lease.userId + && owner.epoch === lease.ownerEpoch + && busySlugsRef.current.get(lease.slug) === lease.token; }, []); - const endPageAction = useCallback((slug: string, key: string) => { - busySlugsRef.current.delete(slug); + const endPageAction = useCallback((lease: PageActionLease) => { + if (busySlugsRef.current.get(lease.slug) !== lease.token) return; + busySlugsRef.current.delete(lease.slug); setPendingBySlug((current) => { - if (current[slug] !== key) return current; + if (current[lease.slug] !== lease.key) return current; const next = { ...current }; - delete next[slug]; + delete next[lease.slug]; return next; }); }, []); + const beginPageAction = useCallback(async ( + page: PageInfo, + key: string, + expectedOwnerEpoch: number, + ): Promise => { + const owner = pageOwnerRef.current; + if (!owner || owner.epoch !== expectedOwnerEpoch || busySlugsRef.current.has(page.slug)) { + return null; + } + const token = `${owner.epoch}:${nextActionTokenRef.current += 1}`; + const lease: PageActionLease = { + slug: page.slug, + key, + token, + userId: owner.userId, + ownerEpoch: owner.epoch, + }; + // A list response captured before this mutation must never overwrite the + // operation's newer result when it eventually arrives. + cancelPendingPageLoad(); + busySlugsRef.current.set(page.slug, token); + setPendingBySlug((current) => ({ ...current, [page.slug]: key })); + const status = await remoteConnectAPI.accountStatus().catch(() => null); + if (!status?.logged_in || status.user_id !== owner.userId || !isPageActionCurrent(lease)) { + if (status) { + adoptPageOwner(status.logged_in ? status.user_id : null, true); + attemptedLoadRef.current = !status.logged_in; + } + endPageAction(lease); + return null; + } + return lease; + }, [adoptPageOwner, cancelPendingPageLoad, endPageAction, isPageActionCurrent]); + + const validatePageAction = useCallback(async (lease: PageActionLease): Promise => { + if (!isPageActionCurrent(lease)) return false; + const status = await remoteConnectAPI.accountStatus().catch(() => null); + if (!status?.logged_in || status.user_id !== lease.userId || !isPageActionCurrent(lease)) { + if (status) { + adoptPageOwner(status.logged_in ? status.user_id : null, true); + attemptedLoadRef.current = !status.logged_in; + } + return false; + } + return true; + }, [adoptPageOwner, isPageActionCurrent]); + const visibilityOptions = useMemo(() => [ { value: 'private', label: t('visibility.private') }, { value: 'relay', label: t('visibility.relay') }, @@ -111,36 +252,65 @@ const PagesScene: React.FC = ({ isActive = true }) => { ), [formatDate]); const loadPages = useCallback(async () => { + if (busySlugsRef.current.size > 0) return; + const requestEpoch = pageLoadEpochRef.current + 1; + pageLoadEpochRef.current = requestEpoch; attemptedLoadRef.current = true; setLoading(true); setLoadError(''); setLoginRequired(false); + let requestedUserId: string | null = null; try { const status = await remoteConnectAPI.accountStatus(); - if (!status.logged_in) { - setPages([]); + if (pageLoadEpochRef.current !== requestEpoch) return; + requestedUserId = status.user_id; + if (!status.logged_in || !status.user_id) { + adoptPageOwner(null, false); setLoginRequired(true); return; } + const ownerEpoch = adoptPageOwner(status.user_id, false); const [nextPages, hint] = await Promise.all([ pageAPI.listPages(), remoteConnectAPI.accountGetCredentialHint().catch(() => null), ]); - setPages(nextPages); + const latestStatus = await remoteConnectAPI.accountStatus().catch(() => null); + if (pageLoadEpochRef.current !== requestEpoch) return; + if (!latestStatus?.logged_in || latestStatus.user_id !== status.user_id) { + // The account changed while this relay request was in flight. Leave + // ownership of the UI to a fresh request for the new account. + adoptPageOwner(latestStatus?.logged_in ? latestStatus.user_id : null, true); + attemptedLoadRef.current = !latestStatus?.logged_in; + return; + } + const currentOwner = pageOwnerRef.current; + if (currentOwner?.userId !== status.user_id || currentOwner.epoch !== ownerEpoch) return; + commitLoadedPages(nextPages); setRelayBaseUrl(hint?.relay_url?.replace(/\/$/, '') ?? ''); } catch (error) { + if (pageLoadEpochRef.current !== requestEpoch) return; log.error('Failed to load published Pages', { error }); const latestStatus = await remoteConnectAPI.accountStatus().catch(() => null); + if (pageLoadEpochRef.current !== requestEpoch) return; + if (requestedUserId !== null + && latestStatus?.logged_in + && latestStatus.user_id !== requestedUserId) { + adoptPageOwner(latestStatus.user_id, true); + attemptedLoadRef.current = false; + return; + } if (latestStatus && !latestStatus.logged_in) { - setPages([]); + adoptPageOwner(null, true); setLoginRequired(true); return; } setLoadError(errorText(error)); } finally { - setLoading(false); + if (pageLoadEpochRef.current === requestEpoch) { + setLoading(false); + } } - }, []); + }, [adoptPageOwner, commitLoadedPages]); useEffect(() => { if (isActive && !attemptedLoadRef.current && !loading) { @@ -148,57 +318,93 @@ const PagesScene: React.FC = ({ isActive = true }) => { } }, [isActive, loadPages, loading]); - const loadVersions = useCallback(async (slug: string) => { - const key = `versions:${slug}`; - if (!beginPageAction(slug, key)) return; + useEffect(() => { + const unlisten = api.listen<{ logged_in: boolean }>( + 'account://login-state', + (payload) => { + const loggedIn = payload?.logged_in === true; + // The event intentionally carries no user id. Treat every transition, + // including a same-user re-login, as a new ownership generation before + // doing any asynchronous status lookup. This immediately removes data + // and actions owned by the previous authenticated session. + adoptPageOwner(null, true, true); + attemptedLoadRef.current = !loggedIn; + if (loggedIn && isActive) { + void loadPages(); + } + }, + ); + return unlisten; + }, [adoptPageOwner, isActive, loadPages]); + + const loadVersions = useCallback(async (page: PageInfo, ownerEpoch: number) => { + const key = `versions:${page.slug}`; + const lease = await beginPageAction(page, key, ownerEpoch); + if (!lease) return; try { - const versions = await pageAPI.listVersions(slug); - setVersionsBySlug((current) => ({ ...current, [slug]: versions })); + const versions = await pageAPI.listVersions(page.slug, page.generation); + if (!await validatePageAction(lease)) return; + setVersionsBySlug((current) => ({ ...current, [page.slug]: versions })); } catch (error) { - log.error('Failed to load Page versions', { slug, error }); + if (!await validatePageAction(lease)) return; + log.error('Failed to load Page versions', { slug: page.slug, error }); notification.error(t('notifications.versionsLoadFailed', { error: errorText(error) })); throw error; } finally { - endPageAction(slug, key); + endPageAction(lease); } - }, [beginPageAction, endPageAction, notification, t]); + }, [beginPageAction, endPageAction, notification, t, validatePageAction]); - const toggleVersions = useCallback(async (slug: string) => { - if (expandedSlugs.has(slug)) { + const toggleVersions = useCallback(async (page: PageInfo, ownerEpoch: number) => { + if (expandedSlugs.has(page.slug)) { setExpandedSlugs((current) => { const next = new Set(current); - next.delete(slug); + next.delete(page.slug); return next; }); return; } - if (!versionsBySlug[slug]) { + if (!versionsBySlug[page.slug]) { try { - await loadVersions(slug); + await loadVersions(page, ownerEpoch); } catch { return; } } - setExpandedSlugs((current) => new Set(current).add(slug)); + if (pageOwnerRef.current?.epoch === ownerEpoch) { + setExpandedSlugs((current) => new Set(current).add(page.slug)); + } }, [expandedSlugs, loadVersions, versionsBySlug]); - const openPage = useCallback(async (page: PageInfo, versionId?: string) => { + const openPage = useCallback(async ( + page: PageInfo, + ownerEpoch: number, + versionId?: string, + ) => { const key = `open:${page.slug}:${versionId ?? 'production'}`; - if (!beginPageAction(page.slug, key)) return; + const lease = await beginPageAction(page, key, ownerEpoch); + if (!lease) return; try { - const link = await pageAPI.createOpenLink(page.slug, versionId); + const link = await pageAPI.createOpenLink(page.slug, page.generation, versionId); + if (!await validatePageAction(lease)) return; await systemAPI.openExternal(link.open_url); } catch (error) { + if (!await validatePageAction(lease)) return; log.error('Failed to open Page', { slug: page.slug, versionId, error }); notification.error(t('notifications.openFailed', { error: errorText(error) })); } finally { - endPageAction(page.slug, key); + endPageAction(lease); } - }, [beginPageAction, endPageAction, notification, t]); + }, [beginPageAction, endPageAction, notification, t, validatePageAction]); - const copyPageLink = useCallback(async (page: PageInfo, version?: PageVersionInfo) => { + const copyPageLink = useCallback(async ( + page: PageInfo, + ownerEpoch: number, + version?: PageVersionInfo, + ) => { const key = `copy:${page.slug}:${version?.version_id ?? 'production'}`; - if (!beginPageAction(page.slug, key)) return; + const lease = await beginPageAction(page, key, ownerEpoch); + if (!lease) return; try { let url = ''; let temporary = false; @@ -206,15 +412,21 @@ const PagesScene: React.FC = ({ isActive = true }) => { if (page.visibility === 'public' && relayBaseUrl && path) { url = `${relayBaseUrl}${path.startsWith('/') ? '' : '/'}${path}`; } else { - const link = await pageAPI.createOpenLink(page.slug, version?.version_id); + const link = await pageAPI.createOpenLink( + page.slug, + page.generation, + version?.version_id, + ); url = link.open_url; temporary = true; } + if (!await validatePageAction(lease)) return; await systemAPI.setClipboard(url); notification.success(temporary ? t('notifications.temporaryLinkCopied') : t('notifications.linkCopied')); } catch (error) { + if (!await validatePageAction(lease)) return; log.error('Failed to copy Page link', { slug: page.slug, versionId: version?.version_id, @@ -222,11 +434,15 @@ const PagesScene: React.FC = ({ isActive = true }) => { }); notification.error(t('notifications.copyFailed', { error: errorText(error) })); } finally { - endPageAction(page.slug, key); + endPageAction(lease); } - }, [beginPageAction, endPageAction, notification, relayBaseUrl, t]); + }, [beginPageAction, endPageAction, notification, relayBaseUrl, t, validatePageAction]); - const changeVisibility = useCallback(async (page: PageInfo, visibility: PageVisibility) => { + const changeVisibility = useCallback(async ( + page: PageInfo, + ownerEpoch: number, + visibility: PageVisibility, + ) => { if (visibility === page.visibility) return; const confirmed = await confirmWarning( t('confirm.visibilityTitle'), @@ -239,40 +455,50 @@ const PagesScene: React.FC = ({ isActive = true }) => { ); if (!confirmed) return; const key = `visibility:${page.slug}`; - if (!beginPageAction(page.slug, key)) return; + const lease = await beginPageAction(page, key, ownerEpoch); + if (!lease) return; try { - const updated = await pageAPI.update(page.slug, { visibility }); - setPages((current) => replacePage(current, updated)); + const updated = await pageAPI.update(page.slug, page.generation, { visibility }); + if (!await validatePageAction(lease)) return; + updateOwnedPages((current) => replacePage(current, updated)); notification.success(t('notifications.visibilityUpdated', { visibility: visibilityLabel(visibility), })); } catch (error) { + if (!await validatePageAction(lease)) return; log.error('Failed to update Page visibility', { slug: page.slug, visibility, error }); notification.error(t('notifications.visibilityFailed', { error: errorText(error) })); } finally { - endPageAction(page.slug, key); + endPageAction(lease); } - }, [beginPageAction, endPageAction, notification, t, visibilityLabel]); + }, [beginPageAction, endPageAction, notification, t, updateOwnedPages, validatePageAction, visibilityLabel]); - const saveTitle = useCallback(async (page: PageInfo) => { + const saveTitle = useCallback(async (page: PageInfo, ownerEpoch: number) => { const title = (titleDrafts[page.slug] ?? page.title).trim(); if (!title || title === page.title) return; const key = `title:${page.slug}`; - if (!beginPageAction(page.slug, key)) return; + const lease = await beginPageAction(page, key, ownerEpoch); + if (!lease) return; try { - const updated = await pageAPI.update(page.slug, { title }); - setPages((current) => replacePage(current, updated)); + const updated = await pageAPI.update(page.slug, page.generation, { title }); + if (!await validatePageAction(lease)) return; + updateOwnedPages((current) => replacePage(current, updated)); setTitleDrafts((current) => ({ ...current, [page.slug]: updated.title })); notification.success(t('notifications.titleUpdated')); } catch (error) { + if (!await validatePageAction(lease)) return; log.error('Failed to update Page title', { slug: page.slug, error }); notification.error(t('notifications.titleFailed', { error: errorText(error) })); } finally { - endPageAction(page.slug, key); + endPageAction(lease); } - }, [beginPageAction, endPageAction, notification, t, titleDrafts]); + }, [beginPageAction, endPageAction, notification, t, titleDrafts, updateOwnedPages, validatePageAction]); - const deployVersion = useCallback(async (page: PageInfo, version: PageVersionInfo) => { + const deployVersion = useCallback(async ( + page: PageInfo, + ownerEpoch: number, + version: PageVersionInfo, + ) => { if (version.deployed) return; const confirmed = await confirmWarning( t('confirm.deployTitle'), @@ -285,10 +511,12 @@ const PagesScene: React.FC = ({ isActive = true }) => { ); if (!confirmed) return; const key = `deploy:${page.slug}:${version.version_id}`; - if (!beginPageAction(page.slug, key)) return; + const lease = await beginPageAction(page, key, ownerEpoch); + if (!lease) return; try { - const updated = await pageAPI.deploy(page.slug, version.version_id); - setPages((current) => replacePage(current, updated)); + const updated = await pageAPI.deploy(page.slug, page.generation, version.version_id); + if (!await validatePageAction(lease)) return; + updateOwnedPages((current) => replacePage(current, updated)); setVersionsBySlug((current) => ({ ...current, [page.slug]: (current[page.slug] ?? []).map((item) => ({ @@ -298,6 +526,7 @@ const PagesScene: React.FC = ({ isActive = true }) => { })); notification.success(t('notifications.deployed', { version: version.version_id })); } catch (error) { + if (!await validatePageAction(lease)) return; log.error('Failed to deploy Page version', { slug: page.slug, versionId: version.version_id, @@ -305,11 +534,11 @@ const PagesScene: React.FC = ({ isActive = true }) => { }); notification.error(t('notifications.deployFailed', { error: errorText(error) })); } finally { - endPageAction(page.slug, key); + endPageAction(lease); } - }, [beginPageAction, endPageAction, notification, t]); + }, [beginPageAction, endPageAction, notification, t, updateOwnedPages, validatePageAction]); - const unpublishPage = useCallback(async (page: PageInfo) => { + const unpublishPage = useCallback(async (page: PageInfo, ownerEpoch: number) => { if (!page.deployed_version_id) return; const confirmed = await confirmWarning( t('confirm.unpublishTitle'), @@ -321,10 +550,12 @@ const PagesScene: React.FC = ({ isActive = true }) => { ); if (!confirmed) return; const key = `unpublish:${page.slug}`; - if (!beginPageAction(page.slug, key)) return; + const lease = await beginPageAction(page, key, ownerEpoch); + if (!lease) return; try { - await pageAPI.unpublish(page.slug); - setPages((current) => current.map((item) => ( + await pageAPI.unpublish(page.slug, page.generation); + if (!await validatePageAction(lease)) return; + updateOwnedPages((current) => current.map((item) => ( item.slug === page.slug ? { ...item, deployed_version_id: null } : item ))); setVersionsBySlug((current) => ({ @@ -333,14 +564,19 @@ const PagesScene: React.FC = ({ isActive = true }) => { })); notification.success(t('notifications.unpublished')); } catch (error) { + if (!await validatePageAction(lease)) return; log.error('Failed to unpublish Page', { slug: page.slug, error }); notification.error(t('notifications.unpublishFailed', { error: errorText(error) })); } finally { - endPageAction(page.slug, key); + endPageAction(lease); } - }, [beginPageAction, endPageAction, notification, t]); + }, [beginPageAction, endPageAction, notification, t, updateOwnedPages, validatePageAction]); - const deleteVersion = useCallback(async (page: PageInfo, version: PageVersionInfo) => { + const deleteVersion = useCallback(async ( + page: PageInfo, + ownerEpoch: number, + version: PageVersionInfo, + ) => { if (version.deployed) return; const confirmed = await confirmDanger( t('confirm.deleteVersionTitle'), @@ -349,9 +585,11 @@ const PagesScene: React.FC = ({ isActive = true }) => { ); if (!confirmed) return; const key = `delete-version:${page.slug}:${version.version_id}`; - if (!beginPageAction(page.slug, key)) return; + const lease = await beginPageAction(page, key, ownerEpoch); + if (!lease) return; try { - await pageAPI.deleteVersion(page.slug, version.version_id); + await pageAPI.deleteVersion(page.slug, page.generation, version.version_id); + if (!await validatePageAction(lease)) return; setVersionsBySlug((current) => ({ ...current, [page.slug]: (current[page.slug] ?? []) @@ -359,6 +597,7 @@ const PagesScene: React.FC = ({ isActive = true }) => { })); notification.success(t('notifications.versionDeleted')); } catch (error) { + if (!await validatePageAction(lease)) return; log.error('Failed to delete Page version', { slug: page.slug, versionId: version.version_id, @@ -366,11 +605,11 @@ const PagesScene: React.FC = ({ isActive = true }) => { }); notification.error(t('notifications.versionDeleteFailed', { error: errorText(error) })); } finally { - endPageAction(page.slug, key); + endPageAction(lease); } - }, [beginPageAction, endPageAction, notification, t]); + }, [beginPageAction, endPageAction, notification, t, validatePageAction]); - const deletePage = useCallback(async (page: PageInfo) => { + const deletePage = useCallback(async (page: PageInfo, ownerEpoch: number) => { const confirmed = await confirmDanger( t('confirm.deletePageTitle'), t('confirm.deletePageMessage', { title: page.title, slug: page.slug }), @@ -378,10 +617,12 @@ const PagesScene: React.FC = ({ isActive = true }) => { ); if (!confirmed) return; const key = `delete-page:${page.slug}`; - if (!beginPageAction(page.slug, key)) return; + const lease = await beginPageAction(page, key, ownerEpoch); + if (!lease) return; try { - await pageAPI.deletePage(page.slug); - setPages((current) => current.filter((item) => item.slug !== page.slug)); + await pageAPI.deletePage(page.slug, page.generation); + if (!await validatePageAction(lease)) return; + updateOwnedPages((current) => current.filter((item) => item.slug !== page.slug)); setVersionsBySlug((current) => { const next = { ...current }; delete next[page.slug]; @@ -389,19 +630,20 @@ const PagesScene: React.FC = ({ isActive = true }) => { }); notification.success(t('notifications.pageDeleted')); } catch (error) { + if (!await validatePageAction(lease)) return; log.error('Failed to delete Page', { slug: page.slug, error }); notification.error(t('notifications.pageDeleteFailed', { error: errorText(error) })); } finally { - endPageAction(page.slug, key); + endPageAction(lease); } - }, [beginPageAction, endPageAction, notification, t]); + }, [beginPageAction, endPageAction, notification, t, updateOwnedPages, validatePageAction]); const refreshButton = ( +
+ )} + {loading && pages.length === 0 ? ( } @@ -482,12 +734,15 @@ const PagesScene: React.FC = ({ isActive = true }) => { value={titleDraft} maxLength={120} disabled={pageBusy} - onChange={(event) => setTitleDrafts((current) => ({ - ...current, - [page.slug]: event.currentTarget.value, - }))} + onChange={(event) => { + const value = event.currentTarget.value; + setTitleDrafts((current) => ({ + ...current, + [page.slug]: value, + })); + }} onKeyDown={(event) => { - if (event.key === 'Enter') void saveTitle(page); + if (event.key === 'Enter') void saveTitle(page, pageOwnerEpoch); }} aria-label={t('titleField.inputAria', { slug: page.slug })} /> @@ -496,7 +751,7 @@ const PagesScene: React.FC = ({ isActive = true }) => { size="small" disabled={pageBusy || !titleDraft.trim() || titleDraft.trim() === page.title} isLoading={pendingAction === `title:${page.slug}`} - onClick={() => void saveTitle(page)} + onClick={() => void saveTitle(page, pageOwnerEpoch)} > {t('actions.saveTitle')} @@ -510,7 +765,11 @@ const PagesScene: React.FC = ({ isActive = true }) => { value={page.visibility} options={visibilityOptions} disabled={pageBusy} - onChange={(value) => void changeVisibility(page, String(value) as PageVisibility)} + onChange={(value) => void changeVisibility( + page, + pageOwnerEpoch, + String(value) as PageVisibility, + )} triggerAriaLabel={t('visibility.changeAria', { title: page.title || page.slug })} />
@@ -520,7 +779,7 @@ const PagesScene: React.FC = ({ isActive = true }) => { @@ -605,7 +864,7 @@ const PagesScene: React.FC = ({ isActive = true }) => { )} - {loginPanel.authorizationUrl && ( + {loginPanel.status === 'pending' && loginPanel.authorizationUrl && (