diff --git a/codex-rs/app-server/src/connection_cleanup.rs b/codex-rs/app-server/src/connection_cleanup.rs new file mode 100644 index 00000000000..201f7fb4ba7 --- /dev/null +++ b/codex-rs/app-server/src/connection_cleanup.rs @@ -0,0 +1,97 @@ +use std::future::Future; +use std::future::pending; + +use tokio::task::JoinError; +use tokio::task::JoinSet; +use tracing::warn; + +pub(crate) struct ConnectionCleanupTasks { + tasks: JoinSet<()>, +} + +impl ConnectionCleanupTasks { + pub(crate) fn new() -> Self { + Self { + tasks: JoinSet::new(), + } + } + + pub(crate) fn spawn(&mut self, future: impl Future + Send + 'static) { + self.tasks.spawn(future); + } + + pub(crate) async fn reap_next(&mut self) { + if self.tasks.is_empty() { + pending::<()>().await; + } + if let Some(result) = self.tasks.join_next().await { + log_cleanup_result(result); + } + } + + pub(crate) async fn drain(&mut self) { + while let Some(result) = self.tasks.join_next().await { + log_cleanup_result(result); + } + } + + pub(crate) fn abort(&mut self) { + self.tasks.abort_all(); + } +} + +fn log_cleanup_result(result: Result<(), JoinError>) { + if let Err(err) = result + && !err.is_cancelled() + { + warn!("connection cleanup task failed: {err}"); + } +} + +#[cfg(test)] +mod tests { + use super::*; + use tokio::sync::oneshot; + use tokio::time::Duration; + use tokio::time::timeout; + + #[tokio::test] + async fn reap_next_waits_when_no_cleanup_tasks_exist() { + let mut tasks = ConnectionCleanupTasks::new(); + + timeout(Duration::from_millis(/*millis*/ 20), tasks.reap_next()) + .await + .expect_err("empty cleanup task set should stay pending"); + } + + #[tokio::test] + async fn reap_next_removes_completed_cleanup_task() { + let mut tasks = ConnectionCleanupTasks::new(); + tasks.spawn(async {}); + + timeout(Duration::from_secs(/*secs*/ 1), tasks.reap_next()) + .await + .expect("completed cleanup task should be reaped"); + + assert!(tasks.tasks.is_empty()); + } + + #[tokio::test] + async fn abort_cancels_blocked_cleanup_task() { + let mut tasks = ConnectionCleanupTasks::new(); + let (started_tx, started_rx) = oneshot::channel(); + let (_release_tx, release_rx) = oneshot::channel::<()>(); + tasks.spawn(async move { + let _ = started_tx.send(()); + let _ = release_rx.await; + }); + + started_rx.await.expect("cleanup task should start"); + tasks.abort(); + timeout(Duration::from_secs(/*secs*/ 1), tasks.drain()) + .await + .expect("aborted cleanup task should drain"); + + assert!(tasks.tasks.is_empty()); + } +} diff --git a/codex-rs/app-server/src/connection_rpc_gate.rs b/codex-rs/app-server/src/connection_rpc_gate.rs index 12fed79b363..fb2aedd352b 100644 --- a/codex-rs/app-server/src/connection_rpc_gate.rs +++ b/codex-rs/app-server/src/connection_rpc_gate.rs @@ -38,12 +38,14 @@ impl ConnectionRpcGate { drop(token); } + pub(crate) async fn close(&self) { + let mut accepting = self.accepting.lock().await; + *accepting = false; + self.tasks.close(); + } + pub(crate) async fn shutdown(&self) { - { - let mut accepting = self.accepting.lock().await; - *accepting = false; - self.tasks.close(); - } + self.close().await; self.tasks.wait().await; } @@ -90,9 +92,9 @@ mod tests { } #[tokio::test] - async fn run_drops_future_without_polling_after_shutdown() { + async fn run_drops_future_without_polling_after_close() { let gate = ConnectionRpcGate::new(); - gate.shutdown().await; + gate.close().await; let polled = Arc::new(AtomicBool::new(/*v*/ false)); let polled_clone = Arc::clone(&polled); @@ -105,6 +107,33 @@ mod tests { assert!(!gate.is_accepting().await); } + #[tokio::test] + async fn close_returns_while_started_run_remains_active() { + let gate = Arc::new(ConnectionRpcGate::new()); + let (started_tx, started_rx) = oneshot::channel(); + let (finish_tx, finish_rx) = oneshot::channel(); + let gate_for_run = Arc::clone(&gate); + let run_task = tokio::spawn(async move { + gate_for_run + .run(async move { + started_tx.send(()).expect("receiver should be open"); + let _ = finish_rx.await; + }) + .await; + }); + + started_rx.await.expect("run should start"); + gate.close().await; + assert!(!gate.is_accepting().await); + assert_eq!(gate.inflight_count(), 1); + + finish_tx + .send(()) + .expect("running future should be waiting"); + run_task.await.expect("run task should complete"); + gate.shutdown().await; + } + #[tokio::test] async fn shutdown_waits_for_started_run_to_finish() { let gate = Arc::new(ConnectionRpcGate::new()); diff --git a/codex-rs/app-server/src/lib.rs b/codex-rs/app-server/src/lib.rs index bb767100c01..1ded6bf4e6e 100644 --- a/codex-rs/app-server/src/lib.rs +++ b/codex-rs/app-server/src/lib.rs @@ -20,6 +20,7 @@ use std::sync::atomic::AtomicBool; use crate::analytics_utils::analytics_events_client_from_config; use crate::config_manager::ConfigManager; +use crate::connection_cleanup::ConnectionCleanupTasks; use crate::message_processor::MessageProcessor; use crate::message_processor::MessageProcessorArgs; use crate::outgoing_message::ConnectionId; @@ -81,6 +82,7 @@ mod command_exec; mod config; mod config_manager; mod config_manager_service; +mod connection_cleanup; mod connection_rpc_gate; mod dynamic_tools; mod error_code; @@ -821,6 +823,7 @@ pub async fn run_main_with_transport_options( let mut thread_created_rx = processor.thread_created_receiver(); let mut running_turn_count_rx = processor.subscribe_running_assistant_turn_count(); let mut connections = HashMap::::new(); + let mut connection_cleanup_tasks = ConnectionCleanupTasks::new(); let mut remote_control_status_rx = remote_control_handle.status_receiver(); let mut remote_control_status = remote_control_status_rx.borrow().clone(); let transport_shutdown_token = transport_shutdown_token.clone(); @@ -908,14 +911,21 @@ pub async fn run_main_with_transport_options( let Some(connection_state) = connections.remove(&connection_id) else { continue; }; - if outbound_control_tx + connection_state.session.rpc_gate.close().await; + let outbound_closed = outbound_control_tx .send(OutboundControlEvent::Closed { connection_id }) .await - .is_err() - { + .is_ok(); + processor.connection_closing(connection_id).await; + let processor = Arc::clone(&processor); + connection_cleanup_tasks.spawn(async move { + processor + .connection_closed(connection_id, &connection_state.session) + .await; + }); + if !outbound_closed { break; } - processor.connection_closed(connection_id, &connection_state.session).await; if shutdown_when_no_connections && connections.is_empty() { break; } @@ -1012,6 +1022,7 @@ pub async fn run_main_with_transport_options( } } } + _ = connection_cleanup_tasks.reap_next() => {} changed = remote_control_status_rx.changed() => { if changed.is_err() { continue; @@ -1064,8 +1075,11 @@ pub async fn run_main_with_transport_options( .map(|connection_state| connection_state.session.rpc_gate.shutdown()), ) .await; + connection_cleanup_tasks.drain().await; processor.drain_background_tasks().await; processor.shutdown_threads().await; + } else { + connection_cleanup_tasks.abort(); } info!("processor task exited (channel closed)"); } diff --git a/codex-rs/app-server/src/message_processor.rs b/codex-rs/app-server/src/message_processor.rs index c383c527157..165705c91cd 100644 --- a/codex-rs/app-server/src/message_processor.rs +++ b/codex-rs/app-server/src/message_processor.rs @@ -729,13 +729,18 @@ impl MessageProcessor { self.thread_processor.shutdown_threads().await; } + pub(crate) async fn connection_closing(&self, connection_id: ConnectionId) { + self.outgoing.connection_closed(connection_id).await; + self.thread_processor.connection_closed(connection_id).await; + } + pub(crate) async fn connection_closed( &self, connection_id: ConnectionId, session_state: &ConnectionSessionState, ) { + tracing::debug!(?connection_id, "connection cleanup started"); session_state.rpc_gate.shutdown().await; - self.outgoing.connection_closed(connection_id).await; self.fs_processor.connection_closed(connection_id).await; self.command_exec_processor .connection_closed(connection_id) @@ -743,7 +748,7 @@ impl MessageProcessor { self.process_exec_processor .connection_closed(connection_id) .await; - self.thread_processor.connection_closed(connection_id).await; + tracing::debug!(?connection_id, "connection cleanup completed"); } pub(crate) fn subscribe_running_assistant_turn_count(&self) -> watch::Receiver { diff --git a/codex-rs/app-server/src/request_serialization.rs b/codex-rs/app-server/src/request_serialization.rs index 0dd167b74dc..77ecfc8f56c 100644 --- a/codex-rs/app-server/src/request_serialization.rs +++ b/codex-rs/app-server/src/request_serialization.rs @@ -311,7 +311,7 @@ mod tests { let key = RequestSerializationQueueKey::Global("test"); let live_gate = gate(); let closed_gate = gate(); - closed_gate.shutdown().await; + closed_gate.close().await; let (tx, mut rx) = mpsc::unbounded_channel(); let (blocked_tx, blocked_rx) = oneshot::channel::<()>(); diff --git a/codex-rs/app-server/tests/suite/v2/connection_handling_websocket.rs b/codex-rs/app-server/tests/suite/v2/connection_handling_websocket.rs index 4eff8d44013..31e4aef2695 100644 --- a/codex-rs/app-server/tests/suite/v2/connection_handling_websocket.rs +++ b/codex-rs/app-server/tests/suite/v2/connection_handling_websocket.rs @@ -1,10 +1,12 @@ use anyhow::Context; use anyhow::Result; use anyhow::bail; +use app_test_support::ChatGptAuthFixture; use app_test_support::DISABLE_PLUGIN_STARTUP_TASKS_ARG; use app_test_support::USE_TEST_KEYRING_STORE_ARG; use app_test_support::create_mock_responses_server_sequence_unchecked; use app_test_support::to_response; +use app_test_support::write_chatgpt_auth; use base64::Engine; use base64::engine::general_purpose::URL_SAFE_NO_PAD; use codex_app_server_protocol::ClientInfo; @@ -14,11 +16,14 @@ use codex_app_server_protocol::JSONRPCMessage; use codex_app_server_protocol::JSONRPCNotification; use codex_app_server_protocol::JSONRPCRequest; use codex_app_server_protocol::JSONRPCResponse; +use codex_app_server_protocol::PluginListMarketplaceKind; +use codex_app_server_protocol::PluginListParams; use codex_app_server_protocol::RequestId; use codex_app_server_protocol::ThreadLoadedListParams; use codex_app_server_protocol::ThreadLoadedListResponse; use codex_app_server_protocol::ThreadStartParams; use codex_app_server_protocol::ThreadStartResponse; +use codex_config::types::AuthCredentialsStoreMode; #[cfg(debug_assertions)] use codex_keyring_store::tests::shared_test_keyring_root; use futures::SinkExt; @@ -31,12 +36,17 @@ use sha2::Sha256; use std::net::SocketAddr; use std::path::Path; use std::process::Stdio; +use std::sync::Arc; +use std::sync::Mutex; +use std::sync::mpsc as std_mpsc; use tempfile::TempDir; use time::OffsetDateTime; use tokio::io::AsyncBufReadExt; use tokio::io::BufReader; use tokio::process::Child; use tokio::process::Command; +use tokio::sync::mpsc; +use tokio::sync::oneshot; use tokio::time::Duration; use tokio::time::Instant; use tokio::time::sleep; @@ -50,6 +60,14 @@ use tokio_tungstenite::tungstenite::client::IntoClientRequest; use tokio_tungstenite::tungstenite::http::HeaderValue; use tokio_tungstenite::tungstenite::http::header::AUTHORIZATION; use tokio_tungstenite::tungstenite::http::header::ORIGIN; +use wiremock::Mock; +use wiremock::MockServer; +use wiremock::Request as WiremockRequest; +use wiremock::Respond; +use wiremock::ResponseTemplate; +use wiremock::matchers::header; +use wiremock::matchers::method; +use wiremock::matchers::path; // macOS and Windows CI can spend tens of seconds starting the app-server test // binary under Bazel before it accepts JSON-RPC or reports its websocket bind @@ -62,6 +80,58 @@ pub(super) const DEFAULT_READ_TIMEOUT: Duration = Duration::from_secs(10); pub(super) type WsClient = WebSocketStream>; type HmacSha256 = Hmac; +#[derive(Clone)] +struct GatedWorkspaceSettingsResponse { + request_started_tx: Arc>>>, + release_rx: Arc>>>, +} + +struct WorkspaceSettingsReleaseGuard { + release_tx: Option>, +} + +impl WorkspaceSettingsReleaseGuard { + fn new(release_tx: std_mpsc::Sender<()>) -> Self { + Self { + release_tx: Some(release_tx), + } + } + + fn release(&mut self) { + if let Some(release_tx) = self.release_tx.take() { + let _ = release_tx.send(()); + } + } +} + +impl Drop for WorkspaceSettingsReleaseGuard { + fn drop(&mut self) { + self.release(); + } +} + +impl Respond for GatedWorkspaceSettingsResponse { + fn respond(&self, _: &WiremockRequest) -> ResponseTemplate { + let request_started_tx = self + .request_started_tx + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner) + .take(); + if let Some(request_started_tx) = request_started_tx { + let _ = request_started_tx.send(()); + } + let release_rx = self + .release_rx + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner) + .take(); + if let Some(release_rx) = release_rx { + let _ = release_rx.recv(); + } + ResponseTemplate::new(200).set_body_string(r#"{"beta_settings":{"enable_plugins":true}}"#) + } +} + #[tokio::test] async fn websocket_transport_routes_per_connection_handshake_and_responses() -> Result<()> { let server = create_mock_responses_server_sequence_unchecked(Vec::new()).await; @@ -107,6 +177,91 @@ async fn websocket_transport_routes_per_connection_handshake_and_responses() -> Ok(()) } +#[tokio::test] +async fn websocket_reinitialize_is_not_blocked_by_disconnected_client_rpc() -> Result<()> { + let server = create_mock_responses_server_sequence_unchecked(Vec::new()).await; + let workspace_server = MockServer::start().await; + let codex_home = TempDir::new()?; + create_config_toml(codex_home.path(), &server.uri(), "never")?; + let config_path = codex_home.path().join("config.toml"); + let config = std::fs::read_to_string(&config_path)?; + std::fs::write( + config_path, + format!( + "chatgpt_base_url = \"{}/backend-api/\"\n{config}\n[features]\nplugins = true\n", + workspace_server.uri() + ), + )?; + write_chatgpt_auth( + codex_home.path(), + ChatGptAuthFixture::new("chatgpt-token") + .account_id("account-123") + .chatgpt_user_id("user-123") + .chatgpt_account_id("account-123") + .plan_type("team"), + AuthCredentialsStoreMode::File, + )?; + let (request_started_tx, request_started_rx) = oneshot::channel(); + let (release_tx, release_rx) = std_mpsc::channel(); + let mut release_guard = WorkspaceSettingsReleaseGuard::new(release_tx); + Mock::given(method("GET")) + .and(path("/backend-api/accounts/account-123/settings")) + .and(header("authorization", "Bearer chatgpt-token")) + .and(header("chatgpt-account-id", "account-123")) + .respond_with(GatedWorkspaceSettingsResponse { + request_started_tx: Arc::new(Mutex::new(Some(request_started_tx))), + release_rx: Arc::new(Mutex::new(Some(release_rx))), + }) + .mount(&workspace_server) + .await; + + let (mut process, bind_addr, mut server_logs) = spawn_websocket_server_with_args_and_logs( + codex_home.path(), + "ws://127.0.0.1:0", + &[], + "codex_app_server::message_processor=debug,codex_app_server_transport=warn", + ) + .await?; + let mut ws1 = connect_websocket(bind_addr).await?; + send_initialize_request(&mut ws1, /*id*/ 1, "blocking_plugin_client").await?; + read_response_for_id(&mut ws1, /*id*/ 1).await?; + + send_request( + &mut ws1, + "plugin/list", + /*id*/ 2, + Some(serde_json::to_value(PluginListParams { + cwds: None, + marketplace_kinds: Some(vec![PluginListMarketplaceKind::Local]), + })?), + ) + .await?; + timeout(DEFAULT_READ_TIMEOUT, request_started_rx) + .await + .context("plugin/list did not start the blocking workspace-settings request")??; + + drop(ws1); + wait_for_server_log(&mut server_logs, "connection cleanup started").await?; + let mut ws2 = connect_websocket(bind_addr).await?; + send_initialize_request(&mut ws2, /*id*/ 3, "replacement_client").await?; + let initialize_result = timeout( + Duration::from_secs(/*secs*/ 5), + read_response_for_id(&mut ws2, /*id*/ 3), + ) + .await; + release_guard.release(); + let initialize_response = initialize_result + .context("replacement initialize was blocked by disconnected client cleanup")??; + assert_eq!(initialize_response.id, RequestId::Integer(3)); + wait_for_server_log(&mut server_logs, "connection cleanup completed").await?; + + process + .kill() + .await + .context("failed to stop websocket app-server process")?; + Ok(()) +} + #[tokio::test] async fn websocket_transport_serves_health_endpoints_on_same_listener() -> Result<()> { let server = create_mock_responses_server_sequence_unchecked(Vec::new()).await; @@ -387,6 +542,18 @@ pub(super) async fn spawn_websocket_server_with_args( listen_url: &str, extra_args: &[String], ) -> Result<(Child, SocketAddr)> { + let (process, bind_addr, _server_logs) = + spawn_websocket_server_with_args_and_logs(codex_home, listen_url, extra_args, "warn") + .await?; + Ok((process, bind_addr)) +} + +async fn spawn_websocket_server_with_args_and_logs( + codex_home: &Path, + listen_url: &str, + extra_args: &[String], + rust_log: &str, +) -> Result<(Child, SocketAddr, mpsc::UnboundedReceiver)> { let program = codex_utils_cargo_bin::cargo_bin("codex-app-server") .context("should find app-server binary")?; let mut cmd = Command::new(program); @@ -399,7 +566,7 @@ pub(super) async fn spawn_websocket_server_with_args( .stdout(Stdio::null()) .stderr(Stdio::piped()) .env("CODEX_LAB_HOME", codex_home) - .env("RUST_LOG", "warn"); + .env("RUST_LOG", rust_log); #[cfg(debug_assertions)] cmd.env( "CODEX_APP_SERVER_TEST_KEYRING_DIR", @@ -415,6 +582,7 @@ pub(super) async fn spawn_websocket_server_with_args( .take() .context("failed to capture websocket app-server stderr")?; let mut stderr_reader = BufReader::new(stderr).lines(); + let (stderr_line_tx, stderr_line_rx) = mpsc::unbounded_channel(); let deadline = Instant::now() + DEFAULT_READ_TIMEOUT; let bind_addr = loop { let line = timeout( @@ -426,6 +594,7 @@ pub(super) async fn spawn_websocket_server_with_args( .context("failed to read websocket app-server stderr")? .context("websocket app-server exited before reporting bound websocket address")?; eprintln!("[websocket app-server stderr] {line}"); + let _ = stderr_line_tx.send(line.clone()); let stripped_line = { let mut stripped = String::with_capacity(line.len()); @@ -457,10 +626,28 @@ pub(super) async fn spawn_websocket_server_with_args( tokio::spawn(async move { while let Ok(Some(line)) = stderr_reader.next_line().await { eprintln!("[websocket app-server stderr] {line}"); + let _ = stderr_line_tx.send(line); } }); - Ok((process, bind_addr)) + Ok((process, bind_addr, stderr_line_rx)) +} + +async fn wait_for_server_log( + server_logs: &mut mpsc::UnboundedReceiver, + expected: &str, +) -> Result<()> { + timeout(DEFAULT_READ_TIMEOUT, async { + while let Some(line) = server_logs.recv().await { + if line.contains(expected) { + return Ok::<(), anyhow::Error>(()); + } + } + bail!("websocket app-server log stream closed before `{expected}`") + }) + .await + .with_context(|| format!("timed out waiting for websocket app-server log `{expected}`"))??; + Ok(()) } pub(super) async fn connect_websocket(bind_addr: SocketAddr) -> Result {