diff --git a/codex-rs/Cargo.lock b/codex-rs/Cargo.lock index 139ae3985f..bb5c8a6ba5 100644 --- a/codex-rs/Cargo.lock +++ b/codex-rs/Cargo.lock @@ -2490,6 +2490,16 @@ dependencies = [ "v8", ] +[[package]] +name = "codex-code-mode-client" +version = "0.0.0" +dependencies = [ + "codex-code-mode-protocol", + "tokio", + "tokio-util", + "tracing", +] + [[package]] name = "codex-code-mode-host" version = "0.0.0" diff --git a/codex-rs/Cargo.toml b/codex-rs/Cargo.toml index 8ce5126823..9b6e4e95a9 100644 --- a/codex-rs/Cargo.toml +++ b/codex-rs/Cargo.toml @@ -21,6 +21,7 @@ members = [ "install-context", "codex-backend-openapi-models", "code-mode", + "code-mode-client", "code-mode-host", "code-mode-protocol", "codex-home", @@ -160,6 +161,7 @@ codex-cloud-config = { path = "cloud-config" } codex-cloud-tasks-client = { path = "cloud-tasks-client" } codex-cloud-tasks-mock-client = { path = "cloud-tasks-mock-client" } codex-code-mode = { path = "code-mode" } +codex-code-mode-client = { path = "code-mode-client" } codex-code-mode-protocol = { path = "code-mode-protocol" } codex-home = { path = "codex-home" } codex-config = { path = "config" } diff --git a/codex-rs/code-mode-client/BUILD.bazel b/codex-rs/code-mode-client/BUILD.bazel new file mode 100644 index 0000000000..071598cc2c --- /dev/null +++ b/codex-rs/code-mode-client/BUILD.bazel @@ -0,0 +1,6 @@ +load("//:defs.bzl", "codex_rust_crate") + +codex_rust_crate( + name = "code-mode-client", + crate_name = "codex_code_mode_client", +) diff --git a/codex-rs/code-mode-client/Cargo.toml b/codex-rs/code-mode-client/Cargo.toml new file mode 100644 index 0000000000..585c6ca261 --- /dev/null +++ b/codex-rs/code-mode-client/Cargo.toml @@ -0,0 +1,19 @@ +[package] +edition.workspace = true +license.workspace = true +name = "codex-code-mode-client" +version.workspace = true + +[lib] +doctest = false +name = "codex_code_mode_client" +path = "src/lib.rs" + +[lints] +workspace = true + +[dependencies] +codex-code-mode-protocol = { workspace = true } +tokio = { workspace = true, features = ["io-std", "io-util", "process", "rt", "sync"] } +tokio-util = { workspace = true, features = ["rt"] } +tracing = { workspace = true } diff --git a/codex-rs/code-mode-client/src/connection.rs b/codex-rs/code-mode-client/src/connection.rs new file mode 100644 index 0000000000..4419c3be4f --- /dev/null +++ b/codex-rs/code-mode-client/src/connection.rs @@ -0,0 +1,416 @@ +use std::collections::HashMap; +use std::process::Stdio; +use std::sync::Arc; +use std::sync::Weak; +use std::sync::atomic::AtomicBool; +use std::sync::atomic::AtomicU64; +use std::sync::atomic::Ordering; + +use codex_code_mode_protocol::CodeModeSessionDelegate; +use codex_code_mode_protocol::ExecuteRequest; +use codex_code_mode_protocol::StartedCell; +use codex_code_mode_protocol::wire::ClientMessage; +use codex_code_mode_protocol::wire::DelegateRequest; +use codex_code_mode_protocol::wire::DelegateRequestId; +use codex_code_mode_protocol::wire::DelegateResponse; +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::read_frame; +use codex_code_mode_protocol::wire::write_frame; +use tokio::io::AsyncBufReadExt; +use tokio::io::BufReader; +use tokio::process::Command; +use tokio::sync::Mutex; +use tokio::sync::mpsc; +use tokio::sync::oneshot; +use tokio_util::sync::CancellationToken; +use tracing::debug; +use tracing::warn; + +use crate::CodeModeHostCommand; + +const IPC_CHANNEL_CAPACITY: usize = 128; + +pub(super) struct Connection { + state: Arc, + cancellation: CancellationToken, +} + +impl Connection { + pub(super) async fn spawn(command: &CodeModeHostCommand) -> Result { + let mut child = host_process(command) + .stdin(Stdio::piped()) + .stdout(Stdio::piped()) + .stderr(Stdio::piped()) + .kill_on_drop(true) + .spawn() + .map_err(|err| { + format!( + "failed to spawn code-mode host {}: {err}", + command.program.display() + ) + })?; + let stdin = child + .stdin + .take() + .ok_or_else(|| "spawned code-mode host has no stdin".to_string())?; + let stdout = child + .stdout + .take() + .ok_or_else(|| "spawned code-mode host has no stdout".to_string())?; + let stderr = child.stderr.take(); + let (outgoing_tx, mut outgoing_rx) = mpsc::channel(IPC_CHANNEL_CAPACITY); + let cancellation = CancellationToken::new(); + let state = Arc::new(ConnectionState::new(outgoing_tx)); + + let writer_state = Arc::downgrade(&state); + let writer_cancellation = cancellation.clone(); + tokio::spawn(async move { + let mut stdin = stdin; + loop { + tokio::select! { + _ = writer_cancellation.cancelled() => break, + message = outgoing_rx.recv() => { + let Some(message) = message else { + break; + }; + if let Err(err) = write_frame(&mut stdin, &message).await { + warn!("failed to write code-mode host message: {err}"); + if let Some(state) = writer_state.upgrade() { + state + .fail(format!("failed to write code-mode host message: {err}")) + .await; + } + break; + } + } + } + } + }); + + let reader_state = Arc::downgrade(&state); + let reader_cancellation = cancellation.clone(); + tokio::spawn(async move { + drive_reader(stdout, reader_state, reader_cancellation).await; + }); + + if let Some(stderr) = stderr { + tokio::spawn(async move { + let mut lines = BufReader::new(stderr).lines(); + loop { + match lines.next_line().await { + Ok(Some(line)) => debug!("code-mode host stderr: {line}"), + Ok(None) => break, + Err(err) => { + warn!("failed to read code-mode host stderr: {err}"); + break; + } + } + } + }); + } + + let supervisor_state = Arc::downgrade(&state); + let supervisor_cancellation = cancellation.clone(); + tokio::spawn(async move { + tokio::select! { + result = child.wait() => { + let reason = match result { + Ok(status) => format!("code-mode host exited with status {status}"), + Err(err) => format!("failed waiting for code-mode host: {err}"), + }; + if let Some(state) = supervisor_state.upgrade() { + state.fail(reason).await; + } + } + _ = supervisor_cancellation.cancelled() => { + let _ = child.start_kill(); + let _ = child.wait().await; + } + } + }); + + Ok(Self { + state, + cancellation, + }) + } + + pub(super) fn is_alive(&self) -> bool { + self.state.alive.load(Ordering::Acquire) + } + + pub(super) async fn request(&self, request: HostRequest) -> Result { + let id = self.state.next_request_id.fetch_add(1, Ordering::Relaxed); + let (response_tx, response_rx) = oneshot::channel(); + self.state.pending.lock().await.insert(id, response_tx); + if let Err(err) = self + .state + .send(ClientMessage::Request { id, request }) + .await + { + self.state.pending.lock().await.remove(&id); + return Err(err); + } + response_rx + .await + .map_err(|_| self.state.failure_message())? + } + + pub(super) async fn execute( + &self, + session_id: SessionId, + request: ExecuteRequest, + ) -> Result { + let id = self.state.next_request_id.fetch_add(1, Ordering::Relaxed); + let (response_tx, response_rx) = oneshot::channel(); + let (initial_tx, initial_rx) = oneshot::channel(); + self.state.pending.lock().await.insert(id, response_tx); + self.state + .initial_responses + .lock() + .await + .insert(id, initial_tx); + if let Err(err) = self + .state + .send(ClientMessage::Request { + id, + request: HostRequest::Execute { + session_id, + request, + }, + }) + .await + { + self.state.pending.lock().await.remove(&id); + self.state.initial_responses.lock().await.remove(&id); + return Err(err); + } + let response = match response_rx.await { + Ok(Ok(response)) => response, + Ok(Err(err)) => { + self.state.initial_responses.lock().await.remove(&id); + return Err(err); + } + Err(_) => { + self.state.initial_responses.lock().await.remove(&id); + return Err(self.state.failure_message()); + } + }; + match response { + HostResponse::ExecutionStarted { cell_id } => { + Ok(StartedCell::from_result_receiver(cell_id, initial_rx)) + } + _ => { + self.state.initial_responses.lock().await.remove(&id); + Err("code-mode host returned an invalid execute response".to_string()) + } + } + } + + pub(super) async fn register_delegate( + &self, + session_id: SessionId, + delegate: Arc, + ) { + self.state + .delegates + .lock() + .await + .insert(session_id, delegate); + } + + pub(super) async fn remove_delegate(&self, session_id: SessionId) { + self.state.delegates.lock().await.remove(&session_id); + } +} + +impl Drop for Connection { + fn drop(&mut self) { + self.cancellation.cancel(); + } +} + +struct ConnectionState { + outgoing_tx: mpsc::Sender, + pending: Mutex>>>, + initial_responses: Mutex< + HashMap< + RequestId, + oneshot::Sender>, + >, + >, + delegates: Mutex>>, + delegate_cancellations: Mutex>, + next_request_id: AtomicU64, + alive: AtomicBool, + failure: std::sync::Mutex>, +} + +impl ConnectionState { + fn new(outgoing_tx: mpsc::Sender) -> Self { + Self { + outgoing_tx, + pending: Mutex::new(HashMap::new()), + initial_responses: Mutex::new(HashMap::new()), + delegates: Mutex::new(HashMap::new()), + delegate_cancellations: Mutex::new(HashMap::new()), + next_request_id: AtomicU64::new(1), + alive: AtomicBool::new(true), + failure: std::sync::Mutex::new(None), + } + } + + async fn send(&self, message: ClientMessage) -> Result<(), String> { + if !self.alive.load(Ordering::Acquire) { + return Err(self.failure_message()); + } + self.outgoing_tx + .send(message) + .await + .map_err(|_| self.failure_message()) + } + + fn failure_message(&self) -> String { + self.failure + .lock() + .ok() + .and_then(|failure| failure.clone()) + .unwrap_or_else(|| "code-mode host connection closed".to_string()) + } + + async fn fail(&self, reason: String) { + if !self.alive.swap(false, Ordering::AcqRel) { + return; + } + warn!(%reason, "code-mode host connection failed"); + if let Ok(mut failure) = self.failure.lock() { + *failure = Some(reason.clone()); + } + for (_, sender) in self.pending.lock().await.drain() { + let _ = sender.send(Err(reason.clone())); + } + self.initial_responses.lock().await.clear(); + for (_, cancellation) in self.delegate_cancellations.lock().await.drain() { + cancellation.cancel(); + } + } +} + +async fn drive_reader( + mut stdout: tokio::process::ChildStdout, + state: Weak, + cancellation: CancellationToken, +) { + loop { + let message = tokio::select! { + _ = cancellation.cancelled() => return, + result = read_frame::<_, HostMessage>(&mut stdout) => result, + }; + match message { + Ok(Some(message)) => { + let Some(state) = state.upgrade() else { + return; + }; + handle_host_message(state, message).await; + } + Ok(None) => { + if let Some(state) = state.upgrade() { + state + .fail("code-mode host closed its stdout".to_string()) + .await; + } + return; + } + Err(err) => { + if let Some(state) = state.upgrade() { + state + .fail(format!("failed to read code-mode host message: {err}")) + .await; + } + return; + } + } + } +} + +async fn handle_host_message(state: Arc, message: HostMessage) { + match message { + HostMessage::Response { id, response } => { + if let Some(sender) = state.pending.lock().await.remove(&id) { + let _ = sender.send(response); + } + } + HostMessage::InitialResponse { id, response } => { + if let Some(sender) = state.initial_responses.lock().await.remove(&id) { + let _ = sender.send(response); + } + } + HostMessage::DelegateRequest { + id, + session_id, + request, + } => { + let delegate = state.delegates.lock().await.get(&session_id).cloned(); + let Some(delegate) = delegate else { + let _ = state + .send(ClientMessage::DelegateResponse { + id, + response: Err(format!("unknown code-mode session {session_id}")), + }) + .await; + return; + }; + let cancellation = CancellationToken::new(); + state + .delegate_cancellations + .lock() + .await + .insert(id, cancellation.clone()); + tokio::spawn(async move { + let response = match request { + DelegateRequest::InvokeTool(invocation) => delegate + .invoke_tool(invocation, cancellation) + .await + .map(DelegateResponse::ToolResult), + DelegateRequest::Notify { + call_id, + cell_id, + text, + } => delegate + .notify(call_id, cell_id, text, cancellation) + .await + .map(|()| DelegateResponse::NotificationDelivered), + }; + state.delegate_cancellations.lock().await.remove(&id); + let _ = state + .send(ClientMessage::DelegateResponse { id, response }) + .await; + }); + } + HostMessage::CancelDelegateRequest { id } => { + if let Some(cancellation) = state.delegate_cancellations.lock().await.remove(&id) { + cancellation.cancel(); + } + } + HostMessage::CellClosed { + session_id, + cell_id, + } => { + if let Some(delegate) = state.delegates.lock().await.get(&session_id).cloned() { + delegate.cell_closed(&cell_id); + } + } + } +} + +fn host_process(command: &CodeModeHostCommand) -> Command { + let mut process = Command::new(&command.program); + process.args(&command.args); + #[cfg(unix)] + process.process_group(0); + process +} diff --git a/codex-rs/code-mode-client/src/lib.rs b/codex-rs/code-mode-client/src/lib.rs new file mode 100644 index 0000000000..f08cc803f4 --- /dev/null +++ b/codex-rs/code-mode-client/src/lib.rs @@ -0,0 +1,235 @@ +use std::path::PathBuf; +use std::sync::Arc; +use std::sync::atomic::AtomicBool; +use std::sync::atomic::Ordering; + +use codex_code_mode_protocol::CellId; +use codex_code_mode_protocol::CodeModeSession; +use codex_code_mode_protocol::CodeModeSessionDelegate; +use codex_code_mode_protocol::CodeModeSessionProvider; +use codex_code_mode_protocol::CodeModeSessionProviderFuture; +use codex_code_mode_protocol::CodeModeSessionResultFuture; +use codex_code_mode_protocol::ExecuteRequest; +use codex_code_mode_protocol::StartedCell; +use codex_code_mode_protocol::WaitOutcome; +use codex_code_mode_protocol::WaitRequest; +use codex_code_mode_protocol::wire::HostRequest; +use codex_code_mode_protocol::wire::HostResponse; +use codex_code_mode_protocol::wire::SessionId; +use tokio::sync::Semaphore; + +const CODE_MODE_HOST_PATH_ENV: &str = "CODEX_CODE_MODE_HOST_PATH"; + +mod connection; + +use connection::Connection; + +#[derive(Clone, Debug)] +pub struct CodeModeHostCommand { + pub program: PathBuf, + pub args: Vec, +} + +impl Default for CodeModeHostCommand { + fn default() -> Self { + Self { + program: default_host_program(), + args: Vec::new(), + } + } +} + +pub struct IpcCodeModeSessionProvider { + command: CodeModeHostCommand, + connection: std::sync::Mutex>>, + spawn_permit: Semaphore, +} + +impl Default for IpcCodeModeSessionProvider { + fn default() -> Self { + Self::new(CodeModeHostCommand::default()) + } +} + +impl IpcCodeModeSessionProvider { + pub fn new(command: CodeModeHostCommand) -> Self { + Self { + command, + connection: std::sync::Mutex::new(None), + spawn_permit: Semaphore::new(/*permits*/ 1), + } + } + + async fn connection(&self) -> Result, String> { + if let Some(connection) = { + let current = self + .connection + .lock() + .map_err(|_| "code-mode host connection lock poisoned".to_string())?; + current + .as_ref() + .filter(|connection| connection.is_alive()) + .cloned() + } { + return Ok(connection); + } + + let _spawn_permit = self + .spawn_permit + .acquire() + .await + .map_err(|_| "code-mode host spawn coordinator closed".to_string())?; + if let Some(connection) = { + let current = self + .connection + .lock() + .map_err(|_| "code-mode host connection lock poisoned".to_string())?; + current + .as_ref() + .filter(|connection| connection.is_alive()) + .cloned() + } { + return Ok(connection); + } + let connection = Arc::new(Connection::spawn(&self.command).await?); + *self + .connection + .lock() + .map_err(|_| "code-mode host connection lock poisoned".to_string())? = + Some(Arc::clone(&connection)); + Ok(connection) + } +} + +impl CodeModeSessionProvider for IpcCodeModeSessionProvider { + fn create_session<'a>( + &'a self, + delegate: Arc, + ) -> CodeModeSessionProviderFuture<'a> { + Box::pin(async move { + let connection = self.connection().await?; + let response = connection.request(HostRequest::CreateSession).await?; + let HostResponse::SessionCreated { session_id } = response else { + return Err( + "code-mode host returned an invalid create-session response".to_string() + ); + }; + connection.register_delegate(session_id, delegate).await; + let session: Arc = Arc::new(IpcCodeModeSession { + connection, + session_id, + shutdown: AtomicBool::new(false), + }); + Ok(session) + }) + } +} + +struct IpcCodeModeSession { + connection: Arc, + session_id: SessionId, + shutdown: AtomicBool, +} + +impl CodeModeSession for IpcCodeModeSession { + fn is_alive(&self) -> bool { + !self.shutdown.load(Ordering::Acquire) && self.connection.is_alive() + } + + fn execute<'a>( + &'a self, + request: ExecuteRequest, + ) -> CodeModeSessionResultFuture<'a, StartedCell> { + Box::pin(async move { + self.ensure_active()?; + self.connection.execute(self.session_id, request).await + }) + } + + fn wait<'a>(&'a self, request: WaitRequest) -> CodeModeSessionResultFuture<'a, WaitOutcome> { + Box::pin(async move { + self.ensure_active()?; + let response = self + .connection + .request(HostRequest::Wait { + session_id: self.session_id, + request, + }) + .await?; + match response { + HostResponse::WaitCompleted { outcome } => Ok(outcome), + _ => Err("code-mode host returned an invalid wait response".to_string()), + } + }) + } + + fn terminate<'a>(&'a self, cell_id: CellId) -> CodeModeSessionResultFuture<'a, WaitOutcome> { + Box::pin(async move { + self.ensure_active()?; + let response = self + .connection + .request(HostRequest::Terminate { + session_id: self.session_id, + cell_id, + }) + .await?; + match response { + HostResponse::WaitCompleted { outcome } => Ok(outcome), + _ => Err("code-mode host returned an invalid terminate response".to_string()), + } + }) + } + + fn shutdown<'a>(&'a self) -> CodeModeSessionResultFuture<'a, ()> { + Box::pin(async move { + if self.shutdown.swap(true, Ordering::AcqRel) { + return Ok(()); + } + let result = self + .connection + .request(HostRequest::ShutdownSession { + session_id: self.session_id, + }) + .await; + self.connection.remove_delegate(self.session_id).await; + match result? { + HostResponse::SessionShutdown => Ok(()), + _ => Err("code-mode host returned an invalid shutdown response".to_string()), + } + }) + } +} + +impl IpcCodeModeSession { + fn ensure_active(&self) -> Result<(), String> { + if self.shutdown.load(Ordering::Acquire) { + Err("code mode session is shutting down".to_string()) + } else { + Ok(()) + } + } +} + +fn default_host_program() -> PathBuf { + if let Some(path) = std::env::var_os(CODE_MODE_HOST_PATH_ENV) { + return PathBuf::from(path); + } + let executable_name = if cfg!(windows) { + "codex-code-mode-host.exe" + } else { + "codex-code-mode-host" + }; + if let Ok(current_exe) = std::env::current_exe() + && let Some(parent) = current_exe.parent() + { + let sibling = parent.join(executable_name); + if sibling.is_file() { + return sibling; + } + } + PathBuf::from(executable_name) +} + +#[cfg(test)] +#[path = "tests.rs"] +mod tests; diff --git a/codex-rs/code-mode-client/src/tests.rs b/codex-rs/code-mode-client/src/tests.rs new file mode 100644 index 0000000000..3b38f03fc3 --- /dev/null +++ b/codex-rs/code-mode-client/src/tests.rs @@ -0,0 +1,53 @@ +use std::path::PathBuf; +use std::sync::Arc; + +use codex_code_mode_protocol::CellId; +use codex_code_mode_protocol::CodeModeNestedToolCall; +use codex_code_mode_protocol::CodeModeSessionDelegate; +use codex_code_mode_protocol::CodeModeSessionProvider; +use codex_code_mode_protocol::NotificationFuture; +use codex_code_mode_protocol::ToolInvocationFuture; +use tokio_util::sync::CancellationToken; + +use super::CodeModeHostCommand; +use super::IpcCodeModeSessionProvider; + +struct TestDelegate; + +impl CodeModeSessionDelegate for TestDelegate { + fn invoke_tool<'a>( + &'a self, + _invocation: CodeModeNestedToolCall, + _cancellation_token: CancellationToken, + ) -> ToolInvocationFuture<'a> { + Box::pin(async { Err("unexpected tool invocation".to_string()) }) + } + + fn notify<'a>( + &'a self, + _call_id: String, + _cell_id: CellId, + _text: String, + _cancellation_token: CancellationToken, + ) -> NotificationFuture<'a> { + Box::pin(async { Ok(()) }) + } + + fn cell_closed(&self, _cell_id: &CellId) {} +} + +#[tokio::test] +async fn create_session_reports_host_spawn_errors() { + let provider = IpcCodeModeSessionProvider::new(CodeModeHostCommand { + program: PathBuf::from("codex-code-mode-host-does-not-exist"), + args: Vec::new(), + }); + + let error = provider + .create_session(Arc::new(TestDelegate)) + .await + .err() + .expect("session creation should fail"); + + assert!(error.contains("failed to spawn code-mode host")); +}