diff --git a/codex-rs/exec-server/src/client.rs b/codex-rs/exec-server/src/client.rs index bb5a3855dc..aae415f047 100644 --- a/codex-rs/exec-server/src/client.rs +++ b/codex-rs/exec-server/src/client.rs @@ -12,7 +12,7 @@ use futures::FutureExt; use futures::future::BoxFuture; use serde_json::Value; use tokio::sync::Mutex; -use tokio::sync::OnceCell; +use tokio::sync::Semaphore; use tokio::sync::mpsc; use tokio::sync::watch; @@ -192,27 +192,62 @@ pub struct ExecServerClient { #[derive(Clone)] pub(crate) struct LazyRemoteExecServerClient { transport_params: ExecServerTransportParams, - client: Arc>, + client: Arc>>, + connect_lock: Arc, } impl LazyRemoteExecServerClient { pub(crate) fn new(transport_params: ExecServerTransportParams) -> Self { Self { transport_params, - client: Arc::new(OnceCell::new()), + client: Arc::new(StdMutex::new(None)), + connect_lock: Arc::new(Semaphore::new(/*permits*/ 1)), } } pub(crate) async fn get(&self) -> Result { + if let Some(client) = self.connected_client() { + return Ok(client); + } + + let _connect_permit = self.connect_lock.acquire().await.map_err(|_| { + ExecServerError::Protocol("exec-server connect lock closed".to_string()) + })?; + if let Some(client) = self.connected_client() { + return Ok(client); + } + + let next_client = match self.cached_client() { + Some(client) + if matches!( + &self.transport_params, + ExecServerTransportParams::WebSocketUrl { .. } + ) => + { + ExecServerClient::connect_for_transport(self.transport_params.clone()).await? + } + Some(client) => return Ok(client), + None => ExecServerClient::connect_for_transport(self.transport_params.clone()).await?, + }; + + let mut cached_client = self + .client + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner); + *cached_client = Some(next_client.clone()); + Ok(next_client) + } + + fn connected_client(&self) -> Option { + self.cached_client() + .filter(|client| !client.is_disconnected()) + } + + fn cached_client(&self) -> Option { self.client - // TODO: Add reconnect/disconnect handling here instead of reusing - // the first successfully initialized connection forever. - .get_or_try_init(|| { - let transport_params = self.transport_params.clone(); - async move { ExecServerClient::connect_for_transport(transport_params).await } - }) - .await - .cloned() + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner) + .clone() } } @@ -424,6 +459,10 @@ impl ExecServerClient { .clone() } + fn is_disconnected(&self) -> bool { + self.inner.disconnected.get().is_some() || self.inner.client.is_disconnected() + } + pub(crate) async fn connect( connection: JsonRpcConnection, options: ExecServerClientConnectOptions, @@ -873,30 +912,38 @@ mod tests { use codex_app_server_protocol::JSONRPCMessage; use codex_app_server_protocol::JSONRPCNotification; use codex_app_server_protocol::JSONRPCResponse; + use futures::SinkExt; + use futures::StreamExt; use pretty_assertions::assert_eq; use std::collections::HashMap; #[cfg(unix)] use std::path::Path; #[cfg(unix)] use std::process::Command; + use std::sync::Arc; use tokio::io::AsyncBufReadExt; use tokio::io::AsyncWrite; use tokio::io::AsyncWriteExt; use tokio::io::BufReader; use tokio::io::duplex; + use tokio::net::TcpListener; + use tokio::net::TcpStream; use tokio::sync::mpsc; use tokio::sync::oneshot; use tokio::time::Duration; #[cfg(unix)] use tokio::time::sleep; use tokio::time::timeout; + use tokio_tungstenite::WebSocketStream; + use tokio_tungstenite::accept_async; + use tokio_tungstenite::tungstenite::Message; use super::ExecServerClient; use super::ExecServerClientConnectOptions; + use super::LazyRemoteExecServerClient; use crate::ProcessId; #[cfg(not(windows))] use crate::client_api::DEFAULT_REMOTE_EXEC_SERVER_INITIALIZE_TIMEOUT; - #[cfg(not(windows))] use crate::client_api::ExecServerTransportParams; use crate::client_api::StdioExecServerCommand; use crate::client_api::StdioExecServerConnectArgs; @@ -937,6 +984,96 @@ mod tests { .expect("json-rpc line should write"); } + async fn accept_websocket(listener: &TcpListener) -> WebSocketStream { + let (stream, _) = listener.accept().await.expect("listener should accept"); + accept_async(stream) + .await + .expect("websocket handshake should succeed") + } + + async fn read_jsonrpc_websocket(websocket: &mut WebSocketStream) -> JSONRPCMessage { + loop { + match timeout(Duration::from_secs(1), websocket.next()) + .await + .expect("json-rpc websocket read should not time out") + .expect("websocket should stay open") + .expect("websocket frame should read") + { + Message::Text(text) => { + return serde_json::from_str(text.as_ref()) + .expect("json-rpc text frame should parse"); + } + Message::Binary(bytes) => { + return serde_json::from_slice(bytes.as_ref()) + .expect("json-rpc binary frame should parse"); + } + Message::Ping(_) | Message::Pong(_) => {} + other => panic!("expected json-rpc websocket frame, got {other:?}"), + } + } + } + + async fn write_jsonrpc_websocket( + websocket: &mut WebSocketStream, + message: JSONRPCMessage, + ) { + let encoded = serde_json::to_string(&message).expect("json-rpc should serialize"); + websocket + .send(Message::Text(encoded.into())) + .await + .expect("json-rpc websocket frame should write"); + } + + async fn complete_websocket_initialize( + websocket: &mut WebSocketStream, + session_id: &str, + expected_resume_session_id: Option<&str>, + ) { + let initialize = read_jsonrpc_websocket(websocket).await; + let request = match initialize { + JSONRPCMessage::Request(request) if request.method == INITIALIZE_METHOD => request, + other => panic!("expected initialize request, got {other:?}"), + }; + let params: crate::protocol::InitializeParams = + serde_json::from_value(request.params.expect("initialize params should exist")) + .expect("initialize params should deserialize"); + assert_eq!( + params.resume_session_id.as_deref(), + expected_resume_session_id + ); + write_jsonrpc_websocket( + websocket, + JSONRPCMessage::Response(JSONRPCResponse { + id: request.id, + result: serde_json::to_value(InitializeResponse { + session_id: session_id.to_string(), + }) + .expect("initialize response should serialize"), + }), + ) + .await; + + let initialized = read_jsonrpc_websocket(websocket).await; + match initialized { + JSONRPCMessage::Notification(notification) + if notification.method == INITIALIZED_METHOD => {} + other => panic!("expected initialized notification, got {other:?}"), + } + } + + async fn wait_for_disconnect(client: &ExecServerClient) { + timeout(Duration::from_secs(1), async { + loop { + if client.is_disconnected() { + return; + } + tokio::task::yield_now().await; + } + }) + .await + .expect("client should observe disconnect"); + } + #[cfg(not(windows))] #[tokio::test] async fn connect_stdio_command_initializes_json_rpc_client() { @@ -1354,6 +1491,57 @@ mod tests { server.await.expect("server task should finish"); } + #[tokio::test] + async fn remote_websocket_client_replaces_disconnected_client_with_fresh_session() { + let listener = TcpListener::bind("127.0.0.1:0") + .await + .expect("listener should bind"); + let websocket_url = format!( + "ws://{}", + listener.local_addr().expect("listener should have address") + ); + let server = tokio::spawn({ + async move { + let mut first = accept_websocket(&listener).await; + complete_websocket_initialize( + &mut first, + "session-1", + /*expected_resume_session_id*/ None, + ) + .await; + first + .close(None) + .await + .expect("first websocket should close"); + + let mut second = accept_websocket(&listener).await; + complete_websocket_initialize( + &mut second, + "session-2", + /*expected_resume_session_id*/ None, + ) + .await; + } + }); + + let client = LazyRemoteExecServerClient::new(ExecServerTransportParams::WebSocketUrl { + websocket_url, + connect_timeout: Duration::from_secs(1), + initialize_timeout: Duration::from_secs(1), + }); + let first = client.get().await.expect("first client should connect"); + wait_for_disconnect(&first).await; + + let (replacement_a, replacement_b) = tokio::join!(client.get(), client.get()); + let replacement_a = replacement_a.expect("first replacement should connect"); + let replacement_b = replacement_b.expect("second replacement should reuse client"); + assert_eq!(replacement_a.session_id().as_deref(), Some("session-2")); + assert_eq!(replacement_b.session_id().as_deref(), Some("session-2")); + assert!(Arc::ptr_eq(&replacement_a.inner, &replacement_b.inner)); + + server.await.expect("server task should finish"); + } + #[tokio::test] async fn wake_notifications_do_not_block_other_sessions() { let (client_stdin, server_reader) = duplex(1 << 20); diff --git a/codex-rs/exec-server/src/rpc.rs b/codex-rs/exec-server/src/rpc.rs index e4f2ff554a..981a1c1a80 100644 --- a/codex-rs/exec-server/src/rpc.rs +++ b/codex-rs/exec-server/src/rpc.rs @@ -312,6 +312,10 @@ impl RpcClient { }) } + pub(crate) fn is_disconnected(&self) -> bool { + *self.disconnected_rx.borrow() + } + pub(crate) async fn call(&self, method: &str, params: &P) -> Result where P: Serialize, diff --git a/codex-rs/tui/src/app/snapshots/codex_tui__app__thread_goal_actions__tests__thread_goal_ephemeral_error_message_renders_snapshot.snap b/codex-rs/tui/src/app/snapshots/codex_tui__app__thread_goal_actions__tests__thread_goal_ephemeral_error_message_renders_snapshot.snap new file mode 100644 index 0000000000..d2faad28d1 --- /dev/null +++ b/codex-rs/tui/src/app/snapshots/codex_tui__app__thread_goal_actions__tests__thread_goal_ephemeral_error_message_renders_snapshot.snap @@ -0,0 +1,9 @@ +--- +source: tui/src/app/thread_goal_actions.rs +expression: terminal.backend() +--- + + +■ Goals need a saved session. This session is temporary. +Run `codex` to start a saved session, or `codex resume` / `/resume` to +reopen one. diff --git a/codex-rs/tui/src/app/thread_goal_actions.rs b/codex-rs/tui/src/app/thread_goal_actions.rs index 89b5465a96..5911801755 100644 --- a/codex-rs/tui/src/app/thread_goal_actions.rs +++ b/codex-rs/tui/src/app/thread_goal_actions.rs @@ -8,9 +8,15 @@ use crate::bottom_pane::SelectionViewParams; use crate::bottom_pane::popup_consts::standard_popup_hint_line; use crate::goal_display::goal_status_label; use crate::goal_display::goal_usage_summary; +use codex_app_server_protocol::ThreadGoal; use codex_app_server_protocol::ThreadGoalStatus; use codex_protocol::ThreadId; +const EPHEMERAL_THREAD_GOAL_ERROR_MESSAGE: &str = concat!( + "Goals need a saved session. This session is temporary.\n", + "Run `codex` to start a saved session, or `codex resume` / `/resume` to reopen one.", +); + impl App { pub(super) async fn open_thread_goal_menu( &mut self, @@ -26,7 +32,7 @@ impl App { Ok(response) => response, Err(err) => { self.chat_widget - .add_error_message(format!("Failed to read thread goal: {err}")); + .add_error_message(thread_goal_error_message("read", &err)); return; } }; @@ -91,7 +97,7 @@ impl App { Ok(response) => response, Err(err) => { self.chat_widget - .add_error_message(format!("Failed to read thread goal: {err}")); + .add_error_message(thread_goal_error_message("read", &err)); return; } }; @@ -111,25 +117,30 @@ impl App { objective: String, mode: ThreadGoalSetMode, ) { - if matches!(mode, ThreadGoalSetMode::ConfirmIfExists) { + let mode = if matches!(mode, ThreadGoalSetMode::ConfirmIfExists) { let result = app_server.thread_goal_get(thread_id).await; if self.current_displayed_thread_id() != Some(thread_id) { return; } match result { - Ok(response) if response.goal.is_some() => { - self.show_replace_thread_goal_confirmation(thread_id, objective); - return; - } - Ok(_) => {} + Ok(response) => match response.goal.as_ref() { + Some(goal) if should_confirm_before_replacing_goal(goal) => { + self.show_replace_thread_goal_confirmation(thread_id, objective); + return; + } + Some(_) => ThreadGoalSetMode::ReplaceExisting, + None => mode, + }, Err(err) => { self.chat_widget - .add_error_message(format!("Failed to read thread goal: {err}")); + .add_error_message(thread_goal_error_message("read", &err)); return; } } - } + } else { + mode + }; let replacing_goal = matches!(mode, ThreadGoalSetMode::ReplaceExisting); if replacing_goal { @@ -140,7 +151,7 @@ impl App { return; } self.chat_widget - .add_error_message(format!("Failed to replace thread goal: {err}")); + .add_error_message(thread_goal_error_message("replace", &err)); return; } } @@ -170,7 +181,7 @@ impl App { Err(err) => { let action = if replacing_goal { "replace" } else { "set" }; self.chat_widget - .add_error_message(format!("Failed to {action} thread goal: {err}")); + .add_error_message(thread_goal_error_message(action, &err)); } } } @@ -200,7 +211,7 @@ impl App { ), Err(err) => self .chat_widget - .add_error_message(format!("Failed to update thread goal: {err}")), + .add_error_message(thread_goal_error_message("update", &err)), } } @@ -228,7 +239,7 @@ impl App { } Err(err) => self .chat_widget - .add_error_message(format!("Failed to clear thread goal: {err}")), + .add_error_message(thread_goal_error_message("clear", &err)), } } @@ -274,3 +285,121 @@ impl App { ); } } + +fn thread_goal_error_message(action: &str, err: &color_eyre::Report) -> String { + if is_ephemeral_thread_goal_error(err) { + EPHEMERAL_THREAD_GOAL_ERROR_MESSAGE.to_string() + } else { + format!("Failed to {action} thread goal: {err}") + } +} + +fn is_ephemeral_thread_goal_error(err: &color_eyre::Report) -> bool { + err.chain().any(|cause| { + let message = cause.to_string(); + message.contains("ephemeral thread does not support goals") + || message.contains("thread goals require a persisted thread; this thread is ephemeral") + }) +} + +fn should_confirm_before_replacing_goal(goal: &ThreadGoal) -> bool { + // Completed goals are terminal, so `/goal ` can start a fresh goal + // without asking the user to confirm replacing already-finished work. + match goal.status { + ThreadGoalStatus::Complete => false, + ThreadGoalStatus::Active + | ThreadGoalStatus::Paused + | ThreadGoalStatus::Blocked + | ThreadGoalStatus::UsageLimited + | ThreadGoalStatus::BudgetLimited => true, + } +} + +#[cfg(test)] +mod tests { + use crate::history_cell::HistoryCell; + use pretty_assertions::assert_eq; + use ratatui::layout::Rect; + + use super::*; + + #[test] + fn thread_goal_error_message_explains_temporary_session() { + let err = color_eyre::eyre::eyre!( + "thread/goal/get failed: ephemeral thread does not support goals: thread-1" + ) + .wrap_err("thread/goal/get failed in TUI"); + + assert_eq!( + thread_goal_error_message("read", &err), + EPHEMERAL_THREAD_GOAL_ERROR_MESSAGE + ); + } + + #[test] + fn thread_goal_ephemeral_error_message_renders_snapshot() { + let err = color_eyre::eyre::eyre!( + "thread/goal/get failed: ephemeral thread does not support goals: thread-1" + ) + .wrap_err("thread/goal/get failed in TUI"); + let cell = crate::history_cell::new_error_event(thread_goal_error_message("read", &err)); + let width = 72; + let height = 6; + let backend = crate::test_backend::VT100Backend::new(width, height); + let mut terminal = + crate::custom_terminal::Terminal::with_options(backend).expect("terminal"); + terminal.set_viewport_area(Rect::new(0, height - 1, width, 1)); + + crate::insert_history::insert_history_lines( + &mut terminal, + cell.display_lines(/*width*/ width), + ) + .expect("insert history lines"); + + insta::assert_snapshot!(terminal.backend()); + } + + #[test] + fn thread_goal_error_message_preserves_generic_failure_context() { + let err = + color_eyre::eyre::eyre!("server disappeared").wrap_err("thread/goal/get failed in TUI"); + + assert_eq!( + thread_goal_error_message("read", &err), + "Failed to read thread goal: thread/goal/get failed in TUI" + ); + } + + #[test] + fn completed_goal_does_not_require_replace_confirmation() { + assert!(!should_confirm_before_replacing_goal(&test_goal( + ThreadGoalStatus::Complete + ))); + } + + #[test] + fn unfinished_goals_require_replace_confirmation() { + for status in [ + ThreadGoalStatus::Active, + ThreadGoalStatus::Paused, + ThreadGoalStatus::Blocked, + ThreadGoalStatus::UsageLimited, + ThreadGoalStatus::BudgetLimited, + ] { + assert!(should_confirm_before_replacing_goal(&test_goal(status))); + } + } + + fn test_goal(status: ThreadGoalStatus) -> ThreadGoal { + ThreadGoal { + thread_id: ThreadId::new().to_string(), + objective: "Finish the thing.".to_string(), + status, + token_budget: None, + tokens_used: 0, + time_used_seconds: 0, + created_at: 1_776_272_400, + updated_at: 1_776_272_460, + } + } +}