diff --git a/apps/codex-plus-launcher/src/main.rs b/apps/codex-plus-launcher/src/main.rs index af73c5067..a6947b83e 100644 --- a/apps/codex-plus-launcher/src/main.rs +++ b/apps/codex-plus-launcher/src/main.rs @@ -161,6 +161,21 @@ async fn activate_existing_codex_app(options: &LaunchOptions) -> anyhow::Result< let helper_port = hooks.select_helper_port(options.helper_port); let settings = hooks.load_settings().await?; let app_dir = hooks.resolve_app_dir(options.app_dir.as_deref(), &settings)?; + let has_pending_recovery = hooks.has_pending_remote_control_session_recoveries(); + let blocking_process_ids = if has_pending_recovery { + codex_plus_core::watcher::find_session_index_cleanup_blocking_processes() + } else { + Vec::new() + }; + if should_finalize_pending_remote_control_recovery(has_pending_recovery, &blocking_process_ids) + { + hooks.run_remote_control_session_recovery().await?; + } else if has_pending_recovery { + let _ = codex_plus_core::diagnostic_log::append_diagnostic_log( + "launcher.remote_control_session_finalization_deferred_existing_app", + json!({"blocking_process_ids": blocking_process_ids}), + ); + } let launch_result = hooks .launch_codex( &app_dir, @@ -215,6 +230,13 @@ async fn activate_existing_codex_app(options: &LaunchOptions) -> anyhow::Result< launch_result.map(|_| ()) } +fn should_finalize_pending_remote_control_recovery( + has_pending_recovery: bool, + blocking_process_ids: &[u32], +) -> bool { + has_pending_recovery && blocking_process_ids.is_empty() +} + fn log_launcher_already_running(debug_port: u16) { let _ = codex_plus_core::diagnostic_log::append_diagnostic_log( "launcher.already_running", @@ -310,6 +332,103 @@ impl LaunchHooks for LauncherHooks { Ok(()) } + fn has_pending_remote_control_session_recoveries(&self) -> bool { + codex_plus_core::paths::default_pending_remote_control_recovery_path().exists() + } + + fn remote_control_session_recovery_is_safe_to_run(&self) -> bool { + codex_plus_core::watcher::find_session_index_cleanup_blocking_processes().is_empty() + } + + async fn run_remote_control_session_recovery(&self) -> anyhow::Result<()> { + let outcomes = tokio::task::spawn_blocking(|| { + let requests = codex_plus_core::remote_control_recovery::load_pending_remote_control_recoveries(None)?; + let settings = codex_plus_core::settings::SettingsStore::default() + .load()?; + let mut outcomes = Vec::with_capacity(requests.len()); + for request in requests { + let current_profile = settings + .relay_profiles + .iter() + .find(|profile| profile.id == request.profile_id); + let request_is_current = settings.active_relay_id == request.profile_id + && current_profile.is_some_and(|profile| { + codex_plus_core::remote_control_recovery::config_generation( + profile, + &request.target_provider, + ) == request.config_generation + }); + if !request_is_current { + outcomes.push(( + request, + codex_plus_data::ProviderSyncResult { + status: codex_plus_data::ProviderSyncStatus::Skipped, + message: "Remote Control session finalization deferred after relay profile changed".to_string(), + target_provider: String::new(), + backup_dir: None, + changed_session_files: 0, + sqlite_rows_updated: 0, + sqlite_provider_rows_updated: 0, + sqlite_user_event_rows_updated: 0, + sqlite_cwd_rows_updated: 0, + sqlite_catalog_rows_inserted: 0, + updated_workspace_roots: 0, + skipped_locked_rollout_files: Vec::new(), + encrypted_content_warning: None, + }, + None, + )); + continue; + } + let result = codex_plus_data::run_remote_control_session_finalization_for_thread_with_target( + None, + &request.thread_id, + &request.target_provider, + ); + let completed = result.status == codex_plus_data::ProviderSyncStatus::Synced; + let completion_error = if completed { + codex_plus_core::remote_control_recovery::complete_pending_remote_control_recovery( + None, + &request.thread_id, + ) + .err() + .map(|error| error.to_string()) + } else { + None + }; + outcomes.push((request, result, completion_error)); + } + Ok::<_, anyhow::Error>(outcomes) + }) + .await + .map_err(|error| anyhow::anyhow!("Remote Control session recovery task failed: {error}"))?; + match outcomes { + Ok(outcomes) => { + for (request, result, completion_error) in outcomes { + let _ = codex_plus_core::diagnostic_log::append_diagnostic_log( + "launcher.remote_control_session_finalization", + json!({ + "thread_id": request.thread_id, + "profile_id": request.profile_id, + "target_provider": request.target_provider, + "config_generation": request.config_generation, + "status": result.status, + "message": result.message, + "completion_error": completion_error + }), + ); + } + } + Err(error) => { + let _ = codex_plus_core::diagnostic_log::append_diagnostic_log( + "launcher.remote_control_session_finalization_failed_nonfatal", + json!({"message": error.to_string()}), + ); + } + } + Ok(()) + } + async fn apply_active_relay_profile( &self, settings: &codex_plus_core::settings::BackendSettings, @@ -514,6 +633,74 @@ impl BridgeDataService for LauncherDataService { .await .map_err(|error| anyhow::anyhow!("thread sort keys task failed: {error}")) } + + async fn recover_remote_control_session(&self, thread_id: String) -> anyhow::Result { + let settings = codex_plus_core::settings::SettingsStore::default() + .load() + .unwrap_or_default(); + let profile = settings.active_relay_profile(); + if !settings.relay_profiles_enabled + || profile.relay_mode != codex_plus_core::settings::RelayMode::Official + || !profile.official_mix_api_key + { + return Ok(json!({ + "status": "skipped", + "message": "Remote Control session recovery is disabled for the active profile" + })); + } + let home = codex_plus_core::codex_sqlite::default_codex_home_dir(); + let target_provider = + codex_plus_core::model_catalog::codex_model_provider_for_relay_profile(&home, &profile); + if target_provider.trim().is_empty() || target_provider == "openai" { + return Ok(json!({ + "status": "skipped", + "message": "Remote Control session recovery requires a non-openai target provider" + })); + } + let candidate_thread_id = thread_id.clone(); + let candidate = tokio::task::spawn_blocking(move || { + codex_plus_data::remote_control_session_recovery_candidate_exists( + None, + &candidate_thread_id, + ) + }) + .await + .map_err(|error| anyhow::anyhow!("Remote Control candidate check failed: {error}"))??; + if !candidate { + return Ok(json!({ + "status": "skipped", + "message": "Remote Control session recovery is waiting for a recent openai thread" + })); + } + let request = codex_plus_core::remote_control_recovery::PendingRemoteControlRecovery { + thread_id: thread_id.clone(), + profile_id: profile.id.clone(), + target_provider: target_provider.clone(), + config_generation: codex_plus_core::remote_control_recovery::config_generation( + &profile, + &target_provider, + ), + created_at: std::time::SystemTime::now() + .duration_since(std::time::UNIX_EPOCH) + .unwrap_or_default() + .as_secs() as i64, + }; + codex_plus_core::remote_control_recovery::enqueue_pending_remote_control_recovery( + None, request, + )?; + tokio::task::spawn_blocking(move || { + serde_json::to_value( + codex_plus_data::run_remote_control_session_catalog_recovery_for_thread_with_target( + None, + &thread_id, + &target_provider, + ), + ) + .map_err(anyhow::Error::from) + }) + .await + .map_err(|error| anyhow::anyhow!("Remote Control session recovery task failed: {error}"))? + } } impl LauncherDataService { @@ -864,6 +1051,40 @@ mod tests { assert!(source.contains("launcher.already_running")); } + #[test] + fn existing_launcher_path_drains_pending_remote_control_recovery_before_activation() { + let source = include_str!("main.rs"); + let start = source + .find("async fn activate_existing_codex_app") + .expect("existing launcher activation function"); + let body = &source[start..]; + let recovery = body + .find( + "let has_pending_recovery = hooks.has_pending_remote_control_session_recoveries()", + ) + .expect("pending recovery guard"); + let launch = body + .find("let launch_result = hooks") + .expect("Codex activation"); + + assert!(recovery < launch); + assert!(body[recovery..launch].contains("find_session_index_cleanup_blocking_processes")); + assert!(body[recovery..launch].contains("should_finalize_pending_remote_control_recovery")); + assert!( + body[recovery..launch].contains("hooks.run_remote_control_session_recovery().await?") + ); + } + + #[test] + fn pending_remote_control_finalization_requires_an_idle_desktop() { + assert!(should_finalize_pending_remote_control_recovery(true, &[])); + assert!(!should_finalize_pending_remote_control_recovery(false, &[])); + assert!(!should_finalize_pending_remote_control_recovery( + true, + &[42] + )); + } + #[test] fn launcher_hooks_forward_runtime_watchdogs_and_computer_use_guard_methods() { let source = include_str!("main.rs"); diff --git a/assets/inject/renderer-inject.js b/assets/inject/renderer-inject.js index 3341b98ea..2189e1f48 100644 --- a/assets/inject/renderer-inject.js +++ b/assets/inject/renderer-inject.js @@ -470,8 +470,9 @@ const codexThreadServiceTierKey = "codexThreadServiceTierOverrides"; const codexThreadServiceTierMaxEntries = 120; const codexThreadServiceTierDraftBindWindowMs = 60 * 1000; - const codexServiceTierRequestOverrideVersion = "5"; - const codexAppServerModelRequestPatchVersion = "4"; + const codexServiceTierRequestOverrideVersion = "8"; + const codexAppServerModelRequestPatchVersion = "5"; + const codexRemoteSessionRecoveryVersion = "4"; const codexPluginMarketplaceUnlockVersion = "15"; const codexThreadScrollMaxEntries = 120; const codexThreadScrollSaveThrottleMs = 120; @@ -3100,9 +3101,10 @@ } function applyCodexServiceTierRequestOverride(method, params, threadIdHint = "") { + const providerParams = applyCodexRemoteSessionProviderOverride(method, params); const override = codexServiceTierOverrideForRequest(method, params, threadIdHint); - if (!override) return params; - const nextParams = { ...(params || {}), serviceTier: override.serviceTier }; + if (!override) return providerParams; + const nextParams = { ...(providerParams || {}), serviceTier: override.serviceTier }; if (Object.prototype.hasOwnProperty.call(nextParams, "service_tier") || override.fastBlocked) { nextParams.service_tier = override.serviceTier; } @@ -3118,8 +3120,264 @@ return nextParams; } + function codexRemoteSessionProviderNormalizationEnabled() { + if (!codexPlusBackendSettings.relayProfilesEnabled) return false; + const profiles = Array.isArray(codexPlusBackendSettings.relayProfiles) + ? codexPlusBackendSettings.relayProfiles + : []; + const activeId = String(codexPlusBackendSettings.activeRelayId || ""); + const profile = profiles.find((item) => String(item?.id || "") === activeId); + if (!profile) return false; + const relayMode = String(profile.relayMode || ""); + return relayMode === "official" && !!profile.officialMixApiKey; + } + + function codexRemoteSessionTargetProvider() { + return String( + codexModelCatalog?.codex_model_provider + || codexModelCatalog?.codexModelProvider + || codexModelCatalog?.model_provider + || codexModelCatalog?.modelProvider + || "" + ).trim(); + } + + function codexRemoteSessionThreadStartMethod(method) { + return [ + "thread/start", + "start-conversation", + "start-thread-for-host", + "thread-prewarm-start", + "prewarm-thread-start-for-host", + ].includes(String(method || "")); + } + + function applyCodexRemoteSessionProviderOverride(method, params) { + if (!codexRemoteSessionThreadStartMethod(method)) return params; + if (!codexRemoteSessionProviderNormalizationEnabled()) return params; + if (!params || typeof params !== "object" || Array.isArray(params)) return params; + const targetProvider = codexRemoteSessionTargetProvider(); + if (!targetProvider || targetProvider === "openai") return params; + const requestedProvider = String(params.modelProvider || params.model_provider || "").trim(); + if (requestedProvider && requestedProvider !== "openai" && requestedProvider !== targetProvider) { + return params; + } + if (requestedProvider === targetProvider && !Object.prototype.hasOwnProperty.call(params, "model_provider")) { + return params; + } + const nextParams = { ...params, modelProvider: targetProvider }; + delete nextParams.model_provider; + sendCodexPlusDiagnostic("remote_session_provider_override_applied", { + method, + from: requestedProvider || "(missing)", + to: targetProvider, + }); + return nextParams; + } + + function codexRemoteSessionStartedThreadId(value) { + const queue = [{ value, depth: 0 }]; + const seen = new WeakSet(); + while (queue.length > 0) { + const current = queue.shift(); + const candidate = current?.value; + if (!candidate || typeof candidate !== "object") continue; + if (seen.has(candidate)) continue; + seen.add(candidate); + const method = String(candidate.method || candidate.type || ""); + if (method === "thread/started") { + const thread = candidate.params?.thread || candidate.thread || candidate.payload?.thread; + const threadId = String(thread?.id || candidate.params?.threadId || candidate.threadId || "").trim(); + if (threadId) return threadId; + } + if (method === "browser-use-session-route-capture") { + const threadId = String( + candidate.params?.conversationId + || candidate.params?.conversation_id + || candidate.conversationId + || candidate.conversation_id + || "" + ).trim(); + if (threadId) return threadId; + } + if (method === "browser-sidebar-browser-use-state") { + const isActive = candidate.params?.isActive ?? candidate.params?.is_active + ?? candidate.isActive ?? candidate.is_active; + if (isActive !== true) continue; + const threadId = String( + candidate.params?.conversationId + || candidate.params?.conversation_id + || candidate.conversationId + || candidate.conversation_id + || "" + ).trim(); + if (threadId) return threadId; + } + if (current.depth >= 4) continue; + for (const key of ["message", "response", "detail", "data", "payload", "params", "request"]) { + const nested = candidate[key]; + if (nested && typeof nested === "object") { + queue.push({ value: nested, depth: current.depth + 1 }); + } + } + } + return ""; + } + + function requestCodexRemoteSessionRecovery(threadId, attempt) { + const payload = { thread_id: threadId }; + const testHook = window.__CODEX_PLUS_TEST_REMOTE_RECOVERY__; + const request = typeof testHook === "function" + ? Promise.resolve(testHook(payload, attempt)) + : postJson("/remote-control-session/recover", payload); + return request.then((result) => { + if (attempt === 0 + || result?.message === "Remote Control session recovery complete" + || result?.message === "Remote Control session catalog recovery complete") { + sendCodexPlusDiagnostic("remote_session_recovery_requested", { + threadId, + attempt, + status: result?.status || "", + message: result?.message || "", + changedSessionFiles: result?.changed_session_files || 0, + catalogRowsInserted: result?.sqlite_catalog_rows_inserted || 0, + }); + } + return result; + }).catch((error) => { + if (attempt === 0) { + sendCodexPlusDiagnostic("remote_session_recovery_failed", { + threadId, + attempt, + errorName: error?.name || "", + errorMessage: error?.message || String(error), + }); + } + return null; + }); + } + + function scheduleCodexRemoteSessionRecovery(threadId) { + if (!codexRemoteSessionProviderNormalizationEnabled()) return false; + const normalizedThreadId = String(threadId || "").trim(); + if (!normalizedThreadId || normalizedThreadId.length > 128) return false; + window.__codexPlusRemoteSessionRecoveryPending = window.__codexPlusRemoteSessionRecoveryPending || new Map(); + const pending = window.__codexPlusRemoteSessionRecoveryPending; + if (pending.has(normalizedThreadId)) return false; + const retryOffsets = [100, 350, 800, 1600, 3000]; + const state = { timer: 0 }; + const finish = () => { + if (state.timer) window.clearTimeout(state.timer); + state.timer = 0; + if (pending.get(normalizedThreadId) === state) pending.delete(normalizedThreadId); + }; + const runAttempt = async (attempt) => { + state.timer = 0; + if (!codexRemoteSessionProviderNormalizationEnabled()) { + finish(); + return; + } + const result = await requestCodexRemoteSessionRecovery(normalizedThreadId, attempt); + const message = String(result?.message || ""); + if (message === "Remote Control session recovery complete" + || message === "Remote Control session catalog recovery complete" + || message === "Remote Control session recovery is disabled for the active profile") { + finish(); + return; + } + const nextAttempt = attempt + 1; + if (nextAttempt >= retryOffsets.length) { + finish(); + return; + } + const nextDelay = retryOffsets[nextAttempt] - retryOffsets[attempt]; + state.timer = window.setTimeout(() => void runAttempt(nextAttempt), nextDelay); + }; + state.timer = window.setTimeout(() => void runAttempt(0), retryOffsets[0]); + pending.set(normalizedThreadId, state); + return true; + } + + function observeCodexRemoteSessionNotification(value) { + const threadId = codexRemoteSessionStartedThreadId(value); + return threadId ? scheduleCodexRemoteSessionRecovery(threadId) : false; + } + + function installCodexRemoteSessionRecoveryListener() { + if (window.__codexPlusRemoteSessionRecoveryInstalled === codexRemoteSessionRecoveryVersion) return true; + if (window.__codexPlusRemoteSessionRecoveryMessageHandler) { + window.removeEventListener("message", window.__codexPlusRemoteSessionRecoveryMessageHandler, true); + } + if (window.__codexPlusRemoteSessionRecoveryViewHandler) { + window.removeEventListener("codex-message-from-view", window.__codexPlusRemoteSessionRecoveryViewHandler, true); + } + const messageHandler = (event) => { + if (event?.source !== window) return false; + const origin = String(event?.origin || ""); + if (origin && origin !== "null" && origin !== window.location.origin) return false; + return observeCodexRemoteSessionNotification(event?.data); + }; + const viewHandler = (event) => observeCodexRemoteSessionNotification(event?.detail); + window.__codexPlusRemoteSessionRecoveryMessageHandler = messageHandler; + window.__codexPlusRemoteSessionRecoveryViewHandler = viewHandler; + window.addEventListener("message", messageHandler, true); + window.addEventListener("codex-message-from-view", viewHandler, true); + window.__codexPlusRemoteSessionRecoveryInstalled = codexRemoteSessionRecoveryVersion; + sendCodexPlusDiagnostic("remote_session_recovery_listener_installed", { + version: codexRemoteSessionRecoveryVersion, + }); + return true; + } + + function installCodexRemoteSessionDispatcherSubscription(dispatcher, assetPrefix = "") { + if (!dispatcher || typeof dispatcher.subscribe !== "function") return false; + if (window.__codexPlusRemoteSessionRecoveryDispatcher === dispatcher + && window.__codexPlusRemoteSessionRecoveryDispatcherVersion === codexRemoteSessionRecoveryVersion) { + return true; + } + if (typeof window.__codexPlusRemoteSessionRecoveryDispatcherUnsubscribe === "function") { + try { + window.__codexPlusRemoteSessionRecoveryDispatcherUnsubscribe(); + } catch { + } + } + const handler = (payload) => { + if (observeCodexRemoteSessionNotification(payload)) return true; + const params = payload && typeof payload === "object" ? payload : {}; + if (observeCodexRemoteSessionNotification({ + method: "thread/started", + params, + })) return true; + return observeCodexRemoteSessionNotification({ + method: "thread/started", + params: { thread: params }, + }); + }; + const browserUseHandler = (payload) => observeCodexRemoteSessionNotification({ + type: "browser-sidebar-browser-use-state", + params: payload && typeof payload === "object" ? payload : {}, + }); + const unsubscribers = [ + dispatcher.subscribe("thread/started", handler), + dispatcher.subscribe("browser-sidebar-browser-use-state", browserUseHandler), + ]; + window.__codexPlusRemoteSessionRecoveryDispatcher = dispatcher; + window.__codexPlusRemoteSessionRecoveryDispatcherHandler = handler; + window.__codexPlusRemoteSessionRecoveryDispatcherUnsubscribe = () => { + for (const unsubscribe of unsubscribers) { + if (typeof unsubscribe !== "function") continue; + try { + unsubscribe(); + } catch { + } + } + }; + window.__codexPlusRemoteSessionRecoveryDispatcherVersion = codexRemoteSessionRecoveryVersion; + sendCodexPlusDiagnostic("remote_session_dispatcher_subscription_installed", { assetPrefix }); + return true; + } + function codexServiceTierRequestOverride(message, skipFetchEnvelope = false) { - if (!codexPlusSettings().serviceTierControls) return message; if (!message || typeof message !== "object") return message; if (!skipFetchEnvelope && message.type === "fetch" && typeof message.url === "string") { const urlPrefix = "vscode://codex/"; @@ -3239,6 +3497,7 @@ dispatcher.dispatchMessage = (type, payload) => { return dispatchCodexPlusMessage(dispatcher, type, payload); }; + installCodexRemoteSessionDispatcherSubscription(dispatcher, assetPrefix); window.__codexServiceTierRequestOverrideInstalled = codexServiceTierRequestOverrideVersion; sendCodexPlusDiagnostic("service_tier_dispatcher_patch_installed", { assetPrefix }); } catch (error) { @@ -3263,6 +3522,9 @@ } codexPlusBackendSettings = { ...codexPlusBackendSettings, ...settings }; codexPlusBackendSettingsLoaded = true; + if (codexRemoteSessionProviderNormalizationEnabled()) { + void loadCodexModelCatalog(); + } refreshCodexPlusBackendToggles(); return true; } catch (_) { @@ -5754,6 +6016,9 @@ const message = codexServiceTierRequestOverride({ ...(payload || {}), type }); const nextType = message?.type || type; const { type: _type, ...nextPayload } = message || {}; + if (nextType === "browser-use-session-route-capture") { + observeCodexRemoteSessionNotification({ type: nextType, params: nextPayload }); + } return dispatcher.__codexServiceTierOriginalDispatchMessage(nextType, nextPayload); } @@ -5765,7 +6030,7 @@ return Array.from(new Set(values.filter((value) => typeof value === "string" && value.trim().length > 0))); } - let codexModelCatalog = { status: "loading", model: "", default_model: "", model_provider: "", provider_name: "", models: [], sources: [], responses_api: { status: "unknown", message: "" } }; + let codexModelCatalog = { status: "loading", model: "", default_model: "", model_provider: "", codex_model_provider: "", provider_name: "", models: [], sources: [], responses_api: { status: "unknown", message: "" } }; let codexModelCatalogLoadedAt = 0; let codexModelCatalogPromise = null; let codexModelWhitelistRefreshTimer = 0; @@ -5775,6 +6040,12 @@ if (window.__CODEX_PLUS_TEST_SERVICE_TIER__) { window.__codexPlusServiceTierTest = { applyServiceTierOverride: (method, params, threadIdHint = "") => applyCodexServiceTierRequestOverride(method, params, threadIdHint), + applyProviderOverride: (method, params) => applyCodexRemoteSessionProviderOverride(method, params), + remoteSessionStartedThreadId: (value) => codexRemoteSessionStartedThreadId(value), + observeRemoteSessionNotification: (value) => observeCodexRemoteSessionNotification(value), + installRemoteSessionRecoveryListener: () => installCodexRemoteSessionRecoveryListener(), + installRemoteSessionDispatcherSubscription: (dispatcher, assetPrefix = "test") => installCodexRemoteSessionDispatcherSubscription(dispatcher, assetPrefix), + dispatchMessage: (dispatcher, type, payload) => dispatchCodexPlusMessage(dispatcher, type, payload), requestOverride: (message) => codexServiceTierRequestOverride(message), diagnostics: () => [...(window.__codexPlusServiceTierTestDiagnostics || [])], statusSummary: (state = {}) => { @@ -5798,6 +6069,7 @@ model: "", default_model: "", model_provider: "", + codex_model_provider: "", provider_name: "", models: [], sources: [], @@ -5807,6 +6079,10 @@ codexModelCatalogLoadedAt = Date.now(); codexModelCatalogPromise = null; }, + setBackendSettings: (settings = {}) => { + codexPlusBackendSettings = { ...codexPlusBackendSettings, ...settings }; + codexPlusBackendSettingsLoaded = true; + }, setServiceTierState: (state = {}) => { codexServiceTierState = { ...codexServiceTierState, ...state }; }, @@ -5844,7 +6120,7 @@ if (!force && codexModelCatalogLoadedAt && Date.now() - codexModelCatalogLoadedAt < 10000) return codexModelCatalog; codexModelCatalogPromise = postJson("/codex-model-catalog", {}) .then(async (result) => { - codexModelCatalog = result && typeof result === "object" ? result : { status: "failed", model: "", default_model: "", model_provider: "", provider_name: "", models: [], sources: [], responses_api: { status: "unknown", message: "" } }; + codexModelCatalog = result && typeof result === "object" ? result : { status: "failed", model: "", default_model: "", model_provider: "", codex_model_provider: "", provider_name: "", models: [], sources: [], responses_api: { status: "unknown", message: "" } }; if ((!codexModelCatalog.models || codexModelCatalog.models.length === 0) && codexModelCatalog.status === "not_configured") { try { const settingsPromise = postJson("/settings/get", {}); @@ -5852,7 +6128,7 @@ const settingsResp = await Promise.race([settingsPromise, timeoutPromise]); if (settingsResp && settingsResp.relayProfiles && Array.isArray(settingsResp.relayProfiles)) { const activeId = settingsResp.activeRelayId || ""; - const profile = settingsResp.relayProfiles.find(p => p.id === activeId) || settingsResp.relayProfiles[0]; + const profile = settingsResp.relayProfiles.find(p => p.id === activeId); if (profile && profile.modelList) { const extraModels = profile.modelList.split(/[\r\n,]+/).map(s => s.trim()).filter(Boolean); if (extraModels.length > 0) { @@ -5872,7 +6148,7 @@ return codexModelCatalog; }) .catch((error) => { - codexModelCatalog = { status: "failed", message: String(error?.message || error), model: "", default_model: "", model_provider: "", provider_name: "", models: [], sources: [], responses_api: { status: "unknown", message: "" } }; + codexModelCatalog = { status: "failed", message: String(error?.message || error), model: "", default_model: "", model_provider: "", codex_model_provider: "", provider_name: "", models: [], sources: [], responses_api: { status: "unknown", message: "" } }; codexModelCatalogLoadedAt = Date.now(); return codexModelCatalog; }) @@ -6234,10 +6510,17 @@ const originalSendRequest = client.__codexPlusModelOriginalSendRequest || client.sendRequest.bind(client); client.__codexPlusModelOriginalSendRequest = originalSendRequest; client.sendRequest = async function codexPlusModelPatchedSendRequest(method, params, options) { - const result = await originalSendRequest(method, params, options); + const requestMethod = appServerModelRequestMethod(String(method || ""), params); + if (codexRemoteSessionThreadStartMethod(requestMethod) + && codexRemoteSessionProviderNormalizationEnabled() + && !codexRemoteSessionTargetProvider()) { + await loadCodexModelCatalog(); + } + const nextParams = applyCodexRemoteSessionProviderOverride(requestMethod, params); + const result = await originalSendRequest(method, nextParams, options); if (!codexPlusModelUnlockEnabled()) return result; if (!codexPlusModelNames().length) await loadCodexModelCatalog(); - return patchAppServerModelResult(appServerModelRequestMethod(String(method || ""), params), result); + return patchAppServerModelResult(requestMethod, result); }; client.__codexPlusModelRequestPatch = codexAppServerModelRequestPatchVersion; return true; @@ -6246,6 +6529,17 @@ const appServerModelRequestPatchMaxMisses = 8; let appServerModelRequestPatchMissCount = 0; let appServerModelRequestPatchDisabled = false; + let appServerModelRequestPatchPromise = null; + let appServerModelRequestPatchRetryTimer = 0; + + function scheduleAppServerModelRequestPatchRetry() { + if (!codexRemoteSessionProviderNormalizationEnabled()) return; + if (appServerModelRequestPatchRetryTimer) return; + appServerModelRequestPatchRetryTimer = window.setTimeout(() => { + appServerModelRequestPatchRetryTimer = 0; + installAppServerModelRequestPatch(); + }, 250); + } function noteAppServerModelRequestPatchMiss(event, detail) { appServerModelRequestPatchMissCount += 1; @@ -6261,6 +6555,10 @@ if (appServerModelRequestPatchMissCount === 1) { sendCodexPlusDiagnostic(event, detail); } + if (codexRemoteSessionProviderNormalizationEnabled()) { + scheduleAppServerModelRequestPatchRetry(); + return; + } if (appServerModelRequestPatchMissCount >= appServerModelRequestPatchMaxMisses && !appServerModelRequestPatchDisabled) { appServerModelRequestPatchDisabled = true; sendCodexPlusDiagnostic("model_app_server_request_patch_skipped", { @@ -6273,12 +6571,12 @@ function installAppServerModelRequestPatch() { if (window.__codexPlusAppServerModelRequestPatchInstalled === codexAppServerModelRequestPatchVersion) return; if (appServerModelRequestPatchDisabled) return; + if (appServerModelRequestPatchPromise) return; const patch = async () => { try { const { modules, candidates, sources, discovery } = await loadAppServerRequestCandidates(); if (modules.length === 0) { - window.__codexPlusAppServerModelRequestPatchInstalled = codexAppServerModelRequestPatchVersion; - sendCodexPlusDiagnostic("model_app_server_request_patch_skipped", { + noteAppServerModelRequestPatchMiss("model_app_server_request_patch_skipped", { reason: "app_server_request_assets_missing", }); return; @@ -6288,6 +6586,8 @@ if (patchAppServerModelRequestClient(candidate)) patchedCount += 1; } if (patchedCount > 0) { + clearTimeout(appServerModelRequestPatchRetryTimer); + appServerModelRequestPatchRetryTimer = 0; appServerModelRequestPatchMissCount = 0; window.__codexPlusAppServerModelRequestPatchInstalled = codexAppServerModelRequestPatchVersion; sendCodexPlusDiagnostic("model_app_server_request_patch_installed", { @@ -6312,14 +6612,20 @@ }); } }; - void patch(); + appServerModelRequestPatchPromise = patch().finally(() => { + appServerModelRequestPatchPromise = null; + }); + void appServerModelRequestPatchPromise; } function ensureCodexModelWhitelistInstalls() { + if (codexPlusModelUnlockEnabled() + || (codexPlusBackendSettingsLoaded && codexRemoteSessionProviderNormalizationEnabled())) { + installAppServerModelRequestPatch(); + } if (!codexPlusModelUnlockEnabled()) return; installModelJsonResponsePatch(); patchAppServerModelMessages(); - installAppServerModelRequestPatch(); } function runCodexModelWhitelistRefreshPass() { @@ -9203,6 +9509,13 @@ function scanLightweight() { installStyle(); installCodexServiceTierDispatcherPatch(); + installCodexRemoteSessionRecoveryListener(); + if (window.__codexPlusRemoteSessionRecoveryDispatcher) { + installCodexRemoteSessionDispatcherSubscription( + window.__codexPlusRemoteSessionRecoveryDispatcher, + "existing-renderer" + ); + } installCodexPlusMenu(); localizeCodexMenus(); scheduleBackendHeartbeat(); diff --git a/crates/codex-plus-core/src/launcher.rs b/crates/codex-plus-core/src/launcher.rs index fc3e985be..5da93be86 100644 --- a/crates/codex-plus-core/src/launcher.rs +++ b/crates/codex-plus-core/src/launcher.rs @@ -145,6 +145,13 @@ pub trait LaunchHooks: Send + Sync { fn select_helper_port(&self, requested: u16) -> u16; async fn load_settings(&self) -> anyhow::Result; async fn run_provider_sync(&self) -> anyhow::Result<()>; + fn has_pending_remote_control_session_recoveries(&self) -> bool { + false + } + fn remote_control_session_recovery_is_safe_to_run(&self) -> bool { + true + } + async fn run_remote_control_session_recovery(&self) -> anyhow::Result<()>; async fn apply_active_relay_profile(&self, _settings: &BackendSettings) -> anyhow::Result<()> { Ok(()) } @@ -285,6 +292,16 @@ where "launcher.after_provider_sync", ); } + if hooks.has_pending_remote_control_session_recoveries() + && hooks.remote_control_session_recovery_is_safe_to_run() + { + hooks.run_remote_control_session_recovery().await?; + } else if hooks.has_pending_remote_control_session_recoveries() { + let _ = crate::diagnostic_log::append_diagnostic_log( + "launcher.remote_control_session_finalization_deferred", + serde_json::json!({"reason": "desktop_writer_active"}), + ); + } crate::dream_skin::sync_default_dream_skin_base_theme( settings.enhancements_enabled && settings.codex_app_dream_skin_enabled @@ -548,6 +565,16 @@ impl LaunchHooks for DefaultLaunchHooks { anyhow::bail!("provider sync requires launcher hooks with codex-plus-data integration") } + async fn run_remote_control_session_recovery(&self) -> anyhow::Result<()> { + anyhow::bail!( + "Remote Control session recovery requires launcher hooks with codex-plus-data integration" + ) + } + + fn remote_control_session_recovery_is_safe_to_run(&self) -> bool { + crate::watcher::find_session_index_cleanup_blocking_processes().is_empty() + } + async fn apply_active_relay_profile(&self, settings: &BackendSettings) -> anyhow::Result<()> { if !settings.relay_profiles_enabled { return Ok(()); diff --git a/crates/codex-plus-core/src/lib.rs b/crates/codex-plus-core/src/lib.rs index 0e0b3ad49..692380039 100644 --- a/crates/codex-plus-core/src/lib.rs +++ b/crates/codex-plus-core/src/lib.rs @@ -34,6 +34,7 @@ pub mod relay_config; pub mod relay_environment; pub mod relay_rotation; pub mod relay_switch; +pub mod remote_control_recovery; pub mod routes; pub mod script_market; pub mod settings; diff --git a/crates/codex-plus-core/src/model_catalog.rs b/crates/codex-plus-core/src/model_catalog.rs index 591546ccf..e526315b4 100644 --- a/crates/codex-plus-core/src/model_catalog.rs +++ b/crates/codex-plus-core/src/model_catalog.rs @@ -76,6 +76,7 @@ pub async fn read_codex_model_catalog() -> Value { fn relay_profile_model_catalog_value(home: &Path, profile: &RelayProfile) -> Value { let models = relay_profile_model_ids(profile); let model = profile.model.trim().to_string(); + let codex_model_provider = codex_model_provider_for_relay_profile(home, profile); let default_model = models.first().cloned().unwrap_or_default(); let provider_name = if profile.name.trim().is_empty() { profile.id.trim() @@ -90,6 +91,7 @@ fn relay_profile_model_catalog_value(home: &Path, profile: &RelayProfile) -> Val "service_tier": config_service_tier_value(home), "model": model, "model_provider": profile.id.trim(), + "codex_model_provider": codex_model_provider, "provider_name": provider_name, "default_model": default_model, "models": models, @@ -109,6 +111,20 @@ fn relay_profile_model_catalog_value(home: &Path, profile: &RelayProfile) -> Val }) } +pub fn codex_model_provider_for_relay_profile(home: &Path, profile: &RelayProfile) -> String { + let profile_config = parse_codex_config(&profile.config_contents); + let profile_provider = string_value(profile_config.root.get("model_provider")); + if !profile_provider.is_empty() { + return profile_provider; + } + + let (live_config, _, error) = load_codex_config(&home.join("config.toml")); + if error.is_some() { + return String::new(); + } + string_value(live_config.root.get("model_provider")) +} + fn relay_profile_model_ids(profile: &RelayProfile) -> Vec { unique_strings( profile diff --git a/crates/codex-plus-core/src/paths.rs b/crates/codex-plus-core/src/paths.rs index 36cc88e2e..a666794be 100644 --- a/crates/codex-plus-core/src/paths.rs +++ b/crates/codex-plus-core/src/paths.rs @@ -6,6 +6,7 @@ const SETTINGS_FILE: &str = "settings.json"; const LATEST_STATUS_FILE: &str = "latest-status.json"; const DIAGNOSTIC_LOG_FILE: &str = "codex-plus.log"; const PENDING_PROVIDER_IMPORT_FILE: &str = "pending-provider-import.json"; +const PENDING_REMOTE_CONTROL_RECOVERY_FILE: &str = "pending-remote-control-recovery.json"; pub fn default_app_state_dir() -> PathBuf { if let Some(home_dir) = directories::BaseDirs::new().map(|dirs| dirs.home_dir().to_path_buf()) { @@ -34,6 +35,10 @@ pub fn default_pending_provider_import_path() -> PathBuf { default_app_state_dir().join(PENDING_PROVIDER_IMPORT_FILE) } +pub fn default_pending_remote_control_recovery_path() -> PathBuf { + default_app_state_dir().join(PENDING_REMOTE_CONTROL_RECOVERY_FILE) +} + fn settings_path_for_tests() -> Option { SETTINGS_PATH_FOR_TESTS .get_or_init(|| Mutex::new(None)) @@ -95,4 +100,11 @@ mod tests { assert!(path.ends_with(".codex-session-delete/pending-provider-import.json")); } + + #[test] + fn default_pending_remote_control_recovery_path_uses_app_state_directory() { + let path = default_pending_remote_control_recovery_path(); + + assert!(path.ends_with(".codex-session-delete/pending-remote-control-recovery.json")); + } } diff --git a/crates/codex-plus-core/src/remote_control_recovery.rs b/crates/codex-plus-core/src/remote_control_recovery.rs new file mode 100644 index 000000000..a87902540 --- /dev/null +++ b/crates/codex-plus-core/src/remote_control_recovery.rs @@ -0,0 +1,216 @@ +use crate::settings::{RelayProfile, atomic_write}; +use serde::{Deserialize, Serialize}; +use sha2::{Digest, Sha256}; +use std::path::{Path, PathBuf}; +use std::sync::{Mutex, OnceLock}; + +const STATE_VERSION: u32 = 1; + +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +#[serde(rename_all = "camelCase")] +pub struct PendingRemoteControlRecovery { + pub thread_id: String, + pub profile_id: String, + pub target_provider: String, + pub config_generation: String, + pub created_at: i64, +} + +#[derive(Debug, Default, Serialize, Deserialize)] +#[serde(rename_all = "camelCase")] +struct PendingRemoteControlRecoveryState { + #[serde(default = "state_version")] + version: u32, + #[serde(default)] + requests: Vec, +} + +fn state_version() -> u32 { + STATE_VERSION +} + +fn state_lock() -> &'static Mutex<()> { + static LOCK: OnceLock> = OnceLock::new(); + LOCK.get_or_init(|| Mutex::new(())) +} + +pub fn config_generation(profile: &RelayProfile, target_provider: &str) -> String { + let mut digest = Sha256::new(); + digest.update(profile.id.trim().as_bytes()); + digest.update([0]); + digest.update(profile.config_contents.as_bytes()); + digest.update([0]); + digest.update(target_provider.trim().as_bytes()); + format!("{:x}", digest.finalize()) +} + +pub fn load_pending_remote_control_recoveries( + path: Option<&Path>, +) -> anyhow::Result> { + let path = pending_path(path); + let _guard = state_lock() + .lock() + .map_err(|_| anyhow::anyhow!("Remote Control recovery state lock poisoned"))?; + Ok(load_state(&path)?.requests) +} + +pub fn enqueue_pending_remote_control_recovery( + path: Option<&Path>, + request: PendingRemoteControlRecovery, +) -> anyhow::Result<()> { + validate_request(&request)?; + let path = pending_path(path); + let _guard = state_lock() + .lock() + .map_err(|_| anyhow::anyhow!("Remote Control recovery state lock poisoned"))?; + let mut state = load_state(&path)?; + // The first observed profile/provider snapshot is authoritative. Retries for the same + // thread must not overwrite it after the user switches relay profiles. + if state + .requests + .iter() + .any(|existing| existing.thread_id == request.thread_id) + { + return Ok(()); + } + state.requests.push(request); + save_state(&path, &state) +} + +pub fn complete_pending_remote_control_recovery( + path: Option<&Path>, + thread_id: &str, +) -> anyhow::Result<()> { + let path = pending_path(path); + let _guard = state_lock() + .lock() + .map_err(|_| anyhow::anyhow!("Remote Control recovery state lock poisoned"))?; + let mut state = load_state(&path)?; + let original_len = state.requests.len(); + state + .requests + .retain(|request| request.thread_id != thread_id); + if state.requests.len() == original_len { + return Ok(()); + } + save_state(&path, &state) +} + +fn pending_path(path: Option<&Path>) -> PathBuf { + path.map(Path::to_path_buf) + .unwrap_or_else(crate::paths::default_pending_remote_control_recovery_path) +} + +fn load_state(path: &Path) -> anyhow::Result { + match std::fs::read_to_string(path) { + Ok(text) => match serde_json::from_str(&text) { + Ok(state) => Ok(state), + Err(_) => { + let timestamp = std::time::SystemTime::now() + .duration_since(std::time::UNIX_EPOCH) + .unwrap_or_default() + .as_nanos(); + let corrupt_path = path.with_file_name(format!( + "{}.corrupt-{}-{timestamp}", + path.file_name() + .and_then(|name| name.to_str()) + .unwrap_or("pending.json"), + std::process::id() + )); + let _ = std::fs::rename(path, corrupt_path); + Ok(PendingRemoteControlRecoveryState { + version: STATE_VERSION, + requests: Vec::new(), + }) + } + }, + Err(error) if error.kind() == std::io::ErrorKind::NotFound => { + Ok(PendingRemoteControlRecoveryState { + version: STATE_VERSION, + requests: Vec::new(), + }) + } + Err(error) => Err(error.into()), + } +} + +fn save_state(path: &Path, state: &PendingRemoteControlRecoveryState) -> anyhow::Result<()> { + if state.requests.is_empty() { + match std::fs::remove_file(path) { + Ok(()) => return Ok(()), + Err(error) if error.kind() == std::io::ErrorKind::NotFound => return Ok(()), + Err(error) => return Err(error.into()), + } + } + atomic_write(path, serde_json::to_string_pretty(state)?.as_bytes()) +} + +fn validate_request(request: &PendingRemoteControlRecovery) -> anyhow::Result<()> { + if request.thread_id.trim().is_empty() || request.thread_id.len() > 128 { + anyhow::bail!("Remote Control recovery requires a valid thread id"); + } + if request.profile_id.trim().is_empty() || request.target_provider.trim().is_empty() { + anyhow::bail!("Remote Control recovery requires profile and provider provenance"); + } + Ok(()) +} + +#[cfg(test)] +mod tests { + use super::*; + use tempfile::tempdir; + + fn request(thread_id: &str) -> PendingRemoteControlRecovery { + PendingRemoteControlRecovery { + thread_id: thread_id.to_string(), + profile_id: "official-mix".to_string(), + target_provider: "custom".to_string(), + config_generation: "generation".to_string(), + created_at: 1, + } + } + + #[test] + fn pending_recovery_state_deduplicates_and_completes_by_thread() { + let dir = tempdir().unwrap(); + let path = dir.path().join("pending.json"); + enqueue_pending_remote_control_recovery(Some(&path), request("one")).unwrap(); + let replacement = request("one"); + enqueue_pending_remote_control_recovery(Some(&path), replacement.clone()).unwrap(); + enqueue_pending_remote_control_recovery(Some(&path), request("two")).unwrap(); + + assert_eq!( + load_pending_remote_control_recoveries(Some(&path)).unwrap(), + vec![request("one"), request("two")] + ); + + complete_pending_remote_control_recovery(Some(&path), "one").unwrap(); + assert_eq!( + load_pending_remote_control_recoveries(Some(&path)).unwrap(), + vec![request("two")] + ); + complete_pending_remote_control_recovery(Some(&path), "two").unwrap(); + assert!(!path.exists()); + } + + #[test] + fn corrupt_pending_recovery_state_is_quarantined() { + let dir = tempdir().unwrap(); + let path = dir.path().join("pending.json"); + std::fs::write(&path, "{broken").unwrap(); + + assert!( + load_pending_remote_control_recoveries(Some(&path)) + .unwrap() + .is_empty() + ); + assert!(!path.exists()); + assert!(std::fs::read_dir(dir.path()).unwrap().any(|entry| { + entry + .unwrap() + .file_name() + .to_string_lossy() + .starts_with("pending.json.corrupt-") + })); + } +} diff --git a/crates/codex-plus-core/src/routes.rs b/crates/codex-plus-core/src/routes.rs index 5c5308510..77d373feb 100644 --- a/crates/codex-plus-core/src/routes.rs +++ b/crates/codex-plus-core/src/routes.rs @@ -120,6 +120,9 @@ pub trait BridgeDataService: Send + Sync { ) -> anyhow::Result; async fn thread_sort_key(&self, session: SessionRef) -> anyhow::Result; async fn thread_sort_keys(&self, sessions: Vec) -> anyhow::Result; + async fn recover_remote_control_session(&self, _thread_id: String) -> anyhow::Result { + anyhow::bail!("Remote Control session recovery is unavailable") + } } pub async fn handle_bridge_request( @@ -265,6 +268,15 @@ pub async fn handle_bridge_request( .thread_sort_keys(sessions_from_payload(&payload)) .await } + "/remote-control-session/recover" => { + let thread_id = payload + .get("thread_id") + .or_else(|| payload.get("threadId")) + .and_then(Value::as_str) + .unwrap_or_default() + .to_string(); + ctx.data.recover_remote_control_session(thread_id).await + } _ => { let _ = crate::diagnostic_log::append_diagnostic_log( "bridge.unknown_path", diff --git a/crates/codex-plus-core/tests/bridge_routes.rs b/crates/codex-plus-core/tests/bridge_routes.rs index 9427b4dd5..c3f4a3575 100644 --- a/crates/codex-plus-core/tests/bridge_routes.rs +++ b/crates/codex-plus-core/tests/bridge_routes.rs @@ -1492,6 +1492,10 @@ impl LaunchHooks for ContextHooks { Ok(()) } + async fn run_remote_control_session_recovery(&self) -> anyhow::Result<()> { + Ok(()) + } + async fn start_helper(&self, _helper_port: u16) -> anyhow::Result<()> { Ok(()) } diff --git a/crates/codex-plus-core/tests/cdp_bridge.rs b/crates/codex-plus-core/tests/cdp_bridge.rs index eb1bcfe8b..79894ee36 100644 --- a/crates/codex-plus-core/tests/cdp_bridge.rs +++ b/crates/codex-plus-core/tests/cdp_bridge.rs @@ -1362,7 +1362,7 @@ fn injection_script_unlocks_custom_model_catalog() { assert!(script.contains("loadAppServerRequestCandidates")); assert!(script.contains("appServerFallbackAssetUrls")); assert!(script.contains("collectAppServerRequestCandidatesFromModule")); - assert!(script.contains("codexAppServerModelRequestPatchVersion = \"4\"")); + assert!(script.contains("codexAppServerModelRequestPatchVersion = \"5\"")); assert!(script.contains("list-models-for-host")); assert!(script.contains("appServerModelRequestMethod")); @@ -1373,6 +1373,7 @@ fn injection_script_unlocks_custom_model_catalog() { assert!(script.contains("model_whitelist_refresh_scheduled")); assert!(script.contains("available_models")); assert!(script.contains("modelWhitelistUnlock")); + assert!(!script.contains("|| settingsResp.relayProfiles[0]")); assert!(script.contains("refreshCodexModelWhitelistFromScan")); assert!(script.contains("codexPlusModelListRequestIds.size === 0")); assert!(!script.contains("function patchReactModelState")); @@ -1631,7 +1632,54 @@ fn injection_script_applies_fast_service_tier_contract() { assert_eq!(cases["legacyStateApi"], true); assert_eq!(cases["currentStateApi"], true); assert_eq!(cases["appServerParamsUnchanged"], true); - assert_eq!(cases["appServerSentCount"], 1); + assert_eq!(cases["appServerSentCount"], 2); + assert_eq!( + cases["providerFromMissing"]["modelProvider"], + "vendor_alpha" + ); + assert_eq!(cases["providerFromOpenAi"]["modelProvider"], "vendor_alpha"); + assert_eq!(cases["providerFromOtherUnchanged"], true); + assert_eq!(cases["nonThreadProviderUnchanged"], true); + assert_eq!( + cases["providerWithServiceTierControlsDisabled"]["modelProvider"], + "vendor_alpha" + ); + assert_eq!(cases["appServerProviderOverride"], "vendor_alpha"); + assert_eq!(cases["directThreadStartedId"], "thread-mobile-direct"); + assert_eq!(cases["nestedThreadStartedId"], "thread-mobile-nested"); + assert_eq!( + cases["browserUseRouteThreadId"], + "thread-mobile-browser-route" + ); + assert_eq!(cases["inactiveBrowserUseUnscheduled"], true); + assert_eq!(cases["remoteRecoveryScheduled"], true); + assert_eq!(cases["remoteRecoveryThreadId"], "thread-mobile-notify"); + assert_eq!(cases["remoteRecoveryCallCountAfterSuccess"], 1); + assert_eq!(cases["remoteRecoveryDispatcherInstalled"], true); + assert_eq!( + cases["remoteRecoveryDispatcherThreadId"], + "thread-mobile-dispatcher" + ); + assert_eq!( + cases["remoteRecoveryBrowserUseDispatcherThreadId"], + "thread-mobile-browser-dispatcher" + ); + assert_eq!( + cases["remoteRecoveryOutboundRouteThreadId"], + "thread-mobile-browser-outbound" + ); + assert_eq!(cases["remoteRecoveryListenerInstalled"], true); + assert_eq!( + cases["remoteRecoveryViewEventThreadId"], + "thread-mobile-view-event" + ); + assert_eq!(cases["remoteRecoveryRetried"], true); + assert_eq!(cases["remoteRecoveryRetryAttempts"], json!([0, 1])); + assert_eq!(cases["missingActiveProviderUnchanged"], true); + assert_eq!(cases["missingActiveRecoveryUnscheduled"], true); + assert_eq!(cases["pureApiProviderUnchanged"], true); + assert_eq!(cases["pureApiRecoveryUnscheduled"], true); + assert_eq!(cases["pureOfficialProviderUnchanged"], true); } fn run_service_tier_contract_harness() -> serde_json::Value { @@ -1669,6 +1717,11 @@ function node() {{ }} globalThis.window = globalThis; window.__CODEX_PLUS_TEST_SERVICE_TIER__ = true; +const windowListeners = new Map(); +window.addEventListener = (type, listener) => windowListeners.set(type, listener); +window.removeEventListener = (type, listener) => {{ + if (windowListeners.get(type) === listener) windowListeners.delete(type); +}}; globalThis.document = {{ scripts: [], documentElement: node(), @@ -1904,6 +1957,160 @@ const appServerParamsUnchanged = appServerCalls[0]?.params === nativeAppServerPa && appServerCalls[0]?.params?.workspaceKind === "project" && appServerCalls[0]?.params?.cwd === "C:/native/work" && appServerCalls[0]?.params?.projectAssignment?.projectId === "C:/native/work"; +api.setBackendSettings({{ + relayProfilesEnabled: true, + activeRelayId: "custom-relay", + relayProfiles: [{{ id: "custom-relay", relayMode: "official", officialMixApiKey: true }}], +}}); +api.setModelCatalog({{ + status: "ok", + model: "gpt-5.6-sol", + default_model: "gpt-5.6-sol", + model_provider: "relay-ms0ihvx9", + codex_model_provider: "vendor_alpha", + models: ["gpt-5.6-sol"], +}}); +localStorage.setItem("codexPlusSettings", JSON.stringify({{ serviceTierControls: false }})); +const providerFromMissing = api.applyProviderOverride("thread/start", {{ cwd: "C:/mobile" }}); +const providerFromOpenAi = api.applyProviderOverride("thread/start", {{ cwd: "C:/mobile", modelProvider: "openai" }}); +const explicitOtherProvider = {{ cwd: "C:/mobile", modelProvider: "other" }}; +const providerFromOther = api.applyProviderOverride("thread/start", explicitOtherProvider); +const nonThreadParams = {{ cwd: "C:/mobile", modelProvider: "openai" }}; +const nonThreadProviderUnchanged = api.applyProviderOverride("turn/start", nonThreadParams) === nonThreadParams; +const providerWithServiceTierControlsDisabled = api.requestOverride({{ + type: "start-conversation", + cwd: "C:/mobile", + modelProvider: "openai", +}}); +await appServerClient.sendRequest("thread/start", {{ cwd: "C:/mobile", modelProvider: "openai" }}, {{ signal: "mobile" }}); +const appServerProviderOverride = appServerCalls[1]?.params?.modelProvider; +const directThreadStartedId = api.remoteSessionStartedThreadId({{ + method: "thread/started", + params: {{ thread: {{ id: "thread-mobile-direct" }} }}, +}}); +const nestedThreadStartedId = api.remoteSessionStartedThreadId({{ + type: "mcp-response", + message: {{ method: "thread/started", params: {{ thread: {{ id: "thread-mobile-nested" }} }} }}, +}}); +const browserUseRouteThreadId = api.remoteSessionStartedThreadId({{ + type: "browser-use-session-route-capture", + conversationId: "thread-mobile-browser-route", +}}); +const inactiveBrowserUseUnscheduled = api.observeRemoteSessionNotification({{ + type: "browser-sidebar-browser-use-state", + conversationId: "thread-mobile-browser-inactive", + isActive: false, +}}) === false; +const remoteRecoveryCalls = []; +const remoteRecoveryDispatcherHandlers = new Map(); +const remoteRecoveryDispatcher = {{ + subscribe(type, callback) {{ + remoteRecoveryDispatcherHandlers.set(type, callback); + return () => remoteRecoveryDispatcherHandlers.delete(type); + }}, +}}; +window.__CODEX_PLUS_TEST_REMOTE_RECOVERY__ = (payload, attempt) => {{ + remoteRecoveryCalls.push({{ payload, attempt }}); + return {{ status: "synced", message: "Remote Control session catalog recovery complete" }}; +}}; +const remoteRecoveryScheduled = api.observeRemoteSessionNotification({{ + response: {{ method: "thread/started", params: {{ thread: {{ id: "thread-mobile-notify" }} }} }}, +}}); +await new Promise((resolve) => setTimeout(resolve, 500)); +const remoteRecoveryCallCountAfterSuccess = remoteRecoveryCalls.length; +const remoteRecoveryDispatcherCalls = []; +window.__CODEX_PLUS_TEST_REMOTE_RECOVERY__ = (payload, attempt) => {{ + remoteRecoveryDispatcherCalls.push({{ payload, attempt }}); + return {{ status: "synced", message: "Remote Control session catalog recovery complete" }}; +}}; +const remoteRecoveryDispatcherInstalled = api.installRemoteSessionDispatcherSubscription(remoteRecoveryDispatcher); +remoteRecoveryDispatcherHandlers.get("thread/started")?.({{ id: "thread-mobile-dispatcher" }}); +await new Promise((resolve) => setTimeout(resolve, 500)); +const remoteRecoveryDispatcherThreadId = remoteRecoveryDispatcherCalls[0]?.payload?.thread_id || ""; +const remoteRecoveryBrowserUseDispatcherCalls = []; +window.__CODEX_PLUS_TEST_REMOTE_RECOVERY__ = (payload, attempt) => {{ + remoteRecoveryBrowserUseDispatcherCalls.push({{ payload, attempt }}); + return {{ status: "synced", message: "Remote Control session catalog recovery complete" }}; +}}; +remoteRecoveryDispatcherHandlers.get("browser-sidebar-browser-use-state")?.({{ + conversationId: "thread-mobile-browser-dispatcher", + isActive: true, +}}); +await new Promise((resolve) => setTimeout(resolve, 500)); +const remoteRecoveryBrowserUseDispatcherThreadId = remoteRecoveryBrowserUseDispatcherCalls[0]?.payload?.thread_id || ""; +const remoteRecoveryOutboundRouteCalls = []; +window.__CODEX_PLUS_TEST_REMOTE_RECOVERY__ = (payload, attempt) => {{ + remoteRecoveryOutboundRouteCalls.push({{ payload, attempt }}); + return {{ status: "synced", message: "Remote Control session catalog recovery complete" }}; +}}; +const outboundDispatcherMessages = []; +const outboundDispatcher = {{ + __codexServiceTierOriginalDispatchMessage(type, payload) {{ + outboundDispatcherMessages.push({{ type, payload }}); + return true; + }}, +}}; +api.dispatchMessage(outboundDispatcher, "browser-use-session-route-capture", {{ + conversationId: "thread-mobile-browser-outbound", +}}); +await new Promise((resolve) => setTimeout(resolve, 500)); +const remoteRecoveryOutboundRouteThreadId = remoteRecoveryOutboundRouteCalls[0]?.payload?.thread_id || ""; +const remoteRecoveryViewEventCalls = []; +window.__CODEX_PLUS_TEST_REMOTE_RECOVERY__ = (payload, attempt) => {{ + remoteRecoveryViewEventCalls.push({{ payload, attempt }}); + return {{ status: "synced", message: "Remote Control session catalog recovery complete" }}; +}}; +const remoteRecoveryListenerInstalled = api.installRemoteSessionRecoveryListener(); +windowListeners.get("codex-message-from-view")?.({{ + detail: {{ + type: "browser-use-session-route-capture", + conversationId: "thread-mobile-view-event", + }}, +}}); +await new Promise((resolve) => setTimeout(resolve, 500)); +const remoteRecoveryViewEventThreadId = remoteRecoveryViewEventCalls[0]?.payload?.thread_id || ""; +const remoteRecoveryRetryCalls = []; +window.__CODEX_PLUS_TEST_REMOTE_RECOVERY__ = (payload, attempt) => {{ + remoteRecoveryRetryCalls.push({{ payload, attempt }}); + if (attempt === 0) {{ + return {{ status: "synced", message: "Remote Control session recovery already up to date" }}; + }} + return {{ status: "synced", message: "Remote Control session recovery complete" }}; +}}; +const remoteRecoveryRetried = api.observeRemoteSessionNotification({{ + response: {{ method: "thread/started", params: {{ thread: {{ id: "thread-mobile-retry" }} }} }}, +}}); +await new Promise((resolve) => setTimeout(resolve, 500)); +const remoteRecoveryRetryAttempts = remoteRecoveryRetryCalls.map((call) => call.attempt); +api.setBackendSettings({{ + relayProfilesEnabled: true, + activeRelayId: "missing", + relayProfiles: [{{ id: "eligible", relayMode: "official", officialMixApiKey: true }}], +}}); +const missingActiveParams = {{ cwd: "C:/mobile", modelProvider: "openai" }}; +const missingActiveProviderUnchanged = api.applyProviderOverride("thread/start", missingActiveParams) === missingActiveParams; +const missingActiveRecoveryUnscheduled = api.observeRemoteSessionNotification({{ + method: "thread/started", + params: {{ thread: {{ id: "thread-mobile-missing-active" }} }}, +}}) === false; +api.setBackendSettings({{ + relayProfilesEnabled: true, + activeRelayId: "pure-api", + relayProfiles: [{{ id: "pure-api", relayMode: "pureApi", officialMixApiKey: true }}], +}}); +const pureApiParams = {{ cwd: "C:/mobile", modelProvider: "openai" }}; +const pureApiProviderUnchanged = api.applyProviderOverride("thread/start", pureApiParams) === pureApiParams; +const pureApiRecoveryUnscheduled = api.observeRemoteSessionNotification({{ + method: "thread/started", + params: {{ thread: {{ id: "thread-mobile-pure-api" }} }}, +}}) === false; +api.setBackendSettings({{ + relayProfilesEnabled: true, + activeRelayId: "official", + relayProfiles: [{{ id: "official", relayMode: "official", officialMixApiKey: false }}], +}}); +const pureOfficialParams = {{ cwd: "C:/mobile", modelProvider: "openai" }}; +const pureOfficialProviderUnchanged = api.applyProviderOverride("thread/start", pureOfficialParams) === pureOfficialParams; process.stdout.write(JSON.stringify({{ supportedFast, unsupportedModel, @@ -1932,6 +2139,32 @@ process.stdout.write(JSON.stringify({{ currentStateApi, appServerParamsUnchanged, appServerSentCount: appServerCalls.length, + providerFromMissing, + providerFromOpenAi, + providerFromOtherUnchanged: providerFromOther === explicitOtherProvider, + nonThreadProviderUnchanged, + providerWithServiceTierControlsDisabled, + appServerProviderOverride, + directThreadStartedId, + nestedThreadStartedId, + browserUseRouteThreadId, + inactiveBrowserUseUnscheduled, + remoteRecoveryScheduled, + remoteRecoveryThreadId: remoteRecoveryCalls[0]?.payload?.thread_id || "", + remoteRecoveryCallCountAfterSuccess, + remoteRecoveryDispatcherInstalled, + remoteRecoveryDispatcherThreadId, + remoteRecoveryBrowserUseDispatcherThreadId, + remoteRecoveryOutboundRouteThreadId, + remoteRecoveryListenerInstalled, + remoteRecoveryViewEventThreadId, + remoteRecoveryRetried, + remoteRecoveryRetryAttempts, + missingActiveProviderUnchanged, + missingActiveRecoveryUnscheduled, + pureApiProviderUnchanged, + pureApiRecoveryUnscheduled, + pureOfficialProviderUnchanged, }})); }}).catch((error) => {{ console.error(error); @@ -1969,7 +2202,7 @@ fn injection_script_leaves_new_threads_to_the_codex_app() { assert!(!script.contains("hotkey-window-projectless-default-enabled")); assert!(script.contains("installCodexServiceTierDispatcherPatch")); assert!(script.contains("installAppServerModelRequestPatch")); - assert!(script.contains("originalSendRequest(method, params, options)")); + assert!(script.contains("originalSendRequest(method, nextParams, options)")); } #[test] diff --git a/crates/codex-plus-core/tests/launcher.rs b/crates/codex-plus-core/tests/launcher.rs index 6799bda60..1ee565dea 100644 --- a/crates/codex-plus-core/tests/launcher.rs +++ b/crates/codex-plus-core/tests/launcher.rs @@ -1187,12 +1187,44 @@ async fn official_mix_responses_profile_starts_fixed_protocol_proxy_without_enha handle.wait_for_codex_exit().await.unwrap(); let events = events.lock().unwrap().clone(); + assert!(!events.contains(&"remote-control-session-recovery".to_string())); + assert!(!events.contains(&"provider-sync".to_string())); assert!(events.contains(&"select-helper:58123".to_string())); assert!(events.contains(&"start-helper:57321".to_string())); assert!(events.contains(&"shutdown-helper:57321".to_string())); assert!(!events.iter().any(|event| event.starts_with("inject:"))); } +#[tokio::test] +async fn pending_remote_control_recovery_runs_without_an_official_mix_profile() { + let temp = tempfile::tempdir().unwrap(); + let app_dir = temp.path().join("Codex.app"); + std::fs::create_dir_all(&app_dir).unwrap(); + let status_store = StatusStore::new(temp.path().join("latest-status.json")); + let events = Arc::new(Mutex::new(Vec::::new())); + let hooks = FakeHooks::new(events.clone()).with_pending_remote_control_session_recoveries(); + + let handle = launch_and_inject_with_hooks( + LaunchOptions { + app_dir: Some(app_dir), + debug_port: 9229, + helper_port: 58123, + status_store, + }, + &hooks, + ) + .await + .unwrap(); + handle.wait_for_codex_exit().await.unwrap(); + + assert!( + events + .lock() + .unwrap() + .contains(&"remote-control-session-recovery".to_string()) + ); +} + #[tokio::test] async fn official_mix_responses_profile_keeps_proxy_when_profile_switching_is_disabled() { let temp = tempfile::tempdir().unwrap(); @@ -1740,6 +1772,7 @@ struct FakeHooks { inject_error: Option, provider_sync_unsupported: bool, plugin_marketplace_error: Option, + has_pending_remote_control_session_recoveries: bool, } impl FakeHooks { @@ -1756,6 +1789,7 @@ impl FakeHooks { inject_error: None, provider_sync_unsupported: false, plugin_marketplace_error: None, + has_pending_remote_control_session_recoveries: false, } } @@ -1789,6 +1823,11 @@ impl FakeHooks { self } + fn with_pending_remote_control_session_recoveries(mut self) -> Self { + self.has_pending_remote_control_session_recoveries = true; + self + } + fn event(&self, event: impl Into) { self.events.lock().unwrap().push(event.into()); } @@ -1829,6 +1868,15 @@ impl LaunchHooks for FakeHooks { Ok(()) } + fn has_pending_remote_control_session_recoveries(&self) -> bool { + self.has_pending_remote_control_session_recoveries + } + + async fn run_remote_control_session_recovery(&self) -> anyhow::Result<()> { + self.event("remote-control-session-recovery"); + Ok(()) + } + async fn apply_active_relay_profile(&self, settings: &BackendSettings) -> anyhow::Result<()> { if !settings.relay_profiles_enabled { return Ok(()); diff --git a/crates/codex-plus-core/tests/model_catalog.rs b/crates/codex-plus-core/tests/model_catalog.rs index 8fdc6e6f4..b42375547 100644 --- a/crates/codex-plus-core/tests/model_catalog.rs +++ b/crates/codex-plus-core/tests/model_catalog.rs @@ -186,7 +186,7 @@ experimental_bearer_token = "relay-key" } #[tokio::test] -async fn model_catalog_uses_active_relay_profile_model_list_for_display() { +async fn model_catalog_uses_active_relay_profile_model_list_and_actual_provider() { let temp = tempfile::tempdir().unwrap(); let codex_home = temp.path().join("codex-home"); std::fs::create_dir_all(&codex_home).unwrap(); @@ -198,27 +198,38 @@ async fn model_catalog_uses_active_relay_profile_model_list_for_display() { std::env::set_var("CODEX_HOME", &codex_home); } - let result = async { - SettingsStore::new(settings_path) - .save(&BackendSettings { - active_relay_id: "relay-a".to_string(), - relay_profiles: vec![RelayProfile { - id: "relay-a".to_string(), - name: "Relay A".to_string(), - model: "qwen3-coder".to_string(), - base_url: "https://example.test/v1".to_string(), - protocol: RelayProtocol::Responses, - relay_mode: RelayMode::MixedApi, - model_list: "deepseek-coder\nqwen3-coder\nclaude-compatible\ngpt-5.6-sol" - .to_string(), - config_contents: "model = \"qwen3-coder\"\n".to_string(), - ..RelayProfile::default() - }], - ..BackendSettings::default() - }) - .unwrap(); - - read_codex_model_catalog().await + let (result, live_fallback_result) = async { + write_config( + &codex_home, + "model = \"qwen3-coder\"\nmodel_provider = \"live_vendor\"\n", + ); + let store = SettingsStore::new(settings_path); + let mut settings = BackendSettings { + active_relay_id: "relay-a".to_string(), + relay_profiles: vec![RelayProfile { + id: "relay-a".to_string(), + name: "Relay A".to_string(), + model: "qwen3-coder".to_string(), + base_url: "https://example.test/v1".to_string(), + protocol: RelayProtocol::Responses, + relay_mode: RelayMode::MixedApi, + model_list: "deepseek-coder\nqwen3-coder\nclaude-compatible\ngpt-5.6-sol" + .to_string(), + config_contents: "model = \"qwen3-coder\"\nmodel_provider = \"vendor_alpha\"\n" + .to_string(), + ..RelayProfile::default() + }], + ..BackendSettings::default() + }; + store.save(&settings).unwrap(); + let result = read_codex_model_catalog().await; + + settings.relay_profiles[0].relay_mode = RelayMode::Official; + settings.relay_profiles[0].official_mix_api_key = false; + settings.relay_profiles[0].config_contents = "model = \"qwen3-coder\"\n".to_string(); + store.save(&settings).unwrap(); + let live_fallback_result = read_codex_model_catalog().await; + (result, live_fallback_result) } .await; @@ -234,6 +245,8 @@ async fn model_catalog_uses_active_relay_profile_model_list_for_display() { assert_eq!(result["status"], "ok"); assert_eq!(result["model_provider"], "relay-a"); + assert_eq!(result["codex_model_provider"], "vendor_alpha"); + assert_eq!(live_fallback_result["codex_model_provider"], "live_vendor"); assert_eq!(result["provider_name"], "Relay A"); assert_eq!(result["default_model"], "qwen3-coder"); assert_eq!( diff --git a/crates/codex-plus-data/src/lib.rs b/crates/codex-plus-data/src/lib.rs index 9f712768d..26aee77f5 100644 --- a/crates/codex-plus-data/src/lib.rs +++ b/crates/codex-plus-data/src/lib.rs @@ -9,8 +9,11 @@ pub use provider_sync::{ ProviderSyncResult, ProviderSyncStatus, ProviderSyncTargetList, ProviderSyncTargetOption, ProviderSyncTargetSource, SessionIndexCleanupApplyError, SessionIndexCleanupCandidate, SessionIndexCleanupPreview, SessionIndexCleanupResult, apply_session_index_cleanup, - load_provider_sync_targets, preview_session_index_cleanup, run_provider_sync, + load_provider_sync_targets, preview_session_index_cleanup, + remote_control_session_recovery_candidate_exists, run_provider_sync, run_provider_sync_with_target, + run_remote_control_session_catalog_recovery_for_thread_with_target, + run_remote_control_session_finalization_for_thread_with_target, }; pub use storage::{ LocalSession, SQLiteStorageAdapter, delete_local_from_paths, diff --git a/crates/codex-plus-data/src/provider_sync.rs b/crates/codex-plus-data/src/provider_sync.rs index b8bb0c839..3b1dd2f4d 100644 --- a/crates/codex-plus-data/src/provider_sync.rs +++ b/crates/codex-plus-data/src/provider_sync.rs @@ -1,15 +1,17 @@ -use rusqlite::{Connection, params_from_iter, types::Value as SqlValue}; +use rusqlite::{Connection, OptionalExtension, params_from_iter, types::Value as SqlValue}; use serde::{Deserialize, Serialize}; use serde_json::{Map, Value, json}; use sha2::{Digest, Sha256}; use std::collections::{HashMap, HashSet}; -use std::fs; +use std::fs::{self, File, OpenOptions}; +use std::io::{Read, Seek, SeekFrom, Write}; use std::path::{Path, PathBuf}; use std::time::{SystemTime, UNIX_EPOCH}; const DEFAULT_PROVIDER: &str = "openai"; const SESSION_DIRS: [&str; 2] = ["sessions", "archived_sessions"]; const BACKUP_KEEP_COUNT: usize = 5; +const REMOTE_CONTROL_CREATION_WINDOW_SECS: i64 = 15 * 60; fn default_codex_home_dir() -> PathBuf { codex_plus_core::codex_home::default_codex_home_dir() @@ -163,6 +165,13 @@ struct CatalogRepairThread { thread_source: Option, } +enum RemoteControlRolloutLookup { + Ready(PathBuf), + Archived, + UnsupportedProvider, + Missing, +} + impl SqliteUpdateCounts { fn total(&self) -> usize { self.provider_rows + self.user_event_rows + self.cwd_rows + self.catalog_insert_rows @@ -180,6 +189,297 @@ pub fn run_provider_sync(codex_home: Option<&Path>) -> ProviderSyncResult { run_provider_sync_with_target(codex_home, None) } +pub fn remote_control_session_recovery_candidate_exists( + codex_home: Option<&Path>, + thread_id: &str, +) -> anyhow::Result { + let thread_id = thread_id.trim(); + if thread_id.is_empty() || thread_id.len() > 128 { + return Ok(false); + } + let home = codex_home + .map(Path::to_path_buf) + .unwrap_or_else(default_codex_home_dir); + let minimum_created_at = now_secs() as i64 - REMOTE_CONTROL_CREATION_WINDOW_SECS; + for path in provider_sync_db_paths(&home) { + if !path.exists() { + continue; + } + let db = Connection::open(path)?; + let columns = table_columns(&db, "threads")?; + if !columns.contains("id") || !columns.contains("model_provider") { + continue; + } + let archived_expr = if columns.contains("archived") { + "COALESCE(archived, 0)" + } else { + "0" + }; + let created_expr = if columns.contains("created_at_ms") { + "CAST(COALESCE(created_at_ms, 0) / 1000 AS INTEGER)" + } else if columns.contains("created_at") { + "CAST(COALESCE(created_at, 0) AS INTEGER)" + } else { + continue; + }; + let sql = format!( + "SELECT 1 FROM threads WHERE id = ?1 AND model_provider = ?2 AND {archived_expr} = 0 AND {created_expr} >= ?3 LIMIT 1" + ); + if db + .query_row( + &sql, + (thread_id, DEFAULT_PROVIDER, minimum_created_at), + |_| Ok(()), + ) + .optional()? + .is_some() + { + return Ok(true); + } + } + Ok(false) +} + +pub fn run_remote_control_session_catalog_recovery_for_thread_with_target( + codex_home: Option<&Path>, + thread_id: &str, + target_provider: &str, +) -> ProviderSyncResult { + let thread_id = thread_id.trim(); + if thread_id.is_empty() || thread_id.len() > 128 { + return result( + ProviderSyncStatus::Skipped, + "Remote Control session recovery requires a valid thread id", + DEFAULT_PROVIDER, + None, + 0, + 0, + ); + } + let target_provider = target_provider.trim(); + if target_provider.is_empty() || target_provider == DEFAULT_PROVIDER { + return result( + ProviderSyncStatus::Skipped, + "Remote Control session recovery requires a non-openai target provider", + target_provider, + None, + 0, + 0, + ); + } + let home = codex_home + .map(Path::to_path_buf) + .unwrap_or_else(default_codex_home_dir); + let lock_dir = home.join("tmp/provider-sync.lock"); + if acquire_lock(&lock_dir).is_err() { + return result( + ProviderSyncStatus::Skipped, + format!("Provider sync lock exists: {}", lock_dir.to_string_lossy()), + target_provider, + None, + 0, + 0, + ); + } + let thread_ids = HashSet::from([thread_id.to_string()]); + let recovery = run_remote_control_catalog_recovery_for_threads( + &provider_sync_db_paths(&home), + target_provider, + &thread_ids, + ); + let _ = release_lock(&lock_dir); + recovery.unwrap_or_else(|error| { + result( + ProviderSyncStatus::Skipped, + format!("Remote Control session catalog recovery skipped: {error}"), + target_provider, + None, + 0, + 0, + ) + }) +} + +pub fn run_remote_control_session_finalization_for_thread_with_target( + codex_home: Option<&Path>, + thread_id: &str, + target_provider: &str, +) -> ProviderSyncResult { + let thread_id = thread_id.trim(); + let target_provider = target_provider.trim(); + if thread_id.is_empty() + || thread_id.len() > 128 + || target_provider.is_empty() + || target_provider == DEFAULT_PROVIDER + { + return result( + ProviderSyncStatus::Skipped, + "Remote Control session finalization requires a thread id and target provider", + target_provider, + None, + 0, + 0, + ); + } + let home = codex_home + .map(Path::to_path_buf) + .unwrap_or_else(default_codex_home_dir); + let lock_dir = home.join("tmp/provider-sync.lock"); + if acquire_lock(&lock_dir).is_err() { + return result( + ProviderSyncStatus::Skipped, + format!("Provider sync lock exists: {}", lock_dir.to_string_lossy()), + target_provider, + None, + 0, + 0, + ); + } + let recovery = (|| -> anyhow::Result { + let sqlite_paths = provider_sync_db_paths(&home); + let rollout_path = match remote_control_rollout_for_thread( + &home, + &sqlite_paths, + thread_id, + target_provider, + )? { + RemoteControlRolloutLookup::Ready(path) => path, + RemoteControlRolloutLookup::Archived => { + return Ok(result( + ProviderSyncStatus::Synced, + "Remote Control session finalization ignored an archived thread", + target_provider, + None, + 0, + 0, + )); + } + RemoteControlRolloutLookup::UnsupportedProvider => { + return Ok(result( + ProviderSyncStatus::Synced, + "Remote Control session finalization ignored a thread owned by another provider", + target_provider, + None, + 0, + 0, + )); + } + RemoteControlRolloutLookup::Missing => { + return Ok(result( + ProviderSyncStatus::Skipped, + "Remote Control session finalization deferred until the thread rollout is available", + target_provider, + None, + 0, + 0, + )); + } + }; + let collected = collect_session_change_for_path( + &rollout_path, + target_provider, + DEFAULT_PROVIDER, + thread_id, + )?; + let rewrite_changes = collected + .changes + .iter() + .filter(|change| change.rewrite_needed) + .cloned() + .collect::>(); + let backup_dir = create_backup(&home, target_provider, &rewrite_changes)?; + let applied = apply_session_changes(&rewrite_changes)?; + if !rollout_file_matches_provider(&rollout_path, thread_id, target_provider)? { + let mut deferred = result( + ProviderSyncStatus::Skipped, + "Remote Control session finalization deferred for a changed or locked rollout", + target_provider, + Some(backup_dir), + applied.changes.len(), + 0, + ); + deferred.skipped_locked_rollout_files = applied.skipped_locked_rollout_files; + return Ok(deferred); + } + let thread_ids = HashSet::from([thread_id.to_string()]); + let catalog_insert_rows = repair_missing_local_thread_catalog_rows_for_threads( + &sqlite_paths, + target_provider, + &thread_ids, + )?; + let mut sqlite_updates = apply_remote_control_recovery_sqlite_updates( + &sqlite_paths, + target_provider, + &thread_ids, + )?; + sqlite_updates.catalog_insert_rows = catalog_insert_rows; + prune_backups(&home)?; + let mut synced = result( + ProviderSyncStatus::Synced, + "Remote Control session finalization complete", + target_provider, + Some(backup_dir), + applied.changes.len(), + sqlite_updates.total(), + ); + synced.sqlite_provider_rows_updated = sqlite_updates.provider_rows; + synced.sqlite_catalog_rows_inserted = sqlite_updates.catalog_insert_rows; + Ok(synced) + })(); + let _ = release_lock(&lock_dir); + recovery.unwrap_or_else(|error| { + result( + ProviderSyncStatus::Skipped, + format!("Remote Control session finalization skipped: {error}"), + target_provider, + None, + 0, + 0, + ) + }) +} + +fn run_remote_control_catalog_recovery_for_threads( + sqlite_paths: &[PathBuf], + target_provider: &str, + requested_thread_ids: &HashSet, +) -> anyhow::Result { + let thread_ids = remote_control_catalog_recovery_thread_ids( + sqlite_paths, + target_provider, + requested_thread_ids, + )?; + if thread_ids.is_empty() { + return Ok(result( + ProviderSyncStatus::Synced, + "Remote Control session catalog already up to date", + target_provider, + None, + 0, + 0, + )); + } + + let catalog_insert_rows = repair_missing_local_thread_catalog_rows_for_threads( + sqlite_paths, + target_provider, + &thread_ids, + )?; + let provider_rows = + apply_remote_control_catalog_updates(sqlite_paths, target_provider, &thread_ids)?; + let mut synced = result( + ProviderSyncStatus::Synced, + "Remote Control session catalog recovery complete", + target_provider, + None, + 0, + provider_rows + catalog_insert_rows, + ); + synced.sqlite_provider_rows_updated = provider_rows; + synced.sqlite_catalog_rows_inserted = catalog_insert_rows; + Ok(synced) +} + pub fn run_provider_sync_with_target( codex_home: Option<&Path>, explicit_target_provider: Option<&str>, @@ -602,6 +902,191 @@ fn collect_session_changes(home: &Path, target_provider: &str) -> anyhow::Result Ok(collected) } +fn remote_control_rollout_for_thread( + home: &Path, + paths: &[PathBuf], + thread_id: &str, + target_provider: &str, +) -> anyhow::Result { + let mut archived_seen = false; + let mut unsupported_seen = false; + let mut candidate_seen = false; + + for path in paths { + if !path.exists() { + continue; + } + let db = Connection::open(path)?; + let columns = table_columns(&db, "threads")?; + if !columns.contains("id") { + continue; + } + let provider_expr = if columns.contains("model_provider") { + "COALESCE(model_provider, '')" + } else { + "''" + }; + let archived_expr = if columns.contains("archived") { + "COALESCE(archived, 0)" + } else { + "0" + }; + let rollout_expr = if columns.contains("rollout_path") { + "COALESCE(rollout_path, '')" + } else { + "''" + }; + let sql = format!( + "SELECT {provider_expr}, {archived_expr}, {rollout_expr} FROM threads WHERE id = ?1" + ); + let mut stmt = db.prepare(&sql)?; + let rows = stmt.query_map([thread_id], |row| { + Ok(( + row.get::<_, String>(0)?, + row.get::<_, i64>(1)?, + row.get::<_, String>(2)?, + )) + })?; + for row in rows { + let (provider, archived, rollout_path) = row?; + candidate_seen = true; + if archived != 0 { + archived_seen = true; + continue; + } + if !provider.is_empty() && provider != DEFAULT_PROVIDER && provider != target_provider { + unsupported_seen = true; + continue; + } + let Some(rollout_path) = resolve_active_rollout_path(home, &rollout_path) else { + continue; + }; + let Some((rollout_thread_id, providers)) = + rollout_provider_state_for_path(&rollout_path)? + else { + continue; + }; + if rollout_thread_id != thread_id { + continue; + } + if providers.is_empty() + || providers + .iter() + .any(|provider| provider != DEFAULT_PROVIDER && provider != target_provider) + { + unsupported_seen = true; + continue; + } + return Ok(RemoteControlRolloutLookup::Ready(rollout_path)); + } + } + + if archived_seen && !unsupported_seen { + Ok(RemoteControlRolloutLookup::Archived) + } else if unsupported_seen { + Ok(RemoteControlRolloutLookup::UnsupportedProvider) + } else if candidate_seen { + Ok(RemoteControlRolloutLookup::Missing) + } else { + Ok(RemoteControlRolloutLookup::Missing) + } +} + +fn resolve_active_rollout_path(home: &Path, value: &str) -> Option { + let raw = value.trim(); + if raw.is_empty() { + return None; + } + let path = PathBuf::from(raw); + let path = if path.is_absolute() { + path + } else { + home.join(path) + }; + let canonical = fs::canonicalize(path).ok()?; + let sessions_root = fs::canonicalize(home.join("sessions")).ok()?; + if !canonical.starts_with(sessions_root) { + return None; + } + Some(canonical) +} + +fn rollout_provider_state_for_path( + path: &Path, +) -> anyhow::Result)>> { + let text = match fs::read_to_string(path) { + Ok(text) => text, + Err(error) if is_locked_io_error(&error) => return Ok(None), + Err(error) => return Err(error.into()), + }; + Ok(rollout_thread_provider_state(&text)) +} + +fn collect_session_change_for_path( + path: &Path, + target_provider: &str, + source_provider: &str, + thread_id: &str, +) -> anyhow::Result { + let mut collected = SessionChanges::default(); + let text = match fs::read_to_string(path) { + Ok(text) => text, + Err(error) if is_locked_io_error(&error) => { + collected + .skipped_locked_rollout_files + .push(path.to_path_buf()); + return Ok(collected); + } + Err(error) => return Err(error.into()), + }; + let rewrite = rewrite_rollout_session_meta_providers_for_threads( + &text, + target_provider, + source_provider, + &HashSet::from([thread_id.to_string()]), + )?; + if rewrite.session_meta_count == 0 || rewrite.thread_id.as_deref() != Some(thread_id) { + return Ok(collected); + } + let has_user_event = text.contains("\"user_message\"") || text.contains("\"user_input\""); + if text.contains("encrypted_content") { + for provider in &rewrite.providers { + *collected + .encrypted_content_counts + .entry(provider.clone()) + .or_insert(0) += 1; + } + } + let original_mtime = fs::metadata(path) + .and_then(|metadata| metadata.modified()) + .ok(); + collected.changes.push(SessionChange { + path: path.to_path_buf(), + original_text: text, + next_text: rewrite.next_text, + original_session_meta_lines: rewrite.original_session_meta_lines, + thread_id: rewrite.thread_id, + cwd: rewrite.cwd, + has_user_event, + rewrite_needed: rewrite.rewrite_needed, + original_mtime, + }); + Ok(collected) +} + +fn rollout_file_matches_provider( + path: &Path, + thread_id: &str, + target_provider: &str, +) -> anyhow::Result { + let Some((rollout_thread_id, providers)) = rollout_provider_state_for_path(path)? else { + return Ok(false); + }; + Ok(rollout_thread_id == thread_id + && !providers.is_empty() + && providers.iter().all(|provider| provider == target_provider)) +} + fn rewrite_rollout_session_meta_providers( text: &str, target_provider: &str, @@ -655,6 +1140,81 @@ fn rewrite_rollout_session_meta_providers( Ok(rewrite) } +fn rewrite_rollout_session_meta_providers_for_threads( + text: &str, + target_provider: &str, + source_provider: &str, + thread_ids: &HashSet, +) -> anyhow::Result { + let rollout_thread_id = text.lines().find_map(|line| { + let record = serde_json::from_str::(line).ok()?; + if record.get("type").and_then(Value::as_str) != Some("session_meta") { + return None; + } + record + .get("payload")? + .get("id")? + .as_str() + .map(ToString::to_string) + }); + if rollout_thread_id + .as_ref() + .is_none_or(|thread_id| !thread_ids.contains(thread_id)) + { + return Ok(RolloutRewrite { + next_text: text.to_string(), + ..RolloutRewrite::default() + }); + } + + let mut rewrite = RolloutRewrite { + thread_id: rollout_thread_id, + ..RolloutRewrite::default() + }; + for segment in text.split_inclusive('\n') { + let (line, line_ending) = split_line_ending(segment); + let mut next_line = line.to_string(); + if !line.trim().is_empty() { + if let Ok(mut record) = serde_json::from_str::(line) { + if record.get("type").and_then(Value::as_str) == Some("session_meta") { + let Some(payload) = record.get_mut("payload").and_then(Value::as_object_mut) + else { + rewrite.next_text.push_str(&next_line); + rewrite.next_text.push_str(line_ending); + continue; + }; + rewrite.session_meta_count += 1; + rewrite.original_session_meta_lines.push(line.to_string()); + if rewrite.cwd.is_none() { + rewrite.cwd = payload + .get("cwd") + .and_then(Value::as_str) + .and_then(to_desktop_workspace_path); + } + let provider = payload + .get("model_provider") + .and_then(Value::as_str) + .map(ToString::to_string); + rewrite + .providers + .push(provider.clone().unwrap_or_else(|| "(missing)".to_string())); + if provider + .as_deref() + .is_none_or(|provider| provider == source_provider) + { + payload.insert("model_provider".to_string(), json!(target_provider)); + next_line = serde_json::to_string(&record)?; + rewrite.rewrite_needed = true; + } + } + } + } + rewrite.next_text.push_str(&next_line); + rewrite.next_text.push_str(line_ending); + } + Ok(rewrite) +} + fn rollout_files(home: &Path) -> anyhow::Result> { let mut files = Vec::new(); for dirname in SESSION_DIRS { @@ -1049,8 +1609,10 @@ fn to_desktop_workspace_path(value: &str) -> Option { } fn is_locked_io_error(error: &std::io::Error) -> bool { - matches!(error.kind(), std::io::ErrorKind::PermissionDenied) - || matches!(error.raw_os_error(), Some(32 | 33)) + matches!( + error.kind(), + std::io::ErrorKind::PermissionDenied | std::io::ErrorKind::WouldBlock + ) || matches!(error.raw_os_error(), Some(32 | 33)) } fn build_encrypted_content_warning( @@ -1173,8 +1735,18 @@ fn create_session_index_cleanup_backup( fn apply_session_changes(changes: &[SessionChange]) -> anyhow::Result { let mut applied = AppliedSessionChanges::default(); for change in changes { - match fs::write(&change.path, &change.next_text) { - Ok(()) => {} + match replace_session_text_if_unchanged( + &change.path, + &change.original_text, + &change.next_text, + ) { + Ok(true) => {} + Ok(false) => { + applied + .skipped_locked_rollout_files + .push(change.path.clone()); + continue; + } Err(error) if is_locked_io_error(&error) => { applied .skipped_locked_rollout_files @@ -1191,12 +1763,57 @@ fn apply_session_changes(changes: &[SessionChange]) -> anyhow::Result anyhow::Result<()> { for change in changes { - fs::write(&change.path, &change.original_text)?; - restore_file_mtime(&change.path, change.original_mtime); + if replace_session_text_if_unchanged( + &change.path, + &change.next_text, + &change.original_text, + )? { + restore_file_mtime(&change.path, change.original_mtime); + } } Ok(()) } +fn replace_session_text_if_unchanged( + path: &Path, + expected_text: &str, + next_text: &str, +) -> std::io::Result { + let mut file = open_session_file_for_update(path)?; + file.try_lock()?; + let mut current_text = String::new(); + file.read_to_string(&mut current_text)?; + if current_text != expected_text { + return Ok(false); + } + + file.seek(SeekFrom::Start(0))?; + file.set_len(0)?; + file.write_all(next_text.as_bytes())?; + file.flush()?; + + file.seek(SeekFrom::Start(0))?; + let mut persisted_text = String::new(); + file.read_to_string(&mut persisted_text)?; + if persisted_text != next_text { + return Err(std::io::Error::other( + "rollout changed while provider metadata was being written", + )); + } + Ok(true) +} + +fn open_session_file_for_update(path: &Path) -> std::io::Result { + let mut options = OpenOptions::new(); + options.read(true).write(true); + #[cfg(windows)] + { + use std::os::windows::fs::OpenOptionsExt; + options.share_mode(0); + } + options.open(path) +} + fn restore_file_mtime(path: &Path, mtime: Option) { let Some(mtime) = mtime else { return }; let Ok(file) = fs::File::options().write(true).open(path) else { @@ -1240,6 +1857,112 @@ fn sqlite_provider_ids(path: &Path) -> anyhow::Result> { Ok(sorted_provider_ids(ids)) } +fn remote_control_catalog_recovery_thread_ids( + paths: &[PathBuf], + target_provider: &str, + requested_thread_ids: &HashSet, +) -> anyhow::Result> { + let mut known_thread_ids = HashSet::new(); + let mut ready_thread_ids = HashSet::new(); + let mut has_local_catalog = false; + for path in paths { + if !path.exists() { + continue; + } + let db = Connection::open(path)?; + let thread_columns = table_columns(&db, "threads")?; + if thread_columns.contains("id") { + let mut stmt = db.prepare("SELECT id FROM threads WHERE COALESCE(id, '') <> ''")?; + for item in stmt.query_map([], |row| row.get::<_, String>(0))? { + let thread_id = item?; + if requested_thread_ids.contains(&thread_id) { + known_thread_ids.insert(thread_id); + } + } + } + + let catalog_columns = table_columns(&db, "local_thread_catalog")?; + if !catalog_columns.contains("thread_id") { + continue; + } + let Some(host_id) = local_catalog_host_id(&db)? else { + continue; + }; + has_local_catalog = true; + let provider_expr = if catalog_columns.contains("model_provider") { + "COALESCE(model_provider, '')" + } else { + "''" + }; + let missing_expr = if catalog_columns.contains("missing_candidate") { + "COALESCE(missing_candidate, 0)" + } else { + "0" + }; + let host_filter = if catalog_columns.contains("host_id") { + " AND host_id = ?1" + } else { + " AND ?1 = ?1" + }; + let sql = format!( + "SELECT thread_id, {provider_expr}, {missing_expr} FROM local_thread_catalog WHERE COALESCE(thread_id, '') <> ''{host_filter}" + ); + let mut stmt = db.prepare(&sql)?; + for item in stmt.query_map([host_id], |row| { + Ok(( + row.get::<_, String>(0)?, + row.get::<_, String>(1)?, + row.get::<_, i64>(2)?, + )) + })? { + let (thread_id, provider, missing_candidate) = item?; + if requested_thread_ids.contains(&thread_id) + && provider == target_provider + && missing_candidate == 0 + { + ready_thread_ids.insert(thread_id); + } + } + } + if !has_local_catalog { + return Ok(HashSet::new()); + } + known_thread_ids.retain(|thread_id| !ready_thread_ids.contains(thread_id)); + Ok(known_thread_ids) +} + +fn rollout_thread_provider_state(text: &str) -> Option<(String, HashSet)> { + let mut thread_id = None; + let mut providers = HashSet::new(); + for segment in text.split_inclusive('\n') { + let (line, _) = split_line_ending(segment); + let Ok(record) = serde_json::from_str::(line) else { + continue; + }; + if record.get("type").and_then(Value::as_str) != Some("session_meta") { + continue; + } + let Some(payload) = record.get("payload").and_then(Value::as_object) else { + continue; + }; + if thread_id.is_none() { + thread_id = payload + .get("id") + .and_then(Value::as_str) + .filter(|id| !id.trim().is_empty()) + .map(ToString::to_string); + } + providers.insert( + payload + .get("model_provider") + .and_then(Value::as_str) + .unwrap_or("(missing)") + .to_string(), + ); + } + thread_id.map(|thread_id| (thread_id, providers)) +} + fn count_sqlite_updates( path: &Path, target_provider: &str, @@ -1373,11 +2096,126 @@ fn apply_sqlite_update_for_paths( Ok(total) } +fn apply_remote_control_recovery_sqlite_updates( + paths: &[PathBuf], + target_provider: &str, + thread_ids: &HashSet, +) -> anyhow::Result { + let mut counts = SqliteUpdateCounts::default(); + for path in paths { + if !path.exists() { + continue; + } + let mut db = Connection::open(path)?; + let thread_columns = table_columns(&db, "threads")?; + let catalog_columns = table_columns(&db, "local_thread_catalog")?; + let local_host_id = if catalog_columns.contains("thread_id") { + local_catalog_host_id(&db)? + } else { + None + }; + let tx = db.transaction()?; + if thread_columns.contains("id") && thread_columns.contains("model_provider") { + for thread_id in thread_ids { + counts.provider_rows += tx.execute( + "UPDATE threads SET model_provider = ?1 WHERE id = ?2 AND model_provider = ?3", + (target_provider, thread_id, DEFAULT_PROVIDER), + )?; + } + } + if catalog_columns.contains("thread_id") + && catalog_columns.contains("model_provider") + && local_host_id.is_some() + { + let host_id = local_host_id.as_deref().unwrap_or("local"); + let host_filter = if catalog_columns.contains("host_id") { + " AND host_id = ?3" + } else { + " AND ?3 = ?3" + }; + for thread_id in thread_ids { + let sql = format!( + "UPDATE local_thread_catalog SET model_provider = ?1 WHERE thread_id = ?2{host_filter} AND model_provider = ?4" + ); + counts.provider_rows += tx.execute( + &sql, + (target_provider, thread_id, host_id, DEFAULT_PROVIDER), + )?; + if catalog_columns.contains("missing_candidate") { + let sql = format!( + "UPDATE local_thread_catalog SET missing_candidate = 0 WHERE thread_id = ?1{} AND COALESCE(missing_candidate, 0) <> 0", + if catalog_columns.contains("host_id") { + " AND host_id = ?2" + } else { + " AND ?2 = ?2" + } + ); + tx.execute(&sql, (thread_id, host_id))?; + } + } + } + tx.commit()?; + } + Ok(counts) +} + +fn apply_remote_control_catalog_updates( + paths: &[PathBuf], + target_provider: &str, + thread_ids: &HashSet, +) -> anyhow::Result { + let mut total = 0; + for path in paths { + if !path.exists() { + continue; + } + let mut db = Connection::open(path)?; + let columns = table_columns(&db, "local_thread_catalog")?; + if !columns.contains("thread_id") || !columns.contains("model_provider") { + continue; + } + let Some(host_id) = local_catalog_host_id(&db)? else { + continue; + }; + let host_filter = if columns.contains("host_id") { + " AND host_id = ?3" + } else { + " AND ?3 = ?3" + }; + let tx = db.transaction()?; + for thread_id in thread_ids { + let sql = format!( + "UPDATE local_thread_catalog SET model_provider = ?1{} WHERE thread_id = ?2{} AND COALESCE(model_provider, '') <> ?1", + if columns.contains("missing_candidate") { + ", missing_candidate = 0" + } else { + "" + }, + host_filter + ); + total += tx.execute(&sql, (target_provider, thread_id, &host_id))?; + if columns.contains("missing_candidate") { + let sql = format!( + "UPDATE local_thread_catalog SET missing_candidate = 0 WHERE thread_id = ?1{} AND COALESCE(missing_candidate, 0) <> 0", + if columns.contains("host_id") { + " AND host_id = ?2" + } else { + " AND ?2 = ?2" + } + ); + tx.execute(&sql, (thread_id, &host_id))?; + } + } + tx.commit()?; + } + Ok(total) +} + fn count_missing_local_thread_catalog_rows( paths: &[PathBuf], target_provider: &str, ) -> anyhow::Result { - let source_threads = collect_catalog_repair_threads(paths, target_provider)?; + let source_threads = collect_catalog_repair_threads(paths, target_provider, None)?; if source_threads.is_empty() { return Ok(0); } @@ -1391,7 +2229,9 @@ fn count_missing_local_thread_catalog_rows( if !catalog_supports_repair(&columns) { continue; } - let host_id = local_catalog_host_id(&db)?; + let Some(host_id) = local_catalog_host_id(&db)? else { + continue; + }; for thread in source_threads.values() { if !local_catalog_contains_thread(&db, &host_id, &thread.id)? { total += 1; @@ -1405,7 +2245,29 @@ fn repair_missing_local_thread_catalog_rows( paths: &[PathBuf], target_provider: &str, ) -> anyhow::Result { - let source_threads = collect_catalog_repair_threads(paths, target_provider)?; + repair_missing_local_thread_catalog_rows_filtered(paths, target_provider, None, true) +} + +fn repair_missing_local_thread_catalog_rows_for_threads( + paths: &[PathBuf], + target_provider: &str, + thread_ids: &HashSet, +) -> anyhow::Result { + repair_missing_local_thread_catalog_rows_filtered( + paths, + target_provider, + Some(thread_ids), + false, + ) +} + +fn repair_missing_local_thread_catalog_rows_filtered( + paths: &[PathBuf], + target_provider: &str, + thread_ids: Option<&HashSet>, + update_full_sync_state: bool, +) -> anyhow::Result { + let source_threads = collect_catalog_repair_threads(paths, target_provider, thread_ids)?; if source_threads.is_empty() { return Ok(0); } @@ -1421,8 +2283,10 @@ fn repair_missing_local_thread_catalog_rows( } let sync_columns = table_columns(&db, "local_thread_catalog_sync_state")?; let metadata_columns = table_columns(&db, "local_thread_catalog_metadata")?; - let host_id = local_catalog_host_id(&db)?; - let mut observation_sequence = local_catalog_max_observation_sequence(&db)?; + let Some(host_id) = local_catalog_host_id(&db)? else { + continue; + }; + let mut observation_sequence = local_catalog_max_observation_sequence(&db, &host_id)?; let insert_columns = local_catalog_insert_columns(&columns); let placeholders = std::iter::repeat_n("?", insert_columns.len()) .collect::>() @@ -1453,13 +2317,15 @@ fn repair_missing_local_thread_catalog_rows( } if inserted > 0 { update_local_catalog_metadata(&tx, &metadata_columns, inserted)?; - update_local_catalog_sync_state( - &tx, - &sync_columns, - &host_id, - observation_sequence, - max_source_updated_at, - )?; + if update_full_sync_state { + update_local_catalog_sync_state( + &tx, + &sync_columns, + &host_id, + observation_sequence, + max_source_updated_at, + )?; + } } tx.commit()?; total_inserted += inserted; @@ -1470,6 +2336,7 @@ fn repair_missing_local_thread_catalog_rows( fn collect_catalog_repair_threads( paths: &[PathBuf], target_provider: &str, + thread_ids: Option<&HashSet>, ) -> anyhow::Result> { let mut threads = HashMap::new(); for path in paths { @@ -1515,6 +2382,9 @@ fn collect_catalog_repair_threads( })?; for item in rows { let thread = item?; + if thread_ids.is_some_and(|thread_ids| !thread_ids.contains(&thread.id)) { + continue; + } let replace = threads .get(&thread.id) .map(|current: &CatalogRepairThread| { @@ -1545,30 +2415,41 @@ fn catalog_supports_repair(columns: &HashSet) -> bool { .all(|column| columns.contains(*column)) } -fn local_catalog_host_id(db: &Connection) -> anyhow::Result { - if !table_columns(db, "local_thread_catalog_hosts")?.contains("host_id") { - return Ok("local".to_string()); +fn local_catalog_host_id(db: &Connection) -> anyhow::Result> { + let columns = table_columns(db, "local_thread_catalog_hosts")?; + if !columns.contains("host_id") { + return Ok(Some("local".to_string())); } - match db.query_row( - "SELECT host_id FROM local_thread_catalog_hosts ORDER BY host_id LIMIT 1", - [], - |row| row.get::<_, String>(0), - ) { - Ok(host_id) if !host_id.trim().is_empty() => Ok(host_id), - Ok(_) | Err(rusqlite::Error::QueryReturnedNoRows) => Ok("local".to_string()), + let query = if columns.contains("host_kind") { + "SELECT host_id FROM local_thread_catalog_hosts WHERE LOWER(COALESCE(host_kind, '')) = 'local' ORDER BY host_id LIMIT 1" + } else { + "SELECT host_id FROM local_thread_catalog_hosts WHERE host_id = 'local' LIMIT 1" + }; + match db.query_row(query, [], |row| row.get::<_, String>(0)) { + Ok(host_id) if !host_id.trim().is_empty() => Ok(Some(host_id)), + Ok(_) | Err(rusqlite::Error::QueryReturnedNoRows) => Ok(None), Err(error) => Err(error.into()), } } -fn local_catalog_max_observation_sequence(db: &Connection) -> anyhow::Result { - if !table_columns(db, "local_thread_catalog")?.contains("observation_sequence") { +fn local_catalog_max_observation_sequence(db: &Connection, host_id: &str) -> anyhow::Result { + let columns = table_columns(db, "local_thread_catalog")?; + if !columns.contains("observation_sequence") { return Ok(0); } - Ok(db.query_row( - "SELECT COALESCE(MAX(observation_sequence), 0) FROM local_thread_catalog", - [], - |row| row.get::<_, i64>(0), - )?) + if columns.contains("host_id") { + Ok(db.query_row( + "SELECT COALESCE(MAX(observation_sequence), 0) FROM local_thread_catalog WHERE host_id = ?1", + [host_id], + |row| row.get::<_, i64>(0), + )?) + } else { + Ok(db.query_row( + "SELECT COALESCE(MAX(observation_sequence), 0) FROM local_thread_catalog", + [], + |row| row.get::<_, i64>(0), + )?) + } } fn local_catalog_contains_thread( diff --git a/crates/codex-plus-data/tests/provider_sync.rs b/crates/codex-plus-data/tests/provider_sync.rs index f53185bda..b4638ad67 100644 --- a/crates/codex-plus-data/tests/provider_sync.rs +++ b/crates/codex-plus-data/tests/provider_sync.rs @@ -1,7 +1,10 @@ use codex_plus_data::{ ProviderSyncStatus, ProviderSyncTargetSource, apply_session_index_cleanup, - load_provider_sync_targets, preview_session_index_cleanup, run_provider_sync, + load_provider_sync_targets, preview_session_index_cleanup, + remote_control_session_recovery_candidate_exists, run_provider_sync, run_provider_sync_with_target, + run_remote_control_session_catalog_recovery_for_thread_with_target, + run_remote_control_session_finalization_for_thread_with_target, }; use rusqlite::Connection; use serde_json::json; @@ -9,7 +12,7 @@ use std::ffi::OsString; use std::fs; use std::path::Path; use std::sync::Mutex; -use std::time::{Duration, SystemTime}; +use std::time::{Duration, Instant, SystemTime, UNIX_EPOCH}; use tempfile::tempdir; static CODEX_HOME_ENV_LOCK: Mutex<()> = Mutex::new(()); @@ -113,6 +116,31 @@ fn create_state_db_with_providers(path: &Path, rows: &[(&str, &str, i64)]) { } } +fn create_remote_control_state_db(path: &Path, rows: &[(&str, &str, i64, &Path)]) { + let db = Connection::open(path).unwrap(); + db.execute( + "CREATE TABLE threads ( + id TEXT PRIMARY KEY, model_provider TEXT, archived INTEGER, has_user_event INTEGER, + cwd TEXT, title TEXT, rollout_path TEXT, source TEXT, created_at_ms INTEGER, + updated_at_ms INTEGER, thread_source TEXT, git_branch TEXT + )", + [], + ) + .unwrap(); + for (id, provider, archived, rollout_path) in rows { + db.execute( + "INSERT INTO threads VALUES (?1, ?2, ?3, 1, 'C:/workspace', ?1, ?4, 'vscode', 100000, 200000, NULL, NULL)", + ( + id, + provider, + archived, + rollout_path.to_string_lossy().to_string(), + ), + ) + .unwrap(); + } +} + fn create_local_thread_catalog_db(path: &Path, rows: &[(&str, &str)]) { let db = Connection::open(path).unwrap(); db.execute( @@ -184,6 +212,46 @@ fn create_local_thread_catalog_db(path: &Path, rows: &[(&str, &str)]) { } } +#[test] +fn remote_control_recovery_candidate_requires_a_recent_unarchived_openai_thread() { + let tmp = tempdir().unwrap(); + let home = tmp.path().join(".codex"); + fs::create_dir_all(&home).unwrap(); + let now_ms = SystemTime::now() + .duration_since(UNIX_EPOCH) + .unwrap() + .as_millis() as i64; + let db = Connection::open(home.join("state_5.sqlite")).unwrap(); + db.execute( + "CREATE TABLE threads ( + id TEXT PRIMARY KEY, + model_provider TEXT, + archived INTEGER, + created_at_ms INTEGER + )", + [], + ) + .unwrap(); + for (id, provider, archived, created_at_ms) in [ + ("recent", "openai", 0, now_ms), + ("stale", "openai", 0, now_ms - 16 * 60 * 1000), + ("archived", "openai", 1, now_ms), + ("custom", "custom", 0, now_ms), + ] { + db.execute( + "INSERT INTO threads VALUES (?1, ?2, ?3, ?4)", + (id, provider, archived, created_at_ms), + ) + .unwrap(); + } + drop(db); + + assert!(remote_control_session_recovery_candidate_exists(Some(&home), "recent").unwrap()); + assert!(!remote_control_session_recovery_candidate_exists(Some(&home), "stale").unwrap()); + assert!(!remote_control_session_recovery_candidate_exists(Some(&home), "archived").unwrap()); + assert!(!remote_control_session_recovery_candidate_exists(Some(&home), "custom").unwrap()); +} + #[test] fn provider_sync_targets_default_to_codex_home_env() { let _lock = CODEX_HOME_ENV_LOCK.lock().unwrap(); @@ -597,6 +665,530 @@ fn provider_sync_repairs_missing_local_thread_catalog_rows_from_threads() { assert_eq!(sync_state.2, 1); } +#[test] +fn remote_control_catalog_recovery_for_thread_does_not_touch_other_candidates() { + let tmp = tempdir().unwrap(); + let home = tmp.path().join(".codex"); + let sqlite_dir = home.join("sqlite"); + fs::create_dir_all(&sqlite_dir).unwrap(); + fs::write(home.join("config.toml"), "model_provider = \"custom\"\n").unwrap(); + + let state_db = home.join("state_5.sqlite"); + let db = Connection::open(&state_db).unwrap(); + db.execute( + "CREATE TABLE threads ( + id TEXT PRIMARY KEY, model_provider TEXT, archived INTEGER, has_user_event INTEGER, + cwd TEXT, title TEXT, rollout_path TEXT, source TEXT, created_at_ms INTEGER, + updated_at_ms INTEGER, thread_source TEXT, git_branch TEXT + )", + [], + ) + .unwrap(); + for id in ["mobile-one", "mobile-two"] { + let rollout = home.join(format!("sessions/rollout-{id}.jsonl")); + write_rollout(&rollout, "openai", id, "C:/workspace"); + db.execute( + "INSERT INTO threads VALUES (?1, 'openai', 0, 1, 'C:/workspace', ?1, ?2, 'vscode', 100000, 200000, NULL, NULL)", + (id, rollout.to_string_lossy().to_string()), + ) + .unwrap(); + } + drop(db); + let catalog_db = sqlite_dir.join("codex-dev.db"); + create_local_thread_catalog_db(&catalog_db, &[]); + + let result = run_remote_control_session_catalog_recovery_for_thread_with_target( + Some(&home), + "mobile-one", + "custom", + ); + + assert_eq!(result.status, ProviderSyncStatus::Synced); + assert_eq!(result.changed_session_files, 0); + assert_eq!(result.sqlite_catalog_rows_inserted, 1); + let db = Connection::open(&state_db).unwrap(); + let providers = ["mobile-one", "mobile-two"] + .into_iter() + .map(|id| { + db.query_row( + "SELECT model_provider FROM threads WHERE id = ?1", + [id], + |row| row.get::<_, String>(0), + ) + .unwrap() + }) + .collect::>(); + assert_eq!(providers, vec!["openai", "openai"]); + let catalog = Connection::open(&catalog_db).unwrap(); + assert_eq!( + catalog + .query_row( + "SELECT COUNT(*) FROM local_thread_catalog WHERE thread_id = 'mobile-one' AND model_provider = 'custom'", + [], + |row| row.get::<_, i64>(0), + ) + .unwrap(), + 1 + ); + assert_eq!( + catalog + .query_row( + "SELECT COUNT(*) FROM local_thread_catalog WHERE thread_id = 'mobile-two'", + [], + |row| row.get::<_, i64>(0), + ) + .unwrap(), + 0 + ); + for id in ["mobile-one", "mobile-two"] { + let rollout = home.join(format!("sessions/rollout-{id}.jsonl")); + let first: serde_json::Value = + serde_json::from_str(fs::read_to_string(rollout).unwrap().lines().next().unwrap()) + .unwrap(); + assert_eq!(first["payload"]["model_provider"], "openai"); + } +} + +#[test] +fn remote_control_catalog_recovery_for_thread_only_repairs_the_local_catalog_host() { + let tmp = tempdir().unwrap(); + let home = tmp.path().join(".codex"); + let sqlite_dir = home.join("sqlite"); + fs::create_dir_all(&sqlite_dir).unwrap(); + fs::write(home.join("config.toml"), "model_provider = \"custom\"\n").unwrap(); + let rollout = home.join("sessions/rollout-mobile.jsonl"); + write_rollout(&rollout, "openai", "mobile", "C:/workspace"); + create_state_db_with_providers(&home.join("state_5.sqlite"), &[("mobile", "openai", 0)]); + + let catalog_db = sqlite_dir.join("codex-dev.db"); + create_local_thread_catalog_db(&catalog_db, &[]); + let before_sync_state = Connection::open(&catalog_db) + .unwrap() + .query_row( + "SELECT watermark_updated_at, initial_build_complete, observation_sequence, last_full_reconciled_at FROM local_thread_catalog_sync_state WHERE host_id = 'local'", + [], + |row| { + Ok(( + row.get::<_, f64>(0)?, + row.get::<_, i64>(1)?, + row.get::<_, i64>(2)?, + row.get::<_, i64>(3)?, + )) + }, + ) + .unwrap(); + let db = Connection::open(&catalog_db).unwrap(); + db.execute( + "INSERT INTO local_thread_catalog_hosts VALUES ('aaa-remote', 'ssh')", + [], + ) + .unwrap(); + db.execute( + "INSERT INTO local_thread_catalog ( + host_id, thread_id, display_title, source_created_at, source_updated_at, cwd, + source_kind, source_detail, model_provider, git_branch, observation_sequence, + missing_candidate, thread_source + ) VALUES ('aaa-remote', 'mobile', 'Remote copy', 100, 100, '/remote', 'cli', '', 'openai', NULL, 1, 0, 'user')", + [], + ) + .unwrap(); + drop(db); + + let result = run_remote_control_session_catalog_recovery_for_thread_with_target( + Some(&home), + "mobile", + "custom", + ); + + assert_eq!(result.status, ProviderSyncStatus::Synced); + assert_eq!(result.changed_session_files, 0); + assert_eq!(result.sqlite_catalog_rows_inserted, 1); + let db = Connection::open(&catalog_db).unwrap(); + let rows = db + .prepare( + "SELECT host_id, model_provider FROM local_thread_catalog WHERE thread_id = 'mobile' ORDER BY host_id", + ) + .unwrap() + .query_map([], |row| { + Ok((row.get::<_, String>(0)?, row.get::<_, String>(1)?)) + }) + .unwrap() + .collect::>>() + .unwrap(); + assert_eq!( + rows, + vec![ + ("aaa-remote".to_string(), "openai".to_string()), + ("local".to_string(), "custom".to_string()), + ] + ); + let state = Connection::open(home.join("state_5.sqlite")).unwrap(); + let provider: String = state + .query_row( + "SELECT model_provider FROM threads WHERE id = 'mobile'", + [], + |row| row.get(0), + ) + .unwrap(); + assert_eq!(provider, "openai"); + let after_sync_state = Connection::open(&catalog_db) + .unwrap() + .query_row( + "SELECT watermark_updated_at, initial_build_complete, observation_sequence, last_full_reconciled_at FROM local_thread_catalog_sync_state WHERE host_id = 'local'", + [], + |row| { + Ok(( + row.get::<_, f64>(0)?, + row.get::<_, i64>(1)?, + row.get::<_, i64>(2)?, + row.get::<_, i64>(3)?, + )) + }, + ) + .unwrap(); + assert_eq!(after_sync_state, before_sync_state); +} + +#[test] +fn remote_control_finalization_uses_only_recorded_rollout_and_preserves_full_sync_state() { + let tmp = tempdir().unwrap(); + let home = tmp.path().join(".codex"); + let sqlite_dir = home.join("sqlite"); + fs::create_dir_all(&sqlite_dir).unwrap(); + let target_rollout = home.join("sessions/rollout-mobile.jsonl"); + let other_rollout = home.join("sessions/rollout-other.jsonl"); + write_rollout(&target_rollout, "openai", "mobile", "C:/workspace"); + write_rollout(&other_rollout, "openai", "other", "C:/workspace"); + create_remote_control_state_db( + &home.join("state_5.sqlite"), + &[ + ("mobile", "openai", 0, &target_rollout), + ("other", "openai", 0, &other_rollout), + ], + ); + let catalog_db = sqlite_dir.join("codex-dev.db"); + create_local_thread_catalog_db(&catalog_db, &[]); + + let before_sync_state = Connection::open(&catalog_db) + .unwrap() + .query_row( + "SELECT watermark_updated_at, initial_build_complete, observation_sequence, last_full_reconciled_at FROM local_thread_catalog_sync_state WHERE host_id = 'local'", + [], + |row| { + Ok(( + row.get::<_, f64>(0)?, + row.get::<_, i64>(1)?, + row.get::<_, i64>(2)?, + row.get::<_, i64>(3)?, + )) + }, + ) + .unwrap(); + + let result = run_remote_control_session_finalization_for_thread_with_target( + Some(&home), + "mobile", + "custom", + ); + + assert_eq!(result.status, ProviderSyncStatus::Synced); + assert_eq!(result.changed_session_files, 1); + assert_eq!(result.sqlite_catalog_rows_inserted, 1); + let target_first: serde_json::Value = serde_json::from_str( + fs::read_to_string(&target_rollout) + .unwrap() + .lines() + .next() + .unwrap(), + ) + .unwrap(); + let other_first: serde_json::Value = serde_json::from_str( + fs::read_to_string(&other_rollout) + .unwrap() + .lines() + .next() + .unwrap(), + ) + .unwrap(); + assert_eq!(target_first["payload"]["model_provider"], "custom"); + assert_eq!(other_first["payload"]["model_provider"], "openai"); + + let state = Connection::open(home.join("state_5.sqlite")).unwrap(); + let providers = ["mobile", "other"] + .into_iter() + .map(|id| { + state + .query_row( + "SELECT model_provider FROM threads WHERE id = ?1", + [id], + |row| row.get::<_, String>(0), + ) + .unwrap() + }) + .collect::>(); + assert_eq!(providers, vec!["custom", "openai"]); + + let catalog = Connection::open(&catalog_db).unwrap(); + assert_eq!( + catalog + .query_row( + "SELECT COUNT(*) FROM local_thread_catalog WHERE thread_id = 'mobile' AND model_provider = 'custom'", + [], + |row| row.get::<_, i64>(0), + ) + .unwrap(), + 1 + ); + assert_eq!( + catalog + .query_row( + "SELECT COUNT(*) FROM local_thread_catalog WHERE thread_id = 'other'", + [], + |row| row.get::<_, i64>(0), + ) + .unwrap(), + 0 + ); + let after_sync_state = catalog + .query_row( + "SELECT watermark_updated_at, initial_build_complete, observation_sequence, last_full_reconciled_at FROM local_thread_catalog_sync_state WHERE host_id = 'local'", + [], + |row| { + Ok(( + row.get::<_, f64>(0)?, + row.get::<_, i64>(1)?, + row.get::<_, i64>(2)?, + row.get::<_, i64>(3)?, + )) + }, + ) + .unwrap(); + assert_eq!(after_sync_state, before_sync_state); +} + +#[test] +fn remote_control_finalization_ignores_archived_and_other_provider_threads() { + let tmp = tempdir().unwrap(); + let home = tmp.path().join(".codex"); + let sqlite_dir = home.join("sqlite"); + fs::create_dir_all(&sqlite_dir).unwrap(); + let archived_rollout = home.join("sessions/rollout-archived.jsonl"); + let other_rollout = home.join("sessions/rollout-other.jsonl"); + write_rollout(&archived_rollout, "openai", "archived", "C:/workspace"); + write_rollout(&other_rollout, "other", "other", "C:/workspace"); + create_remote_control_state_db( + &home.join("state_5.sqlite"), + &[ + ("archived", "openai", 1, &archived_rollout), + ("other", "other", 0, &other_rollout), + ], + ); + let catalog_db = sqlite_dir.join("codex-dev.db"); + create_local_thread_catalog_db(&catalog_db, &[]); + + let archived = run_remote_control_session_finalization_for_thread_with_target( + Some(&home), + "archived", + "custom", + ); + let other = run_remote_control_session_finalization_for_thread_with_target( + Some(&home), + "other", + "custom", + ); + + assert_eq!(archived.status, ProviderSyncStatus::Synced); + assert_eq!(other.status, ProviderSyncStatus::Synced); + assert!(archived.message.contains("archived")); + assert!(other.message.contains("another provider")); + for (id, provider) in [("archived", "openai"), ("other", "other")] { + let first: serde_json::Value = serde_json::from_str( + fs::read_to_string(home.join(format!("sessions/rollout-{id}.jsonl"))) + .unwrap() + .lines() + .next() + .unwrap(), + ) + .unwrap(); + assert_eq!(first["payload"]["model_provider"], provider); + } + assert_eq!( + Connection::open(&catalog_db) + .unwrap() + .query_row("SELECT COUNT(*) FROM local_thread_catalog", [], |row| row + .get::<_, i64>( + 0 + ),) + .unwrap(), + 0 + ); +} + +#[test] +fn remote_control_finalization_defers_when_rollout_changes_after_collection() { + let tmp = tempdir().unwrap(); + let home = tmp.path().join(".codex"); + let sqlite_dir = home.join("sqlite"); + fs::create_dir_all(&sqlite_dir).unwrap(); + fs::write(home.join("config.toml"), "model_provider = \"custom\"\n").unwrap(); + let rollout = home.join("sessions/rollout-mobile.jsonl"); + write_rollout(&rollout, "openai", "mobile", "C:/workspace"); + let state_db = home.join("state_5.sqlite"); + create_remote_control_state_db(&state_db, &[("mobile", "openai", 0, &rollout)]); + let db = Connection::open(&state_db).unwrap(); + db.execute("CREATE TABLE backup_padding (data BLOB)", []) + .unwrap(); + db.execute("INSERT INTO backup_padding VALUES (zeroblob(33554432))", []) + .unwrap(); + drop(db); + let catalog_db = sqlite_dir.join("codex-dev.db"); + create_local_thread_catalog_db(&catalog_db, &[]); + + let backup_root = home.join("backups_state/provider-sync"); + let watched_rollout = rollout.clone(); + let writer = std::thread::spawn(move || { + let deadline = Instant::now() + Duration::from_secs(10); + loop { + let backup_started = backup_root.exists() + && fs::read_dir(&backup_root) + .map(|mut entries| entries.next().is_some()) + .unwrap_or(false); + if backup_started { + let mut file = fs::OpenOptions::new() + .append(true) + .open(&watched_rollout) + .unwrap(); + use std::io::Write as _; + writeln!( + file, + "{}", + json!({"type": "event_msg", "payload": {"type": "task_started"}}) + ) + .unwrap(); + return; + } + assert!(Instant::now() < deadline, "backup did not start in time"); + std::thread::sleep(Duration::from_millis(1)); + } + }); + + let result = run_remote_control_session_finalization_for_thread_with_target( + Some(&home), + "mobile", + "custom", + ); + writer.join().unwrap(); + + assert_eq!(result.status, ProviderSyncStatus::Skipped); + assert_eq!(result.changed_session_files, 0); + assert_eq!(result.skipped_locked_rollout_files.len(), 1); + assert_eq!( + fs::canonicalize(&result.skipped_locked_rollout_files[0]).unwrap(), + fs::canonicalize(&rollout).unwrap() + ); + let text = fs::read_to_string(&rollout).unwrap(); + assert!(text.contains("task_started")); + let first: serde_json::Value = serde_json::from_str(text.lines().next().unwrap()).unwrap(); + assert_eq!(first["payload"]["model_provider"], "openai"); + let state = Connection::open(&state_db).unwrap(); + let provider: String = state + .query_row( + "SELECT model_provider FROM threads WHERE id = 'mobile'", + [], + |row| row.get(0), + ) + .unwrap(); + assert_eq!(provider, "openai"); + let catalog = Connection::open(&catalog_db).unwrap(); + let catalog_rows: i64 = catalog + .query_row( + "SELECT COUNT(*) FROM local_thread_catalog WHERE thread_id = 'mobile'", + [], + |row| row.get(0), + ) + .unwrap(); + assert_eq!(catalog_rows, 0); +} + +#[test] +fn remote_control_finalization_retries_after_catalog_only_partial_commit() { + let tmp = tempdir().unwrap(); + let home = tmp.path().join(".codex"); + let sqlite_dir = home.join("sqlite"); + fs::create_dir_all(&sqlite_dir).unwrap(); + fs::write(home.join("config.toml"), "model_provider = \"custom\"\n").unwrap(); + let rollout = home.join("sessions/rollout-mobile.jsonl"); + write_rollout(&rollout, "openai", "mobile", "C:/workspace"); + let state_db = home.join("state_5.sqlite"); + create_remote_control_state_db(&state_db, &[("mobile", "openai", 0, &rollout)]); + let db = Connection::open(&state_db).unwrap(); + db.execute( + "CREATE TRIGGER fail_remote_recovery BEFORE UPDATE OF model_provider ON threads BEGIN SELECT RAISE(ABORT, 'boom'); END", + [], + ) + .unwrap(); + drop(db); + let catalog_db = sqlite_dir.join("codex-dev.db"); + create_local_thread_catalog_db(&catalog_db, &[]); + + let first = run_remote_control_session_finalization_for_thread_with_target( + Some(&home), + "mobile", + "custom", + ); + + assert_eq!(first.status, ProviderSyncStatus::Skipped); + let first_line: serde_json::Value = serde_json::from_str( + fs::read_to_string(&rollout) + .unwrap() + .lines() + .next() + .unwrap(), + ) + .unwrap(); + assert_eq!(first_line["payload"]["model_provider"], "custom"); + let catalog = Connection::open(&catalog_db).unwrap(); + let catalog_provider: String = catalog + .query_row( + "SELECT model_provider FROM local_thread_catalog WHERE host_id = 'local' AND thread_id = 'mobile'", + [], + |row| row.get(0), + ) + .unwrap(); + assert_eq!(catalog_provider, "custom"); + let state = Connection::open(&state_db).unwrap(); + let state_provider: String = state + .query_row( + "SELECT model_provider FROM threads WHERE id = 'mobile'", + [], + |row| row.get(0), + ) + .unwrap(); + assert_eq!(state_provider, "openai"); + state + .execute("DROP TRIGGER fail_remote_recovery", []) + .unwrap(); + drop(state); + + let second = run_remote_control_session_finalization_for_thread_with_target( + Some(&home), + "mobile", + "custom", + ); + + assert_eq!(second.status, ProviderSyncStatus::Synced); + assert_eq!(second.changed_session_files, 0); + let state = Connection::open(&state_db).unwrap(); + let state_provider: String = state + .query_row( + "SELECT model_provider FROM threads WHERE id = 'mobile'", + [], + |row| row.get(0), + ) + .unwrap(); + assert_eq!(state_provider, "custom"); +} + #[test] fn provider_sync_backup_metadata_contains_reference_fields_and_managed_marker() { let tmp = tempdir().unwrap();