diff --git a/codex-rs/app-server-transport/src/transport/remote_control/mod.rs b/codex-rs/app-server-transport/src/transport/remote_control/mod.rs index a9a7814f4f..32d8bf2ccf 100644 --- a/codex-rs/app-server-transport/src/transport/remote_control/mod.rs +++ b/codex-rs/app-server-transport/src/transport/remote_control/mod.rs @@ -106,6 +106,7 @@ pub(super) struct QueuedServerEnvelope { #[derive(Clone)] pub struct RemoteControlHandle { policy: RemoteControlPolicy, + shutdown_token: CancellationToken, desired_state_tx: Arc>, desired_state_rpc_lock: Arc, desired_state_persistence_lock: Arc, @@ -592,14 +593,23 @@ impl RemoteControlHandle { RemoteControlEnrollmentSelection::ReplaceExisting => {} } - let enrollment = enroll_pairing_server( - &self.auth_manager, - auth, - &remote_control_target, - installation_id, - server_name, - ) - .await?; + // Reused enrollments must still reach durable persistence during shutdown. + let enrollment = tokio::select! { + biased; + _ = self.shutdown_token.cancelled() => { + return Err(io::Error::new( + io::ErrorKind::Interrupted, + "remote control is shutting down", + )); + } + result = enroll_pairing_server( + &self.auth_manager, + auth, + &remote_control_target, + installation_id, + server_name, + ) => result?, + }; Ok((enrollment, true)) } @@ -1005,6 +1015,7 @@ pub async fn start_remote_control( let installation_id_for_log = installation_id.clone(); let server_name_for_log = server_name.clone(); let shutdown_token_for_log = shutdown_token.clone(); + let handle_shutdown_token = shutdown_token.clone(); let join_handle = tokio::spawn(async move { info!( remote_control_url = %remote_control_url_for_log, @@ -1070,6 +1081,7 @@ pub async fn start_remote_control( join_handle, RemoteControlHandle { policy, + shutdown_token: handle_shutdown_token, desired_state_tx, desired_state_rpc_lock, desired_state_persistence_lock, diff --git a/codex-rs/app-server-transport/src/transport/remote_control/tests.rs b/codex-rs/app-server-transport/src/transport/remote_control/tests.rs index 0114a0bad5..e1bfaa9aca 100644 --- a/codex-rs/app-server-transport/src/transport/remote_control/tests.rs +++ b/codex-rs/app-server-transport/src/transport/remote_control/tests.rs @@ -411,6 +411,7 @@ pub(super) fn remote_control_handle_with_current_enrollment( ))); RemoteControlHandle { policy: RemoteControlPolicy::Allowed, + shutdown_token: CancellationToken::new(), desired_state_tx: Arc::new(desired_state_tx), desired_state_rpc_lock: Arc::new(Semaphore::new(1)), desired_state_persistence_lock: Arc::new(Semaphore::new(1)), @@ -424,6 +425,216 @@ pub(super) fn remote_control_handle_with_current_enrollment( } } +#[tokio::test] +async fn durable_enable_reuses_in_memory_enrollment_after_shutdown() { + let listener = TcpListener::bind("127.0.0.1:0") + .await + .expect("listener should bind"); + let remote_control_url = remote_control_url_for_listener(&listener); + let codex_home = TempDir::new().expect("temp dir should create"); + let state_db = remote_control_state_runtime(&codex_home).await; + let mut remote_handle = remote_control_handle_with_current_enrollment( + &remote_control_url, + remote_control_auth_manager(), + ); + remote_handle.state_db = Some(state_db.clone()); + remote_handle + .desired_state_tx + .send_replace(RemoteControlDesiredState::Disabled); + let enrollment = remote_handle + .current_enrollment + .snapshot() + .expect("in-memory enrollment should exist"); + let expected_record = RemoteControlEnrollmentRecord { + websocket_url: enrollment.remote_control_target.websocket_url.clone(), + account_id: enrollment.account_id.clone(), + app_server_client_name: None, + server_id: enrollment.server_id.clone(), + environment_id: enrollment.environment_id.clone(), + server_name: enrollment.server_name.clone(), + remote_control_enabled: Some(true), + }; + assert_eq!( + state_db + .get_remote_control_enrollment( + &expected_record.websocket_url, + &expected_record.account_id, + /*app_server_client_name*/ None, + ) + .await + .expect("enrollment should load"), + None + ); + remote_handle.shutdown_token.cancel(); + + let status = timeout( + Duration::from_secs(5), + remote_handle.enable(/*app_server_client_name*/ None), + ) + .await + .expect("cached enable should complete without network I/O") + .expect("shutdown should not cancel durable enable using in-memory enrollment"); + + assert_eq!( + state_db + .get_remote_control_enrollment( + &expected_record.websocket_url, + &expected_record.account_id, + /*app_server_client_name*/ None, + ) + .await + .expect("enabled enrollment should load"), + Some(expected_record) + ); + assert_eq!( + *remote_handle.desired_state_tx.borrow(), + RemoteControlDesiredState::Enabled { + persistence_preference: Some(true), + } + ); + assert_eq!( + status.environment_id.as_deref(), + Some(enrollment.environment_id.as_str()) + ); + assert_eq!( + remote_handle.current_enrollment.snapshot(), + Some(enrollment) + ); + timeout(Duration::from_millis(100), listener.accept()) + .await + .expect_err("in-memory enrollment should prevent backend contact"); +} + +#[tokio::test] +async fn durable_enable_reuses_persisted_enrollment_after_shutdown() { + let listener = TcpListener::bind("127.0.0.1:0") + .await + .expect("listener should bind"); + let remote_control_url = remote_control_url_for_listener(&listener); + let remote_control_target = normalize_remote_control_url(&remote_control_url) + .expect("remote control target should normalize"); + let codex_home = TempDir::new().expect("temp dir should create"); + let state_db = remote_control_state_runtime(&codex_home).await; + let persisted_enrollment = RemoteControlEnrollmentRecord { + websocket_url: remote_control_target.websocket_url, + account_id: "account_id".to_string(), + app_server_client_name: None, + server_id: "persisted-server-id".to_string(), + environment_id: "persisted-environment-id".to_string(), + server_name: format!("{}-persisted", test_server_name()), + remote_control_enabled: Some(false), + }; + state_db + .upsert_remote_control_enrollment(&persisted_enrollment) + .await + .expect("disabled enrollment should persist"); + let mut remote_handle = remote_control_handle_with_current_enrollment( + &remote_control_url, + remote_control_auth_manager(), + ); + remote_handle.state_db = Some(state_db.clone()); + *remote_handle.current_enrollment.lock().await = None; + remote_handle + .desired_state_tx + .send_replace(RemoteControlDesiredState::Disabled); + remote_handle.shutdown_token.cancel(); + + let status = timeout( + Duration::from_secs(5), + remote_handle.enable(/*app_server_client_name*/ None), + ) + .await + .expect("cached enable should complete without network I/O") + .expect("shutdown should not cancel durable enable using persisted enrollment"); + + assert_eq!( + status.environment_id.as_deref(), + Some(persisted_enrollment.environment_id.as_str()) + ); + assert_eq!( + remote_handle + .current_enrollment + .snapshot() + .map(|enrollment| enrollment.server_id), + Some(persisted_enrollment.server_id.clone()) + ); + assert_eq!( + state_db + .get_remote_control_enrollment( + &persisted_enrollment.websocket_url, + &persisted_enrollment.account_id, + /*app_server_client_name*/ None, + ) + .await + .expect("enabled enrollment should load"), + Some(RemoteControlEnrollmentRecord { + remote_control_enabled: Some(true), + ..persisted_enrollment + }) + ); + assert_eq!( + *remote_handle.desired_state_tx.borrow(), + RemoteControlDesiredState::Enabled { + persistence_preference: Some(true), + } + ); + timeout(Duration::from_millis(100), listener.accept()) + .await + .expect_err("persisted enrollment should prevent backend contact"); +} + +#[tokio::test] +async fn durable_enable_without_cached_enrollment_is_cancelled_after_shutdown() { + let listener = TcpListener::bind("127.0.0.1:0") + .await + .expect("listener should bind"); + let remote_control_url = remote_control_url_for_listener(&listener); + let remote_control_target = normalize_remote_control_url(&remote_control_url) + .expect("remote control target should normalize"); + let codex_home = TempDir::new().expect("temp dir should create"); + let state_db = remote_control_state_runtime(&codex_home).await; + let mut remote_handle = remote_control_handle_with_current_enrollment( + &remote_control_url, + remote_control_auth_manager(), + ); + remote_handle.state_db = Some(state_db.clone()); + *remote_handle.current_enrollment.lock().await = None; + remote_handle + .desired_state_tx + .send_replace(RemoteControlDesiredState::Disabled); + remote_handle.shutdown_token.cancel(); + + let error = timeout( + Duration::from_secs(5), + remote_handle.enable(/*app_server_client_name*/ None), + ) + .await + .expect("shutdown should cancel network enrollment promptly") + .expect_err("enable without cached enrollment should be cancelled"); + + assert_eq!(error.kind(), std::io::ErrorKind::Interrupted); + assert_eq!(error.to_string(), "remote control is shutting down"); + assert_eq!(remote_handle.current_enrollment.snapshot(), None); + assert_eq!( + *remote_handle.desired_state_tx.borrow(), + RemoteControlDesiredState::Disabled + ); + assert_eq!( + state_db + .get_remote_control_enrollment( + &remote_control_target.websocket_url, + "account_id", + /*app_server_client_name*/ None, + ) + .await + .expect("enrollment should load"), + None + ); + timeout(Duration::from_millis(100), listener.accept()) + .await + .expect_err("cancelled enrollment should prevent backend contact"); +} + #[tokio::test] async fn ephemeral_enable_preserves_durable_preference() { let codex_home = TempDir::new().expect("temp dir should create"); diff --git a/codex-rs/app-server-transport/src/transport/remote_control/tests/clients_tests.rs b/codex-rs/app-server-transport/src/transport/remote_control/tests/clients_tests.rs index e88082ecf9..92c9e95358 100644 --- a/codex-rs/app-server-transport/src/transport/remote_control/tests/clients_tests.rs +++ b/codex-rs/app-server-transport/src/transport/remote_control/tests/clients_tests.rs @@ -23,6 +23,7 @@ fn client_management_handle( }); RemoteControlHandle { policy: RemoteControlPolicy::Allowed, + shutdown_token: CancellationToken::new(), desired_state_tx: Arc::new(desired_state_tx), desired_state_rpc_lock: Arc::new(Semaphore::new(1)), desired_state_persistence_lock: Arc::new(Semaphore::new(1)), diff --git a/codex-rs/app-server/src/lib.rs b/codex-rs/app-server/src/lib.rs index 7edccfbf66..5f6c97cecd 100644 --- a/codex-rs/app-server/src/lib.rs +++ b/codex-rs/app-server/src/lib.rs @@ -731,6 +731,8 @@ pub async fn run_main_with_transport_options( } let installation_id = resolve_installation_id(&config.codex_home).await?; let transport_shutdown_token = CancellationToken::new(); + // Remote enrollment must cancel before RPC drain without shutting down telemetry. + let remote_control_shutdown_token = transport_shutdown_token.child_token(); let mut transport_accept_handles = Vec::>::new(); let single_client_mode = matches!(&transport, AppServerTransport::Stdio); @@ -809,7 +811,7 @@ pub async fn run_main_with_transport_options( state_db.clone(), auth_manager.clone(), transport_event_tx.clone(), - transport_shutdown_token.clone(), + remote_control_shutdown_token.clone(), app_server_client_name_rx, remote_control_startup_mode, ) @@ -1051,6 +1053,8 @@ pub async fn run_main_with_transport_options( break "outbound_router_closed"; } if single_client_mode && stdio_closed { + // Pending remote enrollment must stop before RPCs drain. + remote_control_shutdown_token.cancel(); break "stdio_connection_closed"; } } diff --git a/codex-rs/app-server/tests/suite/v2/remote_control.rs b/codex-rs/app-server/tests/suite/v2/remote_control.rs index 17e0d2c5d4..9b0def1370 100644 --- a/codex-rs/app-server/tests/suite/v2/remote_control.rs +++ b/codex-rs/app-server/tests/suite/v2/remote_control.rs @@ -10,6 +10,7 @@ use app_test_support::ChatGptAuthFixture; use app_test_support::DEFAULT_CLIENT_NAME; use app_test_support::MockResponsesConfig; use app_test_support::TestAppServer; +use app_test_support::create_fake_paginated_rollout; use app_test_support::to_response; use app_test_support::write_chatgpt_auth; use codex_app_server::AppServerRuntimeOptions; @@ -37,6 +38,8 @@ use codex_app_server_protocol::RemoteControlPairingStatusResponse; use codex_app_server_protocol::RemoteControlStatusChangedNotification; use codex_app_server_protocol::RemoteControlStatusReadResponse; use codex_app_server_protocol::RequestId; +use codex_app_server_protocol::ThreadResumeParams; +use codex_app_server_protocol::ThreadResumeResponse; use codex_arg0::Arg0DispatchPaths; use codex_config::LoaderOverrides; use codex_config::types::AuthCredentialsStoreMode; @@ -493,6 +496,113 @@ async fn stdio_eof_exits_with_remote_control_connection() -> Result<()> { Ok(()) } +#[tokio::test] +async fn stdio_eof_releases_thread_writer_with_pending_remote_control_enable() -> Result<()> { + let codex_home = TempDir::new()?; + let mut backend = BlockingRemoteControlBackend::start(codex_home.path()).await?; + let config_path = codex_home.path().join("config.toml"); + let config = std::fs::read_to_string(&config_path)?; + // Keep thread initialization from using the enrollment-only backend for unrelated requests. + std::fs::write( + config_path, + format!( + "{config}\n[features]\napps = false\nremote_plugin = false\n\n[analytics]\nenabled = false\n" + ), + )?; + let thread_id = create_fake_paginated_rollout( + codex_home.path(), + "2025-01-01T00-00-00", + "2025-01-01T00:00:00Z", + "owned thread", + Some("mock_provider"), + /*git_info*/ None, + )?; + let mut owner = TestAppServer::builder() + .with_codex_home(codex_home.path()) + .without_auto_env() + .build_initialized() + .await?; + let _: ThreadResumeResponse = owner + .request(|request_id| ClientRequest::ThreadResume { + request_id, + params: ThreadResumeParams { + thread_id: thread_id.clone(), + exclude_turns: true, + ..Default::default() + }, + }) + .await?; + + let secondary_sqlite_home = TempDir::new()?; + let secondary_sqlite_home_path = secondary_sqlite_home.path().to_string_lossy(); + let mut secondary = TestAppServer::builder() + .with_codex_home(codex_home.path()) + .without_auto_env() + .with_env_overrides(&[( + "CODEX_SQLITE_HOME", + Some(secondary_sqlite_home_path.as_ref()), + )]) + .build_initialized() + .await?; + let resume_id = secondary + .send_thread_resume_request(ThreadResumeParams { + thread_id: thread_id.clone(), + exclude_turns: true, + ..Default::default() + }) + .await?; + let error = timeout( + DEFAULT_TIMEOUT, + secondary.read_stream_until_error_message(RequestId::Integer(resume_id)), + ) + .await??; + assert_eq!(error.error.code, -32600); + assert_eq!( + error.error.message, + format!("thread {thread_id} already has an active writer") + ); + + owner.send_remote_control_enable_request().await?; + assert_eq!( + timeout(DEFAULT_TIMEOUT, backend.wait_for_enroll_request()).await??, + "POST /backend-api/wham/remote/control/server/enroll HTTP/1.1" + ); + // Keep enrollment pending while EOF requests teardown of the owning process. + let status = timeout(DEFAULT_TIMEOUT, owner.shutdown_gracefully()) + .await + .context("stdio EOF did not stop the thread writer while enrollment was pending")??; + assert!(status.success()); + + let state_db = StateRuntime::init( + codex_state::SqliteConfig::new_for_testing(codex_home.path().abs()), + "test-provider".to_string(), + ) + .await?; + assert_eq!( + state_db + .get_remote_control_enrollment( + backend.websocket_url(), + "account_id", + Some(DEFAULT_CLIENT_NAME), + ) + .await?, + None + ); + + let resumed: ThreadResumeResponse = secondary + .request(|request_id| ClientRequest::ThreadResume { + request_id, + params: ThreadResumeParams { + thread_id: thread_id.clone(), + exclude_turns: true, + ..Default::default() + }, + }) + .await?; + assert_eq!(resumed.thread.id, thread_id); + Ok(()) +} + #[tokio::test] async fn disable_waits_for_in_flight_durable_enable() -> Result<()> { let codex_home = TempDir::new()?;