From cf0785e99bfbf5db3c63e28cc4630ee42681a9f0 Mon Sep 17 00:00:00 2001 From: pakrym-oai Date: Tue, 10 Mar 2026 19:48:06 -0700 Subject: [PATCH] codex: persist code mode runner sessions --- codex-rs/core/src/state/service.rs | 11 ++ codex-rs/core/src/tools/code_mode.rs | 148 ++++++++++++++----- codex-rs/core/src/tools/code_mode_runner.cjs | 69 +++++---- 3 files changed, 159 insertions(+), 69 deletions(-) diff --git a/codex-rs/core/src/state/service.rs b/codex-rs/core/src/state/service.rs index 5c0a741a12..0c436a921b 100644 --- a/codex-rs/core/src/state/service.rs +++ b/codex-rs/core/src/state/service.rs @@ -15,6 +15,7 @@ use crate::models_manager::manager::ModelsManager; use crate::plugins::PluginsManager; use crate::skills::SkillsManager; use crate::state_db::StateDbHandle; +use crate::tools::code_mode::CodeModeProcess; use crate::tools::network_approval::NetworkApprovalService; use crate::tools::runtimes::ExecveSessionApproval; use crate::tools::sandboxing::ApprovalStore; @@ -31,12 +32,14 @@ use tokio_util::sync::CancellationToken; pub(crate) struct CodeModeStoreService { stored_values: Mutex>, + process: Mutex>, } impl Default for CodeModeStoreService { fn default() -> Self { Self { stored_values: Mutex::new(HashMap::new()), + process: Mutex::new(None), } } } @@ -49,6 +52,14 @@ impl CodeModeStoreService { pub(crate) async fn replace_stored_values(&self, values: HashMap) { *self.stored_values.lock().await = values; } + + pub(crate) async fn store_process(&self, process: CodeModeProcess) { + *self.process.lock().await = Some(process); + } + + pub(crate) async fn take_process(&self) -> Option { + self.process.lock().await.take() + } } pub(crate) struct SessionServices { diff --git a/codex-rs/core/src/tools/code_mode.rs b/codex-rs/core/src/tools/code_mode.rs index ba8dd29e04..955691c784 100644 --- a/codex-rs/core/src/tools/code_mode.rs +++ b/codex-rs/core/src/tools/code_mode.rs @@ -30,6 +30,7 @@ use tokio::io::AsyncBufReadExt; use tokio::io::AsyncReadExt; use tokio::io::AsyncWriteExt; use tokio::io::BufReader; +use tokio::task::JoinHandle; const CODE_MODE_RUNNER_SOURCE: &str = include_str!("code_mode_runner.cjs"); const CODE_MODE_BRIDGE_SOURCE: &str = include_str!("code_mode_bridge.js"); @@ -42,6 +43,37 @@ struct ExecContext { tracker: SharedTurnDiffTracker, } +pub(crate) struct CodeModeProcess { + child: tokio::process::Child, + stdin: tokio::process::ChildStdin, + stdout_lines: tokio::io::Lines>, + stderr_task: Option>, +} + +impl CodeModeProcess { + fn has_exited(&mut self) -> Result { + self.child + .try_wait() + .map(|status| status.is_some()) + .map_err(|err| format!("failed to inspect {PUBLIC_TOOL_NAME} runner: {err}")) + } + + async fn wait_for_exit(&mut self) -> Result { + self.child + .wait() + .await + .map_err(|err| format!("failed to wait for {PUBLIC_TOOL_NAME} runner: {err}")) + } + + async fn stderr(&mut self) -> Result { + self.stderr_task + .take() + .ok_or_else(|| format!("{PUBLIC_TOOL_NAME} stderr collector missing"))? + .await + .map_err(|err| format!("failed to collect {PUBLIC_TOOL_NAME} stderr: {err}")) + } +} + #[derive(Clone, Copy, Debug, Deserialize, Eq, PartialEq, Serialize)] #[serde(rename_all = "snake_case")] enum CodeModeToolKind { @@ -63,12 +95,14 @@ struct EnabledTool { #[derive(Serialize)] #[serde(tag = "type", rename_all = "snake_case")] enum HostToNodeMessage { - Init { + Start { + session_id: String, enabled_tools: Vec, stored_values: HashMap, source: String, }, Response { + session_id: String, id: String, code_mode_result: JsonValue, }, @@ -78,12 +112,14 @@ enum HostToNodeMessage { #[serde(tag = "type", rename_all = "snake_case")] enum NodeToHostMessage { ToolCall { + session_id: String, id: String, name: String, #[serde(default)] input: Option, }, Result { + session_id: String, content_items: Vec, stored_values: HashMap, #[serde(default)] @@ -138,20 +174,33 @@ pub(crate) async fn execute( let enabled_tools = build_enabled_tools(&exec).await; let stored_values = exec.session.services.code_mode_store.stored_values().await; let source = build_source(&code, &enabled_tools).map_err(FunctionCallError::RespondToModel)?; - execute_node(exec, source, enabled_tools, stored_values) - .await - .map_err(FunctionCallError::RespondToModel) + let mut process = match exec.session.services.code_mode_store.take_process().await { + Some(mut process) => { + if matches!(process.has_exited(), Ok(false)) { + process + } else { + spawn_code_mode_process(&exec) + .await + .map_err(FunctionCallError::RespondToModel)? + } + } + None => spawn_code_mode_process(&exec) + .await + .map_err(FunctionCallError::RespondToModel)?, + }; + let result = execute_node(&exec, &mut process, source, enabled_tools, stored_values).await; + if result.is_ok() && matches!(process.has_exited(), Ok(false)) { + exec.session + .services + .code_mode_store + .store_process(process) + .await; + } + result.map_err(FunctionCallError::RespondToModel) } -async fn execute_node( - exec: ExecContext, - source: String, - enabled_tools: Vec, - stored_values: HashMap, -) -> Result { +async fn spawn_code_mode_process(exec: &ExecContext) -> Result { let node_path = resolve_compatible_node(exec.turn.config.js_repl_node_path.as_deref()).await?; - let started_at = std::time::Instant::now(); - let env = create_env(&exec.turn.shell_environment_policy, None); let mut cmd = tokio::process::Command::new(&node_path); cmd.arg("--experimental-vm-modules"); @@ -176,7 +225,7 @@ async fn execute_node( .stderr .take() .ok_or_else(|| format!("{PUBLIC_TOOL_NAME} runner missing stderr"))?; - let mut stdin = child + let stdin = child .stdin .take() .ok_or_else(|| format!("{PUBLIC_TOOL_NAME} runner missing stdin"))?; @@ -188,19 +237,38 @@ async fn execute_node( String::from_utf8_lossy(&buf).trim().to_string() }); + Ok(CodeModeProcess { + child, + stdin, + stdout_lines: BufReader::new(stdout).lines(), + stderr_task: Some(stderr_task), + }) +} + +async fn execute_node( + exec: &ExecContext, + process: &mut CodeModeProcess, + source: String, + enabled_tools: Vec, + stored_values: HashMap, +) -> Result { + let started_at = std::time::Instant::now(); + let session_id = uuid::Uuid::new_v4().to_string(); + write_message( - &mut stdin, - &HostToNodeMessage::Init { - enabled_tools: enabled_tools.clone(), + &mut process.stdin, + &HostToNodeMessage::Start { + session_id: session_id.clone(), + enabled_tools, stored_values, source, }, ) .await?; - let mut stdout_lines = BufReader::new(stdout).lines(); let mut pending_result = None; - while let Some(line) = stdout_lines + while let Some(line) = process + .stdout_lines .next_line() .await .map_err(|err| format!("failed to read {PUBLIC_TOOL_NAME} runner stdout: {err}"))? @@ -212,19 +280,36 @@ async fn execute_node( format!("invalid {PUBLIC_TOOL_NAME} runner message: {err}; line={line}") })?; match message { - NodeToHostMessage::ToolCall { id, name, input } => { + NodeToHostMessage::ToolCall { + session_id: message_session_id, + id, + name, + input, + } => { + if message_session_id != session_id { + return Err(format!( + "unexpected {PUBLIC_TOOL_NAME} runner tool call session id: {message_session_id}" + )); + } let response = HostToNodeMessage::Response { + session_id: message_session_id, id, code_mode_result: call_nested_tool(exec.clone(), name, input).await, }; - write_message(&mut stdin, &response).await?; + write_message(&mut process.stdin, &response).await?; } NodeToHostMessage::Result { + session_id: message_session_id, content_items, stored_values, error_text, max_output_tokens_per_exec_call, } => { + if message_session_id != session_id { + return Err(format!( + "unexpected {PUBLIC_TOOL_NAME} runner result session id: {message_session_id}" + )); + } exec.session .services .code_mode_store @@ -240,20 +325,11 @@ async fn execute_node( } } - drop(stdin); - - let status = child - .wait() - .await - .map_err(|err| format!("failed to wait for {PUBLIC_TOOL_NAME} runner: {err}"))?; - let stderr = stderr_task - .await - .map_err(|err| format!("failed to collect {PUBLIC_TOOL_NAME} stderr: {err}"))?; let wall_time = started_at.elapsed(); - let success = status.success(); - let Some((mut content_items, error_text, max_output_tokens_per_exec_call)) = pending_result else { + let status = process.wait_for_exit().await?; + let stderr = process.stderr().await?; let message = if stderr.is_empty() { format!("{PUBLIC_TOOL_NAME} runner exited without returning a result (status {status})") } else { @@ -262,14 +338,8 @@ async fn execute_node( return Err(message); }; - if !success { - let error_text = error_text.unwrap_or_else(|| { - if stderr.is_empty() { - format!("Process exited with status {status}") - } else { - stderr - } - }); + let success = error_text.is_none(); + if let Some(error_text) = error_text { content_items.push(FunctionCallOutputContentItem::InputText { text: format!("Script error:\n{error_text}"), }); diff --git a/codex-rs/core/src/tools/code_mode_runner.cjs b/codex-rs/core/src/tools/code_mode_runner.cjs index f36fa6f92e..921564436e 100644 --- a/codex-rs/core/src/tools/code_mode_runner.cjs +++ b/codex-rs/core/src/tools/code_mode_runner.cjs @@ -21,11 +21,10 @@ function createProtocol() { let nextId = 0; const pending = new Map(); - let initResolve; - let initReject; - const init = new Promise((resolve, reject) => { - initResolve = resolve; - initReject = reject; + const sessions = new Map(); + let closedResolve; + const closed = new Promise((resolve) => { + closedResolve = resolve; }); rl.on('line', (line) => { @@ -37,35 +36,38 @@ function createProtocol() { try { message = JSON.parse(line); } catch (error) { - initReject(error); + process.stderr.write(`${formatErrorText(error)}\n`); return; } - if (message.type === 'init') { - initResolve(message); + if (message.type === 'start') { + const session = { id: String(message.session_id) }; + sessions.set(session.id, session); + void processSession(protocol, sessions, session, message); return; } if (message.type === 'response') { - const entry = pending.get(message.id); + const entry = pending.get(`${message.session_id}:${message.id}`); if (!entry) { return; } - pending.delete(message.id); + pending.delete(`${message.session_id}:${message.id}`); entry.resolve(message.code_mode_result ?? ''); return; } - initReject(new Error(`Unknown protocol message type: ${message.type}`)); + process.stderr.write(`Unknown protocol message type: ${message.type}\n`); }); rl.on('close', () => { const error = new Error('stdin closed'); - initReject(error); for (const entry of pending.values()) { entry.reject(error); } pending.clear(); + sessions.clear(); + closedResolve(); }); function send(message) { @@ -80,18 +82,20 @@ function createProtocol() { }); } - function request(type, payload) { + function request(sessionId, type, payload) { const id = `msg-${++nextId}`; + const pendingKey = `${sessionId}:${id}`; return new Promise((resolve, reject) => { - pending.set(id, { resolve, reject }); - void send({ type, id, ...payload }).catch((error) => { - pending.delete(id); + pending.set(pendingKey, { resolve, reject }); + void send({ type, session_id: sessionId, id, ...payload }).catch((error) => { + pending.delete(pendingKey); reject(error); }); }); } - return { init, request, send }; + const protocol = { closed, request, send }; + return protocol; } function readContentItems(context) { @@ -112,9 +116,9 @@ function cloneJsonValue(value) { return JSON.parse(JSON.stringify(value)); } -function createToolCaller(protocol) { +function createToolCaller(protocol, sessionId) { return (name, input) => - protocol.request('tool_call', { + protocol.request(sessionId, 'tool_call', { name: String(name), input, }); @@ -348,14 +352,14 @@ function createModuleResolver(context, callTool, enabledTools, state) { }; } -async function runModule(context, request, state, callTool) { +async function runModule(context, start, state, callTool) { const resolveModule = createModuleResolver( context, callTool, - request.enabled_tools ?? [], + start.enabled_tools ?? [], state ); - const mainModule = new SourceTextModule(request.source, { + const mainModule = new SourceTextModule(start.source, { context, identifier: 'exec_main.mjs', importModuleDynamically: async (specifier) => resolveModule(specifier), @@ -365,40 +369,45 @@ async function runModule(context, request, state, callTool) { await mainModule.evaluate(); } -async function main() { - const protocol = createProtocol(); - const request = await protocol.init; +async function processSession(protocol, sessions, session, start) { const state = { maxOutputTokensPerExecCall: DEFAULT_MAX_OUTPUT_TOKENS_PER_EXEC_CALL, - storedValues: cloneJsonValue(request.stored_values ?? {}), + storedValues: cloneJsonValue(start.stored_values ?? {}), }; - const callTool = createToolCaller(protocol); + const callTool = createToolCaller(protocol, session.id); const context = vm.createContext({ __codexContentItems: [], __codex_tool_call: callTool, }); try { - await runModule(context, request, state, callTool); + await runModule(context, start, state, callTool); await protocol.send({ type: 'result', + session_id: session.id, content_items: readContentItems(context), stored_values: state.storedValues, max_output_tokens_per_exec_call: state.maxOutputTokensPerExecCall, }); - process.exit(0); } catch (error) { await protocol.send({ type: 'result', + session_id: session.id, content_items: readContentItems(context), stored_values: state.storedValues, error_text: formatErrorText(error), max_output_tokens_per_exec_call: state.maxOutputTokensPerExecCall, }); - process.exit(1); + } finally { + sessions.delete(session.id); } } +async function main() { + const protocol = createProtocol(); + await protocol.closed; +} + void main().catch(async (error) => { try { process.stderr.write(`${formatErrorText(error)}\n`);