From 2720658f1de332089632316201009784cb3b0edd Mon Sep 17 00:00:00 2001 From: Channing Conger Date: Tue, 16 Jun 2026 05:10:04 +0000 Subject: [PATCH] code-mode: implement and verify stdio session host --- codex-rs/Cargo.lock | 9 + codex-rs/code-mode-host/Cargo.toml | 11 + codex-rs/code-mode-host/src/convert.rs | 126 +++++ codex-rs/code-mode-host/src/main.rs | 323 ++++++++++- codex-rs/code-mode-host/src/peer.rs | 146 +++++ codex-rs/code-mode-host/src/session.rs | 120 ++++ codex-rs/code-mode-host/tests/host.rs | 684 +++++++++++++++++++++++ codex-rs/code-mode/src/cell_actor/mod.rs | 9 +- codex-rs/code-mode/src/runtime/mod.rs | 38 +- 9 files changed, 1427 insertions(+), 39 deletions(-) create mode 100644 codex-rs/code-mode-host/src/convert.rs create mode 100644 codex-rs/code-mode-host/src/peer.rs create mode 100644 codex-rs/code-mode-host/src/session.rs create mode 100644 codex-rs/code-mode-host/tests/host.rs diff --git a/codex-rs/Cargo.lock b/codex-rs/Cargo.lock index ea4e8357a5..2ef511cdfb 100644 --- a/codex-rs/Cargo.lock +++ b/codex-rs/Cargo.lock @@ -2510,6 +2510,15 @@ dependencies = [ [[package]] name = "codex-code-mode-host" version = "0.0.0" +dependencies = [ + "codex-code-mode", + "codex-code-mode-protocol", + "codex-utils-cargo-bin", + "pretty_assertions", + "serde_json", + "tokio", + "tokio-util", +] [[package]] name = "codex-code-mode-protocol" diff --git a/codex-rs/code-mode-host/Cargo.toml b/codex-rs/code-mode-host/Cargo.toml index a2c384d010..6a34b8bb79 100644 --- a/codex-rs/code-mode-host/Cargo.toml +++ b/codex-rs/code-mode-host/Cargo.toml @@ -10,3 +10,14 @@ path = "src/main.rs" [lints] workspace = true + +[dependencies] +codex-code-mode = { workspace = true } +codex-code-mode-protocol = { workspace = true } +serde_json = { workspace = true } +tokio = { workspace = true, features = ["io-std", "io-util", "macros", "process", "rt-multi-thread", "sync", "time"] } +tokio-util = { workspace = true, features = ["rt"] } + +[dev-dependencies] +codex-utils-cargo-bin = { workspace = true } +pretty_assertions = { workspace = true } diff --git a/codex-rs/code-mode-host/src/convert.rs b/codex-rs/code-mode-host/src/convert.rs new file mode 100644 index 0000000000..93ef31beb9 --- /dev/null +++ b/codex-rs/code-mode-host/src/convert.rs @@ -0,0 +1,126 @@ +use std::time::Duration; + +use codex_code_mode::session_runtime as runtime; +use codex_code_mode_protocol::wire; + +pub(super) fn create_cell_request(request: wire::CreateCellRequest) -> runtime::CreateCellRequest { + runtime::CreateCellRequest { + tool_call_id: request.tool_call_id, + enabled_tools: request + .enabled_tools + .into_iter() + .map(|definition| runtime::ToolDefinition { + name: definition.name, + tool_name: runtime::ToolName { + name: definition.tool_name.name, + namespace: definition.tool_name.namespace, + }, + description: definition.description, + kind: tool_kind(definition.kind), + }) + .collect(), + source: request.source, + } +} + +pub(super) fn observe_mode(mode: wire::ObserveMode) -> runtime::ObserveMode { + match mode { + wire::ObserveMode::YieldAfter { duration_ms } => { + runtime::ObserveMode::YieldAfter(Duration::from_millis(duration_ms)) + } + wire::ObserveMode::PendingFrontier => runtime::ObserveMode::PendingFrontier, + } +} + +pub(super) fn runtime_cell_id(cell_id: &wire::CellId) -> runtime::CellId { + runtime::CellId::new(cell_id.as_str()) +} + +pub(super) fn wire_cell_id(cell_id: &runtime::CellId) -> wire::CellId { + wire::CellId::new(cell_id.as_str()) +} + +pub(super) fn cell_event(event: runtime::CellEvent) -> wire::CellEvent { + match event { + runtime::CellEvent::Yielded { content_items } => wire::CellEvent::Yielded { + content_items: content_items.into_iter().map(output_item).collect(), + }, + runtime::CellEvent::Pending { + content_items, + pending_tool_call_ids, + } => wire::CellEvent::Pending { + content_items: content_items.into_iter().map(output_item).collect(), + pending_tool_call_ids, + }, + runtime::CellEvent::Completed { + content_items, + error_text, + } => wire::CellEvent::Completed { + content_items: content_items.into_iter().map(output_item).collect(), + error_text, + }, + runtime::CellEvent::Terminated { content_items } => wire::CellEvent::Terminated { + content_items: content_items.into_iter().map(output_item).collect(), + }, + } +} + +pub(super) fn runtime_error(error: runtime::Error) -> wire::Error { + match error { + runtime::Error::ShuttingDown => wire::Error::ShuttingDown, + runtime::Error::DuplicateCell(cell_id) => wire::Error::DuplicateCell { + cell_id: wire_cell_id(&cell_id), + }, + runtime::Error::MissingCell(cell_id) => wire::Error::MissingCell { + cell_id: wire_cell_id(&cell_id), + }, + runtime::Error::ClosedCell(cell_id) => wire::Error::ClosedCell { + cell_id: wire_cell_id(&cell_id), + }, + runtime::Error::BusyObserver(cell_id) => wire::Error::BusyObserver { + cell_id: wire_cell_id(&cell_id), + }, + runtime::Error::AlreadyTerminating(cell_id) => wire::Error::AlreadyTerminating { + cell_id: wire_cell_id(&cell_id), + }, + runtime::Error::Runtime(message) => wire::Error::Runtime { message }, + } +} + +pub(super) fn nested_tool_call(invocation: runtime::NestedToolCall) -> wire::NestedToolCall { + wire::NestedToolCall { + cell_id: wire_cell_id(&invocation.cell_id), + runtime_tool_call_id: invocation.runtime_tool_call_id, + tool_name: wire::ToolName { + name: invocation.tool_name.name, + namespace: invocation.tool_name.namespace, + }, + tool_kind: match invocation.tool_kind { + runtime::ToolKind::Function => wire::ToolKind::Function, + runtime::ToolKind::Freeform => wire::ToolKind::Freeform, + }, + input: invocation.input, + } +} + +fn tool_kind(kind: wire::ToolKind) -> runtime::ToolKind { + match kind { + wire::ToolKind::Function => runtime::ToolKind::Function, + wire::ToolKind::Freeform => runtime::ToolKind::Freeform, + } +} + +fn output_item(item: runtime::OutputItem) -> wire::OutputItem { + match item { + runtime::OutputItem::Text { text } => wire::OutputItem::Text { text }, + runtime::OutputItem::Image { image_url, detail } => wire::OutputItem::Image { + image_url, + detail: detail.map(|detail| match detail { + runtime::ImageDetail::Auto => wire::ImageDetail::Auto, + runtime::ImageDetail::Low => wire::ImageDetail::Low, + runtime::ImageDetail::High => wire::ImageDetail::High, + runtime::ImageDetail::Original => wire::ImageDetail::Original, + }), + }, + } +} diff --git a/codex-rs/code-mode-host/src/main.rs b/codex-rs/code-mode-host/src/main.rs index f328e4d9d0..3007a39d0e 100644 --- a/codex-rs/code-mode-host/src/main.rs +++ b/codex-rs/code-mode-host/src/main.rs @@ -1 +1,322 @@ -fn main() {} +use std::collections::HashMap; +use std::future::Future; +use std::sync::Arc; +use std::sync::Mutex as StdMutex; +use std::sync::PoisonError; +use std::sync::atomic::AtomicBool; +use std::sync::atomic::AtomicU64; +use std::sync::atomic::Ordering; + +use codex_code_mode_protocol::wire::ClientMessage; +use codex_code_mode_protocol::wire::Error; +use codex_code_mode_protocol::wire::HostMessage; +use codex_code_mode_protocol::wire::HostRequest; +use codex_code_mode_protocol::wire::HostResponse; +use codex_code_mode_protocol::wire::RequestId; +use codex_code_mode_protocol::wire::SessionId; +use codex_code_mode_protocol::wire::WireResult; +use codex_code_mode_protocol::wire::read_frame; +use codex_code_mode_protocol::wire::write_frame; +use tokio::sync::Mutex; +use tokio::sync::mpsc; +use tokio_util::sync::CancellationToken; + +mod convert; +mod peer; +mod session; + +use peer::HostPeer; +use session::HostSession; + +const IPC_CHANNEL_CAPACITY: usize = 128; + +#[tokio::main(flavor = "current_thread")] +async fn main() { + if let Err(err) = run().await { + eprintln!("codex-code-mode-host failed: {err}"); + std::process::exit(1); + } +} + +async fn run() -> Result<(), String> { + let (outgoing_tx, mut outgoing_rx) = mpsc::channel(IPC_CHANNEL_CAPACITY); + let peer = Arc::new(HostPeer::new(outgoing_tx)); + let state = Arc::new(HostState { + sessions: StdMutex::new(HashMap::new()), + pending_requests: Mutex::new(HashMap::new()), + next_session_id: AtomicU64::new(1), + closing: AtomicBool::new(false), + peer: Arc::clone(&peer), + }); + + let writer = tokio::spawn(async move { + let mut stdout = tokio::io::stdout(); + while let Some(message) = outgoing_rx.recv().await { + write_frame(&mut stdout, &message) + .await + .map_err(|err| err.to_string())?; + } + Ok::<(), String>(()) + }); + + let input_result = async { + let mut stdin = tokio::io::stdin(); + while let Some(message) = read_frame::<_, ClientMessage>(&mut stdin) + .await + .map_err(|err| err.to_string())? + { + match message { + ClientMessage::Request { id, request } => { + state.spawn_request(id, request).await; + } + ClientMessage::CancelRequest { id } => state.cancel_request(id).await, + ClientMessage::CallbackResponse { id, response } => { + peer.complete(id, response).await; + } + } + } + Ok::<(), String>(()) + } + .await; + + peer.disconnect().await; + state.disconnect().await; + drop(state); + drop(peer); + let writer_result = writer.await.map_err(|err| err.to_string())?; + input_result?; + writer_result +} + +struct HostState { + sessions: StdMutex>, + pending_requests: Mutex>, + next_session_id: AtomicU64, + closing: AtomicBool, + peer: Arc, +} + +impl HostState { + async fn spawn_request(self: &Arc, request_id: RequestId, request: HostRequest) { + if self.closing.load(Ordering::Acquire) { + self.send_response(request_id, Err(Error::ShuttingDown), None) + .await; + return; + } + let cancellation_token = CancellationToken::new(); + let duplicate = { + let mut pending_requests = self.pending_requests.lock().await; + if let std::collections::hash_map::Entry::Vacant(e) = pending_requests.entry(request_id) + { + e.insert(cancellation_token.clone()); + false + } else { + true + } + }; + if duplicate { + self.send_response( + request_id, + Err(Error::InvalidRequest { + message: format!("duplicate request id {request_id}"), + }), + None, + ) + .await; + return; + } + + let state = Arc::clone(self); + spawn_critical_request_task(async move { + state + .handle_request(request_id, request, cancellation_token) + .await; + state.pending_requests.lock().await.remove(&request_id); + }); + } + + async fn cancel_request(&self, request_id: RequestId) { + if let Some(cancellation_token) = self.pending_requests.lock().await.get(&request_id) { + cancellation_token.cancel(); + } + } + + async fn handle_request( + &self, + request_id: RequestId, + request: HostRequest, + cancellation_token: CancellationToken, + ) { + match request { + HostRequest::CreateSession => { + if cancellation_token.is_cancelled() || self.closing.load(Ordering::Acquire) { + return; + } + let session_id = self.next_session_id.fetch_add(1, Ordering::Relaxed); + let session = HostSession::new(session_id, Arc::clone(&self.peer)); + let inserted = { + let mut sessions = self.sessions.lock().unwrap_or_else(PoisonError::into_inner); + if self.closing.load(Ordering::Acquire) { + false + } else { + sessions.insert(session_id, session); + true + } + }; + if !inserted { + return; + } + self.send_response( + request_id, + Ok(HostResponse::SessionCreated { session_id }), + Some(&cancellation_token), + ) + .await; + } + HostRequest::ShutdownSession { session_id } => { + let session = self + .sessions + .lock() + .unwrap_or_else(PoisonError::into_inner) + .remove(&session_id); + let result = match session { + Some(session) => session + .shutdown() + .await + .map(|()| HostResponse::SessionShutdown) + .map_err(convert::runtime_error), + None => Err(Error::MissingSession { session_id }), + }; + self.send_response(request_id, result, Some(&cancellation_token)) + .await; + } + HostRequest::CreateCell { + session_id, + request, + } => { + let Some(session) = self.session(session_id) else { + self.send_response( + request_id, + Err(Error::MissingSession { session_id }), + Some(&cancellation_token), + ) + .await; + return; + }; + let result = session + .create_cell(convert::create_cell_request(request)) + .await + .map(|cell_id| HostResponse::CellCreated { + cell_id: convert::wire_cell_id(&cell_id), + }) + .map_err(convert::runtime_error); + self.send_response(request_id, result, Some(&cancellation_token)) + .await; + } + HostRequest::Observe { + session_id, + cell_id, + mode, + } => { + let Some(session) = self.session(session_id) else { + self.send_response( + request_id, + Err(Error::MissingSession { session_id }), + Some(&cancellation_token), + ) + .await; + return; + }; + let runtime_cell_id = convert::runtime_cell_id(&cell_id); + let result = session + .observe(&runtime_cell_id, convert::observe_mode(mode)) + .await + .map(convert::cell_event) + .map(|event| HostResponse::Observed { event }) + .map_err(convert::runtime_error); + self.send_response(request_id, result, Some(&cancellation_token)) + .await; + } + HostRequest::Terminate { + session_id, + cell_id, + } => { + let Some(session) = self.session(session_id) else { + self.send_response( + request_id, + Err(Error::MissingSession { session_id }), + Some(&cancellation_token), + ) + .await; + return; + }; + let runtime_cell_id = convert::runtime_cell_id(&cell_id); + let result = session + .terminate(&runtime_cell_id) + .await + .map(convert::cell_event) + .map(|event| HostResponse::Observed { event }) + .map_err(convert::runtime_error); + self.send_response(request_id, result, Some(&cancellation_token)) + .await; + } + } + } + + fn session(&self, session_id: SessionId) -> Option { + self.sessions + .lock() + .unwrap_or_else(PoisonError::into_inner) + .get(&session_id) + .cloned() + } + + async fn send_response( + &self, + request_id: RequestId, + result: Result, + cancellation_token: Option<&CancellationToken>, + ) { + if cancellation_token.is_some_and(CancellationToken::is_cancelled) { + return; + } + self.peer + .send(HostMessage::Response { + id: request_id, + result: WireResult::from_result(result), + }) + .await; + } + + async fn disconnect(&self) { + self.closing.store(true, Ordering::Release); + for cancellation_token in self.pending_requests.lock().await.values() { + cancellation_token.cancel(); + } + let sessions = self + .sessions + .lock() + .unwrap_or_else(PoisonError::into_inner) + .drain() + .map(|(_, session)| session) + .collect::>(); + for session in sessions { + let _ = session.shutdown().await; + } + while !self.pending_requests.lock().await.is_empty() { + tokio::task::yield_now().await; + } + } +} + +fn spawn_critical_request_task(future: impl Future + Send + 'static) { + let task = tokio::spawn(future); + tokio::spawn(async move { + if let Err(err) = task.await + && err.is_panic() + { + eprintln!("code-mode host request task panicked: {err}"); + std::process::exit(1); + } + }); +} diff --git a/codex-rs/code-mode-host/src/peer.rs b/codex-rs/code-mode-host/src/peer.rs new file mode 100644 index 0000000000..067af3e457 --- /dev/null +++ b/codex-rs/code-mode-host/src/peer.rs @@ -0,0 +1,146 @@ +use std::collections::HashMap; +use std::sync::Arc; +use std::sync::atomic::AtomicU64; +use std::sync::atomic::Ordering; + +use codex_code_mode_protocol::wire::CallbackId; +use codex_code_mode_protocol::wire::CallbackRequest; +use codex_code_mode_protocol::wire::CallbackResponse; +use codex_code_mode_protocol::wire::HostMessage; +use codex_code_mode_protocol::wire::SessionId; +use tokio::sync::Mutex; +use tokio::sync::mpsc; +use tokio::sync::oneshot; +use tokio_util::sync::CancellationToken; + +pub(super) struct HostPeer { + outgoing_tx: mpsc::Sender, + pending_callbacks: Mutex>>, + next_callback_id: AtomicU64, + connection_closed: CancellationToken, +} + +impl HostPeer { + pub(super) fn new(outgoing_tx: mpsc::Sender) -> Self { + Self { + outgoing_tx, + pending_callbacks: Mutex::new(HashMap::new()), + next_callback_id: AtomicU64::new(1), + connection_closed: CancellationToken::new(), + } + } + + pub(super) async fn send(&self, message: HostMessage) -> bool { + tokio::select! { + sent = self.outgoing_tx.send(message) => sent.is_ok(), + _ = self.connection_closed.cancelled() => false, + } + } + + pub(super) fn send_nowait(self: &Arc, message: HostMessage) { + match self.outgoing_tx.try_send(message) { + Ok(()) | Err(mpsc::error::TrySendError::Closed(_)) => {} + Err(mpsc::error::TrySendError::Full(message)) => { + let peer = Arc::clone(self); + tokio::spawn(async move { + peer.send(message).await; + }); + } + } + } + + pub(super) async fn call( + self: &Arc, + session_id: SessionId, + request: CallbackRequest, + cancellation_token: CancellationToken, + ) -> Result { + if self.connection_closed.is_cancelled() { + return Err("code-mode client connection closed".to_string()); + } + let id = self.next_callback_id.fetch_add(1, Ordering::Relaxed); + let (response_tx, response_rx) = oneshot::channel(); + self.pending_callbacks.lock().await.insert(id, response_tx); + let mut pending_callback = PendingCallback::new(Arc::clone(self), id); + + let request_sent = tokio::select! { + sent = self.outgoing_tx.send(HostMessage::CallbackRequest { + id, + session_id, + request, + }) => sent.is_ok(), + _ = cancellation_token.cancelled() => false, + _ = self.connection_closed.cancelled() => false, + }; + if !request_sent { + self.pending_callbacks.lock().await.remove(&id); + pending_callback.disarm(); + return Err(if cancellation_token.is_cancelled() { + "code mode callback cancelled".to_string() + } else { + "code-mode client connection closed".to_string() + }); + } + + tokio::select! { + response = response_rx => { + pending_callback.disarm(); + response.map_err(|_| { + "code-mode client closed before returning callback output".to_string() + }) + }, + _ = cancellation_token.cancelled() => { + if self.pending_callbacks.lock().await.remove(&id).is_some() { + self.send(HostMessage::CancelCallback { id }).await; + } + pending_callback.disarm(); + Err("code mode callback cancelled".to_string()) + } + _ = self.connection_closed.cancelled() => { + self.pending_callbacks.lock().await.remove(&id); + pending_callback.disarm(); + Err("code-mode client connection closed".to_string()) + } + } + } + + pub(super) async fn complete(&self, id: CallbackId, response: CallbackResponse) { + if let Some(response_tx) = self.pending_callbacks.lock().await.remove(&id) { + let _ = response_tx.send(response); + } + } + + pub(super) async fn disconnect(&self) { + self.connection_closed.cancel(); + self.pending_callbacks.lock().await.clear(); + } +} + +struct PendingCallback { + peer: Arc, + id: Option, +} + +impl PendingCallback { + fn new(peer: Arc, id: CallbackId) -> Self { + Self { peer, id: Some(id) } + } + + fn disarm(&mut self) { + self.id = None; + } +} + +impl Drop for PendingCallback { + fn drop(&mut self) { + let Some(id) = self.id.take() else { + return; + }; + let peer = Arc::clone(&self.peer); + tokio::spawn(async move { + if peer.pending_callbacks.lock().await.remove(&id).is_some() { + peer.send(HostMessage::CancelCallback { id }).await; + } + }); + } +} diff --git a/codex-rs/code-mode-host/src/session.rs b/codex-rs/code-mode-host/src/session.rs new file mode 100644 index 0000000000..5c37cdac8a --- /dev/null +++ b/codex-rs/code-mode-host/src/session.rs @@ -0,0 +1,120 @@ +use std::sync::Arc; + +use codex_code_mode::SessionRuntime; +use codex_code_mode::session_runtime as runtime; +use codex_code_mode_protocol::wire::CallbackRequest; +use codex_code_mode_protocol::wire::CallbackResponse; +use codex_code_mode_protocol::wire::HostMessage; +use codex_code_mode_protocol::wire::SessionId; +use serde_json::Value as JsonValue; +use tokio_util::sync::CancellationToken; + +use crate::convert; +use crate::peer::HostPeer; + +#[derive(Clone)] +pub(super) struct HostSession { + runtime: Arc>, +} + +impl HostSession { + pub(super) fn new(session_id: SessionId, peer: Arc) -> Self { + let delegate = Arc::new(RemoteDelegate { session_id, peer }); + Self { + runtime: Arc::new(SessionRuntime::new(delegate)), + } + } + + pub(super) async fn create_cell( + &self, + request: runtime::CreateCellRequest, + ) -> Result { + self.runtime.create_cell(request).await + } + + pub(super) async fn observe( + &self, + cell_id: &runtime::CellId, + mode: runtime::ObserveMode, + ) -> Result { + self.runtime.observe(cell_id, mode).await + } + + pub(super) async fn terminate( + &self, + cell_id: &runtime::CellId, + ) -> Result { + self.runtime.terminate(cell_id).await + } + + pub(super) async fn shutdown(&self) -> Result<(), runtime::Error> { + self.runtime.shutdown().await + } +} + +struct RemoteDelegate { + session_id: SessionId, + peer: Arc, +} + +impl runtime::SessionRuntimeDelegate for RemoteDelegate { + async fn invoke_tool( + &self, + invocation: runtime::NestedToolCall, + cancellation_token: CancellationToken, + ) -> Result { + match self + .peer + .call( + self.session_id, + CallbackRequest::InvokeTool { + invocation: convert::nested_tool_call(invocation), + }, + cancellation_token, + ) + .await? + { + CallbackResponse::ToolResult { result } => Ok(result), + CallbackResponse::ToolError { error_text } => Err(error_text), + CallbackResponse::NotificationDelivered + | CallbackResponse::NotificationError { .. } => { + Err("code-mode client returned an invalid tool response".to_string()) + } + } + } + + async fn notify( + &self, + call_id: String, + cell_id: runtime::CellId, + text: String, + cancellation_token: CancellationToken, + ) -> Result<(), String> { + match self + .peer + .call( + self.session_id, + CallbackRequest::Notify { + call_id, + cell_id: convert::wire_cell_id(&cell_id), + text, + }, + cancellation_token, + ) + .await? + { + CallbackResponse::NotificationDelivered => Ok(()), + CallbackResponse::NotificationError { error_text } => Err(error_text), + CallbackResponse::ToolResult { .. } | CallbackResponse::ToolError { .. } => { + Err("code-mode client returned an invalid notification response".to_string()) + } + } + } + + fn cell_closed(&self, cell_id: &runtime::CellId) { + self.peer.send_nowait(HostMessage::CellClosed { + session_id: self.session_id, + cell_id: convert::wire_cell_id(cell_id), + }); + } +} diff --git a/codex-rs/code-mode-host/tests/host.rs b/codex-rs/code-mode-host/tests/host.rs new file mode 100644 index 0000000000..7c50ddd4d8 --- /dev/null +++ b/codex-rs/code-mode-host/tests/host.rs @@ -0,0 +1,684 @@ +#![allow(clippy::expect_used)] + +use std::process::Stdio; +use std::time::Duration; + +use codex_code_mode_protocol::wire::*; +use pretty_assertions::assert_eq; +use serde_json::json; +use tokio::process::Child; +use tokio::process::ChildStdin; +use tokio::process::Command; +use tokio::sync::mpsc; + +struct HostHarness { + child: Child, + stdin: Option, + messages_rx: mpsc::UnboundedReceiver>, +} + +impl HostHarness { + fn spawn() -> Self { + let mut child = Command::new( + codex_utils_cargo_bin::cargo_bin("codex-code-mode-host").expect("resolve host binary"), + ) + .stdin(Stdio::piped()) + .stdout(Stdio::piped()) + .stderr(Stdio::inherit()) + .kill_on_drop(true) + .spawn() + .expect("spawn host"); + let stdin = child.stdin.take().expect("host stdin"); + let mut stdout = child.stdout.take().expect("host stdout"); + let (messages_tx, messages_rx) = mpsc::unbounded_channel(); + tokio::spawn(async move { + loop { + match read_frame(&mut stdout).await { + Ok(Some(message)) => { + if messages_tx.send(Ok(message)).is_err() { + break; + } + } + Ok(None) => break, + Err(err) => { + let _ = messages_tx.send(Err(err.to_string())); + break; + } + } + } + }); + Self { + child, + stdin: Some(stdin), + messages_rx, + } + } + + async fn send(&mut self, message: ClientMessage) { + write_frame(self.stdin.as_mut().expect("host stdin open"), &message) + .await + .expect("write host message"); + } + + async fn recv(&mut self) -> HostMessage { + tokio::time::timeout(Duration::from_secs(5), self.messages_rx.recv()) + .await + .expect("host message timeout") + .expect("host output closed") + .expect("read host message") + } + + async fn assert_no_message(&mut self, duration: Duration) { + assert!( + tokio::time::timeout(duration, self.messages_rx.recv()) + .await + .is_err(), + "received an unexpected host message" + ); + } + + async fn create_session(&mut self) -> SessionId { + self.send(ClientMessage::Request { + id: 1, + request: HostRequest::CreateSession, + }) + .await; + match self.recv().await { + HostMessage::Response { + id: 1, + result: + WireResult::Ok { + value: HostResponse::SessionCreated { session_id }, + }, + } => session_id, + message => panic!("unexpected create-session response: {message:?}"), + } + } + + async fn shutdown_session(&mut self, session_id: SessionId, request_id: RequestId) { + self.send(ClientMessage::Request { + id: request_id, + request: HostRequest::ShutdownSession { session_id }, + }) + .await; + assert_eq!( + self.recv().await, + HostMessage::Response { + id: request_id, + result: WireResult::Ok { + value: HostResponse::SessionShutdown, + }, + } + ); + } + + async fn finish(mut self) { + self.stdin.take(); + let status = tokio::time::timeout(Duration::from_secs(5), self.child.wait()) + .await + .expect("host exit timeout") + .expect("wait for host"); + assert!(status.success(), "host exited with {status}"); + } +} + +fn create_cell_request(source: &str) -> CreateCellRequest { + CreateCellRequest { + tool_call_id: "call-1".to_string(), + enabled_tools: Vec::new(), + source: source.to_string(), + } +} + +async fn create_cell( + host: &mut HostHarness, + session_id: SessionId, + request_id: RequestId, + request: CreateCellRequest, +) -> CellId { + host.send(ClientMessage::Request { + id: request_id, + request: HostRequest::CreateCell { + session_id, + request, + }, + }) + .await; + match host.recv().await { + HostMessage::Response { + id, + result: + WireResult::Ok { + value: HostResponse::CellCreated { cell_id }, + }, + } if id == request_id => cell_id, + message => panic!("unexpected create-cell response: {message:?}"), + } +} + +async fn observe_cell( + host: &mut HostHarness, + session_id: SessionId, + request_id: RequestId, + cell_id: CellId, + mode: ObserveMode, +) { + host.send(ClientMessage::Request { + id: request_id, + request: HostRequest::Observe { + session_id, + cell_id, + mode, + }, + }) + .await; +} + +async fn recv_observation(host: &mut HostHarness, request_id: RequestId) -> CellEvent { + match host.recv().await { + HostMessage::Response { + id, + result: + WireResult::Ok { + value: HostResponse::Observed { event }, + }, + } if id == request_id => event, + message => panic!("unexpected observation response: {message:?}"), + } +} + +async fn recv_terminal_observation( + host: &mut HostHarness, + session_id: SessionId, + request_id: RequestId, + cell_id: &CellId, +) -> CellEvent { + let mut event = None; + let mut cell_closed = false; + for _ in 0..2 { + match host.recv().await { + HostMessage::Response { + id, + result: + WireResult::Ok { + value: HostResponse::Observed { event: response }, + }, + } if id == request_id => event = Some(response), + HostMessage::CellClosed { + session_id: closed_session_id, + cell_id: closed_cell_id, + } => { + assert_eq!(closed_session_id, session_id); + assert_eq!(&closed_cell_id, cell_id); + cell_closed = true; + } + message => panic!("unexpected terminal observation message: {message:?}"), + } + } + assert!(cell_closed); + event.expect("terminal observation response") +} + +#[tokio::test] +async fn yields_resumes_and_closes_over_stdio() { + let mut host = HostHarness::spawn(); + let session_id = host.create_session().await; + let cell_id = create_cell( + &mut host, + session_id, + 2, + create_cell_request(concat!( + "await new Promise(resolve => setTimeout(resolve, 100));", + r#"text("before"); yield_control(); text("after");"#, + )), + ) + .await; + observe_cell( + &mut host, + session_id, + 3, + cell_id.clone(), + ObserveMode::YieldAfter { + duration_ms: 60_000, + }, + ) + .await; + let mut yielded = None; + let mut saw_cell_closed = false; + for _ in 0..2 { + match host.recv().await { + HostMessage::Response { + id: 3, + result: + WireResult::Ok { + value: HostResponse::Observed { event }, + }, + } => yielded = Some(event), + HostMessage::CellClosed { + session_id: closed_session_id, + cell_id: closed_cell_id, + } => { + assert_eq!(closed_session_id, session_id); + assert_eq!(closed_cell_id, cell_id); + saw_cell_closed = true; + } + message => panic!("unexpected initial observation message: {message:?}"), + } + } + assert_eq!( + yielded, + Some(CellEvent::Yielded { + content_items: vec![OutputItem::Text { + text: "before".to_string(), + }], + }) + ); + assert!(saw_cell_closed); + observe_cell( + &mut host, + session_id, + 4, + cell_id.clone(), + ObserveMode::YieldAfter { duration_ms: 1 }, + ) + .await; + assert_eq!( + recv_observation(&mut host, 4).await, + CellEvent::Completed { + content_items: vec![OutputItem::Text { + text: "after".to_string(), + }], + error_text: None, + } + ); + host.shutdown_session(session_id, 5).await; + host.finish().await; +} + +#[tokio::test] +async fn termination_resolves_an_active_observer() { + let mut host = HostHarness::spawn(); + let session_id = host.create_session().await; + let cell_id = create_cell( + &mut host, + session_id, + 2, + create_cell_request("await new Promise(() => {});"), + ) + .await; + observe_cell( + &mut host, + session_id, + 3, + cell_id.clone(), + ObserveMode::YieldAfter { + duration_ms: 60_000, + }, + ) + .await; + host.assert_no_message(Duration::from_millis(100)).await; + + host.send(ClientMessage::Request { + id: 4, + request: HostRequest::Terminate { + session_id, + cell_id: cell_id.clone(), + }, + }) + .await; + + let mut saw_observation_response = false; + let mut saw_termination_response = false; + let mut saw_cell_closed = false; + for _ in 0..3 { + match host.recv().await { + HostMessage::Response { + id: 3, + result: + WireResult::Ok { + value: + HostResponse::Observed { + event: CellEvent::Terminated { content_items }, + }, + }, + } => { + assert!(content_items.is_empty()); + saw_observation_response = true; + } + HostMessage::Response { + id: 4, + result: + WireResult::Ok { + value: + HostResponse::Observed { + event: CellEvent::Terminated { content_items }, + }, + }, + } => { + assert!(content_items.is_empty()); + saw_termination_response = true; + } + HostMessage::CellClosed { + session_id: closed_session_id, + cell_id: closed_cell_id, + } => { + assert_eq!(closed_session_id, session_id); + assert_eq!(closed_cell_id, cell_id); + saw_cell_closed = true; + } + message => panic!("unexpected termination message: {message:?}"), + } + } + assert!(saw_observation_response); + assert!(saw_termination_response); + assert!(saw_cell_closed); + + host.shutdown_session(session_id, 5).await; + host.finish().await; +} + +#[tokio::test] +async fn forwards_tool_and_notification_callbacks() { + let mut host = HostHarness::spawn(); + let session_id = host.create_session().await; + let mut request = create_cell_request( + r#" +notify("note"); +const value = await tools.echo({ value: "input" }); +text(value.value); +"#, + ); + request.enabled_tools.push(ToolDefinition { + name: "echo".to_string(), + tool_name: ToolName { + name: "echo".to_string(), + namespace: None, + }, + description: String::new(), + kind: ToolKind::Function, + }); + let cell_id = create_cell(&mut host, session_id, 2, request).await; + observe_cell( + &mut host, + session_id, + 3, + cell_id.clone(), + ObserveMode::YieldAfter { + duration_ms: 60_000, + }, + ) + .await; + + for _ in 0..2 { + match host.recv().await { + HostMessage::CallbackRequest { + id, + session_id: callback_session_id, + request: CallbackRequest::Notify { text, .. }, + } => { + assert_eq!(callback_session_id, session_id); + assert_eq!(text, "note"); + host.send(ClientMessage::CallbackResponse { + id, + response: CallbackResponse::NotificationDelivered, + }) + .await; + } + HostMessage::CallbackRequest { + id, + session_id: callback_session_id, + request: CallbackRequest::InvokeTool { invocation }, + } => { + assert_eq!(callback_session_id, session_id); + assert_eq!(invocation.input, Some(json!({"value": "input"}))); + host.send(ClientMessage::CallbackResponse { + id, + response: CallbackResponse::ToolResult { + result: json!({"value": "output"}), + }, + }) + .await; + } + message => panic!("unexpected callback message: {message:?}"), + } + } + + assert_eq!( + recv_terminal_observation(&mut host, session_id, 3, &cell_id).await, + CellEvent::Completed { + content_items: vec![OutputItem::Text { + text: "output".to_string(), + }], + error_text: None, + } + ); + host.shutdown_session(session_id, 4).await; + host.finish().await; +} + +#[tokio::test] +async fn pending_frontier_rejects_a_second_observer_and_termination_preempts_the_first() { + let mut host = HostHarness::spawn(); + let session_id = host.create_session().await; + let cell_id = create_cell( + &mut host, + session_id, + 2, + create_cell_request("await new Promise(() => {});"), + ) + .await; + observe_cell( + &mut host, + session_id, + 3, + cell_id.clone(), + ObserveMode::PendingFrontier, + ) + .await; + assert_eq!( + recv_observation(&mut host, 3).await, + CellEvent::Pending { + content_items: Vec::new(), + pending_tool_call_ids: Vec::new(), + } + ); + + host.send(ClientMessage::Request { + id: 4, + request: HostRequest::Observe { + session_id, + cell_id: cell_id.clone(), + mode: ObserveMode::YieldAfter { + duration_ms: 60_000, + }, + }, + }) + .await; + host.assert_no_message(Duration::from_millis(100)).await; + host.send(ClientMessage::Request { + id: 5, + request: HostRequest::Observe { + session_id, + cell_id: cell_id.clone(), + mode: ObserveMode::YieldAfter { + duration_ms: 60_000, + }, + }, + }) + .await; + assert_eq!( + host.recv().await, + HostMessage::Response { + id: 5, + result: WireResult::Err { + error: Error::BusyObserver { + cell_id: cell_id.clone(), + }, + }, + } + ); + + host.send(ClientMessage::Request { + id: 6, + request: HostRequest::Terminate { + session_id, + cell_id: cell_id.clone(), + }, + }) + .await; + let terminated = HostResponse::Observed { + event: CellEvent::Terminated { + content_items: Vec::new(), + }, + }; + let mut response_ids = Vec::new(); + let mut saw_cell_closed = false; + for _ in 0..3 { + match host.recv().await { + HostMessage::Response { + id, + result: WireResult::Ok { value }, + } => { + assert_eq!(value, terminated); + response_ids.push(id); + } + HostMessage::CellClosed { + session_id: closed_session_id, + cell_id: closed_cell_id, + } => { + assert_eq!(closed_session_id, session_id); + assert_eq!(closed_cell_id, cell_id); + saw_cell_closed = true; + } + message => panic!("unexpected termination response: {message:?}"), + } + } + response_ids.sort_unstable(); + assert_eq!(response_ids, vec![4, 6]); + assert!(saw_cell_closed); + host.shutdown_session(session_id, 7).await; + host.finish().await; +} + +#[tokio::test] +async fn shutdown_cancels_callbacks_before_acknowledging_the_session() { + let mut host = HostHarness::spawn(); + let session_id = host.create_session().await; + let mut request = create_cell_request("await tools.echo({});"); + request.enabled_tools.push(ToolDefinition { + name: "echo".to_string(), + tool_name: ToolName { + name: "echo".to_string(), + namespace: None, + }, + description: String::new(), + kind: ToolKind::Function, + }); + let cell_id = create_cell(&mut host, session_id, 2, request).await; + observe_cell( + &mut host, + session_id, + 3, + cell_id.clone(), + ObserveMode::YieldAfter { + duration_ms: 60_000, + }, + ) + .await; + let callback_id = match host.recv().await { + HostMessage::CallbackRequest { + id, + request: CallbackRequest::InvokeTool { .. }, + .. + } => id, + message => panic!("unexpected callback message: {message:?}"), + }; + + host.send(ClientMessage::Request { + id: 4, + request: HostRequest::ShutdownSession { session_id }, + }) + .await; + let mut saw_callback_cancellation = false; + let mut saw_terminated_response = false; + let mut saw_shutdown_response = false; + let mut saw_cell_closed = false; + for _ in 0..4 { + match host.recv().await { + HostMessage::CancelCallback { id } => { + assert_eq!(id, callback_id); + saw_callback_cancellation = true; + } + HostMessage::Response { + id: 3, + result: + WireResult::Ok { + value: + HostResponse::Observed { + event: CellEvent::Terminated { content_items }, + }, + }, + } => { + assert!(content_items.is_empty()); + saw_terminated_response = true; + } + HostMessage::Response { + id: 4, + result: + WireResult::Ok { + value: HostResponse::SessionShutdown, + }, + } => saw_shutdown_response = true, + HostMessage::CellClosed { + session_id: closed_session_id, + cell_id: closed_cell_id, + } => { + assert_eq!(closed_session_id, session_id); + assert_eq!(closed_cell_id, cell_id); + saw_cell_closed = true; + } + message => panic!("unexpected shutdown message: {message:?}"), + } + } + assert!(saw_callback_cancellation); + assert!(saw_terminated_response); + assert!(saw_shutdown_response); + assert!(saw_cell_closed); + host.finish().await; +} + +#[tokio::test] +async fn disconnect_exits_with_a_pending_callback() { + let mut host = HostHarness::spawn(); + let session_id = host.create_session().await; + let mut request = create_cell_request("await tools.echo({});"); + request.enabled_tools.push(ToolDefinition { + name: "echo".to_string(), + tool_name: ToolName { + name: "echo".to_string(), + namespace: None, + }, + description: String::new(), + kind: ToolKind::Function, + }); + let cell_id = create_cell(&mut host, session_id, 2, request).await; + observe_cell( + &mut host, + session_id, + 3, + cell_id, + ObserveMode::YieldAfter { + duration_ms: 60_000, + }, + ) + .await; + assert!(matches!( + host.recv().await, + HostMessage::CallbackRequest { + request: CallbackRequest::InvokeTool { .. }, + .. + } + )); + + host.finish().await; +} diff --git a/codex-rs/code-mode/src/cell_actor/mod.rs b/codex-rs/code-mode/src/cell_actor/mod.rs index 137affe845..cbee3c60f8 100644 --- a/codex-rs/code-mode/src/cell_actor/mod.rs +++ b/codex-rs/code-mode/src/cell_actor/mod.rs @@ -30,7 +30,6 @@ pub(crate) use self::types::CellToolCall; pub(crate) use self::types::CompletionCommit; use self::types::CompletionDelivery; use self::types::ObservationDelivery; -use crate::runtime::PendingRuntimeMode; use crate::runtime::RuntimeCommand; use crate::runtime::RuntimeControlCommand; use crate::runtime::RuntimeEvent; @@ -52,12 +51,8 @@ impl CellActor { ) -> Result<(CellHandle, impl Future + Send + 'static), String> { let (event_tx, event_rx) = mpsc::unbounded_channel(); let (command_tx, command_rx) = mpsc::unbounded_channel(); - let (runtime_tx, runtime_control_tx, runtime_terminate_handle) = spawn_runtime( - stored_values, - runtime_request(request), - event_tx, - PendingRuntimeMode::PauseUntilResumed, - )?; + let (runtime_tx, runtime_control_tx, runtime_terminate_handle) = + spawn_runtime(stored_values, runtime_request(request), event_tx)?; let handle = CellHandle::new(command_tx, Arc::clone(&cell_state)); let task = run_cell( host, diff --git a/codex-rs/code-mode/src/runtime/mod.rs b/codex-rs/code-mode/src/runtime/mod.rs index ce6a27c9ae..ca2536f402 100644 --- a/codex-rs/code-mode/src/runtime/mod.rs +++ b/codex-rs/code-mode/src/runtime/mod.rs @@ -29,13 +29,6 @@ pub(crate) enum RuntimeCommand { Terminate, } -#[derive(Clone, Copy, Debug, PartialEq)] -pub(crate) enum PendingRuntimeMode { - #[cfg(test)] - Continue, - PauseUntilResumed, -} - #[derive(Debug)] pub(crate) enum RuntimeControlCommand { Continue, @@ -69,7 +62,6 @@ pub(crate) fn spawn_runtime( stored_values: HashMap, request: CreateCellRequest, event_tx: mpsc::UnboundedSender, - pending_mode: PendingRuntimeMode, ) -> Result< ( std_mpsc::Sender, @@ -102,7 +94,6 @@ pub(crate) fn spawn_runtime( event_tx, command_rx, control_rx, - pending_mode, isolate_handle_tx, runtime_command_tx, ); @@ -165,7 +156,6 @@ fn run_runtime( event_tx: mpsc::UnboundedSender, command_rx: std_mpsc::Receiver, control_rx: std_mpsc::Receiver, - pending_mode: PendingRuntimeMode, isolate_handle_tx: std_mpsc::SyncSender, runtime_command_tx: std_mpsc::Sender, ) { @@ -221,9 +211,7 @@ fn run_runtime( } let mut pending_promise = pending_promise; - while let Some(command) = - next_runtime_command(&event_tx, &command_rx, &control_rx, pending_mode) - { + while let Some(command) = next_runtime_command(&event_tx, &command_rx, &control_rx) { match command { RuntimeCommand::Terminate => break, RuntimeCommand::ToolResponse { id, result } => { @@ -276,7 +264,6 @@ fn next_runtime_command( event_tx: &mpsc::UnboundedSender, command_rx: &std_mpsc::Receiver, control_rx: &std_mpsc::Receiver, - pending_mode: PendingRuntimeMode, ) -> Option { loop { match command_rx.try_recv() { @@ -286,14 +273,10 @@ fn next_runtime_command( } let _ = event_tx.send(RuntimeEvent::Pending); - match pending_mode { - #[cfg(test)] - PendingRuntimeMode::Continue => return command_rx.recv().ok(), - PendingRuntimeMode::PauseUntilResumed => match control_rx.recv().ok()? { - RuntimeControlCommand::Continue => return command_rx.recv().ok(), - RuntimeControlCommand::Resume => continue, - RuntimeControlCommand::Terminate => return Some(RuntimeCommand::Terminate), - }, + match control_rx.recv().ok()? { + RuntimeControlCommand::Continue => return command_rx.recv().ok(), + RuntimeControlCommand::Resume => continue, + RuntimeControlCommand::Terminate => return Some(RuntimeCommand::Terminate), } } } @@ -331,7 +314,6 @@ mod tests { use tokio::sync::mpsc; use super::CreateCellRequest; - use super::PendingRuntimeMode; use super::RuntimeCommand; use super::RuntimeControlCommand; use super::RuntimeEvent; @@ -349,13 +331,8 @@ mod tests { #[tokio::test] async fn terminate_execution_stops_cpu_bound_module() { let (event_tx, mut event_rx) = mpsc::unbounded_channel(); - let (_runtime_tx, _runtime_control_tx, runtime_terminate_handle) = spawn_runtime( - HashMap::new(), - execute_request("while (true) {}"), - event_tx, - PendingRuntimeMode::Continue, - ) - .unwrap(); + let (_runtime_tx, _runtime_control_tx, runtime_terminate_handle) = + spawn_runtime(HashMap::new(), execute_request("while (true) {}"), event_tx).unwrap(); let started_event = tokio::time::timeout(Duration::from_secs(1), event_rx.recv()) .await @@ -395,7 +372,6 @@ await new Promise(() => {}); "#, ), event_tx, - PendingRuntimeMode::PauseUntilResumed, ) .unwrap();